diff --git a/audit/pkg/security.audit.md b/audit/pkg/security.audit.md index 1c76893..d9d7c38 100644 --- a/audit/pkg/security.audit.md +++ b/audit/pkg/security.audit.md @@ -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 diff --git a/pkg/security/concurrency_test.go b/pkg/security/concurrency_test.go new file mode 100644 index 0000000..cc9c3e9 --- /dev/null +++ b/pkg/security/concurrency_test.go @@ -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) + } +} diff --git a/pkg/security/hooks.go b/pkg/security/hooks.go index 021fa50..a9bae98 100644 --- a/pkg/security/hooks.go +++ b/pkg/security/hooks.go @@ -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 }) } diff --git a/pkg/security/oauth2_methods.go b/pkg/security/oauth2_methods.go index 5b7ccf7..bc4fe9c 100644 --- a/pkg/security/oauth2_methods.go +++ b/pkg/security/oauth2_methods.go @@ -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{ diff --git a/pkg/security/provider.go b/pkg/security/provider.go index 8f638f5..31d69d3 100644 --- a/pkg/security/provider.go +++ b/pkg/security/provider.go @@ -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") diff --git a/pkg/security/providers.go b/pkg/security/providers.go index 059c985..51d9940 100644 --- a/pkg/security/providers.go +++ b/pkg/security/providers.go @@ -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 }