Files
notarius/internal/modules/dnd/extract/enemyevents/grounding_test.go

294 lines
12 KiB
Go

package enemyevents
import (
"crypto/sha256"
"encoding/hex"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
combatturncodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/combatturns"
occurrencecodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/npcoccurrences"
npccodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/npcregistry"
scenecodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/scenedescriptions"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/npcs/identity"
sceneregistry "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/scenedescriptions/registry"
)
func TestReferenceSlotsDescribeRequiredTypedArtifacts(t *testing.T) {
slots := referenceSlots()
if len(slots) != 8 {
t.Fatalf("ReferenceSlots() count = %d, want 8", len(slots))
}
byName := make(map[string]contracts.ReferenceSlot, len(slots))
for index, slot := range slots {
if index > 0 && slots[index-1].Name > slot.Name {
t.Fatalf("ReferenceSlots() is not sorted: %#v", slots)
}
byName[slot.Name] = slot
}
for _, want := range []struct {
name string
kind contracts.ArtifactKind
}{
{NPCRegistryReferenceSlot, dnd.NPCRegistryKind},
{SceneDescriptionReferenceSlot, dnd.SceneDescriptionListKind},
{CombatTurnReferenceSlot, dnd.CombatTurnListKind},
{NPCOccurrenceReferenceSlot, dnd.NPCOccurrenceListKind},
} {
slot, ok := byName[want.name]
if !ok || !slot.Required || slot.MaxBytes != ReferenceMaxBytes || len(slot.AcceptedMediaTypes) != 1 || slot.AcceptedMediaTypes[0] != "application/json" || len(slot.AcceptedArtifactKinds) != 1 || slot.AcceptedArtifactKinds[0] != want.kind {
t.Fatalf("slot %q = %#v", want.name, slot)
}
}
slots[0].AcceptedMediaTypes[0] = "changed"
if referenceSlots()[0].AcceptedMediaTypes[0] == "changed" {
t.Fatal("ReferenceSlots() returned caller-owned storage")
}
}
func TestGroundingProducesExactSourceFreePromptInputs(t *testing.T) {
resolver, err := newGroundingResolver(groundingReferences(t, "Ashfang", dnd.SceneKindCombat))
if err != nil {
t.Fatal(err)
}
resolved, err := resolver.Resolve(contracts.ReferenceSet{})
if err != nil {
t.Fatal(err)
}
inputs := resolved.PromptInputs()
want := map[string]string{
NPCRegistryReferenceSlot: `{"npcs":[{"name":"Ashfang"}]}`,
CombatTurnReferenceSlot: `{"combat_turns":[{"actor":"Ashfang","turn_kind":"turn"},{"actor":"Aria","turn_kind":"reaction"}]}`,
NPCOccurrenceReferenceSlot: `{"npc_occurrences":[{"name":"Ashfang","kind":"combat_opponent"}]}`,
}
if len(inputs) != len(want) {
t.Fatalf("PromptInputs() = %#v", inputs)
}
for name, content := range want {
input, ok := inputs[name]
if !ok || string(input.Content) != content || input.Digest != digest([]byte(content)) || input.OriginURI != "" || input.MediaType != "application/json" {
t.Fatalf("PromptInputs()[%q] = %#v, want %q", name, input, content)
}
for _, forbidden := range []string{"source_ref", "npc-ashfang", "origin", "summary", "combat-session"} {
if strings.Contains(string(input.Content), forbidden) {
t.Fatalf("PromptInputs()[%q] leaked %q: %s", name, forbidden, input.Content)
}
}
}
if _, ok := inputs[SceneDescriptionReferenceSlot]; ok {
t.Fatal("scene descriptions were rendered as prompt grounding")
}
inputs[CombatTurnReferenceSlot] = contracts.NewLLMInputMaterial("changed", "text/plain", []byte("changed"), "", "")
if next := resolved.PromptInputs()[CombatTurnReferenceSlot]; string(next.Content) != want[CombatTurnReferenceSlot] {
t.Fatal("PromptInputs() did not return a defensive copy")
}
}
func TestGroundingResolvesGeneratedReferencesAndSceneEligibility(t *testing.T) {
prepared := groundingReferences(t, "Ashfang", dnd.SceneKindCombat)
resolver, err := newGroundingResolver(prepared)
if err != nil {
t.Fatal(err)
}
generated := groundingReferences(t, "Grimjaw", dnd.SceneKindNarrative)
resolved, err := resolver.Resolve(generated)
if err != nil {
t.Fatal(err)
}
if got := string(resolved.PromptInputs()[NPCRegistryReferenceSlot].Content); got != `{"npcs":[{"name":"Grimjaw"}]}` {
t.Fatalf("generated NPC projection = %s", got)
}
match, err := resolver.SceneMatch(generated, combatChunk())
if err != nil {
t.Fatal(err)
}
if match != (sceneregistry.ChunkMatch{State: sceneregistry.MatchExact, Kind: dnd.SceneKindNarrative}) {
t.Fatalf("generated scene match = %#v", match)
}
match, err = resolver.SceneMatch(generated, &source.Chunk{ID: "other", Ref: combatChunk().Ref})
if err != nil {
t.Fatal(err)
}
if match.State != sceneregistry.MatchMissing {
t.Fatal("missing scene was not reported")
}
mismatched := combatChunk()
mismatched.Ref.EndUnitID++
match, err = resolver.SceneMatch(generated, mismatched)
if err != nil {
t.Fatal(err)
}
if match.State != sceneregistry.MatchMismatched {
t.Fatal("mismatched scene was not reported")
}
static, err := resolver.Resolve(contracts.ReferenceSet{})
if err != nil {
t.Fatal(err)
}
if got := string(static.PromptInputs()[NPCRegistryReferenceSlot].Content); got != `{"npcs":[{"name":"Ashfang"}]}` {
t.Fatalf("static NPC projection changed after generated resolution: %s", got)
}
match, err = resolver.SceneMatch(contracts.ReferenceSet{}, combatChunk())
if err != nil {
t.Fatal(err)
}
if match != (sceneregistry.ChunkMatch{State: sceneregistry.MatchExact, Kind: dnd.SceneKindCombat}) {
t.Fatalf("static scene match = %#v", match)
}
}
func TestGroundingRejectsMissingAndInvalidReferences(t *testing.T) {
valid := groundingReferences(t, "Ashfang", dnd.SceneKindCombat)
for _, test := range []struct {
slot string
want string
}{
{NPCRegistryReferenceSlot, "NPC registry"},
{SceneDescriptionReferenceSlot, "scene descriptions"},
{CombatTurnReferenceSlot, CombatTurnReferenceSlot},
{NPCOccurrenceReferenceSlot, NPCOccurrenceReferenceSlot},
} {
t.Run("missing "+test.slot, func(t *testing.T) {
resolver, err := newGroundingResolver(withoutSlot(valid, test.slot))
if err != nil {
t.Fatal(err)
}
if test.slot == SceneDescriptionReferenceSlot {
_, err = resolver.SceneMatch(contracts.ReferenceSet{}, combatChunk())
} else {
_, err = resolver.Resolve(contracts.ReferenceSet{})
}
if err == nil || !strings.Contains(err.Error(), test.want) {
t.Fatalf("grounding error = %v, want missing %q reference", err, test.slot)
}
})
}
tests := []struct {
name string
set contracts.ReferenceSet
}{
{"multiple items", replaceSlot(valid, CombatTurnReferenceSlot, contracts.ResolvedReferenceSlot{Items: []contracts.ReferenceItem{{MediaType: "application/json"}, {MediaType: "application/json"}}})},
{"malformed durable content", replaceItem(valid, NPCOccurrenceReferenceSlot, contracts.ReferenceItem{MediaType: "application/json", Content: []byte(`{}`)})},
{"wrong media type", replaceItem(valid, CombatTurnReferenceSlot, contracts.ReferenceItem{MediaType: "text/plain", Content: valid.Slots[CombatTurnReferenceSlot].Items[0].Content})},
{"oversize", replaceItem(valid, CombatTurnReferenceSlot, contracts.ReferenceItem{MediaType: "application/json", Content: make([]byte, ReferenceMaxBytes+1)})},
{"reversed combat-turn evidence", reversedCombatTurnReference(t, valid)},
}
resolver, err := newGroundingResolver(valid)
if err != nil {
t.Fatal(err)
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if _, err := resolver.Resolve(test.set); err == nil {
t.Fatal("Resolve() error = nil")
}
})
}
}
func reversedCombatTurnReference(t *testing.T, references contracts.ReferenceSet) contracts.ReferenceSet {
t.Helper()
value := dnd.CombatTurnList{CombatTurns: []dnd.CombatTurn{{
Actor: "Ashfang", TurnKind: dnd.CombatTurnKindTurn,
SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 2, EndUnitID: 1}},
}}}
content, err := combatturncodec.New().EncodeCandidate(value)
if err != nil {
t.Fatalf("EncodeCandidate() error = %v", err)
}
return replaceItem(references, CombatTurnReferenceSlot, newReferenceItem(CombatTurnReferenceSlot, content))
}
func groundingReferences(t *testing.T, enemy string, sceneKind dnd.SceneKind) contracts.ReferenceSet {
t.Helper()
npcContent, err := npccodec.New().Encode(dnd.NPCRegistry{NPCs: []dnd.NPC{{
ID: identity.DeriveID(enemy),
Name: enemy,
SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 1}},
}}})
if err != nil {
t.Fatal(err)
}
sceneContent, err := scenecodec.New().Encode(dnd.SceneDescriptionList{Scenes: []dnd.SceneDescription{{
ID: "combat-scene",
SourceRef: source.SourceRef{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 4},
Kind: sceneKind,
Title: "Combat session",
Summary: "A detailed scene summary that must not reach the prompt.",
}}})
if err != nil {
t.Fatal(err)
}
turnContent, err := combatturncodec.New().Encode(dnd.CombatTurnList{CombatTurns: []dnd.CombatTurn{
{Actor: enemy, TurnKind: dnd.CombatTurnKindTurn, SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 1}}},
{Actor: "Aria", TurnKind: dnd.CombatTurnKindReaction, SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 2, EndUnitID: 2}}},
}})
if err != nil {
t.Fatal(err)
}
occurrenceContent, err := occurrencecodec.New().Encode(dnd.NPCOccurrenceList{Occurrences: []dnd.NPCOccurrence{
{NPCID: identity.DeriveID(enemy), Name: enemy, Kind: dnd.NPCOccurrenceKindCombatOpponent, SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 1}}},
{NPCID: identity.DeriveID("Aria"), Name: "Aria", Kind: dnd.NPCOccurrenceKindCombatAlly, SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 2, EndUnitID: 2}}},
}})
if err != nil {
t.Fatal(err)
}
return contracts.ReferenceSet{Slots: map[string]contracts.ResolvedReferenceSlot{
NPCRegistryReferenceSlot: {Items: []contracts.ReferenceItem{newReferenceItem(NPCRegistryReferenceSlot, npcContent)}},
SceneDescriptionReferenceSlot: {Items: []contracts.ReferenceItem{newReferenceItem(SceneDescriptionReferenceSlot, sceneContent)}},
CombatTurnReferenceSlot: {Items: []contracts.ReferenceItem{newReferenceItem(CombatTurnReferenceSlot, turnContent)}},
NPCOccurrenceReferenceSlot: {Items: []contracts.ReferenceItem{newReferenceItem(NPCOccurrenceReferenceSlot, occurrenceContent)}},
}}
}
func newReferenceItem(slot string, content []byte) contracts.ReferenceItem {
return contracts.ReferenceItem{
SlotName: slot,
MediaType: "application/json",
Content: append([]byte(nil), content...),
Digest: digest(content),
Origin: contracts.ReferenceOrigin{Type: "file", URI: "file:///private/reference.json"},
Producer: contracts.ReferenceProducer{PipelineID: "prior", StepID: "extract", LaneID: "lane", ModuleKey: "dnd/example"},
}
}
func withoutSlot(set contracts.ReferenceSet, name string) contracts.ReferenceSet {
cloned := contracts.ReferenceSet{Slots: make(map[string]contracts.ResolvedReferenceSlot, len(set.Slots)-1)}
for slotName, slot := range set.Slots {
if slotName != name {
cloned.Slots[slotName] = slot
}
}
return cloned
}
func replaceSlot(set contracts.ReferenceSet, name string, value contracts.ResolvedReferenceSlot) contracts.ReferenceSet {
cloned := contracts.ReferenceSet{Slots: make(map[string]contracts.ResolvedReferenceSlot, len(set.Slots))}
for slotName, slot := range set.Slots {
cloned.Slots[slotName] = slot
}
cloned.Slots[name] = value
return cloned
}
func replaceItem(set contracts.ReferenceSet, name string, item contracts.ReferenceItem) contracts.ReferenceSet {
return replaceSlot(set, name, contracts.ResolvedReferenceSlot{Items: []contracts.ReferenceItem{item}})
}
func combatChunk() *source.Chunk {
return &source.Chunk{ID: "combat-scene", Ref: source.SourceRef{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 4}}
}
func digest(content []byte) string {
sum := sha256.Sum256(content)
return "sha256:" + hex.EncodeToString(sum[:])
}