mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +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:
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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