mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-04 20:41: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:
+103
-34
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user