mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-05 13:01:58 +00:00
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:
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user