package spells import ( "context" "encoding/json" "errors" "strings" "testing" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" "gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared" ) func TestExtractReturnsCanonicalOutputFromStructuredResponse(t *testing.T) { client := &fakeSpellsLLMClient{ response: extractionResponse{ SpellCasts: []spellCastResponse{ { Caster: " Aria ", Spell: " Cure Wounds ", Effect: " Heals an injured ally. ", NarrativeDescription: " Aria restores the fighter after the fight. ", SourceRefs: responseSourceRefsInt("transcript", 1, 2), }, }, }, content: []byte(`{"spell_casts":[{"caster":" Aria ","spell":" Cure Wounds ","effect":" Heals an injured ally. ","narrative_description":" Aria restores the fighter after the fight. ","source_refs":[{"source_id":"session-alpha","start_unit_id":1,"end_unit_id":2}]}],"raw_marker":true}`), } extractReq := extractionRequestWithClient(client) result, err := New().Extract(context.Background(), extractReq) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } if len(client.requests) != 1 { t.Fatalf("LLM calls = %d, want 1", len(client.requests)) } llmReq := client.requests[0] if llmReq.StageName != Key { t.Fatalf("StageName = %q, want %q", llmReq.StageName, Key) } if llmReq.PromptID != PromptID || llmReq.PromptVersion != SchemaVersion { t.Fatalf("prompt = %q/%q, want %q/%q", llmReq.PromptID, llmReq.PromptVersion, PromptID, SchemaVersion) } if llmReq.SessionID != "session-123" || llmReq.ProfileID != "profile-spells" { t.Fatalf("session/profile = %q/%q, want session-123/profile-spells", llmReq.SessionID, llmReq.ProfileID) } transcript := llmReq.Inputs["transcript"] if transcript.Name != "transcript" || transcript.MediaType != "application/json" || transcript.Digest != "sha256:chunk" || transcript.OriginURI != "file:///session-alpha.json" { t.Fatalf("transcript metadata = %#v", transcript) } if got := string(transcript.Content); got != string(extractReq.Chunk.Content) { t.Fatalf("transcript content = %q, want chunk content %q", got, extractReq.Chunk.Content) } if result.Output.Payload.MediaType != "application/json" { t.Fatalf("MediaType = %q, want application/json", result.Output.Payload.MediaType) } if result.Output.Schema.ID != ResponseSchemaID || result.Output.Schema.Name != ResponseSchemaName || result.Output.Schema.Version != SchemaVersion { t.Fatalf("schema = %#v, want response schema provenance", result.Output.Schema) } if !json.Valid(result.Output.Schema.JSONSchema) { t.Fatalf("schema JSON is invalid or missing: %s", result.Output.Schema.JSONSchema) } if strings.Contains(string(result.Output.Payload.Content), "raw_marker") { t.Fatalf("content = %q, want canonical payload without raw completion marker", result.Output.Payload.Content) } var payload extractionResponse if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } if len(payload.SpellCasts) != 1 || payload.SpellCasts[0].Spell != " Cure Wounds " { t.Fatalf("payload = %#v, want structured response fields", payload) } if got := payload.SpellCasts[0].SourceRefs[0].SourceID; got != "session-alpha" { t.Fatalf("source_id = %q, want canonical source document ID", got) } } func TestExtractorManifestMetadataIncludesPromptAndSchemaProvenance(t *testing.T) { metadata := New().ManifestMetadata() tests := map[string]string{ "prompt_id": PromptID, "prompt_version": SchemaVersion, "response_schema_key": string(ResponseSchemaKey), "response_schema_id": ResponseSchemaID, "response_schema_name": ResponseSchemaName, "response_schema_version": SchemaVersion, } for key, want := range tests { if metadata[key] != want { t.Fatalf("metadata[%q] = %#v, want %q", key, metadata[key], want) } } for _, key := range []string{"prompt_sha256", "response_schema_sha256"} { value, ok := metadata[key].(string) if !ok || !strings.HasPrefix(value, "sha256:") { t.Fatalf("metadata[%q] = %#v, want sha256 value", key, metadata[key]) } } } func TestExtractPassesReferencesAsPromptInputs(t *testing.T) { client := &fakeSpellsLLMClient{response: extractionResponse{SpellCasts: []spellCastResponse{}}} req := extractionRequestWithClient(client) req.References = contracts.ReferenceSet{ Slots: map[string]contracts.ResolvedReferenceSlot{ "players": { Slot: contracts.ReferenceSlot{Name: "players"}, Items: []contracts.ReferenceItem{ {SlotName: "players", Content: []byte("Alice: Aria Brightmantle")}, }, }, "party": { Slot: contracts.ReferenceSlot{Name: "party"}, Items: []contracts.ReferenceItem{ {SlotName: "party", Content: []byte("Aria Brightmantle: party cleric")}, }, }, "glossary": { Slot: contracts.ReferenceSlot{Name: "glossary"}, Items: []contracts.ReferenceItem{ {SlotName: "glossary", Content: []byte("Brightmantle: local temple name")}, }, }, }, } if _, err := New().Extract(context.Background(), req); err != nil { t.Fatalf("Extract() error = %v, want nil", err) } if len(client.requests) != 1 { t.Fatalf("LLM calls = %d, want 1", len(client.requests)) } request := client.requests[0] if request.PromptID != PromptID || request.PromptVersion != SchemaVersion { t.Fatalf("prompt = %q/%q, want %q/%q", request.PromptID, request.PromptVersion, PromptID, SchemaVersion) } if got := string(request.Inputs["players"].Content); got != "Alice: Aria Brightmantle" { t.Fatalf("players input = %q, want reference content", got) } if got := string(request.Inputs["party"].Content); got != "Aria Brightmantle: party cleric" { t.Fatalf("party input = %q, want reference content", got) } if got := string(request.Inputs["glossary"].Content); got != "Brightmantle: local temple name" { t.Fatalf("glossary input = %q, want reference content", got) } if strings.Contains(string(request.Inputs["transcript"].Content), "Aria Brightmantle: party cleric") { t.Fatalf("transcript input contains reference content") } } func TestPromptInputsMapLegacyRosterReferenceToParty(t *testing.T) { inputs := shared.PromptInputs(spellSourceInput(), contracts.ReferenceSet{ Slots: map[string]contracts.ResolvedReferenceSlot{ "roster": { Slot: contracts.ReferenceSlot{Name: "roster"}, Items: []contracts.ReferenceItem{ {SlotName: "roster", Content: []byte("Legacy roster text")}, }, }, }, }) if got := string(inputs["party"].Content); got != "Legacy roster text" { t.Fatalf("party input = %q, want legacy roster content", got) } if _, ok := inputs["roster"]; ok { t.Fatalf("roster prompt input was present; want only party input") } } func TestExtractReturnsRawOutputForEmptyResponse(t *testing.T) { client := &fakeSpellsLLMClient{response: extractionResponse{SpellCasts: []spellCastResponse{}}} result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } var payload extractionResponse if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } if len(payload.SpellCasts) != 0 { t.Fatalf("SpellCasts = %#v, want none", payload.SpellCasts) } } func TestExtractReturnsCanonicalOutputForMalformedStructuredResponse(t *testing.T) { client := &fakeSpellsLLMClient{response: extractionResponse{}} result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } if string(result.Output.Payload.Content) != `{"spell_casts":null}` { t.Fatalf("content = %s, want canonical structured output", result.Output.Payload.Content) } } func TestExtractWrapsLLMClientError(t *testing.T) { client := &fakeSpellsLLMClient{err: errors.New("provider unavailable")} _, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err == nil { t.Fatal("Extract() error = nil, want LLM error") } if !strings.Contains(err.Error(), "dnd spells") || !strings.Contains(err.Error(), "provider unavailable") { t.Fatalf("Extract() error = %q, want wrapped LLM context", err.Error()) } } func TestExtractRejectsInvalidRequests(t *testing.T) { validClient := &fakeSpellsLLMClient{response: extractionResponse{SpellCasts: []spellCastResponse{}}} validReq := extractionRequestWithClient(validClient) canceledCtx, cancel := context.WithCancel(context.Background()) cancel() tests := []struct { name string extractor *Extractor ctx context.Context req contracts.ExtractionRequest want string }{ {name: "nil extractor", extractor: nil, ctx: context.Background(), req: validReq, want: "extractor"}, {name: "nil context", extractor: New(), ctx: nil, req: validReq, want: "context"}, {name: "canceled context", extractor: New(), ctx: canceledCtx, req: validReq, want: "context"}, {name: "nil source", extractor: New(), ctx: context.Background(), req: contracts.ExtractionRequest{Chunk: validReq.Chunk, LLMClient: validReq.LLMClient}, want: "source"}, {name: "nil chunk", extractor: New(), ctx: context.Background(), req: contracts.ExtractionRequest{Source: validReq.Source, LLMClient: validReq.LLMClient}, want: "chunk"}, {name: "empty chunk units", extractor: New(), ctx: context.Background(), req: emptyChunkRequest(validReq), want: "units"}, {name: "nil LLM client", extractor: New(), ctx: context.Background(), req: contracts.ExtractionRequest{Source: validReq.Source, Chunk: validReq.Chunk}, want: "LLM client"}, {name: "source input mismatches chunk", extractor: New(), ctx: context.Background(), req: mismatchedSourceInputRequest(validReq), want: "must match chunk"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { _, err := tt.extractor.Extract(tt.ctx, tt.req) if err == nil { t.Fatal("Extract() error = nil, want error") } if !strings.Contains(err.Error(), "dnd spells") || !strings.Contains(err.Error(), tt.want) { t.Fatalf("Extract() error = %q, want %q context", err.Error(), tt.want) } }) } } func TestExtractOrdersSpellCastsByEarliestSourceUnit(t *testing.T) { client := &fakeSpellsLLMClient{ response: extractionResponse{ SpellCasts: []spellCastResponse{ { Caster: "Bandit Shaman", Spell: "Fire Bolt", Effect: "Burns.", NarrativeDescription: "Second spell.", SourceRefs: responseSourceRefs("session-alpha", 2, 2), }, { Caster: "Aria", Spell: "Cure Wounds", Effect: "Heals.", NarrativeDescription: "First spell.", SourceRefs: responseSourceRefs("session-alpha", 1, 1), }, { Caster: "Narrator", Spell: "Unknown Spell", Effect: "No cited range.", NarrativeDescription: "This should sort after cited spell casts.", }, }, }, } result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } var payload extractionResponse if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } if len(payload.SpellCasts) != 3 || payload.SpellCasts[0].Spell != "Cure Wounds" || payload.SpellCasts[1].Spell != "Fire Bolt" || payload.SpellCasts[2].Spell != "Unknown Spell" { t.Fatalf("spell order = %#v, want earliest source-unit order with uncited spell last", payload.SpellCasts) } } func TestExtractCanonicalizesSourceRefs(t *testing.T) { client := &fakeSpellsLLMClient{ response: extractionResponse{ SpellCasts: []spellCastResponse{ { Caster: "Aria", Spell: "Cure Wounds", Effect: "Heals.", NarrativeDescription: "Aria heals.", SourceRefs: []shared.SourceRefResponse{ {SourceID: "gameplay_transcript", StartUnitID: shared.UnitRefFromInt(2), EndUnitID: shared.UnitRefFromInt(2)}, {SourceID: "", StartUnitID: shared.UnitRefFromInt(1), EndUnitID: shared.UnitRefFromInt(2)}, {SourceID: "transcript", StartUnitID: shared.UnitRefFromInt(1), EndUnitID: shared.UnitRefFromInt(2)}, }, }, }, }, } result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } var payload extractionResponse if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } refs := payload.SpellCasts[0].SourceRefs if len(refs) != 2 { t.Fatalf("source refs = %#v, want duplicate collapsed", refs) } for _, ref := range refs { if ref.SourceID != "session-alpha" { t.Fatalf("source ref = %#v, want canonical source_id", ref) } } if refs[0].StartUnitID.Int() != 1 || refs[0].EndUnitID.Int() != 2 || refs[1].StartUnitID.Int() != 2 || refs[1].EndUnitID.Int() != 2 { t.Fatalf("source refs = %#v, want sorted unit ranges", refs) } } func TestExtractPreservesInvalidSourceRefsForValidators(t *testing.T) { client := &fakeSpellsLLMClient{ response: extractionResponse{ SpellCasts: []spellCastResponse{ { Caster: "Aria", Spell: "Cure Wounds", Effect: "Heals.", NarrativeDescription: "Aria heals.", SourceRefs: []shared.SourceRefResponse{ {SourceID: "transcript", StartUnitID: shared.UnitRefFromInt(99), EndUnitID: shared.UnitRefFromString("missing")}, }, }, }, }, } result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } var payload map[string][]map[string]any if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } ref := payload["spell_casts"][0]["source_refs"].([]any)[0].(map[string]any) if ref["source_id"] != "session-alpha" || ref["start_unit_id"] != float64(99) || ref["end_unit_id"] != "" { t.Fatalf("source ref = %#v, want source_id canonicalized without unit repair", ref) } } func TestExtractDefensivelyCopiesRawContent(t *testing.T) { client := &fakeSpellsLLMClient{ response: extractionResponse{ SpellCasts: []spellCastResponse{ { Caster: "Aria", Spell: "Cure Wounds", Effect: "Heals.", NarrativeDescription: "Aria heals.", SourceRefs: responseSourceRefs("session-alpha", 1, 2), }, }, }, } result, err := New().Extract(context.Background(), extractionRequestWithClient(client)) if err != nil { t.Fatalf("Extract() error = %v, want nil", err) } client.response.SpellCasts[0].SourceRefs[0].StartUnitID = shared.UnitRefFromInt(99) var payload extractionResponse if err := json.Unmarshal(result.Output.Payload.Content, &payload); err != nil { t.Fatalf("Unmarshal(Content) error = %v, want nil", err) } if got := payload.SpellCasts[0].SourceRefs[0].StartUnitID.String(); got != "1" { t.Fatalf("source ref start = %q, want copied 1", got) } } func extractionRequestWithClient(client contracts.StructuredLLMClient) contracts.ExtractionRequest { req := promptExtractionRequest() req.LLMClient = client req.SourceInput = spellChunkInput(req.Chunk) req.SessionID = "session-123" req.LLMProfile = "profile-spells" return req } const spellTranscriptJSON = `{"id":"session-alpha","segments":[{"id":1,"text":"Aria raises her hand and casts Cure Wounds."}]}` func spellSourceInput() contracts.LLMInputMaterial { return contracts.NewLLMInputMaterial("source", "application/json", []byte(spellTranscriptJSON), "sha256:transcript", "file:///session-alpha.json") } func spellChunkInput(chunk *contracts.SourceChunk) contracts.LLMInputMaterial { return contracts.NewLLMInputMaterial("source", chunk.MediaType, chunk.Content, "sha256:chunk", "file:///session-alpha.json") } func emptyChunkRequest(req contracts.ExtractionRequest) contracts.ExtractionRequest { req.Chunk = &contracts.SourceChunk{ ID: req.Chunk.ID, SourceID: req.Chunk.SourceID, Index: req.Chunk.Index, } return req } func mismatchedSourceInputRequest(req contracts.ExtractionRequest) contracts.ExtractionRequest { req.SourceInput = spellSourceInput() return req } type fakeSpellsLLMClient struct { response extractionResponse content []byte err error requests []contracts.StructuredCompletionRequest } func (client *fakeSpellsLLMClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) { client.requests = append(client.requests, cloneStructuredCompletionRequest(req)) 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 := append([]byte(nil), client.content...) if len(content) == 0 { var err error content, err = json.Marshal(client.response) if err != nil { return contracts.StructuredCompletionResponse{}, err } } return contracts.StructuredCompletionResponse{Content: content}, nil } func cloneStructuredCompletionRequest(req contracts.StructuredCompletionRequest) contracts.StructuredCompletionRequest { req.Inputs = req.Inputs.Clone() req.Vars = cloneVars(req.Vars) return req } func cloneVars(in map[string]any) map[string]any { if len(in) == 0 { return nil } out := make(map[string]any, len(in)) for key, value := range in { out[key] = value } return out }