109 lines
5.1 KiB
Go
109 lines
5.1 KiB
Go
package itemregistry
|
|
|
|
import (
|
|
"context"
|
|
"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/modules/dnd"
|
|
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/items/identity"
|
|
)
|
|
|
|
func TestExtractMapsItemsWithOwnedEvidenceAndDeterministicOrder(t *testing.T) {
|
|
client := &fakeItemsLLMClient{response: extractionResponse{Items: []itemResponse{
|
|
{Name: " Gold Pieces ", SourceRefs: responseSourceRefs(3, 3)},
|
|
{Name: "Star Compass", SourceRefs: []itemSourceRefResponse{{StartUnitID: 2, EndUnitID: 2}, {StartUnitID: 1, EndUnitID: 1}, {StartUnitID: 1, EndUnitID: 1}}},
|
|
}}}
|
|
result, err := newExtractor(t, client).Extract(context.Background(), extractionRequest())
|
|
if err != nil {
|
|
t.Fatalf("Extract() error = %v, want nil", err)
|
|
}
|
|
refs := []source.SourceRef{{SourceID: "session-items", StartUnitID: 1, EndUnitID: 1}, {SourceID: "session-items", StartUnitID: 2, EndUnitID: 2}}
|
|
want := dnd.ItemRegistry{Items: []dnd.Item{
|
|
{ID: identity.DeriveID("Star Compass"), Name: "Star Compass", SourceRefs: refs},
|
|
{ID: identity.DeriveID("Gold Pieces"), Name: "Gold Pieces", SourceRefs: []source.SourceRef{{SourceID: "session-items", StartUnitID: 3, EndUnitID: 3}}},
|
|
}}
|
|
if !reflect.DeepEqual(result.Value, want) {
|
|
t.Fatalf("Value = %#v, want %#v", result.Value, want)
|
|
}
|
|
result.Value.Items[0].SourceRefs[0].StartUnitID = 99
|
|
for _, item := range client.response.Items {
|
|
for _, ref := range item.SourceRefs {
|
|
if ref.StartUnitID == 99 {
|
|
t.Fatal("result source references alias the model response")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestExtractKeepsAliasCandidatesSeparate(t *testing.T) {
|
|
client := &fakeItemsLLMClient{response: extractionResponse{Items: []itemResponse{
|
|
{Name: "Star Compass", SourceRefs: responseSourceRefs(1, 1)},
|
|
{Name: "Compass of the Stars", SourceRefs: responseSourceRefs(2, 2)},
|
|
}}}
|
|
result, err := newExtractor(t, client).Extract(context.Background(), extractionRequest())
|
|
if err != nil || len(result.Value.Items) != 2 || result.Value.Items[0].ID == result.Value.Items[1].ID {
|
|
t.Fatalf("Extract() = %#v, %v; want separate alias candidates", result, err)
|
|
}
|
|
}
|
|
|
|
func TestExtractPreservesInvalidCandidatesForValidators(t *testing.T) {
|
|
client := &fakeItemsLLMClient{content: []byte(`{"items":[{"name":"","source_refs":[{"start_unit_id":0,"end_unit_id":-1}]}]}`)}
|
|
result, err := newExtractor(t, client).Extract(context.Background(), extractionRequest())
|
|
if err != nil {
|
|
t.Fatalf("Extract() error = %v, want nil", err)
|
|
}
|
|
item := result.Value.Items[0]
|
|
if item.ID != "" || item.Name != "" || !reflect.DeepEqual(item.SourceRefs, []source.SourceRef{{SourceID: "session-items", StartUnitID: 0, EndUnitID: -1}}) {
|
|
t.Fatalf("item = %#v, want invalid candidate preserved", item)
|
|
}
|
|
}
|
|
|
|
func TestExtractPassesReferencesWithoutTreatingThemAsEvidence(t *testing.T) {
|
|
client := &fakeItemsLLMClient{response: extractionResponse{Items: []itemResponse{}}}
|
|
req := extractionRequest()
|
|
req.References = contracts.ReferenceSet{Slots: map[string]contracts.ResolvedReferenceSlot{
|
|
"glossary": {Slot: contracts.ReferenceSlot{Name: "glossary"}, Items: []contracts.ReferenceItem{{SlotName: "glossary", Content: []byte("Star Compass: heirloom")}}},
|
|
}}
|
|
if _, err := newExtractor(t, client).Extract(context.Background(), req); err != nil {
|
|
t.Fatalf("Extract() error = %v", err)
|
|
}
|
|
inputs := client.requests[0].Inputs
|
|
if string(inputs["glossary"].Content) != "Star Compass: heirloom" || strings.Contains(string(inputs["transcript"].Content), "heirloom") {
|
|
t.Fatalf("prompt inputs = %#v, want separated reference material", inputs)
|
|
}
|
|
}
|
|
|
|
func TestExtractHandlesCancellationAndFailures(t *testing.T) {
|
|
request := extractionRequest()
|
|
canceled, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
if _, err := newExtractor(t, &fakeItemsLLMClient{}).Extract(canceled, request); err == nil || !errors.Is(err, context.Canceled) || !strings.Contains(err.Error(), "dnd item registry") {
|
|
t.Fatalf("canceled Extract() error = %v, want contextual cancellation", err)
|
|
}
|
|
for _, test := range []struct {
|
|
name string
|
|
extractor *Extractor
|
|
req contracts.TypedExtractionRequest
|
|
want string
|
|
}{
|
|
{name: "nil extractor", extractor: nil, req: request, want: "extractor"},
|
|
{name: "nil client", extractor: &Extractor{}, req: request, want: "LLM client"},
|
|
{name: "preflight", extractor: newExtractor(t, &fakeItemsLLMClient{}), req: mismatchedSourceInputRequest(request), want: "must match chunk"},
|
|
} {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
if _, err := test.extractor.Extract(context.Background(), test.req); err == nil || !strings.Contains(err.Error(), "dnd item registry") || !strings.Contains(err.Error(), test.want) {
|
|
t.Fatalf("Extract() error = %v, want local context", err)
|
|
}
|
|
})
|
|
}
|
|
_, err := newExtractor(t, &fakeItemsLLMClient{err: errors.New("provider unavailable")}).Extract(context.Background(), request)
|
|
if err == nil || !strings.Contains(err.Error(), "dnd item registry") || !strings.Contains(err.Error(), "provider unavailable") {
|
|
t.Fatalf("provider error = %v, want contextual provider error", err)
|
|
}
|
|
}
|