Files

230 lines
10 KiB
Go

package enemyevents
import (
"context"
"encoding/json"
"errors"
"reflect"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
)
func TestExtractMapsEnemyEventsInSourceOrder(t *testing.T) {
client := &fakeEnemyEventsLLMClient{response: extractionResponse{Events: []enemyEventResponse{
{Name: "Ashfang", Kind: "killed", SourceRefs: []enemySourceRefResponse{{StartUnitID: 4, EndUnitID: 4}}},
{Name: "Ashfang", Kind: "engaged", SourceRefs: []enemySourceRefResponse{{StartUnitID: 1, EndUnitID: 1}, {StartUnitID: 1, EndUnitID: 1}}},
{Name: "Orcs", Kind: "fled", SourceRefs: []enemySourceRefResponse{{StartUnitID: 3, EndUnitID: 3}}},
{Name: "One orc", Kind: "captured", SourceRefs: []enemySourceRefResponse{{StartUnitID: 2, EndUnitID: 2}}},
{Name: "Ashfang", Kind: "incapacitated", SourceRefs: []enemySourceRefResponse{{StartUnitID: 2, EndUnitID: 2}}},
}}}
result, err := newEnemyExtractor(t, client).Extract(context.Background(), enemyExtractionRequest(t))
if err != nil {
t.Fatal(err)
}
if got := []dnd.EnemyEventKind{result.Value.Events[0].Kind, result.Value.Events[1].Kind, result.Value.Events[2].Kind, result.Value.Events[3].Kind, result.Value.Events[4].Kind}; !reflect.DeepEqual(got, []dnd.EnemyEventKind{
dnd.EnemyEventKindEngaged,
dnd.EnemyEventKindCaptured,
dnd.EnemyEventKindIncapacitated,
dnd.EnemyEventKindFled,
dnd.EnemyEventKindKilled,
}) {
t.Fatalf("event order = %#v", got)
}
if refs := result.Value.Events[0].SourceRefs; !reflect.DeepEqual(refs, []source.SourceRef{{SourceID: "combat-session", StartUnitID: 1, EndUnitID: 1}}) {
t.Fatalf("canonical source refs = %#v", refs)
}
if len(client.requests) != 1 {
t.Fatalf("LLM calls = %d, want 1", len(client.requests))
}
request := client.requests[0]
if request.StageName != Key || request.PromptID != PromptID || request.PromptVersion != SchemaVersion || request.ProfileID != "enemy-profile" || request.SessionID != "session-123" {
t.Fatalf("LLM request identity = %#v", request)
}
if got := string(request.Inputs[CombatTurnReferenceSlot].Content); !strings.Contains(got, `"actor":"Ashfang"`) || strings.Contains(got, "source_ref") {
t.Fatalf("combat grounding = %s", got)
}
if got := string(request.Inputs[NPCOccurrenceReferenceSlot].Content); !strings.Contains(got, `"kind":"combat_opponent"`) || strings.Contains(got, "Aria") {
t.Fatalf("occurrence grounding = %s", got)
}
}
func TestExtractPreservesSemanticCandidatesAndResponseOwnership(t *testing.T) {
client := &fakeEnemyEventsLLMClient{response: extractionResponse{Events: []enemyEventResponse{{
Name: " ", Kind: "unsupported", SourceRefs: []enemySourceRefResponse{{StartUnitID: 99, EndUnitID: 0}},
}}}}
result, err := newEnemyExtractor(t, client).Extract(context.Background(), enemyExtractionRequest(t))
if err != nil {
t.Fatal(err)
}
event := result.Value.Events[0]
if event.Name != " " || event.Kind != "unsupported" || event.SourceRefs[0] != (source.SourceRef{SourceID: "combat-session", StartUnitID: 99}) {
t.Fatalf("semantic candidate = %#v", event)
}
result.Value.Events[0].SourceRefs[0].StartUnitID = 7
if client.response.Events[0].SourceRefs[0].StartUnitID != 99 {
t.Fatal("mapped event aliases model-owned source references")
}
}
func TestExtractSkipsModelForIneligibleScenes(t *testing.T) {
for _, test := range []struct {
name string
kind dnd.SceneKind
change func(*source.Chunk)
wantWarning bool
}{
{name: "non-combat", kind: dnd.SceneKindNarrative},
{name: "missing", kind: dnd.SceneKindCombat, change: func(chunk *source.Chunk) { chunk.ID = "other" }, wantWarning: true},
{name: "mismatched", kind: dnd.SceneKindCombat, change: func(chunk *source.Chunk) { chunk.Ref.EndUnitID++ }, wantWarning: true},
} {
t.Run(test.name, func(t *testing.T) {
client := &fakeEnemyEventsLLMClient{}
request := enemyExtractionRequest(t)
request.References = groundingReferences(t, "Ashfang", test.kind)
if test.change != nil {
test.change(request.Chunk)
}
result, err := newEnemyExtractor(t, client).Extract(context.Background(), request)
if err != nil {
t.Fatal(err)
}
if len(client.requests) != 0 || result.Value.Events == nil || len(result.Value.Events) != 0 {
t.Fatalf("result = %#v, calls = %d", result, len(client.requests))
}
if test.wantWarning {
if len(result.Warnings) != 1 || result.Warnings[0].ReasonCode != "scene_classification_unavailable" || result.Warnings[0].Scope != SceneDescriptionReferenceSlot {
t.Fatalf("warnings = %#v", result.Warnings)
}
} else if len(result.Warnings) != 0 {
t.Fatalf("warnings = %#v", result.Warnings)
}
})
}
}
func TestExtractReportsRequiredGroundingAndProviderFailures(t *testing.T) {
client := &fakeEnemyEventsLLMClient{}
request := enemyExtractionRequest(t)
request.References = withoutSlot(request.References, NPCRegistryReferenceSlot)
if _, err := (&Extractor{llm: client, grounding: mustGroundingResolver(t, request.References)}).Extract(context.Background(), request); err == nil || !strings.Contains(err.Error(), "NPC registry") {
t.Fatalf("Extract() error = %v, want required grounding context", err)
}
provider := &fakeEnemyEventsLLMClient{err: errors.New("provider unavailable")}
if _, err := newEnemyExtractor(t, provider).Extract(context.Background(), enemyExtractionRequest(t)); err == nil || !strings.Contains(err.Error(), "complete structured output") || !strings.Contains(err.Error(), "provider unavailable") {
t.Fatalf("Extract() error = %v, want provider context", err)
}
}
func TestConstructorSpecOptionsAndSafeMetadata(t *testing.T) {
if _, err := New(nil, Options{}); err == nil || !strings.Contains(err.Error(), "LLM client") {
t.Fatalf("New(nil) error = %v", err)
}
if _, err := New(&fakeEnemyEventsLLMClient{}, Options{}, contracts.ReferenceSet{}, contracts.ReferenceSet{}); err == nil || !strings.Contains(err.Error(), "at most one") {
t.Fatalf("New() error = %v", err)
}
if _, err := DecodeOptions(map[string]any{"unexpected": true}); err == nil || !strings.Contains(err.Error(), "unknown option") {
t.Fatalf("DecodeOptions() error = %v", err)
}
first := ModuleSpec()
first.Requires[0] = "changed"
first.ReferenceSlots[0].AcceptedMediaTypes[0] = "changed"
second := ModuleSpec()
if second.Requires[0] != "chunks" || second.ArtifactKind != dnd.EnemyEventListKind || second.ExecutionClass != contracts.ExecutionClassLLMBacked || second.ReferenceSlots[0].AcceptedMediaTypes[0] == "changed" {
t.Fatalf("ModuleSpec() reused mutable state: %#v", second)
}
registry := pipeline.NewExtractorRegistry()
if err := Register(registry); err != nil {
t.Fatal(err)
}
if spec, ok := registry.Spec(Key); !ok || spec.Key != Key || spec.ArtifactKind != dnd.EnemyEventListKind {
t.Fatalf("registered spec = %#v, %t", spec, ok)
}
extractor := newEnemyExtractor(t, &fakeEnemyEventsLLMClient{}, groundingReferences(t, "Ashfang", dnd.SceneKindCombat))
metadata := extractor.ManifestMetadata()
encoded, err := json.Marshal(metadata)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(encoded), "Ashfang") || metadata["mapping_policy"] != mappingPolicy || metadata["scene_gate_policy"] != sceneGatePolicy {
t.Fatalf("unsafe or incomplete metadata = %s", encoded)
}
fingerprints := extractor.CheckpointFingerprints()
if len(fingerprints) != 4 || fingerprints[0].Value != metadata["prompt_sha256"] || fingerprints[1].Value != metadata["response_schema_sha256"] || fingerprints[2].Value != mappingPolicy || fingerprints[3].Value != sceneGatePolicy {
t.Fatalf("fingerprints = %#v", fingerprints)
}
}
func newEnemyExtractor(t *testing.T, client contracts.StructuredLLMClient, references ...contracts.ReferenceSet) *Extractor {
t.Helper()
if len(references) == 0 {
references = []contracts.ReferenceSet{groundingReferences(t, "Ashfang", dnd.SceneKindCombat)}
}
extractor, err := New(client, Options{}, references...)
if err != nil {
t.Fatal(err)
}
return extractor
}
func mustGroundingResolver(t *testing.T, references contracts.ReferenceSet) *groundingResolver {
t.Helper()
resolver, err := newGroundingResolver(references)
if err != nil {
t.Fatal(err)
}
return resolver
}
func enemyExtractionRequest(t *testing.T) contracts.TypedExtractionRequest {
t.Helper()
chunk := combatChunk()
chunk.Units = []source.SourceUnit{{ID: 1}, {ID: 2}, {ID: 3}, {ID: 4}}
chunk.MediaType = "application/json"
chunk.Content = []byte(`{"id":"combat-scene","units":[1,2,3,4]}`)
return contracts.TypedExtractionRequest{
Source: &source.SourceDocument{ID: "combat-session", Units: append([]source.SourceUnit(nil), chunk.Units...)},
Chunk: chunk,
SourceInput: contracts.NewLLMInputMaterial("source", "application/json", chunk.Content, digest(chunk.Content), "file:///combat-session.json"),
References: groundingReferences(t, "Ashfang", dnd.SceneKindCombat),
LLMProfile: "enemy-profile",
SessionID: "session-123",
}
}
type fakeEnemyEventsLLMClient struct {
response extractionResponse
err error
requests []contracts.StructuredCompletionRequest
}
func (client *fakeEnemyEventsLLMClient) CompleteStructured(_ context.Context, request contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
client.requests = append(client.requests, cloneEnemyRequest(request))
if client.err != nil {
return contracts.StructuredCompletionResponse{}, client.err
}
target, ok := out.(*extractionResponse)
if !ok {
return contracts.StructuredCompletionResponse{}, errors.New("unexpected output target")
}
*target = client.response
content, err := json.Marshal(client.response)
if err != nil {
return contracts.StructuredCompletionResponse{}, err
}
return contracts.StructuredCompletionResponse{Content: content}, nil
}
func cloneEnemyRequest(request contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest {
request.Inputs = request.Inputs.Clone()
return request
}