mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 11:31:57 +00:00
fix(modelregistry): address audit findings
Replace try-lock/sleep scheme with blocking locks, add ErrModelNotFound/ ErrModelExists/ErrInvalidModel sentinels, make RegisterModelWithRules atomic, snapshot in IterateModels, guard defaultRegistry access, cap the pointer-unwrap depth, and recover panics in callbacks and reflection. Security hooks now allow-by-default only on ErrModelNotFound. Add tests.
This commit is contained in:
@@ -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 |
|
| 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 |
|
| 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 |
|
| 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 |
|
| 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
|
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/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/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/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/tracing` | `tracer` global | `tracing.audit.md` finding 5 |
|
||||||
| `pkg/errortracking` | `sentry.Init` mutates process globals | `errortracking.audit.md` finding 2 |
|
| `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 |
|
| `metrics` | 1 | 64 | no |
|
||||||
| `resolvemcp` | 1 | 34 | no |
|
| `resolvemcp` | 1 | 34 | no |
|
||||||
| `logger` | 0 | 0 | — |
|
| `logger` | 0 | 0 | — |
|
||||||
| `modelregistry` | 0 | 0 | — |
|
| `modelregistry` | 1 | ~150 | yes (`-race`) *(added 2026-09-30)* |
|
||||||
| `testmodels` | 0 | 0 | — |
|
| `testmodels` | 0 | 0 | — |
|
||||||
| `tracing` | 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/cache` | `defaultCache *Cache` (`cache.go:10`) | **no** |
|
||||||
| `pkg/config` | `configInstance *Manager` (`manager.go:15`) | **no** |
|
| `pkg/config` | `configInstance *Manager` (`manager.go:15`) | **no** |
|
||||||
| `pkg/tracing` | `tracer` (`tracing.go:19`) | **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` | `globalProvider` (`interfaces.go:50-51`) | **yes** — `globalProviderMu sync.RWMutex` |
|
||||||
|
|
||||||
`pkg/metrics` is the model the others should follow:
|
`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"
|
only **Critical** authorization finding: `GetModel` returns a "registry locked"
|
||||||
error under write-lock contention, which `security/hooks.go:274-294` converts
|
error under write-lock contention, which `security/hooks.go:274-294` converts
|
||||||
into `return nil // model not registered, allow by default`
|
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
|
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.
|
immediately. The package guards a security boundary and has never been tested.
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,22 @@ There are **zero tests** in this package.
|
|||||||
| 11 | Low | Slowness | `os.Getpid()` called on every log line |
|
| 11 | Low | Slowness | `os.Getpid()` called on every log line |
|
||||||
| 12 | Low | Observability | `UpdateLogger` build failure degrades silently to stdlib `log` |
|
| 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
|
## Findings
|
||||||
|
|||||||
@@ -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) |
|
| 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 |
|
| 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
|
## Findings
|
||||||
|
|||||||
@@ -1787,7 +1787,7 @@ func checkModelUpdateAllowed(secCtx SecurityContext) error {
|
|||||||
`checkModelDeleteAllowed` is identical (`:298-318`, fail-open at `:311`). A model
|
`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
|
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
|
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".
|
`return nil` combine into "unknown model ⇒ permitted".
|
||||||
|
|
||||||
**Any authenticated user may perform any operation.** `CheckModelAuthAllowed` is
|
**Any authenticated user may perform any operation.** `CheckModelAuthAllowed` is
|
||||||
|
|||||||
+115
-147
@@ -1,10 +1,12 @@
|
|||||||
package modelregistry
|
package modelregistry
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ModelRules defines the permissions and security settings for a model
|
// ModelRules defines the permissions and security settings for a model
|
||||||
@@ -52,6 +54,21 @@ var defaultRegistry = &DefaultModelRegistry{
|
|||||||
var registries = []*DefaultModelRegistry{defaultRegistry}
|
var registries = []*DefaultModelRegistry{defaultRegistry}
|
||||||
var registriesMutex sync.RWMutex
|
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
|
// NewModelRegistry creates a new model registry
|
||||||
func NewModelRegistry() *DefaultModelRegistry {
|
func NewModelRegistry() *DefaultModelRegistry {
|
||||||
return &DefaultModelRegistry{
|
return &DefaultModelRegistry{
|
||||||
@@ -60,44 +77,19 @@ func NewModelRegistry() *DefaultModelRegistry {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// lockRetryAttempts/lockRetryDelay bound how long the try-lock helpers below
|
// GetDefaultRegistry returns the current default registry.
|
||||||
// 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.
|
|
||||||
func GetDefaultRegistry() *DefaultModelRegistry {
|
func GetDefaultRegistry() *DefaultModelRegistry {
|
||||||
for i := 0; i < lockRetryAttempts; i++ {
|
registriesMutex.RLock()
|
||||||
if registriesMutex.TryRLock() {
|
defer registriesMutex.RUnlock()
|
||||||
defer registriesMutex.RUnlock()
|
|
||||||
return defaultRegistry
|
|
||||||
}
|
|
||||||
time.Sleep(lockRetryDelay)
|
|
||||||
}
|
|
||||||
return defaultRegistry
|
return defaultRegistry
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetDefaultRegistry replaces the default registry. It uses a bounded
|
// SetDefaultRegistry replaces the default registry. A nil registry is ignored.
|
||||||
// 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.
|
|
||||||
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
func SetDefaultRegistry(registry *DefaultModelRegistry) {
|
||||||
acquired := false
|
if registry == nil {
|
||||||
for i := 0; i < lockRetryAttempts; i++ {
|
|
||||||
if registriesMutex.TryLock() {
|
|
||||||
acquired = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
time.Sleep(lockRetryDelay)
|
|
||||||
}
|
|
||||||
if !acquired {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
registriesMutex.Lock()
|
||||||
defer registriesMutex.Unlock()
|
defer registriesMutex.Unlock()
|
||||||
|
|
||||||
foundAt := -1
|
foundAt := -1
|
||||||
@@ -123,99 +115,96 @@ func AddRegistry(registry *DefaultModelRegistry) {
|
|||||||
registries = append(registries, registry)
|
registries = append(registries, registry)
|
||||||
}
|
}
|
||||||
|
|
||||||
// tryLock attempts to acquire the registry's write lock, retrying briefly.
|
// registriesSnapshot returns a copy of the registry list so callers can
|
||||||
// Returns false if it could not be acquired within the bound.
|
// iterate without holding registriesMutex.
|
||||||
func (r *DefaultModelRegistry) tryLock() bool {
|
func registriesSnapshot() []*DefaultModelRegistry {
|
||||||
for i := 0; i < lockRetryAttempts; i++ {
|
registriesMutex.RLock()
|
||||||
if r.mutex.TryLock() {
|
defer registriesMutex.RUnlock()
|
||||||
return true
|
return append([]*DefaultModelRegistry(nil), registries...)
|
||||||
}
|
|
||||||
time.Sleep(lockRetryDelay)
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// tryRLock attempts to acquire the registry's read lock, retrying briefly.
|
// validateModel checks the model is a struct (or pointer/slice/array of one)
|
||||||
// Returns false if it could not be acquired within the bound.
|
// and returns the normalised non-pointer struct value. It takes no locks.
|
||||||
func (r *DefaultModelRegistry) tryRLock() bool {
|
func validateModel(model interface{}) (result interface{}, err error) {
|
||||||
for i := 0; i < lockRetryAttempts; i++ {
|
// Reflection on a pathological type must fail the registration, not crash the process.
|
||||||
if r.mutex.TryRLock() {
|
defer func() {
|
||||||
return true
|
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)
|
modelType := reflect.TypeOf(model)
|
||||||
if modelType == nil {
|
if modelType == nil {
|
||||||
return fmt.Errorf("model cannot be nil")
|
return nil, fmt.Errorf("%w: model cannot be nil", ErrInvalidModel)
|
||||||
}
|
}
|
||||||
|
|
||||||
originalType := modelType
|
originalType := modelType
|
||||||
|
|
||||||
// Unwrap pointers, slices, and arrays to check the underlying type
|
// 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()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate that the underlying type is a struct
|
|
||||||
if modelType.Kind() != reflect.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 a pointer/slice/array was passed, unwrap to the base struct
|
||||||
if originalType != modelType {
|
if originalType != modelType {
|
||||||
// Create a zero value of the struct type
|
|
||||||
model = reflect.New(modelType).Elem().Interface()
|
model = reflect.New(modelType).Elem().Interface()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Additional check: ensure model is not a pointer
|
if finalType := reflect.TypeOf(model); finalType.Kind() == reflect.Pointer {
|
||||||
finalType := reflect.TypeOf(model)
|
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())
|
||||||
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())
|
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
|
r.mutex.Lock()
|
||||||
// Initialize with default rules if not already set
|
defer r.mutex.Unlock()
|
||||||
if _, exists := r.rules[name]; !exists {
|
|
||||||
r.rules[name] = DefaultModelRules()
|
if _, exists := r.models[name]; exists {
|
||||||
|
return fmt.Errorf("%w: %s", ErrModelExists, name)
|
||||||
}
|
}
|
||||||
|
r.models[name] = model
|
||||||
|
r.rules[name] = rules
|
||||||
return nil
|
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) {
|
func (r *DefaultModelRegistry) GetModel(name string) (interface{}, error) {
|
||||||
if !r.tryRLock() {
|
r.mutex.RLock()
|
||||||
return nil, fmt.Errorf("failed to get model %s: registry locked", name)
|
|
||||||
}
|
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
model, exists := r.models[name]
|
model, exists := r.models[name]
|
||||||
if !exists {
|
if !exists {
|
||||||
return nil, fmt.Errorf("model %s not found", name)
|
return nil, fmt.Errorf("%w: %s", ErrModelNotFound, name)
|
||||||
}
|
}
|
||||||
|
|
||||||
return model, nil
|
return model, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
||||||
if !r.tryRLock() {
|
r.mutex.RLock()
|
||||||
return make(map[string]interface{})
|
|
||||||
}
|
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
result := make(map[string]interface{})
|
result := make(map[string]interface{}, len(r.models))
|
||||||
for k, v := range r.models {
|
for k, v := range r.models {
|
||||||
result[k] = v
|
result[k] = v
|
||||||
}
|
}
|
||||||
@@ -225,9 +214,13 @@ func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
|
|||||||
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
||||||
// Try full name first
|
// Try full name first
|
||||||
fullName := fmt.Sprintf("%s.%s", schema, entity)
|
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
|
return model, nil
|
||||||
}
|
}
|
||||||
|
if !errors.Is(err, ErrModelNotFound) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
// Fallback to entity name only
|
// Fallback to entity name only
|
||||||
return r.GetModel(entity)
|
return r.GetModel(entity)
|
||||||
@@ -238,9 +231,8 @@ func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) erro
|
|||||||
r.mutex.Lock()
|
r.mutex.Lock()
|
||||||
defer r.mutex.Unlock()
|
defer r.mutex.Unlock()
|
||||||
|
|
||||||
// Check if model exists
|
|
||||||
if _, exists := r.models[name]; !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
|
r.rules[name] = rules
|
||||||
@@ -253,12 +245,10 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
|
|||||||
r.mutex.RLock()
|
r.mutex.RLock()
|
||||||
defer r.mutex.RUnlock()
|
defer r.mutex.RUnlock()
|
||||||
|
|
||||||
// Check if model exists
|
|
||||||
if _, exists := r.models[name]; !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 {
|
if rules, exists := r.rules[name]; exists {
|
||||||
return rules, nil
|
return rules, nil
|
||||||
}
|
}
|
||||||
@@ -266,84 +256,62 @@ func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
|
|||||||
return DefaultModelRules(), nil
|
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 {
|
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
|
||||||
// First register the model
|
return r.registerLocked(name, model, rules)
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Global convenience functions using the default registry
|
// Global convenience functions using the default registry
|
||||||
|
|
||||||
// RegisterModel registers a model with the default global registry
|
// RegisterModel registers a model with the default global registry
|
||||||
func RegisterModel(model interface{}, name string) error {
|
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
|
// GetModelByName retrieves a model by searching through all registries in order
|
||||||
// Returns the first match found
|
// Returns the first match found
|
||||||
func GetModelByName(name string) (interface{}, error) {
|
func GetModelByName(name string) (interface{}, error) {
|
||||||
registriesMutex.RLock()
|
for _, registry := range registriesSnapshot() {
|
||||||
defer registriesMutex.RUnlock()
|
|
||||||
|
|
||||||
for _, registry := range registries {
|
|
||||||
if model, err := registry.GetModel(name); err == nil {
|
if model, err := registry.GetModel(name); err == nil {
|
||||||
return model, 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{})) {
|
func IterateModels(fn func(name string, model interface{})) {
|
||||||
defaultRegistry.mutex.RLock()
|
for name, model := range GetDefaultRegistry().GetAllModels() {
|
||||||
defer defaultRegistry.mutex.RUnlock()
|
callIsolated(name, model, fn)
|
||||||
|
|
||||||
for name, model := range defaultRegistry.models {
|
|
||||||
fn(name, model)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetModels returns a list of all models from all registries
|
func callIsolated(name string, model interface{}, fn func(name string, model interface{})) {
|
||||||
// Models are collected in registry order, with duplicates included
|
defer func() {
|
||||||
func GetModels() []interface{} {
|
if r := recover(); r != nil {
|
||||||
acquired := false
|
_ = logger.HandlePanic("modelregistry.IterateModels", r, "model", name)
|
||||||
for i := 0; i < lockRetryAttempts; i++ {
|
|
||||||
if registriesMutex.TryRLock() {
|
|
||||||
acquired = true
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
time.Sleep(lockRetryDelay)
|
}()
|
||||||
}
|
fn(name, model)
|
||||||
if !acquired {
|
}
|
||||||
return nil
|
|
||||||
}
|
|
||||||
defer registriesMutex.RUnlock()
|
|
||||||
|
|
||||||
|
// 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{}
|
var models []interface{}
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
for _, registry := range registries {
|
for _, registry := range registriesSnapshot() {
|
||||||
if !registry.tryRLock() {
|
for name, model := range registry.GetAllModels() {
|
||||||
continue
|
|
||||||
}
|
|
||||||
for name, model := range registry.models {
|
|
||||||
// Only add the first occurrence of each model name
|
|
||||||
if !seen[name] {
|
if !seen[name] {
|
||||||
models = append(models, model)
|
models = append(models, model)
|
||||||
seen[name] = true
|
seen[name] = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
registry.mutex.RUnlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return models
|
return models
|
||||||
@@ -351,31 +319,31 @@ func GetModels() []interface{} {
|
|||||||
|
|
||||||
// SetModelRules sets the rules for a specific model in the default registry
|
// SetModelRules sets the rules for a specific model in the default registry
|
||||||
func SetModelRules(name string, rules ModelRules) error {
|
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
|
// GetModelRules retrieves the rules for a specific model from the default registry
|
||||||
func GetModelRules(name string) (ModelRules, error) {
|
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
|
// 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) {
|
func GetModelRulesByName(name string) (ModelRules, error) {
|
||||||
registriesMutex.RLock()
|
for _, registry := range registriesSnapshot() {
|
||||||
defer registriesMutex.RUnlock()
|
rules, err := registry.GetModelRules(name)
|
||||||
|
if err == nil {
|
||||||
for _, registry := range registries {
|
return rules, nil
|
||||||
if _, err := registry.GetModel(name); err == nil {
|
}
|
||||||
// Model found in this registry, get its rules
|
if !errors.Is(err, ErrModelNotFound) {
|
||||||
return registry.GetModelRules(name)
|
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
|
// RegisterModelWithRules registers a model with specific rules in the default registry
|
||||||
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
|
func RegisterModelWithRules(model interface{}, name string, rules ModelRules) error {
|
||||||
return defaultRegistry.RegisterModelWithRules(name, model, rules)
|
return GetDefaultRegistry().RegisterModelWithRules(name, model, rules)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -2,6 +2,7 @@ package security
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
@@ -284,7 +285,10 @@ func checkModelUpdateAllowed(secCtx SecurityContext) error {
|
|||||||
rules, err = modelregistry.GetModelRulesByName(entity)
|
rules, err = modelregistry.GetModelRulesByName(entity)
|
||||||
}
|
}
|
||||||
if err != nil {
|
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 {
|
if !rules.CanUpdate {
|
||||||
@@ -308,7 +312,10 @@ func checkModelDeleteAllowed(secCtx SecurityContext) error {
|
|||||||
rules, err = modelregistry.GetModelRulesByName(entity)
|
rules, err = modelregistry.GetModelRulesByName(entity)
|
||||||
}
|
}
|
||||||
if err != nil {
|
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 {
|
if !rules.CanDelete {
|
||||||
|
|||||||
Reference in New Issue
Block a user