Migrate location registry to shared resolver

This commit is contained in:
2026-08-04 13:30:12 +00:00
parent 7f28899730
commit f5fd115046
2 changed files with 102 additions and 109 deletions

View File

@@ -7,9 +7,6 @@ import (
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"fmt" "fmt"
"mime"
"strings"
"sync"
"gitea.maximumdirect.net/eric/notarius/internal/core/source" "gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
@@ -17,6 +14,7 @@ import (
locationcodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/locations" locationcodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/locations"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/locations/identity" "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/locations/identity"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared/diagnostics" "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared/diagnostics"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared/registryresolver"
) )
const ( const (
@@ -37,38 +35,20 @@ type Registry struct {
lookupByID map[string]int lookupByID map[string]int
} }
// Resolver retains the construction-time registry and memoizes immutable // Resolver selects and memoizes immutable location registry views.
// operation-time registries. It never retains caller-owned reference bytes.
type Resolver struct { type Resolver struct {
seeded *Registry resolver *registryresolver.Resolver[*Registry]
mu sync.Mutex
cache map[string]*Registry
rawCache map[string]*Registry
} }
// NewResolver validates a materialized construction-time location reference. // NewResolver validates a materialized construction-time location reference.
// An empty slot is permitted because a generated reference is supplied only at // An empty slot is permitted because a generated reference is supplied only at
// operation time. // operation time.
func NewResolver(references contracts.ReferenceSet) (*Resolver, error) { func NewResolver(references contracts.ReferenceSet) (*Resolver, error) {
seeded, err := Resolve(constructionReferences(references)) resolver, err := registryresolver.New(registryResolverConfig(), references)
if err != nil { if err != nil {
return nil, err return nil, err
} }
return &Resolver{seeded: seeded, cache: make(map[string]*Registry), rawCache: make(map[string]*Registry)}, nil return &Resolver{resolver: resolver}, nil
}
func constructionReferences(references contracts.ReferenceSet) contracts.ReferenceSet {
slot, ok := references.Slots[ReferenceSlot]
if !ok || len(slot.Items) > 0 {
return references
}
cloned := contracts.ReferenceSet{Slots: make(map[string]contracts.ResolvedReferenceSlot, len(references.Slots))}
for name, value := range references.Slots {
cloned.Slots[name] = value
}
delete(cloned.Slots, ReferenceSlot)
return cloned
} }
// Seeded returns the validated construction-time registry. // Seeded returns the validated construction-time registry.
@@ -76,67 +56,49 @@ func (r *Resolver) Seeded() *Registry {
if r == nil { if r == nil {
return nil return nil
} }
return r.seeded return r.resolver.Seeded()
} }
// Resolve returns the generated operation-time registry when the locations // Resolve returns the generated operation-time registry when the locations
// slot is present, otherwise it returns the construction-time registry. // slot is present, otherwise it returns the construction-time registry.
func (r *Resolver) Resolve(references contracts.ReferenceSet) (*Registry, error) { func (r *Resolver) Resolve(references contracts.ReferenceSet) (*Registry, error) {
if r == nil { if r == nil || r.resolver == nil {
return Resolve(references) return Resolve(references)
} }
if _, ok := references.Slots[ReferenceSlot]; !ok { return r.resolver.Resolve(references)
return r.seeded, nil
}
slot := references.Slots[ReferenceSlot]
rawKey := ""
if len(slot.Items) == 1 {
rawKey = strings.ToLower(strings.TrimSpace(slot.Items[0].MediaType)) + "\x00" + semanticDigest(slot.Items[0].Content)
}
r.mu.Lock()
defer r.mu.Unlock()
if rawKey != "" {
if cached, ok := r.rawCache[rawKey]; ok {
return cached, nil
}
}
resolved, err := Resolve(references)
if err != nil {
return nil, err
}
if sameRegistryIdentity(r.seeded, resolved) {
if rawKey != "" {
r.rawCache[rawKey] = r.seeded
}
return r.seeded, nil
}
if cached, ok := r.cache[resolved.Digest()]; ok {
if rawKey != "" {
r.rawCache[rawKey] = cached
}
return cached, nil
}
r.cache[resolved.Digest()] = resolved
if rawKey != "" {
r.rawCache[rawKey] = resolved
}
return resolved, nil
}
func sameRegistryIdentity(first, second *Registry) bool {
if first == nil || second == nil {
return first == second
}
return first.bound == second.bound && first.digest == second.digest
} }
// Resolve validates an optional location-list reference. An absent reference // Resolve validates an optional location-list reference. An absent reference
// uses the canonical empty projection and has no durable registry identity. // uses the canonical empty projection and has no durable registry identity.
func Resolve(references contracts.ReferenceSet) (*Registry, error) { func Resolve(references contracts.ReferenceSet) (*Registry, error) {
slot, ok := references.Slots[ReferenceSlot] item, present, err := registryresolver.ResolveOptionalSingleItem(references, locationReferenceSpec())
if !ok { if err != nil {
return nil, err
}
if !present {
return emptyRegistry(), nil
}
return loadRegistry(item.Content)
}
func registryResolverConfig() registryresolver.Config[*Registry] {
return registryresolver.Config[*Registry]{
Reference: locationReferenceSpec(),
Absent: func() (*Registry, error) {
return emptyRegistry(), nil
},
Load: loadRegistry,
SemanticIdentity: func(registry *Registry) string {
return registry.Digest()
},
}
}
func locationReferenceSpec() registryresolver.ReferenceSpec {
return registryresolver.ReferenceSpec{SlotName: ReferenceSlot, AcceptedMediaType: locationcodec.MediaType, MaxBytes: MaxBytes}
}
func emptyRegistry() *Registry {
content := []byte(emptyPrompt) content := []byte(emptyPrompt)
projectionDigest := semanticDigest(content) projectionDigest := semanticDigest(content)
return &Registry{ return &Registry{
@@ -145,26 +107,12 @@ func Resolve(references contracts.ReferenceSet) (*Registry, error) {
projectionDigest: projectionDigest, projectionDigest: projectionDigest,
promptInput: contracts.NewLLMInputMaterial(ReferenceSlot, locationcodec.MediaType, content, projectionDigest, ""), promptInput: contracts.NewLLMInputMaterial(ReferenceSlot, locationcodec.MediaType, content, projectionDigest, ""),
lookupByID: map[string]int{}, lookupByID: map[string]int{},
}, nil
}
if len(slot.Items) != 1 {
return nil, fmt.Errorf("reference slot %q must contain exactly one item", ReferenceSlot)
}
item := slot.Items[0]
mediaType, _, err := mime.ParseMediaType(item.MediaType)
if err != nil {
return nil, fmt.Errorf("reference slot %q item media type is invalid", ReferenceSlot)
}
if !strings.EqualFold(mediaType, locationcodec.MediaType) {
return nil, fmt.Errorf("reference slot %q item media type must be %s", ReferenceSlot, locationcodec.MediaType)
}
if len(item.Content) > MaxBytes {
return nil, fmt.Errorf("reference slot %q item is %d bytes, limit %d", ReferenceSlot, len(item.Content), MaxBytes)
} }
}
func loadRegistry(referenceContent []byte) (*Registry, error) {
codec := locationcodec.New() codec := locationcodec.New()
value, err := codec.Decode(item.Content) value, err := codec.Decode(referenceContent)
if err != nil { if err != nil {
return nil, fmt.Errorf("decode location registry: invalid approved location JSON") return nil, fmt.Errorf("decode location registry: invalid approved location JSON")
} }

View File

@@ -100,18 +100,47 @@ func TestResolveRejectsInvalidReferenceInputs(t *testing.T) {
} }
} }
func TestNewResolverValidatesStaticReferenceBeforeOperations(t *testing.T) { func TestResolverValidatesStaticAndOperationReferences(t *testing.T) {
if _, err := NewResolver(referenceSet(item([]byte(`{"locations":[`)))); err == nil || !strings.Contains(err.Error(), "invalid approved") { placeholder, err := NewResolver(referenceSet())
t.Fatalf("NewResolver() error = %v, want malformed static reference failure", err) if err != nil || placeholder.Seeded().Bound() {
t.Fatalf("generated placeholder = %#v, %v; want unbound seed", placeholder, err)
} }
resolver, err := NewResolver(referenceSet(item(encodeList(t, registryFixture()))))
if err != nil || !resolver.Seeded().Bound() || resolver.Seeded().Count() != 2 { validContent := encodeList(t, registryFixture())
t.Fatalf("NewResolver() = %#v, %v", resolver, err) invalidSets := []contracts.ReferenceSet{
referenceSet(item([]byte(`{"locations":[`))),
referenceSet(contracts.ReferenceItem{MediaType: "text/plain", Content: validContent}),
referenceSet(item(make([]byte, MaxBytes+1))),
referenceSet(item(validContent), item(validContent)),
}
for index, references := range invalidSets {
if _, err := NewResolver(references); err == nil {
t.Fatalf("NewResolver(invalid %d) error = nil", index)
}
if _, err := placeholder.Resolve(references); err == nil {
t.Fatalf("Resolve(invalid %d) error = nil", index)
}
}
staticContent := append([]byte(nil), validContent...)
staticReferences := referenceSet(item(staticContent))
seeded, err := NewResolver(staticReferences)
if err != nil {
t.Fatal(err)
}
resolved, err := seeded.Resolve(staticReferences)
if err != nil || resolved != seeded.Seeded() {
t.Fatalf("seed reuse = %p / %p, %v", resolved, seeded.Seeded(), err)
}
staticContent[0] = '['
delete(staticReferences.Slots, ReferenceSlot)
if seeded.Seeded().Count() != 2 || seeded.Seeded().CanonicalBytes()[0] != '{' {
t.Fatalf("seeded registry retained construction references: %#v", seeded.Seeded())
} }
} }
func TestResolverCachesRawAndSemanticallyEquivalentRegistriesConcurrently(t *testing.T) { func TestResolverCachesEquivalentRegistriesConcurrentlyAndIgnoresCallerDigest(t *testing.T) {
resolver, err := NewResolver(contracts.ReferenceSet{}) resolver, err := NewResolver(referenceSet())
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -130,7 +159,7 @@ func TestResolverCachesRawAndSemanticallyEquivalentRegistriesConcurrently(t *tes
} }
spaced := append([]byte("\n "), content...) spaced := append([]byte("\n "), content...)
spaced = append(spaced, '\n') spaced = append(spaced, '\n')
third, err := resolver.Resolve(referenceSet(item(spaced))) third, err := resolver.Resolve(referenceSet(contracts.ReferenceItem{MediaType: "APPLICATION/JSON; charset=utf-8", Content: spaced}))
if err != nil || third != first { if err != nil || third != first {
t.Fatalf("semantic cache Resolve() = %p, %p, %v", first, third, err) t.Fatalf("semantic cache Resolve() = %p, %p, %v", first, third, err)
} }
@@ -153,13 +182,29 @@ func TestResolverCachesRawAndSemanticallyEquivalentRegistriesConcurrently(t *tes
t.Error(err) t.Error(err)
} }
firstSet.Slots[ReferenceSlot].Items[0].Content[0] = '[' sharedDigest := "sha256:" + strings.Repeat("0", 64)
if registry, err := resolver.Resolve(contracts.ReferenceSet{}); err != nil || registry != resolver.Seeded() || registry.Count() != 0 { firstItem := contracts.ReferenceItem{MediaType: locationcodec.MediaType, Content: content, Digest: sharedDigest}
t.Fatalf("caller bytes affected resolver: %#v, %v", registry, err) otherList := registryFixture()
otherList.Locations[0].Name = "Moon Gate"
otherList.Locations[0].ID = identity.DeriveID(otherList.Locations[0].Name, otherList.Locations[0].SourceRefs)
otherContent := encodeList(t, otherList)
otherItem := contracts.ReferenceItem{MediaType: locationcodec.MediaType, Content: otherContent, Digest: sharedDigest}
byDigestFirst, err := resolver.Resolve(referenceSet(firstItem))
if err != nil {
t.Fatal(err)
} }
byDigestOther, err := resolver.Resolve(referenceSet(otherItem))
if err != nil || byDigestFirst == byDigestOther || byDigestFirst.Digest() == byDigestOther.Digest() {
t.Fatalf("caller digest aliased different registries: %p / %p, %v", byDigestFirst, byDigestOther, err)
}
firstSet.Slots[ReferenceSlot].Items[0].Content[0] = '['
if got, ok := first.Lookup(registryFixture().Locations[0].ID); !ok || got.Name != "The Tavern" { if got, ok := first.Lookup(registryFixture().Locations[0].ID); !ok || got.Name != "The Tavern" {
t.Fatalf("cached registry retained caller bytes: %#v, %t", got, ok) t.Fatalf("cached registry retained caller bytes: %#v, %t", got, ok)
} }
if fallback, err := resolver.Resolve(contracts.ReferenceSet{}); err != nil || fallback != resolver.Seeded() || fallback.Bound() {
t.Fatalf("fallback = %#v, %v; want unbound seed", fallback, err)
}
} }
func registryFixture() dnd.LocationList { func registryFixture() dnd.LocationList {