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
+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")