mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
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
313 lines
8.6 KiB
Go
313 lines
8.6 KiB
Go
// 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()
|
|
}
|