mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 03:22:09 +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:
+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;
|
||||
|
||||
Reference in New Issue
Block a user