mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-05 13:01:58 +00:00
refactor(security): move all database access into pkg/security/lookup
pkg/security no longer contains SQL. Every provider calls a store interface
from lookup, implemented by a procedure backend (Postgres stored procedures,
the default there) and a direct backend (dialect-driven SQL for postgres,
sqlite, mysql and mssql with configurable table and column names).
- add sectypes, lookup, lookup/{dialect,procedure,direct,backends,ddl,conformance}
- split totp and providers sub packages out of the core package
- replace SQLNames/TableNames/QueryMode with lookup.Config (see breaking_changes.md)
- direct backend now covers column/row security and API-key login
- move txsettings SQL to lookup.ApplyTxSettings; remove password.go
- move schema scripts under lookup/, add reference DDL per dialect
- add a shared conformance suite; run it on sqlite, and on Postgres in a
podman/docker container (RESOLVESPEC_TEST_CONTAINERS=1)
- fix procedure schema bugs found on real Postgres: duplicate p_data
parameter, JSON null arrays, expires_at timezone casts, passkey list
GROUP BY, missing resolvespec_passkey_login; accept zone-less timestamps
This commit is contained in:
@@ -0,0 +1,93 @@
|
||||
package dialect
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// --- postgres ---------------------------------------------------------------
|
||||
|
||||
type postgres struct{}
|
||||
|
||||
func (postgres) Name() string { return "postgres" }
|
||||
func (postgres) Matches(driver string) bool {
|
||||
return strings.Contains(driver, "pgx") || strings.Contains(driver, "lib/pq") ||
|
||||
strings.Contains(driver, "postgres")
|
||||
}
|
||||
func (postgres) Placeholder(n int) string { return "$" + strconv.Itoa(n) }
|
||||
func (postgres) Quote(ident string) string { return quoteWith(ident, `"`, `"`) }
|
||||
func (postgres) Bool(v bool) any { return v }
|
||||
func (postgres) ScanBool(src any) (bool, error) { return scanBool(src) }
|
||||
func (postgres) ScanTime(src any) (time.Time, error) { return scanTime(src) }
|
||||
func (postgres) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
|
||||
func (postgres) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
|
||||
func (d postgres) InsertReturningID(table string, cols []string, idCol string) Insert {
|
||||
return Insert{SQL: insertSQL(d, table, cols, "", "RETURNING "+d.Quote(idCol), "DEFAULT VALUES"), Strategy: ReturningQuery}
|
||||
}
|
||||
|
||||
// --- sqlite -----------------------------------------------------------------
|
||||
|
||||
type sqlite struct{}
|
||||
|
||||
func (sqlite) Name() string { return "sqlite" }
|
||||
func (sqlite) Matches(driver string) bool {
|
||||
return strings.Contains(driver, "sqlite")
|
||||
}
|
||||
func (sqlite) Placeholder(int) string { return "?" }
|
||||
func (sqlite) Quote(ident string) string { return quoteWith(ident, `"`, `"`) }
|
||||
func (sqlite) Bool(v bool) any { return boolInt(v) }
|
||||
func (sqlite) ScanBool(src any) (bool, error) { return scanBool(src) }
|
||||
func (sqlite) ScanTime(src any) (time.Time, error) { return scanTime(src) }
|
||||
func (sqlite) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
|
||||
func (sqlite) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
|
||||
func (d sqlite) InsertReturningID(table string, cols []string, _ string) Insert {
|
||||
return Insert{SQL: insertSQL(d, table, cols, "", "", "DEFAULT VALUES"), Strategy: LastInsertID}
|
||||
}
|
||||
|
||||
// --- mysql / mariadb ----------------------------------------------------------
|
||||
|
||||
type mysql struct{}
|
||||
|
||||
func (mysql) Name() string { return "mysql" }
|
||||
func (mysql) Matches(driver string) bool {
|
||||
return strings.Contains(driver, "mysql") || strings.Contains(driver, "mariadb")
|
||||
}
|
||||
func (mysql) Placeholder(int) string { return "?" }
|
||||
func (mysql) Quote(ident string) string { return quoteWith(ident, "`", "`") }
|
||||
func (mysql) Bool(v bool) any { return boolInt(v) }
|
||||
func (mysql) ScanBool(src any) (bool, error) { return scanBool(src) }
|
||||
func (mysql) ScanTime(src any) (time.Time, error) { return scanTime(src) }
|
||||
func (mysql) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
|
||||
func (mysql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
|
||||
func (d mysql) InsertReturningID(table string, cols []string, _ string) Insert {
|
||||
// MySQL has no DEFAULT VALUES; an empty column list is spelled "() VALUES ()".
|
||||
s := insertSQL(d, table, cols, "", "", "() VALUES ()")
|
||||
return Insert{SQL: s, Strategy: LastInsertID}
|
||||
}
|
||||
|
||||
// --- mssql (SQL Server) --------------------------------------------------------
|
||||
|
||||
type mssql struct{}
|
||||
|
||||
func (mssql) Name() string { return "mssql" }
|
||||
func (mssql) Matches(driver string) bool {
|
||||
return strings.Contains(driver, "mssql") || strings.Contains(driver, "sqlserver")
|
||||
}
|
||||
func (mssql) Placeholder(n int) string { return "@p" + strconv.Itoa(n) }
|
||||
func (mssql) Quote(ident string) string { return quoteWith(ident, "[", "]") }
|
||||
func (mssql) Bool(v bool) any { return v }
|
||||
func (mssql) ScanBool(src any) (bool, error) { return scanBool(src) }
|
||||
func (mssql) ScanTime(src any) (time.Time, error) { return scanTime(src) }
|
||||
func (mssql) EncodeJSON(v any) (any, error) { return encodeJSON(v) }
|
||||
func (mssql) DecodeJSON(src any, dst any) error { return decodeJSON(src, dst) }
|
||||
func (d mssql) InsertReturningID(table string, cols []string, idCol string) Insert {
|
||||
return Insert{SQL: insertSQL(d, table, cols, "OUTPUT INSERTED."+d.Quote(idCol), "", "DEFAULT VALUES"), Strategy: ReturningQuery}
|
||||
}
|
||||
|
||||
func boolInt(v bool) int64 {
|
||||
if v {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,312 @@
|
||||
// Package dialect holds the per-database adaptors used by the lookup direct backend.
|
||||
// An adaptor supplies only what differs between databases (placeholders, quoting,
|
||||
// booleans, time and JSON handling, insert-returning-id); the backend builds queries from it.
|
||||
// It imports only the standard library.
|
||||
package dialect
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Dialect is the adaptor for one database type.
|
||||
type Dialect interface {
|
||||
// Name is the registry name, e.g. "postgres".
|
||||
Name() string
|
||||
// Matches reports whether a driver type identifier (lowercased "<pkgpath>.<Type>")
|
||||
// belongs to this database. Used by Detect.
|
||||
Matches(driver string) bool
|
||||
// Placeholder returns the bind placeholder for the n-th (1-based) argument.
|
||||
Placeholder(n int) string
|
||||
// Quote quotes an identifier. A dotted name is quoted per part ("schema.table").
|
||||
// Embedded quote characters are escaped, never interpolated raw.
|
||||
Quote(ident string) string
|
||||
// Bool converts a Go bool to the value bound as a boolean column argument.
|
||||
Bool(v bool) any
|
||||
// ScanBool reads a boolean column value that the driver returned as bool, integer, string or bytes.
|
||||
ScanBool(src any) (bool, error)
|
||||
// ScanTime reads a time column value that the driver returned as time.Time, string or bytes.
|
||||
// NULL (nil) yields the zero time.
|
||||
ScanTime(src any) (time.Time, error)
|
||||
// EncodeJSON converts a value to the argument bound to a JSON/TEXT column. A nil
|
||||
// value, map or slice yields nil (SQL NULL).
|
||||
EncodeJSON(v any) (any, error)
|
||||
// DecodeJSON reads a JSON/TEXT column value into dst. NULL and empty values leave dst untouched.
|
||||
DecodeJSON(src any, dst any) error
|
||||
// InsertReturningID builds an INSERT of cols into table and describes how to read the new id.
|
||||
// Arguments are bound positionally in cols order.
|
||||
InsertReturningID(table string, cols []string, idCol string) Insert
|
||||
}
|
||||
|
||||
// InsertStrategy tells how the generated id is read after an Insert.
|
||||
type InsertStrategy int
|
||||
|
||||
const (
|
||||
// ReturningQuery means the statement returns the id as a single row (QueryRow + Scan).
|
||||
ReturningQuery InsertStrategy = iota
|
||||
// LastInsertID means the id is read from sql.Result.LastInsertId after Exec.
|
||||
LastInsertID
|
||||
)
|
||||
|
||||
// Insert is a generated INSERT statement and how to read the id it creates.
|
||||
type Insert struct {
|
||||
SQL string
|
||||
Strategy InsertStrategy
|
||||
}
|
||||
|
||||
// Querier is implemented by *sql.DB, *sql.Tx and *sql.Conn.
|
||||
type Querier interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
}
|
||||
|
||||
// Run executes the insert and returns the generated id.
|
||||
func (i Insert) Run(ctx context.Context, q Querier, args ...any) (int64, error) {
|
||||
switch i.Strategy {
|
||||
case ReturningQuery:
|
||||
var id int64
|
||||
if err := q.QueryRowContext(ctx, i.SQL, args...).Scan(&id); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
case LastInsertID:
|
||||
res, err := q.ExecContext(ctx, i.SQL, args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
return 0, fmt.Errorf("dialect: unknown insert strategy %d", i.Strategy)
|
||||
}
|
||||
|
||||
// Factory creates a Dialect.
|
||||
type Factory func() Dialect
|
||||
|
||||
var (
|
||||
regMu sync.RWMutex
|
||||
registry = map[string]Factory{}
|
||||
)
|
||||
|
||||
// Register adds a dialect under name. Registering a name twice replaces it, so applications
|
||||
// can override a built-in. Adding a database = implementing Dialect and calling Register.
|
||||
func Register(name string, f Factory) {
|
||||
regMu.Lock()
|
||||
defer regMu.Unlock()
|
||||
registry[strings.ToLower(name)] = f
|
||||
}
|
||||
|
||||
// Get returns the dialect registered under name.
|
||||
func Get(name string) (Dialect, error) {
|
||||
regMu.RLock()
|
||||
f, ok := registry[strings.ToLower(name)]
|
||||
regMu.RUnlock()
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("dialect: unknown dialect %q (registered: %s)", name, strings.Join(Names(), ", "))
|
||||
}
|
||||
return f(), nil
|
||||
}
|
||||
|
||||
// Names lists the registered dialect names, sorted.
|
||||
func Names() []string {
|
||||
regMu.RLock()
|
||||
defer regMu.RUnlock()
|
||||
names := make([]string, 0, len(registry))
|
||||
for n := range registry {
|
||||
names = append(names, n)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
// Detect picks the dialect for db from its driver type.
|
||||
func Detect(db *sql.DB) (Dialect, error) {
|
||||
if db == nil {
|
||||
return nil, fmt.Errorf("dialect: nil database")
|
||||
}
|
||||
return DetectDriver(driverID(db.Driver()))
|
||||
}
|
||||
|
||||
// DetectDriver picks the dialect for a driver type identifier (see Dialect.Matches).
|
||||
func DetectDriver(driver string) (Dialect, error) {
|
||||
driver = strings.ToLower(driver)
|
||||
for _, n := range Names() {
|
||||
d, _ := Get(n)
|
||||
if d != nil && d.Matches(driver) {
|
||||
return d, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("dialect: cannot detect a dialect for driver %q; set the dialect explicitly", driver)
|
||||
}
|
||||
|
||||
// driverID builds "<pkgpath>.<Type>" for a driver value, lowercased.
|
||||
func driverID(drv any) string {
|
||||
t := reflect.TypeOf(drv)
|
||||
for t != nil && t.Kind() == reflect.Pointer {
|
||||
t = t.Elem()
|
||||
}
|
||||
if t == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.ToLower(t.PkgPath() + "." + t.Name())
|
||||
}
|
||||
|
||||
func init() {
|
||||
Register("postgres", func() Dialect { return postgres{} })
|
||||
Register("sqlite", func() Dialect { return sqlite{} })
|
||||
Register("mysql", func() Dialect { return mysql{} })
|
||||
Register("mssql", func() Dialect { return mssql{} })
|
||||
}
|
||||
|
||||
// --- shared helpers -------------------------------------------------------
|
||||
|
||||
// quoteWith quotes each dotted part of ident with open/close, doubling embedded close characters.
|
||||
func quoteWith(ident, open, closeq string) string {
|
||||
parts := strings.Split(ident, ".")
|
||||
for i, p := range parts {
|
||||
parts[i] = open + strings.ReplaceAll(p, closeq, closeq+closeq) + closeq
|
||||
}
|
||||
return strings.Join(parts, ".")
|
||||
}
|
||||
|
||||
func scanBool(src any) (bool, error) {
|
||||
switch v := src.(type) {
|
||||
case nil:
|
||||
return false, nil
|
||||
case bool:
|
||||
return v, nil
|
||||
case int64:
|
||||
return v != 0, nil
|
||||
case int:
|
||||
return v != 0, nil
|
||||
case int32:
|
||||
return v != 0, nil
|
||||
case float64:
|
||||
return v != 0, nil
|
||||
case []byte:
|
||||
return scanBool(string(v))
|
||||
case string:
|
||||
switch strings.ToLower(strings.TrimSpace(v)) {
|
||||
case "1", "t", "true", "y", "yes", "on":
|
||||
return true, nil
|
||||
case "", "0", "f", "false", "n", "no", "off":
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
return false, fmt.Errorf("dialect: cannot read %T (%v) as bool", src, src)
|
||||
}
|
||||
|
||||
var timeLayouts = []string{
|
||||
time.RFC3339Nano,
|
||||
"2006-01-02 15:04:05.999999999 -0700",
|
||||
"2006-01-02 15:04:05.999999999-07:00",
|
||||
"2006-01-02 15:04:05.999999999Z07:00",
|
||||
"2006-01-02T15:04:05.999999999",
|
||||
"2006-01-02 15:04:05.999999999",
|
||||
"2006-01-02",
|
||||
}
|
||||
|
||||
func scanTime(src any) (time.Time, error) {
|
||||
switch v := src.(type) {
|
||||
case nil:
|
||||
return time.Time{}, nil
|
||||
case time.Time:
|
||||
return v, nil
|
||||
case []byte:
|
||||
return scanTime(string(v))
|
||||
case string:
|
||||
s := strings.TrimSpace(v)
|
||||
if s == "" {
|
||||
return time.Time{}, nil
|
||||
}
|
||||
// Some drivers append the Go monotonic/zone suffix ("+0000 UTC"); drop it.
|
||||
if i := strings.Index(s, " m="); i >= 0 {
|
||||
s = s[:i]
|
||||
}
|
||||
s = strings.TrimSuffix(s, " UTC")
|
||||
s = strings.TrimSuffix(s, " +0000 +0000")
|
||||
for _, l := range timeLayouts {
|
||||
if t, err := time.Parse(l, s); err == nil {
|
||||
return t, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return time.Time{}, fmt.Errorf("dialect: cannot read %T (%v) as time", src, src)
|
||||
}
|
||||
|
||||
func encodeJSON(v any) (any, error) {
|
||||
if v == nil {
|
||||
return nil, nil
|
||||
}
|
||||
rv := reflect.ValueOf(v)
|
||||
switch rv.Kind() {
|
||||
case reflect.Map, reflect.Slice, reflect.Pointer, reflect.Interface:
|
||||
if rv.IsNil() {
|
||||
return nil, nil
|
||||
}
|
||||
}
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("dialect: encode json: %w", err)
|
||||
}
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
func decodeJSON(src any, dst any) error {
|
||||
var raw []byte
|
||||
switch v := src.(type) {
|
||||
case nil:
|
||||
return nil
|
||||
case string:
|
||||
raw = []byte(v)
|
||||
case []byte:
|
||||
raw = v
|
||||
default:
|
||||
return fmt.Errorf("dialect: cannot read %T as json", src)
|
||||
}
|
||||
if len(strings.TrimSpace(string(raw))) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := json.Unmarshal(raw, dst); err != nil {
|
||||
return fmt.Errorf("dialect: decode json: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// insertSQL assembles "INSERT INTO t (cols) <mid> VALUES (...) <tail>" for a dialect.
|
||||
func insertSQL(d Dialect, table string, cols []string, mid, tail, defaults string) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("INSERT INTO ")
|
||||
b.WriteString(d.Quote(table))
|
||||
if len(cols) == 0 {
|
||||
if mid != "" {
|
||||
b.WriteString(" " + mid)
|
||||
}
|
||||
b.WriteString(" " + defaults)
|
||||
if tail != "" {
|
||||
b.WriteString(" " + tail)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
qc := make([]string, len(cols))
|
||||
ph := make([]string, len(cols))
|
||||
for i, c := range cols {
|
||||
qc[i] = d.Quote(c)
|
||||
ph[i] = d.Placeholder(i + 1)
|
||||
}
|
||||
b.WriteString(" (" + strings.Join(qc, ", ") + ")")
|
||||
if mid != "" {
|
||||
b.WriteString(" " + mid)
|
||||
}
|
||||
b.WriteString(" VALUES (" + strings.Join(ph, ", ") + ")")
|
||||
if tail != "" {
|
||||
b.WriteString(" " + tail)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
@@ -0,0 +1,270 @@
|
||||
package dialect_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/glebarez/go-sqlite"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
func get(t *testing.T, name string) dialect.Dialect {
|
||||
t.Helper()
|
||||
d, err := dialect.Get(name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func TestRegistryHasBuiltins(t *testing.T) {
|
||||
got := strings.Join(dialect.Names(), ",")
|
||||
if got != "mssql,mysql,postgres,sqlite" {
|
||||
t.Errorf("names = %s", got)
|
||||
}
|
||||
if _, err := dialect.Get("oracle"); err == nil {
|
||||
t.Error("unknown dialect should error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterCustomDialect(t *testing.T) {
|
||||
// A new database is added by registering a dialect; no core change needed.
|
||||
dialect.Register("custom", func() dialect.Dialect { return customDialect{get(t, "sqlite")} })
|
||||
d, err := dialect.Get("CUSTOM")
|
||||
if err != nil || d.Name() != "custom" {
|
||||
t.Fatalf("got %v, %v", d, err)
|
||||
}
|
||||
if got, _ := dialect.DetectDriver("example.com/driver.customdriver"); got == nil || got.Name() != "custom" {
|
||||
t.Errorf("custom detection failed: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
type customDialect struct{ dialect.Dialect }
|
||||
|
||||
func (customDialect) Name() string { return "custom" }
|
||||
func (customDialect) Matches(driver string) bool { return strings.Contains(driver, "customdriver") }
|
||||
|
||||
func TestPlaceholderAndQuote(t *testing.T) {
|
||||
cases := []struct {
|
||||
name, ph1, ph3, quote, quoteDotted, quoteEscape string
|
||||
}{
|
||||
{"postgres", "$1", "$3", `"users"`, `"auth"."users"`, `"a""b"`},
|
||||
{"sqlite", "?", "?", `"users"`, `"auth"."users"`, `"a""b"`},
|
||||
{"mysql", "?", "?", "`users`", "`auth`.`users`", "`a``b`"},
|
||||
{"mssql", "@p1", "@p3", "[users]", "[auth].[users]", "[a]]b]"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
d := get(t, c.name)
|
||||
if d.Placeholder(1) != c.ph1 || d.Placeholder(3) != c.ph3 {
|
||||
t.Errorf("%s placeholders: %s %s", c.name, d.Placeholder(1), d.Placeholder(3))
|
||||
}
|
||||
if d.Quote("users") != c.quote {
|
||||
t.Errorf("%s quote: %s", c.name, d.Quote("users"))
|
||||
}
|
||||
if d.Quote("auth.users") != c.quoteDotted {
|
||||
t.Errorf("%s dotted: %s", c.name, d.Quote("auth.users"))
|
||||
}
|
||||
raw := map[string]string{"postgres": `a"b`, "sqlite": `a"b`, "mysql": "a`b", "mssql": "a]b"}[c.name]
|
||||
q := c.quoteEscape
|
||||
if got := d.Quote(raw); got != q {
|
||||
t.Errorf("%s escape: got %s want %s", c.name, got, q)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInsertReturningID(t *testing.T) {
|
||||
cols := []string{"user_id", "name"}
|
||||
cases := []struct {
|
||||
name string
|
||||
wantSQL string
|
||||
strategy dialect.InsertStrategy
|
||||
empty string
|
||||
}{
|
||||
{"postgres", `INSERT INTO "t" ("user_id", "name") VALUES ($1, $2) RETURNING "id"`, dialect.ReturningQuery, `INSERT INTO "t" DEFAULT VALUES RETURNING "id"`},
|
||||
{"sqlite", `INSERT INTO "t" ("user_id", "name") VALUES (?, ?)`, dialect.LastInsertID, `INSERT INTO "t" DEFAULT VALUES`},
|
||||
{"mysql", "INSERT INTO `t` (`user_id`, `name`) VALUES (?, ?)", dialect.LastInsertID, "INSERT INTO `t` () VALUES ()"},
|
||||
{"mssql", `INSERT INTO [t] ([user_id], [name]) OUTPUT INSERTED.[id] VALUES (@p1, @p2)`, dialect.ReturningQuery, `INSERT INTO [t] OUTPUT INSERTED.[id] DEFAULT VALUES`},
|
||||
}
|
||||
for _, c := range cases {
|
||||
ins := get(t, c.name).InsertReturningID("t", cols, "id")
|
||||
if ins.SQL != c.wantSQL || ins.Strategy != c.strategy {
|
||||
t.Errorf("%s: got %q (%d)", c.name, ins.SQL, ins.Strategy)
|
||||
}
|
||||
if e := get(t, c.name).InsertReturningID("t", nil, "id"); e.SQL != c.empty {
|
||||
t.Errorf("%s empty: got %q", c.name, e.SQL)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBoolRoundTrip(t *testing.T) {
|
||||
for _, name := range dialect.Names() {
|
||||
if name == "custom" {
|
||||
continue
|
||||
}
|
||||
d := get(t, name)
|
||||
for _, v := range []bool{true, false} {
|
||||
got, err := d.ScanBool(d.Bool(v))
|
||||
if err != nil || got != v {
|
||||
t.Errorf("%s: Bool(%v) round trip = %v, %v", name, v, got, err)
|
||||
}
|
||||
}
|
||||
for _, c := range []struct {
|
||||
in any
|
||||
want bool
|
||||
}{{int64(1), true}, {int64(0), false}, {"1", true}, {"false", false}, {[]byte("t"), true}, {nil, false}} {
|
||||
if got, err := d.ScanBool(c.in); err != nil || got != c.want {
|
||||
t.Errorf("%s: ScanBool(%v) = %v, %v", name, c.in, got, err)
|
||||
}
|
||||
}
|
||||
if _, err := d.ScanBool("maybe"); err == nil {
|
||||
t.Errorf("%s: ScanBool(maybe) should fail", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanTime(t *testing.T) {
|
||||
d := get(t, "sqlite")
|
||||
want := time.Date(2026, 3, 4, 5, 6, 7, 0, time.UTC)
|
||||
for _, in := range []any{
|
||||
want,
|
||||
"2026-03-04 05:06:07+00:00",
|
||||
"2026-03-04T05:06:07Z",
|
||||
"2026-03-04 05:06:07",
|
||||
[]byte("2026-03-04 05:06:07"),
|
||||
"2026-03-04 05:06:07 +0000 UTC",
|
||||
} {
|
||||
got, err := d.ScanTime(in)
|
||||
if err != nil || !got.Equal(want) {
|
||||
t.Errorf("ScanTime(%v) = %v, %v", in, got, err)
|
||||
}
|
||||
}
|
||||
if got, err := d.ScanTime(nil); err != nil || !got.IsZero() {
|
||||
t.Errorf("ScanTime(nil) = %v, %v", got, err)
|
||||
}
|
||||
if _, err := d.ScanTime("not a time"); err == nil {
|
||||
t.Error("ScanTime(garbage) should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSON(t *testing.T) {
|
||||
for _, name := range []string{"postgres", "sqlite", "mysql", "mssql"} {
|
||||
d := get(t, name)
|
||||
enc, err := d.EncodeJSON([]string{"a", "b"})
|
||||
if err != nil || enc != `["a","b"]` {
|
||||
t.Errorf("%s encode = %v, %v", name, enc, err)
|
||||
}
|
||||
var nilSlice []string
|
||||
var nilMap map[string]any
|
||||
for _, v := range []any{nil, nilSlice, nilMap} {
|
||||
if enc, err := d.EncodeJSON(v); err != nil || enc != nil {
|
||||
t.Errorf("%s: EncodeJSON(%T nil) = %v, %v; want SQL NULL", name, v, enc, err)
|
||||
}
|
||||
}
|
||||
var out []string
|
||||
if err := d.DecodeJSON([]byte(`["x"]`), &out); err != nil || len(out) != 1 || out[0] != "x" {
|
||||
t.Errorf("%s decode bytes = %v, %v", name, out, err)
|
||||
}
|
||||
out = []string{"keep"}
|
||||
if err := d.DecodeJSON(nil, &out); err != nil || out[0] != "keep" {
|
||||
t.Errorf("%s decode NULL should leave dst: %v, %v", name, out, err)
|
||||
}
|
||||
if err := d.DecodeJSON("", &out); err != nil || out[0] != "keep" {
|
||||
t.Errorf("%s decode empty should leave dst: %v, %v", name, out, err)
|
||||
}
|
||||
if err := d.DecodeJSON("{bad", &out); err == nil {
|
||||
t.Errorf("%s decode of invalid json should fail", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectDriver(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"github.com/jackc/pgx/v5/stdlib.driver": "postgres",
|
||||
"github.com/lib/pq.driver": "postgres",
|
||||
"github.com/mattn/go-sqlite3.sqlitedriver": "sqlite",
|
||||
"modernc.org/sqlite.driver": "sqlite",
|
||||
"github.com/go-sql-driver/mysql.mysqldriver": "mysql",
|
||||
"github.com/microsoft/go-mssqldb.driver": "mssql",
|
||||
"github.com/denisenkom/go-mssqldb.driver": "mssql",
|
||||
}
|
||||
for drv, want := range cases {
|
||||
d, err := dialect.DetectDriver(drv)
|
||||
if err != nil || d.Name() != want {
|
||||
t.Errorf("%s: got %v, %v; want %s", drv, d, err, want)
|
||||
}
|
||||
}
|
||||
if _, err := dialect.DetectDriver("example.com/unknown.driver"); err == nil {
|
||||
t.Error("unknown driver should fail with an explicit-dialect hint")
|
||||
}
|
||||
if _, err := dialect.Detect(nil); err == nil {
|
||||
t.Error("nil db should fail")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSQLiteRoundTrip runs the sqlite dialect against a real in-memory database.
|
||||
func TestSQLiteRoundTrip(t *testing.T) {
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
db.SetMaxOpenConns(1)
|
||||
|
||||
d, err := dialect.Detect(db)
|
||||
if err != nil || d.Name() != "sqlite" {
|
||||
t.Fatalf("Detect = %v, %v", d, err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
if _, err := db.ExecContext(ctx, `CREATE TABLE t (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT, active BOOLEAN, scopes TEXT, at TIMESTAMP)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
now := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC)
|
||||
scopes, _ := d.EncodeJSON([]string{"read", "write"})
|
||||
ins := d.InsertReturningID("t", []string{"name", "active", "scopes", "at"}, "id")
|
||||
id, err := ins.Run(ctx, db, "k1", d.Bool(true), scopes, now)
|
||||
if err != nil || id != 1 {
|
||||
t.Fatalf("insert = %d, %v", id, err)
|
||||
}
|
||||
id, err = ins.Run(ctx, db, "k2", d.Bool(false), nil, now)
|
||||
if err != nil || id != 2 {
|
||||
t.Fatalf("second insert = %d, %v", id, err)
|
||||
}
|
||||
|
||||
q := "SELECT active, scopes, at FROM " + d.Quote("t") + " WHERE " + d.Quote("id") + " = " + d.Placeholder(1)
|
||||
var active, scopesRaw, at any
|
||||
if err := db.QueryRowContext(ctx, q, 1).Scan(&active, &scopesRaw, &at); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, err := d.ScanBool(active); err != nil || !b {
|
||||
t.Errorf("active = %v, %v", b, err)
|
||||
}
|
||||
var got []string
|
||||
if err := d.DecodeJSON(scopesRaw, &got); err != nil || len(got) != 2 {
|
||||
t.Errorf("scopes = %v, %v", got, err)
|
||||
}
|
||||
if ts, err := d.ScanTime(at); err != nil || !ts.Equal(now) {
|
||||
t.Errorf("at = %v (%T), %v", ts, at, err)
|
||||
}
|
||||
|
||||
if err := db.QueryRowContext(ctx, q, 2).Scan(&active, &scopesRaw, &at); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if b, _ := d.ScanBool(active); b {
|
||||
t.Error("second row should be inactive")
|
||||
}
|
||||
if scopesRaw != nil {
|
||||
t.Errorf("nil scopes should be NULL, got %v", scopesRaw)
|
||||
}
|
||||
|
||||
// Insert inside a transaction goes through the same Querier.
|
||||
tx, _ := db.BeginTx(ctx, nil)
|
||||
if id, err := d.InsertReturningID("t", nil, "id").Run(ctx, tx); err != nil || id != 3 {
|
||||
t.Errorf("DEFAULT VALUES insert in tx = %d, %v", id, err)
|
||||
}
|
||||
_ = tx.Rollback()
|
||||
}
|
||||
Reference in New Issue
Block a user