From ff76eb8e1f17a4521d71f23c1b121667db3abfc6 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 22:55:26 +0200 Subject: [PATCH] feat(security): stamp transaction-local settings on OnTxBegin in all specs --- audit/single_tran.md | 6 +- pkg/common/TRANSACTIONS.md | 45 ++++++++++ pkg/funcspec/security_adapter.go | 6 ++ pkg/mqttspec/security_hooks.go | 6 ++ pkg/resolvemcp/security_hooks.go | 6 ++ pkg/resolvespec/security_hooks.go | 6 ++ pkg/resolvespec/tx_settings_test.go | 78 +++++++++++++++++ pkg/restheadspec/security_hooks.go | 6 ++ pkg/security/README.md | 2 + pkg/security/provider.go | 4 + pkg/security/txsettings.go | 86 ++++++++++++++++++ pkg/security/txsettings_test.go | 131 ++++++++++++++++++++++++++++ pkg/websocketspec/security_hooks.go | 6 ++ 13 files changed, 386 insertions(+), 2 deletions(-) create mode 100644 pkg/common/TRANSACTIONS.md create mode 100644 pkg/resolvespec/tx_settings_test.go create mode 100644 pkg/security/txsettings.go create mode 100644 pkg/security/txsettings_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 1039c5e..fa0b2b4 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -78,7 +78,7 @@ | 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 5 | DONE | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 6 | DONE | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | -| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | +| 7 | DONE | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | ## Progress - DONE P0: baseline via `dbtrace` on real Postgres (commit `cd96404`): create/read/delete `pooled=0`; update `pooled=1` (re-fetch) = P3 target. websocketspec/mqttspec/resolvemcp not measured. @@ -93,7 +93,9 @@ - DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx"). - DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock). - DONE P6: funcspec. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx` for `SqlQuery` and `SqlQueryList`. `BeforeResponse` now runs in a second short tx (`Tx` is no longer the pool). `BeforeOp` is unchanged (still per statement). A begin/`OnTxBegin`/commit failure answers 500 `transaction_error` / "Transaction failed" (before, it returned with no response); body failures still answer via `sendError`. Tests: `pkg/funcspec/tx_test.go`. -- NEXT: P7. +- DONE (AfterRead, decided by user): restheadspec `AfterRead` now runs in a second short tx. Test: `pkg/restheadspec/read_tx_test.go`. +- DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. +- NEXT: extra tests (create/update for other specs, `dbtrace` `pooled == 0` on real Postgres). ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/common/TRANSACTIONS.md b/pkg/common/TRANSACTIONS.md new file mode 100644 index 0000000..df2f8d2 --- /dev/null +++ b/pkg/common/TRANSACTIONS.md @@ -0,0 +1,45 @@ +# Request Transactions (cheatsheet) + +Every DB statement and every DB-touching hook of one request runs on one transaction. Hooks never get the pool. + +## Rules +- `hookCtx.Tx` is always the open transaction (never `h.db`), except `BeforeHandle`, which runs before any tx and must not touch the DB. +- `OnTxBegin` fires once, first, in every tx the handler opens (incl. the second short tx). +- `OnTxBegin` error or abort: rollback, client gets a generic error, nothing leaked. +- Begin or commit failure: generic error (`transaction_error` / "Transaction failed" in websocketspec, mqttspec, funcspec). +- Transaction-local state (`set_config(..., true)`, RLS GUCs) is only visible on that tx. Set it in `OnTxBegin`. + +## Transactions per operation +| Operation | Tx 1 | Tx 2 (short, after commit) | +|---|---|---| +| read | `BeforeRead`, count, scan, `AfterRead`* | restheadspec: `AfterRead` | +| create | `Before*`, insert | re-fetch, `BeforeScan`, `AfterCreate` | +| update | `Before*`, select, update, `AfterUpdate`† | re-fetch (+ `AfterUpdate` where noted) | +| delete (single/batch) | `BeforeDelete`, select, delete, `AfterDelete` | none | +| funcspec query | `BeforeQuery*`, `BeforeSQLExec`, SQL, `After*` | `BeforeResponse` | + +\* resolvespec, websocketspec, mqttspec, resolvemcp. restheadspec runs `AfterRead` in tx 2. +† restheadspec, websocketspec, mqttspec run `AfterUpdate` in tx 2. resolvespec, resolvemcp run it in tx 1. + +Tx 2 exists so the re-fetch sees trigger changes from the committed write. + +## Per spec +| Spec | `OnTxBegin` | Helper | +|---|---|---| +| resolvespec, restheadspec, websocketspec, resolvemcp, funcspec | own `HookType` = `common.TxHookName` | `Handler.runInTx` | +| mqttspec | re-exports `websocketspec.OnTxBegin` | `Handler.runInTx` | + +- `common.RunRequestTx(ctx, db, TxContext, onBegin, body)`: open tx, `SetTx`, `onBegin`, `body`. +- `common.TxContext`: `SetTx(tx)`; implemented by each spec's `HookContext`. + +## RLS / transaction settings (pkg/security) +- `SecurityList.SetTxSettings(fn)`: `fn(SecurityContext) (map[string]string, error)`; nil disables. +- Every spec's `RegisterSecurityHooks` registers `OnTxBegin` → `security.StampTxSettings`. `fn` is read per call, so set order does not matter. +- Stamps via `set_config(name, value, true)` in name order, before any other SQL. +- Fail closed: `fn` error, invalid name, or non-Postgres driver with a non-empty map aborts the tx. +- Name: dotted identifier (`ns.name`). Value is hex-encoded in SQL, never inlined. +- Low level: `security.ApplyTxSettings(secCtx, tx, map)`. + +## Test notes +- sqlmock + `SetMaxOpenConns(1)`: any pool use inside an open tx blocks and fails. +- restheadspec model-based updates/reads need the bun adapter. diff --git a/pkg/funcspec/security_adapter.go b/pkg/funcspec/security_adapter.go index f0201cc..5648b2b 100644 --- a/pkg/funcspec/security_adapter.go +++ b/pkg/funcspec/security_adapter.go @@ -12,6 +12,12 @@ import ( // Note: funcspec operates on SQL queries directly, so row-level security is not directly applicable // We provide auth enforcement and audit logging for data access tracking func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newFuncSpecSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeQueryList - Auth check before list query execution handler.Hooks().Register(BeforeQueryList, func(hookCtx *HookContext) error { if hookCtx.UserContext == nil || hookCtx.UserContext.UserID == 0 { diff --git a/pkg/mqttspec/security_hooks.go b/pkg/mqttspec/security_hooks.go index a92462c..bcdec15 100644 --- a/pkg/mqttspec/security_hooks.go +++ b/pkg/mqttspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the MQTT handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvemcp/security_hooks.go b/pkg/resolvemcp/security_hooks.go index cb8f2ab..f7d8bd4 100644 --- a/pkg/resolvemcp/security_hooks.go +++ b/pkg/resolvemcp/security_hooks.go @@ -19,6 +19,12 @@ import ( // - Column-level security: sensitive columns masked/hidden in read results. // - Audit logging after each read. func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // BeforeHandle: enforce model-level operation rules (auth check). handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvespec/security_hooks.go b/pkg/resolvespec/security_hooks.go index 450e33e..5db28a5 100644 --- a/pkg/resolvespec/security_hooks.go +++ b/pkg/resolvespec/security_hooks.go @@ -11,6 +11,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/resolvespec/tx_settings_test.go b/pkg/resolvespec/tx_settings_test.go new file mode 100644 index 0000000..42d1918 --- /dev/null +++ b/pkg/resolvespec/tx_settings_test.go @@ -0,0 +1,78 @@ +package resolvespec + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" +) + +// stubProvider satisfies security.SecurityProvider; its methods are never called +// because the model rules are public and no rules are loaded for this test. +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsFirstOnEveryTx(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + base = context.WithValue(base, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleDelete(ctx, w, "7", nil) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestSecurityHooksTxSettingsErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + list, _ := security.NewSecurityList(stubProvider{}) + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"bad name": "1"}, nil + }) + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + base = context.WithValue(base, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}) + h.handleDelete(ctx, w, "7", nil) + + if rec.Code == http.StatusOK { + t.Fatal("a failed stamp must not let the delete proceed") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/restheadspec/security_hooks.go b/pkg/restheadspec/security_hooks.go index b9365b6..e1c18f8 100644 --- a/pkg/restheadspec/security_hooks.go +++ b/pkg/restheadspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil { diff --git a/pkg/security/README.md b/pkg/security/README.md index 0fd0a69..7f11ace 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -1213,6 +1213,8 @@ The main changes: ## Documentation +- [Request transactions and RLS stamping](../common/TRANSACTIONS.md) + | File | Description | |------|-------------| | **QUICK_REFERENCE.md** | Quick reference guide with examples | diff --git a/pkg/security/provider.go b/pkg/security/provider.go index f7a2ee4..547e70b 100644 --- a/pkg/security/provider.go +++ b/pkg/security/provider.go @@ -132,6 +132,10 @@ type SecurityList struct { lastColPrune time.Time lastRowPrune time.Time + // txSettings stamps transaction-local settings at OnTxBegin (see txsettings.go). + txSettingsMu sync.RWMutex + txSettings TxSettingsFunc + // loads collapses concurrent provider calls for the same key (cold-cache stampede). loads singleflight.Group } diff --git a/pkg/security/txsettings.go b/pkg/security/txsettings.go new file mode 100644 index 0000000..ebec0cc --- /dev/null +++ b/pkg/security/txsettings.go @@ -0,0 +1,86 @@ +package security + +import ( + "encoding/hex" + "fmt" + "regexp" + "sort" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +// TxSettingsFunc returns the transaction-local settings (e.g. RLS GUCs such as +// "app.user_id") to stamp on a transaction. It runs once per transaction, at +// OnTxBegin, before any other SQL. Returning an error rolls the transaction back. +type TxSettingsFunc func(secCtx SecurityContext) (map[string]string, error) + +// 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_]*)+$`) + +// SetTxSettings sets the function that provides transaction-local settings for +// every transaction opened by a spec that registered its security hooks with this +// list. Pass nil to disable. May be called before or after RegisterSecurityHooks. +func (m *SecurityList) SetTxSettings(fn TxSettingsFunc) { + m.txSettingsMu.Lock() + defer m.txSettingsMu.Unlock() + m.txSettings = fn +} + +// TxSettings returns the configured TxSettingsFunc, or nil. +func (m *SecurityList) TxSettings() TxSettingsFunc { + m.txSettingsMu.RLock() + defer m.txSettingsMu.RUnlock() + return m.txSettings +} + +// StampTxSettings runs the list's TxSettingsFunc and applies the result to tx as +// transaction-local settings (set_config(name, value, true)). No-op when no +// function is configured or it returns no settings. tx must be the transaction +// itself, never the pool: the settings are lost on any other connection. +func StampTxSettings(secCtx SecurityContext, list *SecurityList, tx common.Database) error { + if list == nil { + return nil + } + fn := list.TxSettings() + if fn == nil { + return nil + } + settings, err := fn(secCtx) + if err != nil { + return err + } + return ApplyTxSettings(secCtx, tx, settings) +} + +// 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(secCtx SecurityContext, 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(secCtx.GetContext(), query); err != nil { + return fmt.Errorf("tx settings: set %s: %w", name, err) + } + } + return nil +} diff --git a/pkg/security/txsettings_test.go b/pkg/security/txsettings_test.go new file mode 100644 index 0000000..1b1bd37 --- /dev/null +++ b/pkg/security/txsettings_test.go @@ -0,0 +1,131 @@ +package security + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" +) + +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) + sc := &mockSecurityContext{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 ApplyTxSettings(sc, 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) + sc := &mockSecurityContext{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 ApplyTxSettings(sc, &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) + sc := &mockSecurityContext{ctx: context.Background()} + + for _, name := range []string{"user_id", "app.x'); DROP", "app..x", "app.x y"} { + if err := ApplyTxSettings(sc, pool, map[string]string{name: "1"}); err == nil { + t.Fatalf("name %q must be rejected", name) + } + } + if err := ApplyTxSettings(sc, pool, nil); err != nil { + t.Fatalf("empty settings must be a no-op: %v", err) + } + if err := ApplyTxSettings(sc, &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 } + +func TestStampTxSettingsUsesConfiguredFunc(t *testing.T) { + pool, mock := txSettingsDB(t) + sc := &mockSecurityContext{ctx: context.Background()} + list, err := NewSecurityList(&mockSecurityProvider{}) + if err != nil { + t.Fatal(err) + } + + // Nil list and unset func are no-ops (no SQL expected). + if err := StampTxSettings(sc, nil, pool); err != nil { + t.Fatal(err) + } + if err := StampTxSettings(sc, list, pool); err != nil { + t.Fatal(err) + } + + list.SetTxSettings(func(SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + if err := pool.RunInTransaction(context.Background(), func(tx common.Database) error { + return StampTxSettings(sc, list, tx) + }); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + + boom := errors.New("no tenant") + list.SetTxSettings(func(SecurityContext) (map[string]string, error) { return nil, boom }) + if err := StampTxSettings(sc, list, pool); !errors.Is(err, boom) { + t.Fatalf("func error must propagate, got %v", err) + } +} diff --git a/pkg/websocketspec/security_hooks.go b/pkg/websocketspec/security_hooks.go index f5596a0..41d5c21 100644 --- a/pkg/websocketspec/security_hooks.go +++ b/pkg/websocketspec/security_hooks.go @@ -10,6 +10,12 @@ import ( // RegisterSecurityHooks registers all security-related hooks with the handler func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList) { + // OnTxBegin: stamp transaction-local settings (e.g. RLS GUCs) before any SQL. + // Looked up per call so SetTxSettings may come after registration. + handler.Hooks().Register(OnTxBegin, func(hookCtx *HookContext) error { + return security.StampTxSettings(newSecurityContext(hookCtx), securityList, hookCtx.Tx) + }) + // Hook 0: BeforeHandle - enforce auth after model resolution handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error { if err := security.CheckModelAuthAllowed(newSecurityContext(hookCtx), hookCtx.Operation); err != nil {