Compare commits

..
5 Commits
Author SHA1 Message Date
Hein Puth (Warkanum) 9c4d916490 Merge pull request #22 from bitechdev/fix/db-connection-bursts
Tests / Unit Tests (push) Failing after 28s
Tests / Race Detector (push) Failing after 29s
Tests / Integration Tests (push) Failing after 28s
Build , Vet Test, and Lint / Build (push) Successful in 1m7s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 1m38s
Build , Vet Test, and Lint / Lint Code (push) Failing after 1m39s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 1m39s
Fix/db connection bursts
2026-09-30 21:45:20 +02:00
warkanum 3e327d0c78 fix(db): reduce per-request connection bursts and add dbtrace
* Throttle async session-activity writes to once per token per minute
* Add singleflight to session lookups, keystore validation and
  column/row security loads to stop cold-cache stampedes
* Preload security rules in BeforeHandle (restheadspec, resolvespec) so
  they no longer need a second connection while the read tx is open
* Add pkg/dbtrace: opt-in per-request DB call counting and pool logging
  (db_trace.* config, RESOLVESPEC_DB_TRACE_* env), wired into testserver
* Add tests for load dedup, activity throttle and dbtrace
2026-09-30 21:44:28 +02:00
warkanum 62cc14c02a feat(resolvespec-js): support extended restheadspec headers
Add HeaderSpecOptions, vector_search, X-Preload-Where, X-Expand,
custom SQL joins/or, spatial/vector filters, response format, flags
and X-Files to buildHeaders, types and README.
2026-09-30 21:37:41 +02:00
Hein 4c5dffc3d1 test(mqttspec): add tests for update behavior with empty strings 2026-09-30 17:15:55 +02:00
Hein 2898b335f8 feat(db): add ApplicationName to connection configuration
* Introduced ApplicationName field to identify clients in DSN
* Set default ApplicationName to "ResolveSpec"
* Updated tests for ApplicationName handling in DSN
2026-09-30 17:12:04 +02:00
31 changed files with 991 additions and 99 deletions
+6 -1
View File
@@ -9,6 +9,7 @@ import (
"github.com/bitechdev/ResolveSpec/pkg/config"
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/middleware"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
@@ -73,8 +74,12 @@ func main() {
queue := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10})
defer queue.Close()
// DB usage logging (off unless db_trace.enabled / RESOLVESPEC_DB_TRACE_ENABLED).
// Inside the queue so queue wait is not counted in request duration.
dbtrace.Configure(dbtrace.FromConfig(cfg.DBTrace))
// Setup routes using new SetupMuxRoutes function (without authentication)
resolvespec.SetupMuxRoutes(r, handler, middleware.Chain(queue.Middleware))
resolvespec.SetupMuxRoutes(r, handler, middleware.Chain(queue.Middleware, dbtrace.Middleware))
// Create server manager
mgr := server.NewManager()
+1 -1
View File
@@ -148,7 +148,7 @@ require (
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/mod v0.38.0 // indirect
golang.org/x/net v0.58.0 // indirect
golang.org/x/sync v0.22.0 // indirect
golang.org/x/sync v0.22.0
golang.org/x/text v0.41.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
+14 -12
View File
@@ -12,6 +12,7 @@ import (
"github.com/uptrace/bun"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
@@ -201,7 +202,7 @@ func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}
err = run()
}
}
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, operation, schema, entity, table, startedAt, err)
return &BunResult{result: result}, err
}
@@ -219,7 +220,7 @@ func (b *BunAdapter) Query(ctx context.Context, dest interface{}, query string,
err = b.getDB().NewRaw(query, args...).Scan(ctx, dest)
}
}
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, operation, schema, entity, table, startedAt, err)
return err
}
@@ -254,6 +255,7 @@ func (b *BunAdapter) RunInTransaction(ctx context.Context, fn func(common.Databa
err = logger.HandlePanic("BunAdapter.RunInTransaction", r)
}
}()
defer dbtrace.TxBegin(ctx)()
run := func() error {
return b.getDB().RunInTx(ctx, &sql.TxOptions{}, func(ctx context.Context, tx bun.Tx) error {
adapter := &BunTxAdapter{tx: tx, driverName: b.driverName, metricsEnabled: b.metricsEnabled}
@@ -1301,7 +1303,7 @@ func (b *BunSelectQuery) Scan(ctx context.Context, dest interface{}) (err error)
if r := recover(); r != nil {
err = logger.HandlePanic("BunSelectQuery.Scan", r)
}
recordQueryMetrics(b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
}()
if dest == nil {
err = fmt.Errorf("destination cannot be nil")
@@ -1347,7 +1349,7 @@ func (b *BunSelectQuery) ScanModel(ctx context.Context) (err error) {
logger.Error("Panic in BunSelectQuery.ScanModel: %v. %s. SQL: %s", r, modelInfo, sqlStr)
err = logger.HandlePanic("BunSelectQuery.ScanModel", r)
}
recordQueryMetrics(b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "SELECT", b.schema, b.entity, b.tableName, startedAt, err)
}()
if b.query.GetModel() == nil {
err = fmt.Errorf("model is nil")
@@ -1391,7 +1393,7 @@ func (b *BunSelectQuery) Count(ctx context.Context) (count int, err error) {
err = logger.HandlePanic("BunSelectQuery.Count", r)
count = 0
}
recordQueryMetrics(b.metricsEnabled, "COUNT", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "COUNT", b.schema, b.entity, b.tableName, startedAt, err)
}()
// If Model() was set, use bun's native Count() which works properly
if b.hasModel {
@@ -1425,7 +1427,7 @@ func (b *BunSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
err = logger.HandlePanic("BunSelectQuery.Exists", r)
exists = false
}
recordQueryMetrics(b.metricsEnabled, "EXISTS", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "EXISTS", b.schema, b.entity, b.tableName, startedAt, err)
}()
exists, err = b.query.Exists(ctx)
if err != nil {
@@ -1512,7 +1514,7 @@ func (b *BunInsertQuery) Exec(ctx context.Context) (res common.Result, err error
startedAt := time.Now()
b.prepareValues()
result, err := b.query.Exec(ctx)
recordQueryMetrics(b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
return &BunResult{result: result}, err
}
@@ -1525,7 +1527,7 @@ func (b *BunInsertQuery) Scan(ctx context.Context, dest interface{}) (err error)
startedAt := time.Now()
b.prepareValues()
err = b.query.Scan(ctx, dest)
recordQueryMetrics(b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "INSERT", b.schema, b.entity, b.tableName, startedAt, err)
return err
}
@@ -1622,7 +1624,7 @@ func (b *BunUpdateQuery) Exec(ctx context.Context) (res common.Result, err error
logger.Error("BunUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(b.metricsEnabled, "UPDATE", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "UPDATE", b.schema, b.entity, b.tableName, startedAt, err)
return &BunResult{result: result}, err
}
@@ -1674,7 +1676,7 @@ func (b *BunDeleteQuery) Exec(ctx context.Context) (res common.Result, err error
logger.Error("BunDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(b.metricsEnabled, "DELETE", b.schema, b.entity, b.tableName, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, "DELETE", b.schema, b.entity, b.tableName, startedAt, err)
return &BunResult{result: result}, err
}
@@ -1730,7 +1732,7 @@ func (b *BunTxAdapter) Exec(ctx context.Context, query string, args ...interface
startedAt := time.Now()
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
result, err := b.tx.ExecContext(ctx, query, args...)
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, operation, schema, entity, table, startedAt, err)
return &BunResult{result: result}, err
}
@@ -1738,7 +1740,7 @@ func (b *BunTxAdapter) Query(ctx context.Context, dest interface{}, query string
startedAt := time.Now()
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
err := b.tx.NewRaw(query, args...).Scan(ctx, dest)
recordQueryMetrics(b.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, b.metricsEnabled, operation, schema, entity, table, startedAt, err)
return err
}
+12 -10
View File
@@ -12,6 +12,7 @@ import (
"gorm.io/gorm/clause"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
@@ -151,7 +152,7 @@ func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{
result = run()
}
}
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, result.Error)
recordQueryMetrics(ctx, g.metricsEnabled, operation, schema, entity, table, startedAt, result.Error)
return &GormResult{result: result}, result.Error
}
@@ -172,7 +173,7 @@ func (g *GormAdapter) Query(ctx context.Context, dest interface{}, query string,
err = run()
}
}
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, g.metricsEnabled, operation, schema, entity, table, startedAt, err)
return err
}
@@ -206,6 +207,7 @@ func (g *GormAdapter) RunInTransaction(ctx context.Context, fn func(common.Datab
err = logger.HandlePanic("GormAdapter.RunInTransaction", r)
}
}()
defer dbtrace.TxBegin(ctx)()
run := func() error {
return g.getDB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
adapter := &GormAdapter{db: tx, dbFactory: g.dbFactory, driverName: g.driverName, metricsEnabled: g.metricsEnabled}
@@ -585,7 +587,7 @@ func (g *GormSelectQuery) Scan(ctx context.Context, dest interface{}) (err error
logger.Error("GormSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
recordQueryMetrics(ctx, g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
return err
}
@@ -616,7 +618,7 @@ func (g *GormSelectQuery) ScanModel(ctx context.Context) (err error) {
logger.Error("GormSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
recordQueryMetrics(ctx, g.metricsEnabled, "SELECT", g.schema, g.entity, g.tableName, startedAt, err)
return err
}
@@ -646,7 +648,7 @@ func (g *GormSelectQuery) Count(ctx context.Context) (count int, err error) {
logger.Error("GormSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "COUNT", g.schema, g.entity, g.tableName, startedAt, err)
recordQueryMetrics(ctx, g.metricsEnabled, "COUNT", g.schema, g.entity, g.tableName, startedAt, err)
return int(count64), err
}
@@ -676,7 +678,7 @@ func (g *GormSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
logger.Error("GormSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "EXISTS", g.schema, g.entity, g.tableName, startedAt, err)
recordQueryMetrics(ctx, g.metricsEnabled, "EXISTS", g.schema, g.entity, g.tableName, startedAt, err)
return count > 0, err
}
@@ -752,7 +754,7 @@ func (g *GormInsertQuery) Exec(ctx context.Context) (res common.Result, err erro
result = run()
}
}
recordQueryMetrics(g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
recordQueryMetrics(ctx, g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
return &GormResult{result: result}, result.Error
}
@@ -790,7 +792,7 @@ func (g *GormInsertQuery) Scan(ctx context.Context, dest interface{}) (err error
}
}
recordQueryMetrics(g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
recordQueryMetrics(ctx, g.metricsEnabled, "INSERT", g.schema, g.entity, g.tableName, startedAt, result.Error)
if result.Error != nil {
return result.Error
}
@@ -937,7 +939,7 @@ func (g *GormUpdateQuery) Exec(ctx context.Context) (res common.Result, err erro
logger.Error("GormUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "UPDATE", g.schema, g.entity, g.tableName, startedAt, result.Error)
recordQueryMetrics(ctx, g.metricsEnabled, "UPDATE", g.schema, g.entity, g.tableName, startedAt, result.Error)
return &GormResult{result: result}, result.Error
}
@@ -999,7 +1001,7 @@ func (g *GormDeleteQuery) Exec(ctx context.Context) (res common.Result, err erro
logger.Error("GormDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
}
recordQueryMetrics(g.metricsEnabled, "DELETE", g.schema, g.entity, g.tableName, startedAt, result.Error)
recordQueryMetrics(ctx, g.metricsEnabled, "DELETE", g.schema, g.entity, g.tableName, startedAt, result.Error)
return &GormResult{result: result}, result.Error
}
+19 -17
View File
@@ -11,6 +11,7 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
@@ -137,10 +138,10 @@ func (p *PgSQLAdapter) Exec(ctx context.Context, query string, args ...interface
}
if err != nil {
logger.Error("PgSQL Exec failed: %v", err)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return nil, common.WrapSQLError(err, query)
}
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
return &PgSQLResult{result: result}, nil
}
@@ -163,13 +164,13 @@ func (p *PgSQLAdapter) Query(ctx context.Context, dest interface{}, query string
}
if err != nil {
logger.Error("PgSQL Query failed: %v", err)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return common.WrapSQLError(err, query)
}
defer rows.Close()
err = scanRows(rows, dest)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return err
}
@@ -196,6 +197,7 @@ func (p *PgSQLAdapter) RunInTransaction(ctx context.Context, fn func(common.Data
}
}()
defer dbtrace.TxBegin(ctx)()
tx, err := p.getDB().BeginTx(ctx, nil)
if err != nil {
return err
@@ -510,20 +512,20 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
if err != nil {
logger.Error("PgSQL SELECT failed: %v", err)
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
return common.WrapSQLError(err, query)
}
defer rows.Close()
err = scanRows(rows, dest)
if err != nil {
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
return err
}
// Apply preloads that use separate queries
err = p.applySubqueryPreloads(ctx, dest)
recordQueryMetrics(p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "SELECT", p.schema, p.entity, p.tableName, startedAt, err)
return err
}
@@ -590,7 +592,7 @@ func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
logger.Error("PgSQL COUNT failed: %v", err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(p.metricsEnabled, "COUNT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "COUNT", p.schema, p.entity, p.tableName, startedAt, err)
return count, err
}
@@ -608,7 +610,7 @@ func (p *PgSQLSelectQuery) Exists(ctx context.Context) (exists bool, err error)
logger.Error("PgSQL EXISTS failed: %v", err)
err = common.WrapSQLError(err, sqlStr)
}
recordQueryMetrics(p.metricsEnabled, "EXISTS", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "EXISTS", p.schema, p.entity, p.tableName, startedAt, err)
return count > 0, err
}
@@ -667,7 +669,7 @@ func (p *PgSQLInsertQuery) Exec(ctx context.Context) (res common.Result, err err
if r := recover(); r != nil {
err = logger.HandlePanic("PgSQLInsertQuery.Exec", r)
}
recordQueryMetrics(p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
}()
if len(p.values) == 0 {
@@ -718,7 +720,7 @@ func (p *PgSQLInsertQuery) Scan(ctx context.Context, dest interface{}) (err erro
if r := recover(); r != nil {
err = logger.HandlePanic("PgSQLInsertQuery.Scan", r)
}
recordQueryMetrics(p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "INSERT", p.schema, p.entity, p.tableName, startedAt, err)
}()
if len(p.values) == 0 {
@@ -868,7 +870,7 @@ func (p *PgSQLUpdateQuery) Exec(ctx context.Context) (res common.Result, err err
if r := recover(); r != nil {
err = logger.HandlePanic("PgSQLUpdateQuery.Exec", r)
}
recordQueryMetrics(p.metricsEnabled, "UPDATE", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "UPDATE", p.schema, p.entity, p.tableName, startedAt, err)
}()
if len(p.sets) == 0 {
@@ -994,7 +996,7 @@ func (p *PgSQLDeleteQuery) Exec(ctx context.Context) (res common.Result, err err
if r := recover(); r != nil {
err = logger.HandlePanic("PgSQLDeleteQuery.Exec", r)
}
recordQueryMetrics(p.metricsEnabled, "DELETE", p.schema, p.entity, p.tableName, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, "DELETE", p.schema, p.entity, p.tableName, startedAt, err)
}()
query := fmt.Sprintf("DELETE FROM %s", p.tableName) //nolint:gosec // G201: table identifier is internal/validated; values use placeholders
@@ -1094,10 +1096,10 @@ func (p *PgSQLTxAdapter) Exec(ctx context.Context, query string, args ...interfa
result, err := p.tx.ExecContext(ctx, query, args...)
if err != nil {
logger.Error("PgSQL Tx Exec failed: %v", err)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return nil, common.WrapSQLError(err, query)
}
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, nil)
return &PgSQLResult{result: result}, nil
}
@@ -1108,13 +1110,13 @@ func (p *PgSQLTxAdapter) Query(ctx context.Context, dest interface{}, query stri
rows, err := p.tx.QueryContext(ctx, query, args...)
if err != nil {
logger.Error("PgSQL Tx Query failed: %v", err)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return common.WrapSQLError(err, query)
}
defer rows.Close()
err = scanRows(rows, dest)
recordQueryMetrics(p.metricsEnabled, operation, schema, entity, table, startedAt, err)
recordQueryMetrics(ctx, p.metricsEnabled, operation, schema, entity, table, startedAt, err)
return err
}
@@ -1,18 +1,21 @@
package database
import (
"context"
"reflect"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/metrics"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
)
const maxMetricFallbackEntityLength = 120
func recordQueryMetrics(enabled bool, operation, schema, entity, table string, startedAt time.Time, err error) {
func recordQueryMetrics(ctx context.Context, enabled bool, operation, schema, entity, table string, startedAt time.Time, err error) {
dbtrace.Query(ctx)
if !enabled {
return
}
+10
View File
@@ -16,6 +16,7 @@ type Config struct {
CORS CORSConfig `mapstructure:"cors"`
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
DBManager DBManagerConfig `mapstructure:"dbmanager"`
DBTrace DBTraceConfig `mapstructure:"db_trace"`
Paths PathsConfig `mapstructure:"paths"`
Extensions map[string]interface{} `mapstructure:"extensions"`
}
@@ -142,6 +143,15 @@ type CORSConfig struct {
MaxAge int `mapstructure:"max_age"`
}
// DBTraceConfig controls database usage logging (off by default).
// Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG.
type DBTraceConfig struct {
Enabled bool `mapstructure:"enabled"` // per-request DB call logging
MinCalls int `mapstructure:"min_calls"` // log requests with at least this many DB calls
MinDuration time.Duration `mapstructure:"min_duration"` // or that took at least this long (0 = off)
PoolLog bool `mapstructure:"pool_log"` // log pool changes on each metrics publish
}
// ErrorTrackingConfig holds error tracking configuration
type ErrorTrackingConfig struct {
Enabled bool `mapstructure:"enabled"`
+4
View File
@@ -55,6 +55,10 @@ type DBConnectionConfig struct {
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
Schema string `mapstructure:"schema"` // Default schema
// ApplicationName identifies this client to the server (postgres
// application_name, mssql app name, mongodb appName)
ApplicationName string `mapstructure:"application_name"`
// SQLite specific
FilePath string `mapstructure:"filepath"`
+7
View File
@@ -167,6 +167,7 @@ func (m *Manager) SetConfig(cfg *Config) error {
m.v.Set("cors", cfg.CORS)
m.v.Set("event_broker", cfg.EventBroker)
m.v.Set("dbmanager", cfg.DBManager)
m.v.Set("db_trace", cfg.DBTrace)
m.v.Set("paths", cfg.Paths)
m.v.Set("extensions", cfg.Extensions)
@@ -282,6 +283,12 @@ func setDefaults(v *viper.Viper) {
v.SetDefault("database.url", "")
// Database Manager defaults
// DB trace defaults (disabled)
v.SetDefault("db_trace.enabled", false)
v.SetDefault("db_trace.min_calls", 5)
v.SetDefault("db_trace.min_duration", "0s")
v.SetDefault("db_trace.pool_log", false)
v.SetDefault("dbmanager.default_connection", "default")
v.SetDefault("dbmanager.max_open_conns", 25)
v.SetDefault("dbmanager.max_idle_conns", 5)
+1
View File
@@ -271,6 +271,7 @@ db, _ := mgr.GetDefaultDatabase()
| `database` | string | Database name |
| `sslmode` | string | SSL mode (postgres/mssql): `disable`, `require`, etc. |
| `schema` | string | Default schema (postgres/mssql) |
| `application_name` | string | Client name shown by the server (postgres `application_name`, mssql `app name`, mongodb `appName`); defaults to `ResolveSpec` |
| `filepath` | string | File path (sqlite only) |
| `auth_source` | string | Auth source (mongodb) |
| `replica_set` | string | Replica set name (mongodb) |
+22
View File
@@ -72,6 +72,9 @@ type ManagerConfig struct {
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
}
// DefaultApplicationName is used when a connection does not set ApplicationName.
const DefaultApplicationName = "ResolveSpec"
// ConnectionConfig defines configuration for a single database connection
type ConnectionConfig struct {
// Name is the unique name of this connection
@@ -95,6 +98,10 @@ type ConnectionConfig struct {
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
Schema string `mapstructure:"schema"` // Default schema
// ApplicationName identifies this client to the server (postgres
// application_name, mssql app name, mongodb appName)
ApplicationName string `mapstructure:"application_name"`
// SQLite specific
FilePath string `mapstructure:"filepath"`
@@ -225,6 +232,10 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
cc.ConnMaxIdleTime = &idleTime
}
if cc.ApplicationName == "" {
cc.ApplicationName = DefaultApplicationName
}
// Default timeouts
if cc.ConnectTimeout == 0 {
cc.ConnectTimeout = 10 * time.Second
@@ -347,6 +358,9 @@ func (cc *ConnectionConfig) buildPostgresDSN() string {
if cc.Schema != "" {
q.Set("search_path", cc.Schema)
}
if cc.ApplicationName != "" {
q.Set("application_name", cc.ApplicationName)
}
u := url.URL{
Scheme: "postgres",
@@ -405,6 +419,9 @@ func (cc *ConnectionConfig) buildMSSQLDSN() string {
if cc.Schema != "" {
q.Set("schema", cc.Schema)
}
if cc.ApplicationName != "" {
q.Set("app name", cc.ApplicationName)
}
if cc.ConnectTimeout > 0 {
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
q.Set("connection timeout", sec)
@@ -437,6 +454,9 @@ func (cc *ConnectionConfig) buildMongoDSN() string {
if cc.ReadPreference != "" {
q.Set("readPreference", cc.ReadPreference)
}
if cc.ApplicationName != "" {
q.Set("appName", cc.ApplicationName)
}
u := url.URL{
Scheme: "mongodb",
@@ -480,6 +500,7 @@ func FromConfig(cfg config.DBManagerConfig) ManagerConfig {
Database: connCfg.Database,
SSLMode: connCfg.SSLMode,
Schema: connCfg.Schema,
ApplicationName: connCfg.ApplicationName,
FilePath: connCfg.FilePath,
AuthSource: connCfg.AuthSource,
ReplicaSet: connCfg.ReplicaSet,
@@ -519,6 +540,7 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
func (cc *ConnectionConfig) GetApplicationName() string { return cc.ApplicationName }
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
+30
View File
@@ -107,3 +107,33 @@ func TestSQLiteMemoryPoolPinned(t *testing.T) {
}
}
}
func TestApplicationNameInDSN(t *testing.T) {
cc := ConnectionConfig{Host: "h", Port: 1, Database: "d", ApplicationName: "my app&x=y"}
for name, c := range map[string]struct{ dsn, key string }{
"postgres": {cc.buildPostgresDSN(), "application_name"},
"mssql": {cc.buildMSSQLDSN(), "app name"},
"mongo": {cc.buildMongoDSN(), "appName"},
} {
u, err := url.Parse(c.dsn)
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if got := u.Query().Get(c.key); got != cc.ApplicationName {
t.Errorf("%s: %s = %q, want %q", name, c.key, got, cc.ApplicationName)
}
}
}
func TestApplicationNameDefault(t *testing.T) {
cc := ConnectionConfig{Type: DatabaseTypePostgreSQL, Host: "h", Database: "d"}
cc.ApplyDefaults(nil)
if cc.ApplicationName != "ResolveSpec" {
t.Errorf("ApplicationName = %q, want ResolveSpec", cc.ApplicationName)
}
cc = ConnectionConfig{Type: DatabaseTypePostgreSQL, Host: "h", Database: "d", ApplicationName: "x"}
cc.ApplyDefaults(nil)
if cc.ApplicationName != "x" {
t.Errorf("explicit ApplicationName overwritten: %q", cc.ApplicationName)
}
}
+12
View File
@@ -2,7 +2,10 @@ package dbmanager
import (
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
)
@@ -129,6 +132,15 @@ func (m *connectionManager) PublishMetrics() {
// sql.DBStats values are cumulative, so add only the growth since
// the last publish to keep these true counters.
prev := lastPublished.swap(name, connStats)
if dbtrace.PoolLogEnabled() {
opened := connStats.TotalOpened - prev.TotalOpened
waits := connStats.WaitCount - prev.WaitCount
if opened > 0 || waits > 0 || connStats.InUse > 0 {
logger.Info("dbtrace pool %s: open=%d in_use=%d idle=%d max=%d opened=+%d waits=+%d wait_time=+%s",
name, connStats.OpenConnections, connStats.InUse, connStats.Idle, connStats.MaxOpenConnections,
opened, waits, (connStats.WaitDuration - prev.WaitDuration).Round(time.Millisecond))
}
}
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
+6
View File
@@ -123,6 +123,12 @@ func buildPGXConfig(cfg ConnectionConfig) (*pgx.ConnConfig, error) {
}
cc.DialFunc = newDialFunc(cfg.GetConnectTimeout())
// Also applies to caller-supplied DSNs that do not set application_name.
if name := cfg.GetApplicationName(); name != "" {
if _, set := cc.RuntimeParams["application_name"]; !set {
cc.RuntimeParams["application_name"] = name
}
}
if cfg.GetQueryTimeout() > 0 {
if _, set := cc.RuntimeParams["statement_timeout"]; !set {
cc.RuntimeParams["statement_timeout"] = fmt.Sprintf("%d", cfg.GetQueryTimeout().Milliseconds())
+1
View File
@@ -59,6 +59,7 @@ type ConnectionConfig interface {
GetConnMaxLifetime() *time.Duration
GetConnMaxIdleTime() *time.Duration
GetReadPreference() string
GetApplicationName() string
GetRetryAttempts() int
GetRetryDelay() time.Duration
GetRetryMaxDelay() time.Duration
+28
View File
@@ -0,0 +1,28 @@
# dbtrace
Per-request DB call counting + pool logging. Off by default.
## Enable
| Config (`db_trace.*`) | Env | Meaning |
|---|---|---|
| `enabled` | `RESOLVESPEC_DB_TRACE_ENABLED` | per-request logging |
| `min_calls` (5) | `..._MIN_CALLS` | log if tx+pooled+raw >= N |
| `min_duration` (0) | `..._MIN_DURATION` | or request took >= D |
| `pool_log` | `..._POOL_LOG` | log pool stats on each dbmanager metrics publish |
Wire: `dbtrace.Configure(dbtrace.FromConfig(cfg.DBTrace))` and wrap handlers with `dbtrace.Middleware` (outside the auth middleware).
## Log fields
- `tx` transactions begun · `tx_queries` adapter queries inside `RunInTransaction` (share the tx connection)
- `pooled` adapter queries outside a tx (each takes a pool connection)
- `raw` direct `*sql.DB` calls, with kinds: `auth.session`, `auth.activity`, `security.column`, `security.row`, `probe.pg_proc`, `keystore.validate`
- Connections used ≈ `tx + pooled + raw`
## Pool log
`dbtrace pool <name>: open in_use idle max opened=+N waits=+N wait_time=+D` — `opened`/`waits` are deltas since last publish.
## Limits
- tx attribution is per request, assumes sequential use of a request's context
- `auth.activity` runs detached after the response: not in the request's log line
- `BeginTx`/`CommitTx` (manual tx) not counted; only `RunInTransaction`
- Raw counters cover the hot paths listed above only (not login/OAuth/passkey/TOTP)
+167
View File
@@ -0,0 +1,167 @@
// Package dbtrace counts database calls per request and logs the heavy ones.
// It is off until Configure enables it; the hot-path helpers are no-ops for
// contexts without a tracker.
//
// Connection use per request ≈ tx + pooled + raw (each pooled/raw call takes
// its own pool connection; calls inside RunInTransaction share the tx's).
package dbtrace
import (
"context"
"fmt"
"net/http"
"sort"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/bitechdev/ResolveSpec/pkg/config"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// Options controls tracing.
type Options struct {
Enabled bool
MinCalls int // log when tx+pooled+raw >= MinCalls
MinDuration time.Duration // or when the request took at least this (0 = off)
PoolLog bool // log pool deltas on each dbmanager metrics publish
}
// FromConfig maps the application config to Options.
func FromConfig(c config.DBTraceConfig) Options {
return Options{Enabled: c.Enabled, MinCalls: c.MinCalls, MinDuration: c.MinDuration, PoolLog: c.PoolLog}
}
var (
opts atomic.Pointer[Options]
)
// Configure sets the active options. Safe to call at any time.
func Configure(o Options) {
if o.MinCalls <= 0 {
o.MinCalls = 1
}
opts.Store(&o)
}
// Enabled reports whether per-request tracing is on.
func Enabled() bool {
o := opts.Load()
return o != nil && o.Enabled
}
// PoolLogEnabled reports whether pool logging is on.
func PoolLogEnabled() bool {
o := opts.Load()
return o != nil && o.PoolLog
}
// Tracker holds one request's counters.
type Tracker struct {
pooled atomic.Int32 // adapter calls outside a transaction
inTx atomic.Int32 // adapter calls inside RunInTransaction
tx atomic.Int32 // transactions begun
txDepth atomic.Int32
raw atomic.Int32 // direct *sql.DB calls (auth, security, keystore)
mu sync.Mutex
rawKind map[string]int
}
type ctxKey struct{}
// Start attaches a new Tracker to ctx.
func Start(ctx context.Context) (context.Context, *Tracker) {
t := &Tracker{}
return context.WithValue(ctx, ctxKey{}, t), t
}
// From returns the Tracker on ctx, or nil.
func From(ctx context.Context) *Tracker {
if ctx == nil {
return nil
}
t, _ := ctx.Value(ctxKey{}).(*Tracker)
return t
}
// Query counts one adapter query.
func Query(ctx context.Context) {
t := From(ctx)
if t == nil {
return
}
if t.txDepth.Load() > 0 {
t.inTx.Add(1)
} else {
t.pooled.Add(1)
}
}
// TxBegin counts a transaction and returns a func to call when it ends.
// Queries between the two are attributed to the transaction's connection.
func TxBegin(ctx context.Context) func() {
t := From(ctx)
if t == nil {
return func() {}
}
t.tx.Add(1)
t.txDepth.Add(1)
return func() { t.txDepth.Add(-1) }
}
// Raw counts one direct *sql.DB call, labelled by what issued it.
func Raw(ctx context.Context, kind string) {
t := From(ctx)
if t == nil {
return
}
t.raw.Add(1)
t.mu.Lock()
if t.rawKind == nil {
t.rawKind = make(map[string]int)
}
t.rawKind[kind]++
t.mu.Unlock()
}
func (t *Tracker) total() int { return int(t.tx.Load() + t.pooled.Load() + t.raw.Load()) }
// Summary renders the counters, e.g. "tx=1 tx_queries=3 pooled=2 raw=1 [auth.session=1]".
func (t *Tracker) Summary() string {
var b strings.Builder
fmt.Fprintf(&b, "tx=%d tx_queries=%d pooled=%d raw=%d", t.tx.Load(), t.inTx.Load(), t.pooled.Load(), t.raw.Load())
t.mu.Lock()
defer t.mu.Unlock()
if len(t.rawKind) > 0 {
kinds := make([]string, 0, len(t.rawKind))
for k, n := range t.rawKind {
kinds = append(kinds, fmt.Sprintf("%s=%d", k, n))
}
sort.Strings(kinds)
fmt.Fprintf(&b, " [%s]", strings.Join(kinds, " "))
}
return b.String()
}
// Middleware tracks every request and logs those over the configured
// thresholds once the handler returns. Options are read per request, so
// Configure can toggle it at runtime. Raw calls made by detached goroutines
// after the response (e.g. session activity) are not included in the log line.
func Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
o := opts.Load()
if o == nil || !o.Enabled {
next.ServeHTTP(w, r)
return
}
ctx, t := Start(r.Context())
start := time.Now()
next.ServeHTTP(w, r.WithContext(ctx))
elapsed := time.Since(start)
if t.total() >= o.MinCalls || (o.MinDuration > 0 && elapsed >= o.MinDuration) {
logger.Info("dbtrace: %s %s %s duration=%s", r.Method, r.URL.Path, t.Summary(), elapsed.Round(time.Millisecond))
}
})
}
+64
View File
@@ -0,0 +1,64 @@
package dbtrace
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestTrackerCounts(t *testing.T) {
ctx, tr := Start(context.Background())
Query(ctx) // pooled
end := TxBegin(ctx)
Query(ctx)
Query(ctx)
end()
Query(ctx) // pooled again
Raw(ctx, "auth.session")
Raw(ctx, "auth.session")
want := "tx=1 tx_queries=2 pooled=2 raw=2 [auth.session=2]"
if got := tr.Summary(); got != want {
t.Fatalf("summary = %q, want %q", got, want)
}
if tr.total() != 5 {
t.Fatalf("total = %d, want 5", tr.total())
}
}
func TestHelpersNoopWithoutTracker(t *testing.T) {
ctx := context.Background()
Query(ctx)
Raw(ctx, "x")
TxBegin(ctx)()
if From(ctx) != nil {
t.Fatal("unexpected tracker")
}
}
func TestMiddlewareDisabledAddsNoTracker(t *testing.T) {
Configure(Options{Enabled: false})
var seen *Tracker
h := Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { seen = From(r.Context()) }))
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/x", nil))
if seen != nil {
t.Fatal("tracker present while disabled")
}
}
func TestMiddlewareEnabledTracks(t *testing.T) {
Configure(Options{Enabled: true, MinCalls: 2})
defer Configure(Options{})
var summary string
h := Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
Raw(r.Context(), "a")
Query(r.Context())
summary = From(r.Context()).Summary()
}))
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/x", nil))
if !strings.Contains(summary, "pooled=1 raw=1") {
t.Fatalf("summary = %q", summary)
}
}
+62
View File
@@ -771,3 +771,65 @@ func TestHandler_HandleIncomingMessage_ValidMessage(t *testing.T) {
// Should not panic or error
handler.handleIncomingMessage("spec/test-client/request", payload)
}
func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) {
newHook := func(id string, data map[string]interface{}) *HookContext {
return &HookContext{
Context: context.Background(),
TableName: "users",
Model: &TestUser{},
ModelPtr: &TestUser{},
Schema: "public",
Entity: "users",
ID: id,
Data: data,
Options: &common.RequestOptions{},
}
}
seed := func(t *testing.T, db *gorm.DB) {
require.NoError(t, db.Create(&TestUser{ID: 1, Name: "Original", Email: "orig@example.com", Status: "active"}).Error)
}
t.Run("empty string clears only that field", func(t *testing.T) {
handler, db := setupTestHandler(t)
seed(t, db)
_, err := handler.update(newHook("1", map[string]interface{}{"name": ""}))
require.NoError(t, err)
var got TestUser
require.NoError(t, db.First(&got, 1).Error)
assert.Equal(t, "", got.Name)
assert.Equal(t, "orig@example.com", got.Email)
assert.Equal(t, "active", got.Status)
})
t.Run("absent keys are untouched", func(t *testing.T) {
handler, db := setupTestHandler(t)
seed(t, db)
_, err := handler.update(newHook("1", map[string]interface{}{"status": "inactive"}))
require.NoError(t, err)
var got TestUser
require.NoError(t, db.First(&got, 1).Error)
assert.Equal(t, "Original", got.Name)
assert.Equal(t, "orig@example.com", got.Email)
assert.Equal(t, "inactive", got.Status)
})
t.Run("disallowNulls skips null but applies empty string", func(t *testing.T) {
handler, db := setupTestHandler(t)
handler.SetDisallowNulls(true)
seed(t, db)
_, err := handler.update(newHook("1", map[string]interface{}{"name": nil, "status": ""}))
require.NoError(t, err)
var got TestUser
require.NoError(t, db.First(&got, 1).Error)
assert.Equal(t, "Original", got.Name)
assert.Equal(t, "", got.Status)
assert.Equal(t, "orig@example.com", got.Email)
})
}
+6
View File
@@ -22,6 +22,12 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return nil
})
// Hook 0b: BeforeHandle - preload security rules before the handler opens its
// transaction (BeforeRead runs inside it and would need a second connection).
handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error {
return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation)
})
// Hook 1: BeforeRead - Load security rules
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
+6
View File
@@ -21,6 +21,12 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
return nil
})
// Hook 0b: BeforeHandle - preload security rules before the handler opens its
// transaction (BeforeRead runs inside it and would need a second connection).
handler.Hooks().Register(BeforeHandle, func(hookCtx *HookContext) error {
return security.PreloadSecurityRules(newSecurityContext(hookCtx), securityList, hookCtx.Operation)
})
// Hook 1: BeforeRead - Load security rules
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
secCtx := newSecurityContext(hookCtx)
+13
View File
@@ -241,6 +241,19 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
return loadSecurityRules(secCtx, securityList)
}
// PreloadSecurityRules loads column/row security rules into the SecurityList
// cache for read operations. Call it from a BeforeHandle hook, i.e. before the
// handler opens its transaction, so the provider queries do not need a second
// pooled connection while the transaction holds one. Later LoadSecurityRules
// calls in the same request are then cache hits. Non-read operations and
// models with security disabled are skipped.
func PreloadSecurityRules(secCtx SecurityContext, securityList *SecurityList, operation string) error {
if operation != "read" || IsModelSecurityDisabled(secCtx) {
return nil
}
return loadSecurityRules(secCtx, securityList)
}
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
// This allows other packages to apply row-level security using the generic interface
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
+23
View File
@@ -12,6 +12,8 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/cache"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"golang.org/x/sync/singleflight"
)
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
@@ -51,6 +53,9 @@ type DatabaseKeyStore struct {
capability *dbCapability
cache *cache.Cache
cacheTTL time.Duration
// validateLoads collapses concurrent key lookups for the same key
validateLoads singleflight.Group
}
// NewDatabaseKeyStore creates a DatabaseKeyStore with optional configuration.
@@ -237,6 +242,24 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT
}
}
// Concurrent misses for the same key share one database lookup.
v, err, _ := ks.validateLoads.Do(cacheKey+"|"+string(keyType), func() (any, error) {
return ks.validateKeyLoad(ctx, hash, cacheKey, keyType)
})
if err != nil {
return nil, err
}
key, _ := v.(*UserKey)
if key == nil {
return nil, errors.New("invalid or expired key")
}
cp := *key
return &cp, nil
}
// validateKeyLoad validates against the database and fills the cache.
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
dbtrace.Raw(ctx, "keystore.validate")
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) {
key, err := ks.validateKeyDirect(ctx, hash, keyType)
if err != nil {
+12 -2
View File
@@ -15,6 +15,7 @@ import (
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
"golang.org/x/sync/singleflight"
)
type ColumnSecurity struct {
@@ -130,6 +131,9 @@ type SecurityList struct {
rowSecExpiry map[string]time.Time
lastColPrune time.Time
lastRowPrune time.Time
// loads collapses concurrent provider calls for the same key (cold-cache stampede).
loads singleflight.Group
}
const (
@@ -479,10 +483,13 @@ func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSch
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
colSecList, err := m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
v, err, _ := m.loads.Do("col:"+secKey, func() (any, error) {
return m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
})
if err != nil {
return fmt.Errorf("GetColumnSecurity failed: %v", err)
}
colSecList, _ := v.([]ColumnSecurity)
if colSecList == nil {
colSecList = make([]ColumnSecurity, 0)
}
@@ -552,10 +559,13 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchem
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
record, err := m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
v, err, _ := m.loads.Do("row:"+secKey, func() (any, error) {
return m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
})
if err != nil {
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
}
record, _ := v.(RowSecurity)
now := time.Now()
m.RowSecurityMutex.Lock()
+97 -43
View File
@@ -12,7 +12,9 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/cache"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"golang.org/x/sync/singleflight"
)
// Production-Ready Authenticators
@@ -71,6 +73,39 @@ const maxAuthTokens = 4
// sessionActivityTimeout bounds the detached last-activity update.
const sessionActivityTimeout = 5 * time.Second
// sessionActivityInterval is the minimum gap between last-activity writes for
// one session token. Requests inside it skip the write.
const sessionActivityInterval = time.Minute
// activityThrottle remembers when each token's activity was last written.
type activityThrottle struct {
mu sync.Mutex
last map[string]time.Time
lastPrune time.Time
}
// allow reports whether token is due an activity write, and if so records it.
func (t *activityThrottle) allow(token string, now time.Time) bool {
t.mu.Lock()
defer t.mu.Unlock()
if t.last == nil {
t.last = make(map[string]time.Time)
}
if prev, ok := t.last[token]; ok && now.Sub(prev) < sessionActivityInterval {
return false
}
t.last[token] = now
if now.Sub(t.lastPrune) > sessionActivityInterval {
t.lastPrune = now
for k, v := range t.last {
if now.Sub(v) >= sessionActivityInterval {
delete(t.last, k)
}
}
}
return true
}
// DatabaseAuthenticator provides session-based authentication with database storage
// All database operations go through stored procedures for security and consistency
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
@@ -94,6 +129,10 @@ type DatabaseAuthenticator struct {
// activityWG tracks in-flight asynchronous session activity updates
activityWG sync.WaitGroup
// activityLimit throttles those updates to one per token per interval
activityLimit activityThrottle
// sessionLoads collapses concurrent session lookups for the same token
sessionLoads singleflight.Group
// Cookie session support (optional, gated by enableCookieSession)
enableCookieSession bool
@@ -421,63 +460,75 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
cacheKey := fmt.Sprintf("auth:session:%s", token)
// Use cache.GetOrSet to get from cache or load from database
var userCtx UserContext
err := a.cache.GetOrSet(r.Context(), cacheKey, &userCtx, a.cacheTTL, func() (any, error) {
// This function is called only if cache miss
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
return a.sessionDirect(r.Context(), token)
}
// Concurrent misses for the same token share one database lookup.
v, err, _ := a.sessionLoads.Do(cacheKey, func() (any, error) {
var loaded UserContext
err := a.cache.GetOrSet(r.Context(), cacheKey, &loaded, a.cacheTTL, func() (any, error) {
// This function is called only if cache miss
dbtrace.Raw(r.Context(), "auth.session")
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
return a.sessionDirect(r.Context(), token)
}
var success bool
var errorMsg sql.NullString
var userJSON sql.NullString
var success bool
var errorMsg sql.NullString
var userJSON sql.NullString
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
}
return nil, fmt.Errorf("invalid or expired session")
}
if !userJSON.Valid {
return nil, fmt.Errorf("no user data in session")
}
// Parse UserContext
var user 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
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
return nil, err
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
}
return nil, fmt.Errorf("invalid or expired session")
}
if !userJSON.Valid {
return nil, fmt.Errorf("no user data in session")
}
// Parse UserContext
var user 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
return loaded, nil
})
if err != nil {
lastErr = err
continue // Try next token
}
userCtx, _ := v.(UserContext)
// Authentication succeeded with this token
// Update last activity timestamp asynchronously
activityCtx := userCtx
// Detach from the request (it is cancelled when the handler returns) but
// keep a deadline, and never let a panic here take the process down.
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
a.activityWG.Add(1)
go func(ctx context.Context, token string) {
defer a.activityWG.Done()
defer cancel()
defer logger.CatchPanic("updateSessionActivity")()
a.updateSessionActivity(ctx, token, &activityCtx)
}(detached, token)
if a.activityLimit.allow(token, time.Now()) {
activityCtx := userCtx
// Detach from the request (it is cancelled when the handler returns) but
// keep a deadline, and never let a panic here take the process down.
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
a.activityWG.Add(1)
go func(ctx context.Context, token string) {
defer a.activityWG.Done()
defer cancel()
defer logger.CatchPanic("updateSessionActivity")()
a.updateSessionActivity(ctx, token, &activityCtx)
}(detached, token)
}
return &userCtx, nil
}
@@ -513,6 +564,7 @@ func (a *DatabaseAuthenticator) ClearUserCache(userID int) error {
// updateSessionActivity updates the last activity timestamp for the session
func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) {
dbtrace.Raw(ctx, "auth.activity")
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.SessionUpdate) {
_ = a.updateSessionActivityDirect(ctx, sessionToken)
return
@@ -852,6 +904,7 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
return nil, ErrDirectModeUnsupported
}
dbtrace.Raw(ctx, "security.column")
var rules []ColumnSecurity
@@ -968,6 +1021,7 @@ func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userRe
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
return RowSecurity{}, ErrDirectModeUnsupported
}
dbtrace.Raw(ctx, "security.row")
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
// hand back the full *UserContext so non-DB providers can inspect claims;
+3
View File
@@ -10,6 +10,8 @@ import (
"strings"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
)
// QueryMode selects how a provider talks to the database: via the configured
@@ -59,6 +61,7 @@ func probeFunctionExists(ctx context.Context, db *sql.DB, procName string) bool
if db == nil {
return false
}
dbtrace.Raw(ctx, "probe.pg_proc")
var exists bool
defer func() {
// Guard against any unexpected panic from a misbehaving driver.
+40
View File
@@ -0,0 +1,40 @@
package security
import (
"context"
"sync"
"testing"
"time"
)
func TestConcurrentColdLoadsShareOneProviderCall(t *testing.T) {
p := &slowProvider{delay: 100 * time.Millisecond}
sl, _ := NewSecurityList(p)
var wg sync.WaitGroup
for i := 0; i < 10; i++ {
wg.Add(2)
go func() { defer wg.Done(); _ = sl.LoadColumnSecurity(context.Background(), 1, "s", "t", false) }()
go func() { defer wg.Done(); _, _ = sl.LoadRowSecurity(context.Background(), 1, "s", "t", false) }()
}
wg.Wait()
if got := p.calls.Load(); got != 2 {
t.Fatalf("provider calls = %d, want 2 (one column, one row)", got)
}
}
func TestActivityThrottle(t *testing.T) {
var th activityThrottle
now := time.Now()
if !th.allow("a", now) {
t.Fatal("first call must be allowed")
}
if th.allow("a", now.Add(sessionActivityInterval/2)) {
t.Fatal("call inside interval must be skipped")
}
if !th.allow("b", now) {
t.Fatal("other token must be allowed")
}
if !th.allow("a", now.Add(sessionActivityInterval)) {
t.Fatal("call after interval must be allowed")
}
}
+19
View File
@@ -106,6 +106,25 @@ await client.delete('public', 'users', '42');
| `X-Fetch-RowNumber` | `fetch_row_number` | string |
| `X-CQL-SEL-{col}` | `computedColumns` | expression |
| `X-Custom-SQL-W` | `customOperators` | SQL AND-joined |
| `X-Preload-Where` | `preload[].where` | applies to all preloads in `X-Preload`; differing wheres go to `X-Preload-{n}` + `X-Preload-{n}-Where` |
| `X-Expand` | `expand` | `Rel:col1,col2` pipe-separated (LEFT JOIN) |
| `X-Custom-SQL-Join` | `custom_sql_joins` | JOIN clauses, pipe-separated |
| `X-Custom-SQL-Or` | `custom_sql_or` | SQL OR-joined |
| `X-SearchCols` | `search_columns` | comma-separated |
| `X-AdvSQL-{col}` | `advanced_sql` | column -> SQL |
| `X-SpatialFilter-{col}` | `filters` (`st_dwithin`, `st_*`, `bbox`) | JSON `{op,value,logic}` |
| `X-VectorFilter-{col}` | `filters` (`l2_within`, `cosine_within`, `ip_within`) | JSON `{op,value,logic}` |
| `X-Vector-Search-{col}` / `-Vector` / `-As` / `-Dir` | `vector_search` | metric / JSON array / alias / asc\|desc |
| `X-Clean-JSON` | `clean_json` | bool |
| `X-Distinct` | `distinct` | bool |
| `X-SkipCount` / `X-SkipCache` | `skip_count` / `skip_cache` | bool |
| `X-PKRow` | `pk_row` | string |
| `X-SimpleApi` / `X-DetailApi` / `X-Syncfusion` | `response_format` | `simple` \| `detail` \| `syncfusion` |
| `X-Single-Record-As-Object` | `single_record_as_object` | bool (server default true) |
| `X-Transaction-Atomic` | `atomic_transaction` | bool |
| `X-Files` | `xfiles` | JSON, sent as `ZIP_` base64 |
Extended fields live on `HeaderSpecOptions` (extends `Options`); `vector_search` is on `Options`.
### Utility Functions
@@ -2,6 +2,101 @@ import { describe, it, expect, vi, beforeEach } from 'vitest';
import { buildHeaders, encodeHeaderValue, decodeHeaderValue, HeaderSpecClient, getHeaderSpecClient } from '../headerspec/client';
import type { Options, ClientConfig, APIResponse } from '../common/types';
describe('buildHeaders (extended restheadspec options)', () => {
it('should set X-Preload-Where when all preloads share one where', () => {
const h = buildHeaders({
preload: [
{ relation: 'Items', columns: ['id'], where: 'active = true' },
{ relation: 'Tags', where: 'active = true' },
],
});
expect(h['X-Preload']).toBe('Items:id|Tags');
expect(h['X-Preload-Where']).toBe('active = true');
});
it('should use numbered headers for mixed where clauses', () => {
const h = buildHeaders({
preload: [
{ relation: 'Items', where: 'a = 1' },
{ relation: 'Category' },
{ relation: 'Tags', where: 'b = 2' },
],
});
expect(h['X-Preload']).toBe('Category');
expect(h['X-Preload-Where']).toBeUndefined();
expect(h['X-Preload-1']).toBe('Items');
expect(h['X-Preload-1-Where']).toBe('a = 1');
expect(h['X-Preload-2']).toBe('Tags');
expect(h['X-Preload-2-Where']).toBe('b = 2');
});
it('should set expand, joins, or-sql, search cols, advsql', () => {
const h = buildHeaders({
expand: [{ relation: 'Dept', columns: ['id', 'name'] }, { relation: 'Role' }],
custom_sql_joins: ['LEFT JOIN a ON a.id = b.id', 'INNER JOIN c ON c.id = b.cid'],
custom_sql_or: ['x = 1', 'y = 2'],
search_columns: ['name', 'email'],
advanced_sql: { total: 'a + b' },
});
expect(h['X-Expand']).toBe('Dept:id,name|Role');
expect(h['X-Custom-SQL-Join']).toBe('LEFT JOIN a ON a.id = b.id|INNER JOIN c ON c.id = b.cid');
expect(h['X-Custom-SQL-Or']).toBe('x = 1 OR y = 2');
expect(h['X-SearchCols']).toBe('name,email');
expect(h['X-AdvSQL-total']).toBe('a + b');
});
it('should set boolean flags, pk row and response format', () => {
const h = buildHeaders({
clean_json: true,
distinct: true,
skip_count: true,
skip_cache: false,
atomic_transaction: true,
single_record_as_object: false,
pk_row: '42',
response_format: 'detail',
});
expect(h['X-Clean-JSON']).toBe('true');
expect(h['X-Distinct']).toBe('true');
expect(h['X-SkipCount']).toBe('true');
expect(h['X-SkipCache']).toBe('false');
expect(h['X-Transaction-Atomic']).toBe('true');
expect(h['X-Single-Record-As-Object']).toBe('false');
expect(h['X-PKRow']).toBe('42');
expect(h['X-DetailApi']).toBe('true');
});
it('should set spatial and vector filters as JSON', () => {
const h = buildHeaders({
filters: [
{ column: 'geom', operator: 'st_dwithin', value: { geom: 'POINT(0 0)', distance: 5 }, logic_operator: 'OR' },
{ column: 'emb', operator: 'cosine_within', value: { vector: [1, 2], distance: 0.3 } },
],
});
expect(JSON.parse(h['X-SpatialFilter-geom'])).toEqual({
op: 'st_dwithin', value: { geom: 'POINT(0 0)', distance: 5 }, logic: 'or',
});
expect(JSON.parse(h['X-VectorFilter-emb']).op).toBe('cosine_within');
});
it('should set vector search headers', () => {
const h = buildHeaders({
vector_search: { column: 'emb', vector: [0.1, 0.2], metric: 'cosine', as: 'dist', direction: 'desc' },
});
expect(h['X-Vector-Search-emb']).toBe('cosine');
expect(h['X-Vector-Search-Vector']).toBe('[0.1,0.2]');
expect(h['X-Vector-Search-As']).toBe('dist');
expect(h['X-Vector-Search-Dir']).toBe('desc');
});
it('should encode X-Files as ZIP_ JSON', () => {
const xf = { tablename: 'users', prefix: 'USR', limit: 10 };
const h = buildHeaders({ xfiles: xf });
expect(h['X-Files'].startsWith('ZIP_')).toBe(true);
expect(JSON.parse(decodeHeaderValue(h['X-Files']))).toEqual(xf);
});
});
describe('buildHeaders', () => {
it('should set X-Select-Fields for columns', () => {
const h = buildHeaders({ columns: ['id', 'name', 'email'] });
+86 -1
View File
@@ -5,7 +5,11 @@ export type Operator =
| 'like' | 'ilike' | 'in'
| 'contains' | 'startswith' | 'endswith'
| 'between' | 'between_inclusive'
| 'is_null' | 'is_not_null';
| 'is_null' | 'is_not_null'
// PostGIS spatial (sent via X-SpatialFilter-{col})
| 'st_dwithin' | 'bbox'
// pgvector similarity (sent via X-VectorFilter-{col})
| 'l2_within' | 'cosine_within' | 'ip_within';
export type Operation = 'read' | 'create' | 'update' | 'delete';
export type SortDirection = 'asc' | 'desc' | 'ASC' | 'DESC';
@@ -61,6 +65,54 @@ export interface ComputedColumn {
expression: string;
}
export type VectorMetric = 'l2' | 'cosine' | 'ip';
export type ResponseFormat = 'simple' | 'detail' | 'syncfusion';
/** pgvector KNN search: order by distance between `column` and `vector`. */
export interface VectorSearchOption {
column: string;
vector: number[];
metric?: VectorMetric;
/** Distance column alias. Default `_distance` */
as?: string;
direction?: 'asc' | 'desc';
}
/** LEFT JOIN expansion of a relation (X-Expand). */
export interface ExpandOption {
relation: string;
columns?: string[];
}
/** X-Files configuration (Go restheadspec XFiles). Sent as a single JSON header. */
export interface XFiles {
tablename?: string;
schema?: string;
primarykey?: string;
foreignkey?: string;
relatedkey?: string;
sort?: string[];
prefix?: string;
editable?: boolean;
recursive?: boolean;
expand?: boolean;
rownumber?: boolean;
skipcount?: boolean;
offset?: number;
limit?: number;
columns?: string[];
omit_columns?: string[];
cql_columns?: string[];
sql_joins?: string[];
sql_or?: string[];
sql_and?: string[];
parenttables?: XFiles[];
childtables?: XFiles[];
filter_fields?: { field: string; value: string; operator: string }[];
cursor_forward?: string;
cursor_backward?: string;
}
export interface Options {
preload?: PreloadOption[];
columns?: string[];
@@ -75,6 +127,39 @@ export interface Options {
cursor_forward?: string;
cursor_backward?: string;
fetch_row_number?: string;
vector_search?: VectorSearchOption;
}
/** Options only available to the header-based (restheadspec) protocol. */
export interface HeaderSpecOptions extends Options {
/** X-Expand: LEFT JOIN relations */
expand?: ExpandOption[];
/** X-Custom-SQL-Join: raw JOIN clauses */
custom_sql_joins?: string[];
/** X-Custom-SQL-Or: raw SQL, OR-combined */
custom_sql_or?: string[];
/** X-SearchCols: columns for multi-column search */
search_columns?: string[];
/** X-AdvSQL-{col}: column -> SQL expression */
advanced_sql?: Record<string, string>;
/** X-Clean-JSON */
clean_json?: boolean;
/** X-Distinct */
distinct?: boolean;
/** X-SkipCount: skip total count query */
skip_count?: boolean;
/** X-SkipCache */
skip_cache?: boolean;
/** X-PKRow: primary key value of a row to fetch */
pk_row?: string;
/** X-SimpleApi / X-DetailApi / X-Syncfusion */
response_format?: ResponseFormat;
/** X-Single-Record-As-Object (server default true) */
single_record_as_object?: boolean;
/** X-Transaction-Atomic */
atomic_transaction?: boolean;
/** X-Files: single JSON configuration */
xfiles?: XFiles;
}
export interface RequestBody {
+121 -11
View File
@@ -5,7 +5,7 @@ import type {
ClientConfig,
CustomOperator,
FilterOption,
Options,
HeaderSpecOptions,
PreloadOption,
SortOption,
} from "../common/types";
@@ -59,8 +59,15 @@ function decodeBase64(str: string): string {
* - X-Fetch-RowNumber: row number fetch
* - X-CQL-SEL-{col}: computed columns
* - X-Custom-SQL-W: custom operators (AND)
* - X-Preload-Where: where for X-Preload (extra where groups use X-Preload-{n}[-Where])
* - X-SpatialFilter-{col} / X-VectorFilter-{col}: JSON {op,value,logic}
* - X-Vector-Search-{col|vector|as|dir}: pgvector KNN
* - X-Expand, X-Custom-SQL-Join, X-Custom-SQL-Or, X-SearchCols, X-AdvSQL-{col}
* - X-Clean-JSON, X-Distinct, X-SkipCount, X-SkipCache, X-PKRow
* - X-SimpleApi / X-DetailApi / X-Syncfusion, X-Single-Record-As-Object
* - X-Transaction-Atomic, X-Files
*/
export function buildHeaders(options: Options): Record<string, string> {
export function buildHeaders(options: HeaderSpecOptions): Record<string, string> {
const headers: Record<string, string> = {};
// Column selection
@@ -79,6 +86,17 @@ export function buildHeaders(options: Options): Record<string, string> {
const op = mapOperatorToHeaderOp(filter.operator);
const valueStr = formatFilterValue(filter);
const geoPrefix = geoFilterHeader(filter.operator);
if (geoPrefix) {
const payload: Record<string, unknown> = {
op: filter.operator,
value: filter.value,
};
if (logicOp === "OR") payload.logic = "or";
headers[`${geoPrefix}${filter.column}`] = JSON.stringify(payload);
continue;
}
if (filter.operator === "eq" && logicOp === "AND") {
// Simple field filter shorthand
headers[`X-FieldFilter-${filter.column}`] = valueStr;
@@ -117,13 +135,94 @@ export function buildHeaders(options: Options): Record<string, string> {
// Preload
if (options.preload?.length) {
const parts = options.preload.map((p: PreloadOption) => {
if (p.columns?.length) {
return `${p.relation}:${p.columns.join(",")}`;
// Go applies X-Preload-Where to every preload in the matching X-Preload header,
// so preloads are grouped by where clause.
const groups = new Map<string, string[]>();
for (const p of options.preload) {
const spec = p.columns?.length
? `${p.relation}:${p.columns.join(",")}`
: p.relation;
const where = p.where ?? "";
groups.set(where, [...(groups.get(where) ?? []), spec]);
}
let n = 0;
for (const [where, specs] of groups) {
if (!where) {
headers["X-Preload"] = specs.join("|");
} else if (!groups.has("") && n === 0) {
// X-Preload-Where would also apply to a where-less X-Preload, so only use it alone
headers["X-Preload"] = specs.join("|");
headers["X-Preload-Where"] = where;
n++;
} else {
n++;
headers[`X-Preload-${n}`] = specs.join("|");
headers[`X-Preload-${n}-Where`] = where;
}
return p.relation;
});
headers["X-Preload"] = parts.join("|");
}
}
// Expand (LEFT JOIN)
if (options.expand?.length) {
headers["X-Expand"] = options.expand
.map((e) =>
e.columns?.length ? `${e.relation}:${e.columns.join(",")}` : e.relation,
)
.join("|");
}
if (options.custom_sql_joins?.length) {
headers["X-Custom-SQL-Join"] = options.custom_sql_joins.join("|");
}
if (options.custom_sql_or?.length) {
headers["X-Custom-SQL-Or"] = options.custom_sql_or.join(" OR ");
}
if (options.search_columns?.length) {
headers["X-SearchCols"] = options.search_columns.join(",");
}
if (options.advanced_sql) {
for (const [col, sql] of Object.entries(options.advanced_sql)) {
headers[`X-AdvSQL-${col}`] = sql;
}
}
// pgvector KNN search
if (options.vector_search) {
const vs = options.vector_search;
headers[`X-Vector-Search-${vs.column}`] = vs.metric ?? "l2";
headers["X-Vector-Search-Vector"] = JSON.stringify(vs.vector);
if (vs.as) headers["X-Vector-Search-As"] = vs.as;
if (vs.direction) headers["X-Vector-Search-Dir"] = vs.direction;
}
// Flags
const flags: [string, boolean | undefined][] = [
["X-Clean-JSON", options.clean_json],
["X-Distinct", options.distinct],
["X-SkipCount", options.skip_count],
["X-SkipCache", options.skip_cache],
["X-Transaction-Atomic", options.atomic_transaction],
["X-Single-Record-As-Object", options.single_record_as_object],
];
for (const [name, val] of flags) {
if (val !== undefined) headers[name] = String(val);
}
if (options.pk_row) {
headers["X-PKRow"] = options.pk_row;
}
if (options.response_format) {
const formatHeaders = {
simple: "X-SimpleApi",
detail: "X-DetailApi",
syncfusion: "X-Syncfusion",
} as const;
headers[formatHeaders[options.response_format]] = "true";
}
if (options.xfiles) {
headers["X-Files"] = encodeHeaderValue(JSON.stringify(options.xfiles));
}
// Fetch row number
@@ -149,6 +248,17 @@ export function buildHeaders(options: Options): Record<string, string> {
return headers;
}
const VECTOR_OPS = new Set(["l2_within", "cosine_within", "ip_within"]);
function geoFilterHeader(operator: string): string | null {
const op = operator.toLowerCase();
if (VECTOR_OPS.has(op) || op.endsWith("_within")) return "X-VectorFilter-";
if (op.startsWith("st_") || op === "bbox" || op === "&&") {
return "X-SpatialFilter-";
}
return null;
}
function mapOperatorToHeaderOp(operator: string): string {
switch (operator) {
case "eq":
@@ -280,7 +390,7 @@ export class HeaderSpecClient {
schema: string,
entity: string,
id?: string,
options?: Options,
options?: HeaderSpecOptions,
): Promise<APIResponse<T>> {
const url = this.buildUrl(schema, entity, id);
const optHeaders = options ? buildHeaders(options) : {};
@@ -294,7 +404,7 @@ export class HeaderSpecClient {
schema: string,
entity: string,
data: any,
options?: Options,
options?: HeaderSpecOptions,
): Promise<APIResponse<T>> {
const url = this.buildUrl(schema, entity);
const optHeaders = options ? buildHeaders(options) : {};
@@ -310,7 +420,7 @@ export class HeaderSpecClient {
entity: string,
id: string,
data: any,
options?: Options,
options?: HeaderSpecOptions,
): Promise<APIResponse<T>> {
const url = this.buildUrl(schema, entity, id);
const optHeaders = options ? buildHeaders(options) : {};