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:
2026-09-30 21:44:28 +02:00
parent 62cc14c02a
commit 3e327d0c78
20 changed files with 544 additions and 87 deletions
+13
View File
@@ -241,6 +241,19 @@ func LoadSecurityRules(secCtx SecurityContext, securityList *SecurityList) error
return loadSecurityRules(secCtx, securityList)
}
// PreloadSecurityRules loads column/row security rules into the SecurityList
// cache for read operations. Call it from a BeforeHandle hook, i.e. before the
// handler opens its transaction, so the provider queries do not need a second
// pooled connection while the transaction holds one. Later LoadSecurityRules
// calls in the same request are then cache hits. Non-read operations and
// models with security disabled are skipped.
func PreloadSecurityRules(secCtx SecurityContext, securityList *SecurityList, operation string) error {
if operation != "read" || IsModelSecurityDisabled(secCtx) {
return nil
}
return loadSecurityRules(secCtx, securityList)
}
// ApplyRowSecurity is a public wrapper for applyRowSecurity that accepts a SecurityContext
// This allows other packages to apply row-level security using the generic interface
func ApplyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error {
+23
View File
@@ -12,6 +12,8 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/cache"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"golang.org/x/sync/singleflight"
)
// DatabaseKeyStoreOptions configures DatabaseKeyStore.
@@ -51,6 +53,9 @@ type DatabaseKeyStore struct {
capability *dbCapability
cache *cache.Cache
cacheTTL time.Duration
// validateLoads collapses concurrent key lookups for the same key
validateLoads singleflight.Group
}
// NewDatabaseKeyStore creates a DatabaseKeyStore with optional configuration.
@@ -237,6 +242,24 @@ func (ks *DatabaseKeyStore) ValidateKey(ctx context.Context, rawKey string, keyT
}
}
// Concurrent misses for the same key share one database lookup.
v, err, _ := ks.validateLoads.Do(cacheKey+"|"+string(keyType), func() (any, error) {
return ks.validateKeyLoad(ctx, hash, cacheKey, keyType)
})
if err != nil {
return nil, err
}
key, _ := v.(*UserKey)
if key == nil {
return nil, errors.New("invalid or expired key")
}
cp := *key
return &cp, nil
}
// validateKeyLoad validates against the database and fills the cache.
func (ks *DatabaseKeyStore) validateKeyLoad(ctx context.Context, hash, cacheKey string, keyType KeyType) (*UserKey, error) {
dbtrace.Raw(ctx, "keystore.validate")
if !ks.capability.ShouldUseProcedure(ctx, ks.queryMode, ks.getDB(), ks.sqlNames.ValidateKey) {
key, err := ks.validateKeyDirect(ctx, hash, keyType)
if err != nil {
+12 -2
View File
@@ -15,6 +15,7 @@ import (
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
"golang.org/x/sync/singleflight"
)
type ColumnSecurity struct {
@@ -130,6 +131,9 @@ type SecurityList struct {
rowSecExpiry map[string]time.Time
lastColPrune time.Time
lastRowPrune time.Time
// loads collapses concurrent provider calls for the same key (cold-cache stampede).
loads singleflight.Group
}
const (
@@ -479,10 +483,13 @@ func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSch
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
colSecList, err := m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
v, err, _ := m.loads.Do("col:"+secKey, func() (any, error) {
return m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
})
if err != nil {
return fmt.Errorf("GetColumnSecurity failed: %v", err)
}
colSecList, _ := v.([]ColumnSecurity)
if colSecList == nil {
colSecList = make([]ColumnSecurity, 0)
}
@@ -552,10 +559,13 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchem
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
record, err := m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
v, err, _ := m.loads.Do("row:"+secKey, func() (any, error) {
return m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
})
if err != nil {
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
}
record, _ := v.(RowSecurity)
now := time.Now()
m.RowSecurityMutex.Lock()
+97 -43
View File
@@ -12,7 +12,9 @@ import (
"time"
"github.com/bitechdev/ResolveSpec/pkg/cache"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"golang.org/x/sync/singleflight"
)
// Production-Ready Authenticators
@@ -71,6 +73,39 @@ const maxAuthTokens = 4
// sessionActivityTimeout bounds the detached last-activity update.
const sessionActivityTimeout = 5 * time.Second
// sessionActivityInterval is the minimum gap between last-activity writes for
// one session token. Requests inside it skip the write.
const sessionActivityInterval = time.Minute
// activityThrottle remembers when each token's activity was last written.
type activityThrottle struct {
mu sync.Mutex
last map[string]time.Time
lastPrune time.Time
}
// allow reports whether token is due an activity write, and if so records it.
func (t *activityThrottle) allow(token string, now time.Time) bool {
t.mu.Lock()
defer t.mu.Unlock()
if t.last == nil {
t.last = make(map[string]time.Time)
}
if prev, ok := t.last[token]; ok && now.Sub(prev) < sessionActivityInterval {
return false
}
t.last[token] = now
if now.Sub(t.lastPrune) > sessionActivityInterval {
t.lastPrune = now
for k, v := range t.last {
if now.Sub(v) >= sessionActivityInterval {
delete(t.last, k)
}
}
}
return true
}
// DatabaseAuthenticator provides session-based authentication with database storage
// All database operations go through stored procedures for security and consistency
// Procedure names are configurable via SQLNames (see DefaultSQLNames for defaults)
@@ -94,6 +129,10 @@ type DatabaseAuthenticator struct {
// activityWG tracks in-flight asynchronous session activity updates
activityWG sync.WaitGroup
// activityLimit throttles those updates to one per token per interval
activityLimit activityThrottle
// sessionLoads collapses concurrent session lookups for the same token
sessionLoads singleflight.Group
// Cookie session support (optional, gated by enableCookieSession)
enableCookieSession bool
@@ -421,63 +460,75 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
cacheKey := fmt.Sprintf("auth:session:%s", token)
// Use cache.GetOrSet to get from cache or load from database
var userCtx UserContext
err := a.cache.GetOrSet(r.Context(), cacheKey, &userCtx, a.cacheTTL, func() (any, error) {
// This function is called only if cache miss
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
return a.sessionDirect(r.Context(), token)
}
// Concurrent misses for the same token share one database lookup.
v, err, _ := a.sessionLoads.Do(cacheKey, func() (any, error) {
var loaded UserContext
err := a.cache.GetOrSet(r.Context(), cacheKey, &loaded, a.cacheTTL, func() (any, error) {
// This function is called only if cache miss
dbtrace.Raw(r.Context(), "auth.session")
if !a.capability.ShouldUseProcedure(r.Context(), a.queryMode, a.getDB(), a.sqlNames.Session) {
return a.sessionDirect(r.Context(), token)
}
var success bool
var errorMsg sql.NullString
var userJSON sql.NullString
var success bool
var errorMsg sql.NullString
var userJSON sql.NullString
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
query := fmt.Sprintf(`SELECT p_success, p_error, p_user::text FROM %s($1, $2)`, a.sqlNames.Session)
return db.QueryRowContext(r.Context(), query, token, reference).Scan(&success, &errorMsg, &userJSON) //nolint:gosec // G701: identifier comes from trusted config, values are bound parameters
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
}
return nil, fmt.Errorf("invalid or expired session")
}
if !userJSON.Valid {
return nil, fmt.Errorf("no user data in session")
}
// Parse UserContext
var user UserContext
if err := json.Unmarshal([]byte(userJSON.String), &user); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
return &user, nil
})
if err != nil {
return nil, fmt.Errorf("session query failed: %w", err)
return nil, err
}
if !success {
if errorMsg.Valid {
return nil, fmt.Errorf("%s", errorMsg.String)
}
return nil, fmt.Errorf("invalid or expired session")
}
if !userJSON.Valid {
return nil, fmt.Errorf("no user data in session")
}
// Parse UserContext
var user UserContext
if err := json.Unmarshal([]byte(userJSON.String), &user); err != nil {
return nil, fmt.Errorf("failed to parse user context: %w", err)
}
return &user, nil
return loaded, nil
})
if err != nil {
lastErr = err
continue // Try next token
}
userCtx, _ := v.(UserContext)
// Authentication succeeded with this token
// Update last activity timestamp asynchronously
activityCtx := userCtx
// Detach from the request (it is cancelled when the handler returns) but
// keep a deadline, and never let a panic here take the process down.
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
a.activityWG.Add(1)
go func(ctx context.Context, token string) {
defer a.activityWG.Done()
defer cancel()
defer logger.CatchPanic("updateSessionActivity")()
a.updateSessionActivity(ctx, token, &activityCtx)
}(detached, token)
if a.activityLimit.allow(token, time.Now()) {
activityCtx := userCtx
// Detach from the request (it is cancelled when the handler returns) but
// keep a deadline, and never let a panic here take the process down.
detached, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), sessionActivityTimeout)
a.activityWG.Add(1)
go func(ctx context.Context, token string) {
defer a.activityWG.Done()
defer cancel()
defer logger.CatchPanic("updateSessionActivity")()
a.updateSessionActivity(ctx, token, &activityCtx)
}(detached, token)
}
return &userCtx, nil
}
@@ -513,6 +564,7 @@ func (a *DatabaseAuthenticator) ClearUserCache(userID int) error {
// updateSessionActivity updates the last activity timestamp for the session
func (a *DatabaseAuthenticator) updateSessionActivity(ctx context.Context, sessionToken string, userCtx *UserContext) {
dbtrace.Raw(ctx, "auth.activity")
if !a.capability.ShouldUseProcedure(ctx, a.queryMode, a.getDB(), a.sqlNames.SessionUpdate) {
_ = a.updateSessionActivityDirect(ctx, sessionToken)
return
@@ -852,6 +904,7 @@ func (p *DatabaseColumnSecurityProvider) GetColumnSecurity(ctx context.Context,
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.ColumnSecurity) {
return nil, ErrDirectModeUnsupported
}
dbtrace.Raw(ctx, "security.column")
var rules []ColumnSecurity
@@ -968,6 +1021,7 @@ func (p *DatabaseRowSecurityProvider) GetRowSecurity(ctx context.Context, userRe
if !p.capability.ShouldUseProcedure(ctx, p.queryMode, p.getDB(), p.sqlNames.RowSecurity) {
return RowSecurity{}, ErrDirectModeUnsupported
}
dbtrace.Raw(ctx, "security.row")
// resolvespec_row_security's p_user_id is a scalar integer. GetUserRef() may
// hand back the full *UserContext so non-DB providers can inspect claims;
+3
View File
@@ -10,6 +10,8 @@ import (
"strings"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
)
// QueryMode selects how a provider talks to the database: via the configured
@@ -59,6 +61,7 @@ func probeFunctionExists(ctx context.Context, db *sql.DB, procName string) bool
if db == nil {
return false
}
dbtrace.Raw(ctx, "probe.pg_proc")
var exists bool
defer func() {
// Guard against any unexpected panic from a misbehaving driver.
+40
View File
@@ -0,0 +1,40 @@
package security
import (
"context"
"sync"
"testing"
"time"
)
func TestConcurrentColdLoadsShareOneProviderCall(t *testing.T) {
p := &slowProvider{delay: 100 * time.Millisecond}
sl, _ := NewSecurityList(p)
var wg sync.WaitGroup
for i := 0; i < 10; i++ {
wg.Add(2)
go func() { defer wg.Done(); _ = sl.LoadColumnSecurity(context.Background(), 1, "s", "t", false) }()
go func() { defer wg.Done(); _, _ = sl.LoadRowSecurity(context.Background(), 1, "s", "t", false) }()
}
wg.Wait()
if got := p.calls.Load(); got != 2 {
t.Fatalf("provider calls = %d, want 2 (one column, one row)", got)
}
}
func TestActivityThrottle(t *testing.T) {
var th activityThrottle
now := time.Now()
if !th.allow("a", now) {
t.Fatal("first call must be allowed")
}
if th.allow("a", now.Add(sessionActivityInterval/2)) {
t.Fatal("call inside interval must be skipped")
}
if !th.allow("b", now) {
t.Fatal("other token must be allowed")
}
if !th.allow("a", now.Add(sessionActivityInterval)) {
t.Fatal("call after interval must be allowed")
}
}