fix(security): remove lock contention, cache rules, stop leaked goroutines

Load column/row security without holding locks across provider calls,
honour pOverwrite with a 30s TTL and pruning, cap tokens per
Authorization header, detach session activity updates with a timeout,
stop OAuth2 cleanup goroutines via Close, move nil-map checks inside
locks, and drop O(n^2) string building. Update audit status.
This commit is contained in:
Hein
2026-09-30 13:31:03 +02:00
parent 97fe88b3a6
commit 164ba2b240
6 changed files with 315 additions and 53 deletions
+14
View File
@@ -128,6 +128,20 @@ this; the gap is that the primary login path never got the same treatment.
| 32 | **Low** | security | `Authenticate` may return `(nil, nil)` through the callback, and the caller dereferences it |
| 33 | **Low** | security | `requestPasswordReset` returns the raw reset token to its caller |
## Resolution status (2026-09-30)
Only the thread-locking, data-race and slowness findings have been addressed so far
(#6, #9, #20, #21, #26, #30); every other finding is untouched.
- **#6** — Fixed: `LoadColumnSecurity`/`LoadRowSecurity` no longer hold a mutex across the provider call (load first, then publish under the lock), the provider call gets a 10 s deadline derived from the request context, and `pOverwrite` is honoured. Results are cached for 30 s (`securityCacheTTL`; revocations take up to that long to apply) and expired entries are pruned on write after a further 30 s grace. Expiry is tracked in side maps, so the exported `ColumnSecurity`/`RowSecurity` maps keep their shape. Duplicate cold-key queries are not collapsed (no singleflight)
- **#9** — Partly fixed: `Authenticate` rejects more than 4 comma-separated tokens (`maxAuthTokens`) with `too many authorization tokens`, and splits with `SplitN` so a huge header is not fully split. Not done: aborting the loop on the first hard failure, per-token rate limiting, and dropping the header from the `Warn` (finding 7)
- **#14** — Partly fixed as a side effect of #6: entries now expire and are pruned. The unstable `%v` row-security key (session token and claims inside the key) is unchanged
- **#20** — Fixed: the activity update runs on `context.WithoutCancel(r.Context())` with a 5 s timeout and recovers panics; the existing `activityWG` tracks it. Not done: coalescing into one batched flush
- **#21** — Fixed: each `OAuth2Provider` has a stop channel and `cleanupStates` exits on it (and recovers panics); replacing a provider stops the old one; new `DatabaseAuthenticator.Close()` stops all of them and waits for in-flight activity updates. `Close` is not yet called from the server shutdown path
- **#26** — Fixed: the nil-map checks in `ApplyColumnSecurity`, `ColumSecurityApplyOnRecord` and `GetRowSecurityTemplate` now run inside the lock. Error messages are unchanged
- **#30** — Fixed: `splitTag` uses `strings.FieldsFunc`; `maskString` uses a `strings.Builder` (its off-by-one offsets, finding 25, are unchanged)
- Tests: `pkg/security/concurrency_test.go` (run with `-race`).
---
## 1. Critical — the password is never verified, in either query mode
+147
View File
@@ -0,0 +1,147 @@
package security
import (
"context"
"net/http"
"net/http/httptest"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
)
// slowProvider embeds a nil SecurityProvider; only the two load methods are used.
type slowProvider struct {
SecurityProvider
calls atomic.Int32
active atomic.Int32
maxSeen atomic.Int32
delay time.Duration
}
func (p *slowProvider) enter() {
p.calls.Add(1)
n := p.active.Add(1)
for {
m := p.maxSeen.Load()
if n <= m || p.maxSeen.CompareAndSwap(m, n) {
break
}
}
time.Sleep(p.delay)
p.active.Add(-1)
}
func (p *slowProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
p.enter()
return []ColumnSecurity{{Schema: schema, Tablename: table}}, nil
}
func (p *slowProvider) GetRowSecurity(ctx context.Context, ref any, schema, table string) (RowSecurity, error) {
p.enter()
return RowSecurity{Schema: schema, Tablename: table}, nil
}
func TestLoadDoesNotHoldLockAcrossProvider(t *testing.T) {
p := &slowProvider{delay: 100 * time.Millisecond}
sl, _ := NewSecurityList(p)
var wg sync.WaitGroup
start := time.Now()
for i := 0; i < 8; i++ {
wg.Add(2)
go func(i int) { defer wg.Done(); _ = sl.LoadColumnSecurity(context.Background(), i, "s", "t", false) }(i)
go func(i int) { defer wg.Done(); _, _ = sl.LoadRowSecurity(context.Background(), i, "s", "t", false) }(i)
}
wg.Wait()
if el := time.Since(start); el > 500*time.Millisecond {
t.Fatalf("loads serialised: %v", el)
}
if p.maxSeen.Load() < 2 {
t.Fatal("provider calls never overlapped")
}
}
func TestLoadCachesAndHonoursOverwrite(t *testing.T) {
p := &slowProvider{}
sl, _ := NewSecurityList(p)
ctx := context.Background()
for i := 0; i < 3; i++ {
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
_, _ = sl.LoadRowSecurity(ctx, 1, "s", "t", false)
}
if got := p.calls.Load(); got != 2 {
t.Fatalf("expected 2 provider calls (cached), got %d", got)
}
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", true)
_, _ = sl.LoadRowSecurity(ctx, 1, "s", "t", true)
if got := p.calls.Load(); got != 4 {
t.Fatalf("overwrite should reload: got %d calls", got)
}
}
func TestLoadExpiryAndPrune(t *testing.T) {
p := &slowProvider{}
sl, _ := NewSecurityList(p)
ctx := context.Background()
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
sl.ColumnSecurityMutex.Lock()
sl.colSecExpiry["s.t@1"] = time.Now().Add(-time.Second)
sl.ColumnSecurityMutex.Unlock()
_ = sl.LoadColumnSecurity(ctx, 1, "s", "t", false)
if p.calls.Load() != 2 {
t.Fatal("expired entry should reload")
}
sl.ColumnSecurityMutex.Lock()
sl.colSecExpiry["s.old@9"] = time.Now().Add(-time.Hour)
sl.ColumnSecurity["s.old@9"] = nil
sl.lastColPrune = time.Time{}
sl.ColumnSecurityMutex.Unlock()
_ = sl.LoadColumnSecurity(ctx, 2, "s", "t", false)
sl.ColumnSecurityMutex.RLock()
_, ok := sl.ColumnSecurity["s.old@9"]
sl.ColumnSecurityMutex.RUnlock()
if ok {
t.Fatal("stale entry not pruned")
}
}
func TestAuthenticateRejectsTooManyTokens(t *testing.T) {
db, _, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
defer db.Close()
a := NewDatabaseAuthenticator(db)
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.Header.Set("Authorization", "a,b,c,d,e,f,g,h")
if _, err := a.Authenticate(r); err == nil || err.Error() != "too many authorization tokens" {
t.Fatalf("got %v", err)
}
}
func TestOAuth2CleanupStopsOnClose(t *testing.T) {
db, _, _ := sqlmock.New()
defer db.Close()
a := NewDatabaseAuthenticator(db).WithOAuth2(OAuth2Config{ClientID: "x", ProviderName: "p"})
p := a.oauth2Providers["p"]
if err := a.Close(); err != nil {
t.Fatal(err)
}
_ = a.Close() // idempotent
select {
case <-p.stopCh:
default:
t.Fatal("stop channel not closed")
}
}
func TestSplitTagDropsEmpty(t *testing.T) {
got := splitTag("a,,b,c,", ',')
if len(got) != 3 || got[0] != "a" || got[2] != "c" {
t.Fatalf("got %v", got)
}
}
+2 -16
View File
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"reflect"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
@@ -429,20 +430,5 @@ func extractSQLName(tag string) string {
}
func splitTag(tag string, sep rune) []string {
var parts []string
var current string
for _, ch := range tag {
if ch == sep {
if current != "" {
parts = append(parts, current)
current = ""
}
} else {
current += string(ch)
}
}
if current != "" {
parts = append(parts, current)
}
return parts
return strings.FieldsFunc(tag, func(r rune) bool { return r == sep })
}
+32 -1
View File
@@ -12,6 +12,8 @@ import (
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"golang.org/x/oauth2"
)
@@ -39,6 +41,8 @@ type OAuth2Provider struct {
providerName string
states map[string]time.Time // state -> expiry time
statesMutex sync.RWMutex
stopCh chan struct{} // closed to stop cleanupStates
stopOnce sync.Once
}
// WithOAuth2 configures OAuth2 support for the DatabaseAuthenticator
@@ -68,6 +72,7 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
userInfoParser: cfg.UserInfoParser,
providerName: cfg.ProviderName,
states: make(map[string]time.Time),
stopCh: make(chan struct{}),
}
// Initialize providers map if needed
@@ -77,6 +82,9 @@ func (a *DatabaseAuthenticator) WithOAuth2(cfg OAuth2Config) *DatabaseAuthentica
}
// Register provider
if old := a.oauth2Providers[cfg.ProviderName]; old != nil {
old.stop() // replaced provider: stop its cleanup goroutine
}
a.oauth2Providers[cfg.ProviderName] = provider
a.oauth2ProvidersMutex.Unlock()
@@ -335,10 +343,16 @@ func (p *OAuth2Provider) validateState(state string) bool {
// cleanupStates removes expired states periodically
func (p *OAuth2Provider) cleanupStates() {
defer logger.CatchPanic("OAuth2Provider.cleanupStates")()
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for range ticker.C {
for {
select {
case <-p.stopCh:
return
case <-ticker.C:
}
p.statesMutex.Lock()
now := time.Now()
for state, expiry := range p.states {
@@ -350,6 +364,23 @@ func (p *OAuth2Provider) cleanupStates() {
}
}
// stop terminates the cleanup goroutine; safe to call more than once.
func (p *OAuth2Provider) stop() {
p.stopOnce.Do(func() { close(p.stopCh) })
}
// Close stops the background OAuth2 state cleanup goroutines and waits for
// in-flight session activity updates. It is safe to call more than once.
func (a *DatabaseAuthenticator) Close() error {
a.oauth2ProvidersMutex.RLock()
for _, p := range a.oauth2Providers {
p.stop()
}
a.oauth2ProvidersMutex.RUnlock()
a.activityWG.Wait()
return nil
}
// defaultOAuth2UserInfoParser parses standard OAuth2 user info claims
func defaultOAuth2UserInfoParser(userInfo map[string]any) (*UserContext, error) {
ctx := &UserContext{
+103 -34
View File
@@ -6,6 +6,7 @@ import (
"reflect"
"strings"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/reflection"
@@ -58,8 +59,25 @@ type SecurityList struct {
ColumnSecurity map[string][]ColumnSecurity
RowSecurityMutex sync.RWMutex
RowSecurity map[string]RowSecurity
// Expiry bookkeeping for the two caches above; guarded by the same mutexes.
colSecExpiry map[string]time.Time
rowSecExpiry map[string]time.Time
lastColPrune time.Time
lastRowPrune time.Time
}
const (
// securityCacheTTL is how long loaded rules are served without re-querying
// the provider. Revoked rules take up to this long to take effect.
securityCacheTTL = 30 * time.Second
// securityLoadTimeout bounds a single provider call.
securityLoadTimeout = 10 * time.Second
// securityPruneGrace keeps expired entries around long enough that a
// request that loaded them can still read them.
securityPruneGrace = securityCacheTTL
)
// NewSecurityList creates a new security list with the given provider
func NewSecurityList(provider SecurityProvider) (*SecurityList, error) {
if provider == nil {
@@ -85,7 +103,8 @@ const SECURITY_CONTEXT_KEY CONTEXT_KEY = "SecurityList"
func maskString(pString string, maskStart, maskEnd int, maskChar string, invert bool) string {
strLen := len(pString)
middleIndex := (strLen / 2)
newStr := ""
var newStr strings.Builder
newStr.Grow(strLen)
if maskStart == 0 && maskEnd == 0 {
maskStart = strLen
maskEnd = strLen
@@ -101,32 +120,29 @@ func maskString(pString string, maskStart, maskEnd int, maskChar string, invert
}
for index, char := range pString {
if invert && index >= middleIndex-maskStart && index <= middleIndex {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if invert && index <= middleIndex+maskEnd && index >= middleIndex {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if !invert && index <= maskStart {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
if !invert && index >= strLen-1-maskEnd {
newStr += maskChar
newStr.WriteString(maskChar)
continue
}
newStr += string(char)
newStr.WriteRune(char)
}
return newStr
return newStr.String()
}
func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newRecord reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) ([]string, error) {
cols := make([]string, 0)
if m.ColumnSecurity == nil {
return cols, fmt.Errorf("security not initialized")
}
if prevRecord.Type() != newRecord.Type() {
logger.Error("prev:%s and new:%s record type mismatch", prevRecord.Type(), newRecord.Type())
@@ -136,6 +152,10 @@ func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newR
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
if m.ColumnSecurity == nil {
return cols, fmt.Errorf("security not initialized")
}
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
if !ok || colsecList == nil {
return cols, fmt.Errorf("no column security data")
@@ -301,13 +321,13 @@ func setColSecValue(fieldsrc reflect.Value, colsec ColumnSecurity, fieldTypeName
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (reflect.Value, error) {
defer logger.CatchPanic("ApplyColumnSecurity")()
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
if m.ColumnSecurity == nil {
return records, fmt.Errorf("security not initialized")
}
m.ColumnSecurityMutex.RLock()
defer m.ColumnSecurityMutex.RUnlock()
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
if !ok || colsecList == nil {
return records, fmt.Errorf("nocolumn security data")
@@ -372,25 +392,49 @@ func (m *SecurityList) LoadColumnSecurity(ctx context.Context, pUserID int, pSch
return fmt.Errorf("security provider not set")
}
m.ColumnSecurityMutex.Lock()
defer m.ColumnSecurityMutex.Unlock()
if m.ColumnSecurity == nil {
m.ColumnSecurity = make(map[string][]ColumnSecurity, 0)
}
secKey := fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)
if pOverwrite || m.ColumnSecurity[secKey] == nil {
m.ColumnSecurity[secKey] = make([]ColumnSecurity, 0)
if !pOverwrite {
m.ColumnSecurityMutex.RLock()
exp, ok := m.colSecExpiry[secKey]
fresh := ok && m.ColumnSecurity[secKey] != nil && time.Now().Before(exp)
m.ColumnSecurityMutex.RUnlock()
if fresh {
return nil
}
}
// Call the provider to load security rules
colSecList, err := m.provider.GetColumnSecurity(ctx, pUserID, pSchema, pTablename)
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
colSecList, err := m.provider.GetColumnSecurity(loadCtx, pUserID, pSchema, pTablename)
if err != nil {
return fmt.Errorf("GetColumnSecurity failed: %v", err)
}
if colSecList == nil {
colSecList = make([]ColumnSecurity, 0)
}
now := time.Now()
m.ColumnSecurityMutex.Lock()
defer m.ColumnSecurityMutex.Unlock()
if m.ColumnSecurity == nil {
m.ColumnSecurity = make(map[string][]ColumnSecurity)
}
if m.colSecExpiry == nil {
m.colSecExpiry = make(map[string]time.Time)
}
m.ColumnSecurity[secKey] = colSecList
m.colSecExpiry[secKey] = now.Add(securityCacheTTL)
if now.Sub(m.lastColPrune) > securityCacheTTL {
m.lastColPrune = now
for k, exp := range m.colSecExpiry {
if now.Sub(exp) > securityPruneGrace {
delete(m.colSecExpiry, k)
delete(m.ColumnSecurity, k)
}
}
}
return nil
}
@@ -421,34 +465,59 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchem
return RowSecurity{}, fmt.Errorf("security provider not set")
}
m.RowSecurityMutex.Lock()
defer m.RowSecurityMutex.Unlock()
if m.RowSecurity == nil {
m.RowSecurity = make(map[string]RowSecurity, 0)
}
secKey := fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)
// Call the provider to load security rules
record, err := m.provider.GetRowSecurity(ctx, pUserRef, pSchema, pTablename)
if !pOverwrite {
m.RowSecurityMutex.RLock()
exp, ok := m.rowSecExpiry[secKey]
cached, present := m.RowSecurity[secKey]
m.RowSecurityMutex.RUnlock()
if ok && present && time.Now().Before(exp) {
return cached, nil
}
}
// Query the provider without holding any lock.
loadCtx, cancel := context.WithTimeout(ctx, securityLoadTimeout)
defer cancel()
record, err := m.provider.GetRowSecurity(loadCtx, pUserRef, pSchema, pTablename)
if err != nil {
return RowSecurity{}, fmt.Errorf("GetRowSecurity failed: %v", err)
}
now := time.Now()
m.RowSecurityMutex.Lock()
defer m.RowSecurityMutex.Unlock()
if m.RowSecurity == nil {
m.RowSecurity = make(map[string]RowSecurity)
}
if m.rowSecExpiry == nil {
m.rowSecExpiry = make(map[string]time.Time)
}
m.RowSecurity[secKey] = record
m.rowSecExpiry[secKey] = now.Add(securityCacheTTL)
if now.Sub(m.lastRowPrune) > securityCacheTTL {
m.lastRowPrune = now
for k, exp := range m.rowSecExpiry {
if now.Sub(exp) > securityPruneGrace {
delete(m.rowSecExpiry, k)
delete(m.RowSecurity, k)
}
}
}
return record, nil
}
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (RowSecurity, error) {
defer logger.CatchPanic("GetRowSecurityTemplate")()
m.RowSecurityMutex.RLock()
defer m.RowSecurityMutex.RUnlock()
if m.RowSecurity == nil {
return RowSecurity{}, fmt.Errorf("security not initialized")
}
m.RowSecurityMutex.RLock()
defer m.RowSecurityMutex.RUnlock()
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)]
if !ok {
return RowSecurity{}, fmt.Errorf("no row security data")
+17 -2
View File
@@ -64,6 +64,13 @@ func (a *HeaderAuthenticator) Authenticate(r *http.Request) (*UserContext, error
}, nil
}
// maxAuthTokens caps the comma-separated credentials tried per request so one
// request cannot drive unbounded session lookups.
const maxAuthTokens = 4
// sessionActivityTimeout bounds the detached last-activity update.
const sessionActivityTimeout = 5 * time.Second
// 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)
@@ -368,7 +375,10 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
} else {
// Parse Authorization header which may contain multiple comma-separated tokens
// Format: "Token abc, Token def" or "Bearer abc" or just "abc"
rawTokens := strings.Split(sessionToken, ",")
rawTokens := strings.SplitN(sessionToken, ",", maxAuthTokens+2)
if len(rawTokens) > maxAuthTokens {
return nil, fmt.Errorf("too many authorization tokens")
}
for _, token := range rawTokens {
token = strings.TrimSpace(token)
// Remove "Bearer " prefix if present
@@ -448,11 +458,16 @@ func (a *DatabaseAuthenticator) Authenticate(r *http.Request) (*UserContext, err
// 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)
}(r.Context(), token)
}(detached, token)
return &userCtx, nil
}