diff --git a/audit/pkg/_CROSS-CUTTING.audit.md b/audit/pkg/_CROSS-CUTTING.audit.md index 1650dfe..c7ba220 100644 --- a/audit/pkg/_CROSS-CUTTING.audit.md +++ b/audit/pkg/_CROSS-CUTTING.audit.md @@ -23,7 +23,7 @@ per-package audits reference this file rather than restating them. | X5 | **Medium** | locking | Unsynchronized package-level mutable globals are the dominant concurrency pattern | | X6 | **Medium** | security | Insecure-by-default transport across the board: `sslmode: disable`, `WithInsecure()`, no TLS in cache configs | | X7 | **Medium** | panic handling | Panic handling is inconsistent and, where it exists, tends to fail open | -| X8 | **Medium** | security | `logger.Warn`/`Error` forward every message to Sentry unscrubbed, and error strings routinely embed attacker data | +| X8 | **Medium** | security | `logger.Warn`/`Error` forward every message to Sentry unscrubbed, and error strings routinely embed attacker data *(partly fixed 2026-09-30: redaction and rate limiting added in `pkg/logger`; call sites still embed attacker data)* | | X9 | **Low** | testing | Test coverage is extremely uneven: 5 packages have no test file at all | The table is ordered by severity; the sections below are in ID order, since other @@ -67,7 +67,7 @@ every one of them on the first run: | `pkg/cache` | `defaultCache` read/written by concurrent request handlers | `cache.audit.md` finding 3 | | `pkg/config` | `*viper.Viper` has no internal lock; `configInstance` singleton | `config.audit.md` findings 1, 2 | | `pkg/logger` | `Logger`, `errorTracker` globals | `logger.audit.md` finding 1 | -| `pkg/modelregistry` | `defaultRegistry` read by 6 functions without the lock | `modelregistry.audit.md` findings 2, 8 | +| `pkg/modelregistry` | `defaultRegistry` read by 6 functions without the lock | `modelregistry.audit.md` findings 2, 8 *(fixed 2026-09-30)* | | `pkg/tracing` | `tracer` global | `tracing.audit.md` finding 5 | | `pkg/errortracking` | `sentry.Init` mutates process globals | `errortracking.audit.md` finding 2 | @@ -160,7 +160,7 @@ The test bodies that exist but are never executed by CI: | `metrics` | 1 | 64 | no | | `resolvemcp` | 1 | 34 | no | | `logger` | 0 | 0 | — | -| `modelregistry` | 0 | 0 | — | +| `modelregistry` | 1 | ~150 | yes (`-race`) *(added 2026-09-30)* | | `testmodels` | 0 | 0 | — | | `tracing` | 0 | 0 | — | @@ -308,7 +308,7 @@ package-level variables, and most guard it with nothing: | `pkg/cache` | `defaultCache *Cache` (`cache.go:10`) | **no** | | `pkg/config` | `configInstance *Manager` (`manager.go:15`) | **no** | | `pkg/tracing` | `tracer` (`tracing.go:19`) | **no** | -| `pkg/modelregistry` | `defaultRegistry` | partially — `TryLock` with a retry/`time.Sleep` loop, and 6 functions read it unlocked | +| `pkg/modelregistry` | `defaultRegistry` | **yes** *(fixed 2026-09-30)* — guarded by `registriesMutex`; all access via `GetDefaultRegistry()` | | `pkg/metrics` | `globalProvider` (`interfaces.go:50-51`) | **yes** — `globalProviderMu sync.RWMutex` | `pkg/metrics` is the model the others should follow: @@ -524,7 +524,7 @@ codebase invites it by logging the header contents as the diagnostic. only **Critical** authorization finding: `GetModel` returns a "registry locked" error under write-lock contention, which `security/hooks.go:274-294` converts into `return nil // model not registered, allow by default` -(`modelregistry.audit.md` finding 1). A twenty-line test that registers a model +(`modelregistry.audit.md` finding 1; *fixed 2026-09-30, regression tests added*). A twenty-line test that registers a model from one goroutine while reading it from another would demonstrate the fail-open immediately. The package guards a security boundary and has never been tested. diff --git a/audit/pkg/logger.audit.md b/audit/pkg/logger.audit.md index f6cf763..45a0ed7 100644 --- a/audit/pkg/logger.audit.md +++ b/audit/pkg/logger.audit.md @@ -32,6 +32,22 @@ There are **zero tests** in this package. | 11 | Low | Slowness | `os.Getpid()` called on every log line | | 12 | Low | Observability | `UpdateLogger` build failure degrades silently to stdlib `log` | +## Resolution status (2026-09-30) + +- **#1** — Fixed (earlier race work): `stateMu` RWMutex with `getLogger`/`swapLogger`/`getErrorTracker`; the exported `Logger` var is kept for compatibility +- **#2** — Fixed: messages are scrubbed before `CaptureMessage` (URL credentials, `password=`/`token=`/`secret=`/`api_key=` values, `Bearer`/`Basic` tokens). Local logs are unchanged. Sentry `BeforeSend` and structured-field allowlisting are not done +- **#3** — Partly fixed: global token bucket (burst 50, 20/s) plus per-severity/template dedup (1s, 1024 keys). Panics are not limited. The `error_tracking.sample_rate` default (Sentry maps 0 to 1.0) is still unset in `config/manager.go` +- **#4** — Partly fixed: `CatchPanicRethrow` added. `pkg/security/provider.go:302` and `:443` still use the swallowing `CatchPanic`; left for the security audit pass +- **#5** — Fixed: stack captured with `runtime.Stack` into a 16 KiB buffer. Per-fingerprint panic rate limiting not done +- **#6** — Fixed: `Info`/`Debug` format first and fall back with `log.Printf("%s", ...)`. `gosec` was enabled separately +- **#7** — Fixed: CR/LF and other control characters are escaped on the stdlib fallback path +- **#8** — Fixed: `Info`/`Debug` strip `context.Context` args +- **#9** — Fixed: the replaced logger is synced on `UpdateLogger` +- **#10** — Fixed: `logger.Sync()` added. Not yet called from the server shutdown path +- **#11** — Fixed: PID cached in a package var +- **#12** — Partly fixed: `UpdateLoggerE` returns the build error and a failed build keeps the previous logger. `Init` still returns nothing +- Tests: `pkg/logger/logger_test.go` (run with `-race`). + --- ## Findings diff --git a/audit/pkg/modelregistry.audit.md b/audit/pkg/modelregistry.audit.md index c7ff5ca..88470f0 100644 --- a/audit/pkg/modelregistry.audit.md +++ b/audit/pkg/modelregistry.audit.md @@ -38,6 +38,34 @@ There are **no tests** in this package and no `-race` coverage of it anywhere. | 12 | Low | Panic | Package has no `recover` anywhere, and calls a caller-supplied callback under a lock (see 6) | | 13 | Low | Security | `DefaultModelRules()` grants `CanRead/Update/Create/Delete: true` — registration without explicit rules is fully mutable | +## Resolution (2026-09-30) + +Fixed in `pkg/modelregistry/model_registry.go`, `pkg/security/hooks.go`, and new +`pkg/modelregistry/model_registry_test.go` (passes under `-race`). + +| # | Status | What changed | +|---|--------|--------------| +| 1 | **Fixed** | Added sentinels `ErrModelNotFound`, `ErrModelExists`, `ErrInvalidModel` (wrapped, `errors.Is`-friendly). `checkModelUpdateAllowed`/`checkModelDeleteAllowed` now allow-by-default **only** on `ErrModelNotFound`; any other error denies. Lookups can no longer return a "locked" error at all. | +| 2 | **Fixed** | `GetDefaultRegistry` uses a plain `RLock`; no unsynchronised fallback. | +| 3 | **Fixed** | `SetDefaultRegistry` uses a blocking `Lock` (cannot silently no-op); a nil registry is ignored. | +| 4 | **Fixed** | `RegisterModelWithRules` and `RegisterModel` share `registerLocked`, which writes model + rules under one lock acquisition. | +| 5 | **Fixed** | `GetAllModels`/`GetModels` use blocking locks and can no longer return empty/partial results due to contention. Signatures unchanged (`GetAllModels` is used through interfaces by resolvespec/restheadspec/openapi). | +| 6 | **Fixed** | `IterateModels` iterates a snapshot; the callback runs with no lock held (regression test re-enters the registry). | +| 7 | **Fixed** | Try-lock/sleep helpers and `lockRetry*` constants removed. | +| 8 | **Fixed** | All package-level functions go through `GetDefaultRegistry()` / `registriesSnapshot()`; `defaultRegistry` is only touched under `registriesMutex`. | +| 9 | **Fixed** | One discipline: blocking locks, snapshot-and-release, documented lock order (`registriesMutex` before a registry's mutex). | +| 10 | **Fixed** | Reflection/validation (`validateModel`) runs before the write lock is taken. | +| 11 | **Fixed** | Unwrap loop capped at 16 levels; `type T *T` now returns `ErrInvalidModel` (tested). | +| 12 | **Fixed** | Sentinel errors added. `IterateModels` recovers a callback panic per model, logs it via `logger.HandlePanic` with the model name, and continues; `validateModel` recovers reflection panics and returns `ErrInvalidModel` so registration fails closed. No lock is held during either, so the registry cannot be wedged. | +| 13 | **Accepted (decision)** | Allow-by-default retained deliberately: `DefaultModelRules()` still grants read/update/create/delete. Callers wanting restrictions must use `RegisterModelWithRules`/`SetModelRules`. | + +Tests added: sentinel errors, recursive pointer type, pointer normalisation, atomic +`RegisterModelWithRules` (concurrent reader never sees permissive rules), re-entrant `IterateModels`, +cross-registry `GetModelRulesByName`, and a concurrent `-race` stress test. + +Not changed: the `pkg/security` middleware-wiring question (context fast-path) remains tracked in +`audit/pkg/security.audit.md`. + --- ## Findings diff --git a/audit/pkg/security.audit.md b/audit/pkg/security.audit.md index d46b4d3..1c76893 100644 --- a/audit/pkg/security.audit.md +++ b/audit/pkg/security.audit.md @@ -1787,7 +1787,7 @@ func checkModelUpdateAllowed(secCtx SecurityContext) error { `checkModelDeleteAllowed` is identical (`:298-318`, fail-open at `:311`). A model served by the spec handler but absent from the registry — or present under a name the two lookups do not produce — is fully writable. This is the consuming side of -`modelregistry.audit.md` finding 1: the registry's lookup failure and this +`modelregistry.audit.md` finding 1 (*registry side fixed 2026-09-30: `checkModelUpdateAllowed`/`checkModelDeleteAllowed` now allow only on `ErrModelNotFound`*): the registry's lookup failure and this `return nil` combine into "unknown model ⇒ permitted". **Any authenticated user may perform any operation.** `CheckModelAuthAllowed` is diff --git a/pkg/modelregistry/model_registry.go b/pkg/modelregistry/model_registry.go index c848474..a600a25 100644 --- a/pkg/modelregistry/model_registry.go +++ b/pkg/modelregistry/model_registry.go @@ -1,10 +1,12 @@ package modelregistry import ( + "errors" "fmt" "reflect" "sync" - "time" + + "github.com/bitechdev/ResolveSpec/pkg/logger" ) // ModelRules defines the permissions and security settings for a model @@ -52,6 +54,21 @@ var defaultRegistry = &DefaultModelRegistry{ var registries = []*DefaultModelRegistry{defaultRegistry} var registriesMutex sync.RWMutex +// Sentinel errors so callers (notably the security layer) can distinguish +// "not registered" from every other failure with errors.Is. +var ( + ErrModelNotFound = errors.New("model not found") + ErrModelExists = errors.New("model already registered") + ErrInvalidModel = errors.New("invalid model") +) + +// maxUnwrapDepth bounds pointer/slice/array unwrapping so a recursive type +// (type T *T) cannot spin forever. +const maxUnwrapDepth = 16 + +// Lock ordering: registriesMutex is always taken before a registry's mutex, +// never the reverse. No caller-supplied code ever runs while a lock is held. + // NewModelRegistry creates a new model registry func NewModelRegistry() *DefaultModelRegistry { return &DefaultModelRegistry{ @@ -60,44 +77,19 @@ func NewModelRegistry() *DefaultModelRegistry { } } -// lockRetryAttempts/lockRetryDelay bound how long the try-lock helpers below -// will spin before giving up, so a contended registriesMutex can never hang -// a caller of GetDefaultRegistry/SetDefaultRegistry. -const ( - lockRetryAttempts = 20 - lockRetryDelay = 1 * time.Millisecond -) - -// GetDefaultRegistry returns the current default registry. It uses a -// bounded TryRLock instead of a blocking RLock so it can never hang; -// if the lock can't be acquired in time it falls back to the last known -// value without synchronization. +// GetDefaultRegistry returns the current default registry. func GetDefaultRegistry() *DefaultModelRegistry { - for i := 0; i < lockRetryAttempts; i++ { - if registriesMutex.TryRLock() { - defer registriesMutex.RUnlock() - return defaultRegistry - } - time.Sleep(lockRetryDelay) - } + registriesMutex.RLock() + defer registriesMutex.RUnlock() return defaultRegistry } -// SetDefaultRegistry replaces the default registry. It uses a bounded -// TryLock instead of a blocking Lock so it can never hang; if the lock -// can't be acquired in time the call is a no-op. +// SetDefaultRegistry replaces the default registry. A nil registry is ignored. func SetDefaultRegistry(registry *DefaultModelRegistry) { - acquired := false - for i := 0; i < lockRetryAttempts; i++ { - if registriesMutex.TryLock() { - acquired = true - break - } - time.Sleep(lockRetryDelay) - } - if !acquired { + if registry == nil { return } + registriesMutex.Lock() defer registriesMutex.Unlock() foundAt := -1 @@ -123,99 +115,96 @@ func AddRegistry(registry *DefaultModelRegistry) { registries = append(registries, registry) } -// tryLock attempts to acquire the registry's write lock, retrying briefly. -// Returns false if it could not be acquired within the bound. -func (r *DefaultModelRegistry) tryLock() bool { - for i := 0; i < lockRetryAttempts; i++ { - if r.mutex.TryLock() { - return true - } - time.Sleep(lockRetryDelay) - } - return false +// registriesSnapshot returns a copy of the registry list so callers can +// iterate without holding registriesMutex. +func registriesSnapshot() []*DefaultModelRegistry { + registriesMutex.RLock() + defer registriesMutex.RUnlock() + return append([]*DefaultModelRegistry(nil), registries...) } -// tryRLock attempts to acquire the registry's read lock, retrying briefly. -// Returns false if it could not be acquired within the bound. -func (r *DefaultModelRegistry) tryRLock() bool { - for i := 0; i < lockRetryAttempts; i++ { - if r.mutex.TryRLock() { - return true +// validateModel checks the model is a struct (or pointer/slice/array of one) +// and returns the normalised non-pointer struct value. It takes no locks. +func validateModel(model interface{}) (result interface{}, err error) { + // Reflection on a pathological type must fail the registration, not crash the process. + defer func() { + if r := recover(); r != nil { + result = nil + err = fmt.Errorf("%w: %v", ErrInvalidModel, logger.HandlePanic("modelregistry.validateModel", r)) } - time.Sleep(lockRetryDelay) - } - return false -} + }() -func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error { - if !r.tryLock() { - return fmt.Errorf("failed to register model %s: registry locked", name) - } - defer r.mutex.Unlock() - - if _, exists := r.models[name]; exists { - return fmt.Errorf("model %s already registered", name) - } - - // Validate that model is a non-pointer struct modelType := reflect.TypeOf(model) if modelType == nil { - return fmt.Errorf("model cannot be nil") + return nil, fmt.Errorf("%w: model cannot be nil", ErrInvalidModel) } originalType := modelType // Unwrap pointers, slices, and arrays to check the underlying type - for modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array { + for depth := 0; modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array; depth++ { + if depth >= maxUnwrapDepth { + return nil, fmt.Errorf("%w: type %s nests deeper than %d levels", ErrInvalidModel, originalType.String(), maxUnwrapDepth) + } modelType = modelType.Elem() } - // Validate that the underlying type is a struct if modelType.Kind() != reflect.Struct { - return fmt.Errorf("model must be a struct or pointer to struct, got %s", originalType.String()) + return nil, fmt.Errorf("%w: model must be a struct or pointer to struct, got %s", ErrInvalidModel, originalType.String()) } // If a pointer/slice/array was passed, unwrap to the base struct if originalType != modelType { - // Create a zero value of the struct type model = reflect.New(modelType).Elem().Interface() } - // Additional check: ensure model is not a pointer - finalType := reflect.TypeOf(model) - if finalType.Kind() == reflect.Pointer { - return fmt.Errorf("model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", finalType.Elem().Name()) + if finalType := reflect.TypeOf(model); finalType.Kind() == reflect.Pointer { + return nil, fmt.Errorf("%w: model must be a non-pointer struct, got pointer to %s. Use MyModel{} instead of &MyModel{}", ErrInvalidModel, finalType.Elem().Name()) + } + return model, nil +} + +// registerLocked validates the model outside the lock, then writes the model +// and its rules under a single lock acquisition so no reader can observe the +// model without its final rules. +func (r *DefaultModelRegistry) registerLocked(name string, model interface{}, rules ModelRules) error { + model, err := validateModel(model) + if err != nil { + return err } - r.models[name] = model - // Initialize with default rules if not already set - if _, exists := r.rules[name]; !exists { - r.rules[name] = DefaultModelRules() + r.mutex.Lock() + defer r.mutex.Unlock() + + if _, exists := r.models[name]; exists { + return fmt.Errorf("%w: %s", ErrModelExists, name) } + r.models[name] = model + r.rules[name] = rules return nil } +func (r *DefaultModelRegistry) RegisterModel(name string, model interface{}) error { + return r.registerLocked(name, model, DefaultModelRules()) +} + func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) { - if !r.tryRLock() { - return nil, fmt.Errorf("failed to get model %s: registry locked", name) - } + r.mutex.RLock() defer r.mutex.RUnlock() model, exists := r.models[name] if !exists { - return nil, fmt.Errorf("model %s not found", name) + return nil, fmt.Errorf("%w: %s", ErrModelNotFound, name) } return model, nil } func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} { - if !r.tryRLock() { - return make(map[string]interface{}) - } + r.mutex.RLock() defer r.mutex.RUnlock() - result := make(map[string]interface{}) + result := make(map[string]interface{}, len(r.models)) for k, v := range r.models { result[k] = v } @@ -225,9 +214,13 @@ func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} { func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) { // Try full name first fullName := fmt.Sprintf("%s.%s", schema, entity) - if model, err := r.GetModel(fullName); err == nil { + model, err := r.GetModel(fullName) + if err == nil { return model, nil } + if !errors.Is(err, ErrModelNotFound) { + return nil, err + } // Fallback to entity name only return r.GetModel(entity) @@ -238,9 +231,8 @@ func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) erro r.mutex.Lock() defer r.mutex.Unlock() - // Check if model exists if _, exists := r.models[name]; !exists { - return fmt.Errorf("model %s not found", name) + return fmt.Errorf("%w: %s", ErrModelNotFound, name) } r.rules[name] = rules @@ -253,12 +245,10 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) { r.mutex.RLock() defer r.mutex.RUnlock() - // Check if model exists if _, exists := r.models[name]; !exists { - return ModelRules{}, fmt.Errorf("model %s not found", name) + return ModelRules{}, fmt.Errorf("%w: %s", ErrModelNotFound, name) } - // Return rules if set, otherwise return default rules if rules, exists := r.rules[name]; exists { return rules, nil } @@ -266,84 +256,62 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) { return DefaultModelRules(), nil } -// RegisterModelWithRules registers a model with specific rules +// RegisterModelWithRules registers a model with specific rules atomically func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error { - // First register the model - if err := r.RegisterModel(name, model); err != nil { - return err - } - - // Then set the rules (we need to lock again for rules) - r.mutex.Lock() - defer r.mutex.Unlock() - r.rules[name] = rules - - return nil + return r.registerLocked(name, model, rules) } // Global convenience functions using the default registry // RegisterModel registers a model with the default global registry func RegisterModel(model interface{}, name string) error { - return defaultRegistry.RegisterModel(name, model) + return GetDefaultRegistry().RegisterModel(name, model) } // GetModelByName retrieves a model by searching through all registries in order // Returns the first match found func GetModelByName(name string) (interface{}, error) { - registriesMutex.RLock() - defer registriesMutex.RUnlock() - - for _, registry := range registries { + for _, registry := range registriesSnapshot() { if model, err := registry.GetModel(name); err == nil { return model, nil } } - return nil, fmt.Errorf("model %s not found in any registry", name) + return nil, fmt.Errorf("%w: %s (in any registry)", ErrModelNotFound, name) } -// IterateModels iterates over all models in the default global registry +// IterateModels iterates over all models in the default global registry. +// It iterates over a snapshot, so fn may safely call back into the registry. +// A panic in fn is recovered and logged with the model name, and iteration +// continues with the remaining models. func IterateModels(fn func(name string, model interface{})) { - defaultRegistry.mutex.RLock() - defer defaultRegistry.mutex.RUnlock() - - for name, model := range defaultRegistry.models { - fn(name, model) + for name, model := range GetDefaultRegistry().GetAllModels() { + callIsolated(name, model, fn) } } -// GetModels returns a list of all models from all registries -// Models are collected in registry order, with duplicates included -func GetModels() []interface{} { - acquired := false - for i := 0; i < lockRetryAttempts; i++ { - if registriesMutex.TryRLock() { - acquired = true - break +func callIsolated(name string, model interface{}, fn func(name string, model interface{})) { + defer func() { + if r := recover(); r != nil { + _ = logger.HandlePanic("modelregistry.IterateModels", r, "model", name) } - time.Sleep(lockRetryDelay) - } - if !acquired { - return nil - } - defer registriesMutex.RUnlock() + }() + fn(name, model) +} +// GetModels returns a list of all models from all registries. +// Only the first occurrence of each model name is included. +func GetModels() []interface{} { var models []interface{} seen := make(map[string]bool) - for _, registry := range registries { - if !registry.tryRLock() { - continue - } - for name, model := range registry.models { - // Only add the first occurrence of each model name + for _, registry := range registriesSnapshot() { + for name, model := range registry.GetAllModels() { if !seen[name] { models = append(models, model) seen[name] = true } } - registry.mutex.RUnlock() } return models @@ -351,31 +319,31 @@ func GetModels() []interface{} { // SetModelRules sets the rules for a specific model in the default registry func SetModelRules(name string, rules ModelRules) error { - return defaultRegistry.SetModelRules(name, rules) + return GetDefaultRegistry().SetModelRules(name, rules) } // GetModelRules retrieves the rules for a specific model from the default registry func GetModelRules(name string) (ModelRules, error) { - return defaultRegistry.GetModelRules(name) + return GetDefaultRegistry().GetModelRules(name) } // GetModelRulesByName retrieves the rules for a model by searching through all registries in order -// Returns the first match found +// Returns the first match found. The error wraps ErrModelNotFound when no registry has the model. func GetModelRulesByName(name string) (ModelRules, error) { - registriesMutex.RLock() - defer registriesMutex.RUnlock() - - for _, registry := range registries { - if _, err := registry.GetModel(name); err == nil { - // Model found in this registry, get its rules - return registry.GetModelRules(name) + for _, registry := range registriesSnapshot() { + rules, err := registry.GetModelRules(name) + if err == nil { + return rules, nil + } + if !errors.Is(err, ErrModelNotFound) { + return ModelRules{}, err } } - return ModelRules{}, fmt.Errorf("model %s not found in any registry", name) + return ModelRules{}, fmt.Errorf("%w: %s (in any registry)", ErrModelNotFound, name) } // RegisterModelWithRules registers a model with specific rules in the default registry func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error { - return defaultRegistry.RegisterModelWithRules(name, model, rules) + return GetDefaultRegistry().RegisterModelWithRules(name, model, rules) } diff --git a/pkg/modelregistry/model_registry_test.go b/pkg/modelregistry/model_registry_test.go new file mode 100644 index 0000000..a3a0c7d --- /dev/null +++ b/pkg/modelregistry/model_registry_test.go @@ -0,0 +1,161 @@ +package modelregistry + +import ( + "errors" + "sync" + "testing" +) + +type testModel struct{ ID int } + +type recursivePtr *recursivePtr + +func TestSentinelErrors(t *testing.T) { + r := NewModelRegistry() + if _, err := r.GetModel("nope"); !errors.Is(err, ErrModelNotFound) { + t.Fatalf("GetModel: want ErrModelNotFound, got %v", err) + } + if _, err := r.GetModelRules("nope"); !errors.Is(err, ErrModelNotFound) { + t.Fatalf("GetModelRules: want ErrModelNotFound, got %v", err) + } + if err := r.SetModelRules("nope", ModelRules{}); !errors.Is(err, ErrModelNotFound) { + t.Fatalf("SetModelRules: want ErrModelNotFound, got %v", err) + } + if err := r.RegisterModel("a", testModel{}); err != nil { + t.Fatal(err) + } + if err := r.RegisterModel("a", testModel{}); !errors.Is(err, ErrModelExists) { + t.Fatalf("want ErrModelExists, got %v", err) + } + if err := r.RegisterModel("b", nil); !errors.Is(err, ErrInvalidModel) { + t.Fatalf("want ErrInvalidModel, got %v", err) + } + if err := r.RegisterModel("c", 42); !errors.Is(err, ErrInvalidModel) { + t.Fatalf("want ErrInvalidModel, got %v", err) + } +} + +func TestRecursivePointerTypeRejected(t *testing.T) { + var x recursivePtr + if err := NewModelRegistry().RegisterModel("r", x); !errors.Is(err, ErrInvalidModel) { + t.Fatalf("want ErrInvalidModel, got %v", err) + } +} + +func TestPointerNormalised(t *testing.T) { + r := NewModelRegistry() + if err := r.RegisterModel("p", &testModel{}); err != nil { + t.Fatal(err) + } + m, _ := r.GetModel("p") + if _, ok := m.(testModel); !ok { + t.Fatalf("want testModel value, got %T", m) + } +} + +// Rules must never be observable as permissive for a restrictively registered model. +func TestRegisterModelWithRulesAtomic(t *testing.T) { + for i := 0; i < 200; i++ { + r := NewModelRegistry() + var wg sync.WaitGroup + stop := make(chan struct{}) + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-stop: + return + default: + } + if rules, err := r.GetModelRules("m"); err == nil && rules.CanDelete { + t.Error("observed permissive rules for restrictive model") + return + } + } + }() + if err := r.RegisterModelWithRules("m", testModel{}, ModelRules{CanRead: true}); err != nil { + t.Fatal(err) + } + close(stop) + wg.Wait() + } +} + +func TestIterateModelsCallbackMayReenter(t *testing.T) { + prev := GetDefaultRegistry() + reg := NewModelRegistry() + SetDefaultRegistry(reg) + defer SetDefaultRegistry(prev) + + if err := RegisterModel(testModel{}, "iter.a"); err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { + defer close(done) + IterateModels(func(name string, _ interface{}) { + _ = RegisterModel(testModel{}, name+".copy") // would deadlock if lock held + }) + }() + <-done + if _, err := reg.GetModel("iter.a.copy"); err != nil { + t.Fatal(err) + } +} + +func TestGetModelRulesByNameAcrossRegistries(t *testing.T) { + extra := NewModelRegistry() + if err := extra.RegisterModelWithRules("x.only", testModel{}, ModelRules{CanRead: true}); err != nil { + t.Fatal(err) + } + AddRegistry(extra) + rules, err := GetModelRulesByName("x.only") + if err != nil || rules.CanDelete { + t.Fatalf("rules=%+v err=%v", rules, err) + } + if _, err := GetModelRulesByName("x.missing"); !errors.Is(err, ErrModelNotFound) { + t.Fatalf("want ErrModelNotFound, got %v", err) + } +} + +func TestConcurrentAccessRace(t *testing.T) { + r := NewModelRegistry() + _ = r.RegisterModel("seed", testModel{}) + var wg sync.WaitGroup + for i := 0; i < 16; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + for j := 0; j < 100; j++ { + _ = r.SetModelRules("seed", ModelRules{CanRead: j%2 == 0}) + _, _ = r.GetModelRules("seed") + _ = r.GetAllModels() + _ = GetDefaultRegistry() + _ = GetModels() + } + }(i) + } + wg.Wait() +} + +func TestIterateModelsRecoversCallbackPanic(t *testing.T) { + prev := GetDefaultRegistry() + SetDefaultRegistry(NewModelRegistry()) + defer SetDefaultRegistry(prev) + + _ = RegisterModel(testModel{}, "p.a") + _ = RegisterModel(testModel{}, "p.b") + calls := 0 + IterateModels(func(string, interface{}) { + calls++ + panic("boom") + }) + if calls != 2 { + t.Fatalf("want both models visited, got %d", calls) + } + // Registry must still be usable (no lock left held). + if err := RegisterModel(testModel{}, "p.c"); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/security/hooks.go b/pkg/security/hooks.go index fdb7f9e..021fa50 100644 --- a/pkg/security/hooks.go +++ b/pkg/security/hooks.go @@ -2,6 +2,7 @@ package security import ( "context" + "errors" "fmt" "reflect" @@ -284,7 +285,10 @@ func checkModelUpdateAllowed(secCtx SecurityContext) error { rules, err = modelregistry.GetModelRulesByName(entity) } if err != nil { - return nil // model not registered, allow by default + if errors.Is(err, modelregistry.ErrModelNotFound) { + return nil // model not registered, allow by default + } + return err } } if !rules.CanUpdate { @@ -308,7 +312,10 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error { rules, err = modelregistry.GetModelRulesByName(entity) } if err != nil { - return nil // model not registered, allow by default + if errors.Is(err, modelregistry.ErrModelNotFound) { + return nil // model not registered, allow by default + } + return err } } if !rules.CanDelete {