mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 13:56:29 +00:00
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
This commit is contained in:
@@ -9,6 +9,7 @@ import (
|
|||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/middleware"
|
"github.com/bitechdev/ResolveSpec/pkg/middleware"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
@@ -73,8 +74,12 @@ func main() {
|
|||||||
queue := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10})
|
queue := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10})
|
||||||
defer queue.Close()
|
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)
|
// 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
|
// Create server manager
|
||||||
mgr := server.NewManager()
|
mgr := server.NewManager()
|
||||||
|
|||||||
@@ -148,7 +148,7 @@ require (
|
|||||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||||
golang.org/x/mod v0.38.0 // indirect
|
golang.org/x/mod v0.38.0 // indirect
|
||||||
golang.org/x/net v0.58.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
|
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/api v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"github.com/uptrace/bun"
|
"github.com/uptrace/bun"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
@@ -201,7 +202,7 @@ func (b *BunAdapter) Exec(ctx context.Context, query string, args ...interface{}
|
|||||||
err = run()
|
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
|
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)
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -254,6 +255,7 @@ func (b *BunAdapter) RunInTransaction(ctx context.Context, fn func(common.Databa
|
|||||||
err = logger.HandlePanic("BunAdapter.RunInTransaction", r)
|
err = logger.HandlePanic("BunAdapter.RunInTransaction", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
defer dbtrace.TxBegin(ctx)()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return b.getDB().RunInTx(ctx, &sql.TxOptions{}, func(ctx context.Context, tx bun.Tx) 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}
|
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 {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("BunSelectQuery.Scan", r)
|
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 {
|
if dest == nil {
|
||||||
err = fmt.Errorf("destination cannot be 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)
|
logger.Error("Panic in BunSelectQuery.ScanModel: %v. %s. SQL: %s", r, modelInfo, sqlStr)
|
||||||
err = logger.HandlePanic("BunSelectQuery.ScanModel", r)
|
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 {
|
if b.query.GetModel() == nil {
|
||||||
err = fmt.Errorf("model is 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)
|
err = logger.HandlePanic("BunSelectQuery.Count", r)
|
||||||
count = 0
|
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 Model() was set, use bun's native Count() which works properly
|
||||||
if b.hasModel {
|
if b.hasModel {
|
||||||
@@ -1425,7 +1427,7 @@ func (b *BunSelectQuery) Exists(ctx context.Context) (exists bool, err error) {
|
|||||||
err = logger.HandlePanic("BunSelectQuery.Exists", r)
|
err = logger.HandlePanic("BunSelectQuery.Exists", r)
|
||||||
exists = false
|
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)
|
exists, err = b.query.Exists(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1512,7 +1514,7 @@ func (b *BunInsertQuery) Exec(ctx context.Context) (res common.Result, err error
|
|||||||
startedAt := time.Now()
|
startedAt := time.Now()
|
||||||
b.prepareValues()
|
b.prepareValues()
|
||||||
result, err := b.query.Exec(ctx)
|
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
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1525,7 +1527,7 @@ func (b *BunInsertQuery) Scan(ctx context.Context, dest interface{}) (err error)
|
|||||||
startedAt := time.Now()
|
startedAt := time.Now()
|
||||||
b.prepareValues()
|
b.prepareValues()
|
||||||
err = b.query.Scan(ctx, dest)
|
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
|
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)
|
logger.Error("BunUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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)
|
logger.Error("BunDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1730,7 +1732,7 @@ func (b *BunTxAdapter) Exec(ctx context.Context, query string, args ...interface
|
|||||||
startedAt := time.Now()
|
startedAt := time.Now()
|
||||||
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
result, err := b.tx.ExecContext(ctx, query, args...)
|
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
|
return &BunResult{result: result}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1738,7 +1740,7 @@ func (b *BunTxAdapter) Query(ctx context.Context, dest interface{}, query string
|
|||||||
startedAt := time.Now()
|
startedAt := time.Now()
|
||||||
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
operation, schema, entity, table := metricTargetFromRawQuery(query, b.driverName)
|
||||||
err := b.tx.NewRaw(query, args...).Scan(ctx, dest)
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"gorm.io/gorm/clause"
|
"gorm.io/gorm/clause"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
@@ -151,7 +152,7 @@ func (g *GormAdapter) Exec(ctx context.Context, query string, args ...interface{
|
|||||||
result = run()
|
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
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -172,7 +173,7 @@ func (g *GormAdapter) Query(ctx context.Context, dest interface{}, query string,
|
|||||||
err = run()
|
err = run()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
recordQueryMetrics(g.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
recordQueryMetrics(ctx, g.metricsEnabled, operation, schema, entity, table, startedAt, err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -206,6 +207,7 @@ func (g *GormAdapter) RunInTransaction(ctx context.Context, fn func(common.Datab
|
|||||||
err = logger.HandlePanic("GormAdapter.RunInTransaction", r)
|
err = logger.HandlePanic("GormAdapter.RunInTransaction", r)
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
defer dbtrace.TxBegin(ctx)()
|
||||||
run := func() error {
|
run := func() error {
|
||||||
return g.getDB().WithContext(ctx).Transaction(func(tx *gorm.DB) 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}
|
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)
|
logger.Error("GormSelectQuery.Scan failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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)
|
logger.Error("GormSelectQuery.ScanModel failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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)
|
logger.Error("GormSelectQuery.Count failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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)
|
logger.Error("GormSelectQuery.Exists failed. SQL: %s. Error: %v", sqlStr, err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
return count > 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -752,7 +754,7 @@ func (g *GormInsertQuery) Exec(ctx context.Context) (res common.Result, err erro
|
|||||||
result = run()
|
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
|
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 {
|
if result.Error != nil {
|
||||||
return result.Error
|
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)
|
logger.Error("GormUpdateQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
||||||
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
|
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
|
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)
|
logger.Error("GormDeleteQuery.Exec failed. SQL: %s. Error: %v", sqlStr, result.Error)
|
||||||
return &GormResult{result: result}, common.WrapSQLError(result.Error, sqlStr)
|
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
|
return &GormResult{result: result}, result.Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
@@ -137,10 +138,10 @@ func (p *PgSQLAdapter) Exec(ctx context.Context, query string, args ...interface
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Exec failed: %v", err)
|
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)
|
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
|
return &PgSQLResult{result: result}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -163,13 +164,13 @@ func (p *PgSQLAdapter) Query(ctx context.Context, dest interface{}, query string
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Query failed: %v", err)
|
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)
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
err = scanRows(rows, dest)
|
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
|
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)
|
tx, err := p.getDB().BeginTx(ctx, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -510,20 +512,20 @@ func (p *PgSQLSelectQuery) Scan(ctx context.Context, dest interface{}) (err erro
|
|||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL SELECT failed: %v", err)
|
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)
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
err = scanRows(rows, dest)
|
err = scanRows(rows, dest)
|
||||||
if err != nil {
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply preloads that use separate queries
|
// Apply preloads that use separate queries
|
||||||
err = p.applySubqueryPreloads(ctx, dest)
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -590,7 +592,7 @@ func (p *PgSQLSelectQuery) Count(ctx context.Context) (count int, err error) {
|
|||||||
logger.Error("PgSQL COUNT failed: %v", err)
|
logger.Error("PgSQL COUNT failed: %v", err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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)
|
logger.Error("PgSQL EXISTS failed: %v", err)
|
||||||
err = common.WrapSQLError(err, sqlStr)
|
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
|
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 {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLInsertQuery.Exec", r)
|
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 {
|
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 {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLInsertQuery.Scan", r)
|
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 {
|
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 {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLUpdateQuery.Exec", r)
|
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 {
|
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 {
|
if r := recover(); r != nil {
|
||||||
err = logger.HandlePanic("PgSQLDeleteQuery.Exec", r)
|
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
|
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...)
|
result, err := p.tx.ExecContext(ctx, query, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Tx Exec failed: %v", err)
|
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)
|
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
|
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...)
|
rows, err := p.tx.QueryContext(ctx, query, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("PgSQL Tx Query failed: %v", err)
|
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)
|
return common.WrapSQLError(err, query)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
defer rows.Close()
|
||||||
|
|
||||||
err = scanRows(rows, dest)
|
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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,18 +1,21 @@
|
|||||||
package database
|
package database
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"reflect"
|
"reflect"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
)
|
)
|
||||||
|
|
||||||
const maxMetricFallbackEntityLength = 120
|
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 {
|
if !enabled {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ type Config struct {
|
|||||||
CORS CORSConfig `mapstructure:"cors"`
|
CORS CORSConfig `mapstructure:"cors"`
|
||||||
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||||
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
||||||
|
DBTrace DBTraceConfig `mapstructure:"db_trace"`
|
||||||
Paths PathsConfig `mapstructure:"paths"`
|
Paths PathsConfig `mapstructure:"paths"`
|
||||||
Extensions map[string]interface{} `mapstructure:"extensions"`
|
Extensions map[string]interface{} `mapstructure:"extensions"`
|
||||||
}
|
}
|
||||||
@@ -142,6 +143,15 @@ type CORSConfig struct {
|
|||||||
MaxAge int `mapstructure:"max_age"`
|
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
|
// ErrorTrackingConfig holds error tracking configuration
|
||||||
type ErrorTrackingConfig struct {
|
type ErrorTrackingConfig struct {
|
||||||
Enabled bool `mapstructure:"enabled"`
|
Enabled bool `mapstructure:"enabled"`
|
||||||
|
|||||||
@@ -167,6 +167,7 @@ func (m *Manager) SetConfig(cfg *Config) error {
|
|||||||
m.v.Set("cors", cfg.CORS)
|
m.v.Set("cors", cfg.CORS)
|
||||||
m.v.Set("event_broker", cfg.EventBroker)
|
m.v.Set("event_broker", cfg.EventBroker)
|
||||||
m.v.Set("dbmanager", cfg.DBManager)
|
m.v.Set("dbmanager", cfg.DBManager)
|
||||||
|
m.v.Set("db_trace", cfg.DBTrace)
|
||||||
m.v.Set("paths", cfg.Paths)
|
m.v.Set("paths", cfg.Paths)
|
||||||
m.v.Set("extensions", cfg.Extensions)
|
m.v.Set("extensions", cfg.Extensions)
|
||||||
|
|
||||||
@@ -282,6 +283,12 @@ func setDefaults(v *viper.Viper) {
|
|||||||
v.SetDefault("database.url", "")
|
v.SetDefault("database.url", "")
|
||||||
|
|
||||||
// Database Manager defaults
|
// 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.default_connection", "default")
|
||||||
v.SetDefault("dbmanager.max_open_conns", 25)
|
v.SetDefault("dbmanager.max_open_conns", 25)
|
||||||
v.SetDefault("dbmanager.max_idle_conns", 5)
|
v.SetDefault("dbmanager.max_idle_conns", 5)
|
||||||
|
|||||||
@@ -2,7 +2,10 @@ package dbmanager
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"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"
|
||||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
"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
|
// sql.DBStats values are cumulative, so add only the growth since
|
||||||
// the last publish to keep these true counters.
|
// the last publish to keep these true counters.
|
||||||
prev := lastPublished.swap(name, connStats)
|
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))
|
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
|
||||||
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
||||||
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -22,6 +22,12 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
return nil
|
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
|
// Hook 1: BeforeRead - Load security rules
|
||||||
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
|||||||
@@ -21,6 +21,12 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
return nil
|
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
|
// Hook 1: BeforeRead - Load security rules
|
||||||
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeRead, func(hookCtx *HookContext) error {
|
||||||
secCtx := newSecurityContext(hookCtx)
|
secCtx := newSecurityContext(hookCtx)
|
||||||
|
|||||||
@@ -241,6 +241,19 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
|
|||||||
return loadSecurityRules(secCtx, securityList)
|
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
|
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
|
||||||
// This allows other packages to apply row-level security using the generic interface
|
// This allows other packages to apply row-level security using the generic interface
|
||||||
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
|
"golang.org/x/sync/singleflight"
|
||||||
)
|
)
|
||||||
|
|
||||||
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
|
||||||
@@ -51,6 +53,9 @@ type DatabaseKeyStore struct {
|
|||||||
capability *dbCapability
|
capability *dbCapability
|
||||||
cache *cache.Cache
|
cache *cache.Cache
|
||||||
cacheTTL time.Duration
|
cacheTTL time.Duration
|
||||||
|
|
||||||
|
// validateLoads collapses concurrent key lookups for the same key
|
||||||
|
validateLoads singleflight.Group
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewDatabaseKeyStore creates a DatabaseKeyStore with optional configuration.
|
// 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) {
|
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) {
|
||||||
key, err := ks.validateKeyDirect(ctx, hash, keyType)
|
key, err := ks.validateKeyDirect(ctx, hash, keyType)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
|
|
||||||
"github.com/tidwall/gjson"
|
"github.com/tidwall/gjson"
|
||||||
"github.com/tidwall/sjson"
|
"github.com/tidwall/sjson"
|
||||||
|
"golang.org/x/sync/singleflight"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ColumnSecurity struct {
|
type ColumnSecurity struct {
|
||||||
@@ -130,6 +131,9 @@ type SecurityList struct {
|
|||||||
rowSecExpiry map[string]time.Time
|
rowSecExpiry map[string]time.Time
|
||||||
lastColPrune time.Time
|
lastColPrune time.Time
|
||||||
lastRowPrune time.Time
|
lastRowPrune time.Time
|
||||||
|
|
||||||
|
// loads collapses concurrent provider calls for the same key (cold-cache stampede).
|
||||||
|
loads singleflight.Group
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -479,10 +483,13 @@ func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSch
|
|||||||
// Query the provider without holding any lock.
|
// Query the provider without holding any lock.
|
||||||
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
|
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
|
||||||
defer cancel()
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("GetColumnSecurity failed: %v", err)
|
return fmt.Errorf("GetColumnSecurity failed: %v", err)
|
||||||
}
|
}
|
||||||
|
colSecList, _ := v.([]ColumnSecurity)
|
||||||
if colSecList == nil {
|
if colSecList == nil {
|
||||||
colSecList = make([]ColumnSecurity, 0)
|
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.
|
// Query the provider without holding any lock.
|
||||||
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
|
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
|
||||||
defer cancel()
|
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 {
|
if err != nil {
|
||||||
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
|
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
|
||||||
}
|
}
|
||||||
|
record, _ := v.(RowSecurity)
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
m.RowSecurityMutex.Lock()
|
m.RowSecurityMutex.Lock()
|
||||||
|
|||||||
+97
-43
@@ -12,7 +12,9 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
"github.com/bitechdev/ResolveSpec/pkg/cache"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"golang.org/x/sync/singleflight"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Production-Ready Authenticators
|
// Production-Ready Authenticators
|
||||||
@@ -71,6 +73,39 @@ const maxAuthTokens = 4
|
|||||||
// sessionActivityTimeout bounds the detached last-activity update.
|
// sessionActivityTimeout bounds the detached last-activity update.
|
||||||
const sessionActivityTimeout = 5 * time.Second
|
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
|
// DatabaseAuthenticator provides session-based authentication with database storage
|
||||||
// All database operations go through stored procedures for security and consistency
|
// All database operations go through stored procedures for security and consistency
|
||||||
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
|
// 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 tracks in-flight asynchronous session activity updates
|
||||||
activityWG sync.WaitGroup
|
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)
|
// Cookie session support (optional, gated by enableCookieSession)
|
||||||
enableCookieSession bool
|
enableCookieSession bool
|
||||||
@@ -421,63 +460,75 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
|
|||||||
cacheKey := fmt.Sprintf("auth:session:%s", token)
|
cacheKey := fmt.Sprintf("auth:session:%s", token)
|
||||||
|
|
||||||
// Use cache.GetOrSet to get from cache or load from database
|
// Use cache.GetOrSet to get from cache or load from database
|
||||||
var userCtx UserContext
|
// Concurrent misses for the same token share one database lookup.
|
||||||
err := a.cache.GetOrSet(r.Context(), cacheKey, &userCtx, a.cacheTTL, func() (any, error) {
|
v, err, _ := a.sessionLoads.Do(cacheKey, func() (any, error) {
|
||||||
// This function is called only if cache miss
|
var loaded UserContext
|
||||||
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
|
err := a.cache.GetOrSet(r.Context(), cacheKey, &loaded, a.cacheTTL, func() (any, error) {
|
||||||
return a.sessionDirect(r.Context(), token)
|
// 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 success bool
|
||||||
var errorMsg sql.NullString
|
var errorMsg sql.NullString
|
||||||
var userJSON sql.NullString
|
var userJSON sql.NullString
|
||||||
|
|
||||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
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)
|
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
|
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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("session query failed: %w", err)
|
return nil, err
|
||||||
}
|
}
|
||||||
|
return loaded, nil
|
||||||
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 {
|
if err != nil {
|
||||||
lastErr = err
|
lastErr = err
|
||||||
continue // Try next token
|
continue // Try next token
|
||||||
}
|
}
|
||||||
|
userCtx, _ := v.(UserContext)
|
||||||
|
|
||||||
// Authentication succeeded with this token
|
// Authentication succeeded with this token
|
||||||
// Update last activity timestamp asynchronously
|
// Update last activity timestamp asynchronously
|
||||||
activityCtx := userCtx
|
if a.activityLimit.allow(token, time.Now()) {
|
||||||
// Detach from the request (it is cancelled when the handler returns) but
|
activityCtx := userCtx
|
||||||
// keep a deadline, and never let a panic here take the process down.
|
// Detach from the request (it is cancelled when the handler returns) but
|
||||||
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
|
// keep a deadline, and never let a panic here take the process down.
|
||||||
a.activityWG.Add(1)
|
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
|
||||||
go func(ctx context.Context, token string) {
|
a.activityWG.Add(1)
|
||||||
defer a.activityWG.Done()
|
go func(ctx context.Context, token string) {
|
||||||
defer cancel()
|
defer a.activityWG.Done()
|
||||||
defer logger.CatchPanic("updateSessionActivity")()
|
defer cancel()
|
||||||
a.updateSessionActivity(ctx, token, &activityCtx)
|
defer logger.CatchPanic("updateSessionActivity")()
|
||||||
}(detached, token)
|
a.updateSessionActivity(ctx, token, &activityCtx)
|
||||||
|
}(detached, token)
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
return &userCtx, nil
|
return &userCtx, nil
|
||||||
}
|
}
|
||||||
@@ -513,6 +564,7 @@ func (a *DatabaseAuthenticator) ClearUserCache(userID int) error {
|
|||||||
|
|
||||||
// updateSessionActivity updates the last activity timestamp for the session
|
// updateSessionActivity updates the last activity timestamp for the session
|
||||||
func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) {
|
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) {
|
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.SessionUpdate) {
|
||||||
_ = a.updateSessionActivityDirect(ctx, sessionToken)
|
_ = a.updateSessionActivityDirect(ctx, sessionToken)
|
||||||
return
|
return
|
||||||
@@ -852,6 +904,7 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
|
|||||||
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
|
||||||
return nil, ErrDirectModeUnsupported
|
return nil, ErrDirectModeUnsupported
|
||||||
}
|
}
|
||||||
|
dbtrace.Raw(ctx, "security.column")
|
||||||
|
|
||||||
var rules []ColumnSecurity
|
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) {
|
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
|
||||||
return RowSecurity{}, ErrDirectModeUnsupported
|
return RowSecurity{}, ErrDirectModeUnsupported
|
||||||
}
|
}
|
||||||
|
dbtrace.Raw(ctx, "security.row")
|
||||||
|
|
||||||
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
|
// 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;
|
// hand back the full *UserContext so non-DB providers can inspect claims;
|
||||||
|
|||||||
@@ -10,6 +10,8 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||||
)
|
)
|
||||||
|
|
||||||
// QueryMode selects how a provider talks to the database: via the configured
|
// 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 {
|
if db == nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
dbtrace.Raw(ctx, "probe.pg_proc")
|
||||||
var exists bool
|
var exists bool
|
||||||
defer func() {
|
defer func() {
|
||||||
// Guard against any unexpected panic from a misbehaving driver.
|
// Guard against any unexpected panic from a misbehaving driver.
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user