294 lines
12 KiB
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[:])
|
|
}
|