mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-04 04:21:57 +00:00
feat(security): stamp transaction-local settings on OnTxBegin in all specs
This commit is contained in:
@@ -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 |
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user