package locationregistry import ( "context" "encoding/json" "errors" "testing" "gitea.maximumdirect.net/eric/notarius/internal/core/source" "gitea.maximumdirect.net/eric/notarius/internal/framework/contracts" ) func extractionRequest() contracts.TypedExtractionRequest { doc := sourceDocument() chunk := &source.Chunk{ ID: "session-locations:chunk:0", SourceID: doc.ID, Index: 0, Ref: source.SourceRef{SourceID: doc.ID, StartUnitID: 1, EndUnitID: 3}, Content: []byte(`{"units":[1,2,3]}`), MediaType: "application/json", Units: append([]source.SourceUnit(nil), doc.Units...), } return contracts.TypedExtractionRequest{ Source: doc, Chunk: chunk, SourceInput: contracts.NewLLMInputMaterial("source", chunk.MediaType, chunk.Content, "sha256:chunk", "file:///session-locations.json"), SessionID: "location-session", LLMProfile: "location-profile", } } func sourceDocument() *source.SourceDocument { return &source.SourceDocument{ID: "session-locations", Kind: "transcript", Format: "application/json", Digest: "sha256:test", Units: []source.SourceUnit{ {ID: 1, Kind: "transcript_segment", Text: "The party enters the Old Mill."}, {ID: 2, Kind: "transcript_segment", Text: "They leave the old road behind."}, {ID: 3, Kind: "transcript_segment", Text: "The tavern is quiet."}, }} } func responseSourceRefs(startUnitID, endUnitID int) []locationSourceRefResponse { return []locationSourceRefResponse{{StartUnitID: startUnitID, EndUnitID: endUnitID}} } func newExtractor(t *testing.T, client contracts.StructuredLLMClient, references ...contracts.ReferenceSet) *Extractor { t.Helper() extractor, err := New(client, Options{}, references...) if err != nil { t.Fatalf("New() error = %v, want nil", err) } return extractor } func mismatchedSourceInputRequest(req contracts.TypedExtractionRequest) contracts.TypedExtractionRequest { req.SourceInput = contracts.NewLLMInputMaterial("source", "application/json", []byte(`{"different":true}`), "sha256:other", "file:///other.json") return req } type fakeLocationsLLMClient struct { response extractionResponse content []byte err error requests []contracts.StructuredCompletionRequest } func (client *fakeLocationsLLMClient) CompleteStructured(_ 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") } content := append([]byte(nil), client.content...) if len(content) != 0 { if err := json.Unmarshal(content, target); err != nil { return contracts.StructuredCompletionResponse{}, err } } else { *target = client.response 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() return req }