Files
ResolveSpec/pkg/modelregistry/model_registry_test.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

162 lines
4.1 KiB
Go

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)
}
}