Files
ResolveSpec/pkg/modelregistry/model_registry.go
T
Hein 97fe88b3a6 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.
2026-09-30 13:26:55 +02:00

350 lines
11 KiB
Go

package modelregistry
import (
"errors"
"fmt"
"reflect"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// ModelRules defines the permissions and security settings for a model
type ModelRules struct {
CanPublicRead bool // Whether the model can be read (GET operations)
CanPublicUpdate bool // Whether the model can be updated (PUT/PATCH operations)
CanPublicCreate bool // Whether the model can be created (POST operations)
CanPublicDelete bool // Whether the model can be deleted (DELETE operations)
CanRead bool // Whether the model can be read (GET operations)
CanUpdate bool // Whether the model can be updated (PUT/PATCH operations)
CanCreate bool // Whether the model can be created (POST operations)
CanDelete bool // Whether the model can be deleted (DELETE operations)
SecurityDisabled bool // Whether security checks are disabled for this model
}
// DefaultModelRules returns the default rules for a model (all operations allowed, security enabled)
func DefaultModelRules() ModelRules {
return ModelRules{
CanRead: true,
CanUpdate: true,
CanCreate: true,
CanDelete: true,
CanPublicRead: false,
CanPublicUpdate: false,
CanPublicCreate: false,
CanPublicDelete: false,
SecurityDisabled: false,
}
}
// DefaultModelRegistry implements ModelRegistry interface
type DefaultModelRegistry struct {
models map[string]interface{}
rules map[string]ModelRules
mutex sync.RWMutex
}
// Global default registry instance
var defaultRegistry = &DefaultModelRegistry{
models: make(map[string]interface{}),
rules: make(map[string]ModelRules),
}
// Global list of registries (searched in order)
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{
models: make(map[string]interface{}),
rules: make(map[string]ModelRules),
}
}
// GetDefaultRegistry returns the current default registry.
func GetDefaultRegistry() *DefaultModelRegistry {
registriesMutex.RLock()
defer registriesMutex.RUnlock()
return defaultRegistry
}
// SetDefaultRegistry replaces the default registry. A nil registry is ignored.
func SetDefaultRegistry(registry *DefaultModelRegistry) {
if registry == nil {
return
}
registriesMutex.Lock()
defer registriesMutex.Unlock()
foundAt := -1
for idx, r := range registries {
if r == defaultRegistry {
foundAt = idx
break
}
}
defaultRegistry = registry
if foundAt >= 0 {
registries[foundAt] = registry
} else {
registries = append([]*DefaultModelRegistry{registry}, registries...)
}
}
// AddRegistry adds a registry to the global list of registries
// Registries are searched in the order they were added
func AddRegistry(registry *DefaultModelRegistry) {
registriesMutex.Lock()
defer registriesMutex.Unlock()
registries = append(registries, registry)
}
// 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...)
}
// 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))
}
}()
modelType := reflect.TypeOf(model)
if modelType == nil {
return nil, fmt.Errorf("%w: model cannot be nil", ErrInvalidModel)
}
originalType := modelType
// Unwrap pointers, slices, and arrays to check the underlying type
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()
}
if modelType.Kind() != reflect.Struct {
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 {
model = reflect.New(modelType).Elem().Interface()
}
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.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) {
r.mutex.RLock()
defer r.mutex.RUnlock()
model, exists := r.models[name]
if !exists {
return nil, fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
return model, nil
}
func (r *DefaultModelRegistry) GetAllModels() map[string]interface{} {
r.mutex.RLock()
defer r.mutex.RUnlock()
result := make(map[string]interface{}, len(r.models))
for k, v := range r.models {
result[k] = v
}
return result
}
func (r *DefaultModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
// Try full name first
fullName := fmt.Sprintf("%s.%s", schema, entity)
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)
}
// SetModelRules sets the rules for a specific model
func (r *DefaultModelRegistry) SetModelRules(name string, rules ModelRules) error {
r.mutex.Lock()
defer r.mutex.Unlock()
if _, exists := r.models[name]; !exists {
return fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
r.rules[name] = rules
return nil
}
// GetModelRules retrieves the rules for a specific model
// Returns default rules if model exists but rules are not set
func (r *DefaultModelRegistry) GetModelRules(name string) (ModelRules, error) {
r.mutex.RLock()
defer r.mutex.RUnlock()
if _, exists := r.models[name]; !exists {
return ModelRules{}, fmt.Errorf("%w: %s", ErrModelNotFound, name)
}
if rules, exists := r.rules[name]; exists {
return rules, nil
}
return DefaultModelRules(), nil
}
// RegisterModelWithRules registers a model with specific rules atomically
func (r *DefaultModelRegistry) RegisterModelWithRules(name string, model interface{}, rules ModelRules) error {
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 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) {
for _, registry := range registriesSnapshot() {
if model, err := registry.GetModel(name); err == nil {
return model, nil
}
}
return nil, fmt.Errorf("%w: %s (in any registry)", ErrModelNotFound, name)
}
// 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{})) {
for name, model := range GetDefaultRegistry().GetAllModels() {
callIsolated(name, model, fn)
}
}
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)
}
}()
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 registriesSnapshot() {
for name, model := range registry.GetAllModels() {
if !seen[name] {
models = append(models, model)
seen[name] = true
}
}
}
return models
}
// SetModelRules sets the rules for a specific model in the default registry
func SetModelRules(name string, rules ModelRules) error {
return GetDefaultRegistry().SetModelRules(name, rules)
}
// GetModelRules retrieves the rules for a specific model from the default registry
func GetModelRules(name string) (ModelRules, error) {
return GetDefaultRegistry().GetModelRules(name)
}
// GetModelRulesByName retrieves the rules for a model by searching through all registries in order
// Returns the first match found. The error wraps ErrModelNotFound when no registry has the model.
func GetModelRulesByName(name string) (ModelRules, error) {
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("%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 GetDefaultRegistry().RegisterModelWithRules(name, model, rules)
}