mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-03 12:01:58 +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:
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user