mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 03:22:09 +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,162 @@
|
||||
// Package backends assembles a lookup.Provider: it builds the procedure and direct stores for
|
||||
// one database and routes every operation to one of them according to lookup.Config.
|
||||
// It lives apart from package lookup because both backends import lookup.
|
||||
package backends
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/direct"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure"
|
||||
)
|
||||
|
||||
// Options are the settings that are not naming or mode.
|
||||
type Options struct {
|
||||
// DBFactory is called to obtain a fresh *sql.DB when the current one has been closed.
|
||||
// Nil disables reconnecting.
|
||||
DBFactory func() (*sql.DB, error)
|
||||
// UpgradePasswordHash rewrites a legacy cleartext password as bcrypt after a successful
|
||||
// direct-mode login. Off by default.
|
||||
UpgradePasswordHash bool
|
||||
// NoGroupTables skips the group membership table when loading direct-mode policy rules.
|
||||
NoGroupTables bool
|
||||
}
|
||||
|
||||
// New builds a Provider for db. cfg is merged with the defaults and validated; the dialect
|
||||
// is cfg.Dialect or detected from the driver. Every operation's mode is resolved up front so
|
||||
// an impossible combination (procedure mode on a non-Postgres dialect) fails here, not on the
|
||||
// first request.
|
||||
func New(db *sql.DB, cfg lookup.Config, opts Options) (*lookup.Provider, error) {
|
||||
if db == nil {
|
||||
return nil, fmt.Errorf("backends: nil database")
|
||||
}
|
||||
res, err := cfg.Resolve()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
d, err := cfg.ResolveDialect(db)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, op := range lookup.AllOps() {
|
||||
if _, err := res.EffectiveMode(op, d.Name()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
c := &chooser{cfg: res, dialect: d.Name(), procs: res.Procs}
|
||||
run := procedure.NewDB(db, opts.DBFactory, c.resetProbes)
|
||||
c.db = run
|
||||
|
||||
base, err := direct.NewBase(run, d, res.Schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p := procedure.NewPasskey(run, res.Procs)
|
||||
return &lookup.Provider{
|
||||
Auth: &authRouter{c: c,
|
||||
proc: procedure.NewAuth(run, res.Procs),
|
||||
direct: direct.NewAuth(base, direct.AuthOptions{UpgradePasswordHash: opts.UpgradePasswordHash})},
|
||||
Keys: &keysRouter{c: c,
|
||||
proc: procedure.NewKeys(run, res.Procs),
|
||||
direct: direct.NewKeys(base)},
|
||||
OAuthClient: &oauthClientRouter{c: c,
|
||||
proc: procedure.NewOAuthClients(run, res.Procs),
|
||||
direct: direct.NewOAuthClients(base)},
|
||||
OAuthUser: &oauthUserRouter{c: c,
|
||||
proc: procedure.NewOAuthUsers(run, res.Procs),
|
||||
direct: direct.NewOAuthUsers(base)},
|
||||
Passkey: &passkeyRouter{c: c, proc: p, direct: direct.NewPasskey(base)},
|
||||
TOTP: &totpRouter{c: c,
|
||||
proc: procedure.NewTOTP(run, res.Procs),
|
||||
direct: direct.NewTOTP(base)},
|
||||
Policy: &policyRouter{c: c,
|
||||
proc: procedure.NewPolicy(run, res.Procs),
|
||||
direct: direct.NewPolicy(base, direct.PolicyOptions{NoGroups: opts.NoGroupTables})},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Failed returns a Provider whose every operation returns err. Constructors that cannot
|
||||
// return an error use it so a bad configuration fails closed on first use.
|
||||
func Failed(err error) *lookup.Provider {
|
||||
c := &chooser{fail: err}
|
||||
return &lookup.Provider{
|
||||
Auth: &authRouter{c: c},
|
||||
Keys: &keysRouter{c: c},
|
||||
OAuthClient: &oauthClientRouter{c: c},
|
||||
OAuthUser: &oauthUserRouter{c: c},
|
||||
Passkey: &passkeyRouter{c: c},
|
||||
TOTP: &totpRouter{c: c},
|
||||
Policy: &policyRouter{c: c},
|
||||
}
|
||||
}
|
||||
|
||||
// chooser decides per operation whether the procedure or the direct store runs.
|
||||
type chooser struct {
|
||||
cfg *lookup.Resolved
|
||||
dialect string
|
||||
procs lookup.ProcNames
|
||||
db *procedure.DB
|
||||
probes sync.Map // proc name -> bool
|
||||
fail error // set by Failed: every operation returns it
|
||||
}
|
||||
|
||||
func (c *chooser) resetProbes() {
|
||||
c.probes.Range(func(k, _ any) bool { c.probes.Delete(k); return true })
|
||||
}
|
||||
|
||||
// useProc reports whether op should call the stored procedure proc. In auto mode on Postgres
|
||||
// the catalog is probed once per procedure (cached until a reconnect).
|
||||
func (c *chooser) useProc(ctx context.Context, op lookup.Op, proc string) (bool, error) {
|
||||
if c.fail != nil {
|
||||
return false, c.fail
|
||||
}
|
||||
m, err := c.cfg.EffectiveMode(op, c.dialect)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
switch m {
|
||||
case lookup.ModeProcedure:
|
||||
return true, nil
|
||||
case lookup.ModeAuto:
|
||||
if v, ok := c.probes.Load(proc); ok {
|
||||
return v.(bool), nil
|
||||
}
|
||||
exists := probeProcedure(ctx, c.db.Get(), proc)
|
||||
c.probes.Store(proc, exists)
|
||||
return exists, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// probeProcedure asks the Postgres catalog whether a function exists. Any failure counts as
|
||||
// "does not exist" so the probe can never block an operation.
|
||||
func probeProcedure(ctx context.Context, db *sql.DB, proc string) (exists bool) {
|
||||
if db == nil {
|
||||
return false
|
||||
}
|
||||
defer func() { _ = recover() }()
|
||||
dbtrace.Raw(ctx, "probe.pg_proc")
|
||||
if err := db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM pg_proc WHERE proname = $1 LIMIT 1)`, proc).Scan(&exists); err != nil {
|
||||
return false
|
||||
}
|
||||
return exists
|
||||
}
|
||||
|
||||
// pick returns the store that should serve op.
|
||||
func pick[T any](c *chooser, ctx context.Context, op lookup.Op, proc string, p, d T) (T, error) {
|
||||
use, err := c.useProc(ctx, op, proc)
|
||||
if err != nil {
|
||||
var zero T
|
||||
return zero, err
|
||||
}
|
||||
if use {
|
||||
return p, nil
|
||||
}
|
||||
return d, nil
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
package backends
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
_ "github.com/glebarez/go-sqlite"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
func sqliteDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
ddl, err := ddl.SQL("sqlite")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.Exec(ddl); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func TestSQLiteDefaultsToDirect(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
p, err := New(sqliteDB(t), lookup.Config{}, Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reg, err := p.Auth.Register(ctx, sectypes.RegisterRequest{Username: "a", Email: "a@x.io", Password: "pw"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := p.Auth.Session(ctx, reg.Token, "authenticate"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if on, err := p.TOTP.Status(ctx, reg.User.UserID); err != nil || on {
|
||||
t.Fatalf("%v %v", on, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcedureModeRejectedOnSQLite(t *testing.T) {
|
||||
_, err := New(sqliteDB(t), lookup.Config{Overrides: map[lookup.Op]lookup.Mode{lookup.OpLogin: lookup.ModeProcedure}}, Options{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomSchemaAndUnknownDialect(t *testing.T) {
|
||||
if _, err := New(sqliteDB(t), lookup.Config{Dialect: "nosuch"}, Options{}); err == nil {
|
||||
t.Fatal("unknown dialect accepted")
|
||||
}
|
||||
bad := lookup.Config{Schema: lookup.Schema{lookup.EntityUsers: {Name: "x; drop"}}}
|
||||
if _, err := New(sqliteDB(t), bad, Options{}); err == nil {
|
||||
t.Fatal("unsafe schema accepted")
|
||||
}
|
||||
if _, err := New(nil, lookup.Config{}, Options{}); err == nil {
|
||||
t.Fatal("nil db accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostgresDefaultsToProcedure(t *testing.T) {
|
||||
db, mock, _ := sqlmock.New()
|
||||
defer db.Close()
|
||||
p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres}, Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true))
|
||||
if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on {
|
||||
t.Fatalf("%v %v", on, err)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostgresAutoProbesOnce(t *testing.T) {
|
||||
db, mock, _ := sqlmock.New()
|
||||
defer db.Close()
|
||||
p, err := New(db, lookup.Config{Dialect: lookup.DialectPostgres, Mode: lookup.ModeAuto}, Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mock.ExpectQuery("pg_proc").WithArgs("resolvespec_totp_get_status").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"e"}).AddRow(true))
|
||||
for i := 0; i < 2; i++ { // second call must reuse the cached probe
|
||||
mock.ExpectQuery("resolvespec_totp_get_status").WithArgs(1).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_enabled"}).AddRow(true, nil, true))
|
||||
if on, err := p.TOTP.Status(context.Background(), 1); err != nil || !on {
|
||||
t.Fatalf("%v %v", on, err)
|
||||
}
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailed(t *testing.T) {
|
||||
p := Failed(errors.New("boom"))
|
||||
if _, err := p.Auth.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "boom" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if _, err := p.Policy.RowSecurity(context.Background(), 1, "s", "t"); err == nil {
|
||||
t.Fatal("want error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package backends
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"os"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
_ "github.com/microsoft/go-mssqldb"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/conformance"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
// Every backend/dialect runs the same suite. SQLite runs always. The others run only when a
|
||||
// DSN is set, and only against a database you are happy to add rows to (all names the suite
|
||||
// creates carry a unique "cf<hex>_" prefix and are removed afterwards):
|
||||
//
|
||||
// RESOLVESPEC_TEST_PG_DSN Postgres with the procedure schema installed
|
||||
// (lookup/database_schema.sql + keystore_schema.sql): procedure mode
|
||||
// RESOLVESPEC_TEST_PG_DIRECT_DSN Postgres for direct mode; ddl/postgres.sql is applied if the
|
||||
// tables are missing. Do not point it at the procedure schema:
|
||||
// the column types differ.
|
||||
// RESOLVESPEC_TEST_MYSQL_DSN MySQL (needs a "mysql" database/sql driver linked into the test binary)
|
||||
// RESOLVESPEC_TEST_MSSQL_DSN SQL Server (driver "sqlserver")
|
||||
func TestConformance(t *testing.T) {
|
||||
t.Run("sqlite/direct", func(t *testing.T) {
|
||||
runConformance(t, sqliteDB(t), "sqlite", lookup.Config{Mode: lookup.ModeDirect}, false)
|
||||
})
|
||||
t.Run("sqlite/default", func(t *testing.T) {
|
||||
runConformance(t, sqliteDB(t), "sqlite", lookup.Config{}, false)
|
||||
})
|
||||
|
||||
real := []struct {
|
||||
name, env, driver, dialect string
|
||||
cfg lookup.Config
|
||||
applyDDL bool
|
||||
}{
|
||||
{"postgres/procedure", "RESOLVESPEC_TEST_PG_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeProcedure}, false},
|
||||
{"postgres/direct", "RESOLVESPEC_TEST_PG_DIRECT_DSN", "pgx", "postgres", lookup.Config{Mode: lookup.ModeDirect}, true},
|
||||
{"mysql/direct", "RESOLVESPEC_TEST_MYSQL_DSN", "mysql", "mysql", lookup.Config{}, true},
|
||||
{"mssql/direct", "RESOLVESPEC_TEST_MSSQL_DSN", "sqlserver", "mssql", lookup.Config{}, true},
|
||||
}
|
||||
for _, r := range real {
|
||||
t.Run(r.name, func(t *testing.T) {
|
||||
dsn := os.Getenv(r.env)
|
||||
if dsn == "" {
|
||||
t.Skipf("%s not set", r.env)
|
||||
}
|
||||
runOnServer(t, r.driver, dsn, r.dialect, r.cfg, r.applyDDL)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// runOnServer opens dsn, optionally applies the reference DDL, and runs the suite.
|
||||
func runOnServer(t *testing.T, driver, dsn, dialectName string, cfg lookup.Config, applyDDL bool) {
|
||||
t.Helper()
|
||||
if !slices.Contains(sql.Drivers(), driver) {
|
||||
t.Skipf("database/sql driver %q is not linked into this test binary", driver)
|
||||
}
|
||||
db, err := sql.Open(driver, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
if err := db.Ping(); err != nil {
|
||||
t.Fatalf("ping: %v", err)
|
||||
}
|
||||
if applyDDL {
|
||||
stmts, err := ddl.Statements(dialectName)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, s := range stmts {
|
||||
if _, err := db.Exec(s); err != nil {
|
||||
t.Fatalf("apply ddl: %v\n%s", err, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
runConformance(t, db, dialectName, cfg, true)
|
||||
}
|
||||
|
||||
func runConformance(t *testing.T, db *sql.DB, dialectName string, cfg lookup.Config, shared bool) {
|
||||
t.Helper()
|
||||
d, err := dialect.Get(dialectName)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg.Dialect = dialectName
|
||||
p, err := New(db, cfg, Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var b [3]byte
|
||||
_, _ = rand.Read(b[:])
|
||||
env := conformance.Env{Provider: p, DB: db, Dialect: d, Prefix: "cf" + hex.EncodeToString(b[:]) + "_"}
|
||||
if shared {
|
||||
env.Cleanup = func(t *testing.T) { cleanup(t, db, d, env.Prefix) }
|
||||
}
|
||||
conformance.Run(t, env)
|
||||
}
|
||||
|
||||
// cleanup deletes the rows a conformance run created, by prefix. Child rows go with their user
|
||||
// through the foreign keys; tables without one are cleaned explicitly.
|
||||
func cleanup(t *testing.T, db *sql.DB, d dialect.Dialect, prefix string) {
|
||||
t.Helper()
|
||||
like := prefix + "%"
|
||||
for _, q := range []struct{ table, col string }{
|
||||
{"oauth_codes", "code"},
|
||||
{"oauth_clients", "client_id"},
|
||||
{"token_blacklist", "token"},
|
||||
{"sec_column_rules", "schema_name"},
|
||||
{"sec_row_rules", "schema_name"},
|
||||
{"users", "username"},
|
||||
} {
|
||||
if _, err := db.Exec(fmt.Sprintf("DELETE FROM %s WHERE %s LIKE %s", q.table, q.col, d.Placeholder(1)), like); err != nil {
|
||||
t.Logf("cleanup %s: %v", q.table, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
package backends
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
// Container tests start a throwaway database server with podman or docker (whichever is
|
||||
// installed, podman first) and run the conformance suite against it. They pull an image, so
|
||||
// they only run when RESOLVESPEC_TEST_CONTAINERS=1 and not with -short. The container is
|
||||
// removed when the test ends.
|
||||
|
||||
const containerPassword = "Resolve_Spec_1"
|
||||
|
||||
func containerRuntime(t *testing.T) string {
|
||||
t.Helper()
|
||||
if testing.Short() {
|
||||
t.Skip("container tests are skipped with -short")
|
||||
}
|
||||
if os.Getenv("RESOLVESPEC_TEST_CONTAINERS") != "1" {
|
||||
t.Skip("set RESOLVESPEC_TEST_CONTAINERS=1 to run tests that start a podman/docker container")
|
||||
}
|
||||
for _, rt := range []string{"podman", "docker"} {
|
||||
if p, err := exec.LookPath(rt); err == nil {
|
||||
return p
|
||||
}
|
||||
}
|
||||
t.Skip("neither podman nor docker found in PATH")
|
||||
return ""
|
||||
}
|
||||
|
||||
func run(t *testing.T, timeout time.Duration, name string, args ...string) string {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
var out, errb bytes.Buffer
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
cmd.Stdout, cmd.Stderr = &out, &errb
|
||||
if err := cmd.Run(); err != nil {
|
||||
t.Fatalf("%s %s: %v\n%s", name, strings.Join(args, " "), err, errb.String())
|
||||
}
|
||||
return strings.TrimSpace(out.String())
|
||||
}
|
||||
|
||||
// startContainer runs image publishing containerPort on a random localhost port and returns
|
||||
// the host port. The container is force-removed on cleanup.
|
||||
func startContainer(t *testing.T, rt, image, containerPort string, env map[string]string) string {
|
||||
t.Helper()
|
||||
args := []string{"run", "-d", "--rm", "-p", "127.0.0.1::" + containerPort}
|
||||
for k, v := range env {
|
||||
args = append(args, "-e", k+"="+v)
|
||||
}
|
||||
args = append(args, image)
|
||||
id := run(t, 10*time.Minute, rt, args...) // first run may pull the image
|
||||
t.Cleanup(func() { _ = exec.Command(rt, "rm", "-f", id).Run() })
|
||||
|
||||
// "127.0.0.1:49153" (docker may print one line per address family)
|
||||
out := run(t, 30*time.Second, rt, "port", id, containerPort)
|
||||
line := strings.Fields(out)[len(strings.Fields(out))-1]
|
||||
for _, l := range strings.Split(out, "\n") {
|
||||
if strings.HasPrefix(strings.TrimSpace(l), "127.0.0.1") || strings.Contains(l, " 127.0.0.1:") {
|
||||
line = l[strings.LastIndex(l, " ")+1:]
|
||||
break
|
||||
}
|
||||
}
|
||||
_, port, err := net.SplitHostPort(line)
|
||||
if err != nil {
|
||||
t.Fatalf("cannot parse published port %q: %v", out, err)
|
||||
}
|
||||
return port
|
||||
}
|
||||
|
||||
// waitReady retries until the server accepts queries or the deadline passes.
|
||||
func waitReady(t *testing.T, driver, dsn string, d time.Duration) *sql.DB {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(d)
|
||||
var last error
|
||||
for time.Now().Before(deadline) {
|
||||
db, err := sql.Open(driver, dsn)
|
||||
if err == nil {
|
||||
if last = db.Ping(); last == nil {
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return db
|
||||
}
|
||||
_ = db.Close()
|
||||
} else {
|
||||
last = err
|
||||
}
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Fatalf("database did not become ready within %s: %v", d, last)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestConformancePostgresContainer(t *testing.T) {
|
||||
rt := containerRuntime(t)
|
||||
port := startContainer(t, rt, "docker.io/library/postgres:16-alpine", "5432", map[string]string{"POSTGRES_PASSWORD": containerPassword})
|
||||
dsn := func(db string) string {
|
||||
return fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/%s?sslmode=disable", containerPassword, port, db)
|
||||
}
|
||||
admin := waitReady(t, "pgx", dsn("postgres"), 90*time.Second)
|
||||
// The official image restarts once during init: make sure the second start is the one we use.
|
||||
time.Sleep(2 * time.Second)
|
||||
admin = waitReady(t, "pgx", dsn("postgres"), 60*time.Second)
|
||||
for _, name := range []string{"cf_proc", "cf_direct"} {
|
||||
if _, err := admin.Exec("CREATE DATABASE " + name); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("procedure", func(t *testing.T) {
|
||||
db, err := sql.Open("pgx", dsn("cf_proc"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
for _, f := range []string{"../database_schema.sql", "../keystore_schema.sql"} {
|
||||
b, err := os.ReadFile(f)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.Exec(string(b)); err != nil {
|
||||
t.Fatalf("apply %s: %v", f, err)
|
||||
}
|
||||
}
|
||||
runConformance(t, db, "postgres", lookup.Config{Mode: lookup.ModeProcedure}, true)
|
||||
})
|
||||
t.Run("direct", func(t *testing.T) {
|
||||
runOnServer(t, "pgx", dsn("cf_direct"), "postgres", lookup.Config{Mode: lookup.ModeDirect}, true)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package backends
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Routers: each store method asks the chooser which backend serves that operation.
|
||||
|
||||
type authRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.AuthStore
|
||||
}
|
||||
|
||||
var _ lookup.AuthStore = (*authRouter)(nil)
|
||||
|
||||
func (r *authRouter) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogin, r.c.procs.Login, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Login(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRegister, r.c.procs.Register, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Register(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLogout, r.c.procs.Logout, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Logout(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpSession, r.c.procs.Session, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Session(ctx, token, reference)
|
||||
}
|
||||
|
||||
func (r *authRouter) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpTouchSession, r.c.procs.SessionUpdate, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.TouchSession(ctx, token, user)
|
||||
}
|
||||
|
||||
func (r *authRouter) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpRefresh, r.c.procs.RefreshToken, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Refresh(ctx, refreshToken)
|
||||
}
|
||||
|
||||
func (r *authRouter) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpLoginAPIKey, r.c.procs.LoginAPIKey, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.LoginAPIKey(ctx, rawKey, claims)
|
||||
}
|
||||
|
||||
func (r *authRouter) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogin, r.c.procs.JWTLogin, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.JWTLogin(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpJWTLogout, r.c.procs.JWTLogout, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.JWTLogout(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetRequest, r.c.procs.PasswordResetRequest, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.ResetRequest(ctx, req)
|
||||
}
|
||||
|
||||
func (r *authRouter) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
|
||||
st, err := pick[lookup.AuthStore](r.c, ctx, lookup.OpResetComplete, r.c.procs.PasswordResetComplete, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.ResetComplete(ctx, req)
|
||||
}
|
||||
|
||||
type keysRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.KeyStore
|
||||
}
|
||||
|
||||
var _ lookup.KeyStore = (*keysRouter)(nil)
|
||||
|
||||
func (r *keysRouter) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
|
||||
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyCreate, r.c.procs.KeystoreCreateKey, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Create(ctx, req, keyHash)
|
||||
}
|
||||
|
||||
func (r *keysRouter) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
|
||||
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyList, r.c.procs.KeystoreGetUserKeys, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.List(ctx, userID, keyType)
|
||||
}
|
||||
|
||||
func (r *keysRouter) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
|
||||
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyDelete, r.c.procs.KeystoreDeleteKey, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return st.Delete(ctx, userID, keyID)
|
||||
}
|
||||
|
||||
func (r *keysRouter) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
|
||||
st, err := pick[lookup.KeyStore](r.c, ctx, lookup.OpKeyValidate, r.c.procs.KeystoreValidateKey, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Validate(ctx, keyHash, keyType)
|
||||
}
|
||||
|
||||
type oauthClientRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.OAuthClientStore
|
||||
}
|
||||
|
||||
var _ lookup.OAuthClientStore = (*oauthClientRouter)(nil)
|
||||
|
||||
func (r *oauthClientRouter) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRegisterClient, r.c.procs.OAuthRegisterClient, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.RegisterClient(ctx, client)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthGetClient, r.c.procs.OAuthGetClient, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.GetClient(ctx, clientID)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthSaveCode, r.c.procs.OAuthSaveCode, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.SaveCode(ctx, code)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthExchangeCode, r.c.procs.OAuthExchangeCode, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.ExchangeCode(ctx, code)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthIntrospect, r.c.procs.OAuthIntrospect, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Introspect(ctx, token)
|
||||
}
|
||||
|
||||
func (r *oauthClientRouter) Revoke(ctx context.Context, token string) error {
|
||||
st, err := pick[lookup.OAuthClientStore](r.c, ctx, lookup.OpOAuthRevoke, r.c.procs.OAuthRevoke, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Revoke(ctx, token)
|
||||
}
|
||||
|
||||
type oauthUserRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.OAuthUserStore
|
||||
}
|
||||
|
||||
var _ lookup.OAuthUserStore = (*oauthUserRouter)(nil)
|
||||
|
||||
func (r *oauthUserRouter) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
|
||||
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetOrCreateUser, r.c.procs.OAuthGetOrCreateUser, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return st.GetOrCreateUser(ctx, user, provider)
|
||||
}
|
||||
|
||||
func (r *oauthUserRouter) CreateSession(ctx context.Context, session lookup.OAuthSession) error {
|
||||
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthCreateSession, r.c.procs.OAuthCreateSession, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.CreateSession(ctx, session)
|
||||
}
|
||||
|
||||
func (r *oauthUserRouter) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
|
||||
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetRefreshToken, r.c.procs.OAuthGetRefreshToken, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.GetByRefreshToken(ctx, refreshToken)
|
||||
}
|
||||
|
||||
func (r *oauthUserRouter) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
|
||||
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthUpdateRefreshToken, r.c.procs.OAuthUpdateRefreshToken, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.UpdateRefreshToken(ctx, userID, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken, expiresAt)
|
||||
}
|
||||
|
||||
func (r *oauthUserRouter) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
|
||||
st, err := pick[lookup.OAuthUserStore](r.c, ctx, lookup.OpOAuthGetUser, r.c.procs.OAuthGetUser, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.GetUser(ctx, userID)
|
||||
}
|
||||
|
||||
type passkeyRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.PasskeyStore
|
||||
}
|
||||
|
||||
var _ lookup.PasskeyStore = (*passkeyRouter)(nil)
|
||||
|
||||
func (r *passkeyRouter) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyStore, r.c.procs.PasskeyStoreCredential, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return st.Store(ctx, rec)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) Get(ctx context.Context, credentialID string) (int, uint32, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyGet, r.c.procs.PasskeyGetCredential, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
return st.Get(ctx, credentialID)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyUpdateCounter, r.c.procs.PasskeyUpdateCounter, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return st.UpdateCounter(ctx, credentialID, newCounter)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyList, r.c.procs.PasskeyGetUserCredentials, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.List(ctx, userID)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) Delete(ctx context.Context, userID int, credentialID string) error {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyDelete, r.c.procs.PasskeyDeleteCredential, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Delete(ctx, userID, credentialID)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) Rename(ctx context.Context, userID int, credentialID, name string) error {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyRename, r.c.procs.PasskeyUpdateName, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Rename(ctx, userID, credentialID, name)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyByUsername, r.c.procs.PasskeyGetCredsByUsername, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return st.ByUsername(ctx, username)
|
||||
}
|
||||
|
||||
func (r *passkeyRouter) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
st, err := pick[lookup.PasskeyStore](r.c, ctx, lookup.OpPasskeyLogin, r.c.procs.PasskeyLogin, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.Login(ctx, userID, claims)
|
||||
}
|
||||
|
||||
type totpRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.TOTPStore
|
||||
}
|
||||
|
||||
var _ lookup.TOTPStore = (*totpRouter)(nil)
|
||||
|
||||
func (r *totpRouter) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPEnable, r.c.procs.TOTPEnable, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Enable(ctx, userID, secret, hashedCodes)
|
||||
}
|
||||
|
||||
func (r *totpRouter) Disable(ctx context.Context, userID int) error {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPDisable, r.c.procs.TOTPDisable, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.Disable(ctx, userID)
|
||||
}
|
||||
|
||||
func (r *totpRouter) Status(ctx context.Context, userID int) (bool, error) {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPStatus, r.c.procs.TOTPGetStatus, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return st.Status(ctx, userID)
|
||||
}
|
||||
|
||||
func (r *totpRouter) Secret(ctx context.Context, userID int) (string, error) {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPSecret, r.c.procs.TOTPGetSecret, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return st.Secret(ctx, userID)
|
||||
}
|
||||
|
||||
func (r *totpRouter) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPRegenerateBackup, r.c.procs.TOTPRegenerateBackup, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return st.RegenerateBackupCodes(ctx, userID, hashedCodes)
|
||||
}
|
||||
|
||||
func (r *totpRouter) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
|
||||
st, err := pick[lookup.TOTPStore](r.c, ctx, lookup.OpTOTPValidateBackupCode, r.c.procs.TOTPValidateBackupCode, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return st.ValidateBackupCode(ctx, userID, codeHash)
|
||||
}
|
||||
|
||||
type policyRouter struct {
|
||||
c *chooser
|
||||
proc, direct lookup.PolicyStore
|
||||
}
|
||||
|
||||
var _ lookup.PolicyStore = (*policyRouter)(nil)
|
||||
|
||||
func (r *policyRouter) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
|
||||
st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpColumnSecurity, r.c.procs.ColumnSecurity, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return st.ColumnSecurity(ctx, userID, schema, table)
|
||||
}
|
||||
|
||||
func (r *policyRouter) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
|
||||
st, err := pick[lookup.PolicyStore](r.c, ctx, lookup.OpRowSecurity, r.c.procs.RowSecurity, r.proc, r.direct)
|
||||
if err != nil {
|
||||
return sectypes.RowSecurity{}, err
|
||||
}
|
||||
return st.RowSecurity(ctx, userRef, schema, table)
|
||||
}
|
||||
@@ -0,0 +1,564 @@
|
||||
// Package conformance is the shared behavioural suite every lookup backend must pass.
|
||||
// It only uses the store interfaces, so the same cases run against the direct backend on
|
||||
// every dialect and against the procedure backend on Postgres. Error messages are not
|
||||
// asserted (backends word them differently), only whether an operation succeeds or fails
|
||||
// and the values it returns.
|
||||
//
|
||||
// The suite names everything it creates with Env.Prefix and never assumes empty tables, so
|
||||
// it can run against a shared database. Env.Cleanup, when set, removes the prefixed rows.
|
||||
package conformance
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Env is one backend under test.
|
||||
type Env struct {
|
||||
Provider *lookup.Provider
|
||||
// DB and Dialect are used only to seed policy rules, which have no store method.
|
||||
DB *sql.DB
|
||||
Dialect dialect.Dialect
|
||||
// Prefix makes every created name unique to this run.
|
||||
Prefix string
|
||||
// Cleanup removes rows whose names start with Prefix. Optional.
|
||||
Cleanup func(t *testing.T)
|
||||
}
|
||||
|
||||
// Run executes the suite.
|
||||
func Run(t *testing.T, env Env) {
|
||||
if env.Cleanup != nil {
|
||||
t.Cleanup(func() { env.Cleanup(t) })
|
||||
}
|
||||
s := &suite{Env: env}
|
||||
t.Run("AuthSessionLifecycle", s.authSessionLifecycle)
|
||||
t.Run("AuthRejectsBadCredentials", s.authRejectsBadCredentials)
|
||||
t.Run("RegisterIgnoresPrivileges", s.registerIgnoresPrivileges)
|
||||
t.Run("RegisterRejectsDuplicates", s.registerRejectsDuplicates)
|
||||
t.Run("PasswordReset", s.passwordReset)
|
||||
t.Run("JWT", s.jwt)
|
||||
t.Run("Keys", s.keys)
|
||||
t.Run("LoginAPIKey", s.loginAPIKey)
|
||||
t.Run("OAuthClientAndCodes", s.oauthClientAndCodes)
|
||||
t.Run("OAuthIntrospectRevoke", s.oauthIntrospectRevoke)
|
||||
t.Run("OAuthUsers", s.oauthUsers)
|
||||
t.Run("Passkey", s.passkey)
|
||||
t.Run("TOTP", s.totp)
|
||||
t.Run("Policy", s.policy)
|
||||
}
|
||||
|
||||
type suite struct{ Env }
|
||||
|
||||
var ctx = context.Background()
|
||||
|
||||
func (s *suite) name(n string) string { return s.Prefix + n }
|
||||
|
||||
func (s *suite) register(t *testing.T, n string) *sectypes.LoginResponse {
|
||||
t.Helper()
|
||||
resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{
|
||||
Username: s.name(n), Email: s.name(n) + "@example.test", Password: "pw-" + n,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("register %s: %v", n, err)
|
||||
}
|
||||
if resp == nil || resp.User == nil || resp.Token == "" || resp.User.UserID == 0 {
|
||||
t.Fatalf("register %s: incomplete response %+v", n, resp)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
func rejected(t *testing.T, what string, err error) {
|
||||
t.Helper()
|
||||
if err == nil {
|
||||
t.Fatalf("%s: expected an error", what)
|
||||
}
|
||||
}
|
||||
|
||||
// notOK asserts an operation did not validate: it either failed or returned false.
|
||||
func notOK(t *testing.T, what string, ok bool, err error) {
|
||||
t.Helper()
|
||||
if err == nil && ok {
|
||||
t.Fatalf("%s: accepted", what)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) authSessionLifecycle(t *testing.T) {
|
||||
a := s.Provider.Auth
|
||||
reg := s.register(t, "life")
|
||||
|
||||
login, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("life"), Password: "pw-life",
|
||||
Claims: map[string]any{"ip_address": "10.0.0.1", "user_agent": "conformance"}})
|
||||
if err != nil || login.Token == "" || login.User.UserName != s.name("life") {
|
||||
t.Fatalf("login: %+v %v", login, err)
|
||||
}
|
||||
if login.Token == reg.Token {
|
||||
t.Fatal("login reused the registration session")
|
||||
}
|
||||
|
||||
u, err := a.Session(ctx, login.Token, "authenticate")
|
||||
if err != nil || u.UserName != s.name("life") || u.UserID != reg.User.UserID {
|
||||
t.Fatalf("session: %+v %v", u, err)
|
||||
}
|
||||
if err := a.TouchSession(ctx, login.Token, u); err != nil {
|
||||
t.Fatalf("touch: %v", err)
|
||||
}
|
||||
_, err = a.Session(ctx, s.name("no-such-token"), "authenticate")
|
||||
rejected(t, "unknown session", err)
|
||||
|
||||
ref, err := a.Refresh(ctx, login.Token)
|
||||
if err != nil || ref.Token == "" || ref.Token == login.Token {
|
||||
t.Fatalf("refresh: %+v %v", ref, err)
|
||||
}
|
||||
_, err = a.Session(ctx, login.Token, "")
|
||||
rejected(t, "session after refresh", err)
|
||||
_, err = a.Refresh(ctx, login.Token)
|
||||
rejected(t, "second refresh of the same token", err)
|
||||
|
||||
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err != nil {
|
||||
t.Fatalf("logout: %v", err)
|
||||
}
|
||||
_, err = a.Session(ctx, ref.Token, "")
|
||||
rejected(t, "session after logout", err)
|
||||
}
|
||||
|
||||
func (s *suite) authRejectsBadCredentials(t *testing.T) {
|
||||
a := s.Provider.Auth
|
||||
s.register(t, "creds")
|
||||
_, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds"), Password: "wrong"})
|
||||
rejected(t, "wrong password", err)
|
||||
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("creds")})
|
||||
rejected(t, "empty password", err)
|
||||
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("nobody"), Password: "pw"})
|
||||
rejected(t, "unknown user", err)
|
||||
}
|
||||
|
||||
func (s *suite) registerIgnoresPrivileges(t *testing.T) {
|
||||
resp, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{
|
||||
Username: s.name("priv"), Email: s.name("priv") + "@example.test", Password: "x",
|
||||
UserLevel: 99, Roles: []string{"admin"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if resp.User.UserLevel != 0 || len(resp.User.Roles) != 0 {
|
||||
t.Fatalf("client-supplied privileges honoured: %+v", resp.User)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) registerRejectsDuplicates(t *testing.T) {
|
||||
s.register(t, "dup")
|
||||
_, err := s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup"), Email: s.name("dup2") + "@example.test", Password: "x"})
|
||||
rejected(t, "duplicate username", err)
|
||||
_, err = s.Provider.Auth.Register(ctx, sectypes.RegisterRequest{Username: s.name("dup2"), Email: s.name("dup") + "@example.test", Password: "x"})
|
||||
rejected(t, "duplicate email", err)
|
||||
}
|
||||
|
||||
func (s *suite) passwordReset(t *testing.T) {
|
||||
a := s.Provider.Auth
|
||||
reg := s.register(t, "reset")
|
||||
|
||||
if r, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("nobody") + "@example.test"}); err != nil || (r != nil && r.Token != "") {
|
||||
t.Fatalf("unknown email must succeed without a token (user enumeration): %+v %v", r, err)
|
||||
}
|
||||
req, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: s.name("reset") + "@example.test"})
|
||||
if err != nil || req == nil || req.Token == "" {
|
||||
t.Fatalf("reset request: %+v %v", req, err)
|
||||
}
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: "bogus", NewPassword: "x"}); err == nil {
|
||||
t.Fatal("bogus reset token accepted")
|
||||
}
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "new-pw"}); err != nil {
|
||||
t.Fatalf("reset complete: %v", err)
|
||||
}
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: req.Token, NewPassword: "again"}); err == nil {
|
||||
t.Fatal("reset token reused")
|
||||
}
|
||||
_, err = a.Session(ctx, reg.Token, "")
|
||||
rejected(t, "session surviving a password reset", err)
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "new-pw"}); err != nil {
|
||||
t.Fatalf("login with new password: %v", err)
|
||||
}
|
||||
_, err = a.Login(ctx, sectypes.LoginRequest{Username: s.name("reset"), Password: "pw-reset"})
|
||||
rejected(t, "old password after reset", err)
|
||||
}
|
||||
|
||||
func (s *suite) jwt(t *testing.T) {
|
||||
a := s.Provider.Auth
|
||||
reg := s.register(t, "jwt")
|
||||
resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "pw-jwt"})
|
||||
if err != nil || resp.Token == "" || resp.User.UserID != reg.User.UserID {
|
||||
t.Fatalf("jwt login: %+v %v", resp, err)
|
||||
}
|
||||
_, err = a.JWTLogin(ctx, sectypes.LoginRequest{Username: s.name("jwt"), Password: "bad"})
|
||||
rejected(t, "jwt login with wrong password", err)
|
||||
if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: s.name("jwt-tok"), UserID: reg.User.UserID}); err != nil {
|
||||
t.Fatalf("jwt logout: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) createKey(t *testing.T, uid int, typ sectypes.KeyType, raw string, exp *time.Time) *sectypes.UserKey {
|
||||
t.Helper()
|
||||
k, err := s.Provider.Keys.Create(ctx, sectypes.CreateKeyRequest{UserID: uid, KeyType: typ, Name: s.name("key"),
|
||||
Scopes: []string{"read"}, ExpiresAt: exp}, sectypes.HashKey(raw))
|
||||
if err != nil || k == nil || k.ID == 0 {
|
||||
t.Fatalf("create key: %+v %v", k, err)
|
||||
}
|
||||
return k
|
||||
}
|
||||
|
||||
func (s *suite) keys(t *testing.T) {
|
||||
k := s.Provider.Keys
|
||||
uid := s.register(t, "keys").User.UserID
|
||||
raw := s.name("raw-keys")
|
||||
created := s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, raw, nil)
|
||||
s.createKey(t, uid, sectypes.KeyTypeJWTSecret, s.name("raw-keys-jwt"), nil)
|
||||
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, s.name("raw-keys-old"), ptr(time.Now().Add(-time.Hour)))
|
||||
|
||||
all, err := k.List(ctx, uid, "")
|
||||
if err != nil || len(all) != 2 {
|
||||
t.Fatalf("list must hide expired keys: %d %v", len(all), err)
|
||||
}
|
||||
one, err := k.List(ctx, uid, sectypes.KeyTypeHeaderAPI)
|
||||
if err != nil || len(one) != 1 || one[0].ID != created.ID || len(one[0].Scopes) != 1 {
|
||||
t.Fatalf("typed list: %+v %v", one, err)
|
||||
}
|
||||
|
||||
got, err := k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeHeaderAPI)
|
||||
if err != nil || got.UserID != uid {
|
||||
t.Fatalf("validate: %+v %v", got, err)
|
||||
}
|
||||
_, err = k.Validate(ctx, sectypes.HashKey(raw), sectypes.KeyTypeGenericAPI)
|
||||
rejected(t, "wrong key type", err)
|
||||
_, err = k.Validate(ctx, sectypes.HashKey(s.name("raw-keys-old")), "")
|
||||
rejected(t, "expired key", err)
|
||||
_, err = k.Validate(ctx, sectypes.HashKey(s.name("unknown")), "")
|
||||
rejected(t, "unknown key", err)
|
||||
|
||||
_, err = k.Delete(ctx, uid+1_000_000, created.ID)
|
||||
rejected(t, "deleting another user's key", err)
|
||||
if _, err := k.Delete(ctx, uid, created.ID); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
_, err = k.Delete(ctx, uid, created.ID)
|
||||
rejected(t, "deleting twice", err)
|
||||
_, err = k.Validate(ctx, sectypes.HashKey(raw), "")
|
||||
rejected(t, "deleted key", err)
|
||||
}
|
||||
|
||||
func (s *suite) loginAPIKey(t *testing.T) {
|
||||
a := s.Provider.Auth
|
||||
uid := s.register(t, "apikey").User.UserID
|
||||
good, generic, jwtKey, off, old := s.name("ak-good"), s.name("ak-generic"), s.name("ak-jwt"), s.name("ak-off"), s.name("ak-old")
|
||||
s.createKey(t, uid, sectypes.KeyTypeHeaderAPI, good, nil)
|
||||
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, generic, nil)
|
||||
s.createKey(t, uid, sectypes.KeyTypeJWTSecret, jwtKey, nil)
|
||||
inactive := s.createKey(t, uid, sectypes.KeyTypeGenericAPI, off, nil)
|
||||
if _, err := s.Provider.Keys.Delete(ctx, uid, inactive.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s.createKey(t, uid, sectypes.KeyTypeGenericAPI, old, ptr(time.Now().Add(-time.Hour)))
|
||||
|
||||
for _, raw := range []string{good, generic} {
|
||||
resp, err := a.LoginAPIKey(ctx, raw, map[string]any{"ip_address": "10.0.0.2"})
|
||||
if err != nil || resp.User.UserName != s.name("apikey") || resp.Token == "" {
|
||||
t.Fatalf("api key login: %+v %v", resp, err)
|
||||
}
|
||||
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
|
||||
t.Fatalf("session from api key login: %v", err)
|
||||
}
|
||||
}
|
||||
for _, raw := range []string{"", s.name("ak-missing"), jwtKey, off, old} {
|
||||
_, err := a.LoginAPIKey(ctx, raw, nil)
|
||||
if !errors.Is(err, lookup.ErrInvalidAPIKey) {
|
||||
t.Fatalf("key %q: want ErrInvalidAPIKey, got %v", raw, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) oauthClientAndCodes(t *testing.T) {
|
||||
c := s.Provider.OAuthClient
|
||||
cid := s.name("client")
|
||||
reg, err := c.RegisterClient(ctx, §ypes.OAuthServerClient{ClientID: cid, RedirectURIs: []string{"https://app.example.test/cb"}, ClientName: "App"})
|
||||
if err != nil || reg.ClientID != cid {
|
||||
t.Fatalf("register client: %+v %v", reg, err)
|
||||
}
|
||||
got, err := c.GetClient(ctx, cid)
|
||||
if err != nil || got.ClientName != "App" || len(got.RedirectURIs) != 1 || got.RedirectURIs[0] != "https://app.example.test/cb" {
|
||||
t.Fatalf("get client: %+v %v", got, err)
|
||||
}
|
||||
_, err = c.GetClient(ctx, s.name("no-client"))
|
||||
rejected(t, "unknown client", err)
|
||||
|
||||
code := §ypes.OAuthCode{Code: s.name("code1"), ClientID: cid, RedirectURI: "https://app.example.test/cb",
|
||||
CodeChallenge: "challenge", SessionToken: s.name("sess"), Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)}
|
||||
if err := c.SaveCode(ctx, code); err != nil {
|
||||
t.Fatalf("save code: %v", err)
|
||||
}
|
||||
ex, err := c.ExchangeCode(ctx, code.Code)
|
||||
if err != nil || ex.Code != code.Code || ex.ClientID != cid || ex.SessionToken != code.SessionToken || len(ex.Scopes) != 1 {
|
||||
t.Fatalf("exchange: %+v %v", ex, err)
|
||||
}
|
||||
_, err = c.ExchangeCode(ctx, code.Code)
|
||||
rejected(t, "code reuse", err)
|
||||
|
||||
expired := *code
|
||||
expired.Code, expired.ExpiresAt = s.name("code2"), time.Now().Add(-time.Minute)
|
||||
if err := c.SaveCode(ctx, &expired); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = c.ExchangeCode(ctx, expired.Code)
|
||||
rejected(t, "expired code", err)
|
||||
}
|
||||
|
||||
func (s *suite) oauthIntrospectRevoke(t *testing.T) {
|
||||
c := s.Provider.OAuthClient
|
||||
reg := s.register(t, "intro")
|
||||
info, err := c.Introspect(ctx, reg.Token)
|
||||
if err != nil || !info.Active || info.Username != s.name("intro") {
|
||||
t.Fatalf("introspect: %+v %v", info, err)
|
||||
}
|
||||
if err := c.Revoke(ctx, reg.Token); err != nil {
|
||||
t.Fatalf("revoke: %v", err)
|
||||
}
|
||||
if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active {
|
||||
t.Fatalf("revoked token still active: %+v %v", info, err)
|
||||
}
|
||||
if err := c.Revoke(ctx, s.name("unknown-token")); err != nil {
|
||||
t.Fatalf("revoking an unknown token must succeed (RFC 7009): %v", err)
|
||||
}
|
||||
if info, err := c.Introspect(ctx, s.name("unknown-token")); err != nil || info.Active {
|
||||
t.Fatalf("unknown token: %+v %v", info, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) oauthUsers(t *testing.T) {
|
||||
o := s.Provider.OAuthUser
|
||||
id, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github")
|
||||
if err != nil || id == 0 {
|
||||
t.Fatalf("get or create: %d %v", id, err)
|
||||
}
|
||||
again, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: s.name("gh"), Email: s.name("gh") + "@example.test", RemoteID: s.name("remote-1")}, "github")
|
||||
if err != nil || again != id {
|
||||
t.Fatalf("second login must return the same user: %d %v", again, err)
|
||||
}
|
||||
|
||||
exp := time.Now().Add(time.Hour)
|
||||
sess := lookup.OAuthSession{SessionToken: s.name("os1"), UserID: id, AccessToken: "a1", RefreshToken: s.name("or1"), TokenType: "Bearer", ExpiresAt: exp, Provider: "github"}
|
||||
if err := o.CreateSession(ctx, sess); err != nil {
|
||||
t.Fatalf("create session: %v", err)
|
||||
}
|
||||
ref, err := o.GetByRefreshToken(ctx, sess.RefreshToken)
|
||||
if err != nil || ref.UserID != id || ref.AccessToken != "a1" {
|
||||
t.Fatalf("by refresh token: %+v %v", ref, err)
|
||||
}
|
||||
_, err = o.GetByRefreshToken(ctx, s.name("or-missing"))
|
||||
rejected(t, "unknown refresh token", err)
|
||||
if err := o.UpdateRefreshToken(ctx, id, sess.RefreshToken, s.name("os2"), "a2", s.name("or2"), exp); err != nil {
|
||||
t.Fatalf("update refresh token: %v", err)
|
||||
}
|
||||
if _, err := o.GetByRefreshToken(ctx, s.name("or2")); err != nil {
|
||||
t.Fatalf("rotated refresh token not found: %v", err)
|
||||
}
|
||||
u, err := o.GetUser(ctx, id)
|
||||
if err != nil || u.UserName != s.name("gh") {
|
||||
t.Fatalf("get user: %+v %v", u, err)
|
||||
}
|
||||
_, err = o.GetUser(ctx, id+1_000_000)
|
||||
rejected(t, "unknown user", err)
|
||||
}
|
||||
|
||||
func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) }
|
||||
|
||||
func (s *suite) passkey(t *testing.T) {
|
||||
p := s.Provider.Passkey
|
||||
reg := s.register(t, "pk")
|
||||
uid := reg.User.UserID
|
||||
c1, c2 := b64(s.name("cred1")), b64(s.name("cred2"))
|
||||
|
||||
rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: c1, PublicKey: b64("pubkey"), AttestationType: "none",
|
||||
Transports: []string{"usb", "nfc"}, Name: "Key 1"}
|
||||
if id, err := p.Store(ctx, rec); err != nil || id == 0 {
|
||||
t.Fatalf("store: %d %v", id, err)
|
||||
}
|
||||
_, err := p.Store(ctx, rec)
|
||||
rejected(t, "duplicate credential", err)
|
||||
rec.CredentialID, rec.Name = c2, "Key 2"
|
||||
if _, err := p.Store(ctx, rec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
owner, count, err := p.Get(ctx, c1)
|
||||
if err != nil || owner != uid || count != 0 {
|
||||
t.Fatalf("get: %d %d %v", owner, count, err)
|
||||
}
|
||||
_, _, err = p.Get(ctx, b64(s.name("missing")))
|
||||
rejected(t, "unknown credential", err)
|
||||
|
||||
if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || clone {
|
||||
t.Fatalf("advance counter: clone=%v %v", clone, err)
|
||||
}
|
||||
if clone, err := p.UpdateCounter(ctx, c1, 5); err != nil || !clone {
|
||||
t.Fatalf("replayed counter must raise a clone warning: clone=%v %v", clone, err)
|
||||
}
|
||||
|
||||
list, err := p.List(ctx, uid)
|
||||
if err != nil || len(list) != 2 {
|
||||
t.Fatalf("list: %d %v", len(list), err)
|
||||
}
|
||||
if err := p.Rename(ctx, uid, c1, "Renamed"); err != nil {
|
||||
t.Fatalf("rename: %v", err)
|
||||
}
|
||||
rejected(t, "renaming another user's credential", p.Rename(ctx, uid+1_000_000, c1, "x"))
|
||||
|
||||
gotID, refs, err := p.ByUsername(ctx, s.name("pk"))
|
||||
if err != nil || gotID != uid || len(refs) != 2 {
|
||||
t.Fatalf("by username: %d %+v %v", gotID, refs, err)
|
||||
}
|
||||
_, _, err = p.ByUsername(ctx, s.name("ghost"))
|
||||
rejected(t, "unknown username", err)
|
||||
|
||||
resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "10.0.0.3"})
|
||||
if err != nil || resp.Token == "" || resp.User.UserName != s.name("pk") {
|
||||
t.Fatalf("passkey login: %+v %v", resp, err)
|
||||
}
|
||||
if _, err := s.Provider.Auth.Session(ctx, resp.Token, ""); err != nil {
|
||||
t.Fatalf("session from passkey login: %v", err)
|
||||
}
|
||||
|
||||
rejected(t, "deleting another user's credential", p.Delete(ctx, uid+1_000_000, c1))
|
||||
if err := p.Delete(ctx, uid, c1); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
rejected(t, "deleting twice", p.Delete(ctx, uid, c1))
|
||||
}
|
||||
|
||||
func (s *suite) totp(t *testing.T) {
|
||||
st := s.Provider.TOTP
|
||||
uid := s.register(t, "totp").User.UserID
|
||||
|
||||
if on, err := st.Status(ctx, uid); err != nil || on {
|
||||
t.Fatalf("initial status: %v %v", on, err)
|
||||
}
|
||||
_, err := st.Secret(ctx, uid)
|
||||
rejected(t, "secret without 2FA", err)
|
||||
|
||||
if err := st.Enable(ctx, uid, "SECRET", []string{s.name("h1"), s.name("h2")}); err != nil {
|
||||
t.Fatalf("enable: %v", err)
|
||||
}
|
||||
if on, _ := st.Status(ctx, uid); !on {
|
||||
t.Fatal("not enabled")
|
||||
}
|
||||
if sec, err := st.Secret(ctx, uid); err != nil || sec != "SECRET" {
|
||||
t.Fatalf("secret: %q %v", sec, err)
|
||||
}
|
||||
|
||||
if ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1")); err != nil || !ok {
|
||||
t.Fatalf("backup code: %v %v", ok, err)
|
||||
}
|
||||
ok, err := st.ValidateBackupCode(ctx, uid, s.name("h1"))
|
||||
notOK(t, "backup code reuse", ok, err)
|
||||
ok, err = st.ValidateBackupCode(ctx, uid, s.name("nope"))
|
||||
notOK(t, "unknown backup code", ok, err)
|
||||
|
||||
if err := st.RegenerateBackupCodes(ctx, uid, []string{s.name("n1")}); err != nil {
|
||||
t.Fatalf("regenerate: %v", err)
|
||||
}
|
||||
ok, err = st.ValidateBackupCode(ctx, uid, s.name("h2"))
|
||||
notOK(t, "old backup code after regenerate", ok, err)
|
||||
if ok, err := st.ValidateBackupCode(ctx, uid, s.name("n1")); err != nil || !ok {
|
||||
t.Fatalf("new backup code: %v %v", ok, err)
|
||||
}
|
||||
|
||||
if err := st.Disable(ctx, uid); err != nil {
|
||||
t.Fatalf("disable: %v", err)
|
||||
}
|
||||
if on, _ := st.Status(ctx, uid); on {
|
||||
t.Fatal("still enabled after disable")
|
||||
}
|
||||
}
|
||||
|
||||
// seed inserts one row with dialect placeholders. Values are bound, booleans converted.
|
||||
func (s *suite) seed(t *testing.T, table string, cols []string, vals ...any) {
|
||||
t.Helper()
|
||||
ph := make([]string, len(vals))
|
||||
args := make([]any, len(vals))
|
||||
for i, v := range vals {
|
||||
ph[i] = s.Dialect.Placeholder(i + 1)
|
||||
if b, ok := v.(bool); ok {
|
||||
v = s.Dialect.Bool(b)
|
||||
}
|
||||
args[i] = v
|
||||
}
|
||||
q := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", table, strings.Join(cols, ", "), strings.Join(ph, ", "))
|
||||
if _, err := s.DB.ExecContext(ctx, q, args...); err != nil {
|
||||
t.Fatalf("seed %s: %v", table, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *suite) policy(t *testing.T) {
|
||||
p := s.Provider.Policy
|
||||
u1 := s.register(t, "pol1").User.UserID
|
||||
u2 := s.register(t, "pol2").User.UserID
|
||||
group := 7_000_000 + u1
|
||||
schema, users, orders, secret := s.name("pub"), "Users", "orders", "secret"
|
||||
|
||||
s.seed(t, "sec_group_members", []string{"group_id", "user_id"}, group, u1)
|
||||
colCols := []string{"user_id", "group_id", "schema_name", "table_name", "column_path", "access_type", "is_active"}
|
||||
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, users, "email", "mask", true)
|
||||
s.seed(t, "sec_column_rules", colCols, nil, group, schema, strings.ToLower(users), "profile.ssn", "hide", true)
|
||||
s.seed(t, "sec_column_rules", colCols, u2, nil, schema, strings.ToLower(users), "other", "hide", true)
|
||||
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, strings.ToLower(users), "inactive", "hide", false)
|
||||
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, orders, "x", "hide", true)
|
||||
s.seed(t, "sec_column_rules", colCols, u1, nil, schema, "users_archive", "y", "hide", true)
|
||||
|
||||
rules, err := p.ColumnSecurity(ctx, u1, schema, "users")
|
||||
if err != nil || len(rules) != 2 {
|
||||
t.Fatalf("column rules (user + group, exact table, active only): %d %v %+v", len(rules), err, rules)
|
||||
}
|
||||
paths := map[string]bool{}
|
||||
for _, r := range rules {
|
||||
paths[strings.Join(r.Path, ".")] = true
|
||||
}
|
||||
if !paths["email"] || !paths["profile.ssn"] {
|
||||
t.Fatalf("paths: %v", paths)
|
||||
}
|
||||
if r, err := p.ColumnSecurity(ctx, u2, schema, "users"); err != nil || len(r) != 1 {
|
||||
t.Fatalf("other user's rules: %d %v", len(r), err)
|
||||
}
|
||||
if r, err := p.ColumnSecurity(ctx, u1+u2+1_000_000, schema, "users"); err != nil || len(r) != 0 {
|
||||
t.Fatalf("no rules must be empty, not an error: %d %v", len(r), err)
|
||||
}
|
||||
|
||||
rowCols := []string{"user_id", "group_id", "schema_name", "table_name", "template", "has_block", "is_active"}
|
||||
s.seed(t, "sec_row_rules", rowCols, u1, nil, schema, orders, "owner_id = {UserID}", false, true)
|
||||
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "region = 1", false, true)
|
||||
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, orders, "ignored = 1", false, false)
|
||||
s.seed(t, "sec_row_rules", rowCols, u2, nil, schema, secret, nil, true, true)
|
||||
s.seed(t, "sec_row_rules", rowCols, nil, group, schema, secret, "x = 1", false, true)
|
||||
|
||||
rs, err := p.RowSecurity(ctx, u1, schema, orders)
|
||||
if err != nil || rs.HasBlock || !strings.Contains(rs.Template, "owner_id = {UserID}") || !strings.Contains(rs.Template, "region = 1") || strings.Contains(rs.Template, "ignored") {
|
||||
t.Fatalf("row template: %+v %v", rs, err)
|
||||
}
|
||||
if rs, err := p.RowSecurity(ctx, u2, schema, secret); err != nil || !rs.HasBlock {
|
||||
t.Fatalf("blocking rule must win: %+v %v", rs, err)
|
||||
}
|
||||
if rs, err := p.RowSecurity(ctx, u1+u2+1_000_000, schema, orders); err != nil || rs.HasBlock || rs.Template != "" {
|
||||
t.Fatalf("no rules: %+v %v", rs, err)
|
||||
}
|
||||
if _, err := p.RowSecurity(ctx, "not-a-number", schema, orders); err == nil {
|
||||
t.Fatal("non-numeric user reference accepted (must fail closed)")
|
||||
}
|
||||
}
|
||||
|
||||
func ptr[T any](v T) *T { return &v }
|
||||
@@ -0,0 +1,44 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
// FromDatabase extracts the *sql.DB and the dialect name from an application's
|
||||
// common.Database (bun, gorm or pgsql adapter), so callers do not have to dig the
|
||||
// connection out or set Config.Dialect by hand. The returned name is the adapter's
|
||||
// normalised DriverName ("postgres", "sqlite", "mssql", "mysql") and is empty when
|
||||
// the adapter reports a driver the dialect registry does not know; set
|
||||
// Config.Dialect explicitly in that case.
|
||||
//
|
||||
// Transaction adapters do not expose a *sql.DB and are rejected.
|
||||
func FromDatabase(db common.Database) (*sql.DB, string, error) {
|
||||
if db == nil {
|
||||
return nil, "", fmt.Errorf("lookup: nil database")
|
||||
}
|
||||
p, ok := db.(common.SQLDBProvider)
|
||||
if !ok {
|
||||
return nil, "", fmt.Errorf("lookup: %T does not expose a *sql.DB (transaction adapter or unsupported adapter)", db)
|
||||
}
|
||||
sqlDB := p.SQLDB()
|
||||
if sqlDB == nil {
|
||||
return nil, "", fmt.Errorf("lookup: %T has no *sql.DB", db)
|
||||
}
|
||||
name := db.DriverName()
|
||||
if _, err := dialect.Get(name); err != nil {
|
||||
name = ""
|
||||
}
|
||||
return sqlDB, name, nil
|
||||
}
|
||||
|
||||
// ResolveDialect returns the dialect for db: the configured one, or detected from the driver.
|
||||
func (c Config) ResolveDialect(db *sql.DB) (dialect.Dialect, error) {
|
||||
if c.Dialect != "" {
|
||||
return dialect.Get(c.Dialect)
|
||||
}
|
||||
return dialect.Detect(db)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,84 @@
|
||||
package lookup_test
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||
"github.com/uptrace/bun/driver/sqliteshim"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
func TestFromDatabase(t *testing.T) {
|
||||
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:?cache=shared")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer sqldb.Close()
|
||||
|
||||
gdb, err := gorm.Open(sqlite.Open("file::memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
cases := map[string]struct {
|
||||
db common.Database
|
||||
same *sql.DB // expected handle; nil = just must be non-nil
|
||||
}{
|
||||
"pgsql": {database.NewPgSQLAdapter(sqldb, "sqlite"), sqldb},
|
||||
"bun": {database.NewBunAdapter(bun.NewDB(sqldb, sqlitedialect.New())), sqldb},
|
||||
"gorm": {database.NewGormAdapter(gdb), nil},
|
||||
}
|
||||
for name, c := range cases {
|
||||
got, dialectName, err := lookup.FromDatabase(c.db)
|
||||
if err != nil {
|
||||
t.Errorf("%s: %v", name, err)
|
||||
continue
|
||||
}
|
||||
if got == nil || (c.same != nil && got != c.same) {
|
||||
t.Errorf("%s: unexpected *sql.DB %v", name, got)
|
||||
}
|
||||
if dialectName != "sqlite" {
|
||||
t.Errorf("%s: dialect = %q, want sqlite", name, dialectName)
|
||||
}
|
||||
if err := got.Ping(); err != nil {
|
||||
t.Errorf("%s: handle not usable: %v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromDatabaseRejects(t *testing.T) {
|
||||
if _, _, err := lookup.FromDatabase(nil); err == nil {
|
||||
t.Error("nil database should fail")
|
||||
}
|
||||
// A database that does not expose a *sql.DB (the embedded interface is nil; only the type matters).
|
||||
if _, _, err := lookup.FromDatabase(struct{ common.Database }{}); err == nil {
|
||||
t.Error("adapter without SQLDB should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveDialect(t *testing.T) {
|
||||
sqldb, err := sql.Open(sqliteshim.ShimName, "file::memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer sqldb.Close()
|
||||
|
||||
d, err := lookup.Config{}.ResolveDialect(sqldb)
|
||||
if err != nil || d.Name() != "sqlite" {
|
||||
t.Errorf("detected = %v, %v", d, err)
|
||||
}
|
||||
d, err = lookup.Config{Dialect: "mysql"}.ResolveDialect(sqldb)
|
||||
if err != nil || d.Name() != "mysql" {
|
||||
t.Errorf("explicit dialect should win: %v, %v", d, err)
|
||||
}
|
||||
if _, err := (lookup.Config{Dialect: "oracle"}).Resolve(); err == nil {
|
||||
t.Error("unknown dialect should fail Resolve")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
// Package ddl holds the reference table schemas for the lookup direct backend, one per
|
||||
// dialect. They use the lookup.DefaultSchema table and column names; copy and adapt them
|
||||
// when you override names through lookup.Config.Schema.
|
||||
//
|
||||
// The Postgres file creates tables only. The stored-procedure schema
|
||||
// (lookup/database_schema.sql) is a separate script with native bytea / text[] columns and
|
||||
// must not be combined with it.
|
||||
package ddl
|
||||
|
||||
import (
|
||||
"embed"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
//go:embed postgres.sql sqlite.sql mysql.sql mssql.sql
|
||||
var files embed.FS
|
||||
|
||||
// SQL returns the schema script for a dialect name ("postgres", "sqlite", "mysql", "mssql").
|
||||
func SQL(dialect string) (string, error) {
|
||||
b, err := files.ReadFile(dialect + ".sql")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("ddl: no reference schema for dialect %q", dialect)
|
||||
}
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// Statements returns the schema as separate statements, for drivers that reject
|
||||
// multi-statement execution. Comment-only lines are dropped.
|
||||
func Statements(dialect string) ([]string, error) {
|
||||
s, err := SQL(dialect)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out []string
|
||||
var cur strings.Builder
|
||||
for _, line := range strings.Split(s, "\n") {
|
||||
t := strings.TrimSpace(line)
|
||||
if t == "" || strings.HasPrefix(t, "--") {
|
||||
continue
|
||||
}
|
||||
cur.WriteString(line)
|
||||
cur.WriteString("\n")
|
||||
if strings.HasSuffix(t, ";") {
|
||||
out = append(out, strings.TrimSpace(cur.String()))
|
||||
cur.Reset()
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package ddl_test
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
_ "github.com/glebarez/go-sqlite"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
|
||||
)
|
||||
|
||||
func TestStatementsAllDialects(t *testing.T) {
|
||||
for _, d := range []string{"postgres", "sqlite", "mysql", "mssql"} {
|
||||
st, err := ddl.Statements(d)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(st) < 12 {
|
||||
t.Errorf("%s: %d statements", d, len(st))
|
||||
}
|
||||
all := strings.Join(st, "\n")
|
||||
for _, tbl := range []string{"users", "user_sessions", "user_keys", "oauth_codes", "sec_group_members", "sec_column_rules", "sec_row_rules"} {
|
||||
if !strings.Contains(all, tbl+" (") {
|
||||
t.Errorf("%s: missing table %s", d, tbl)
|
||||
}
|
||||
}
|
||||
}
|
||||
if _, err := ddl.SQL("oracle"); err == nil {
|
||||
t.Error("unknown dialect must error")
|
||||
}
|
||||
}
|
||||
|
||||
// Every table and column the default schema names must exist in the sqlite reference DDL.
|
||||
func TestSQLiteMatchesDefaultSchema(t *testing.T) {
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
db.SetMaxOpenConns(1)
|
||||
s, _ := ddl.SQL("sqlite")
|
||||
if _, err := db.Exec(s); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.Exec(s); err != nil {
|
||||
t.Fatalf("script must be re-runnable: %v", err)
|
||||
}
|
||||
sc := lookup.DefaultSchema()
|
||||
for _, tbl := range sc {
|
||||
for _, c := range tbl.Columns {
|
||||
if _, err := db.Exec("SELECT " + c + " FROM " + tbl.Name + " WHERE 1=0"); err != nil {
|
||||
t.Errorf("%s.%s: %v", tbl.Name, c, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
-- Reference schema for the lookup direct backend: Microsoft SQL Server 2016+. Direct backend only. Run each statement separately (see ddl.Statements).
|
||||
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
|
||||
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
|
||||
|
||||
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
|
||||
-- roles: comma-separated roles
|
||||
|
||||
IF OBJECT_ID(N'users', N'U') IS NULL
|
||||
CREATE TABLE users (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
username NVARCHAR(255) NOT NULL UNIQUE,
|
||||
email NVARCHAR(255) NOT NULL UNIQUE,
|
||||
password NVARCHAR(255),
|
||||
user_level INT DEFAULT 0,
|
||||
roles NVARCHAR(500),
|
||||
is_active BIT DEFAULT 1,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
updated_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
last_login_at DATETIME2,
|
||||
program_user_id INT DEFAULT 0,
|
||||
program_user_table NVARCHAR(255) DEFAULT '',
|
||||
remote_id NVARCHAR(255),
|
||||
auth_provider NVARCHAR(50),
|
||||
totp_secret NVARCHAR(255),
|
||||
totp_enabled BIT DEFAULT 0,
|
||||
totp_enabled_at DATETIME2
|
||||
);
|
||||
|
||||
|
||||
IF OBJECT_ID(N'user_sessions', N'U') IS NULL
|
||||
CREATE TABLE user_sessions (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
session_token NVARCHAR(450) NOT NULL UNIQUE,
|
||||
user_id INT NOT NULL,
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
last_activity_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
ip_address NVARCHAR(45),
|
||||
user_agent NVARCHAR(MAX),
|
||||
access_token NVARCHAR(MAX),
|
||||
refresh_token NVARCHAR(MAX),
|
||||
token_type NVARCHAR(50) DEFAULT 'Bearer',
|
||||
auth_provider NVARCHAR(50),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_user_id' AND object_id = OBJECT_ID(N'user_sessions'))
|
||||
CREATE INDEX idx_user_sessions_user_id ON user_sessions(user_id);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_sessions_expires_at' AND object_id = OBJECT_ID(N'user_sessions'))
|
||||
CREATE INDEX idx_user_sessions_expires_at ON user_sessions(expires_at);
|
||||
|
||||
|
||||
IF OBJECT_ID(N'token_blacklist', N'U') IS NULL
|
||||
CREATE TABLE token_blacklist (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
token NVARCHAR(500) NOT NULL,
|
||||
user_id INT,
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- code_hash: SHA-256 hex of the backup code
|
||||
|
||||
IF OBJECT_ID(N'user_totp_backup_codes', N'U') IS NULL
|
||||
CREATE TABLE user_totp_backup_codes (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
code_hash NVARCHAR(64) NOT NULL,
|
||||
used BIT DEFAULT 0,
|
||||
used_at DATETIME2,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_user_id' AND object_id = OBJECT_ID(N'user_totp_backup_codes'))
|
||||
CREATE INDEX idx_totp_user_id ON user_totp_backup_codes(user_id);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_totp_code_hash' AND object_id = OBJECT_ID(N'user_totp_backup_codes'))
|
||||
CREATE INDEX idx_totp_code_hash ON user_totp_backup_codes(code_hash);
|
||||
|
||||
|
||||
-- credential_id: base64 text
|
||||
-- public_key: base64 text
|
||||
-- aaguid: base64 text
|
||||
-- transports: JSON-encoded array
|
||||
|
||||
IF OBJECT_ID(N'user_passkey_credentials', N'U') IS NULL
|
||||
CREATE TABLE user_passkey_credentials (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
credential_id VARCHAR(900) NOT NULL UNIQUE,
|
||||
public_key NVARCHAR(MAX) NOT NULL,
|
||||
attestation_type NVARCHAR(50) DEFAULT 'none',
|
||||
aaguid NVARCHAR(MAX),
|
||||
sign_count INT DEFAULT 0,
|
||||
clone_warning BIT DEFAULT 0,
|
||||
transports NVARCHAR(MAX),
|
||||
backup_eligible BIT DEFAULT 0,
|
||||
backup_state BIT DEFAULT 0,
|
||||
name NVARCHAR(255),
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
last_used_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_passkey_user_id' AND object_id = OBJECT_ID(N'user_passkey_credentials'))
|
||||
CREATE INDEX idx_passkey_user_id ON user_passkey_credentials(user_id);
|
||||
|
||||
|
||||
-- token_hash: SHA-256 hex of the raw token
|
||||
|
||||
IF OBJECT_ID(N'user_password_resets', N'U') IS NULL
|
||||
CREATE TABLE user_password_resets (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
token_hash NVARCHAR(64) NOT NULL UNIQUE,
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
used BIT DEFAULT 0,
|
||||
used_at DATETIME2,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_user_id' AND object_id = OBJECT_ID(N'user_password_resets'))
|
||||
CREATE INDEX idx_pw_reset_user_id ON user_password_resets(user_id);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_pw_reset_expires_at' AND object_id = OBJECT_ID(N'user_password_resets'))
|
||||
CREATE INDEX idx_pw_reset_expires_at ON user_password_resets(expires_at);
|
||||
|
||||
|
||||
-- redirect_uris: JSON-encoded array
|
||||
-- grant_types: JSON-encoded array
|
||||
-- allowed_scopes: JSON-encoded array
|
||||
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
|
||||
|
||||
IF OBJECT_ID(N'oauth_clients', N'U') IS NULL
|
||||
CREATE TABLE oauth_clients (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
client_id NVARCHAR(255) NOT NULL UNIQUE,
|
||||
redirect_uris NVARCHAR(MAX) NOT NULL,
|
||||
client_name NVARCHAR(255),
|
||||
grant_types NVARCHAR(MAX),
|
||||
allowed_scopes NVARCHAR(MAX),
|
||||
client_secret_hash NVARCHAR(MAX),
|
||||
token_endpoint_auth_method NVARCHAR(30) DEFAULT 'none',
|
||||
is_active BIT DEFAULT 1,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
||||
);
|
||||
|
||||
|
||||
-- scopes: JSON-encoded array
|
||||
|
||||
IF OBJECT_ID(N'oauth_codes', N'U') IS NULL
|
||||
CREATE TABLE oauth_codes (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
code NVARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id NVARCHAR(255) NOT NULL,
|
||||
redirect_uri NVARCHAR(MAX) NOT NULL,
|
||||
client_state NVARCHAR(MAX),
|
||||
code_challenge NVARCHAR(255) NOT NULL,
|
||||
code_challenge_method NVARCHAR(10) DEFAULT 'S256',
|
||||
session_token NVARCHAR(MAX) NOT NULL,
|
||||
refresh_token NVARCHAR(MAX),
|
||||
scopes NVARCHAR(MAX),
|
||||
expires_at DATETIME2 NOT NULL,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME()
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_oauth_codes_expires' AND object_id = OBJECT_ID(N'oauth_codes'))
|
||||
CREATE INDEX idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
IF OBJECT_ID(N'user_keys', N'U') IS NULL
|
||||
CREATE TABLE user_keys (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
key_type NVARCHAR(50) NOT NULL,
|
||||
key_hash NVARCHAR(64) NOT NULL UNIQUE,
|
||||
name NVARCHAR(255) NOT NULL DEFAULT '',
|
||||
scopes NVARCHAR(MAX),
|
||||
meta NVARCHAR(MAX),
|
||||
expires_at DATETIME2,
|
||||
created_at DATETIME2 DEFAULT SYSUTCDATETIME(),
|
||||
last_used_at DATETIME2,
|
||||
is_active BIT DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_user_id' AND object_id = OBJECT_ID(N'user_keys'))
|
||||
CREATE INDEX idx_user_keys_user_id ON user_keys(user_id);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_user_keys_key_type' AND object_id = OBJECT_ID(N'user_keys'))
|
||||
CREATE INDEX idx_user_keys_key_type ON user_keys(key_type);
|
||||
|
||||
|
||||
-- Optional: omit to use per-user rules only.
|
||||
|
||||
IF OBJECT_ID(N'sec_group_members', N'U') IS NULL
|
||||
CREATE TABLE sec_group_members (
|
||||
group_id INT NOT NULL,
|
||||
user_id INT NOT NULL,
|
||||
PRIMARY KEY (group_id, user_id),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- column_path: dot path under the table: col or col.sub.field
|
||||
-- access_type: mask, hide, read, ...
|
||||
-- extra_filters: JSON object
|
||||
|
||||
IF OBJECT_ID(N'sec_column_rules', N'U') IS NULL
|
||||
CREATE TABLE sec_column_rules (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT,
|
||||
group_id INT,
|
||||
schema_name NVARCHAR(255) NOT NULL,
|
||||
table_name NVARCHAR(255) NOT NULL,
|
||||
column_path NVARCHAR(255) NOT NULL,
|
||||
access_type NVARCHAR(50) NOT NULL,
|
||||
mask_start INT DEFAULT 0,
|
||||
mask_end INT DEFAULT 0,
|
||||
mask_invert BIT DEFAULT 0,
|
||||
mask_char NVARCHAR(10) DEFAULT '*',
|
||||
extra_filters NVARCHAR(MAX),
|
||||
is_active BIT NOT NULL DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_column_rules_table' AND object_id = OBJECT_ID(N'sec_column_rules'))
|
||||
CREATE INDEX idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
|
||||
|
||||
|
||||
-- template: SQL fragment, e.g. user_id = {UserID}
|
||||
|
||||
IF OBJECT_ID(N'sec_row_rules', N'U') IS NULL
|
||||
CREATE TABLE sec_row_rules (
|
||||
id INT IDENTITY(1,1) PRIMARY KEY,
|
||||
user_id INT,
|
||||
group_id INT,
|
||||
schema_name NVARCHAR(255) NOT NULL,
|
||||
table_name NVARCHAR(255) NOT NULL,
|
||||
template NVARCHAR(MAX),
|
||||
has_block BIT NOT NULL DEFAULT 0,
|
||||
is_active BIT NOT NULL DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = N'idx_sec_row_rules_table' AND object_id = OBJECT_ID(N'sec_row_rules'))
|
||||
CREATE INDEX idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
-- Reference schema for the lookup direct backend: MySQL 8.0.16+ / MariaDB 10.2+. Direct backend only. Run each statement separately (see ddl.Statements) unless multiStatements=true.
|
||||
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
|
||||
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
|
||||
|
||||
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
|
||||
-- roles: comma-separated roles
|
||||
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
email VARCHAR(255) NOT NULL UNIQUE,
|
||||
password VARCHAR(255),
|
||||
user_level INT DEFAULT 0,
|
||||
roles VARCHAR(500),
|
||||
is_active TINYINT(1) DEFAULT 1,
|
||||
created_at DATETIME NULL,
|
||||
updated_at DATETIME NULL,
|
||||
last_login_at DATETIME,
|
||||
program_user_id INT DEFAULT 0,
|
||||
program_user_table VARCHAR(255) DEFAULT '',
|
||||
remote_id VARCHAR(255),
|
||||
auth_provider VARCHAR(50),
|
||||
totp_secret VARCHAR(255),
|
||||
totp_enabled TINYINT(1) DEFAULT 0,
|
||||
totp_enabled_at DATETIME
|
||||
);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||
user_id INT NOT NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
created_at DATETIME NULL,
|
||||
last_activity_at DATETIME NULL,
|
||||
ip_address VARCHAR(45),
|
||||
user_agent TEXT,
|
||||
access_token TEXT,
|
||||
refresh_token TEXT,
|
||||
token_type VARCHAR(50) DEFAULT 'Bearer',
|
||||
auth_provider VARCHAR(50),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
INDEX idx_user_sessions_user_id (user_id),
|
||||
INDEX idx_user_sessions_expires_at (expires_at)
|
||||
);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS token_blacklist (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
token VARCHAR(500) NOT NULL,
|
||||
user_id INT,
|
||||
expires_at DATETIME NOT NULL,
|
||||
created_at DATETIME NULL,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- code_hash: SHA-256 hex of the backup code
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
code_hash VARCHAR(64) NOT NULL,
|
||||
used TINYINT(1) DEFAULT 0,
|
||||
used_at DATETIME,
|
||||
created_at DATETIME NULL,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
INDEX idx_totp_user_id (user_id),
|
||||
INDEX idx_totp_code_hash (code_hash)
|
||||
);
|
||||
|
||||
|
||||
-- credential_id: base64 text
|
||||
-- public_key: base64 text
|
||||
-- aaguid: base64 text
|
||||
-- transports: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
credential_id VARCHAR(1400) CHARACTER SET ascii NOT NULL UNIQUE,
|
||||
public_key TEXT NOT NULL,
|
||||
attestation_type VARCHAR(50) DEFAULT 'none',
|
||||
aaguid TEXT,
|
||||
sign_count INT DEFAULT 0,
|
||||
clone_warning TINYINT(1) DEFAULT 0,
|
||||
transports TEXT,
|
||||
backup_eligible TINYINT(1) DEFAULT 0,
|
||||
backup_state TINYINT(1) DEFAULT 0,
|
||||
name VARCHAR(255),
|
||||
created_at DATETIME NULL,
|
||||
last_used_at DATETIME NULL,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
INDEX idx_passkey_user_id (user_id)
|
||||
);
|
||||
|
||||
|
||||
-- token_hash: SHA-256 hex of the raw token
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_password_resets (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
expires_at DATETIME NOT NULL,
|
||||
created_at DATETIME NULL,
|
||||
used TINYINT(1) DEFAULT 0,
|
||||
used_at DATETIME,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
INDEX idx_pw_reset_user_id (user_id),
|
||||
INDEX idx_pw_reset_expires_at (expires_at)
|
||||
);
|
||||
|
||||
|
||||
-- redirect_uris: JSON-encoded array
|
||||
-- grant_types: JSON-encoded array
|
||||
-- allowed_scopes: JSON-encoded array
|
||||
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
client_id VARCHAR(255) NOT NULL UNIQUE,
|
||||
redirect_uris TEXT NOT NULL,
|
||||
client_name VARCHAR(255),
|
||||
grant_types TEXT,
|
||||
allowed_scopes TEXT,
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active TINYINT(1) DEFAULT 1,
|
||||
created_at DATETIME NULL
|
||||
);
|
||||
|
||||
|
||||
-- scopes: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
code VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
redirect_uri TEXT NOT NULL,
|
||||
client_state TEXT,
|
||||
code_challenge VARCHAR(255) NOT NULL,
|
||||
code_challenge_method VARCHAR(10) DEFAULT 'S256',
|
||||
session_token TEXT NOT NULL,
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at DATETIME NOT NULL,
|
||||
created_at DATETIME NULL,
|
||||
INDEX idx_oauth_codes_expires (expires_at)
|
||||
);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_keys (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT NOT NULL,
|
||||
key_type VARCHAR(50) NOT NULL,
|
||||
key_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
scopes TEXT,
|
||||
meta TEXT,
|
||||
expires_at DATETIME,
|
||||
created_at DATETIME NULL,
|
||||
last_used_at DATETIME,
|
||||
is_active TINYINT(1) DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
INDEX idx_user_keys_user_id (user_id),
|
||||
INDEX idx_user_keys_key_type (key_type)
|
||||
);
|
||||
|
||||
|
||||
-- Optional: omit to use per-user rules only.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_group_members (
|
||||
group_id INT NOT NULL,
|
||||
user_id INT NOT NULL,
|
||||
PRIMARY KEY (group_id, user_id),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- column_path: dot path under the table: col or col.sub.field
|
||||
-- access_type: mask, hide, read, ...
|
||||
-- extra_filters: JSON object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_column_rules (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT,
|
||||
group_id INT,
|
||||
schema_name VARCHAR(255) NOT NULL,
|
||||
table_name VARCHAR(255) NOT NULL,
|
||||
column_path VARCHAR(255) NOT NULL,
|
||||
access_type VARCHAR(50) NOT NULL,
|
||||
mask_start INT DEFAULT 0,
|
||||
mask_end INT DEFAULT 0,
|
||||
mask_invert TINYINT(1) DEFAULT 0,
|
||||
mask_char VARCHAR(10) DEFAULT '*',
|
||||
extra_filters TEXT,
|
||||
is_active TINYINT(1) NOT NULL DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)),
|
||||
INDEX idx_sec_column_rules_table (schema_name, table_name)
|
||||
);
|
||||
|
||||
|
||||
-- template: SQL fragment, e.g. user_id = {UserID}
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_row_rules (
|
||||
id INT AUTO_INCREMENT PRIMARY KEY,
|
||||
user_id INT,
|
||||
group_id INT,
|
||||
schema_name VARCHAR(255) NOT NULL,
|
||||
table_name VARCHAR(255) NOT NULL,
|
||||
template TEXT,
|
||||
has_block TINYINT(1) NOT NULL DEFAULT 0,
|
||||
is_active TINYINT(1) NOT NULL DEFAULT 1,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL)),
|
||||
INDEX idx_sec_row_rules_table (schema_name, table_name)
|
||||
);
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
-- Reference schema for the lookup direct backend: PostgreSQL (tables only, no stored procedures). Use with lookup.Config{Mode: lookup.ModeDirect}.
|
||||
-- Do not combine with the procedure schema (database_schema.sql): that schema stores credential ids as bytea and
|
||||
-- list columns as text[], this one stores base64 / JSON text which is what the direct backend reads and writes.
|
||||
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
|
||||
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
|
||||
|
||||
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
|
||||
-- roles: comma-separated roles
|
||||
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id SERIAL PRIMARY KEY,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
email VARCHAR(255) NOT NULL UNIQUE,
|
||||
password VARCHAR(255),
|
||||
user_level INTEGER DEFAULT 0,
|
||||
roles VARCHAR(500),
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
last_login_at TIMESTAMP,
|
||||
program_user_id INTEGER DEFAULT 0,
|
||||
program_user_table VARCHAR(255) DEFAULT '',
|
||||
remote_id VARCHAR(255),
|
||||
auth_provider VARCHAR(50),
|
||||
totp_secret VARCHAR(255),
|
||||
totp_enabled BOOLEAN DEFAULT false,
|
||||
totp_enabled_at TIMESTAMP
|
||||
);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||
id SERIAL PRIMARY KEY,
|
||||
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||
user_id INTEGER NOT NULL,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
last_activity_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
ip_address VARCHAR(45),
|
||||
user_agent TEXT,
|
||||
access_token TEXT,
|
||||
refresh_token TEXT,
|
||||
token_type VARCHAR(50) DEFAULT 'Bearer',
|
||||
auth_provider VARCHAR(50),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_refresh_token ON user_sessions(refresh_token);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS token_blacklist (
|
||||
id SERIAL PRIMARY KEY,
|
||||
token VARCHAR(500) NOT NULL,
|
||||
user_id INTEGER,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- code_hash: SHA-256 hex of the backup code
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
code_hash VARCHAR(64) NOT NULL,
|
||||
used BOOLEAN DEFAULT false,
|
||||
used_at TIMESTAMP,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_totp_user_id ON user_totp_backup_codes(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_totp_code_hash ON user_totp_backup_codes(code_hash);
|
||||
|
||||
|
||||
-- credential_id: base64 text
|
||||
-- public_key: base64 text
|
||||
-- aaguid: base64 text
|
||||
-- transports: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
credential_id TEXT NOT NULL UNIQUE,
|
||||
public_key TEXT NOT NULL,
|
||||
attestation_type VARCHAR(50) DEFAULT 'none',
|
||||
aaguid TEXT,
|
||||
sign_count INTEGER DEFAULT 0,
|
||||
clone_warning BOOLEAN DEFAULT false,
|
||||
transports TEXT,
|
||||
backup_eligible BOOLEAN DEFAULT false,
|
||||
backup_state BOOLEAN DEFAULT false,
|
||||
name VARCHAR(255),
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
last_used_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id);
|
||||
|
||||
|
||||
-- token_hash: SHA-256 hex of the raw token
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_password_resets (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
used BOOLEAN DEFAULT false,
|
||||
used_at TIMESTAMP,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at);
|
||||
|
||||
|
||||
-- redirect_uris: JSON-encoded array
|
||||
-- grant_types: JSON-encoded array
|
||||
-- allowed_scopes: JSON-encoded array
|
||||
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
id SERIAL PRIMARY KEY,
|
||||
client_id VARCHAR(255) NOT NULL UNIQUE,
|
||||
redirect_uris TEXT NOT NULL,
|
||||
client_name VARCHAR(255),
|
||||
grant_types TEXT,
|
||||
allowed_scopes TEXT,
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
|
||||
-- scopes: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
id SERIAL PRIMARY KEY,
|
||||
code VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
redirect_uri TEXT NOT NULL,
|
||||
client_state TEXT,
|
||||
code_challenge VARCHAR(255) NOT NULL,
|
||||
code_challenge_method VARCHAR(10) DEFAULT 'S256',
|
||||
session_token TEXT NOT NULL,
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_keys (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL,
|
||||
key_type VARCHAR(50) NOT NULL,
|
||||
key_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
scopes TEXT,
|
||||
meta TEXT,
|
||||
expires_at TIMESTAMP,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
last_used_at TIMESTAMP,
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
|
||||
|
||||
|
||||
-- Optional: omit to use per-user rules only.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_group_members (
|
||||
group_id INTEGER NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
PRIMARY KEY (group_id, user_id),
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
|
||||
-- column_path: dot path under the table: col or col.sub.field
|
||||
-- access_type: mask, hide, read, ...
|
||||
-- extra_filters: JSON object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_column_rules (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER,
|
||||
group_id INTEGER,
|
||||
schema_name TEXT NOT NULL,
|
||||
table_name TEXT NOT NULL,
|
||||
column_path TEXT NOT NULL,
|
||||
access_type VARCHAR(50) NOT NULL,
|
||||
mask_start INTEGER DEFAULT 0,
|
||||
mask_end INTEGER DEFAULT 0,
|
||||
mask_invert BOOLEAN DEFAULT false,
|
||||
mask_char VARCHAR(10) DEFAULT '*',
|
||||
extra_filters TEXT,
|
||||
is_active BOOLEAN NOT NULL DEFAULT true,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
|
||||
|
||||
|
||||
-- template: SQL fragment, e.g. user_id = {UserID}
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_row_rules (
|
||||
id SERIAL PRIMARY KEY,
|
||||
user_id INTEGER,
|
||||
group_id INTEGER,
|
||||
schema_name TEXT NOT NULL,
|
||||
table_name TEXT NOT NULL,
|
||||
template TEXT,
|
||||
has_block BOOLEAN NOT NULL DEFAULT false,
|
||||
is_active BOOLEAN NOT NULL DEFAULT true,
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
-- Reference schema for the lookup direct backend: SQLite. Direct backend only.
|
||||
-- Table and column names are the lookup.DefaultSchema defaults, override them with lookup.Config.Schema.
|
||||
-- Generated by the project, edit freely for your deployment (types, collations, extra columns).
|
||||
|
||||
-- password: bcrypt hash (nullable for OAuth2 users), legacy cleartext is accepted at login
|
||||
-- roles: comma-separated roles
|
||||
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
email VARCHAR(255) NOT NULL UNIQUE,
|
||||
password VARCHAR(255),
|
||||
user_level INTEGER DEFAULT 0,
|
||||
roles VARCHAR(500),
|
||||
is_active BOOLEAN DEFAULT 1,
|
||||
created_at TIMESTAMP,
|
||||
updated_at TIMESTAMP,
|
||||
last_login_at TIMESTAMP,
|
||||
program_user_id INTEGER DEFAULT 0,
|
||||
program_user_table VARCHAR(255) DEFAULT '',
|
||||
remote_id VARCHAR(255),
|
||||
auth_provider VARCHAR(50),
|
||||
totp_secret VARCHAR(255),
|
||||
totp_enabled BOOLEAN DEFAULT 0,
|
||||
totp_enabled_at TIMESTAMP
|
||||
);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_sessions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_token VARCHAR(500) NOT NULL UNIQUE,
|
||||
user_id INTEGER NOT NULL,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP,
|
||||
last_activity_at TIMESTAMP,
|
||||
ip_address VARCHAR(45),
|
||||
user_agent TEXT,
|
||||
access_token TEXT,
|
||||
refresh_token TEXT,
|
||||
token_type VARCHAR(50) DEFAULT 'Bearer',
|
||||
auth_provider VARCHAR(50)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_user_id ON user_sessions(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_expires_at ON user_sessions(expires_at);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_sessions_refresh_token ON user_sessions(refresh_token);
|
||||
|
||||
|
||||
CREATE TABLE IF NOT EXISTS token_blacklist (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
token VARCHAR(500) NOT NULL,
|
||||
user_id INTEGER,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
|
||||
-- code_hash: SHA-256 hex of the backup code
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_totp_backup_codes (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
code_hash VARCHAR(64) NOT NULL,
|
||||
used BOOLEAN DEFAULT 0,
|
||||
used_at TIMESTAMP,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_totp_user_id ON user_totp_backup_codes(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_totp_code_hash ON user_totp_backup_codes(code_hash);
|
||||
|
||||
|
||||
-- credential_id: base64 text
|
||||
-- public_key: base64 text
|
||||
-- aaguid: base64 text
|
||||
-- transports: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_passkey_credentials (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
credential_id TEXT NOT NULL UNIQUE,
|
||||
public_key TEXT NOT NULL,
|
||||
attestation_type VARCHAR(50) DEFAULT 'none',
|
||||
aaguid TEXT,
|
||||
sign_count INTEGER DEFAULT 0,
|
||||
clone_warning BOOLEAN DEFAULT 0,
|
||||
transports TEXT,
|
||||
backup_eligible BOOLEAN DEFAULT 0,
|
||||
backup_state BOOLEAN DEFAULT 0,
|
||||
name VARCHAR(255),
|
||||
created_at TIMESTAMP,
|
||||
last_used_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_passkey_user_id ON user_passkey_credentials(user_id);
|
||||
|
||||
|
||||
-- token_hash: SHA-256 hex of the raw token
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_password_resets (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
token_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP,
|
||||
used BOOLEAN DEFAULT 0,
|
||||
used_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_pw_reset_user_id ON user_password_resets(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_pw_reset_expires_at ON user_password_resets(expires_at);
|
||||
|
||||
|
||||
-- redirect_uris: JSON-encoded array
|
||||
-- grant_types: JSON-encoded array
|
||||
-- allowed_scopes: JSON-encoded array
|
||||
-- client_secret_hash: SHA-256 hex of the confidential-client secret, NULL for public clients
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
client_id VARCHAR(255) NOT NULL UNIQUE,
|
||||
redirect_uris TEXT NOT NULL,
|
||||
client_name VARCHAR(255),
|
||||
grant_types TEXT,
|
||||
allowed_scopes TEXT,
|
||||
client_secret_hash TEXT,
|
||||
token_endpoint_auth_method VARCHAR(30) DEFAULT 'none',
|
||||
is_active BOOLEAN DEFAULT 1,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
|
||||
-- scopes: JSON-encoded array
|
||||
|
||||
CREATE TABLE IF NOT EXISTS oauth_codes (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
code VARCHAR(255) NOT NULL UNIQUE,
|
||||
client_id VARCHAR(255) NOT NULL,
|
||||
redirect_uri TEXT NOT NULL,
|
||||
client_state TEXT,
|
||||
code_challenge VARCHAR(255) NOT NULL,
|
||||
code_challenge_method VARCHAR(10) DEFAULT 'S256',
|
||||
session_token TEXT NOT NULL,
|
||||
refresh_token TEXT,
|
||||
scopes TEXT,
|
||||
expires_at TIMESTAMP NOT NULL,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_oauth_codes_expires ON oauth_codes(expires_at);
|
||||
|
||||
|
||||
-- key_hash: SHA-256 hex
|
||||
-- scopes: JSON-encoded array
|
||||
-- meta: JSON-encoded object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_keys (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
key_type VARCHAR(50) NOT NULL,
|
||||
key_hash VARCHAR(64) NOT NULL UNIQUE,
|
||||
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
scopes TEXT,
|
||||
meta TEXT,
|
||||
expires_at TIMESTAMP,
|
||||
created_at TIMESTAMP,
|
||||
last_used_at TIMESTAMP,
|
||||
is_active BOOLEAN DEFAULT 1
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
|
||||
|
||||
|
||||
-- Optional: omit to use per-user rules only.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_group_members (
|
||||
group_id INTEGER NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
PRIMARY KEY (group_id, user_id)
|
||||
);
|
||||
|
||||
|
||||
-- column_path: dot path under the table: col or col.sub.field
|
||||
-- access_type: mask, hide, read, ...
|
||||
-- extra_filters: JSON object
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_column_rules (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER,
|
||||
group_id INTEGER,
|
||||
schema_name TEXT NOT NULL,
|
||||
table_name TEXT NOT NULL,
|
||||
column_path TEXT NOT NULL,
|
||||
access_type VARCHAR(50) NOT NULL,
|
||||
mask_start INTEGER DEFAULT 0,
|
||||
mask_end INTEGER DEFAULT 0,
|
||||
mask_invert BOOLEAN DEFAULT 0,
|
||||
mask_char VARCHAR(10) DEFAULT '*',
|
||||
extra_filters TEXT,
|
||||
is_active BOOLEAN NOT NULL DEFAULT 1,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_sec_column_rules_table ON sec_column_rules(schema_name, table_name);
|
||||
|
||||
|
||||
-- template: SQL fragment, e.g. user_id = {UserID}
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sec_row_rules (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER,
|
||||
group_id INTEGER,
|
||||
schema_name TEXT NOT NULL,
|
||||
table_name TEXT NOT NULL,
|
||||
template TEXT,
|
||||
has_block BOOLEAN NOT NULL DEFAULT 0,
|
||||
is_active BOOLEAN NOT NULL DEFAULT 1,
|
||||
CHECK ((user_id IS NULL AND group_id IS NOT NULL) OR (user_id IS NOT NULL AND group_id IS NULL))
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_sec_row_rules_table ON sec_row_rules(schema_name, table_name);
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
package lookup_test
|
||||
|
||||
import (
|
||||
"os/exec"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// lookup and its backends sit above sectypes and below security: they must not import the
|
||||
// core package (security imports lookup) or any sibling sub package.
|
||||
func TestNoUpwardImports(t *testing.T) {
|
||||
for _, pkg := range []string{".", "./procedure", "./direct", "./conformance"} {
|
||||
checkNoUpwardImports(t, pkg)
|
||||
}
|
||||
}
|
||||
|
||||
func checkNoUpwardImports(t *testing.T, pkg string) {
|
||||
t.Helper()
|
||||
out, err := exec.Command("go", "list", "-deps", "-f", "{{.ImportPath}}", pkg).Output()
|
||||
if err != nil {
|
||||
t.Skipf("go list unavailable: %v", err)
|
||||
}
|
||||
for _, p := range strings.Fields(string(out)) {
|
||||
if strings.Contains(p, "uptrace/bun") || strings.Contains(p, "gorm.io") {
|
||||
t.Errorf("%s must not depend on an ORM, found %s", pkg, p)
|
||||
}
|
||||
if !strings.Contains(p, "/pkg/security") || strings.HasSuffix(p, "/lookup") ||
|
||||
strings.HasSuffix(p, "/sectypes") || strings.HasSuffix(p, "/lookup/dialect") ||
|
||||
strings.HasSuffix(p, "/lookup/procedure") || strings.HasSuffix(p, "/lookup/direct") || strings.HasSuffix(p, "/lookup/conformance") {
|
||||
continue
|
||||
}
|
||||
t.Errorf("%s must not import %s", pkg, p)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -0,0 +1,608 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
const sessionLifetime = 24 * time.Hour
|
||||
|
||||
// AuthOptions tunes Auth.
|
||||
type AuthOptions struct {
|
||||
// UpgradePasswordHash rewrites legacy cleartext passwords as bcrypt on a successful login.
|
||||
UpgradePasswordHash bool
|
||||
}
|
||||
|
||||
// Auth implements lookup.AuthStore on the tables. Passwords are verified with bcrypt;
|
||||
// legacy cleartext rows are accepted at login and only rewritten when UpgradePasswordHash
|
||||
// is set. Registration never honours client-supplied user_level or roles. Multi-step writes
|
||||
// (login, register, refresh, reset) run in one transaction.
|
||||
type Auth struct {
|
||||
*Base
|
||||
opts AuthOptions
|
||||
}
|
||||
|
||||
var _ lookup.AuthStore = (*Auth)(nil)
|
||||
|
||||
// NewAuth creates the direct AuthStore.
|
||||
func NewAuth(b *Base, opts AuthOptions) *Auth { return &Auth{Base: b, opts: opts} }
|
||||
|
||||
// GenerateSessionToken returns "sess_<64 hex>_<unix>".
|
||||
func GenerateSessionToken() (string, error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("sess_%s_%d", hex.EncodeToString(buf), time.Now().Unix()), nil
|
||||
}
|
||||
|
||||
// ParseRoles splits the comma-separated roles column.
|
||||
func ParseRoles(s string) []string {
|
||||
if s == "" {
|
||||
return []string{}
|
||||
}
|
||||
return strings.Split(s, ",")
|
||||
}
|
||||
|
||||
func claimStrings(claims map[string]any) (ip, ua string) {
|
||||
if claims == nil {
|
||||
return "", ""
|
||||
}
|
||||
if v, ok := claims["ip_address"].(string); ok {
|
||||
ip = v
|
||||
}
|
||||
if v, ok := claims["user_agent"].(string); ok {
|
||||
ua = v
|
||||
}
|
||||
return ip, ua
|
||||
}
|
||||
|
||||
func sha256Hex(s string) string {
|
||||
h := sha256.Sum256([]byte(s))
|
||||
return hex.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
// userRow is the users columns every session-bearing response needs.
|
||||
type userRow struct {
|
||||
id int
|
||||
username sql.NullString
|
||||
email sql.NullString
|
||||
roles sql.NullString
|
||||
programUserTable sql.NullString
|
||||
userLevel sql.NullInt64
|
||||
programUserID sql.NullInt64
|
||||
}
|
||||
|
||||
func (u *userRow) context(sessionID string) *sectypes.UserContext {
|
||||
return §ypes.UserContext{
|
||||
UserID: u.id,
|
||||
UserName: u.username.String,
|
||||
Email: u.email.String,
|
||||
UserLevel: int(u.userLevel.Int64),
|
||||
SessionID: sessionID,
|
||||
Roles: ParseRoles(u.roles.String),
|
||||
ProgramUserID: int(u.programUserID.Int64),
|
||||
ProgramUserTable: u.programUserTable.String,
|
||||
}
|
||||
}
|
||||
|
||||
// insertSession writes a session row and stamps the user's last login.
|
||||
func (a *Auth) insertSession(ctx context.Context, q Querier, token string, userID int64, expiresAt time.Time, ip, ua string, now time.Time) error {
|
||||
err := a.Insert(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsToken, token),
|
||||
Set(lookup.SessionsUserID, userID),
|
||||
Set(lookup.SessionsExpiresAt, expiresAt),
|
||||
Set(lookup.SessionsIPAddress, ip),
|
||||
Set(lookup.SessionsUserAgent, ua),
|
||||
Set(lookup.SessionsLastActivityAt, now),
|
||||
Set(lookup.SessionsCreatedAt, now),
|
||||
).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return a.touchLastLogin(ctx, q, userID, now)
|
||||
}
|
||||
|
||||
func (a *Auth) touchLastLogin(ctx context.Context, q Querier, userID int64, now time.Time) error {
|
||||
_, err := a.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
return err
|
||||
}
|
||||
|
||||
// Login implements lookup.AuthStore.
|
||||
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
var userID int
|
||||
var email, roles, programUserTable, storedPassword sql.NullString
|
||||
var userLevel, programUserID sql.NullInt64
|
||||
|
||||
err := a.do(func(q Querier) error {
|
||||
return a.From(lookup.EntityUsers).
|
||||
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
|
||||
lookup.UsersProgramUserID, lookup.UsersProgramUserTable, lookup.UsersPassword).
|
||||
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
BurnPasswordCheck(req.Password)
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
|
||||
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
|
||||
if !ok {
|
||||
if storedPassword.String == "" {
|
||||
BurnPasswordCheck(req.Password)
|
||||
}
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
if needsRehash && a.opts.UpgradePasswordHash {
|
||||
a.upgradePasswordHash(ctx, userID, req.Password)
|
||||
}
|
||||
|
||||
token, err := GenerateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
}
|
||||
now := a.Now()
|
||||
ip, ua := claimStrings(req.Claims)
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
return a.insertSession(ctx, q, token, int64(userID), now.Add(sessionLifetime), ip, ua, now)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
|
||||
return §ypes.LoginResponse{
|
||||
Token: token,
|
||||
User: §ypes.UserContext{
|
||||
UserID: userID,
|
||||
UserName: req.Username,
|
||||
Email: email.String,
|
||||
UserLevel: int(userLevel.Int64),
|
||||
Roles: ParseRoles(roles.String),
|
||||
SessionID: token,
|
||||
ProgramUserID: int(programUserID.Int64),
|
||||
ProgramUserTable: programUserTable.String,
|
||||
},
|
||||
ExpiresIn: int64(sessionLifetime.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// upgradePasswordHash replaces a legacy cleartext password with a bcrypt hash. Failure is
|
||||
// logged and ignored: the login itself already succeeded.
|
||||
func (a *Auth) upgradePasswordHash(ctx context.Context, userID int, password string) {
|
||||
h, err := HashPassword(password)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
err = a.do(func(q Querier) error {
|
||||
_, err := a.Update(lookup.EntityUsers).
|
||||
Set(Set(lookup.UsersPassword, h), Set(lookup.UsersUpdatedAt, a.Now())).
|
||||
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Register implements lookup.AuthStore.
|
||||
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
|
||||
if req.Username == "" {
|
||||
return nil, fmt.Errorf("username is required")
|
||||
}
|
||||
if req.Email == "" {
|
||||
return nil, fmt.Errorf("email is required")
|
||||
}
|
||||
if req.Password == "" {
|
||||
return nil, fmt.Errorf("password is required")
|
||||
}
|
||||
hash, err := HashPassword(req.Password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
token, err := GenerateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
}
|
||||
|
||||
// Privileges are never taken from the request: self-registration always creates an
|
||||
// unprivileged user.
|
||||
const userLevel = 0
|
||||
now := a.Now()
|
||||
ip, ua := claimStrings(req.Claims)
|
||||
|
||||
var userID int64
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
exists, err := a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersUsername, req.Username)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return lookup.ErrUsernameExists
|
||||
}
|
||||
exists, err = a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersEmail, req.Email)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return lookup.ErrEmailExists
|
||||
}
|
||||
userID, err = a.Insert(lookup.EntityUsers).Set(
|
||||
Set(lookup.UsersUsername, req.Username),
|
||||
Set(lookup.UsersEmail, req.Email),
|
||||
Set(lookup.UsersPassword, hash),
|
||||
Set(lookup.UsersUserLevel, userLevel),
|
||||
Set(lookup.UsersRoles, ""),
|
||||
Set(lookup.UsersIsActive, true),
|
||||
Set(lookup.UsersCreatedAt, now),
|
||||
Set(lookup.UsersUpdatedAt, now),
|
||||
Set(lookup.UsersProgramUserID, 0),
|
||||
Set(lookup.UsersProgramUserTable, ""),
|
||||
).ExecID(ctx, q, lookup.UsersID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return a.insertSession(ctx, q, token, userID, now.Add(sessionLifetime), ip, ua, now)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, lookup.ErrUsernameExists) || errors.Is(err, lookup.ErrEmailExists) {
|
||||
return nil, err
|
||||
}
|
||||
return nil, fmt.Errorf("register query failed: %w", err)
|
||||
}
|
||||
|
||||
return §ypes.LoginResponse{
|
||||
Token: token,
|
||||
User: §ypes.UserContext{
|
||||
UserID: int(userID),
|
||||
UserName: req.Username,
|
||||
Email: req.Email,
|
||||
UserLevel: userLevel,
|
||||
Roles: ParseRoles(""),
|
||||
SessionID: token,
|
||||
},
|
||||
ExpiresIn: int64(sessionLifetime.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Logout implements lookup.AuthStore.
|
||||
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
token := strings.TrimPrefix(strings.TrimPrefix(req.Token, "Bearer "), "bearer ")
|
||||
var rows int64
|
||||
err := a.do(func(q Querier) error {
|
||||
var err error
|
||||
rows, err = a.Delete(lookup.EntityUserSessions).
|
||||
Where(Eq(lookup.SessionsToken, token), Eq(lookup.SessionsUserID, req.UserID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("logout query failed: %w", err)
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("session not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// sessionUser selects the user behind a live session token.
|
||||
func (a *Auth) sessionUser(ctx context.Context, q Querier, token string, extra ...lookup.Column) (*userRow, []any, error) {
|
||||
var u userRow
|
||||
dest := []any{&u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable}
|
||||
cols := []lookup.Column{lookup.SessionsUserID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
|
||||
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable}
|
||||
extras := make([]any, len(extra))
|
||||
for i, c := range extra {
|
||||
cols = append(cols, c)
|
||||
extras[i] = new(sql.NullString)
|
||||
dest = append(dest, extras[i])
|
||||
}
|
||||
err := a.From(lookup.EntityUserSessions).Cols(cols...).
|
||||
Join(lookup.EntityUsers, EqCol(lookup.SessionsUserID, lookup.UsersID)).
|
||||
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, a.Now()), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, dest...)
|
||||
return &u, extras, err
|
||||
}
|
||||
|
||||
// Session implements lookup.AuthStore. reference is only meaningful to the procedure backend.
|
||||
func (a *Auth) Session(ctx context.Context, token, _ string) (*sectypes.UserContext, error) {
|
||||
var u *userRow
|
||||
err := a.do(func(q Querier) error {
|
||||
var err error
|
||||
u, _, err = a.sessionUser(ctx, q, token)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("invalid or expired session")
|
||||
}
|
||||
return nil, fmt.Errorf("session query failed: %w", err)
|
||||
}
|
||||
return u.context(token), nil
|
||||
}
|
||||
|
||||
// TouchSession implements lookup.AuthStore.
|
||||
func (a *Auth) TouchSession(ctx context.Context, token string, _ *sectypes.UserContext) error {
|
||||
return a.do(func(q Querier) error {
|
||||
now := a.Now()
|
||||
_, err := a.Update(lookup.EntityUserSessions).Set(Set(lookup.SessionsLastActivityAt, now)).
|
||||
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, now)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// Refresh implements lookup.AuthStore: the old session is replaced by a new one.
|
||||
func (a *Auth) Refresh(ctx context.Context, oldToken string) (*sectypes.LoginResponse, error) {
|
||||
var u *userRow
|
||||
var extras []any
|
||||
err := a.do(func(q Querier) error {
|
||||
var err error
|
||||
u, extras, err = a.sessionUser(ctx, q, oldToken, lookup.SessionsIPAddress, lookup.SessionsUserAgent)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("invalid or expired refresh token")
|
||||
}
|
||||
return nil, fmt.Errorf("refresh token query failed: %w", err)
|
||||
}
|
||||
ip := extras[0].(*sql.NullString).String
|
||||
ua := extras[1].(*sql.NullString).String
|
||||
|
||||
newToken, err := GenerateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
}
|
||||
now := a.Now()
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
err := a.Insert(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsToken, newToken),
|
||||
Set(lookup.SessionsUserID, u.id),
|
||||
Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)),
|
||||
Set(lookup.SessionsIPAddress, ip),
|
||||
Set(lookup.SessionsUserAgent, ua),
|
||||
Set(lookup.SessionsLastActivityAt, now),
|
||||
Set(lookup.SessionsCreatedAt, now),
|
||||
).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, oldToken)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("refresh token generation failed: %w", err)
|
||||
}
|
||||
return §ypes.LoginResponse{
|
||||
Token: newToken,
|
||||
User: u.context(newToken),
|
||||
ExpiresIn: int64(sessionLifetime.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// apiKeyTypes are the key types accepted by LoginAPIKey.
|
||||
var apiKeyTypes = []any{string(sectypes.KeyTypeHeaderAPI), string(sectypes.KeyTypeGenericAPI)}
|
||||
|
||||
// LoginAPIKey implements lookup.AuthStore. Unknown, expired, inactive and wrong-type keys
|
||||
// (and inactive users) all return lookup.ErrInvalidAPIKey; the raw key is never logged.
|
||||
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
if rawKey == "" {
|
||||
return nil, lookup.ErrInvalidAPIKey
|
||||
}
|
||||
now := a.Now()
|
||||
var keyID int64
|
||||
var u userRow
|
||||
err := a.do(func(q Querier) error {
|
||||
return a.From(lookup.EntityUserKeys).
|
||||
Cols(lookup.KeysID, lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
|
||||
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
|
||||
Join(lookup.EntityUsers, EqCol(lookup.KeysUserID, lookup.UsersID)).
|
||||
Where(
|
||||
Eq(lookup.KeysKeyHash, sectypes.HashKey(rawKey)),
|
||||
In(lookup.KeysKeyType, apiKeyTypes...),
|
||||
Eq(lookup.KeysIsActive, true),
|
||||
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, now)),
|
||||
Eq(lookup.UsersIsActive, true),
|
||||
).QueryRow(ctx, q, &keyID, &u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, lookup.ErrInvalidAPIKey
|
||||
}
|
||||
return nil, fmt.Errorf("api key login query failed: %w", err)
|
||||
}
|
||||
|
||||
token, err := GenerateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
}
|
||||
ip, ua := claimStrings(claims)
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
if err := a.insertSession(ctx, q, token, int64(u.id), now.Add(sessionLifetime), ip, ua, now); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := a.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, keyID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("api key login query failed: %w", err)
|
||||
}
|
||||
return §ypes.LoginResponse{
|
||||
Token: token,
|
||||
User: u.context(token),
|
||||
ExpiresIn: int64(sessionLifetime.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// JWTLogin implements lookup.AuthStore (mirrors resolvespec_jwt_login). The token is a
|
||||
// placeholder until JWT signing is wired in.
|
||||
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
var userID int
|
||||
var email, roles, storedPassword sql.NullString
|
||||
var userLevel sql.NullInt64
|
||||
err := a.do(func(q Querier) error {
|
||||
return a.From(lookup.EntityUsers).
|
||||
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, lookup.UsersPassword).
|
||||
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &storedPassword)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
BurnPasswordCheck(req.Password)
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
|
||||
if !ok {
|
||||
if storedPassword.String == "" {
|
||||
BurnPasswordCheck(req.Password)
|
||||
}
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
if needsRehash && a.opts.UpgradePasswordHash {
|
||||
a.upgradePasswordHash(ctx, userID, req.Password)
|
||||
}
|
||||
expiresAt := a.Now().Add(sessionLifetime)
|
||||
return §ypes.LoginResponse{
|
||||
Token: fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()),
|
||||
User: §ypes.UserContext{
|
||||
UserID: userID,
|
||||
UserName: req.Username,
|
||||
Email: email.String,
|
||||
UserLevel: int(userLevel.Int64),
|
||||
Roles: ParseRoles(roles.String),
|
||||
},
|
||||
ExpiresIn: int64(sessionLifetime.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// JWTLogout implements lookup.AuthStore: the token goes on the blacklist.
|
||||
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
now := a.Now()
|
||||
err := a.do(func(q Querier) error {
|
||||
return a.Insert(lookup.EntityTokenBlacklist).Set(
|
||||
Set(lookup.BlacklistToken, req.Token),
|
||||
Set(lookup.BlacklistUserID, req.UserID),
|
||||
Set(lookup.BlacklistExpiresAt, now.Add(sessionLifetime)),
|
||||
Set(lookup.BlacklistCreatedAt, now),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("logout query failed: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResetRequest implements lookup.AuthStore. An unknown user yields a generic empty success
|
||||
// so accounts cannot be enumerated.
|
||||
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
|
||||
if req.Email == "" && req.Username == "" {
|
||||
return nil, fmt.Errorf("email or username is required")
|
||||
}
|
||||
var userID int
|
||||
err := a.do(func(q Querier) error {
|
||||
lookupCol, val := lookup.UsersUsername, req.Username
|
||||
if req.Email != "" {
|
||||
lookupCol, val = lookup.UsersEmail, req.Email
|
||||
}
|
||||
return a.From(lookup.EntityUsers).Cols(lookup.UsersID).
|
||||
Where(Eq(lookupCol, val), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return §ypes.PasswordResetResponse{Token: "", ExpiresIn: 0}, nil
|
||||
}
|
||||
return nil, fmt.Errorf("password reset request query failed: %w", err)
|
||||
}
|
||||
|
||||
raw := make([]byte, 32)
|
||||
if _, err := rand.Read(raw); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate reset token: %w", err)
|
||||
}
|
||||
rawToken := hex.EncodeToString(raw)
|
||||
now := a.Now()
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
if _, err := a.Delete(lookup.EntityUserPasswordResets).
|
||||
Where(Eq(lookup.ResetsUserID, userID), Eq(lookup.ResetsUsed, false)).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
return a.Insert(lookup.EntityUserPasswordResets).Set(
|
||||
Set(lookup.ResetsUserID, userID),
|
||||
Set(lookup.ResetsTokenHash, sha256Hex(rawToken)),
|
||||
Set(lookup.ResetsExpiresAt, now.Add(time.Hour)),
|
||||
Set(lookup.ResetsCreatedAt, now),
|
||||
Set(lookup.ResetsUsed, false),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("password reset request query failed: %w", err)
|
||||
}
|
||||
return §ypes.PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil
|
||||
}
|
||||
|
||||
// ResetComplete implements lookup.AuthStore: sets the new password, ends every session of
|
||||
// the user and consumes the reset token, atomically.
|
||||
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
|
||||
if req.Token == "" {
|
||||
return fmt.Errorf("token is required")
|
||||
}
|
||||
if req.NewPassword == "" {
|
||||
return fmt.Errorf("new_password is required")
|
||||
}
|
||||
newHash, err := HashPassword(req.NewPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tokenHash := sha256Hex(req.Token)
|
||||
|
||||
now := a.Now()
|
||||
var resetID, userID int
|
||||
var expiresAt time.Time
|
||||
err = a.do(func(q Querier) error {
|
||||
return a.From(lookup.EntityUserPasswordResets).
|
||||
Cols(lookup.ResetsID, lookup.ResetsUserID, lookup.ResetsExpiresAt).
|
||||
Where(Eq(lookup.ResetsTokenHash, tokenHash), Eq(lookup.ResetsUsed, false)).
|
||||
QueryRow(ctx, q, &resetID, &userID, a.timeDest(&expiresAt))
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return fmt.Errorf("invalid or expired token")
|
||||
}
|
||||
return fmt.Errorf("password reset complete query failed: %w", err)
|
||||
}
|
||||
if !expiresAt.After(now) {
|
||||
return fmt.Errorf("invalid or expired token")
|
||||
}
|
||||
|
||||
err = a.tx(ctx, func(q Querier) error {
|
||||
if _, err := a.Update(lookup.EntityUsers).
|
||||
Set(Set(lookup.UsersPassword, newHash), Set(lookup.UsersUpdatedAt, now)).
|
||||
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsUserID, userID)).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := a.Update(lookup.EntityUserPasswordResets).
|
||||
Set(Set(lookup.ResetsUsed, true), Set(lookup.ResetsUsedAt, now)).
|
||||
Where(Eq(lookup.ResetsID, resetID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("password reset complete query failed: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,268 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
func newAuth(t *testing.T, opts AuthOptions) (*Auth, *sql.DB) {
|
||||
db := newTestDB(t)
|
||||
return NewAuth(newTestBase(t, db, nil), opts), db
|
||||
}
|
||||
|
||||
func registerUser(t *testing.T, a *Auth, name string) *sectypes.LoginResponse {
|
||||
t.Helper()
|
||||
resp, err := a.Register(context.Background(), sectypes.RegisterRequest{Username: name, Email: name + "@x.io", Password: "pw-" + name})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
func TestRegisterLoginSessionFlow(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, db := newAuth(t, AuthOptions{})
|
||||
|
||||
reg, err := a.Register(ctx, sectypes.RegisterRequest{
|
||||
Username: "ann", Email: "ann@x.io", Password: "secret",
|
||||
UserLevel: 99, Roles: []string{"admin"}, // must be ignored
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if reg.User.UserLevel != 0 || len(reg.User.Roles) != 0 {
|
||||
t.Fatalf("register honoured privileges: %+v", reg.User)
|
||||
}
|
||||
var stored string
|
||||
if err := db.QueryRow(`SELECT password FROM users WHERE username='ann'`).Scan(&stored); err != nil || !strings.HasPrefix(stored, "$2") {
|
||||
t.Fatalf("password not bcrypt: %q %v", stored, err)
|
||||
}
|
||||
|
||||
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "ann", Email: "other@x.io", Password: "x"}); !errors.Is(err, lookup.ErrUsernameExists) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "bob", Email: "ann@x.io", Password: "x"}); !errors.Is(err, lookup.ErrEmailExists) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
var n int
|
||||
_ = db.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&n)
|
||||
if n != 1 {
|
||||
t.Fatalf("failed register left a row: %d", n)
|
||||
}
|
||||
|
||||
login, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "secret", Claims: map[string]any{"ip_address": "1.2.3.4", "user_agent": "ua"}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.HasPrefix(login.Token, "sess_") || login.ExpiresIn != 86400 || login.User.Email != "ann@x.io" {
|
||||
t.Fatalf("login: %+v", login)
|
||||
}
|
||||
var ip string
|
||||
_ = db.QueryRow(`SELECT ip_address FROM user_sessions WHERE session_token=?`, login.Token).Scan(&ip)
|
||||
if ip != "1.2.3.4" {
|
||||
t.Fatalf("ip %q", ip)
|
||||
}
|
||||
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "wrong"}); err == nil || err.Error() != "invalid credentials" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "nobody", Password: "x"}); err == nil || err.Error() != "invalid credentials" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
u, err := a.Session(ctx, login.Token, "authenticate")
|
||||
if err != nil || u.UserName != "ann" || u.SessionID != login.Token {
|
||||
t.Fatalf("session: %+v %v", u, err)
|
||||
}
|
||||
if err := a.TouchSession(ctx, login.Token, u); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := a.Session(ctx, "nope", "authenticate"); err == nil || err.Error() != "invalid or expired session" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
ref, err := a.Refresh(ctx, login.Token)
|
||||
if err != nil || ref.Token == login.Token {
|
||||
t.Fatalf("refresh: %+v %v", ref, err)
|
||||
}
|
||||
if _, err := a.Session(ctx, login.Token, ""); err == nil {
|
||||
t.Fatal("old session still valid after refresh")
|
||||
}
|
||||
if _, err := a.Refresh(ctx, login.Token); err == nil || err.Error() != "invalid or expired refresh token" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: "Bearer " + ref.Token, UserID: ref.User.UserID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err == nil || err.Error() != "session not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpiredSessionRejected(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, _ := newAuth(t, AuthOptions{})
|
||||
resp := registerUser(t, a, "eve")
|
||||
a.Now = func() time.Time { return time.Now().Add(48 * time.Hour) }
|
||||
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
|
||||
t.Fatal("expired session accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLegacyPasswordUpgradeIsOptIn(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
for _, upgrade := range []bool{false, true} {
|
||||
a, db := newAuth(t, AuthOptions{UpgradePasswordHash: upgrade})
|
||||
_, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active) VALUES ('old','o@x.io','clear',1,'a,b',1)`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := a.Login(ctx, sectypes.LoginRequest{Username: "old", Password: "clear"})
|
||||
if err != nil || len(resp.User.Roles) != 2 {
|
||||
t.Fatalf("login: %+v %v", resp, err)
|
||||
}
|
||||
var stored string
|
||||
_ = db.QueryRow(`SELECT password FROM users WHERE username='old'`).Scan(&stored)
|
||||
if got := strings.HasPrefix(stored, "$2"); got != upgrade {
|
||||
t.Fatalf("upgrade=%v stored=%q", upgrade, stored)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInactiveUserCannotLogin(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, db := newAuth(t, AuthOptions{})
|
||||
resp := registerUser(t, a, "ian")
|
||||
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ian", Password: "pw-ian"}); err == nil {
|
||||
t.Fatal("inactive login accepted")
|
||||
}
|
||||
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
|
||||
t.Fatal("inactive session accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginAPIKey(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, db := newAuth(t, AuthOptions{})
|
||||
reg := registerUser(t, a, "kim")
|
||||
insert := func(raw, typ string, active int, expires any) {
|
||||
t.Helper()
|
||||
_, err := db.Exec(`INSERT INTO user_keys (user_id, key_type, key_hash, name, is_active, expires_at) VALUES (?,?,?,?,?,?)`,
|
||||
reg.User.UserID, typ, sectypes.HashKey(raw), "k", active, expires)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
insert("good", "header_api", 1, nil)
|
||||
insert("generic", "api", 1, nil)
|
||||
insert("jwt", "jwt_secret", 1, nil)
|
||||
insert("off", "api", 0, nil)
|
||||
insert("old", "api", 1, time.Now().Add(-time.Hour))
|
||||
|
||||
for _, k := range []string{"good", "generic"} {
|
||||
resp, err := a.LoginAPIKey(ctx, k, map[string]any{"ip_address": "9.9.9.9"})
|
||||
if err != nil || resp.User.UserName != "kim" || !strings.HasPrefix(resp.Token, "sess_") {
|
||||
t.Fatalf("%s: %+v %v", k, resp, err)
|
||||
}
|
||||
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
var used sql.NullString
|
||||
_ = db.QueryRow(`SELECT last_used_at FROM user_keys WHERE key_hash = ?`, sectypes.HashKey("good")).Scan(&used)
|
||||
if !used.Valid {
|
||||
t.Fatal("last_used_at not stamped")
|
||||
}
|
||||
for _, k := range []string{"", "missing", "jwt", "off", "old"} {
|
||||
if _, err := a.LoginAPIKey(ctx, k, nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
|
||||
t.Fatalf("%q: got %v", k, err)
|
||||
}
|
||||
}
|
||||
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
|
||||
if _, err := a.LoginAPIKey(ctx, "good", nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
|
||||
t.Fatalf("inactive user: got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasswordReset(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, _ := newAuth(t, AuthOptions{})
|
||||
reg := registerUser(t, a, "rae")
|
||||
|
||||
empty, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "none@x.io"})
|
||||
if err != nil || empty.Token != "" {
|
||||
t.Fatalf("enumeration leak: %+v %v", empty, err)
|
||||
}
|
||||
if _, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{}); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
|
||||
r1, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "rae@x.io"})
|
||||
if err != nil || r1.Token == "" {
|
||||
t.Fatal(err)
|
||||
}
|
||||
r2, _ := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Username: "rae"})
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r1.Token, NewPassword: "n"}); err == nil {
|
||||
t.Fatal("superseded token accepted")
|
||||
}
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "newpw"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "again"}); err == nil {
|
||||
t.Fatal("token reused")
|
||||
}
|
||||
if _, err := a.Session(ctx, reg.Token, ""); err == nil {
|
||||
t.Fatal("sessions survived reset")
|
||||
}
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "rae", Password: "newpw"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTLoginLogout(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, db := newAuth(t, AuthOptions{})
|
||||
reg := registerUser(t, a, "jay")
|
||||
resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: "jay", Password: "pw-jay"})
|
||||
if err != nil || !strings.HasPrefix(resp.Token, "token_") {
|
||||
t.Fatalf("%+v %v", resp, err)
|
||||
}
|
||||
if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: "tok", UserID: reg.User.UserID}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
_ = db.QueryRow(`SELECT COUNT(*) FROM token_blacklist WHERE token='tok'`).Scan(&n)
|
||||
if n != 1 {
|
||||
t.Fatal("token not blacklisted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomSchemaNames(t *testing.T) {
|
||||
db := newTestDB(t, `
|
||||
CREATE TABLE app_users (uid INTEGER PRIMARY KEY AUTOINCREMENT, login TEXT, email TEXT, password TEXT,
|
||||
user_level INTEGER, roles TEXT, is_active INTEGER, created_at DATETIME, updated_at DATETIME,
|
||||
last_login_at DATETIME, program_user_id INTEGER, program_user_table TEXT, remote_id TEXT, auth_provider TEXT,
|
||||
totp_secret TEXT, totp_enabled INTEGER, totp_enabled_at DATETIME);`)
|
||||
schema := lookup.Schema{lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"id": "uid", "username": "login"}}}
|
||||
a := NewAuth(newTestBase(t, db, schema), AuthOptions{})
|
||||
ctx := context.Background()
|
||||
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "zed", Email: "z@x.io", Password: "p"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "zed", Password: "p"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var login string
|
||||
if err := db.QueryRow(`SELECT login FROM app_users`).Scan(&login); err != nil || login != "zed" {
|
||||
t.Fatalf("%q %v", login, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,462 @@
|
||||
// Package direct is the table-backed implementation of the lookup stores. SQL is built
|
||||
// from the configured Schema (table and column names) and Dialect (placeholders, quoting,
|
||||
// booleans, insert-returning-id); no statement is written per database and no ORM is used.
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
|
||||
// procedure.Runner (and procedure.DB) satisfy it.
|
||||
type Runner interface {
|
||||
Run(run func(*sql.DB) error) error
|
||||
}
|
||||
|
||||
// Querier is implemented by *sql.DB and *sql.Tx.
|
||||
type Querier interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
}
|
||||
|
||||
// Base is the state shared by every direct store: the runner, dialect, schema and clock.
|
||||
type Base struct {
|
||||
run Runner
|
||||
d dialect.Dialect
|
||||
schema lookup.Schema
|
||||
// Now is the clock; tests replace it.
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
// NewBase creates the shared state. The schema is merged with the defaults and validated.
|
||||
func NewBase(run Runner, d dialect.Dialect, schema lookup.Schema) (*Base, error) {
|
||||
if run == nil {
|
||||
return nil, fmt.Errorf("direct: nil runner")
|
||||
}
|
||||
if d == nil {
|
||||
return nil, fmt.Errorf("direct: nil dialect")
|
||||
}
|
||||
merged := lookup.DefaultSchema().Merge(schema)
|
||||
if err := merged.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Base{run: run, d: d, schema: merged, Now: time.Now}, nil
|
||||
}
|
||||
|
||||
// Dialect returns the dialect in use.
|
||||
func (b *Base) Dialect() dialect.Dialect { return b.d }
|
||||
|
||||
// do runs fn against the database without a transaction.
|
||||
func (b *Base) do(fn func(q Querier) error) error {
|
||||
return b.run.Run(func(db *sql.DB) error { return fn(db) })
|
||||
}
|
||||
|
||||
// tx runs fn in one transaction; an error rolls back.
|
||||
func (b *Base) tx(ctx context.Context, fn func(q Querier) error) error {
|
||||
return b.run.Run(func(db *sql.DB) error {
|
||||
tx, err := db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fn(tx); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
})
|
||||
}
|
||||
|
||||
// tableRef returns the (possibly schema-qualified) physical table name of an entity.
|
||||
func (b *Base) tableRef(e lookup.Entity) string {
|
||||
t := b.schema[e]
|
||||
name := t.Name
|
||||
if name == "" {
|
||||
name = string(e)
|
||||
}
|
||||
if t.Schema != "" {
|
||||
return t.Schema + "." + name
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// colName returns the physical column name of a logical column.
|
||||
func (b *Base) colName(c lookup.Column) string {
|
||||
if t, ok := b.schema[c.Entity]; ok {
|
||||
if n := t.Columns[c.Name]; n != "" {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return c.Name
|
||||
}
|
||||
|
||||
// arg converts a Go value to a bind argument (booleans go through the dialect).
|
||||
func (b *Base) arg(v any) any {
|
||||
if bv, ok := v.(bool); ok {
|
||||
return b.d.Bool(bv)
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
// timeDest scans a time column through the dialect, so drivers returning strings work.
|
||||
type timeDest struct {
|
||||
d dialect.Dialect
|
||||
v *time.Time
|
||||
}
|
||||
|
||||
func (t timeDest) Scan(src any) error {
|
||||
v, err := t.d.ScanTime(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*t.v = v
|
||||
return nil
|
||||
}
|
||||
|
||||
type boolDest struct {
|
||||
d dialect.Dialect
|
||||
v *bool
|
||||
}
|
||||
|
||||
func (t boolDest) Scan(src any) error {
|
||||
v, err := t.d.ScanBool(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*t.v = v
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *Base) timeDest(v *time.Time) sql.Scanner { return timeDest{d: b.d, v: v} }
|
||||
func (b *Base) boolDest(v *bool) sql.Scanner { return boolDest{d: b.d, v: v} }
|
||||
|
||||
// --- query builder --------------------------------------------------------
|
||||
|
||||
// builder accumulates bind arguments and renders column references.
|
||||
type builder struct {
|
||||
b *Base
|
||||
args []any
|
||||
aliases map[lookup.Entity]string
|
||||
nalias int
|
||||
}
|
||||
|
||||
func (bl *builder) ph(v any) string {
|
||||
bl.args = append(bl.args, bl.b.arg(v))
|
||||
return bl.b.d.Placeholder(len(bl.args))
|
||||
}
|
||||
|
||||
// col renders a column; with aliases set (select queries) it is qualified by its table alias.
|
||||
func (bl *builder) col(c lookup.Column) string {
|
||||
name := bl.b.d.Quote(bl.b.colName(c))
|
||||
if bl.aliases != nil {
|
||||
if a, ok := bl.aliases[c.Entity]; ok {
|
||||
return a + "." + name
|
||||
}
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// Cond renders one boolean condition.
|
||||
type Cond func(*builder) string
|
||||
|
||||
// Eq is `col = value`.
|
||||
func Eq(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " = " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// EqFold is a case-insensitive `LOWER(col) = value` match (the value is lowered in Go).
|
||||
func EqFold(c lookup.Column, v string) Cond {
|
||||
return func(bl *builder) string { return "LOWER(" + bl.col(c) + ") = " + bl.ph(strings.ToLower(v)) }
|
||||
}
|
||||
|
||||
// Ne is `col <> value`.
|
||||
func Ne(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " <> " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// Gt is `col > value`.
|
||||
func Gt(c lookup.Column, v any) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
|
||||
}
|
||||
|
||||
// IsNull is `col IS NULL`.
|
||||
func IsNull(c lookup.Column) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
|
||||
}
|
||||
|
||||
// EqCol is `a = b` between two columns (join conditions).
|
||||
func EqCol(a, c lookup.Column) Cond {
|
||||
return func(bl *builder) string { return bl.col(a) + " = " + bl.col(c) }
|
||||
}
|
||||
|
||||
// In is `col IN (v...)`; an empty list renders a condition that is never true.
|
||||
func In(c lookup.Column, vs ...any) Cond {
|
||||
return func(bl *builder) string {
|
||||
if len(vs) == 0 {
|
||||
return "1 = 0"
|
||||
}
|
||||
ph := make([]string, len(vs))
|
||||
for i, v := range vs {
|
||||
ph[i] = bl.ph(v)
|
||||
}
|
||||
return bl.col(c) + " IN (" + strings.Join(ph, ", ") + ")"
|
||||
}
|
||||
}
|
||||
|
||||
// InSelect is `col IN (subselect)`; the subselect's arguments share the outer numbering.
|
||||
func InSelect(c lookup.Column, sub *Select) Cond {
|
||||
return func(bl *builder) string { return bl.col(c) + " IN (" + sub.render(bl) + ")" }
|
||||
}
|
||||
|
||||
// Or joins conditions with OR inside parentheses.
|
||||
func Or(cs ...Cond) Cond { return joinConds("OR", cs) }
|
||||
|
||||
// And joins conditions with AND inside parentheses.
|
||||
func And(cs ...Cond) Cond { return joinConds("AND", cs) }
|
||||
|
||||
func joinConds(op string, cs []Cond) Cond {
|
||||
return func(bl *builder) string {
|
||||
parts := make([]string, len(cs))
|
||||
for i, c := range cs {
|
||||
parts[i] = c(bl)
|
||||
}
|
||||
return "(" + strings.Join(parts, " "+op+" ") + ")"
|
||||
}
|
||||
}
|
||||
|
||||
func (bl *builder) where(cs []Cond) string {
|
||||
if len(cs) == 0 {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, len(cs))
|
||||
for i, c := range cs {
|
||||
parts[i] = c(bl)
|
||||
}
|
||||
return " WHERE " + strings.Join(parts, " AND ")
|
||||
}
|
||||
|
||||
// Stmt is a rendered statement.
|
||||
type Stmt struct {
|
||||
SQL string
|
||||
Args []any
|
||||
}
|
||||
|
||||
// Select builds a SELECT.
|
||||
type Select struct {
|
||||
b *Base
|
||||
from lookup.Entity
|
||||
joins []join
|
||||
cols []lookup.Column
|
||||
conds []Cond
|
||||
order []lookup.Column
|
||||
}
|
||||
|
||||
type join struct {
|
||||
e lookup.Entity
|
||||
on Cond
|
||||
}
|
||||
|
||||
// From starts a SELECT on e.
|
||||
func (b *Base) From(e lookup.Entity) *Select { return &Select{b: b, from: e} }
|
||||
|
||||
// Cols sets the selected columns.
|
||||
func (s *Select) Cols(cs ...lookup.Column) *Select { s.cols = cs; return s }
|
||||
|
||||
// Join adds `JOIN e ON on`.
|
||||
func (s *Select) Join(e lookup.Entity, on Cond) *Select {
|
||||
s.joins = append(s.joins, join{e: e, on: on})
|
||||
return s
|
||||
}
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (s *Select) Where(cs ...Cond) *Select { s.conds = append(s.conds, cs...); return s }
|
||||
|
||||
// OrderBy adds ascending order columns.
|
||||
func (s *Select) OrderBy(cs ...lookup.Column) *Select { s.order = append(s.order, cs...); return s }
|
||||
|
||||
// Build renders the statement.
|
||||
func (s *Select) Build() Stmt {
|
||||
bl := &builder{b: s.b}
|
||||
sqlText := s.render(bl)
|
||||
return Stmt{SQL: sqlText, Args: bl.args}
|
||||
}
|
||||
|
||||
// render writes the select into bl, giving every table a fresh alias so a subselect cannot
|
||||
// clash with the statement around it.
|
||||
func (s *Select) render(bl *builder) string {
|
||||
saved := bl.aliases
|
||||
defer func() { bl.aliases = saved }()
|
||||
bl.aliases = map[lookup.Entity]string{}
|
||||
alias := func() string { a := fmt.Sprintf("t%d", bl.nalias); bl.nalias++; return a }
|
||||
bl.aliases[s.from] = alias()
|
||||
for _, j := range s.joins {
|
||||
bl.aliases[j.e] = alias()
|
||||
}
|
||||
sel := make([]string, len(s.cols))
|
||||
for i, c := range s.cols {
|
||||
sel[i] = bl.col(c)
|
||||
}
|
||||
var sb strings.Builder
|
||||
sb.WriteString("SELECT " + strings.Join(sel, ", "))
|
||||
sb.WriteString(" FROM " + s.b.d.Quote(s.b.tableRef(s.from)) + " " + bl.aliases[s.from])
|
||||
for _, j := range s.joins {
|
||||
sb.WriteString(" JOIN " + s.b.d.Quote(s.b.tableRef(j.e)) + " " + bl.aliases[j.e] + " ON " + j.on(bl))
|
||||
}
|
||||
sb.WriteString(bl.where(s.conds))
|
||||
if len(s.order) > 0 {
|
||||
o := make([]string, len(s.order))
|
||||
for i, c := range s.order {
|
||||
o[i] = bl.col(c)
|
||||
}
|
||||
sb.WriteString(" ORDER BY " + strings.Join(o, ", "))
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// QueryRow runs the select and scans the first row into dest.
|
||||
func (s *Select) QueryRow(ctx context.Context, q Querier, dest ...any) error {
|
||||
st := s.Build()
|
||||
return q.QueryRowContext(ctx, st.SQL, st.Args...).Scan(dest...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
}
|
||||
|
||||
// Query runs the select.
|
||||
func (s *Select) Query(ctx context.Context, q Querier) (*sql.Rows, error) {
|
||||
st := s.Build()
|
||||
return q.QueryContext(ctx, st.SQL, st.Args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
}
|
||||
|
||||
// Exists reports whether the select returns at least one row.
|
||||
func (s *Select) Exists(ctx context.Context, q Querier) (bool, error) {
|
||||
s.cols = []lookup.Column{s.firstCol()}
|
||||
rows, err := s.Query(ctx, q)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
ok := rows.Next()
|
||||
return ok, rows.Err()
|
||||
}
|
||||
|
||||
func (s *Select) firstCol() lookup.Column {
|
||||
if len(s.cols) > 0 {
|
||||
return s.cols[0]
|
||||
}
|
||||
return lookup.Column{Entity: s.from, Name: lookup.FirstColumn(s.from)}
|
||||
}
|
||||
|
||||
// Assignment is one `col = value` of an UPDATE or INSERT.
|
||||
type Assignment struct {
|
||||
Col lookup.Column
|
||||
Val any
|
||||
}
|
||||
|
||||
// Set builds an Assignment.
|
||||
func Set(c lookup.Column, v any) Assignment { return Assignment{Col: c, Val: v} }
|
||||
|
||||
// Update builds an UPDATE.
|
||||
type Update struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
sets []Assignment
|
||||
conds []Cond
|
||||
}
|
||||
|
||||
// Update starts an UPDATE of e.
|
||||
func (b *Base) Update(e lookup.Entity) *Update { return &Update{b: b, e: e} }
|
||||
|
||||
// Set adds assignments.
|
||||
func (u *Update) Set(as ...Assignment) *Update { u.sets = append(u.sets, as...); return u }
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (u *Update) Where(cs ...Cond) *Update { u.conds = append(u.conds, cs...); return u }
|
||||
|
||||
// Exec runs the update and returns the affected row count.
|
||||
func (u *Update) Exec(ctx context.Context, q Querier) (int64, error) {
|
||||
bl := &builder{b: u.b}
|
||||
set := make([]string, len(u.sets))
|
||||
for i, a := range u.sets {
|
||||
set[i] = bl.col(a.Col) + " = " + bl.ph(a.Val)
|
||||
}
|
||||
sqlText := "UPDATE " + u.b.d.Quote(u.b.tableRef(u.e)) + " SET " + strings.Join(set, ", ") + bl.where(u.conds)
|
||||
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
}
|
||||
|
||||
// Delete builds a DELETE.
|
||||
type Delete struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
conds []Cond
|
||||
}
|
||||
|
||||
// Delete starts a DELETE on e.
|
||||
func (b *Base) Delete(e lookup.Entity) *Delete { return &Delete{b: b, e: e} }
|
||||
|
||||
// Where adds AND-ed conditions.
|
||||
func (d *Delete) Where(cs ...Cond) *Delete { d.conds = append(d.conds, cs...); return d }
|
||||
|
||||
// Exec runs the delete and returns the affected row count.
|
||||
func (d *Delete) Exec(ctx context.Context, q Querier) (int64, error) {
|
||||
bl := &builder{b: d.b}
|
||||
sqlText := "DELETE FROM " + d.b.d.Quote(d.b.tableRef(d.e)) + bl.where(d.conds)
|
||||
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
}
|
||||
|
||||
// Insert builds an INSERT.
|
||||
type Insert struct {
|
||||
b *Base
|
||||
e lookup.Entity
|
||||
sets []Assignment
|
||||
}
|
||||
|
||||
// Insert starts an INSERT into e.
|
||||
func (b *Base) Insert(e lookup.Entity) *Insert { return &Insert{b: b, e: e} }
|
||||
|
||||
// Set adds assignments.
|
||||
func (i *Insert) Set(as ...Assignment) *Insert { i.sets = append(i.sets, as...); return i }
|
||||
|
||||
func (i *Insert) colsAndArgs() ([]string, []any) {
|
||||
cols := make([]string, len(i.sets))
|
||||
args := make([]any, len(i.sets))
|
||||
for n, a := range i.sets {
|
||||
cols[n] = i.b.colName(a.Col)
|
||||
args[n] = i.b.arg(a.Val)
|
||||
}
|
||||
return cols, args
|
||||
}
|
||||
|
||||
// Exec runs the insert.
|
||||
func (i *Insert) Exec(ctx context.Context, q Querier) error {
|
||||
cols, args := i.colsAndArgs()
|
||||
ph := make([]string, len(cols))
|
||||
qc := make([]string, len(cols))
|
||||
for n, c := range cols {
|
||||
qc[n] = i.b.d.Quote(c)
|
||||
ph[n] = i.b.d.Placeholder(n + 1)
|
||||
}
|
||||
sqlText := "INSERT INTO " + i.b.d.Quote(i.b.tableRef(i.e)) + " (" + strings.Join(qc, ", ") + ") VALUES (" + strings.Join(ph, ", ") + ")"
|
||||
_, err := q.ExecContext(ctx, sqlText, args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
|
||||
return err
|
||||
}
|
||||
|
||||
// ExecID runs the insert and returns the generated value of idCol, using the dialect's
|
||||
// insert-returning-id strategy.
|
||||
func (i *Insert) ExecID(ctx context.Context, q Querier, idCol lookup.Column) (int64, error) {
|
||||
cols, args := i.colsAndArgs()
|
||||
ins := i.b.d.InsertReturningID(i.b.tableRef(i.e), cols, i.b.colName(idCol))
|
||||
return ins.Run(ctx, q, args...)
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
)
|
||||
|
||||
type noRun struct{}
|
||||
|
||||
func (noRun) Run(func(*sql.DB) error) error { return nil }
|
||||
|
||||
func TestBuilderRendersPerDialect(t *testing.T) {
|
||||
schema := lookup.Schema{lookup.EntityUserSessions: {Schema: "auth", Name: "sessions"}}
|
||||
cases := map[string]string{
|
||||
"postgres": `SELECT t0."session_token" FROM "auth"."sessions" t0`,
|
||||
"mysql": "SELECT t0.`session_token` FROM `auth`.`sessions` t0",
|
||||
"mssql": `SELECT t0.[session_token] FROM [auth].[sessions] t0`,
|
||||
}
|
||||
tails := map[string]string{
|
||||
"postgres": ` WHERE t0."user_id" = $1 AND t0."session_token" IN ($2, $3)`,
|
||||
"mysql": " WHERE t0.`user_id` = ? AND t0.`session_token` IN (?, ?)",
|
||||
"mssql": ` WHERE t0.[user_id] = @p1 AND t0.[session_token] IN (@p2, @p3)`,
|
||||
}
|
||||
for name, head := range cases {
|
||||
d, err := dialect.Get(name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := NewBase(noRun{}, d, schema)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
st := b.From(lookup.EntityUserSessions).Cols(lookup.SessionsToken).
|
||||
Where(Eq(lookup.SessionsUserID, 7), In(lookup.SessionsToken, "a", "b")).Build()
|
||||
if st.SQL != head+tails[name] || len(st.Args) != 3 {
|
||||
t.Errorf("%s:\n got %s\n want %s", name, st.SQL, head+tails[name])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuilderBoolsGoThroughDialect(t *testing.T) {
|
||||
d, _ := dialect.Get("sqlite")
|
||||
b, _ := NewBase(noRun{}, d, nil)
|
||||
st := b.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersIsActive, true)).Build()
|
||||
if st.Args[0] != d.Bool(true) {
|
||||
t.Fatalf("bool not converted: %#v", st.Args[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuilderSubselectSharesArguments(t *testing.T) {
|
||||
d, _ := dialect.Get("postgres")
|
||||
b, _ := NewBase(noRun{}, d, nil)
|
||||
sub := b.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, 5))
|
||||
st := b.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesID).
|
||||
Where(Eq(lookup.RowRulesTableName, "t"), InSelect(lookup.RowRulesGroupID, sub), Eq(lookup.RowRulesSchemaName, "s")).Build()
|
||||
want := `SELECT t0."id" FROM "sec_row_rules" t0 WHERE t0."table_name" = $1 AND t0."group_id" IN (SELECT t1."group_id" FROM "sec_group_members" t1 WHERE t1."user_id" = $2) AND t0."schema_name" = $3`
|
||||
if st.SQL != want || len(st.Args) != 3 {
|
||||
t.Fatalf("got %s\nwant %s", st.SQL, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaRejectsUnsafeIdentifiers(t *testing.T) {
|
||||
d, _ := dialect.Get("postgres")
|
||||
bad := lookup.Schema{lookup.EntityUsers: {Name: `users"; DROP TABLE x; --`}}
|
||||
if _, err := NewBase(noRun{}, d, bad); err == nil {
|
||||
t.Fatal("unsafe table name accepted")
|
||||
}
|
||||
bad = lookup.Schema{lookup.EntityUsers: {Columns: map[string]string{"username": "a b"}}}
|
||||
if _, err := NewBase(noRun{}, d, bad); err == nil {
|
||||
t.Fatal("unsafe column name accepted")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
_ "github.com/glebarez/go-sqlite"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T, extraDDL ...string) *sql.DB {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", ":memory:")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
ref, err := ddl.SQL("sqlite")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, s := range append([]string{ref}, extraDDL...) {
|
||||
if _, err := db.Exec(s); err != nil {
|
||||
t.Fatalf("ddl: %v", err)
|
||||
}
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestBase(t *testing.T, db *sql.DB, schema lookup.Schema) *Base {
|
||||
t.Helper()
|
||||
d, err := dialect.Detect(db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := NewBase(procedure.NewDB(db, nil, nil), d, schema)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Keys implements lookup.KeyStore on the user keys table. scopes and meta are stored as
|
||||
// JSON through the dialect (native JSON column or TEXT).
|
||||
type Keys struct{ *Base }
|
||||
|
||||
var _ lookup.KeyStore = (*Keys)(nil)
|
||||
|
||||
// NewKeys creates the direct KeyStore.
|
||||
func NewKeys(b *Base) *Keys { return &Keys{Base: b} }
|
||||
|
||||
// Create implements lookup.KeyStore.
|
||||
func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
|
||||
scopes, err := k.d.EncodeJSON(req.Scopes)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal scopes: %w", err)
|
||||
}
|
||||
meta, err := k.d.EncodeJSON(req.Meta)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal meta: %w", err)
|
||||
}
|
||||
now := k.Now()
|
||||
var id int64
|
||||
err = k.do(func(q Querier) error {
|
||||
var err error
|
||||
id, err = k.Insert(lookup.EntityUserKeys).Set(
|
||||
Set(lookup.KeysUserID, req.UserID),
|
||||
Set(lookup.KeysKeyType, string(req.KeyType)),
|
||||
Set(lookup.KeysKeyHash, keyHash),
|
||||
Set(lookup.KeysName, req.Name),
|
||||
Set(lookup.KeysScopes, scopes),
|
||||
Set(lookup.KeysMeta, meta),
|
||||
Set(lookup.KeysExpiresAt, req.ExpiresAt),
|
||||
Set(lookup.KeysCreatedAt, now),
|
||||
Set(lookup.KeysIsActive, true),
|
||||
).ExecID(ctx, q, lookup.KeysID)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create key query failed: %w", err)
|
||||
}
|
||||
return §ypes.UserKey{
|
||||
ID: id,
|
||||
UserID: req.UserID,
|
||||
KeyType: req.KeyType,
|
||||
KeyHash: keyHash,
|
||||
Name: req.Name,
|
||||
Scopes: req.Scopes,
|
||||
Meta: req.Meta,
|
||||
ExpiresAt: req.ExpiresAt,
|
||||
CreatedAt: now,
|
||||
IsActive: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// keyScan holds the destinations for one key row.
|
||||
type keyScan struct {
|
||||
k sectypes.UserKey
|
||||
keyType string
|
||||
scopes, meta any
|
||||
expiresAt, created, lastU time.Time
|
||||
active bool
|
||||
}
|
||||
|
||||
func (k *Keys) keyCols(withLastUsed bool) []lookup.Column {
|
||||
cols := []lookup.Column{lookup.KeysID, lookup.KeysUserID, lookup.KeysKeyType, lookup.KeysName, lookup.KeysScopes,
|
||||
lookup.KeysMeta, lookup.KeysExpiresAt, lookup.KeysCreatedAt, lookup.KeysIsActive}
|
||||
if withLastUsed {
|
||||
cols = append(cols, lookup.KeysLastUsedAt)
|
||||
}
|
||||
return cols
|
||||
}
|
||||
|
||||
func (k *Keys) dest(s *keyScan, withLastUsed bool) []any {
|
||||
d := []any{&s.k.ID, &s.k.UserID, &s.keyType, &s.k.Name, &s.scopes, &s.meta,
|
||||
k.timeDest(&s.expiresAt), k.timeDest(&s.created), k.boolDest(&s.active)}
|
||||
if withLastUsed {
|
||||
d = append(d, k.timeDest(&s.lastU))
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func (k *Keys) finish(s *keyScan) sectypes.UserKey {
|
||||
out := s.k
|
||||
out.KeyType = sectypes.KeyType(s.keyType)
|
||||
out.CreatedAt = s.created
|
||||
out.IsActive = s.active
|
||||
_ = k.d.DecodeJSON(s.scopes, &out.Scopes)
|
||||
_ = k.d.DecodeJSON(s.meta, &out.Meta)
|
||||
if !s.expiresAt.IsZero() {
|
||||
t := s.expiresAt
|
||||
out.ExpiresAt = &t
|
||||
}
|
||||
if !s.lastU.IsZero() {
|
||||
t := s.lastU
|
||||
out.LastUsedAt = &t
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// List implements lookup.KeyStore: active, non-expired keys; an empty keyType means all types.
|
||||
func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
|
||||
keys := []sectypes.UserKey{}
|
||||
conds := []Cond{
|
||||
Eq(lookup.KeysUserID, userID),
|
||||
Eq(lookup.KeysIsActive, true),
|
||||
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())),
|
||||
}
|
||||
if keyType != "" {
|
||||
conds = append(conds, Eq(lookup.KeysKeyType, string(keyType)))
|
||||
}
|
||||
err := k.do(func(q Querier) error {
|
||||
keys = keys[:0]
|
||||
rows, err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(true)...).Where(conds...).OrderBy(lookup.KeysID).Query(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var s keyScan
|
||||
if err := rows.Scan(k.dest(&s, true)...); err != nil {
|
||||
return err
|
||||
}
|
||||
keys = append(keys, k.finish(&s))
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get user keys query failed: %w", err)
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// Delete implements lookup.KeyStore: soft-deletes the key after checking ownership and
|
||||
// returns its hash.
|
||||
func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
|
||||
var keyHash string
|
||||
err := k.tx(ctx, func(q Querier) error {
|
||||
match := []Cond{Eq(lookup.KeysID, keyID), Eq(lookup.KeysUserID, userID), Eq(lookup.KeysIsActive, true)}
|
||||
if err := k.From(lookup.EntityUserKeys).Cols(lookup.KeysKeyHash).Where(match...).QueryRow(ctx, q, &keyHash); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysIsActive, false)).Where(match...).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", errors.New("key not found or already deleted")
|
||||
}
|
||||
return "", fmt.Errorf("delete key query failed: %w", err)
|
||||
}
|
||||
return keyHash, nil
|
||||
}
|
||||
|
||||
// Validate implements lookup.KeyStore: finds an active, non-expired key by hash (optionally of
|
||||
// one type) and stamps last_used_at.
|
||||
func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
|
||||
conds := []Cond{
|
||||
Eq(lookup.KeysKeyHash, keyHash),
|
||||
Eq(lookup.KeysIsActive, true),
|
||||
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())),
|
||||
}
|
||||
if keyType != "" {
|
||||
conds = append(conds, Eq(lookup.KeysKeyType, string(keyType)))
|
||||
}
|
||||
var s keyScan
|
||||
now := k.Now()
|
||||
err := k.tx(ctx, func(q Querier) error {
|
||||
s = keyScan{}
|
||||
if err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(false)...).Where(conds...).QueryRow(ctx, q, k.dest(&s, false)...); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, s.k.ID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, errors.New("invalid or expired key")
|
||||
}
|
||||
return nil, fmt.Errorf("validate key query failed: %w", err)
|
||||
}
|
||||
out := k.finish(&s)
|
||||
out.KeyHash = keyHash
|
||||
out.LastUsedAt = &now
|
||||
return &out, nil
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
func TestKeysLifecycle(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newTestDB(t)
|
||||
k := NewKeys(newTestBase(t, db, nil))
|
||||
_, _ = db.Exec(`INSERT INTO users (username,email,password,is_active) VALUES ('u','u@x.io','x',1)`)
|
||||
|
||||
exp := time.Now().Add(time.Hour)
|
||||
created, err := k.Create(ctx, sectypes.CreateKeyRequest{
|
||||
UserID: 1, KeyType: sectypes.KeyTypeHeaderAPI, Name: "ci",
|
||||
Scopes: []string{"read", "write"}, Meta: map[string]any{"env": "prod"}, ExpiresAt: &exp,
|
||||
}, sectypes.HashKey("raw1"))
|
||||
if err != nil || created.ID == 0 {
|
||||
t.Fatalf("%+v %v", created, err)
|
||||
}
|
||||
// nil scopes/meta must not store the JSON text "null"
|
||||
if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeJWTSecret, Name: "j"}, sectypes.HashKey("raw2")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeGenericAPI, Name: "old", ExpiresAt: ptr(time.Now().Add(-time.Hour))}, sectypes.HashKey("raw3")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
all, err := k.List(ctx, 1, "")
|
||||
if err != nil || len(all) != 2 {
|
||||
t.Fatalf("list all: %d %v", len(all), err)
|
||||
}
|
||||
one, _ := k.List(ctx, 1, sectypes.KeyTypeHeaderAPI)
|
||||
if len(one) != 1 || one[0].Name != "ci" || len(one[0].Scopes) != 2 || one[0].Meta["env"] != "prod" || one[0].ExpiresAt == nil {
|
||||
t.Fatalf("list typed: %+v", one)
|
||||
}
|
||||
if other, _ := k.List(ctx, 2, ""); len(other) != 0 {
|
||||
t.Fatal("other user's keys listed")
|
||||
}
|
||||
|
||||
got, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeHeaderAPI)
|
||||
if err != nil || got.UserID != 1 || got.KeyHash != sectypes.HashKey("raw1") || got.LastUsedAt == nil {
|
||||
t.Fatalf("%+v %v", got, err)
|
||||
}
|
||||
if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeGenericAPI); err == nil || err.Error() != "invalid or expired key" {
|
||||
t.Fatalf("wrong type: %v", err)
|
||||
}
|
||||
if _, err := k.Validate(ctx, sectypes.HashKey("raw3"), ""); err == nil {
|
||||
t.Fatal("expired key validated")
|
||||
}
|
||||
|
||||
if _, err := k.Delete(ctx, 2, created.ID); err == nil || err.Error() != "key not found or already deleted" {
|
||||
t.Fatalf("foreign delete: %v", err)
|
||||
}
|
||||
hash, err := k.Delete(ctx, 1, created.ID)
|
||||
if err != nil || hash != sectypes.HashKey("raw1") {
|
||||
t.Fatalf("%q %v", hash, err)
|
||||
}
|
||||
if _, err := k.Delete(ctx, 1, created.ID); err == nil {
|
||||
t.Fatal("double delete succeeded")
|
||||
}
|
||||
if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), ""); err == nil {
|
||||
t.Fatal("deleted key validated")
|
||||
}
|
||||
}
|
||||
|
||||
func ptr[T any](v T) *T { return &v }
|
||||
@@ -0,0 +1,380 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// nullIfEmpty keeps optional TEXT columns (e.g. client_secret_hash of a public client) NULL
|
||||
// rather than "".
|
||||
func nullIfEmpty(s string) any {
|
||||
if s == "" {
|
||||
return nil
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// OAuthClients implements lookup.OAuthClientStore. Array columns (redirect_uris, grant_types,
|
||||
// allowed_scopes, scopes) are JSON through the dialect.
|
||||
type OAuthClients struct{ *Base }
|
||||
|
||||
var _ lookup.OAuthClientStore = (*OAuthClients)(nil)
|
||||
|
||||
// NewOAuthClients creates the direct OAuthClientStore.
|
||||
func NewOAuthClients(b *Base) *OAuthClients { return &OAuthClients{Base: b} }
|
||||
|
||||
// RegisterClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
|
||||
grantTypes := client.GrantTypes
|
||||
if len(grantTypes) == 0 {
|
||||
grantTypes = []string{"authorization_code"}
|
||||
}
|
||||
allowedScopes := client.AllowedScopes
|
||||
if len(allowedScopes) == 0 {
|
||||
allowedScopes = []string{"openid", "profile", "email"}
|
||||
}
|
||||
authMethod := client.TokenEndpointAuthMethod
|
||||
if authMethod == "" {
|
||||
authMethod = "none"
|
||||
}
|
||||
redirects, err := o.d.EncodeJSON(client.RedirectURIs)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal redirect_uris: %w", err)
|
||||
}
|
||||
if redirects == nil { // the column is NOT NULL
|
||||
redirects = "[]"
|
||||
}
|
||||
grants, err := o.d.EncodeJSON(grantTypes)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal grant_types: %w", err)
|
||||
}
|
||||
scopes, err := o.d.EncodeJSON(allowedScopes)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
|
||||
}
|
||||
|
||||
err = o.do(func(q Querier) error {
|
||||
return o.Insert(lookup.EntityOAuthClients).Set(
|
||||
Set(lookup.OAuthClientsClientID, client.ClientID),
|
||||
Set(lookup.OAuthClientsRedirectURIs, redirects),
|
||||
Set(lookup.OAuthClientsClientName, client.ClientName),
|
||||
Set(lookup.OAuthClientsGrantTypes, grants),
|
||||
Set(lookup.OAuthClientsAllowedScopes, scopes),
|
||||
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
|
||||
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
|
||||
Set(lookup.OAuthClientsIsActive, true),
|
||||
Set(lookup.OAuthClientsCreatedAt, o.Now()),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||
}
|
||||
return §ypes.OAuthServerClient{
|
||||
ClientID: client.ClientID,
|
||||
RedirectURIs: client.RedirectURIs,
|
||||
ClientName: client.ClientName,
|
||||
GrantTypes: grantTypes,
|
||||
AllowedScopes: allowedScopes,
|
||||
ClientSecretHash: client.ClientSecretHash,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||
var redirects, grants, scopes any
|
||||
var name, secret, method sql.NullString
|
||||
err := o.do(func(q Querier) error {
|
||||
return o.From(lookup.EntityOAuthClients).
|
||||
Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes,
|
||||
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod).
|
||||
Where(Eq(lookup.OAuthClientsClientID, clientID), Eq(lookup.OAuthClientsIsActive, true)).
|
||||
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("client not found")
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||
}
|
||||
res := §ypes.OAuthServerClient{
|
||||
ClientID: clientID,
|
||||
ClientName: name.String,
|
||||
ClientSecretHash: secret.String,
|
||||
TokenEndpointAuthMethod: method.String,
|
||||
}
|
||||
_ = o.d.DecodeJSON(redirects, &res.RedirectURIs)
|
||||
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
|
||||
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// SaveCode implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
|
||||
scopes, err := o.d.EncodeJSON(code.Scopes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal scopes: %w", err)
|
||||
}
|
||||
method := code.CodeChallengeMethod
|
||||
if method == "" {
|
||||
method = "S256"
|
||||
}
|
||||
return o.do(func(q Querier) error {
|
||||
return o.Insert(lookup.EntityOAuthCodes).Set(
|
||||
Set(lookup.OAuthCodesCode, code.Code),
|
||||
Set(lookup.OAuthCodesClientID, code.ClientID),
|
||||
Set(lookup.OAuthCodesRedirectURI, code.RedirectURI),
|
||||
Set(lookup.OAuthCodesClientState, code.ClientState),
|
||||
Set(lookup.OAuthCodesCodeChallenge, code.CodeChallenge),
|
||||
Set(lookup.OAuthCodesCodeChallengeMethod, method),
|
||||
Set(lookup.OAuthCodesSessionToken, code.SessionToken),
|
||||
Set(lookup.OAuthCodesRefreshToken, code.RefreshToken),
|
||||
Set(lookup.OAuthCodesScopes, scopes),
|
||||
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
|
||||
Set(lookup.OAuthCodesCreatedAt, o.Now()),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
}
|
||||
|
||||
// ExchangeCode implements lookup.OAuthClientStore: the code is consumed in a transaction and
|
||||
// only the caller whose delete removes the row gets it, so a code cannot be redeemed twice.
|
||||
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||
var res sectypes.OAuthCode
|
||||
var state, refresh sql.NullString
|
||||
var scopes any
|
||||
err := o.tx(ctx, func(q Querier) error {
|
||||
err := o.From(lookup.EntityOAuthCodes).
|
||||
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
|
||||
lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken,
|
||||
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes).
|
||||
Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())).
|
||||
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
|
||||
&res.SessionToken, &refresh, &scopes)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, err := o.Delete(lookup.EntityOAuthCodes).Where(Eq(lookup.OAuthCodesCode, code)).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return sql.ErrNoRows
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("invalid or expired code")
|
||||
}
|
||||
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||
}
|
||||
res.Code = code
|
||||
res.ClientState = state.String
|
||||
res.RefreshToken = refresh.String
|
||||
_ = o.d.DecodeJSON(scopes, &res.Scopes)
|
||||
return &res, nil
|
||||
}
|
||||
|
||||
// Introspect implements lookup.OAuthClientStore (RFC 7662). An unknown or expired token is
|
||||
// {active:false}, not an error.
|
||||
func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
|
||||
var info sectypes.OAuthTokenInfo
|
||||
var userID int
|
||||
var username, email, roles sql.NullString
|
||||
var level sql.NullInt64
|
||||
var exp, iat time.Time
|
||||
err := o.do(func(q Querier) error {
|
||||
return o.From(lookup.EntityUserSessions).
|
||||
Cols(lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
|
||||
lookup.SessionsExpiresAt, lookup.SessionsCreatedAt).
|
||||
Join(lookup.EntityUsers, EqCol(lookup.UsersID, lookup.SessionsUserID)).
|
||||
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, o.Now()), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, &userID, &username, &email, &level, &roles, o.timeDest(&exp), o.timeDest(&iat))
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return §ypes.OAuthTokenInfo{Active: false}, nil
|
||||
}
|
||||
return nil, fmt.Errorf("failed to introspect token: %w", err)
|
||||
}
|
||||
info.Active = true
|
||||
info.Sub = fmt.Sprintf("%d", userID)
|
||||
info.Username = username.String
|
||||
info.Email = email.String
|
||||
info.UserLevel = int(level.Int64)
|
||||
info.Roles = ParseRoles(roles.String)
|
||||
if !exp.IsZero() {
|
||||
info.Exp = exp.Unix()
|
||||
}
|
||||
if !iat.IsZero() {
|
||||
info.Iat = iat.Unix()
|
||||
}
|
||||
return &info, nil
|
||||
}
|
||||
|
||||
// Revoke implements lookup.OAuthClientStore (RFC 7009): the session is deleted; an unknown
|
||||
// token is not an error.
|
||||
func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
|
||||
return o.do(func(q Querier) error {
|
||||
_, err := o.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, token)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// OAuthUsers implements lookup.OAuthUserStore.
|
||||
type OAuthUsers struct{ *Base }
|
||||
|
||||
var _ lookup.OAuthUserStore = (*OAuthUsers)(nil)
|
||||
|
||||
// NewOAuthUsers creates the direct OAuthUserStore.
|
||||
func NewOAuthUsers(b *Base) *OAuthUsers { return &OAuthUsers{Base: b} }
|
||||
|
||||
// GetOrCreateUser implements lookup.OAuthUserStore: select by email, then update or insert,
|
||||
// in one transaction (no upsert).
|
||||
func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
|
||||
roles := strings.Join(user.Roles, ",")
|
||||
var userID int
|
||||
err := o.tx(ctx, func(q Querier) error {
|
||||
now := o.Now()
|
||||
var remoteID, authProvider sql.NullString
|
||||
err := o.From(lookup.EntityUsers).Cols(lookup.UsersID, lookup.UsersRemoteID, lookup.UsersAuthProvider).
|
||||
Where(Eq(lookup.UsersEmail, user.Email)).QueryRow(ctx, q, &userID, &remoteID, &authProvider)
|
||||
if err == nil {
|
||||
// remote_id and auth_provider are only filled when still unset.
|
||||
sets := []Assignment{Set(lookup.UsersLastLoginAt, now), Set(lookup.UsersUpdatedAt, now)}
|
||||
if !remoteID.Valid {
|
||||
sets = append(sets, Set(lookup.UsersRemoteID, user.RemoteID))
|
||||
}
|
||||
if !authProvider.Valid {
|
||||
sets = append(sets, Set(lookup.UsersAuthProvider, provider))
|
||||
}
|
||||
_, err := o.Update(lookup.EntityUsers).Set(sets...).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
return err
|
||||
}
|
||||
if !errors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
id, err := o.Insert(lookup.EntityUsers).Set(
|
||||
Set(lookup.UsersUsername, user.UserName),
|
||||
Set(lookup.UsersEmail, user.Email),
|
||||
Set(lookup.UsersPassword, nil),
|
||||
Set(lookup.UsersUserLevel, user.UserLevel),
|
||||
Set(lookup.UsersRoles, roles),
|
||||
Set(lookup.UsersIsActive, true),
|
||||
Set(lookup.UsersCreatedAt, now),
|
||||
Set(lookup.UsersUpdatedAt, now),
|
||||
Set(lookup.UsersLastLoginAt, now),
|
||||
Set(lookup.UsersRemoteID, user.RemoteID),
|
||||
Set(lookup.UsersAuthProvider, provider),
|
||||
).ExecID(ctx, q, lookup.UsersID)
|
||||
userID = int(id)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to get or create user: %w", err)
|
||||
}
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
// CreateSession implements lookup.OAuthUserStore: insert, or update when the token exists.
|
||||
func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error {
|
||||
return o.tx(ctx, func(q Querier) error {
|
||||
now := o.Now()
|
||||
exists, err := o.From(lookup.EntityUserSessions).Cols(lookup.SessionsID).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
_, err := o.Update(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsAccessToken, s.AccessToken),
|
||||
Set(lookup.SessionsRefreshToken, s.RefreshToken),
|
||||
Set(lookup.SessionsTokenType, s.TokenType),
|
||||
Set(lookup.SessionsExpiresAt, s.ExpiresAt),
|
||||
Set(lookup.SessionsLastActivityAt, now),
|
||||
).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exec(ctx, q)
|
||||
return err
|
||||
}
|
||||
return o.Insert(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsToken, s.SessionToken),
|
||||
Set(lookup.SessionsUserID, s.UserID),
|
||||
Set(lookup.SessionsExpiresAt, s.ExpiresAt),
|
||||
Set(lookup.SessionsCreatedAt, now),
|
||||
Set(lookup.SessionsLastActivityAt, now),
|
||||
Set(lookup.SessionsAccessToken, s.AccessToken),
|
||||
Set(lookup.SessionsRefreshToken, s.RefreshToken),
|
||||
Set(lookup.SessionsTokenType, s.TokenType),
|
||||
Set(lookup.SessionsAuthProvider, s.Provider),
|
||||
).Exec(ctx, q)
|
||||
})
|
||||
}
|
||||
|
||||
// GetByRefreshToken implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
|
||||
var s lookup.OAuthRefreshSession
|
||||
var access, tokenType sql.NullString
|
||||
err := o.do(func(q Querier) error {
|
||||
return o.From(lookup.EntityUserSessions).
|
||||
Cols(lookup.SessionsUserID, lookup.SessionsAccessToken, lookup.SessionsTokenType, lookup.SessionsExpiresAt).
|
||||
Where(Eq(lookup.SessionsRefreshToken, refreshToken), Gt(lookup.SessionsExpiresAt, o.Now())).
|
||||
QueryRow(ctx, q, &s.UserID, &access, &tokenType, o.timeDest(&s.Expiry))
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("refresh token not found or expired")
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
|
||||
}
|
||||
s.AccessToken = access.String
|
||||
s.TokenType = tokenType.String
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
// UpdateRefreshToken implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
|
||||
var rows int64
|
||||
err := o.do(func(q Querier) error {
|
||||
var err error
|
||||
rows, err = o.Update(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsToken, newSessionToken),
|
||||
Set(lookup.SessionsAccessToken, newAccessToken),
|
||||
Set(lookup.SessionsRefreshToken, newRefreshToken),
|
||||
Set(lookup.SessionsExpiresAt, expiresAt),
|
||||
Set(lookup.SessionsLastActivityAt, o.Now()),
|
||||
).Where(Eq(lookup.SessionsUserID, userID), Eq(lookup.SessionsRefreshToken, oldRefreshToken)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update session: %w", err)
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("session not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUser implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
|
||||
var u userRow
|
||||
err := o.do(func(q Querier) error {
|
||||
return o.From(lookup.EntityUsers).
|
||||
Cols(lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
|
||||
lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
|
||||
Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
return nil, fmt.Errorf("failed to get user data: %w", err)
|
||||
}
|
||||
u.id = userID
|
||||
return u.context(""), nil
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
func TestOAuthClientAndCodeFlow(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newTestDB(t)
|
||||
b := newTestBase(t, db, nil)
|
||||
c := NewOAuthClients(b)
|
||||
|
||||
reg, err := c.RegisterClient(ctx, §ypes.OAuthServerClient{ClientID: "cid", RedirectURIs: []string{"https://a/cb"}, ClientName: "App"})
|
||||
if err != nil || reg.TokenEndpointAuthMethod != "none" || len(reg.GrantTypes) != 1 || len(reg.AllowedScopes) != 3 {
|
||||
t.Fatalf("%+v %v", reg, err)
|
||||
}
|
||||
got, err := c.GetClient(ctx, "cid")
|
||||
if err != nil || got.ClientName != "App" || got.RedirectURIs[0] != "https://a/cb" || got.ClientSecretHash != "" {
|
||||
t.Fatalf("%+v %v", got, err)
|
||||
}
|
||||
if _, err := c.GetClient(ctx, "nope"); err == nil || err.Error() != "client not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
_, _ = db.Exec(`UPDATE oauth_clients SET is_active = 0`)
|
||||
if _, err := c.GetClient(ctx, "cid"); err == nil {
|
||||
t.Fatal("inactive client returned")
|
||||
}
|
||||
|
||||
code := §ypes.OAuthCode{Code: "c1", ClientID: "cid", RedirectURI: "https://a/cb", CodeChallenge: "ch",
|
||||
SessionToken: "st", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)}
|
||||
if err := c.SaveCode(ctx, code); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ex, err := c.ExchangeCode(ctx, "c1")
|
||||
if err != nil || ex.Code != "c1" || ex.CodeChallengeMethod != "S256" || ex.SessionToken != "st" || len(ex.Scopes) != 1 {
|
||||
t.Fatalf("%+v %v", ex, err)
|
||||
}
|
||||
if _, err := c.ExchangeCode(ctx, "c1"); err == nil || err.Error() != "invalid or expired code" {
|
||||
t.Fatalf("code reused: %v", err)
|
||||
}
|
||||
code.Code, code.ExpiresAt = "c2", time.Now().Add(-time.Minute)
|
||||
_ = c.SaveCode(ctx, code)
|
||||
if _, err := c.ExchangeCode(ctx, "c2"); err == nil {
|
||||
t.Fatal("expired code exchanged")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthIntrospectRevoke(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, db := newAuth(t, AuthOptions{})
|
||||
reg := registerUser(t, a, "oli")
|
||||
_, _ = db.Exec(`UPDATE users SET roles='r1,r2', user_level=3`)
|
||||
c := NewOAuthClients(a.Base)
|
||||
|
||||
info, err := c.Introspect(ctx, reg.Token)
|
||||
if err != nil || !info.Active || info.Username != "oli" || info.UserLevel != 3 || len(info.Roles) != 2 || info.Exp == 0 || info.Iat == 0 || info.Sub != "1" {
|
||||
t.Fatalf("%+v %v", info, err)
|
||||
}
|
||||
if err := c.Revoke(ctx, reg.Token); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active {
|
||||
t.Fatalf("%+v %v", info, err)
|
||||
}
|
||||
if err := c.Revoke(ctx, "unknown"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthUsers(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newTestDB(t)
|
||||
o := NewOAuthUsers(newTestBase(t, db, nil))
|
||||
|
||||
id, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-1", Roles: []string{"a"}}, "github")
|
||||
if err != nil || id == 0 {
|
||||
t.Fatalf("%d %v", id, err)
|
||||
}
|
||||
// second login: same user; existing remote_id/auth_provider are kept
|
||||
id2, err := o.GetOrCreateUser(ctx, §ypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-2"}, "google")
|
||||
if err != nil || id2 != id {
|
||||
t.Fatalf("%d %v", id2, err)
|
||||
}
|
||||
var remote, prov string
|
||||
_ = db.QueryRow(`SELECT remote_id, auth_provider FROM users WHERE id=?`, id).Scan(&remote, &prov)
|
||||
if remote != "r-1" || prov != "github" {
|
||||
t.Fatalf("overwrote: %q %q", remote, prov)
|
||||
}
|
||||
|
||||
exp := time.Now().Add(time.Hour)
|
||||
s := lookup.OAuthSession{SessionToken: "s1", UserID: id, AccessToken: "a1", RefreshToken: "r1", TokenType: "Bearer", ExpiresAt: exp, Provider: "github"}
|
||||
if err := o.CreateSession(ctx, s); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s.AccessToken = "a1b" // same token: updated, not duplicated
|
||||
if err := o.CreateSession(ctx, s); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
_ = db.QueryRow(`SELECT COUNT(*) FROM user_sessions`).Scan(&n)
|
||||
if n != 1 {
|
||||
t.Fatalf("sessions: %d", n)
|
||||
}
|
||||
|
||||
ref, err := o.GetByRefreshToken(ctx, "r1")
|
||||
if err != nil || ref.UserID != id || ref.AccessToken != "a1b" || ref.TokenType != "Bearer" || ref.Expiry.IsZero() {
|
||||
t.Fatalf("%+v %v", ref, err)
|
||||
}
|
||||
if _, err := o.GetByRefreshToken(ctx, "zzz"); err == nil {
|
||||
t.Fatal("unknown refresh token accepted")
|
||||
}
|
||||
if err := o.UpdateRefreshToken(ctx, id, "r1", "s2", "a2", "r2", exp); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := o.UpdateRefreshToken(ctx, id, "r1", "s3", "a3", "r3", exp); err == nil || err.Error() != "session not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
u, err := o.GetUser(ctx, id)
|
||||
if err != nil || u.UserName != "gh" || u.UserID != id {
|
||||
t.Fatalf("%+v %v", u, err)
|
||||
}
|
||||
if _, err := o.GetUser(ctx, 999); err == nil || err.Error() != "user not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Passkey implements lookup.PasskeyStore. credential_id, public_key and aaguid are base64
|
||||
// TEXT (not native bytea) and transports is JSON, so one schema works on every dialect.
|
||||
type Passkey struct{ *Base }
|
||||
|
||||
var _ lookup.PasskeyStore = (*Passkey)(nil)
|
||||
|
||||
// NewPasskey creates the direct PasskeyStore.
|
||||
func NewPasskey(b *Base) *Passkey { return &Passkey{Base: b} }
|
||||
|
||||
// Store implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
|
||||
transports, err := p.d.EncodeJSON(rec.Transports)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to marshal transports: %w", err)
|
||||
}
|
||||
var id int64
|
||||
err = p.tx(ctx, func(q Querier) error {
|
||||
exists, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyID).
|
||||
Where(Eq(lookup.PasskeyCredentialID, rec.CredentialID)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return fmt.Errorf("credential already exists")
|
||||
}
|
||||
userExists, err := p.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersID, rec.UserID)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !userExists {
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
now := p.Now()
|
||||
id, err = p.Insert(lookup.EntityUserPasskeyCredentials).Set(
|
||||
Set(lookup.PasskeyUserID, rec.UserID),
|
||||
Set(lookup.PasskeyCredentialID, rec.CredentialID),
|
||||
Set(lookup.PasskeyPublicKey, rec.PublicKey),
|
||||
Set(lookup.PasskeyAttestationType, rec.AttestationType),
|
||||
Set(lookup.PasskeyAAGUID, ""),
|
||||
Set(lookup.PasskeySignCount, int64(rec.SignCount)),
|
||||
Set(lookup.PasskeyTransports, transports),
|
||||
Set(lookup.PasskeyBackupEligible, rec.BackupEligible),
|
||||
Set(lookup.PasskeyBackupState, rec.BackupState),
|
||||
Set(lookup.PasskeyName, rec.Name),
|
||||
Set(lookup.PasskeyCreatedAt, now),
|
||||
Set(lookup.PasskeyLastUsedAt, now),
|
||||
).ExecID(ctx, q, lookup.PasskeyID)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// Get implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Get(ctx context.Context, credentialID string) (int, uint32, error) {
|
||||
var userID int
|
||||
var count int64
|
||||
err := p.do(func(q Querier) error {
|
||||
return p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyUserID, lookup.PasskeySignCount).
|
||||
Where(Eq(lookup.PasskeyCredentialID, credentialID)).QueryRow(ctx, q, &userID, &count)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, 0, fmt.Errorf("credential not found")
|
||||
}
|
||||
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
|
||||
}
|
||||
return userID, uint32(count), nil //nolint:gosec // sign counters are stored from uint32 values
|
||||
}
|
||||
|
||||
// UpdateCounter implements lookup.PasskeyStore. A counter that did not advance flags the
|
||||
// credential as possibly cloned and leaves the stored counter unchanged.
|
||||
func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
|
||||
var clone bool
|
||||
err := p.tx(ctx, func(q Querier) error {
|
||||
match := Eq(lookup.PasskeyCredentialID, credentialID)
|
||||
var old int64
|
||||
if err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeySignCount).Where(match).QueryRow(ctx, q, &old); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return err
|
||||
}
|
||||
if int64(newCounter) <= old {
|
||||
clone = true
|
||||
_, err := p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyCloneWarning, true)).Where(match).Exec(ctx, q)
|
||||
return err
|
||||
}
|
||||
_, err := p.Update(lookup.EntityUserPasskeyCredentials).
|
||||
Set(Set(lookup.PasskeySignCount, int64(newCounter)), Set(lookup.PasskeyLastUsedAt, p.Now())).
|
||||
Where(match).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
return clone, err
|
||||
}
|
||||
|
||||
// List implements lookup.PasskeyStore, newest first.
|
||||
func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
|
||||
var out []sectypes.PasskeyCredential
|
||||
err := p.do(func(q Querier) error {
|
||||
out = make([]sectypes.PasskeyCredential, 0)
|
||||
rows, err := p.From(lookup.EntityUserPasskeyCredentials).
|
||||
Cols(lookup.PasskeyID, lookup.PasskeyUserID, lookup.PasskeyCredentialID, lookup.PasskeyPublicKey,
|
||||
lookup.PasskeyAttestationType, lookup.PasskeyAAGUID, lookup.PasskeySignCount, lookup.PasskeyCloneWarning,
|
||||
lookup.PasskeyTransports, lookup.PasskeyBackupEligible, lookup.PasskeyBackupState, lookup.PasskeyName,
|
||||
lookup.PasskeyCreatedAt, lookup.PasskeyLastUsedAt).
|
||||
Where(Eq(lookup.PasskeyUserID, userID)).OrderBy(lookup.PasskeyCreatedAt).Query(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var id, uid int
|
||||
var credB64, pubB64 string
|
||||
var attestation, aaguidB64, name sql.NullString
|
||||
var count sql.NullInt64
|
||||
var clone, eligible, state bool
|
||||
var transports any
|
||||
var created, last time.Time
|
||||
if err := rows.Scan(&id, &uid, &credB64, &pubB64, &attestation, &aaguidB64, &count, p.boolDest(&clone),
|
||||
&transports, p.boolDest(&eligible), p.boolDest(&state), &name, p.timeDest(&created), p.timeDest(&last)); err != nil {
|
||||
return err
|
||||
}
|
||||
credID, err := base64.StdEncoding.DecodeString(credB64)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
pub, err := base64.StdEncoding.DecodeString(pubB64)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
aaguid, _ := base64.StdEncoding.DecodeString(aaguidB64.String)
|
||||
c := sectypes.PasskeyCredential{
|
||||
ID: fmt.Sprintf("%d", id),
|
||||
UserID: uid,
|
||||
CredentialID: credID,
|
||||
PublicKey: pub,
|
||||
AttestationType: attestation.String,
|
||||
AAGUID: aaguid,
|
||||
SignCount: uint32(count.Int64), //nolint:gosec // stored from uint32 values
|
||||
CloneWarning: clone,
|
||||
BackupEligible: eligible,
|
||||
BackupState: state,
|
||||
Name: name.String,
|
||||
CreatedAt: created,
|
||||
LastUsedAt: last,
|
||||
}
|
||||
_ = p.d.DecodeJSON(transports, &c.Transports)
|
||||
out = append(out, c)
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||
}
|
||||
// newest first
|
||||
for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 {
|
||||
out[i], out[j] = out[j], out[i]
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Delete implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error {
|
||||
var rows int64
|
||||
err := p.do(func(q Querier) error {
|
||||
var err error
|
||||
rows, err = p.Base.Delete(lookup.EntityUserPasskeyCredentials).
|
||||
Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Rename implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error {
|
||||
var rows int64
|
||||
err := p.do(func(q Querier) error {
|
||||
var err error
|
||||
rows, err = p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyName, name)).
|
||||
Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("credential not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ByUsername implements lookup.PasskeyStore.
|
||||
func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
|
||||
var userID int
|
||||
var creds []lookup.PasskeyCredentialRef
|
||||
err := p.do(func(q Querier) error {
|
||||
creds = make([]lookup.PasskeyCredentialRef, 0)
|
||||
if err := p.From(lookup.EntityUsers).Cols(lookup.UsersID).
|
||||
Where(Eq(lookup.UsersUsername, username), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID); err != nil {
|
||||
return err
|
||||
}
|
||||
rows, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyCredentialID, lookup.PasskeyTransports).
|
||||
Where(Eq(lookup.PasskeyUserID, userID)).Query(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var ref lookup.PasskeyCredentialRef
|
||||
var transports any
|
||||
if err := rows.Scan(&ref.CredentialID, &transports); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = p.d.DecodeJSON(transports, &ref.Transports)
|
||||
creds = append(creds, ref)
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, nil, fmt.Errorf("user not found")
|
||||
}
|
||||
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||
}
|
||||
return userID, creds, nil
|
||||
}
|
||||
|
||||
// Login implements lookup.PasskeyStore: it creates the session for a user whose passkey
|
||||
// assertion was already verified.
|
||||
func (p *Passkey) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
var u userRow
|
||||
err := p.do(func(q Querier) error {
|
||||
return p.From(lookup.EntityUsers).
|
||||
Cols(lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
|
||||
lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
|
||||
Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersIsActive, true)).
|
||||
QueryRow(ctx, q, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
return nil, fmt.Errorf("passkey login query failed: %w", err)
|
||||
}
|
||||
u.id = userID
|
||||
token, err := GenerateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
}
|
||||
now := p.Now()
|
||||
ip, ua := claimStrings(claims)
|
||||
err = p.tx(ctx, func(q Querier) error {
|
||||
if err := p.Insert(lookup.EntityUserSessions).Set(
|
||||
Set(lookup.SessionsToken, token),
|
||||
Set(lookup.SessionsUserID, userID),
|
||||
Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)),
|
||||
Set(lookup.SessionsIPAddress, ip),
|
||||
Set(lookup.SessionsUserAgent, ua),
|
||||
Set(lookup.SessionsLastActivityAt, now),
|
||||
Set(lookup.SessionsCreatedAt, now),
|
||||
).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := p.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("passkey login query failed: %w", err)
|
||||
}
|
||||
return §ypes.LoginResponse{Token: token, User: u.context(token), ExpiresIn: int64(sessionLifetime.Seconds())}, nil
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) }
|
||||
|
||||
func TestPasskeyLifecycle(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, _ := newAuth(t, AuthOptions{})
|
||||
reg := registerUser(t, a, "pat")
|
||||
uid := reg.User.UserID
|
||||
p := NewPasskey(a.Base)
|
||||
|
||||
rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: b64("cred1"), PublicKey: b64("pk"), AttestationType: "none",
|
||||
Transports: []string{"usb", "nfc"}, Name: "Key 1"}
|
||||
id, err := p.Store(ctx, rec)
|
||||
if err != nil || id == 0 {
|
||||
t.Fatalf("%d %v", id, err)
|
||||
}
|
||||
if _, err := p.Store(ctx, rec); err == nil || err.Error() != "credential already exists" {
|
||||
t.Fatalf("dup: %v", err)
|
||||
}
|
||||
rec.CredentialID, rec.UserID = b64("cred2"), 999
|
||||
if _, err := p.Store(ctx, rec); err == nil || err.Error() != "user not found" {
|
||||
t.Fatalf("no user: %v", err)
|
||||
}
|
||||
rec.UserID, rec.Name = uid, "Key 2"
|
||||
if _, err := p.Store(ctx, rec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
owner, count, err := p.Get(ctx, b64("cred1"))
|
||||
if err != nil || owner != uid || count != 0 {
|
||||
t.Fatalf("%d %d %v", owner, count, err)
|
||||
}
|
||||
if _, _, err := p.Get(ctx, b64("zzz")); err == nil || err.Error() != "credential not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || clone {
|
||||
t.Fatalf("%v %v", clone, err)
|
||||
}
|
||||
if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || !clone {
|
||||
t.Fatalf("replayed counter must flag clone: %v %v", clone, err)
|
||||
}
|
||||
if _, count, _ = p.Get(ctx, b64("cred1")); count != 5 {
|
||||
t.Fatalf("counter changed on clone: %d", count)
|
||||
}
|
||||
if _, err := p.UpdateCounter(ctx, b64("missing"), 1); err == nil {
|
||||
t.Fatal("expected not found")
|
||||
}
|
||||
|
||||
list, err := p.List(ctx, uid)
|
||||
if err != nil || len(list) != 2 {
|
||||
t.Fatalf("%d %v", len(list), err)
|
||||
}
|
||||
for _, c := range list {
|
||||
if string(c.CredentialID) == "cred1" {
|
||||
if !c.CloneWarning || c.SignCount != 5 || len(c.Transports) != 2 || c.Name != "Key 1" {
|
||||
t.Fatalf("%+v", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if err := p.Rename(ctx, uid, b64("cred1"), "Renamed"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.Rename(ctx, uid+1, b64("cred1"), "x"); err == nil {
|
||||
t.Fatal("renamed another user's credential")
|
||||
}
|
||||
|
||||
gotID, refs, err := p.ByUsername(ctx, "pat")
|
||||
if err != nil || gotID != uid || len(refs) != 2 {
|
||||
t.Fatalf("%d %+v %v", gotID, refs, err)
|
||||
}
|
||||
if _, _, err := p.ByUsername(ctx, "ghost"); err == nil || err.Error() != "user not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "1.1.1.1"})
|
||||
if err != nil || resp.User.UserName != "pat" || resp.ExpiresIn != 86400 {
|
||||
t.Fatalf("%+v %v", resp, err)
|
||||
}
|
||||
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := p.Delete(ctx, uid+1, b64("cred1")); err == nil {
|
||||
t.Fatal("deleted another user's credential")
|
||||
}
|
||||
if err := p.Delete(ctx, uid, b64("cred1")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := p.Delete(ctx, uid, b64("cred1")); err == nil || err.Error() != "credential not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTOTPLifecycle(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
a, _ := newAuth(t, AuthOptions{})
|
||||
uid := registerUser(t, a, "tom").User.UserID
|
||||
s := NewTOTP(a.Base)
|
||||
|
||||
if on, err := s.Status(ctx, uid); err != nil || on {
|
||||
t.Fatalf("%v %v", on, err)
|
||||
}
|
||||
if _, err := s.Secret(ctx, uid); err == nil || err.Error() != "TOTP not enabled for user" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if err := s.RegenerateBackupCodes(ctx, uid, []string{"h"}); err == nil {
|
||||
t.Fatal("regenerate without 2FA")
|
||||
}
|
||||
if err := s.Enable(ctx, 999, "S", nil); err == nil || err.Error() != "user not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
if err := s.Enable(ctx, uid, "SECRET", []string{"h1", "h2"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if on, _ := s.Status(ctx, uid); !on {
|
||||
t.Fatal("not enabled")
|
||||
}
|
||||
if sec, err := s.Secret(ctx, uid); err != nil || sec != "SECRET" {
|
||||
t.Fatalf("%q %v", sec, err)
|
||||
}
|
||||
|
||||
if ok, err := s.ValidateBackupCode(ctx, uid, "h1"); err != nil || !ok {
|
||||
t.Fatalf("%v %v", ok, err)
|
||||
}
|
||||
if _, err := s.ValidateBackupCode(ctx, uid, "h1"); err == nil || err.Error() != "backup code already used" {
|
||||
t.Fatalf("reuse: %v", err)
|
||||
}
|
||||
if ok, err := s.ValidateBackupCode(ctx, uid, "nope"); err != nil || ok {
|
||||
t.Fatalf("%v %v", ok, err)
|
||||
}
|
||||
if err := s.RegenerateBackupCodes(ctx, uid, []string{"n1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ok, _ := s.ValidateBackupCode(ctx, uid, "h2"); ok {
|
||||
t.Fatal("old code survived regenerate")
|
||||
}
|
||||
if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); !ok {
|
||||
t.Fatal("new code rejected")
|
||||
}
|
||||
|
||||
if err := s.Disable(ctx, uid); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if on, _ := s.Status(ctx, uid); on {
|
||||
t.Fatal("still enabled")
|
||||
}
|
||||
if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); ok {
|
||||
t.Fatal("codes survived disable")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// bcrypt only considers the first 72 bytes of input; longer passwords are
|
||||
// rejected rather than silently truncated.
|
||||
const maxPasswordBytes = 72
|
||||
|
||||
var errPasswordTooLong = errors.New("password must be at most 72 bytes")
|
||||
|
||||
// HashPassword returns the bcrypt hash of password.
|
||||
func HashPassword(password string) (string, error) {
|
||||
if len(password) > maxPasswordBytes {
|
||||
return "", errPasswordTooLong
|
||||
}
|
||||
h, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(h), nil
|
||||
}
|
||||
|
||||
func isBcryptHash(s string) bool {
|
||||
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
||||
}
|
||||
|
||||
// VerifyPassword checks supplied against the stored value. A stored bcrypt hash is compared
|
||||
// with bcrypt. A legacy cleartext value (written before hashing was implemented) is compared
|
||||
// in constant time and, on a match, needsRehash is true so the caller can upgrade the row.
|
||||
// An empty stored value (e.g. an OAuth2-only user) never matches.
|
||||
func VerifyPassword(stored, supplied string) (ok, needsRehash bool) {
|
||||
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
|
||||
return false, false
|
||||
}
|
||||
if isBcryptHash(stored) {
|
||||
return bcrypt.CompareHashAndPassword([]byte(stored), []byte(supplied)) == nil, false
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(stored), []byte(supplied)) == 1 {
|
||||
return true, true
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
|
||||
var (
|
||||
dummyHashOnce sync.Once
|
||||
dummyHash string
|
||||
)
|
||||
|
||||
// BurnPasswordCheck spends roughly one bcrypt comparison so an unknown username costs about
|
||||
// the same as a wrong password.
|
||||
func BurnPasswordCheck(supplied string) {
|
||||
dummyHashOnce.Do(func() {
|
||||
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
|
||||
dummyHash = string(h)
|
||||
})
|
||||
if len(supplied) > maxPasswordBytes {
|
||||
supplied = supplied[:maxPasswordBytes]
|
||||
}
|
||||
_ = bcrypt.CompareHashAndPassword([]byte(dummyHash), []byte(supplied))
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package direct
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestVerifyPasswordEdgeCases(t *testing.T) {
|
||||
h, _ := HashPassword("pw")
|
||||
if ok, _ := VerifyPassword(h, "pw"); !ok {
|
||||
t.Error("bcrypt match failed")
|
||||
}
|
||||
if ok, _ := VerifyPassword("", "pw"); ok {
|
||||
t.Error("empty stored must not match")
|
||||
}
|
||||
if ok, _ := VerifyPassword("pw", ""); ok {
|
||||
t.Error("empty supplied must not match")
|
||||
}
|
||||
if _, err := HashPassword(string(make([]byte, 73))); err == nil {
|
||||
t.Error("73-byte password must be rejected")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// PolicyOptions tunes Policy.
|
||||
type PolicyOptions struct {
|
||||
// NoGroups skips the group membership table: only rules addressed to the user directly
|
||||
// apply. Use it when the sec_group_members table is not deployed.
|
||||
NoGroups bool
|
||||
}
|
||||
|
||||
// Policy implements lookup.PolicyStore on the rule tables.
|
||||
//
|
||||
// Applicable rules are the active rules whose user_id is the caller plus the rules of every
|
||||
// group the caller belongs to; schema and table match case-insensitively and exactly (never a
|
||||
// prefix). Column security returns the union of the matching rules. Row security: any
|
||||
// applicable has_block rule wins, otherwise the templates are combined with AND, each in
|
||||
// parentheses. No rule is an empty result; failures are errors so callers fail closed.
|
||||
type Policy struct {
|
||||
*Base
|
||||
opts PolicyOptions
|
||||
}
|
||||
|
||||
var _ lookup.PolicyStore = (*Policy)(nil)
|
||||
|
||||
// NewPolicy creates the direct PolicyStore.
|
||||
func NewPolicy(b *Base, opts PolicyOptions) *Policy { return &Policy{Base: b, opts: opts} }
|
||||
|
||||
// applicable restricts a rule query to the rules that apply to userID.
|
||||
func (p *Policy) applicable(userCol, groupCol lookup.Column, userID int64) Cond {
|
||||
if p.opts.NoGroups {
|
||||
return Eq(userCol, userID)
|
||||
}
|
||||
members := p.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, userID))
|
||||
return Or(Eq(userCol, userID), InSelect(groupCol, members))
|
||||
}
|
||||
|
||||
// ColumnSecurity implements lookup.PolicyStore.
|
||||
func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
|
||||
var rules []sectypes.ColumnSecurity
|
||||
err := p.do(func(q Querier) error {
|
||||
rules = nil
|
||||
rows, err := p.From(lookup.EntitySecColumnRules).
|
||||
Cols(lookup.ColRulesID, lookup.ColRulesColumnPath, lookup.ColRulesAccessType, lookup.ColRulesMaskStart,
|
||||
lookup.ColRulesMaskEnd, lookup.ColRulesMaskInvert, lookup.ColRulesMaskChar, lookup.ColRulesExtraFilters).
|
||||
Where(
|
||||
Eq(lookup.ColRulesIsActive, true),
|
||||
EqFold(lookup.ColRulesSchemaName, schema),
|
||||
EqFold(lookup.ColRulesTableName, table),
|
||||
p.applicable(lookup.ColRulesUserID, lookup.ColRulesGroupID, int64(userID)),
|
||||
).OrderBy(lookup.ColRulesID).Query(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var id int
|
||||
var path, access string
|
||||
var start, end sql.NullInt64
|
||||
var invert sql.NullBool
|
||||
var maskChar sql.NullString
|
||||
var extra any
|
||||
var inv any
|
||||
if err := rows.Scan(&id, &path, &access, &start, &end, &inv, &maskChar, &extra); err != nil {
|
||||
return err
|
||||
}
|
||||
if inv != nil {
|
||||
b, err := p.d.ScanBool(inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
invert = sql.NullBool{Bool: b, Valid: true}
|
||||
}
|
||||
rule := sectypes.ColumnSecurity{
|
||||
ID: id,
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
Path: strings.Split(path, "."),
|
||||
Accesstype: access,
|
||||
UserID: userID,
|
||||
MaskStart: int(start.Int64),
|
||||
MaskEnd: int(end.Int64),
|
||||
MaskInvert: invert.Bool,
|
||||
MaskChar: "*",
|
||||
Control: schema + "." + table + "." + path,
|
||||
}
|
||||
if maskChar.Valid && maskChar.String != "" {
|
||||
rule.MaskChar = maskChar.String
|
||||
}
|
||||
if err := p.d.DecodeJSON(extra, &rule.ExtraFilters); err != nil {
|
||||
return err
|
||||
}
|
||||
rules = append(rules, rule)
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load column security: %w", err)
|
||||
}
|
||||
return rules, nil
|
||||
}
|
||||
|
||||
// numericUser reduces a user reference to the integer the rule tables key on. Structured
|
||||
// values are rejected, and so are non-numeric strings: a reference that cannot be matched
|
||||
// must fail closed rather than silently load no rules.
|
||||
func numericUser(ref any) (int64, error) {
|
||||
switch v := ref.(type) {
|
||||
case *sectypes.UserContext:
|
||||
if v == nil {
|
||||
return 0, fmt.Errorf("row security: nil user context")
|
||||
}
|
||||
return int64(v.UserID), nil
|
||||
case sectypes.UserContext:
|
||||
return int64(v.UserID), nil
|
||||
case int:
|
||||
return int64(v), nil
|
||||
case int8:
|
||||
return int64(v), nil
|
||||
case int16:
|
||||
return int64(v), nil
|
||||
case int32:
|
||||
return int64(v), nil
|
||||
case int64:
|
||||
return v, nil
|
||||
case uint:
|
||||
return int64(v), nil //nolint:gosec // user ids fit int64
|
||||
case uint8:
|
||||
return int64(v), nil
|
||||
case uint16:
|
||||
return int64(v), nil
|
||||
case uint32:
|
||||
return int64(v), nil
|
||||
case uint64:
|
||||
return int64(v), nil //nolint:gosec // user ids fit int64
|
||||
case string:
|
||||
n, err := strconv.ParseInt(strings.TrimSpace(v), 10, 64)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("row security: user reference %q is not a numeric id", v)
|
||||
}
|
||||
return n, nil
|
||||
case nil:
|
||||
return 0, fmt.Errorf("row security: no user reference")
|
||||
}
|
||||
return 0, fmt.Errorf("row security: unsupported user reference type %T", ref)
|
||||
}
|
||||
|
||||
// RowSecurity implements lookup.PolicyStore.
|
||||
func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
|
||||
uid, err := numericUser(userRef)
|
||||
if err != nil {
|
||||
return sectypes.RowSecurity{}, err
|
||||
}
|
||||
var templates []string
|
||||
block := false
|
||||
err = p.do(func(q Querier) error {
|
||||
templates, block = nil, false
|
||||
rows, err := p.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesTemplate, lookup.RowRulesHasBlock).
|
||||
Where(
|
||||
Eq(lookup.RowRulesIsActive, true),
|
||||
EqFold(lookup.RowRulesSchemaName, schema),
|
||||
EqFold(lookup.RowRulesTableName, table),
|
||||
p.applicable(lookup.RowRulesUserID, lookup.RowRulesGroupID, uid),
|
||||
).OrderBy(lookup.RowRulesID).Query(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var tpl sql.NullString
|
||||
var hb bool
|
||||
if err := rows.Scan(&tpl, p.boolDest(&hb)); err != nil {
|
||||
return err
|
||||
}
|
||||
if hb {
|
||||
block = true
|
||||
}
|
||||
if t := strings.TrimSpace(tpl.String); t != "" {
|
||||
templates = append(templates, "("+t+")")
|
||||
}
|
||||
}
|
||||
return rows.Err()
|
||||
})
|
||||
if err != nil {
|
||||
return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
|
||||
}
|
||||
rs := sectypes.RowSecurity{Schema: schema, Tablename: table, UserID: userRef, HasBlock: block}
|
||||
if !block {
|
||||
rs.Template = strings.Join(templates, " AND ")
|
||||
}
|
||||
return rs, nil
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPolicyColumnAndRowSecurity(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newTestDB(t)
|
||||
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{})
|
||||
exec := func(q string, args ...any) {
|
||||
t.Helper()
|
||||
if _, err := db.Exec(q, args...); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
exec(`INSERT INTO sec_group_members (group_id, user_id) VALUES (10, 1), (10, 3)`)
|
||||
|
||||
// column rules: user 1 direct, group 10 (user 1 and 3), another user, inactive, other table, prefix table
|
||||
exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, mask_start, mask_end, mask_invert, mask_char, extra_filters, is_active)
|
||||
VALUES (1, NULL, 'Public', 'Users', 'email', 'mask', 2, 1, 1, '#', '{"k":"v"}', 1),
|
||||
(NULL, 10, 'public', 'users', 'profile.ssn', 'hide', NULL, NULL, NULL, NULL, NULL, 1),
|
||||
(2, NULL, 'public', 'users', 'other', 'hide', 0, 0, 0, '*', NULL, 1),
|
||||
(1, NULL, 'public', 'users', 'off', 'hide', 0, 0, 0, '*', NULL, 0),
|
||||
(1, NULL, 'public', 'orders', 'x', 'hide', 0, 0, 0, '*', NULL, 1),
|
||||
(1, NULL, 'public', 'users_archive', 'y', 'hide', 0, 0, 0, '*', NULL, 1)`)
|
||||
|
||||
rules, err := p.ColumnSecurity(ctx, 1, "public", "users")
|
||||
if err != nil || len(rules) != 2 {
|
||||
t.Fatalf("%d %v %+v", len(rules), err, rules)
|
||||
}
|
||||
m := rules[0]
|
||||
if m.Accesstype != "mask" || m.MaskStart != 2 || m.MaskEnd != 1 || !m.MaskInvert || m.MaskChar != "#" ||
|
||||
m.ExtraFilters["k"] != "v" || len(m.Path) != 1 || m.Path[0] != "email" || m.UserID != 1 {
|
||||
t.Fatalf("%+v", m)
|
||||
}
|
||||
h := rules[1]
|
||||
if len(h.Path) != 2 || h.Path[1] != "ssn" || h.MaskChar != "*" || h.Accesstype != "hide" {
|
||||
t.Fatalf("%+v", h)
|
||||
}
|
||||
if r, err := p.ColumnSecurity(ctx, 3, "public", "users"); err != nil || len(r) != 1 {
|
||||
t.Fatalf("group member: %d %v", len(r), err)
|
||||
}
|
||||
if r, err := p.ColumnSecurity(ctx, 99, "public", "users"); err != nil || len(r) != 0 {
|
||||
t.Fatalf("no rules must be empty: %d %v", len(r), err)
|
||||
}
|
||||
|
||||
// row rules
|
||||
exec(`INSERT INTO sec_row_rules (user_id, group_id, schema_name, table_name, template, has_block, is_active) VALUES
|
||||
(1, NULL, 'public', 'orders', 'owner_id = {UserID}', 0, 1),
|
||||
(NULL, 10, 'public', 'orders', 'region = 1', 0, 1),
|
||||
(NULL, 10, 'public', 'orders', 'ignored', 0, 0),
|
||||
(3, NULL, 'public', 'secret', NULL, 1, 1),
|
||||
(NULL, 10, 'public', 'secret', 'x = 1', 0, 1)`)
|
||||
rs, err := p.RowSecurity(ctx, 1, "public", "orders")
|
||||
if err != nil || rs.Template != "(owner_id = {UserID}) AND (region = 1)" || rs.HasBlock {
|
||||
t.Fatalf("%+v %v", rs, err)
|
||||
}
|
||||
if rs, err := p.RowSecurity(ctx, "3", "PUBLIC", "Secret"); err != nil || !rs.HasBlock || rs.Template != "" {
|
||||
t.Fatalf("block must win: %+v %v", rs, err)
|
||||
}
|
||||
if rs, err := p.RowSecurity(ctx, 99, "public", "orders"); err != nil || rs.Template != "" || rs.HasBlock {
|
||||
t.Fatalf("%+v %v", rs, err)
|
||||
}
|
||||
for _, bad := range []any{nil, "abc", []int{1}, 1.5} {
|
||||
if _, err := p.RowSecurity(ctx, bad, "public", "orders"); err == nil {
|
||||
t.Fatalf("user ref %#v accepted", bad)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPolicyNoGroups(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := newTestDB(t)
|
||||
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{NoGroups: true})
|
||||
_, _ = db.Exec(`DROP TABLE sec_group_members`)
|
||||
_, _ = db.Exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, is_active) VALUES
|
||||
(1, NULL, 's', 't', 'a', 'hide', 1), (NULL, 5, 's', 't', 'b', 'hide', 1)`)
|
||||
rules, err := p.ColumnSecurity(ctx, 1, "s", "t")
|
||||
if err != nil || len(rules) != 1 || rules[0].Path[0] != "a" {
|
||||
t.Fatalf("%+v %v", rules, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPolicyFailsClosedOnMissingTable(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{})
|
||||
_, _ = db.Exec(`DROP TABLE sec_row_rules`)
|
||||
if _, err := p.RowSecurity(context.Background(), 1, "s", "t"); err == nil {
|
||||
t.Fatal("expected error for missing table")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
// TOTP implements lookup.TOTPStore: the secret and enabled flag live on the users table,
|
||||
// backup code hashes in their own table.
|
||||
type TOTP struct{ *Base }
|
||||
|
||||
var _ lookup.TOTPStore = (*TOTP)(nil)
|
||||
|
||||
// NewTOTP creates the direct TOTPStore.
|
||||
func NewTOTP(b *Base) *TOTP { return &TOTP{Base: b} }
|
||||
|
||||
func (t *TOTP) replaceBackupCodes(ctx context.Context, q Querier, userID int, hashed []string) error {
|
||||
if _, err := t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
now := t.Now()
|
||||
for _, h := range hashed {
|
||||
if err := t.Insert(lookup.EntityUserTOTPBackupCodes).Set(
|
||||
Set(lookup.BackupCodesUserID, userID),
|
||||
Set(lookup.BackupCodesCodeHash, h),
|
||||
Set(lookup.BackupCodesUsed, false),
|
||||
Set(lookup.BackupCodesCreatedAt, now),
|
||||
).Exec(ctx, q); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Enable implements lookup.TOTPStore.
|
||||
func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
|
||||
return t.tx(ctx, func(q Querier) error {
|
||||
n, err := t.Update(lookup.EntityUsers).Set(
|
||||
Set(lookup.UsersTOTPSecret, secret),
|
||||
Set(lookup.UsersTOTPEnabled, true),
|
||||
Set(lookup.UsersTOTPEnabledAt, t.Now()),
|
||||
).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
return t.replaceBackupCodes(ctx, q, userID, hashedCodes)
|
||||
})
|
||||
}
|
||||
|
||||
// Disable implements lookup.TOTPStore.
|
||||
func (t *TOTP) Disable(ctx context.Context, userID int) error {
|
||||
return t.tx(ctx, func(q Querier) error {
|
||||
n, err := t.Update(lookup.EntityUsers).Set(Set(lookup.UsersTOTPSecret, nil), Set(lookup.UsersTOTPEnabled, false)).
|
||||
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return fmt.Errorf("user not found")
|
||||
}
|
||||
_, err = t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// Status implements lookup.TOTPStore.
|
||||
func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) {
|
||||
var enabled bool
|
||||
err := t.do(func(q Querier) error {
|
||||
return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)).
|
||||
QueryRow(ctx, q, t.boolDest(&enabled))
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, fmt.Errorf("user not found")
|
||||
}
|
||||
return false, fmt.Errorf("get 2FA status query failed: %w", err)
|
||||
}
|
||||
return enabled, nil
|
||||
}
|
||||
|
||||
// Secret implements lookup.TOTPStore.
|
||||
func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) {
|
||||
var secret sql.NullString
|
||||
var enabled bool
|
||||
err := t.do(func(q Querier) error {
|
||||
return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPSecret, lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)).
|
||||
QueryRow(ctx, q, &secret, t.boolDest(&enabled))
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return "", fmt.Errorf("user not found")
|
||||
}
|
||||
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||
}
|
||||
if !enabled {
|
||||
return "", fmt.Errorf("TOTP not enabled for user")
|
||||
}
|
||||
return secret.String, nil
|
||||
}
|
||||
|
||||
// RegenerateBackupCodes implements lookup.TOTPStore.
|
||||
func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
|
||||
return t.tx(ctx, func(q Querier) error {
|
||||
ok, err := t.From(lookup.EntityUsers).Cols(lookup.UsersID).
|
||||
Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersTOTPEnabled, true)).Exists(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("user not found or TOTP not enabled")
|
||||
}
|
||||
return t.replaceBackupCodes(ctx, q, userID, hashedCodes)
|
||||
})
|
||||
}
|
||||
|
||||
// ValidateBackupCode implements lookup.TOTPStore. An unknown code is (false, nil); a used
|
||||
// code is an error. The code is consumed with a conditional update so it cannot be spent twice.
|
||||
func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
|
||||
var valid bool
|
||||
err := t.tx(ctx, func(q Querier) error {
|
||||
var id int64
|
||||
var used bool
|
||||
err := t.From(lookup.EntityUserTOTPBackupCodes).Cols(lookup.BackupCodesID, lookup.BackupCodesUsed).
|
||||
Where(Eq(lookup.BackupCodesUserID, userID), Eq(lookup.BackupCodesCodeHash, codeHash)).
|
||||
QueryRow(ctx, q, &id, t.boolDest(&used))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if used {
|
||||
return fmt.Errorf("backup code already used")
|
||||
}
|
||||
n, err := t.Update(lookup.EntityUserTOTPBackupCodes).
|
||||
Set(Set(lookup.BackupCodesUsed, true), Set(lookup.BackupCodesUsedAt, t.Now())).
|
||||
Where(Eq(lookup.BackupCodesID, id), Eq(lookup.BackupCodesUsed, false)).Exec(ctx, q)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return fmt.Errorf("backup code already used")
|
||||
}
|
||||
valid = true
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return valid, nil
|
||||
}
|
||||
@@ -0,0 +1,265 @@
|
||||
-- Keystore schema for per-user auth keys
|
||||
-- Apply alongside database_schema.sql (requires the users table)
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_keys (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
key_type VARCHAR(50) NOT NULL,
|
||||
key_hash VARCHAR(64) NOT NULL UNIQUE, -- SHA-256 hex digest (64 chars)
|
||||
name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
scopes TEXT, -- JSON array, e.g. '["read","write"]'
|
||||
meta JSONB,
|
||||
expires_at TIMESTAMP,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
last_used_at TIMESTAMP,
|
||||
is_active BOOLEAN DEFAULT true
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_user_id ON user_keys(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_key_hash ON user_keys(key_hash);
|
||||
CREATE INDEX IF NOT EXISTS idx_user_keys_key_type ON user_keys(key_type);
|
||||
|
||||
-- resolvespec_keystore_get_user_keys
|
||||
-- Returns all active, non-expired keys for a user.
|
||||
-- Pass empty p_key_type to return all key types.
|
||||
CREATE OR REPLACE FUNCTION resolvespec_keystore_get_user_keys(
|
||||
p_user_id INTEGER,
|
||||
p_key_type TEXT DEFAULT ''
|
||||
)
|
||||
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_keys JSONB)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_keys JSONB;
|
||||
BEGIN
|
||||
SELECT COALESCE(
|
||||
jsonb_agg(
|
||||
jsonb_build_object(
|
||||
'id', k.id,
|
||||
'user_id', k.user_id,
|
||||
'key_type', k.key_type,
|
||||
'name', k.name,
|
||||
'scopes', CASE WHEN k.scopes IS NOT NULL THEN k.scopes::jsonb ELSE '[]'::jsonb END,
|
||||
'meta', COALESCE(k.meta, '{}'::jsonb),
|
||||
'expires_at', k.expires_at,
|
||||
'created_at', k.created_at,
|
||||
'last_used_at', k.last_used_at,
|
||||
'is_active', k.is_active
|
||||
)
|
||||
),
|
||||
'[]'::jsonb
|
||||
)
|
||||
INTO v_keys
|
||||
FROM user_keys k
|
||||
WHERE k.user_id = p_user_id
|
||||
AND k.is_active = true
|
||||
AND (k.expires_at IS NULL OR k.expires_at > NOW())
|
||||
AND (p_key_type = '' OR k.key_type = p_key_type);
|
||||
|
||||
RETURN QUERY SELECT true, NULL::TEXT, v_keys;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||
END;
|
||||
$$;
|
||||
|
||||
-- resolvespec_keystore_create_key
|
||||
-- Inserts a new key row. key_hash is provided by the caller (Go hashes the raw key).
|
||||
-- Returns the created key record (without key_hash).
|
||||
CREATE OR REPLACE FUNCTION resolvespec_keystore_create_key(
|
||||
p_request JSONB
|
||||
)
|
||||
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key JSONB)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_id BIGINT;
|
||||
v_created_at TIMESTAMP;
|
||||
v_key JSONB;
|
||||
BEGIN
|
||||
INSERT INTO user_keys (user_id, key_type, key_hash, name, scopes, meta, expires_at)
|
||||
VALUES (
|
||||
(p_request->>'user_id')::INTEGER,
|
||||
p_request->>'key_type',
|
||||
p_request->>'key_hash',
|
||||
COALESCE(p_request->>'name', ''),
|
||||
p_request->>'scopes',
|
||||
p_request->'meta',
|
||||
CASE WHEN p_request->>'expires_at' IS NOT NULL
|
||||
THEN (p_request->>'expires_at')::timestamptz::timestamp
|
||||
ELSE NULL
|
||||
END
|
||||
)
|
||||
RETURNING id, created_at INTO v_id, v_created_at;
|
||||
|
||||
v_key := jsonb_build_object(
|
||||
'id', v_id,
|
||||
'user_id', (p_request->>'user_id')::INTEGER,
|
||||
'key_type', p_request->>'key_type',
|
||||
'name', COALESCE(p_request->>'name', ''),
|
||||
'scopes', CASE WHEN p_request->>'scopes' IS NOT NULL
|
||||
THEN (p_request->>'scopes')::jsonb
|
||||
ELSE '[]'::jsonb END,
|
||||
'meta', COALESCE(p_request->'meta', '{}'::jsonb),
|
||||
'expires_at', p_request->>'expires_at',
|
||||
'created_at', v_created_at,
|
||||
'is_active', true
|
||||
);
|
||||
|
||||
RETURN QUERY SELECT true, NULL::TEXT, v_key;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||
END;
|
||||
$$;
|
||||
|
||||
-- resolvespec_keystore_delete_key
|
||||
-- Soft-deletes a key (is_active = false) after verifying ownership.
|
||||
-- Returns p_key_hash so the caller can invalidate cache entries without a separate query.
|
||||
CREATE OR REPLACE FUNCTION resolvespec_keystore_delete_key(
|
||||
p_user_id INTEGER,
|
||||
p_key_id BIGINT
|
||||
)
|
||||
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key_hash TEXT)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_hash TEXT;
|
||||
BEGIN
|
||||
UPDATE user_keys
|
||||
SET is_active = false
|
||||
WHERE id = p_key_id AND user_id = p_user_id AND is_active = true
|
||||
RETURNING key_hash INTO v_hash;
|
||||
|
||||
IF NOT FOUND THEN
|
||||
RETURN QUERY SELECT false, 'key not found or already deleted'::TEXT, NULL::TEXT;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
RETURN QUERY SELECT true, NULL::TEXT, v_hash;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, NULL::TEXT;
|
||||
END;
|
||||
$$;
|
||||
|
||||
-- resolvespec_keystore_validate_key
|
||||
-- Looks up a key by its SHA-256 hash, checks active status and expiry,
|
||||
-- updates last_used_at, and returns the key record.
|
||||
-- p_key_type can be empty to accept any key type.
|
||||
CREATE OR REPLACE FUNCTION resolvespec_keystore_validate_key(
|
||||
p_key_hash TEXT,
|
||||
p_key_type TEXT DEFAULT ''
|
||||
)
|
||||
RETURNS TABLE(p_success BOOLEAN, p_error TEXT, p_key JSONB)
|
||||
LANGUAGE plpgsql AS $$
|
||||
DECLARE
|
||||
v_key_rec user_keys%ROWTYPE;
|
||||
v_key JSONB;
|
||||
BEGIN
|
||||
SELECT * INTO v_key_rec
|
||||
FROM user_keys
|
||||
WHERE key_hash = p_key_hash
|
||||
AND is_active = true
|
||||
AND (expires_at IS NULL OR expires_at > NOW())
|
||||
AND (p_key_type = '' OR key_type = p_key_type);
|
||||
|
||||
IF NOT FOUND THEN
|
||||
RETURN QUERY SELECT false, 'invalid or expired key'::TEXT, NULL::JSONB;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
UPDATE user_keys SET last_used_at = NOW() WHERE id = v_key_rec.id;
|
||||
|
||||
v_key := jsonb_build_object(
|
||||
'id', v_key_rec.id,
|
||||
'user_id', v_key_rec.user_id,
|
||||
'key_type', v_key_rec.key_type,
|
||||
'name', v_key_rec.name,
|
||||
'scopes', CASE WHEN v_key_rec.scopes IS NOT NULL
|
||||
THEN v_key_rec.scopes::jsonb
|
||||
ELSE '[]'::jsonb END,
|
||||
'meta', COALESCE(v_key_rec.meta, '{}'::jsonb),
|
||||
'expires_at', v_key_rec.expires_at,
|
||||
'created_at', v_key_rec.created_at,
|
||||
'last_used_at', NOW(),
|
||||
'is_active', v_key_rec.is_active
|
||||
);
|
||||
|
||||
RETURN QUERY SELECT true, NULL::TEXT, v_key;
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM, NULL::JSONB;
|
||||
END;
|
||||
$$;
|
||||
|
||||
-- resolvespec_login_api_key - Exchanges a raw API key for a session
|
||||
-- Input: p_request jsonb {api_key: string, claims: object}
|
||||
-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb)
|
||||
-- Requires pgcrypto (digest) and the user_keys table (keystore_schema.sql).
|
||||
-- Only header_api / api keys are accepted. Unknown, expired and inactive keys
|
||||
-- all return the same generic error.
|
||||
CREATE OR REPLACE FUNCTION resolvespec_login_api_key(p_request jsonb)
|
||||
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
|
||||
DECLARE
|
||||
v_raw_key TEXT;
|
||||
v_key_id BIGINT;
|
||||
v_user_id INTEGER;
|
||||
v_username TEXT;
|
||||
v_email TEXT;
|
||||
v_user_level INTEGER;
|
||||
v_roles TEXT;
|
||||
v_program_user_id INTEGER;
|
||||
v_program_user_table TEXT;
|
||||
v_session_token TEXT;
|
||||
v_expires_at TIMESTAMP;
|
||||
v_ip_address TEXT;
|
||||
v_user_agent TEXT;
|
||||
BEGIN
|
||||
v_raw_key := p_request->>'api_key';
|
||||
v_ip_address := p_request->'claims'->>'ip_address';
|
||||
v_user_agent := p_request->'claims'->>'user_agent';
|
||||
|
||||
IF v_raw_key IS NULL OR v_raw_key = '' THEN
|
||||
RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
SELECT k.id, u.id, u.username, u.email, u.user_level, u.roles, u.program_user_id, u.program_user_table
|
||||
INTO v_key_id, v_user_id, v_username, v_email, v_user_level, v_roles, v_program_user_id, v_program_user_table
|
||||
FROM user_keys k
|
||||
JOIN users u ON u.id = k.user_id
|
||||
WHERE k.key_hash = encode(digest(v_raw_key, 'sha256'), 'hex')
|
||||
AND k.key_type IN ('header_api', 'api')
|
||||
AND k.is_active = true
|
||||
AND (k.expires_at IS NULL OR k.expires_at > now())
|
||||
AND u.is_active = true;
|
||||
|
||||
IF NOT FOUND THEN
|
||||
RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text;
|
||||
v_expires_at := now() + interval '24 hours';
|
||||
|
||||
INSERT INTO user_sessions (session_token, user_id, expires_at, ip_address, user_agent, last_activity_at)
|
||||
VALUES (v_session_token, v_user_id, v_expires_at, v_ip_address, v_user_agent, now());
|
||||
|
||||
UPDATE user_keys SET last_used_at = now() WHERE id = v_key_id;
|
||||
UPDATE users SET last_login_at = now() WHERE id = v_user_id;
|
||||
|
||||
RETURN QUERY SELECT
|
||||
true,
|
||||
NULL::text,
|
||||
jsonb_build_object(
|
||||
'token', v_session_token,
|
||||
'user', jsonb_build_object(
|
||||
'user_id', v_user_id,
|
||||
'user_name', v_username,
|
||||
'email', v_email,
|
||||
'user_level', v_user_level,
|
||||
'roles', string_to_array(COALESCE(v_roles, ''), ','),
|
||||
'session_id', v_session_token,
|
||||
'program_user_id', COALESCE(v_program_user_id, 0),
|
||||
'program_user_table', COALESCE(v_program_user_table, '')
|
||||
),
|
||||
'expires_in', 86400
|
||||
);
|
||||
EXCEPTION WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, 'invalid api key'::text, NULL::jsonb;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
@@ -0,0 +1,208 @@
|
||||
// Package lookup owns every database read and write the security package needs.
|
||||
// pkg/security itself contains no SQL: it calls the store interfaces defined here.
|
||||
//
|
||||
// Each store has a procedure implementation (stored procedures, the Postgres default)
|
||||
// and a direct implementation (tables through a dialect-driven query builder). Which
|
||||
// one runs is decided per operation by Config.EffectiveMode.
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// AuthStore covers sessions, login, registration and password reset.
|
||||
type AuthStore interface {
|
||||
Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error)
|
||||
Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error)
|
||||
Logout(ctx context.Context, req sectypes.LogoutRequest) error
|
||||
// Session resolves a session token to its user. reference says where the token came
|
||||
// from ("authenticate", "cookie", "refresh"); the procedure backend passes it through.
|
||||
Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error)
|
||||
// TouchSession records last activity for a session token. user is the context the
|
||||
// session resolved to; the procedure backend passes it to the update procedure.
|
||||
TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error
|
||||
Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error)
|
||||
// LoginAPIKey logs in with a raw header/generic API key. Unknown, expired, inactive and
|
||||
// wrong-type keys all return the same error.
|
||||
LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error)
|
||||
JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error)
|
||||
JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error
|
||||
ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error)
|
||||
ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error
|
||||
}
|
||||
|
||||
// ErrInvalidAPIKey is the single error LoginAPIKey returns for unknown, expired, inactive
|
||||
// and wrong-type keys, so callers cannot tell them apart.
|
||||
var ErrInvalidAPIKey = errors.New("invalid api key")
|
||||
|
||||
// KeyStore persists per-user auth keys. Hashing and raw-key generation happen in Go,
|
||||
// so the store only sees key hashes.
|
||||
type KeyStore interface {
|
||||
Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error)
|
||||
List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error)
|
||||
// Delete soft-deletes a key after verifying ownership and returns its hash so callers can
|
||||
// invalidate caches. The hash is empty when the backend cannot report it.
|
||||
Delete(ctx context.Context, userID int, keyID int64) (keyHash string, err error)
|
||||
Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error)
|
||||
}
|
||||
|
||||
// OAuthClientStore persists the OAuth2 authorization server state (RFC 7591 clients,
|
||||
// authorization codes, token introspection and revocation).
|
||||
type OAuthClientStore interface {
|
||||
RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error)
|
||||
GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error)
|
||||
SaveCode(ctx context.Context, code *sectypes.OAuthCode) error
|
||||
// ExchangeCode atomically consumes an authorization code.
|
||||
ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error)
|
||||
Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error)
|
||||
Revoke(ctx context.Context, token string) error
|
||||
}
|
||||
|
||||
// OAuthSession is the session row written after an OAuth2 client login.
|
||||
type OAuthSession struct {
|
||||
SessionToken string
|
||||
UserID int
|
||||
AccessToken string
|
||||
RefreshToken string
|
||||
TokenType string
|
||||
ExpiresAt time.Time
|
||||
Provider string
|
||||
}
|
||||
|
||||
// OAuthRefreshSession is the stored token state needed to refresh an OAuth2 login.
|
||||
type OAuthRefreshSession struct {
|
||||
UserID int `json:"user_id"`
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
Expiry time.Time `json:"expiry"`
|
||||
}
|
||||
|
||||
// OAuthUserStore persists users and sessions created through OAuth2 client login.
|
||||
type OAuthUserStore interface {
|
||||
GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error)
|
||||
CreateSession(ctx context.Context, session OAuthSession) error
|
||||
GetByRefreshToken(ctx context.Context, refreshToken string) (*OAuthRefreshSession, error)
|
||||
UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error
|
||||
GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error)
|
||||
}
|
||||
|
||||
// PasskeyCredentialRecord is a credential as persisted: byte fields are base64 text.
|
||||
type PasskeyCredentialRecord struct {
|
||||
UserID int
|
||||
CredentialID string // base64
|
||||
PublicKey string // base64
|
||||
AttestationType string
|
||||
SignCount uint32
|
||||
Transports []string
|
||||
BackupEligible bool
|
||||
BackupState bool
|
||||
Name string
|
||||
}
|
||||
|
||||
// PasskeyCredentialRef is a credential id and transports, as returned for a username lookup.
|
||||
type PasskeyCredentialRef struct {
|
||||
CredentialID string `json:"credential_id"`
|
||||
Transports []string `json:"transports"`
|
||||
}
|
||||
|
||||
// PasskeyStore persists WebAuthn credentials.
|
||||
type PasskeyStore interface {
|
||||
Store(ctx context.Context, rec PasskeyCredentialRecord) (int64, error)
|
||||
// Get returns the owner and signature counter of a credential.
|
||||
Get(ctx context.Context, credentialID string) (userID int, signCount uint32, err error)
|
||||
UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (cloneWarning bool, err error)
|
||||
List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error)
|
||||
Delete(ctx context.Context, userID int, credentialID string) error
|
||||
Rename(ctx context.Context, userID int, credentialID, name string) error
|
||||
ByUsername(ctx context.Context, username string) (userID int, creds []PasskeyCredentialRef, err error)
|
||||
Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error)
|
||||
}
|
||||
|
||||
// TOTPStore persists two-factor state. Backup codes arrive already hashed.
|
||||
type TOTPStore interface {
|
||||
Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error
|
||||
Disable(ctx context.Context, userID int) error
|
||||
Status(ctx context.Context, userID int) (bool, error)
|
||||
Secret(ctx context.Context, userID int) (string, error)
|
||||
RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error
|
||||
ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error)
|
||||
}
|
||||
|
||||
// PolicyStore loads column and row security rules. No rules is an empty result, never
|
||||
// an error; failures are errors so callers fail closed.
|
||||
type PolicyStore interface {
|
||||
ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error)
|
||||
RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error)
|
||||
}
|
||||
|
||||
// Provider bundles every store. Security constructors take a Provider.
|
||||
type Provider struct {
|
||||
Auth AuthStore
|
||||
Keys KeyStore
|
||||
OAuthClient OAuthClientStore
|
||||
OAuthUser OAuthUserStore
|
||||
Passkey PasskeyStore
|
||||
TOTP TOTPStore
|
||||
Policy PolicyStore
|
||||
}
|
||||
|
||||
// Config selects dialect, mode and naming. The zero value is valid: dialect detected from
|
||||
// the driver, default mode per dialect, default procedure/table/column names.
|
||||
type Config struct {
|
||||
// Dialect names a registered dialect ("postgres", "sqlite", "mysql", "mssql", or one added
|
||||
// with dialect.Register). Empty = detect from the driver.
|
||||
Dialect string
|
||||
// Mode is the default mode for every operation. See ModeDefault.
|
||||
Mode Mode
|
||||
// Overrides sets the mode per operation, e.g. direct for OpSession, procedure for OpLogin.
|
||||
Overrides map[Op]Mode
|
||||
// Procs overrides procedure names; empty fields keep the default.
|
||||
Procs ProcNames
|
||||
// Schema overrides table and column names; missing entries keep the default.
|
||||
Schema Schema
|
||||
}
|
||||
|
||||
// Resolved is a Config merged with defaults and validated.
|
||||
type Resolved struct {
|
||||
Config
|
||||
Procs ProcNames
|
||||
Schema Schema
|
||||
}
|
||||
|
||||
// Resolve merges c with the defaults and validates the result.
|
||||
func (c Config) Resolve() (*Resolved, error) {
|
||||
if !c.Mode.valid() {
|
||||
return nil, fmt.Errorf("lookup: invalid mode %q", c.Mode)
|
||||
}
|
||||
for op, m := range c.Overrides {
|
||||
if !m.valid() {
|
||||
return nil, fmt.Errorf("lookup: invalid mode %q for %s", m, op)
|
||||
}
|
||||
}
|
||||
if c.Dialect != "" {
|
||||
if _, err := dialect.Get(c.Dialect); err != nil {
|
||||
return nil, fmt.Errorf("lookup: %w", err)
|
||||
}
|
||||
}
|
||||
procs := DefaultProcNames().Merge(c.Procs)
|
||||
if err := procs.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema := DefaultSchema().Merge(c.Schema)
|
||||
if err := schema.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Resolved{Config: c, Procs: procs, Schema: schema}, nil
|
||||
}
|
||||
|
||||
// Registration conflicts reported by AuthStore.Register in direct mode.
|
||||
var (
|
||||
ErrUsernameExists = errors.New("username already exists")
|
||||
ErrEmailExists = errors.New("email already exists")
|
||||
)
|
||||
@@ -0,0 +1,165 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
ddlpkg "github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
|
||||
)
|
||||
|
||||
// TestDefaultSchemaMatchesSQLiteDDL keeps the default schema in step with the reference DDL:
|
||||
// every table and column in the DDL must be a known logical column and vice versa.
|
||||
func TestDefaultSchemaMatchesSQLiteDDL(t *testing.T) {
|
||||
ddl, err := ddlpkg.SQL("sqlite")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tableRe := regexp.MustCompile(`(?s)CREATE TABLE IF NOT EXISTS (\w+) \((.*?)\n\);`)
|
||||
colRe := regexp.MustCompile(`^\s*(\w+)\s+[A-Z]+`)
|
||||
def := DefaultSchema()
|
||||
seen := map[Entity]bool{}
|
||||
for _, m := range tableRe.FindAllStringSubmatch(string(ddl), -1) {
|
||||
e := Entity(m[1])
|
||||
tbl, ok := def[e]
|
||||
if !ok {
|
||||
t.Errorf("DDL table %s has no entity", e)
|
||||
continue
|
||||
}
|
||||
seen[e] = true
|
||||
cols := map[string]bool{}
|
||||
for _, line := range strings.Split(m[2], "\n") {
|
||||
if cm := colRe.FindStringSubmatch(line); cm != nil && cm[1] != "PRIMARY" && cm[1] != "CHECK" && cm[1] != "FOREIGN" {
|
||||
cols[cm[1]] = true
|
||||
if _, ok := tbl.Columns[cm[1]]; !ok {
|
||||
t.Errorf("DDL column %s.%s is not a logical column", e, cm[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
for c := range tbl.Columns {
|
||||
if !cols[c] {
|
||||
t.Errorf("logical column %s.%s is not in the DDL", e, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
for e := range def {
|
||||
if !seen[e] {
|
||||
t.Errorf("entity %s is not in the DDL", e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultSchemaValid(t *testing.T) {
|
||||
if err := DefaultSchema().Validate(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := DefaultProcNames().Validate(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaMergeAndLookup(t *testing.T) {
|
||||
cfg := Config{Schema: Schema{
|
||||
EntityUsers: {Name: "app_users", Schema: "auth", Columns: map[string]string{"username": "login_name"}},
|
||||
}}
|
||||
r, err := cfg.Resolve()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := r.Schema.TableName(EntityUsers); got != "app_users" {
|
||||
t.Errorf("table = %q", got)
|
||||
}
|
||||
if got := r.Schema.SchemaName(EntityUsers); got != "auth" {
|
||||
t.Errorf("schema = %q", got)
|
||||
}
|
||||
if got := r.Schema.Col(UsersUsername); got != "login_name" {
|
||||
t.Errorf("username col = %q", got)
|
||||
}
|
||||
if got := r.Schema.Col(UsersEmail); got != "email" {
|
||||
t.Errorf("email col = %q, want default", got)
|
||||
}
|
||||
// Merge must not mutate the defaults.
|
||||
if DefaultSchema().Col(UsersUsername) != "username" {
|
||||
t.Error("default schema was mutated")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaZeroValueUsesDefaults(t *testing.T) {
|
||||
var s Schema
|
||||
if s.TableName(EntityUserKeys) != "user_keys" || s.Col(KeysKeyHash) != "key_hash" {
|
||||
t.Error("zero schema should fall back to defaults")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveRejectsBadConfig(t *testing.T) {
|
||||
bad := map[string]Config{
|
||||
"table injection": {Schema: Schema{EntityUsers: {Name: "users; DROP TABLE users"}}},
|
||||
"column injection": {Schema: Schema{EntityUsers: {Columns: map[string]string{"id": "id) --"}}}},
|
||||
"schema injection": {Schema: Schema{EntityUsers: {Schema: "a.b"}}},
|
||||
"unknown entity": {Schema: Schema{"nope": {Name: "x"}}},
|
||||
"unknown column": {Schema: Schema{EntityUsers: {Columns: map[string]string{"nope": "x"}}}},
|
||||
"proc injection": {Procs: ProcNames{Login: "f(); --"}},
|
||||
"bad mode": {Mode: "sometimes"},
|
||||
"bad override": {Overrides: map[Op]Mode{OpLogin: "x"}},
|
||||
"bad dialect": {Dialect: "oracle"},
|
||||
}
|
||||
for name, cfg := range bad {
|
||||
if _, err := cfg.Resolve(); err == nil {
|
||||
t.Errorf("%s: expected error", name)
|
||||
}
|
||||
}
|
||||
// A single schema qualifier is allowed on table and procedure names.
|
||||
ok := Config{
|
||||
Schema: Schema{EntityUsers: {Name: "auth.users"}},
|
||||
Procs: ProcNames{Login: "auth.resolvespec_login"},
|
||||
}
|
||||
if _, err := ok.Resolve(); err != nil {
|
||||
t.Errorf("qualified names should be valid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcNamesMerge(t *testing.T) {
|
||||
m := DefaultProcNames().Merge(ProcNames{Login: "custom_login"})
|
||||
if m.Login != "custom_login" {
|
||||
t.Errorf("Login = %q", m.Login)
|
||||
}
|
||||
if m.Register != "resolvespec_register" {
|
||||
t.Errorf("Register = %q, want default", m.Register)
|
||||
}
|
||||
if DefaultProcNames().LoginAPIKey != "resolvespec_login_api_key" {
|
||||
t.Error("LoginAPIKey default missing")
|
||||
}
|
||||
if DefaultProcNames().KeystoreValidateKey != "resolvespec_keystore_validate_key" {
|
||||
t.Error("keystore defaults missing")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEffectiveMode(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
cfg Config
|
||||
dialect string
|
||||
want Mode
|
||||
wantErr bool
|
||||
}{
|
||||
{"pg default", Config{}, DialectPostgres, ModeProcedure, false},
|
||||
{"sqlite default", Config{}, DialectSQLite, ModeDirect, false},
|
||||
{"mysql default", Config{}, DialectMySQL, ModeDirect, false},
|
||||
{"pg direct", Config{Mode: ModeDirect}, DialectPostgres, ModeDirect, false},
|
||||
{"pg auto probes", Config{Mode: ModeAuto}, DialectPostgres, ModeAuto, false},
|
||||
{"sqlite auto is direct", Config{Mode: ModeAuto}, DialectSQLite, ModeDirect, false},
|
||||
{"sqlite procedure rejected", Config{Mode: ModeProcedure}, DialectSQLite, "", true},
|
||||
{"override wins", Config{Overrides: map[Op]Mode{OpSession: ModeDirect}}, DialectPostgres, ModeDirect, false},
|
||||
{"override only for its op", Config{Overrides: map[Op]Mode{OpSession: ModeDirect}}, DialectPostgres, ModeProcedure, false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
op := OpLogin
|
||||
if strings.HasPrefix(c.name, "override wins") {
|
||||
op = OpSession
|
||||
}
|
||||
got, err := c.cfg.EffectiveMode(op, c.dialect)
|
||||
if (err != nil) != c.wantErr || got != c.want {
|
||||
t.Errorf("%s: got (%q, %v), want (%q, err=%v)", c.name, got, err, c.want, c.wantErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package lookup
|
||||
|
||||
import "fmt"
|
||||
|
||||
// Dialect names understood by Config.Dialect. The dialect package (step 2) owns the
|
||||
// implementations; the names are defined here so Config can be validated without it.
|
||||
const (
|
||||
DialectPostgres = "postgres"
|
||||
DialectSQLite = "sqlite"
|
||||
DialectMySQL = "mysql"
|
||||
DialectMSSQL = "mssql"
|
||||
)
|
||||
|
||||
// Mode selects how a store talks to the database.
|
||||
type Mode string
|
||||
|
||||
const (
|
||||
// ModeDefault (the zero value) resolves to ModeProcedure on Postgres and ModeDirect elsewhere.
|
||||
ModeDefault Mode = ""
|
||||
// ModeProcedure always calls the configured stored procedure; a missing procedure is an error.
|
||||
ModeProcedure Mode = "procedure"
|
||||
// ModeDirect always works on the tables through the dialect builder.
|
||||
ModeDirect Mode = "direct"
|
||||
// ModeAuto probes the procedure once per operation on Postgres (cached) and uses it when
|
||||
// present, otherwise direct. Other dialects resolve to ModeDirect.
|
||||
ModeAuto Mode = "auto"
|
||||
)
|
||||
|
||||
func (m Mode) valid() bool {
|
||||
switch m {
|
||||
case ModeDefault, ModeProcedure, ModeDirect, ModeAuto:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Op names one store operation so its mode can be overridden individually.
|
||||
type Op string
|
||||
|
||||
const (
|
||||
OpLogin Op = "login"
|
||||
OpRegister Op = "register"
|
||||
OpLogout Op = "logout"
|
||||
OpSession Op = "session"
|
||||
OpTouchSession Op = "touch_session"
|
||||
OpRefresh Op = "refresh"
|
||||
OpLoginAPIKey Op = "login_api_key"
|
||||
OpJWTLogin Op = "jwt_login"
|
||||
OpJWTLogout Op = "jwt_logout"
|
||||
OpResetRequest Op = "reset_request"
|
||||
OpResetComplete Op = "reset_complete"
|
||||
|
||||
OpKeyCreate Op = "key_create"
|
||||
OpKeyList Op = "key_list"
|
||||
OpKeyDelete Op = "key_delete"
|
||||
OpKeyValidate Op = "key_validate"
|
||||
|
||||
OpOAuthRegisterClient Op = "oauth_register_client"
|
||||
OpOAuthGetClient Op = "oauth_get_client"
|
||||
OpOAuthSaveCode Op = "oauth_save_code"
|
||||
OpOAuthExchangeCode Op = "oauth_exchange_code"
|
||||
OpOAuthIntrospect Op = "oauth_introspect"
|
||||
OpOAuthRevoke Op = "oauth_revoke"
|
||||
|
||||
OpOAuthGetOrCreateUser Op = "oauth_get_or_create_user"
|
||||
OpOAuthCreateSession Op = "oauth_create_session"
|
||||
OpOAuthGetRefreshToken Op = "oauth_get_refresh_token"
|
||||
OpOAuthUpdateRefreshToken Op = "oauth_update_refresh_token"
|
||||
OpOAuthGetUser Op = "oauth_get_user"
|
||||
|
||||
OpPasskeyStore Op = "passkey_store"
|
||||
OpPasskeyGet Op = "passkey_get"
|
||||
OpPasskeyUpdateCounter Op = "passkey_update_counter"
|
||||
OpPasskeyList Op = "passkey_list"
|
||||
OpPasskeyDelete Op = "passkey_delete"
|
||||
OpPasskeyRename Op = "passkey_rename"
|
||||
OpPasskeyByUsername Op = "passkey_by_username"
|
||||
OpPasskeyLogin Op = "passkey_login"
|
||||
|
||||
OpTOTPEnable Op = "totp_enable"
|
||||
OpTOTPDisable Op = "totp_disable"
|
||||
OpTOTPStatus Op = "totp_status"
|
||||
OpTOTPSecret Op = "totp_secret"
|
||||
OpTOTPRegenerateBackup Op = "totp_regenerate_backup"
|
||||
OpTOTPValidateBackupCode Op = "totp_validate_backup_code"
|
||||
|
||||
OpColumnSecurity Op = "column_security"
|
||||
OpRowSecurity Op = "row_security"
|
||||
)
|
||||
|
||||
// EffectiveMode resolves the mode for one operation: a per-operation override wins over
|
||||
// Config.Mode, and ModeDefault is replaced by the dialect default. The result is
|
||||
// ModeProcedure, ModeDirect or ModeAuto; ModeAuto only survives on Postgres, where the
|
||||
// caller must probe the procedure. dialect is the resolved dialect name.
|
||||
func (c Config) EffectiveMode(op Op, dialect string) (Mode, error) {
|
||||
m := c.Mode
|
||||
if o, ok := c.Overrides[op]; ok && o != ModeDefault {
|
||||
m = o
|
||||
}
|
||||
if !m.valid() {
|
||||
return "", fmt.Errorf("lookup: invalid mode %q for %s", m, op)
|
||||
}
|
||||
pg := dialect == DialectPostgres
|
||||
switch m {
|
||||
case ModeDefault:
|
||||
if pg {
|
||||
return ModeProcedure, nil
|
||||
}
|
||||
return ModeDirect, nil
|
||||
case ModeAuto:
|
||||
if pg {
|
||||
return ModeAuto, nil
|
||||
}
|
||||
return ModeDirect, nil
|
||||
case ModeProcedure:
|
||||
if !pg {
|
||||
return "", fmt.Errorf("lookup: procedure mode for %s requires the postgres dialect, got %q", op, dialect)
|
||||
}
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// AllOps lists every operation, so callers can resolve or validate modes up front.
|
||||
func AllOps() []Op {
|
||||
return []Op{
|
||||
OpLogin,
|
||||
OpRegister,
|
||||
OpLogout,
|
||||
OpSession,
|
||||
OpTouchSession,
|
||||
OpRefresh,
|
||||
OpLoginAPIKey,
|
||||
OpJWTLogin,
|
||||
OpJWTLogout,
|
||||
OpResetRequest,
|
||||
OpResetComplete,
|
||||
OpKeyCreate,
|
||||
OpKeyList,
|
||||
OpKeyDelete,
|
||||
OpKeyValidate,
|
||||
OpOAuthRegisterClient,
|
||||
OpOAuthGetClient,
|
||||
OpOAuthSaveCode,
|
||||
OpOAuthExchangeCode,
|
||||
OpOAuthIntrospect,
|
||||
OpOAuthRevoke,
|
||||
OpOAuthGetOrCreateUser,
|
||||
OpOAuthCreateSession,
|
||||
OpOAuthGetRefreshToken,
|
||||
OpOAuthUpdateRefreshToken,
|
||||
OpOAuthGetUser,
|
||||
OpPasskeyStore,
|
||||
OpPasskeyGet,
|
||||
OpPasskeyUpdateCounter,
|
||||
OpPasskeyList,
|
||||
OpPasskeyDelete,
|
||||
OpPasskeyRename,
|
||||
OpPasskeyByUsername,
|
||||
OpPasskeyLogin,
|
||||
OpTOTPEnable,
|
||||
OpTOTPDisable,
|
||||
OpTOTPStatus,
|
||||
OpTOTPSecret,
|
||||
OpTOTPRegenerateBackup,
|
||||
OpTOTPValidateBackupCode,
|
||||
OpColumnSecurity,
|
||||
OpRowSecurity,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,318 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Auth implements lookup.AuthStore with stored procedures.
|
||||
type Auth struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.AuthStore = (*Auth)(nil)
|
||||
|
||||
// NewAuth creates the procedure-backed AuthStore.
|
||||
func NewAuth(run Runner, procs lookup.ProcNames) *Auth { return &Auth{run: run, procs: procs} }
|
||||
|
||||
// callData runs "SELECT p_success, p_error, p_data::text FROM proc($1::jsonb)".
|
||||
func (a *Auth) callData(ctx context.Context, proc, queryErrOp string, arg any) (sql.NullString, error) {
|
||||
var success bool
|
||||
var errorMsg, dataJSON sql.NullString
|
||||
err := a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, proc) //nolint:gosec // G201: identifier comes from validated config
|
||||
return db.QueryRowContext(ctx, query, arg).Scan(&success, &errorMsg, &dataJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return sql.NullString{}, fmt.Errorf("%s query failed: %w", queryErrOp, err)
|
||||
}
|
||||
if !success {
|
||||
return sql.NullString{}, failure(errorMsg, queryErrOp+" failed")
|
||||
}
|
||||
return dataJSON, nil
|
||||
}
|
||||
|
||||
// Login implements lookup.AuthStore.
|
||||
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal login request: %w", err)
|
||||
}
|
||||
data, err := a.callData(ctx, a.procs.Login, "login", string(reqJSON))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var response sectypes.LoginResponse
|
||||
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse login response: %w", err)
|
||||
}
|
||||
return &response, nil
|
||||
}
|
||||
|
||||
// Register implements lookup.AuthStore.
|
||||
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
|
||||
reqJSON, err := json.Marshal(req) //nolint:gosec // G117: intentional: field must be serialized
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal register request: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errorMsg, dataJSON sql.NullString
|
||||
err = a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.Register)
|
||||
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("register query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "registration failed")
|
||||
}
|
||||
var response sectypes.LoginResponse
|
||||
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse register response: %w", err)
|
||||
}
|
||||
return &response, nil
|
||||
}
|
||||
|
||||
// Logout implements lookup.AuthStore.
|
||||
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
reqJSON, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal logout request: %w", err)
|
||||
}
|
||||
_, err = a.callData(ctx, a.procs.Logout, "logout", string(reqJSON))
|
||||
return err
|
||||
}
|
||||
|
||||
// Session implements lookup.AuthStore.
|
||||
func (a *Auth) Session(ctx context.Context, token, reference string) (*sectypes.UserContext, error) {
|
||||
var success bool
|
||||
var errorMsg, userJSON sql.NullString
|
||||
err := a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
||||
return db.QueryRowContext(ctx, query, token, reference).Scan(&success, &errorMsg, &userJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("session query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "invalid or expired session")
|
||||
}
|
||||
if !userJSON.Valid {
|
||||
return nil, fmt.Errorf("no user data in session")
|
||||
}
|
||||
var user sectypes.UserContext
|
||||
if err := json.Unmarshal([]byte(userJSON.String), &user); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// TouchSession implements lookup.AuthStore.
|
||||
func (a *Auth) TouchSession(ctx context.Context, token string, user *sectypes.UserContext) error {
|
||||
userJSON, err := json.Marshal(user)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var success bool
|
||||
var errorMsg, updatedUserJSON sql.NullString
|
||||
return a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.SessionUpdate)
|
||||
return db.QueryRowContext(ctx, query, token, string(userJSON)).Scan(&success, &errorMsg, &updatedUserJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
|
||||
})
|
||||
}
|
||||
|
||||
// Refresh implements lookup.AuthStore.
|
||||
func (a *Auth) Refresh(ctx context.Context, refreshToken string) (*sectypes.LoginResponse, error) {
|
||||
// Get the current session to pass to refresh.
|
||||
var success bool
|
||||
var errorMsg, userJSON sql.NullString
|
||||
err := a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.procs.Session)
|
||||
return db.QueryRowContext(ctx, query, refreshToken, "refresh").Scan(&success, &errorMsg, &userJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("refresh token query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "invalid refresh token")
|
||||
}
|
||||
|
||||
var newSuccess bool
|
||||
var newErrorMsg, newUserJSON sql.NullString
|
||||
err = a.run.Run(func(db *sql.DB) error {
|
||||
refreshQuery := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2::jsonb)`, a.procs.RefreshToken)
|
||||
return db.QueryRowContext(ctx, refreshQuery, refreshToken, userJSON).Scan(&newSuccess, &newErrorMsg, &newUserJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("refresh token generation failed: %w", err)
|
||||
}
|
||||
if !newSuccess {
|
||||
return nil, failure(newErrorMsg, "failed to refresh token")
|
||||
}
|
||||
|
||||
var userCtx sectypes.UserContext
|
||||
if err := json.Unmarshal([]byte(newUserJSON.String), &userCtx); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||
}
|
||||
|
||||
// A resolvespec_refresh_token implementation that issues its own rotating
|
||||
// refresh token (independent of the access/session token) returns it
|
||||
// under claims.refresh_token, since UserContext has no dedicated field
|
||||
// for it. Surface that into LoginResponse.RefreshToken so callers don't
|
||||
// need to reach into User.Claims themselves. claims.expires_in
|
||||
// (seconds) similarly overrides the default access-token ExpiresIn when
|
||||
// the procedure provides a real value. Implementations that don't set
|
||||
// these claims keep today's behavior unchanged (empty RefreshToken,
|
||||
// 24h ExpiresIn default).
|
||||
resp := §ypes.LoginResponse{
|
||||
Token: userCtx.SessionID, // New session token from stored procedure
|
||||
User: &userCtx,
|
||||
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||
}
|
||||
if rt, ok := userCtx.Claims["refresh_token"].(string); ok && rt != "" {
|
||||
resp.RefreshToken = rt
|
||||
}
|
||||
if expiresIn, ok := userCtx.Claims["expires_in"].(float64); ok && expiresIn > 0 {
|
||||
resp.ExpiresIn = int64(expiresIn)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// LoginAPIKey implements lookup.AuthStore. Unknown, expired and inactive keys all return
|
||||
// lookup.ErrInvalidAPIKey; the raw key is never logged.
|
||||
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
if rawKey == "" {
|
||||
return nil, lookup.ErrInvalidAPIKey
|
||||
}
|
||||
reqJSON, err := json.Marshal(map[string]any{"api_key": rawKey, "claims": claims})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal api key login request: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errorMsg, dataJSON sql.NullString
|
||||
err = a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, a.procs.LoginAPIKey)
|
||||
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("api key login query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, lookup.ErrInvalidAPIKey
|
||||
}
|
||||
var response sectypes.LoginResponse
|
||||
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse api key login response: %w", err)
|
||||
}
|
||||
return &response, nil
|
||||
}
|
||||
|
||||
// JWTLogin implements lookup.AuthStore. The password is verified inside the procedure;
|
||||
// the hash is never returned. The token is a placeholder until JWT signing is wired in.
|
||||
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
var userJSON []byte
|
||||
err := a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_user FROM %s($1, $2)`, a.procs.JWTLogin)
|
||||
return db.QueryRowContext(ctx, query, req.Username, req.Password).Scan(&success, &errorMsg, &userJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "invalid credentials")
|
||||
}
|
||||
var user struct {
|
||||
ID int `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Email string `json:"email"`
|
||||
UserLevel int `json:"user_level"`
|
||||
Roles string `json:"roles"`
|
||||
}
|
||||
if err := json.Unmarshal(userJSON, &user); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user data: %w", err)
|
||||
}
|
||||
roles := []string{}
|
||||
if user.Roles != "" {
|
||||
roles = strings.Split(user.Roles, ",")
|
||||
}
|
||||
expiresAt := time.Now().Add(24 * time.Hour)
|
||||
return §ypes.LoginResponse{
|
||||
Token: fmt.Sprintf("token_%d_%d", user.ID, expiresAt.Unix()),
|
||||
User: §ypes.UserContext{
|
||||
UserID: user.ID,
|
||||
UserName: user.Username,
|
||||
Email: user.Email,
|
||||
UserLevel: user.UserLevel,
|
||||
Roles: roles,
|
||||
},
|
||||
ExpiresIn: int64(24 * time.Hour.Seconds()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// JWTLogout implements lookup.AuthStore.
|
||||
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
err := a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, a.procs.JWTLogout)
|
||||
return db.QueryRowContext(ctx, query, req.Token, req.UserID).Scan(&success, &errorMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("logout query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errorMsg, "logout failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResetRequest implements lookup.AuthStore.
|
||||
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
|
||||
reqJSON, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal password reset request: %w", err)
|
||||
}
|
||||
data, err := a.callData(ctx, a.procs.PasswordResetRequest, "password reset request", string(reqJSON))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var response sectypes.PasswordResetResponse
|
||||
if data.Valid && data.String != "" {
|
||||
if err := json.Unmarshal([]byte(data.String), &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse password reset response: %w", err)
|
||||
}
|
||||
}
|
||||
return &response, nil
|
||||
}
|
||||
|
||||
// ResetComplete implements lookup.AuthStore.
|
||||
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
|
||||
reqJSON, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal password reset complete request: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
err = a.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1::jsonb)`, a.procs.PasswordResetComplete)
|
||||
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("password reset complete query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errorMsg, "password reset failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// normalizeTimes rewrites the zone-less timestamps the procedures emit (Postgres `timestamp`
|
||||
// columns serialise as "2026-01-02T03:04:05.123456") as UTC RFC 3339, so the standard
|
||||
// time.Time decoder accepts them. It handles a JSON object or an array of objects and only
|
||||
// touches string fields whose name ends in "_at" or is "expiry". Anything else, including
|
||||
// input that is not valid JSON, is returned unchanged.
|
||||
func normalizeTimes(raw []byte) []byte {
|
||||
var v any
|
||||
if err := json.Unmarshal(raw, &v); err != nil {
|
||||
return raw
|
||||
}
|
||||
switch x := v.(type) {
|
||||
case map[string]any:
|
||||
fixTimes(x)
|
||||
case []any:
|
||||
for _, e := range x {
|
||||
if m, ok := e.(map[string]any); ok {
|
||||
fixTimes(m)
|
||||
}
|
||||
}
|
||||
default:
|
||||
return raw
|
||||
}
|
||||
out, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return raw
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func fixTimes(m map[string]any) {
|
||||
for k, v := range m {
|
||||
s, ok := v.(string)
|
||||
if !ok || !(strings.HasSuffix(k, "_at") || k == "expiry") {
|
||||
continue
|
||||
}
|
||||
if _, err := time.Parse(time.RFC3339Nano, s); err == nil {
|
||||
continue
|
||||
}
|
||||
for _, layout := range []string{"2006-01-02T15:04:05.999999999", "2006-01-02 15:04:05.999999999"} {
|
||||
if t, err := time.Parse(layout, s); err == nil {
|
||||
m[k] = t.UTC().Format(time.RFC3339Nano)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Keys implements lookup.KeyStore with the resolvespec_keystore_* procedures.
|
||||
type Keys struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.KeyStore = (*Keys)(nil)
|
||||
|
||||
// NewKeys creates the procedure-backed KeyStore.
|
||||
func NewKeys(run Runner, procs lookup.ProcNames) *Keys { return &Keys{run: run, procs: procs} }
|
||||
|
||||
// orDefault returns the procedure's error message when it is non-empty, otherwise def.
|
||||
func orDefault(s sql.NullString, def string) string {
|
||||
if s.Valid && s.String != "" {
|
||||
return s.String
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
// Create implements lookup.KeyStore.
|
||||
func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
|
||||
type createRequest struct {
|
||||
UserID int `json:"user_id"`
|
||||
KeyType sectypes.KeyType `json:"key_type"`
|
||||
KeyHash string `json:"key_hash"`
|
||||
Name string `json:"name"`
|
||||
Scopes []string `json:"scopes,omitempty"`
|
||||
Meta map[string]any `json:"meta,omitempty"`
|
||||
ExpiresAt *time.Time `json:"expires_at,omitempty"`
|
||||
}
|
||||
reqJSON, err := json.Marshal(createRequest{
|
||||
UserID: req.UserID,
|
||||
KeyType: req.KeyType,
|
||||
KeyHash: keyHash,
|
||||
Name: req.Name,
|
||||
Scopes: req.Scopes,
|
||||
Meta: req.Meta,
|
||||
ExpiresAt: req.ExpiresAt,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal create key request: %w", err)
|
||||
}
|
||||
|
||||
var success bool
|
||||
var errorMsg, keyJSON sql.NullString
|
||||
err = k.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1::jsonb)`, k.procs.KeystoreCreateKey)
|
||||
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &keyJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create key procedure failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, errors.New(orDefault(errorMsg, "create key failed"))
|
||||
}
|
||||
key, err := decodeKey([]byte(keyJSON.String))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse created key: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// List implements lookup.KeyStore.
|
||||
func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
|
||||
var success bool
|
||||
var errorMsg, keysJSON sql.NullString
|
||||
err := k.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_keys::text FROM %s($1, $2)`, k.procs.KeystoreGetUserKeys)
|
||||
return db.QueryRowContext(ctx, query, userID, string(keyType)).Scan(&success, &errorMsg, &keysJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get user keys procedure failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, errors.New(orDefault(errorMsg, "get user keys failed"))
|
||||
}
|
||||
var keys []sectypes.UserKey
|
||||
if keysJSON.Valid && keysJSON.String != "" && keysJSON.String != "[]" {
|
||||
var raw []json.RawMessage
|
||||
if err := json.Unmarshal([]byte(keysJSON.String), &raw); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user keys: %w", err)
|
||||
}
|
||||
for _, r := range raw {
|
||||
k, err := decodeKey(r)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user keys: %w", err)
|
||||
}
|
||||
keys = append(keys, *k)
|
||||
}
|
||||
}
|
||||
if keys == nil {
|
||||
keys = []sectypes.UserKey{}
|
||||
}
|
||||
return keys, nil
|
||||
}
|
||||
|
||||
// Delete implements lookup.KeyStore. The procedure returns the key hash.
|
||||
func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
|
||||
var success bool
|
||||
var errorMsg, keyHash sql.NullString
|
||||
err := k.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_key_hash FROM %s($1, $2)`, k.procs.KeystoreDeleteKey)
|
||||
return db.QueryRowContext(ctx, query, userID, keyID).Scan(&success, &errorMsg, &keyHash)
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("delete key procedure failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return "", errors.New(orDefault(errorMsg, "delete key failed"))
|
||||
}
|
||||
return keyHash.String, nil
|
||||
}
|
||||
|
||||
// Validate implements lookup.KeyStore.
|
||||
func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
|
||||
var success bool
|
||||
var errorMsg, keyJSON sql.NullString
|
||||
err := k.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_key::text FROM %s($1, $2)`, k.procs.KeystoreValidateKey)
|
||||
return db.QueryRowContext(ctx, query, keyHash, string(keyType)).Scan(&success, &errorMsg, &keyJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("validate key procedure failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, errors.New(orDefault(errorMsg, "invalid or expired key"))
|
||||
}
|
||||
key, err := decodeKey([]byte(keyJSON.String))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse validated key: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// decodeKey reads one key record from a key procedure.
|
||||
func decodeKey(raw []byte) (*sectypes.UserKey, error) {
|
||||
var k sectypes.UserKey
|
||||
if err := json.Unmarshal(normalizeTimes(raw), &k); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &k, nil
|
||||
}
|
||||
@@ -0,0 +1,315 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// OAuthUsers implements lookup.OAuthUserStore with the resolvespec_oauth_* procedures.
|
||||
type OAuthUsers struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.OAuthUserStore = (*OAuthUsers)(nil)
|
||||
|
||||
// NewOAuthUsers creates the procedure-backed OAuthUserStore.
|
||||
func NewOAuthUsers(run Runner, procs lookup.ProcNames) *OAuthUsers {
|
||||
return &OAuthUsers{run: run, procs: procs}
|
||||
}
|
||||
|
||||
// GetOrCreateUser implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
|
||||
userJSON, err := json.Marshal(map[string]any{
|
||||
"username": user.UserName,
|
||||
"email": user.Email,
|
||||
"remote_id": user.RemoteID,
|
||||
"user_level": user.UserLevel,
|
||||
"roles": user.Roles,
|
||||
"auth_provider": provider,
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to marshal user data: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
var userID sql.NullInt64
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_user_id
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthGetOrCreateUser), userJSON).Scan(&success, &errMsg, &userID)
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to get or create user: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return 0, failure(errMsg, "failed to get or create user")
|
||||
}
|
||||
if !userID.Valid {
|
||||
return 0, fmt.Errorf("user ID not returned")
|
||||
}
|
||||
return int(userID.Int64), nil
|
||||
}
|
||||
|
||||
// CreateSession implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error {
|
||||
sessionJSON, err := json.Marshal(map[string]any{
|
||||
"session_token": s.SessionToken,
|
||||
"user_id": s.UserID,
|
||||
"access_token": s.AccessToken,
|
||||
"refresh_token": s.RefreshToken,
|
||||
"token_type": s.TokenType,
|
||||
"expires_at": s.ExpiresAt,
|
||||
"auth_provider": s.Provider,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal session data: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthCreateSession), sessionJSON).Scan(&success, &errMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create session: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errMsg, "failed to create session")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetByRefreshToken implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
var data []byte
|
||||
err := o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_data::text
|
||||
FROM %s($1)
|
||||
`, o.procs.OAuthGetRefreshToken), refreshToken).Scan(&success, &errMsg, &data)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errMsg, "invalid or expired refresh token")
|
||||
}
|
||||
var session lookup.OAuthRefreshSession
|
||||
if err := json.Unmarshal(normalizeTimes(data), &session); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse session data: %w", err)
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
// UpdateRefreshToken implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
|
||||
updateJSON, err := json.Marshal(map[string]any{
|
||||
"user_id": userID,
|
||||
"old_refresh_token": oldRefreshToken,
|
||||
"new_session_token": newSessionToken,
|
||||
"new_access_token": newAccessToken,
|
||||
"new_refresh_token": newRefreshToken,
|
||||
"expires_at": expiresAt,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal update data: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthUpdateRefreshToken), updateJSON).Scan(&success, &errMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update session: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errMsg, "failed to update session")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUser implements lookup.OAuthUserStore.
|
||||
func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
var data []byte
|
||||
err := o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_data::text
|
||||
FROM %s($1)
|
||||
`, o.procs.OAuthGetUser), userID).Scan(&success, &errMsg, &data)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get user data: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errMsg, "failed to get user data")
|
||||
}
|
||||
var userCtx sectypes.UserContext
|
||||
if err := json.Unmarshal(data, &userCtx); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse user context: %w", err)
|
||||
}
|
||||
return &userCtx, nil
|
||||
}
|
||||
|
||||
// OAuthClients implements lookup.OAuthClientStore with the resolvespec_oauth_* server procedures.
|
||||
type OAuthClients struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.OAuthClientStore = (*OAuthClients)(nil)
|
||||
|
||||
// NewOAuthClients creates the procedure-backed OAuthClientStore.
|
||||
func NewOAuthClients(run Runner, procs lookup.ProcNames) *OAuthClients {
|
||||
return &OAuthClients{run: run, procs: procs}
|
||||
}
|
||||
|
||||
// callData runs a `(p_success, p_error, p_data)` procedure with one argument.
|
||||
func (o *OAuthClients) callData(ctx context.Context, proc string, arg any) (data []byte, ok bool, errMsg sql.NullString, err error) {
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_data::text
|
||||
FROM %s($1)
|
||||
`, proc), arg).Scan(&ok, &errMsg, &data)
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// callNoData runs a `(p_success, p_error)` procedure with one argument.
|
||||
func (o *OAuthClients) callNoData(ctx context.Context, proc string, arg any) (ok bool, errMsg sql.NullString, err error) {
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error
|
||||
FROM %s($1)
|
||||
`, proc), arg).Scan(&ok, &errMsg)
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// RegisterClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
|
||||
input, err := json.Marshal(client)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal client: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
var data []byte
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error, p_data::text
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthRegisterClient), input).Scan(&success, &errMsg, &data)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to register client: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errMsg, "failed to register client")
|
||||
}
|
||||
var result sectypes.OAuthServerClient
|
||||
if err := json.Unmarshal(data, &result); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse registered client: %w", err)
|
||||
}
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// GetClient implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
|
||||
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthGetClient, clientID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get client: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return nil, failure(errMsg, "client not found")
|
||||
}
|
||||
var result sectypes.OAuthServerClient
|
||||
if err := json.Unmarshal(data, &result); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse client: %w", err)
|
||||
}
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// SaveCode implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
|
||||
input, err := json.Marshal(code) //nolint:gosec // G117: intentional: field must be serialized
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal code: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errMsg sql.NullString
|
||||
err = o.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT p_success, p_error
|
||||
FROM %s($1::jsonb)
|
||||
`, o.procs.OAuthSaveCode), input).Scan(&success, &errMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to save code: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errMsg, "failed to save code")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExchangeCode implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
|
||||
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthExchangeCode, code)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to exchange code: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return nil, failure(errMsg, "invalid or expired code")
|
||||
}
|
||||
var result sectypes.OAuthCode
|
||||
if err := json.Unmarshal(normalizeTimes(data), &result); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse code data: %w", err)
|
||||
}
|
||||
result.Code = code
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// Introspect implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
|
||||
data, ok, errMsg, err := o.callData(ctx, o.procs.OAuthIntrospect, token)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to introspect token: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return nil, failure(errMsg, "introspection failed")
|
||||
}
|
||||
var result sectypes.OAuthTokenInfo
|
||||
if err := json.Unmarshal(data, &result); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse token info: %w", err)
|
||||
}
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// Revoke implements lookup.OAuthClientStore.
|
||||
func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
|
||||
ok, errMsg, err := o.callNoData(ctx, o.procs.OAuthRevoke, token)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to revoke token: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return failure(errMsg, "failed to revoke token")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,282 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Passkey implements lookup.PasskeyStore with the resolvespec_passkey_* procedures.
|
||||
// Credential ids cross the lookup interface as base64 text; the procedures that take a
|
||||
// bytea credential id receive the decoded bytes.
|
||||
type Passkey struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.PasskeyStore = (*Passkey)(nil)
|
||||
|
||||
// NewPasskey creates the procedure-backed PasskeyStore.
|
||||
func NewPasskey(run Runner, procs lookup.ProcNames) *Passkey {
|
||||
return &Passkey{run: run, procs: procs}
|
||||
}
|
||||
|
||||
func decodeCredentialID(b64 string) ([]byte, error) {
|
||||
id, err := base64.StdEncoding.DecodeString(b64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid credential ID: %w", err)
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// Store implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
|
||||
credJSON, err := json.Marshal(map[string]any{
|
||||
"user_id": rec.UserID,
|
||||
"credential_id": rec.CredentialID,
|
||||
"public_key": rec.PublicKey,
|
||||
"attestation_type": rec.AttestationType,
|
||||
"sign_count": rec.SignCount,
|
||||
"transports": rec.Transports,
|
||||
"backup_eligible": rec.BackupEligible,
|
||||
"backup_state": rec.BackupState,
|
||||
"name": rec.Name,
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to marshal credential data: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
var credentialID sql.NullInt64
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential_id FROM %s($1::jsonb)`, p.procs.PasskeyStoreCredential)
|
||||
return db.QueryRowContext(ctx, query, string(credJSON)).Scan(&success, &errorMsg, &credentialID)
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to store credential: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return 0, failure(errorMsg, "failed to store credential")
|
||||
}
|
||||
return credentialID.Int64, nil
|
||||
}
|
||||
|
||||
// Get implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Get(ctx context.Context, credentialID string) (int, uint32, error) {
|
||||
raw, err := decodeCredentialID(credentialID)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
var success bool
|
||||
var errorMsg, credentialJSON sql.NullString
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_credential::text FROM %s($1)`, p.procs.PasskeyGetCredential)
|
||||
return db.QueryRowContext(ctx, query, raw).Scan(&success, &errorMsg, &credentialJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return 0, 0, failure(errorMsg, "credential not found")
|
||||
}
|
||||
var cred struct {
|
||||
UserID int `json:"user_id"`
|
||||
SignCount uint32 `json:"sign_count"`
|
||||
}
|
||||
if err := json.Unmarshal(normalizeTimes([]byte(credentialJSON.String)), &cred); err != nil {
|
||||
return 0, 0, fmt.Errorf("failed to parse credential: %w", err)
|
||||
}
|
||||
return cred.UserID, cred.SignCount, nil
|
||||
}
|
||||
|
||||
// UpdateCounter implements lookup.PasskeyStore. Like the code it replaces, it only reports
|
||||
// an error when the query itself fails; the procedure's success flag is not checked.
|
||||
func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
|
||||
raw, err := decodeCredentialID(credentialID)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
var cloneWarning sql.NullBool
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_clone_warning FROM %s($1, $2)`, p.procs.PasskeyUpdateCounter)
|
||||
return db.QueryRowContext(ctx, query, raw, newCounter).Scan(&success, &errorMsg, &cloneWarning)
|
||||
})
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return cloneWarning.Valid && cloneWarning.Bool, nil
|
||||
}
|
||||
|
||||
// List implements lookup.PasskeyStore.
|
||||
func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
|
||||
var success bool
|
||||
var errorMsg, credentialsJSON sql.NullString
|
||||
err := p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetUserCredentials)
|
||||
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &credentialsJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "failed to get credentials")
|
||||
}
|
||||
|
||||
var rawCreds []struct {
|
||||
ID int `json:"id"`
|
||||
UserID int `json:"user_id"`
|
||||
CredentialID string `json:"credential_id"`
|
||||
PublicKey string `json:"public_key"`
|
||||
AttestationType string `json:"attestation_type"`
|
||||
AAGUID string `json:"aaguid"`
|
||||
SignCount uint32 `json:"sign_count"`
|
||||
CloneWarning bool `json:"clone_warning"`
|
||||
Transports []string `json:"transports"`
|
||||
BackupEligible bool `json:"backup_eligible"`
|
||||
BackupState bool `json:"backup_state"`
|
||||
Name string `json:"name"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
LastUsedAt time.Time `json:"last_used_at"`
|
||||
}
|
||||
if err := json.Unmarshal(normalizeTimes([]byte(credentialsJSON.String)), &rawCreds); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||
}
|
||||
|
||||
credentials := make([]sectypes.PasskeyCredential, 0, len(rawCreds))
|
||||
for i := range rawCreds {
|
||||
raw := rawCreds[i]
|
||||
credID, err := base64.StdEncoding.DecodeString(raw.CredentialID)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
pubKey, err := base64.StdEncoding.DecodeString(raw.PublicKey)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
aaguid, _ := base64.StdEncoding.DecodeString(raw.AAGUID)
|
||||
credentials = append(credentials, sectypes.PasskeyCredential{
|
||||
ID: fmt.Sprintf("%d", raw.ID),
|
||||
UserID: raw.UserID,
|
||||
CredentialID: credID,
|
||||
PublicKey: pubKey,
|
||||
AttestationType: raw.AttestationType,
|
||||
AAGUID: aaguid,
|
||||
SignCount: raw.SignCount,
|
||||
CloneWarning: raw.CloneWarning,
|
||||
Transports: raw.Transports,
|
||||
BackupEligible: raw.BackupEligible,
|
||||
BackupState: raw.BackupState,
|
||||
Name: raw.Name,
|
||||
CreatedAt: raw.CreatedAt,
|
||||
LastUsedAt: raw.LastUsedAt,
|
||||
})
|
||||
}
|
||||
return credentials, nil
|
||||
}
|
||||
|
||||
// Delete implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error {
|
||||
raw, err := decodeCredentialID(credentialID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2)`, p.procs.PasskeyDeleteCredential)
|
||||
return db.QueryRowContext(ctx, query, userID, raw).Scan(&success, &errorMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete credential: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errorMsg, "failed to delete credential")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Rename implements lookup.PasskeyStore.
|
||||
func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error {
|
||||
raw, err := decodeCredentialID(credentialID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3)`, p.procs.PasskeyUpdateName)
|
||||
return db.QueryRowContext(ctx, query, userID, raw, name).Scan(&success, &errorMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update credential name: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errorMsg, "failed to update credential name")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ByUsername implements lookup.PasskeyStore.
|
||||
func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
|
||||
var success bool
|
||||
var errorMsg, credentialsJSON sql.NullString
|
||||
var userID sql.NullInt64
|
||||
err := p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_user_id, p_credentials::text FROM %s($1)`, p.procs.PasskeyGetCredsByUsername)
|
||||
return db.QueryRowContext(ctx, query, username).Scan(&success, &errorMsg, &userID, &credentialsJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return 0, nil, failure(errorMsg, "failed to get credentials")
|
||||
}
|
||||
var creds []lookup.PasskeyCredentialRef
|
||||
if err := json.Unmarshal([]byte(credentialsJSON.String), &creds); err != nil {
|
||||
return 0, nil, fmt.Errorf("failed to parse credentials: %w", err)
|
||||
}
|
||||
return int(userID.Int64), creds, nil
|
||||
}
|
||||
|
||||
// Login implements lookup.PasskeyStore: it creates the session for a user whose passkey
|
||||
// assertion was already verified.
|
||||
func (p *Passkey) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
|
||||
reqData := map[string]any{"user_id": userID}
|
||||
if claims != nil {
|
||||
if ip, ok := claims["ip_address"].(string); ok {
|
||||
reqData["ip_address"] = ip
|
||||
}
|
||||
if ua, ok := claims["user_agent"].(string); ok {
|
||||
reqData["user_agent"] = ua
|
||||
}
|
||||
}
|
||||
reqJSON, err := json.Marshal(reqData)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal passkey login request: %w", err)
|
||||
}
|
||||
var success bool
|
||||
var errorMsg, dataJSON sql.NullString
|
||||
err = p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_data::text FROM %s($1::jsonb)`, p.procs.PasskeyLogin)
|
||||
return db.QueryRowContext(ctx, query, string(reqJSON)).Scan(&success, &errorMsg, &dataJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("passkey login query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "passkey login failed")
|
||||
}
|
||||
var response sectypes.LoginResponse
|
||||
if err := json.Unmarshal([]byte(dataJSON.String), &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse passkey login response: %w", err)
|
||||
}
|
||||
return &response, nil
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// Policy implements lookup.PolicyStore with the column and row security procedures.
|
||||
type Policy struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.PolicyStore = (*Policy)(nil)
|
||||
|
||||
// NewPolicy creates the procedure-backed PolicyStore.
|
||||
func NewPolicy(run Runner, procs lookup.ProcNames) *Policy { return &Policy{run: run, procs: procs} }
|
||||
|
||||
// ColumnSecurity implements lookup.PolicyStore.
|
||||
func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
var rulesJSON []byte
|
||||
err := p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_rules FROM %s($1, $2, $3)`, p.procs.ColumnSecurity)
|
||||
return db.QueryRowContext(ctx, query, userID, schema, table).Scan(&success, &errorMsg, &rulesJSON)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load column security: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return nil, failure(errorMsg, "failed to load column security")
|
||||
}
|
||||
|
||||
type securityRecord struct {
|
||||
Control string `json:"control"`
|
||||
Accesstype string `json:"accesstype"`
|
||||
JSONValue string `json:"jsonvalue"`
|
||||
}
|
||||
var records []securityRecord
|
||||
if err := json.Unmarshal(rulesJSON, &records); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse security rules: %w", err)
|
||||
}
|
||||
|
||||
var rules []sectypes.ColumnSecurity
|
||||
for _, rec := range records {
|
||||
parts := strings.Split(rec.Control, ".")
|
||||
if len(parts) < 3 {
|
||||
continue
|
||||
}
|
||||
rules = append(rules, sectypes.ColumnSecurity{
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
Path: parts[2:],
|
||||
Accesstype: rec.Accesstype,
|
||||
UserID: userID,
|
||||
})
|
||||
}
|
||||
return rules, nil
|
||||
}
|
||||
|
||||
// RowSecurity implements lookup.PolicyStore. userRef is unwrapped to a scalar user id
|
||||
// because the procedure's p_user_id is an integer.
|
||||
func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
|
||||
switch v := userRef.(type) {
|
||||
case *sectypes.UserContext:
|
||||
if v != nil {
|
||||
userRef = v.UserID
|
||||
}
|
||||
case sectypes.UserContext:
|
||||
userRef = v.UserID
|
||||
}
|
||||
|
||||
var template sql.NullString
|
||||
var hasBlock sql.NullBool
|
||||
err := p.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_template, p_block FROM %s($1, $2, $3)`, p.procs.RowSecurity)
|
||||
return db.QueryRowContext(ctx, query, schema, table, userRef).Scan(&template, &hasBlock)
|
||||
})
|
||||
if err != nil {
|
||||
return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
|
||||
}
|
||||
return sectypes.RowSecurity{
|
||||
Schema: schema,
|
||||
Tablename: table,
|
||||
UserID: userRef,
|
||||
Template: template.String,
|
||||
HasBlock: hasBlock.Bool,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
// Package procedure is the stored-procedure backend of lookup: it calls the
|
||||
// resolvespec_* functions (names from lookup.ProcNames) and keeps their
|
||||
// p_success / p_error / p_data contract. Error texts match the ones the security
|
||||
// package returned before the extraction.
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
|
||||
type Runner interface {
|
||||
Run(run func(*sql.DB) error) error
|
||||
}
|
||||
|
||||
// RunFunc adapts a function to Runner. The security package passes its own
|
||||
// reconnecting helper this way.
|
||||
type RunFunc func(run func(*sql.DB) error) error
|
||||
|
||||
// Run implements Runner.
|
||||
func (f RunFunc) Run(run func(*sql.DB) error) error { return f(run) }
|
||||
|
||||
// DB is a standalone Runner over a *sql.DB with an optional reconnect factory.
|
||||
type DB struct {
|
||||
mu sync.RWMutex
|
||||
db *sql.DB
|
||||
factory func() (*sql.DB, error)
|
||||
onReconnect func()
|
||||
}
|
||||
|
||||
// NewDB wraps db. factory (optional) is called to obtain a fresh handle when the current
|
||||
// one is closed; onReconnect (optional) runs after a successful reconnect, e.g. to reset
|
||||
// cached procedure probes.
|
||||
func NewDB(db *sql.DB, factory func() (*sql.DB, error), onReconnect func()) *DB {
|
||||
return &DB{db: db, factory: factory, onReconnect: onReconnect}
|
||||
}
|
||||
|
||||
// Get returns the current handle.
|
||||
func (d *DB) Get() *sql.DB {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
return d.db
|
||||
}
|
||||
|
||||
func (d *DB) reconnect() error {
|
||||
if d.factory == nil {
|
||||
return fmt.Errorf("no db factory configured for reconnect")
|
||||
}
|
||||
newDB, err := d.factory()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
d.mu.Lock()
|
||||
d.db = newDB
|
||||
d.mu.Unlock()
|
||||
if d.onReconnect != nil {
|
||||
d.onReconnect()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Run implements Runner.
|
||||
func (d *DB) Run(run func(*sql.DB) error) error {
|
||||
db := d.Get()
|
||||
if db == nil {
|
||||
return fmt.Errorf("database connection is nil")
|
||||
}
|
||||
err := run(db)
|
||||
if IsClosed(err) {
|
||||
if reconnErr := d.reconnect(); reconnErr == nil {
|
||||
err = run(d.Get())
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// IsClosed reports whether err indicates the *sql.DB has been closed.
|
||||
func IsClosed(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "sql: database is closed")
|
||||
}
|
||||
|
||||
// failure builds the error for p_success = false: the procedure's own message when it
|
||||
// returned one, otherwise def.
|
||||
func failure(errMsg sql.NullString, def string) error {
|
||||
if errMsg.Valid {
|
||||
return fmt.Errorf("%s", errMsg.String)
|
||||
}
|
||||
return fmt.Errorf("%s", def)
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"regexp"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
func newMock(t *testing.T) (*DB, sqlmock.Sqlmock) {
|
||||
t.Helper()
|
||||
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return NewDB(db, nil, nil), mock
|
||||
}
|
||||
|
||||
func q(s string) string { return regexp.QuoteMeta(s) }
|
||||
|
||||
func TestAuthLoginCallsProcedureWithJSON(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
procs := lookup.DefaultProcNames()
|
||||
procs.Login = "custom_login"
|
||||
a := NewAuth(run, procs)
|
||||
|
||||
mock.ExpectQuery(q("SELECT p_success, p_error, p_data::text FROM custom_login($1::jsonb)")).
|
||||
WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
|
||||
AddRow(true, nil, `{"token":"sess_1","user":{"user_id":7,"user_name":"bob"}}`))
|
||||
|
||||
resp, err := a.Login(context.Background(), sectypes.LoginRequest{Username: "bob", Password: "x"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if resp.Token != "sess_1" || resp.User == nil || resp.User.UserID != 7 {
|
||||
t.Fatalf("unexpected response: %+v", resp)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLoginFailureUsesProcedureMessage(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
a := NewAuth(run, lookup.DefaultProcNames())
|
||||
|
||||
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, "bad credentials", nil))
|
||||
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() != "bad credentials" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
|
||||
mock.ExpectQuery("resolvespec_login").WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
|
||||
if _, err := a.Login(context.Background(), sectypes.LoginRequest{}); err == nil || err.Error() == "" {
|
||||
t.Fatalf("expected default error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerReconnectsOnClosedDB(t *testing.T) {
|
||||
first, _, err := sqlmock.New()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = first.Close()
|
||||
|
||||
second, mock, err := sqlmock.New()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer second.Close()
|
||||
|
||||
reconnected := false
|
||||
run := NewDB(first, func() (*sql.DB, error) { return second, nil }, func() { reconnected = true })
|
||||
mock.ExpectQuery("SELECT 1").WillReturnRows(sqlmock.NewRows([]string{"x"}).AddRow(1))
|
||||
|
||||
err = run.Run(func(db *sql.DB) error {
|
||||
var x int
|
||||
return db.QueryRow("SELECT 1").Scan(&x)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !reconnected || run.Get() != second {
|
||||
t.Fatal("expected reconnect to the new handle")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerNoFactoryReturnsClosedError(t *testing.T) {
|
||||
db, _, _ := sqlmock.New()
|
||||
_ = db.Close()
|
||||
err := NewDB(db, nil, nil).Run(func(db *sql.DB) error { return db.QueryRow("SELECT 1").Scan(new(int)) })
|
||||
if !IsClosed(err) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyGetDecodesCredentialID(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||
raw := []byte{1, 2, 3, 4}
|
||||
|
||||
mock.ExpectQuery("resolvespec_passkey_get_credential").WithArgs(raw).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_credential"}).
|
||||
AddRow(true, nil, `{"user_id":9,"sign_count":4}`))
|
||||
uid, count, err := p.Get(context.Background(), base64.StdEncoding.EncodeToString(raw))
|
||||
if err != nil || uid != 9 || count != 4 {
|
||||
t.Fatalf("got %d %d %v", uid, count, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyInvalidBase64(t *testing.T) {
|
||||
run, _ := newMock(t)
|
||||
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||
if _, _, err := p.Get(context.Background(), "***"); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if err := p.Delete(context.Background(), 1, "***"); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyUpdateCounterReportsCloneWarning(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||
id := base64.StdEncoding.EncodeToString([]byte("abc"))
|
||||
|
||||
mock.ExpectQuery("resolvespec_passkey_update_counter").WithArgs([]byte("abc"), uint32(5)).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_clone_warning"}).AddRow(true, nil, true))
|
||||
warn, err := p.UpdateCounter(context.Background(), id, 5)
|
||||
if err != nil || !warn {
|
||||
t.Fatalf("got %v %v", warn, err)
|
||||
}
|
||||
|
||||
mock.ExpectQuery("resolvespec_passkey_update_counter").WillReturnError(errors.New("boom"))
|
||||
if _, err := p.UpdateCounter(context.Background(), id, 6); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyByUsername(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
p := NewPasskey(run, lookup.DefaultProcNames())
|
||||
mock.ExpectQuery("resolvespec_passkey_get_credentials_by_username").WithArgs("bob").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id", "p_credentials"}).
|
||||
AddRow(true, nil, 3, `[{"credential_id":"YWJj","transports":["usb"]}]`))
|
||||
uid, refs, err := p.ByUsername(context.Background(), "bob")
|
||||
if err != nil || uid != 3 || len(refs) != 1 || refs[0].CredentialID != "YWJj" || refs[0].Transports[0] != "usb" {
|
||||
t.Fatalf("got %d %+v %v", uid, refs, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthUsersGetOrCreateUser(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
o := NewOAuthUsers(run, lookup.DefaultProcNames())
|
||||
|
||||
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, 11))
|
||||
id, err := o.GetOrCreateUser(context.Background(), §ypes.UserContext{UserName: "u", Email: "e"}, "github")
|
||||
if err != nil || id != 11 {
|
||||
t.Fatalf("got %d %v", id, err)
|
||||
}
|
||||
|
||||
mock.ExpectQuery("resolvespec_oauth_getorcreateuser").WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_user_id"}).AddRow(true, nil, nil))
|
||||
if _, err := o.GetOrCreateUser(context.Background(), §ypes.UserContext{}, "github"); err == nil || err.Error() != "user ID not returned" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthUsersRefreshRoundTrip(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
o := NewOAuthUsers(run, lookup.DefaultProcNames())
|
||||
|
||||
mock.ExpectQuery("resolvespec_oauth_getrefreshtoken").WithArgs("r1").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).
|
||||
AddRow(true, nil, `{"user_id":2,"access_token":"a","token_type":"Bearer","expiry":"2030-01-01T00:00:00Z"}`))
|
||||
s, err := o.GetByRefreshToken(context.Background(), "r1")
|
||||
if err != nil || s.UserID != 2 || s.AccessToken != "a" {
|
||||
t.Fatalf("got %+v %v", s, err)
|
||||
}
|
||||
|
||||
mock.ExpectQuery("resolvespec_oauth_updaterefreshtoken").WithArgs(sqlmock.AnyArg()).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(false, "session not found"))
|
||||
err = o.UpdateRefreshToken(context.Background(), 2, "r1", "s2", "a2", "r2", time.Now())
|
||||
if err == nil || err.Error() != "session not found" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthClientsExchangeCodeSetsCode(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
c := NewOAuthClients(run, lookup.DefaultProcNames())
|
||||
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("abc").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, `{"client_id":"cid"}`))
|
||||
code, err := c.ExchangeCode(context.Background(), "abc")
|
||||
if err != nil || code.Code != "abc" || code.ClientID != "cid" {
|
||||
t.Fatalf("got %+v %v", code, err)
|
||||
}
|
||||
|
||||
mock.ExpectQuery("resolvespec_oauth_exchange_code").WithArgs("zzz").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, nil, nil))
|
||||
if _, err := c.ExchangeCode(context.Background(), "zzz"); err == nil || err.Error() != "invalid or expired code" {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthClientsRevoke(t *testing.T) {
|
||||
run, mock := newMock(t)
|
||||
c := NewOAuthClients(run, lookup.DefaultProcNames())
|
||||
mock.ExpectQuery("resolvespec_oauth_revoke").WithArgs("t").
|
||||
WillReturnRows(sqlmock.NewRows([]string{"p_success", "p_error"}).AddRow(true, nil))
|
||||
if err := c.Revoke(context.Background(), "t"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package procedure
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
// TOTP implements lookup.TOTPStore with the resolvespec_totp_* procedures.
|
||||
type TOTP struct {
|
||||
run Runner
|
||||
procs lookup.ProcNames
|
||||
}
|
||||
|
||||
var _ lookup.TOTPStore = (*TOTP)(nil)
|
||||
|
||||
// NewTOTP creates the procedure-backed TOTPStore.
|
||||
func NewTOTP(run Runner, procs lookup.ProcNames) *TOTP { return &TOTP{run: run, procs: procs} }
|
||||
|
||||
// exec runs a "p_success, p_error" procedure and maps failure to an error.
|
||||
func (t *TOTP) exec(ctx context.Context, query, op, def string, args ...any) error {
|
||||
var success bool
|
||||
var errorMsg sql.NullString
|
||||
err := t.run.Run(func(db *sql.DB) error {
|
||||
return db.QueryRowContext(ctx, query, args...).Scan(&success, &errorMsg)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s query failed: %w", op, err)
|
||||
}
|
||||
if !success {
|
||||
return failure(errorMsg, def)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Enable implements lookup.TOTPStore.
|
||||
func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
|
||||
codesJSON, err := json.Marshal(hashedCodes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||
}
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2, $3::jsonb)`, t.procs.TOTPEnable)
|
||||
return t.exec(ctx, query, "enable 2FA", "failed to enable 2FA", userID, secret, string(codesJSON))
|
||||
}
|
||||
|
||||
// Disable implements lookup.TOTPStore.
|
||||
func (t *TOTP) Disable(ctx context.Context, userID int) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1)`, t.procs.TOTPDisable)
|
||||
return t.exec(ctx, query, "disable 2FA", "failed to disable 2FA", userID)
|
||||
}
|
||||
|
||||
// Status implements lookup.TOTPStore.
|
||||
func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) {
|
||||
var success, enabled bool
|
||||
var errorMsg sql.NullString
|
||||
err := t.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_enabled FROM %s($1)`, t.procs.TOTPGetStatus)
|
||||
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &enabled)
|
||||
})
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("get 2FA status query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return false, failure(errorMsg, "failed to get 2FA status")
|
||||
}
|
||||
return enabled, nil
|
||||
}
|
||||
|
||||
// Secret implements lookup.TOTPStore.
|
||||
func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) {
|
||||
var success bool
|
||||
var errorMsg, secret sql.NullString
|
||||
err := t.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_secret FROM %s($1)`, t.procs.TOTPGetSecret)
|
||||
return db.QueryRowContext(ctx, query, userID).Scan(&success, &errorMsg, &secret)
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
return "", failure(errorMsg, "failed to get 2FA secret")
|
||||
}
|
||||
if !secret.Valid {
|
||||
return "", fmt.Errorf("2FA secret not found")
|
||||
}
|
||||
return secret.String, nil
|
||||
}
|
||||
|
||||
// RegenerateBackupCodes implements lookup.TOTPStore.
|
||||
func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
|
||||
codesJSON, err := json.Marshal(hashedCodes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal backup codes: %w", err)
|
||||
}
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error FROM %s($1, $2::jsonb)`, t.procs.TOTPRegenerateBackup)
|
||||
return t.exec(ctx, query, "regenerate backup codes", "failed to regenerate backup codes", userID, string(codesJSON))
|
||||
}
|
||||
|
||||
// ValidateBackupCode implements lookup.TOTPStore. A failure without a message means
|
||||
// "not valid", not an error.
|
||||
func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
|
||||
var success, valid bool
|
||||
var errorMsg sql.NullString
|
||||
err := t.run.Run(func(db *sql.DB) error {
|
||||
query := fmt.Sprintf(`SELECT p_success, p_error, p_valid FROM %s($1, $2)`, t.procs.TOTPValidateBackupCode)
|
||||
return db.QueryRowContext(ctx, query, userID, codeHash).Scan(&success, &errorMsg, &valid)
|
||||
})
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("validate backup code query failed: %w", err)
|
||||
}
|
||||
if !success {
|
||||
if errorMsg.Valid {
|
||||
return false, fmt.Errorf("%s", errorMsg.String)
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
return valid, nil
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// ProcNames holds the stored procedure (function) names used by the procedure backend.
|
||||
// It replaces security.SQLNames and security.KeyStoreSQLNames. Zero fields mean "default"
|
||||
// when merged with DefaultProcNames.
|
||||
type ProcNames struct {
|
||||
|
||||
// Auth procedures (DatabaseAuthenticator)
|
||||
Login string // default: "resolvespec_login"
|
||||
Register string // default: "resolvespec_register"
|
||||
Logout string // default: "resolvespec_logout"
|
||||
Session string // default: "resolvespec_session"
|
||||
SessionUpdate string // default: "resolvespec_session_update"
|
||||
RefreshToken string // default: "resolvespec_refresh_token"
|
||||
LoginAPIKey string // default: "resolvespec_login_api_key"
|
||||
|
||||
// JWT procedures (JWTAuthenticator)
|
||||
JWTLogin string // default: "resolvespec_jwt_login"
|
||||
JWTLogout string // default: "resolvespec_jwt_logout"
|
||||
|
||||
// Security policy procedures
|
||||
ColumnSecurity string // default: "resolvespec_column_security"
|
||||
RowSecurity string // default: "resolvespec_row_security"
|
||||
|
||||
// TOTP procedures (DatabaseTwoFactorProvider)
|
||||
TOTPEnable string // default: "resolvespec_totp_enable"
|
||||
TOTPDisable string // default: "resolvespec_totp_disable"
|
||||
TOTPGetStatus string // default: "resolvespec_totp_get_status"
|
||||
TOTPGetSecret string // default: "resolvespec_totp_get_secret"
|
||||
TOTPRegenerateBackup string // default: "resolvespec_totp_regenerate_backup_codes"
|
||||
TOTPValidateBackupCode string // default: "resolvespec_totp_validate_backup_code"
|
||||
|
||||
// Passkey procedures (DatabasePasskeyProvider)
|
||||
PasskeyStoreCredential string // default: "resolvespec_passkey_store_credential"
|
||||
PasskeyGetCredsByUsername string // default: "resolvespec_passkey_get_credentials_by_username"
|
||||
PasskeyGetCredential string // default: "resolvespec_passkey_get_credential"
|
||||
PasskeyUpdateCounter string // default: "resolvespec_passkey_update_counter"
|
||||
PasskeyGetUserCredentials string // default: "resolvespec_passkey_get_user_credentials"
|
||||
PasskeyDeleteCredential string // default: "resolvespec_passkey_delete_credential"
|
||||
PasskeyUpdateName string // default: "resolvespec_passkey_update_name"
|
||||
PasskeyLogin string // default: "resolvespec_passkey_login"
|
||||
|
||||
// Password reset procedures (DatabaseAuthenticator)
|
||||
PasswordResetRequest string // default: "resolvespec_password_reset_request"
|
||||
PasswordResetComplete string // default: "resolvespec_password_reset"
|
||||
|
||||
// OAuth2 procedures (DatabaseAuthenticator OAuth2 methods)
|
||||
OAuthGetOrCreateUser string // default: "resolvespec_oauth_getorcreateuser"
|
||||
OAuthCreateSession string // default: "resolvespec_oauth_createsession"
|
||||
OAuthGetRefreshToken string // default: "resolvespec_oauth_getrefreshtoken"
|
||||
OAuthUpdateRefreshToken string // default: "resolvespec_oauth_updaterefreshtoken"
|
||||
OAuthGetUser string // default: "resolvespec_oauth_getuser"
|
||||
|
||||
// OAuth2 server procedures (OAuthServer persistence)
|
||||
OAuthRegisterClient string // default: "resolvespec_oauth_register_client"
|
||||
OAuthGetClient string // default: "resolvespec_oauth_get_client"
|
||||
OAuthSaveCode string // default: "resolvespec_oauth_save_code"
|
||||
OAuthExchangeCode string // default: "resolvespec_oauth_exchange_code"
|
||||
OAuthIntrospect string // default: "resolvespec_oauth_introspect"
|
||||
OAuthRevoke string // default: "resolvespec_oauth_revoke"
|
||||
|
||||
// Keystore procedures (KeyStore)
|
||||
KeystoreGetUserKeys string // default: "resolvespec_keystore_get_user_keys"
|
||||
KeystoreCreateKey string // default: "resolvespec_keystore_create_key"
|
||||
KeystoreDeleteKey string // default: "resolvespec_keystore_delete_key"
|
||||
KeystoreValidateKey string // default: "resolvespec_keystore_validate_key"
|
||||
}
|
||||
|
||||
// DefaultProcNames returns the default resolvespec_* procedure names.
|
||||
func DefaultProcNames() ProcNames {
|
||||
return ProcNames{ //nolint:gosec // G101: false positive: identifiers, not credentials
|
||||
Login: "resolvespec_login",
|
||||
Register: "resolvespec_register",
|
||||
Logout: "resolvespec_logout",
|
||||
Session: "resolvespec_session",
|
||||
SessionUpdate: "resolvespec_session_update",
|
||||
RefreshToken: "resolvespec_refresh_token",
|
||||
LoginAPIKey: "resolvespec_login_api_key",
|
||||
JWTLogin: "resolvespec_jwt_login",
|
||||
JWTLogout: "resolvespec_jwt_logout",
|
||||
ColumnSecurity: "resolvespec_column_security",
|
||||
RowSecurity: "resolvespec_row_security",
|
||||
TOTPEnable: "resolvespec_totp_enable",
|
||||
TOTPDisable: "resolvespec_totp_disable",
|
||||
TOTPGetStatus: "resolvespec_totp_get_status",
|
||||
TOTPGetSecret: "resolvespec_totp_get_secret",
|
||||
TOTPRegenerateBackup: "resolvespec_totp_regenerate_backup_codes",
|
||||
TOTPValidateBackupCode: "resolvespec_totp_validate_backup_code",
|
||||
PasskeyStoreCredential: "resolvespec_passkey_store_credential",
|
||||
PasskeyGetCredsByUsername: "resolvespec_passkey_get_credentials_by_username",
|
||||
PasskeyGetCredential: "resolvespec_passkey_get_credential",
|
||||
PasskeyUpdateCounter: "resolvespec_passkey_update_counter",
|
||||
PasskeyGetUserCredentials: "resolvespec_passkey_get_user_credentials",
|
||||
PasskeyDeleteCredential: "resolvespec_passkey_delete_credential",
|
||||
PasskeyUpdateName: "resolvespec_passkey_update_name",
|
||||
PasskeyLogin: "resolvespec_passkey_login",
|
||||
PasswordResetRequest: "resolvespec_password_reset_request",
|
||||
PasswordResetComplete: "resolvespec_password_reset",
|
||||
OAuthGetOrCreateUser: "resolvespec_oauth_getorcreateuser",
|
||||
OAuthCreateSession: "resolvespec_oauth_createsession",
|
||||
OAuthGetRefreshToken: "resolvespec_oauth_getrefreshtoken",
|
||||
OAuthUpdateRefreshToken: "resolvespec_oauth_updaterefreshtoken",
|
||||
OAuthGetUser: "resolvespec_oauth_getuser",
|
||||
OAuthRegisterClient: "resolvespec_oauth_register_client",
|
||||
OAuthGetClient: "resolvespec_oauth_get_client",
|
||||
OAuthSaveCode: "resolvespec_oauth_save_code",
|
||||
OAuthExchangeCode: "resolvespec_oauth_exchange_code",
|
||||
OAuthIntrospect: "resolvespec_oauth_introspect",
|
||||
OAuthRevoke: "resolvespec_oauth_revoke",
|
||||
KeystoreGetUserKeys: "resolvespec_keystore_get_user_keys",
|
||||
KeystoreCreateKey: "resolvespec_keystore_create_key",
|
||||
KeystoreDeleteKey: "resolvespec_keystore_delete_key",
|
||||
KeystoreValidateKey: "resolvespec_keystore_validate_key",
|
||||
}
|
||||
}
|
||||
|
||||
// Merge returns a copy of p with every non-empty field of override applied.
|
||||
func (p ProcNames) Merge(override ProcNames) ProcNames {
|
||||
merged := p
|
||||
mv, ov := reflect.ValueOf(&merged).Elem(), reflect.ValueOf(override)
|
||||
for i := 0; i < ov.NumField(); i++ {
|
||||
if v := ov.Field(i).String(); v != "" {
|
||||
mv.Field(i).SetString(v)
|
||||
}
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
// Validate checks that every name is a safe (optionally schema-qualified) identifier.
|
||||
func (p ProcNames) Validate() error {
|
||||
v, t := reflect.ValueOf(p), reflect.TypeOf(p)
|
||||
for i := 0; i < v.NumField(); i++ {
|
||||
if name := v.Field(i).String(); !validQualifiedIdent(name) {
|
||||
return fmt.Errorf("lookup: invalid procedure name %q for %s", name, t.Field(i).Name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sort"
|
||||
)
|
||||
|
||||
var (
|
||||
identRe = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`)
|
||||
qualifiedIdentRe = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*(\.[a-zA-Z_][a-zA-Z0-9_]*)?$`)
|
||||
)
|
||||
|
||||
func validIdent(s string) bool { return identRe.MatchString(s) }
|
||||
func validQualifiedIdent(s string) bool { return qualifiedIdentRe.MatchString(s) }
|
||||
|
||||
// Entity identifies one table the direct backend reads or writes.
|
||||
type Entity string
|
||||
|
||||
const (
|
||||
EntityUsers Entity = "users"
|
||||
EntityUserSessions Entity = "user_sessions"
|
||||
EntityTokenBlacklist Entity = "token_blacklist"
|
||||
EntityUserTOTPBackupCodes Entity = "user_totp_backup_codes"
|
||||
EntityUserPasskeyCredentials Entity = "user_passkey_credentials"
|
||||
EntityUserPasswordResets Entity = "user_password_resets"
|
||||
EntityOAuthClients Entity = "oauth_clients"
|
||||
EntityOAuthCodes Entity = "oauth_codes"
|
||||
EntityUserKeys Entity = "user_keys"
|
||||
EntitySecGroupMembers Entity = "sec_group_members"
|
||||
EntitySecColumnRules Entity = "sec_column_rules"
|
||||
EntitySecRowRules Entity = "sec_row_rules"
|
||||
)
|
||||
|
||||
// Column is a typed key naming one logical column of an entity. The physical column
|
||||
// name is looked up in the Schema, so every column of every entity is configurable.
|
||||
type Column struct {
|
||||
Entity Entity
|
||||
Name string
|
||||
}
|
||||
|
||||
func (c Column) String() string { return string(c.Entity) + "." + c.Name }
|
||||
|
||||
func col(e Entity, name string) Column { return Column{Entity: e, Name: name} }
|
||||
|
||||
// Logical columns. The names are the default physical column names.
|
||||
var (
|
||||
UsersID = col(EntityUsers, "id")
|
||||
UsersUsername = col(EntityUsers, "username")
|
||||
UsersEmail = col(EntityUsers, "email")
|
||||
UsersPassword = col(EntityUsers, "password")
|
||||
UsersUserLevel = col(EntityUsers, "user_level")
|
||||
UsersRoles = col(EntityUsers, "roles")
|
||||
UsersIsActive = col(EntityUsers, "is_active")
|
||||
UsersCreatedAt = col(EntityUsers, "created_at")
|
||||
UsersUpdatedAt = col(EntityUsers, "updated_at")
|
||||
UsersLastLoginAt = col(EntityUsers, "last_login_at")
|
||||
UsersProgramUserID = col(EntityUsers, "program_user_id")
|
||||
UsersProgramUserTable = col(EntityUsers, "program_user_table")
|
||||
UsersRemoteID = col(EntityUsers, "remote_id")
|
||||
UsersAuthProvider = col(EntityUsers, "auth_provider")
|
||||
UsersTOTPSecret = col(EntityUsers, "totp_secret")
|
||||
UsersTOTPEnabled = col(EntityUsers, "totp_enabled")
|
||||
UsersTOTPEnabledAt = col(EntityUsers, "totp_enabled_at")
|
||||
|
||||
SessionsID = col(EntityUserSessions, "id")
|
||||
SessionsToken = col(EntityUserSessions, "session_token")
|
||||
SessionsUserID = col(EntityUserSessions, "user_id")
|
||||
SessionsExpiresAt = col(EntityUserSessions, "expires_at")
|
||||
SessionsCreatedAt = col(EntityUserSessions, "created_at")
|
||||
SessionsLastActivityAt = col(EntityUserSessions, "last_activity_at")
|
||||
SessionsIPAddress = col(EntityUserSessions, "ip_address")
|
||||
SessionsUserAgent = col(EntityUserSessions, "user_agent")
|
||||
SessionsAccessToken = col(EntityUserSessions, "access_token")
|
||||
SessionsRefreshToken = col(EntityUserSessions, "refresh_token")
|
||||
SessionsTokenType = col(EntityUserSessions, "token_type")
|
||||
SessionsAuthProvider = col(EntityUserSessions, "auth_provider")
|
||||
|
||||
BlacklistID = col(EntityTokenBlacklist, "id")
|
||||
BlacklistToken = col(EntityTokenBlacklist, "token")
|
||||
BlacklistUserID = col(EntityTokenBlacklist, "user_id")
|
||||
BlacklistExpiresAt = col(EntityTokenBlacklist, "expires_at")
|
||||
BlacklistCreatedAt = col(EntityTokenBlacklist, "created_at")
|
||||
|
||||
BackupCodesID = col(EntityUserTOTPBackupCodes, "id")
|
||||
BackupCodesUserID = col(EntityUserTOTPBackupCodes, "user_id")
|
||||
BackupCodesCodeHash = col(EntityUserTOTPBackupCodes, "code_hash")
|
||||
BackupCodesUsed = col(EntityUserTOTPBackupCodes, "used")
|
||||
BackupCodesUsedAt = col(EntityUserTOTPBackupCodes, "used_at")
|
||||
BackupCodesCreatedAt = col(EntityUserTOTPBackupCodes, "created_at")
|
||||
|
||||
PasskeyID = col(EntityUserPasskeyCredentials, "id")
|
||||
PasskeyUserID = col(EntityUserPasskeyCredentials, "user_id")
|
||||
PasskeyCredentialID = col(EntityUserPasskeyCredentials, "credential_id")
|
||||
PasskeyPublicKey = col(EntityUserPasskeyCredentials, "public_key")
|
||||
PasskeyAttestationType = col(EntityUserPasskeyCredentials, "attestation_type")
|
||||
PasskeyAAGUID = col(EntityUserPasskeyCredentials, "aaguid")
|
||||
PasskeySignCount = col(EntityUserPasskeyCredentials, "sign_count")
|
||||
PasskeyCloneWarning = col(EntityUserPasskeyCredentials, "clone_warning")
|
||||
PasskeyTransports = col(EntityUserPasskeyCredentials, "transports")
|
||||
PasskeyBackupEligible = col(EntityUserPasskeyCredentials, "backup_eligible")
|
||||
PasskeyBackupState = col(EntityUserPasskeyCredentials, "backup_state")
|
||||
PasskeyName = col(EntityUserPasskeyCredentials, "name")
|
||||
PasskeyCreatedAt = col(EntityUserPasskeyCredentials, "created_at")
|
||||
PasskeyLastUsedAt = col(EntityUserPasskeyCredentials, "last_used_at")
|
||||
|
||||
ResetsID = col(EntityUserPasswordResets, "id")
|
||||
ResetsUserID = col(EntityUserPasswordResets, "user_id")
|
||||
ResetsTokenHash = col(EntityUserPasswordResets, "token_hash")
|
||||
ResetsExpiresAt = col(EntityUserPasswordResets, "expires_at")
|
||||
ResetsCreatedAt = col(EntityUserPasswordResets, "created_at")
|
||||
ResetsUsed = col(EntityUserPasswordResets, "used")
|
||||
ResetsUsedAt = col(EntityUserPasswordResets, "used_at")
|
||||
|
||||
OAuthClientsID = col(EntityOAuthClients, "id")
|
||||
OAuthClientsClientID = col(EntityOAuthClients, "client_id")
|
||||
OAuthClientsRedirectURIs = col(EntityOAuthClients, "redirect_uris")
|
||||
OAuthClientsClientName = col(EntityOAuthClients, "client_name")
|
||||
OAuthClientsGrantTypes = col(EntityOAuthClients, "grant_types")
|
||||
OAuthClientsAllowedScopes = col(EntityOAuthClients, "allowed_scopes")
|
||||
OAuthClientsClientSecretHash = col(EntityOAuthClients, "client_secret_hash")
|
||||
OAuthClientsTokenEndpointAuthMethod = col(EntityOAuthClients, "token_endpoint_auth_method")
|
||||
OAuthClientsIsActive = col(EntityOAuthClients, "is_active")
|
||||
OAuthClientsCreatedAt = col(EntityOAuthClients, "created_at")
|
||||
|
||||
OAuthCodesID = col(EntityOAuthCodes, "id")
|
||||
OAuthCodesCode = col(EntityOAuthCodes, "code")
|
||||
OAuthCodesClientID = col(EntityOAuthCodes, "client_id")
|
||||
OAuthCodesRedirectURI = col(EntityOAuthCodes, "redirect_uri")
|
||||
OAuthCodesClientState = col(EntityOAuthCodes, "client_state")
|
||||
OAuthCodesCodeChallenge = col(EntityOAuthCodes, "code_challenge")
|
||||
OAuthCodesCodeChallengeMethod = col(EntityOAuthCodes, "code_challenge_method")
|
||||
OAuthCodesSessionToken = col(EntityOAuthCodes, "session_token")
|
||||
OAuthCodesRefreshToken = col(EntityOAuthCodes, "refresh_token")
|
||||
OAuthCodesScopes = col(EntityOAuthCodes, "scopes")
|
||||
OAuthCodesExpiresAt = col(EntityOAuthCodes, "expires_at")
|
||||
OAuthCodesCreatedAt = col(EntityOAuthCodes, "created_at")
|
||||
|
||||
KeysID = col(EntityUserKeys, "id")
|
||||
KeysUserID = col(EntityUserKeys, "user_id")
|
||||
KeysKeyType = col(EntityUserKeys, "key_type")
|
||||
KeysKeyHash = col(EntityUserKeys, "key_hash")
|
||||
KeysName = col(EntityUserKeys, "name")
|
||||
KeysScopes = col(EntityUserKeys, "scopes")
|
||||
KeysMeta = col(EntityUserKeys, "meta")
|
||||
KeysExpiresAt = col(EntityUserKeys, "expires_at")
|
||||
KeysCreatedAt = col(EntityUserKeys, "created_at")
|
||||
KeysLastUsedAt = col(EntityUserKeys, "last_used_at")
|
||||
KeysIsActive = col(EntityUserKeys, "is_active")
|
||||
|
||||
GroupMembersGroupID = col(EntitySecGroupMembers, "group_id")
|
||||
GroupMembersUserID = col(EntitySecGroupMembers, "user_id")
|
||||
|
||||
ColRulesID = col(EntitySecColumnRules, "id")
|
||||
ColRulesUserID = col(EntitySecColumnRules, "user_id")
|
||||
ColRulesGroupID = col(EntitySecColumnRules, "group_id")
|
||||
ColRulesSchemaName = col(EntitySecColumnRules, "schema_name")
|
||||
ColRulesTableName = col(EntitySecColumnRules, "table_name")
|
||||
ColRulesColumnPath = col(EntitySecColumnRules, "column_path")
|
||||
ColRulesAccessType = col(EntitySecColumnRules, "access_type")
|
||||
ColRulesMaskStart = col(EntitySecColumnRules, "mask_start")
|
||||
ColRulesMaskEnd = col(EntitySecColumnRules, "mask_end")
|
||||
ColRulesMaskInvert = col(EntitySecColumnRules, "mask_invert")
|
||||
ColRulesMaskChar = col(EntitySecColumnRules, "mask_char")
|
||||
ColRulesExtraFilters = col(EntitySecColumnRules, "extra_filters")
|
||||
ColRulesIsActive = col(EntitySecColumnRules, "is_active")
|
||||
|
||||
RowRulesID = col(EntitySecRowRules, "id")
|
||||
RowRulesUserID = col(EntitySecRowRules, "user_id")
|
||||
RowRulesGroupID = col(EntitySecRowRules, "group_id")
|
||||
RowRulesSchemaName = col(EntitySecRowRules, "schema_name")
|
||||
RowRulesTableName = col(EntitySecRowRules, "table_name")
|
||||
RowRulesTemplate = col(EntitySecRowRules, "template")
|
||||
RowRulesHasBlock = col(EntitySecRowRules, "has_block")
|
||||
RowRulesIsActive = col(EntitySecRowRules, "is_active")
|
||||
)
|
||||
|
||||
// allColumns lists every logical column; it defines the default schema.
|
||||
var allColumns = []Column{
|
||||
UsersID, UsersUsername, UsersEmail, UsersPassword, UsersUserLevel, UsersRoles, UsersIsActive,
|
||||
UsersCreatedAt, UsersUpdatedAt, UsersLastLoginAt, UsersProgramUserID, UsersProgramUserTable,
|
||||
UsersRemoteID, UsersAuthProvider, UsersTOTPSecret, UsersTOTPEnabled, UsersTOTPEnabledAt,
|
||||
SessionsID, SessionsToken, SessionsUserID, SessionsExpiresAt, SessionsCreatedAt, SessionsLastActivityAt,
|
||||
SessionsIPAddress, SessionsUserAgent, SessionsAccessToken, SessionsRefreshToken, SessionsTokenType, SessionsAuthProvider,
|
||||
BlacklistID, BlacklistToken, BlacklistUserID, BlacklistExpiresAt, BlacklistCreatedAt,
|
||||
BackupCodesID, BackupCodesUserID, BackupCodesCodeHash, BackupCodesUsed, BackupCodesUsedAt, BackupCodesCreatedAt,
|
||||
PasskeyID, PasskeyUserID, PasskeyCredentialID, PasskeyPublicKey, PasskeyAttestationType, PasskeyAAGUID,
|
||||
PasskeySignCount, PasskeyCloneWarning, PasskeyTransports, PasskeyBackupEligible, PasskeyBackupState,
|
||||
PasskeyName, PasskeyCreatedAt, PasskeyLastUsedAt,
|
||||
ResetsID, ResetsUserID, ResetsTokenHash, ResetsExpiresAt, ResetsCreatedAt, ResetsUsed, ResetsUsedAt,
|
||||
OAuthClientsID, OAuthClientsClientID, OAuthClientsRedirectURIs, OAuthClientsClientName, OAuthClientsGrantTypes,
|
||||
OAuthClientsAllowedScopes, OAuthClientsClientSecretHash, OAuthClientsTokenEndpointAuthMethod,
|
||||
OAuthClientsIsActive, OAuthClientsCreatedAt,
|
||||
OAuthCodesID, OAuthCodesCode, OAuthCodesClientID, OAuthCodesRedirectURI, OAuthCodesClientState,
|
||||
OAuthCodesCodeChallenge, OAuthCodesCodeChallengeMethod, OAuthCodesSessionToken, OAuthCodesRefreshToken,
|
||||
OAuthCodesScopes, OAuthCodesExpiresAt, OAuthCodesCreatedAt,
|
||||
KeysID, KeysUserID, KeysKeyType, KeysKeyHash, KeysName, KeysScopes, KeysMeta, KeysExpiresAt,
|
||||
KeysCreatedAt, KeysLastUsedAt, KeysIsActive,
|
||||
GroupMembersGroupID, GroupMembersUserID,
|
||||
ColRulesID, ColRulesUserID, ColRulesGroupID, ColRulesSchemaName, ColRulesTableName, ColRulesColumnPath,
|
||||
ColRulesAccessType, ColRulesMaskStart, ColRulesMaskEnd, ColRulesMaskInvert, ColRulesMaskChar,
|
||||
ColRulesExtraFilters, ColRulesIsActive,
|
||||
RowRulesID, RowRulesUserID, RowRulesGroupID, RowRulesSchemaName, RowRulesTableName, RowRulesTemplate,
|
||||
RowRulesHasBlock, RowRulesIsActive,
|
||||
}
|
||||
|
||||
// Table maps one entity to a physical table and its columns.
|
||||
type Table struct {
|
||||
// Schema optionally qualifies the table (schema.table). Empty = unqualified.
|
||||
Schema string
|
||||
// Name is the physical table name. Empty = default (the entity name).
|
||||
Name string
|
||||
// Columns maps logical column name -> physical column name. Missing = default.
|
||||
Columns map[string]string
|
||||
}
|
||||
|
||||
// Schema maps every entity to its physical table and columns. The zero value is
|
||||
// valid and means "all defaults"; use DefaultSchema for the explicit baseline.
|
||||
type Schema map[Entity]Table
|
||||
|
||||
// DefaultSchema returns the baseline schema: every entity and column under its default name.
|
||||
func DefaultSchema() Schema {
|
||||
s := Schema{}
|
||||
for _, c := range allColumns {
|
||||
t, ok := s[c.Entity]
|
||||
if !ok {
|
||||
t = Table{Name: string(c.Entity), Columns: map[string]string{}}
|
||||
}
|
||||
t.Columns[c.Name] = c.Name
|
||||
s[c.Entity] = t
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// Merge returns a copy of s with every non-empty field of override applied.
|
||||
// Unknown entities or columns in override are kept so Validate can report them.
|
||||
func (s Schema) Merge(override Schema) Schema {
|
||||
merged := Schema{}
|
||||
for e, t := range s {
|
||||
merged[e] = cloneTable(t)
|
||||
}
|
||||
for e, ot := range override {
|
||||
t, ok := merged[e]
|
||||
if !ok {
|
||||
merged[e] = cloneTable(ot)
|
||||
continue
|
||||
}
|
||||
if ot.Schema != "" {
|
||||
t.Schema = ot.Schema
|
||||
}
|
||||
if ot.Name != "" {
|
||||
t.Name = ot.Name
|
||||
}
|
||||
for k, v := range ot.Columns {
|
||||
if v != "" {
|
||||
t.Columns[k] = v
|
||||
}
|
||||
}
|
||||
merged[e] = t
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func cloneTable(t Table) Table {
|
||||
c := t
|
||||
c.Columns = make(map[string]string, len(t.Columns))
|
||||
for k, v := range t.Columns {
|
||||
c.Columns[k] = v
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// Validate checks that the schema only names known entities and columns and that
|
||||
// every identifier is safe. It is meant to run on the merged (default + override) schema.
|
||||
func (s Schema) Validate() error {
|
||||
known := map[Column]bool{}
|
||||
for _, c := range allColumns {
|
||||
known[c] = true
|
||||
}
|
||||
entities := make([]string, 0, len(s))
|
||||
for e := range s {
|
||||
entities = append(entities, string(e))
|
||||
}
|
||||
sort.Strings(entities)
|
||||
for _, en := range entities {
|
||||
e := Entity(en)
|
||||
t := s[e]
|
||||
if firstKnownColumn(e) == "" {
|
||||
return fmt.Errorf("lookup: unknown entity %q", e)
|
||||
}
|
||||
if !validQualifiedIdent(t.Name) {
|
||||
return fmt.Errorf("lookup: invalid table name %q for %s", t.Name, e)
|
||||
}
|
||||
if t.Schema != "" && !validIdent(t.Schema) {
|
||||
return fmt.Errorf("lookup: invalid schema name %q for %s", t.Schema, e)
|
||||
}
|
||||
names := make([]string, 0, len(t.Columns))
|
||||
for n := range t.Columns {
|
||||
names = append(names, n)
|
||||
}
|
||||
sort.Strings(names)
|
||||
for _, n := range names {
|
||||
if !known[Column{Entity: e, Name: n}] {
|
||||
return fmt.Errorf("lookup: unknown column %q for %s", n, e)
|
||||
}
|
||||
if !validIdent(t.Columns[n]) {
|
||||
return fmt.Errorf("lookup: invalid column name %q for %s.%s", t.Columns[n], e, n)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func firstKnownColumn(e Entity) string {
|
||||
for _, c := range allColumns {
|
||||
if c.Entity == e {
|
||||
return c.Name
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// TableName returns the physical table name of an entity (unqualified, unquoted).
|
||||
func (s Schema) TableName(e Entity) string {
|
||||
if t, ok := s[e]; ok && t.Name != "" {
|
||||
return t.Name
|
||||
}
|
||||
return string(e)
|
||||
}
|
||||
|
||||
// SchemaName returns the optional schema qualifier of an entity.
|
||||
func (s Schema) SchemaName(e Entity) string { return s[e].Schema }
|
||||
|
||||
// Col returns the physical column name for a logical column (unquoted).
|
||||
func (s Schema) Col(c Column) string {
|
||||
if t, ok := s[c.Entity]; ok {
|
||||
if n := t.Columns[c.Name]; n != "" {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return c.Name
|
||||
}
|
||||
|
||||
// FirstColumn returns the name of the first logical column of an entity (used by the
|
||||
// direct backend for existence checks).
|
||||
func FirstColumn(e Entity) string { return firstKnownColumn(e) }
|
||||
@@ -0,0 +1,47 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sort"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
// settingNameRE matches a custom GUC name: two or more dot-separated identifiers.
|
||||
var settingNameRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)+$`)
|
||||
|
||||
// ApplyTxSettings sets each entry as a transaction-local setting on tx, in name
|
||||
// order. Postgres only; any other driver with a non-empty map is an error so a
|
||||
// missing RLS stamp fails closed.
|
||||
func ApplyTxSettings(ctx context.Context, tx common.Database, settings map[string]string) error {
|
||||
if len(settings) == 0 {
|
||||
return nil
|
||||
}
|
||||
if tx == nil {
|
||||
return fmt.Errorf("tx settings: no transaction")
|
||||
}
|
||||
if drv := tx.DriverName(); drv != "postgres" && drv != "pgsql" {
|
||||
return fmt.Errorf("tx settings: unsupported driver %q", drv)
|
||||
}
|
||||
names := make([]string, 0, len(settings))
|
||||
for name := range settings {
|
||||
if !settingNameRE.MatchString(name) {
|
||||
return fmt.Errorf("tx settings: invalid setting name %q", name)
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
sort.Strings(names)
|
||||
for _, name := range names {
|
||||
// The value is hex-encoded so it needs no quoting and cannot be read as a
|
||||
// bind placeholder by any adapter.
|
||||
query := fmt.Sprintf("SELECT set_config('%s', convert_from(decode('%s', 'hex'), 'UTF8'), true)",
|
||||
name, hex.EncodeToString([]byte(settings[name])))
|
||||
if _, err := tx.Exec(ctx, query); err != nil {
|
||||
return fmt.Errorf("tx settings: set %s: %w", name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package lookup_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
|
||||
)
|
||||
|
||||
func txSettingsDB(t *testing.T) (common.Database, sqlmock.Sqlmock) {
|
||||
t.Helper()
|
||||
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return database.NewPgSQLAdapter(db), mock
|
||||
}
|
||||
|
||||
func TestApplyTxSettingsStampsInNameOrderOnTx(t *testing.T) {
|
||||
pool, mock := txSettingsDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectExec(`set_config\('app\.tenant', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectExec(`set_config\('app\.user_id', .*decode\('`).WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectCommit()
|
||||
|
||||
err := pool.RunInTransaction(context.Background(), func(tx common.Database) error {
|
||||
return lookup.ApplyTxSettings(ctx, tx, map[string]string{"app.user_id": "7", "app.tenant": "o'x?"})
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyTxSettingsValueIsNeverInlined(t *testing.T) {
|
||||
pool, mock := txSettingsDB(t)
|
||||
ctx := context.Background()
|
||||
var seen string
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectExec(`set_config`).WillReturnResult(sqlmock.NewResult(0, 0))
|
||||
mock.ExpectCommit()
|
||||
_ = pool.RunInTransaction(context.Background(), func(tx common.Database) error {
|
||||
// Capture via a wrapper so the raw statement can be inspected.
|
||||
return lookup.ApplyTxSettings(ctx, &queryRecorder{Database: tx, got: &seen}, map[string]string{"app.v": "'; DROP TABLE x; --?"})
|
||||
})
|
||||
if strings.Contains(seen, "DROP") || strings.Contains(seen, "?") {
|
||||
t.Fatalf("value leaked into SQL text: %s", seen)
|
||||
}
|
||||
}
|
||||
|
||||
type queryRecorder struct {
|
||||
common.Database
|
||||
got *string
|
||||
}
|
||||
|
||||
func (q *queryRecorder) Exec(ctx context.Context, query string, args ...interface{}) (common.Result, error) {
|
||||
*q.got = query
|
||||
return q.Database.Exec(ctx, query, args...)
|
||||
}
|
||||
|
||||
func TestApplyTxSettingsRejectsBadNameAndDriver(t *testing.T) {
|
||||
pool, _ := txSettingsDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, name := range []string{"user_id", "app.x'); DROP", "app..x", "app.x y"} {
|
||||
if err := lookup.ApplyTxSettings(ctx, pool, map[string]string{name: "1"}); err == nil {
|
||||
t.Fatalf("name %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
if err := lookup.ApplyTxSettings(ctx, pool, nil); err != nil {
|
||||
t.Fatalf("empty settings must be a no-op: %v", err)
|
||||
}
|
||||
if err := lookup.ApplyTxSettings(ctx, &driverStub{Database: pool, name: "sqlite"}, map[string]string{"app.x": "1"}); err == nil {
|
||||
t.Fatal("non-postgres driver must fail closed")
|
||||
}
|
||||
}
|
||||
|
||||
type driverStub struct {
|
||||
common.Database
|
||||
name string
|
||||
}
|
||||
|
||||
func (d *driverStub) DriverName() string { return d.name }
|
||||
Reference in New Issue
Block a user