diff --git a/cmd/testserver/main.go b/cmd/testserver/main.go index 5268637..5828064 100644 --- a/cmd/testserver/main.go +++ b/cmd/testserver/main.go @@ -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() diff --git a/go.mod b/go.mod index 8c72e95..e7f12ba 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/pkg/common/adapters/database/bun.go b/pkg/common/adapters/database/bun.go index de65185..9fcf939 100644 --- a/pkg/common/adapters/database/bun.go +++ b/pkg/common/adapters/database/bun.go @@ -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 } diff --git a/pkg/common/adapters/database/gorm.go b/pkg/common/adapters/database/gorm.go index 5e3f034..a681447 100644 --- a/pkg/common/adapters/database/gorm.go +++ b/pkg/common/adapters/database/gorm.go @@ -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 } diff --git a/pkg/common/adapters/database/pgsql.go b/pkg/common/adapters/database/pgsql.go index d0795dc..96b3b5c 100644 --- a/pkg/common/adapters/database/pgsql.go +++ b/pkg/common/adapters/database/pgsql.go @@ -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 } diff --git a/pkg/common/adapters/database/query_metrics.go b/pkg/common/adapters/database/query_metrics.go index 80ee6ce..861ac41 100644 --- a/pkg/common/adapters/database/query_metrics.go +++ b/pkg/common/adapters/database/query_metrics.go @@ -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 } diff --git a/pkg/config/config.go b/pkg/config/config.go index da230c6..9bf5338 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -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"` diff --git a/pkg/config/manager.go b/pkg/config/manager.go index 57614f6..d66fedc 100644 --- a/pkg/config/manager.go +++ b/pkg/config/manager.go @@ -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) diff --git a/pkg/dbmanager/metrics.go b/pkg/dbmanager/metrics.go index ef96bae..7219f5c 100644 --- a/pkg/dbmanager/metrics.go +++ b/pkg/dbmanager/metrics.go @@ -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)) diff --git a/pkg/dbtrace/README.md b/pkg/dbtrace/README.md new file mode 100644 index 0000000..5a204b0 --- /dev/null +++ b/pkg/dbtrace/README.md @@ -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 : 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) diff --git a/pkg/dbtrace/dbtrace.go b/pkg/dbtrace/dbtrace.go new file mode 100644 index 0000000..b7b8f7e --- /dev/null +++ b/pkg/dbtrace/dbtrace.go @@ -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)) + } + }) +} diff --git a/pkg/dbtrace/dbtrace_test.go b/pkg/dbtrace/dbtrace_test.go new file mode 100644 index 0000000..4e7ce5c --- /dev/null +++ b/pkg/dbtrace/dbtrace_test.go @@ -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) + } +} diff --git a/pkg/resolvespec/security_hooks.go b/pkg/resolvespec/security_hooks.go index c1148fa..450e33e 100644 --- a/pkg/resolvespec/security_hooks.go +++ b/pkg/resolvespec/security_hooks.go @@ -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) diff --git a/pkg/restheadspec/security_hooks.go b/pkg/restheadspec/security_hooks.go index 07bc555..b9365b6 100644 --- a/pkg/restheadspec/security_hooks.go +++ b/pkg/restheadspec/security_hooks.go @@ -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) diff --git a/pkg/security/hooks.go b/pkg/security/hooks.go index 414b3bf..23809d1 100644 --- a/pkg/security/hooks.go +++ b/pkg/security/hooks.go @@ -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 { diff --git a/pkg/security/keystore_database.go b/pkg/security/keystore_database.go index 70d250a..7dda23c 100644 --- a/pkg/security/keystore_database.go +++ b/pkg/security/keystore_database.go @@ -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 { diff --git a/pkg/security/provider.go b/pkg/security/provider.go index d3acd08..f7a2ee4 100644 --- a/pkg/security/provider.go +++ b/pkg/security/provider.go @@ -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() diff --git a/pkg/security/providers.go b/pkg/security/providers.go index 8f117f3..5c66809 100644 --- a/pkg/security/providers.go +++ b/pkg/security/providers.go @@ -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; diff --git a/pkg/security/query_mode.go b/pkg/security/query_mode.go index 9c040d1..0a32bb6 100644 --- a/pkg/security/query_mode.go +++ b/pkg/security/query_mode.go @@ -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. diff --git a/pkg/security/stampede_test.go b/pkg/security/stampede_test.go new file mode 100644 index 0000000..0f9ee7b --- /dev/null +++ b/pkg/security/stampede_test.go @@ -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") + } +} diff --git a/resolvespec-js/README.md b/resolvespec-js/README.md index 74ecffb..e07131a 100644 --- a/resolvespec-js/README.md +++ b/resolvespec-js/README.md @@ -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 diff --git a/resolvespec-js/src/__tests__/headerspec.test.ts b/resolvespec-js/src/__tests__/headerspec.test.ts index 9efd47c..021f40d 100644 --- a/resolvespec-js/src/__tests__/headerspec.test.ts +++ b/resolvespec-js/src/__tests__/headerspec.test.ts @@ -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'] }); diff --git a/resolvespec-js/src/common/types.ts b/resolvespec-js/src/common/types.ts index fd148d0..c81a108 100644 --- a/resolvespec-js/src/common/types.ts +++ b/resolvespec-js/src/common/types.ts @@ -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; + /** 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 { diff --git a/resolvespec-js/src/headerspec/client.ts b/resolvespec-js/src/headerspec/client.ts index e08a55b..e1eb6cb 100644 --- a/resolvespec-js/src/headerspec/client.ts +++ b/resolvespec-js/src/headerspec/client.ts @@ -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 { +export function buildHeaders(options: HeaderSpecOptions): Record { const headers: Record = {}; // Column selection @@ -79,6 +86,17 @@ export function buildHeaders(options: Options): Record { const op = mapOperatorToHeaderOp(filter.operator); const valueStr = formatFilterValue(filter); + const geoPrefix = geoFilterHeader(filter.operator); + if (geoPrefix) { + const payload: Record = { + 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 { // 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(); + 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 { 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> { 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> { 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> { const url = this.buildUrl(schema, entity, id); const optHeaders = options ? buildHeaders(options) : {};