93 lines
3.3 KiB
Go
93 lines
3.3 KiB
Go
package itemregistry
|
|
|
|
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-items: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-items.json"),
|
|
SessionID: "item-session", LLMProfile: "item-profile",
|
|
}
|
|
}
|
|
|
|
func sourceDocument() *source.SourceDocument {
|
|
return &source.SourceDocument{ID: "session-items", Kind: "transcript", Format: "application/json", Digest: "sha256:test", Units: []source.SourceUnit{
|
|
{ID: 1, Kind: "transcript_segment", Text: "The party finds a rope."},
|
|
{ID: 2, Kind: "transcript_segment", Text: "They recover the Star Compass."},
|
|
{ID: 3, Kind: "transcript_segment", Text: "The chest contains gold pieces."},
|
|
}}
|
|
}
|
|
|
|
func responseSourceRefs(startUnitID, endUnitID int) []itemSourceRefResponse {
|
|
return []itemSourceRefResponse{{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 fakeItemsLLMClient struct {
|
|
response extractionResponse
|
|
content []byte
|
|
err error
|
|
requests []contracts.StructuredCompletionRequest
|
|
}
|
|
|
|
func (client *fakeItemsLLMClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
|
|
client.requests = append(client.requests, cloneStructuredCompletionRequest(req))
|
|
if err := ctx.Err(); err != nil {
|
|
return contracts.StructuredCompletionResponse{}, err
|
|
}
|
|
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
|
|
}
|