Prepare D&D enemy event grounding references
This commit is contained in:
278
internal/modules/dnd/extract/enemyevents/grounding.go
Normal file
278
internal/modules/dnd/extract/enemyevents/grounding.go
Normal file
@@ -0,0 +1,278 @@
|
||||
// Package enemyevents prepares validated D&D combat grounding for enemy-event
|
||||
// extraction.
|
||||
package enemyevents
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"mime"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"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"
|
||||
interactioncodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/npcinteractions"
|
||||
npcregistry "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/npcs/registry"
|
||||
sceneregistry "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/scenedescriptions/registry"
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared"
|
||||
)
|
||||
|
||||
const (
|
||||
NPCRegistryReferenceSlot = npcregistry.ReferenceSlot
|
||||
SceneDescriptionReferenceSlot = sceneregistry.ReferenceSlot
|
||||
CombatTurnReferenceSlot = "combat_turns"
|
||||
NPCInteractionReferenceSlot = "npc_interactions"
|
||||
ReferenceMaxBytes = 1048576
|
||||
)
|
||||
|
||||
var referenceSlotDescriptions = shared.ReferenceSlotDescriptions{
|
||||
Glossary: "Optional campaign glossary reference material used only for disambiguation.",
|
||||
Party: "Optional party roster reference material used only for disambiguation.",
|
||||
Players: "Optional player list reference material used only for disambiguation.",
|
||||
Roster: "Deprecated alias for party roster reference material used only for disambiguation.",
|
||||
}
|
||||
|
||||
func referenceSlots() []contracts.ReferenceSlot {
|
||||
slots := shared.ReferenceSlots(referenceSlotDescriptions)
|
||||
slots = append(slots,
|
||||
contracts.ReferenceSlot{
|
||||
Name: NPCRegistryReferenceSlot,
|
||||
Description: "Required normalized NPC registry used only for enemy-subject grounding, never as event evidence.",
|
||||
Required: true,
|
||||
AcceptedMediaTypes: []string{combatturncodec.MediaType},
|
||||
AcceptedArtifactKinds: []contracts.ArtifactKind{dnd.NPCListKind},
|
||||
MaxBytes: ReferenceMaxBytes,
|
||||
},
|
||||
contracts.ReferenceSlot{
|
||||
Name: SceneDescriptionReferenceSlot,
|
||||
Description: "Required scene descriptions used only to determine exact combat eligibility, never as event evidence.",
|
||||
Required: true,
|
||||
AcceptedMediaTypes: []string{combatturncodec.MediaType},
|
||||
AcceptedArtifactKinds: []contracts.ArtifactKind{dnd.SceneDescriptionListKind},
|
||||
MaxBytes: ReferenceMaxBytes,
|
||||
},
|
||||
contracts.ReferenceSlot{
|
||||
Name: CombatTurnReferenceSlot,
|
||||
Description: "Required combat-turn artifact used only as source-free enemy-event grounding.",
|
||||
Required: true,
|
||||
AcceptedMediaTypes: []string{combatturncodec.MediaType},
|
||||
AcceptedArtifactKinds: []contracts.ArtifactKind{dnd.CombatTurnListKind},
|
||||
MaxBytes: ReferenceMaxBytes,
|
||||
},
|
||||
contracts.ReferenceSlot{
|
||||
Name: NPCInteractionReferenceSlot,
|
||||
Description: "Required NPC-interaction artifact used only as source-free enemy-event grounding.",
|
||||
Required: true,
|
||||
AcceptedMediaTypes: []string{interactioncodec.MediaType},
|
||||
AcceptedArtifactKinds: []contracts.ArtifactKind{dnd.NPCInteractionListKind},
|
||||
MaxBytes: ReferenceMaxBytes,
|
||||
},
|
||||
)
|
||||
sort.Slice(slots, func(left, right int) bool { return slots[left].Name < slots[right].Name })
|
||||
return contracts.CloneReferenceSlots(slots)
|
||||
}
|
||||
|
||||
// groundingResolver retains only validated, compact construction-time views.
|
||||
// Per-operation generated references are decoded when supplied and never
|
||||
// become static metadata.
|
||||
type groundingResolver struct {
|
||||
npcs *npcregistry.Resolver
|
||||
scenes *sceneregistry.Resolver
|
||||
|
||||
combatTurns *contracts.LLMInputMaterial
|
||||
npcInteractions *contracts.LLMInputMaterial
|
||||
}
|
||||
|
||||
type grounding struct {
|
||||
npcInput contracts.LLMInputMaterial
|
||||
combatTurnInput contracts.LLMInputMaterial
|
||||
npcInteractionInput contracts.LLMInputMaterial
|
||||
sceneEligibilityView *sceneregistry.Registry
|
||||
}
|
||||
|
||||
func newGroundingResolver(references contracts.ReferenceSet) (*groundingResolver, error) {
|
||||
npcs, err := npcregistry.NewResolver(references)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prepare NPC registry grounding: %w", err)
|
||||
}
|
||||
scenes, err := sceneregistry.NewResolver(references)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prepare scene eligibility: %w", err)
|
||||
}
|
||||
combatTurns, err := prepareCombatTurnInput(references)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
npcInteractions, err := prepareNPCInteractionInput(references)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &groundingResolver{
|
||||
npcs: npcs,
|
||||
scenes: scenes,
|
||||
combatTurns: combatTurns,
|
||||
npcInteractions: npcInteractions,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *groundingResolver) Resolve(references contracts.ReferenceSet) (grounding, error) {
|
||||
if r == nil {
|
||||
return grounding{}, fmt.Errorf("grounding resolver must not be nil")
|
||||
}
|
||||
npcs, err := r.npcs.Resolve(references)
|
||||
if err != nil {
|
||||
return grounding{}, fmt.Errorf("resolve NPC registry grounding: %w", err)
|
||||
}
|
||||
if !npcs.Bound() {
|
||||
return grounding{}, fmt.Errorf("NPC registry reference is required")
|
||||
}
|
||||
scenes, err := r.scenes.Resolve(references)
|
||||
if err != nil {
|
||||
return grounding{}, fmt.Errorf("resolve scene eligibility: %w", err)
|
||||
}
|
||||
if !scenes.Bound() {
|
||||
return grounding{}, fmt.Errorf("scene descriptions reference is required")
|
||||
}
|
||||
combatTurns, err := resolveInput(references, CombatTurnReferenceSlot, r.combatTurns, prepareCombatTurnInput)
|
||||
if err != nil {
|
||||
return grounding{}, err
|
||||
}
|
||||
npcInteractions, err := resolveInput(references, NPCInteractionReferenceSlot, r.npcInteractions, prepareNPCInteractionInput)
|
||||
if err != nil {
|
||||
return grounding{}, err
|
||||
}
|
||||
return grounding{
|
||||
npcInput: npcs.PromptInput(),
|
||||
combatTurnInput: combatTurns,
|
||||
npcInteractionInput: npcInteractions,
|
||||
sceneEligibilityView: scenes,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func resolveInput(references contracts.ReferenceSet, slot string, seeded *contracts.LLMInputMaterial, prepare func(contracts.ReferenceSet) (*contracts.LLMInputMaterial, error)) (contracts.LLMInputMaterial, error) {
|
||||
if _, ok := references.Slots[slot]; ok {
|
||||
value, err := prepare(references)
|
||||
if err != nil {
|
||||
return contracts.LLMInputMaterial{}, err
|
||||
}
|
||||
if value == nil {
|
||||
return contracts.LLMInputMaterial{}, fmt.Errorf("%s reference is required", slot)
|
||||
}
|
||||
return value.Clone(), nil
|
||||
}
|
||||
if seeded == nil {
|
||||
return contracts.LLMInputMaterial{}, fmt.Errorf("%s reference is required", slot)
|
||||
}
|
||||
return seeded.Clone(), nil
|
||||
}
|
||||
|
||||
func (g grounding) PromptInputs() contracts.LLMInputSet {
|
||||
return contracts.LLMInputSet{
|
||||
NPCRegistryReferenceSlot: g.npcInput.Clone(),
|
||||
CombatTurnReferenceSlot: g.combatTurnInput.Clone(),
|
||||
NPCInteractionReferenceSlot: g.npcInteractionInput.Clone(),
|
||||
}
|
||||
}
|
||||
|
||||
func (g grounding) SceneMatch(chunk *source.Chunk) sceneregistry.ChunkMatch {
|
||||
return g.sceneEligibilityView.Match(chunk)
|
||||
}
|
||||
|
||||
func prepareCombatTurnInput(references contracts.ReferenceSet) (*contracts.LLMInputMaterial, error) {
|
||||
item, ok, err := referenceItem(references, CombatTurnReferenceSlot, combatturncodec.MediaType)
|
||||
if err != nil || !ok {
|
||||
return nil, err
|
||||
}
|
||||
value, err := combatturncodec.New().Decode(item.Content)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode combat-turn grounding: invalid approved combat-turn JSON")
|
||||
}
|
||||
content, err := json.Marshal(struct {
|
||||
CombatTurns []combatTurnProjection `json:"combat_turns"`
|
||||
}{CombatTurns: projectCombatTurns(value.CombatTurns)})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode combat-turn grounding: %w", err)
|
||||
}
|
||||
return newPromptInput(CombatTurnReferenceSlot, content), nil
|
||||
}
|
||||
|
||||
func prepareNPCInteractionInput(references contracts.ReferenceSet) (*contracts.LLMInputMaterial, error) {
|
||||
item, ok, err := referenceItem(references, NPCInteractionReferenceSlot, interactioncodec.MediaType)
|
||||
if err != nil || !ok {
|
||||
return nil, err
|
||||
}
|
||||
value, err := interactioncodec.New().Decode(item.Content)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode NPC-interaction grounding: invalid approved NPC-interaction JSON")
|
||||
}
|
||||
content, err := json.Marshal(struct {
|
||||
Interactions []npcInteractionProjection `json:"npc_interactions"`
|
||||
}{Interactions: projectNPCInteractions(value.Interactions)})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode NPC-interaction grounding: %w", err)
|
||||
}
|
||||
return newPromptInput(NPCInteractionReferenceSlot, content), nil
|
||||
}
|
||||
|
||||
func referenceItem(references contracts.ReferenceSet, slotName, expectedMediaType string) (contracts.ReferenceItem, bool, error) {
|
||||
slot, ok := references.Slots[slotName]
|
||||
if !ok {
|
||||
return contracts.ReferenceItem{}, false, nil
|
||||
}
|
||||
if len(slot.Items) != 1 {
|
||||
return contracts.ReferenceItem{}, false, fmt.Errorf("reference slot %q must contain exactly one item", slotName)
|
||||
}
|
||||
item := slot.Items[0]
|
||||
mediaType, _, err := mime.ParseMediaType(item.MediaType)
|
||||
if err != nil {
|
||||
return contracts.ReferenceItem{}, false, fmt.Errorf("reference slot %q item media type is invalid", slotName)
|
||||
}
|
||||
if !strings.EqualFold(mediaType, expectedMediaType) {
|
||||
return contracts.ReferenceItem{}, false, fmt.Errorf("reference slot %q item media type must be %s", slotName, expectedMediaType)
|
||||
}
|
||||
if len(item.Content) > ReferenceMaxBytes {
|
||||
return contracts.ReferenceItem{}, false, fmt.Errorf("reference slot %q item is %d bytes, limit %d", slotName, len(item.Content), ReferenceMaxBytes)
|
||||
}
|
||||
return item, true, nil
|
||||
}
|
||||
|
||||
type combatTurnProjection struct {
|
||||
Actor string `json:"actor"`
|
||||
TurnKind dnd.CombatTurnKind `json:"turn_kind"`
|
||||
}
|
||||
|
||||
func projectCombatTurns(turns []dnd.CombatTurn) []combatTurnProjection {
|
||||
if turns == nil {
|
||||
return nil
|
||||
}
|
||||
projection := make([]combatTurnProjection, len(turns))
|
||||
for index, turn := range turns {
|
||||
projection[index] = combatTurnProjection{Actor: turn.Actor, TurnKind: turn.TurnKind}
|
||||
}
|
||||
return projection
|
||||
}
|
||||
|
||||
type npcInteractionProjection struct {
|
||||
Name string `json:"name"`
|
||||
Kind dnd.NPCInteractionKind `json:"kind"`
|
||||
}
|
||||
|
||||
func projectNPCInteractions(interactions []dnd.NPCInteraction) []npcInteractionProjection {
|
||||
projection := make([]npcInteractionProjection, 0, len(interactions))
|
||||
for _, interaction := range interactions {
|
||||
if interaction.Kind == dnd.NPCInteractionKindCombatOpponent {
|
||||
projection = append(projection, npcInteractionProjection{Name: interaction.Name, Kind: interaction.Kind})
|
||||
}
|
||||
}
|
||||
return projection
|
||||
}
|
||||
|
||||
func newPromptInput(name string, content []byte) *contracts.LLMInputMaterial {
|
||||
sum := sha256.Sum256(content)
|
||||
material := contracts.NewLLMInputMaterial(name, combatturncodec.MediaType, content, "sha256:"+hex.EncodeToString(sum[:]), "")
|
||||
return &material
|
||||
}
|
||||
258
internal/modules/dnd/extract/enemyevents/grounding_test.go
Normal file
258
internal/modules/dnd/extract/enemyevents/grounding_test.go
Normal file
@@ -0,0 +1,258 @@
|
||||
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"
|
||||
interactioncodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/npcinteractions"
|
||||
npccodec "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/codec/npcs"
|
||||
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.NPCListKind},
|
||||
{SceneDescriptionReferenceSlot, dnd.SceneDescriptionListKind},
|
||||
{CombatTurnReferenceSlot, dnd.CombatTurnListKind},
|
||||
{NPCInteractionReferenceSlot, dnd.NPCInteractionListKind},
|
||||
} {
|
||||
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"}]}`,
|
||||
NPCInteractionReferenceSlot: `{"npc_interactions":[{"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)
|
||||
}
|
||||
if resolved.SceneMatch(combatChunk()) != (sceneregistry.ChunkMatch{State: sceneregistry.MatchExact, Kind: dnd.SceneKindNarrative}) {
|
||||
t.Fatalf("generated scene match = %#v", resolved.SceneMatch(combatChunk()))
|
||||
}
|
||||
if resolved.SceneMatch(&source.Chunk{ID: "other", Ref: combatChunk().Ref}).State != sceneregistry.MatchMissing {
|
||||
t.Fatal("missing scene was not reported")
|
||||
}
|
||||
mismatched := combatChunk()
|
||||
mismatched.Ref.EndUnitID++
|
||||
if resolved.SceneMatch(mismatched).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)
|
||||
}
|
||||
if static.SceneMatch(combatChunk()) != (sceneregistry.ChunkMatch{State: sceneregistry.MatchExact, Kind: dnd.SceneKindCombat}) {
|
||||
t.Fatalf("static scene match = %#v", static.SceneMatch(combatChunk()))
|
||||
}
|
||||
}
|
||||
|
||||
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},
|
||||
{NPCInteractionReferenceSlot, NPCInteractionReferenceSlot},
|
||||
} {
|
||||
t.Run("missing "+test.slot, func(t *testing.T) {
|
||||
resolver, err := newGroundingResolver(withoutSlot(valid, test.slot))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := resolver.Resolve(contracts.ReferenceSet{}); err == nil || !strings.Contains(err.Error(), test.want) {
|
||||
t.Fatalf("Resolve() 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, NPCInteractionReferenceSlot, 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)})},
|
||||
}
|
||||
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 groundingReferences(t *testing.T, enemy string, sceneKind dnd.SceneKind) contracts.ReferenceSet {
|
||||
t.Helper()
|
||||
npcContent, err := npccodec.New().Encode(dnd.NPCList{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)
|
||||
}
|
||||
interactionContent, err := interactioncodec.New().Encode(dnd.NPCInteractionList{Interactions: []dnd.NPCInteraction{
|
||||
{Name: enemy, Kind: dnd.NPCInteractionKindCombatOpponent, SourceRefs: []source.SourceRef{{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 1}}},
|
||||
{Name: "Aria", Kind: dnd.NPCInteractionKindCombatAlly, 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)}},
|
||||
NPCInteractionReferenceSlot: {Items: []contracts.ReferenceItem{newReferenceItem(NPCInteractionReferenceSlot, interactionContent)}},
|
||||
}}
|
||||
}
|
||||
|
||||
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[:])
|
||||
}
|
||||
Reference in New Issue
Block a user