299 lines
9.0 KiB
Go
299 lines
9.0 KiB
Go
package contracts_test
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"testing"
|
|
|
|
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
|
|
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
|
|
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
|
)
|
|
|
|
var _ contracts.InputAdapter = compositionAdapter{}
|
|
var _ contracts.Chunker = compositionChunker{}
|
|
var _ contracts.Extractor = compositionExtractor{}
|
|
var _ contracts.Merger = compositionMerger{}
|
|
var _ contracts.Normalizer = compositionNormalizer{}
|
|
var _ contracts.Validator = compositionValidator{}
|
|
var _ contracts.StructuredLLMClient = compositionLLMClient{}
|
|
var _ contracts.OutputEncoder = compositionOutputEncoder{}
|
|
|
|
func TestContractsComposeAcrossPackages(t *testing.T) {
|
|
ctx := context.Background()
|
|
adapter := compositionAdapter{}
|
|
chunker := compositionChunker{}
|
|
extractor := compositionExtractor{}
|
|
merger := compositionMerger{}
|
|
normalizer := compositionNormalizer{}
|
|
encoder := compositionOutputEncoder{}
|
|
|
|
doc, err := adapter.Parse(ctx, contracts.ParseRequest{SourceID: "source-1"})
|
|
if err != nil {
|
|
t.Fatalf("Parse() error = %v, want nil", err)
|
|
}
|
|
if err := source.ValidateDocument(doc); err != nil {
|
|
t.Fatalf("ValidateDocument() error = %v, want nil", err)
|
|
}
|
|
|
|
chunking, err := chunker.Chunk(ctx, contracts.ChunkRequest{
|
|
Source: doc,
|
|
LLMClient: compositionLLMClient{},
|
|
Metadata: map[string]any{"max_units": 2},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Chunk() error = %v, want nil", err)
|
|
}
|
|
if len(chunking.Chunks) != 1 {
|
|
t.Fatalf("len(Chunks) = %d, want 1", len(chunking.Chunks))
|
|
}
|
|
|
|
extraction, err := extractor.Extract(ctx, contracts.ExtractionRequest{
|
|
Source: doc,
|
|
Chunk: &chunking.Chunks[0],
|
|
AmbientContext: map[string]any{"synopsis": "example synopsis"},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Extract() error = %v, want nil", err)
|
|
}
|
|
if extraction.Output.Payload.MediaType != "application/json" {
|
|
t.Fatalf("extract media type = %q, want application/json", extraction.Output.Payload.MediaType)
|
|
}
|
|
|
|
merge, err := merger.Merge(ctx, contracts.MergeRequest{
|
|
Source: doc,
|
|
LaneID: "generic-lane",
|
|
ExtractOutputs: []contracts.ExtractOutput{extraction.Output},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Merge() error = %v, want nil", err)
|
|
}
|
|
if string(merge.Output.Payload.Content) != `{"value":"example"}` {
|
|
t.Fatalf("merge output = %s, want extract payload", merge.Output.Payload.Content)
|
|
}
|
|
|
|
normalize, err := normalizer.Normalize(ctx, contracts.NormalizeRequest{
|
|
Source: doc,
|
|
LaneID: "generic-lane",
|
|
MergeOutput: merge.Output,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Normalize() error = %v, want nil", err)
|
|
}
|
|
if string(normalize.Output.Payload.Content) != `{"value":"example"}` {
|
|
t.Fatalf("normalize output = %s, want merge payload", normalize.Output.Payload.Content)
|
|
}
|
|
|
|
output, err := encoder.Encode(ctx, contracts.OutputRequest{
|
|
Manifest: artifacts.RunManifest{RunID: "run-1"},
|
|
NormalizeOutputs: []contracts.NormalizeOutput{normalize.Output},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Encode() error = %v, want nil", err)
|
|
}
|
|
if len(output.Files) != 1 {
|
|
t.Fatalf("len(Files) = %d, want 1", len(output.Files))
|
|
}
|
|
if output.Files[0].ContentType != "application/json" {
|
|
t.Fatalf("ContentType = %q, want application/json", output.Files[0].ContentType)
|
|
}
|
|
if len(output.Files[0].Bytes) == 0 {
|
|
t.Fatal("len(Bytes) = 0, want encoded bytes")
|
|
}
|
|
}
|
|
|
|
type compositionAdapter struct{}
|
|
|
|
func (adapter compositionAdapter) Key() string {
|
|
return "generic-input"
|
|
}
|
|
|
|
func (adapter compositionAdapter) Parse(ctx context.Context, req contracts.ParseRequest) (*source.SourceDocument, error) {
|
|
return &source.SourceDocument{
|
|
ID: req.SourceID,
|
|
Kind: "document",
|
|
Format: "text/plain",
|
|
Digest: "sha256:abc123",
|
|
Units: []source.SourceUnit{
|
|
{ID: 1, Kind: "unit", Text: "First source unit."},
|
|
{ID: 2, Kind: "unit", Text: "Second source unit."},
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
type compositionChunker struct{}
|
|
|
|
func (chunker compositionChunker) Key() string {
|
|
return "generic-chunker"
|
|
}
|
|
|
|
func (chunker compositionChunker) ReferenceSlots() []contracts.ReferenceSlot {
|
|
return nil
|
|
}
|
|
|
|
func (chunker compositionChunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) {
|
|
if req.Source == nil {
|
|
return contracts.ChunkResult{}, errors.New("source document is required")
|
|
}
|
|
if req.LLMClient == nil {
|
|
return contracts.ChunkResult{}, errors.New("structured llm client is required")
|
|
}
|
|
|
|
return contracts.ChunkResult{
|
|
Chunks: []contracts.SourceChunk{
|
|
{
|
|
ID: req.Source.ID + ":chunk:0",
|
|
SourceID: req.Source.ID,
|
|
Index: 0,
|
|
StartUnitID: req.Source.Units[0].ID,
|
|
EndUnitID: req.Source.Units[len(req.Source.Units)-1].ID,
|
|
Content: []byte(`{"units":[{"id":1,"kind":"unit","text":"First source unit."},{"id":2,"kind":"unit","text":"Second source unit."}]}`),
|
|
MediaType: "application/json",
|
|
Units: append([]source.SourceUnit(nil), req.Source.Units...),
|
|
Metadata: map[string]any{"strategy": "whole-document"},
|
|
},
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
type compositionLLMClient struct{}
|
|
|
|
func (client compositionLLMClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
|
|
return contracts.StructuredCompletionResponse{}, nil
|
|
}
|
|
|
|
type compositionExtractor struct{}
|
|
|
|
func (extractor compositionExtractor) Key() string {
|
|
return "generic-extractor"
|
|
}
|
|
|
|
func (extractor compositionExtractor) ReferenceSlots() []contracts.ReferenceSlot {
|
|
return nil
|
|
}
|
|
|
|
func (extractor compositionExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) {
|
|
if req.Source == nil {
|
|
return contracts.ExtractionResult{}, errors.New("source document is required")
|
|
}
|
|
if req.AmbientContext["synopsis"] == "" {
|
|
return contracts.ExtractionResult{}, errors.New("ambient synopsis is required")
|
|
}
|
|
|
|
return contracts.ExtractionResult{
|
|
Output: contracts.ExtractOutput{
|
|
Schema: contracts.ResponseSchema{ID: "schema-id", Name: "schema-name", Version: "v1"},
|
|
Payload: contracts.RawPayload{
|
|
Content: []byte(`{"value":"example"}`),
|
|
MediaType: "application/json",
|
|
},
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
type compositionMerger struct{}
|
|
|
|
func (merger compositionMerger) Key() string {
|
|
return "generic-merger"
|
|
}
|
|
|
|
func (merger compositionMerger) Merge(ctx context.Context, req contracts.MergeRequest) (contracts.MergeResult, error) {
|
|
output := req.ExtractOutputs[0]
|
|
return contracts.MergeResult{Output: contracts.MergeOutput{
|
|
LaneID: req.LaneID,
|
|
MergerKey: merger.Key(),
|
|
SourceID: output.SourceID,
|
|
Schema: output.Schema,
|
|
Payload: cloneCompositionPayload(output.Payload),
|
|
}}, nil
|
|
}
|
|
|
|
type compositionNormalizer struct{}
|
|
|
|
func (normalizer compositionNormalizer) Key() string {
|
|
return "generic-normalizer"
|
|
}
|
|
|
|
func (normalizer compositionNormalizer) ReferenceSlots() []contracts.ReferenceSlot {
|
|
return nil
|
|
}
|
|
|
|
func (normalizer compositionNormalizer) Normalize(ctx context.Context, req contracts.NormalizeRequest) (contracts.NormalizeResult, error) {
|
|
return contracts.NormalizeResult{Output: contracts.NormalizeOutput{
|
|
LaneID: req.LaneID,
|
|
NormalizerKey: normalizer.Key(),
|
|
SourceID: req.MergeOutput.SourceID,
|
|
Schema: req.MergeOutput.Schema,
|
|
Payload: cloneCompositionPayload(req.MergeOutput.Payload),
|
|
}}, nil
|
|
}
|
|
|
|
func cloneCompositionPayload(payload contracts.RawPayload) contracts.RawPayload {
|
|
return contracts.RawPayload{
|
|
Content: append([]byte(nil), payload.Content...),
|
|
MediaType: payload.MediaType,
|
|
Metadata: cloneCompositionMetadata(payload.Metadata),
|
|
Warnings: append([]contracts.Warning(nil), payload.Warnings...),
|
|
}
|
|
}
|
|
|
|
func cloneCompositionMetadata(metadata map[string]any) map[string]any {
|
|
if len(metadata) == 0 {
|
|
return nil
|
|
}
|
|
out := make(map[string]any, len(metadata))
|
|
for key, value := range metadata {
|
|
out[key] = value
|
|
}
|
|
return out
|
|
}
|
|
|
|
type compositionValidator struct{}
|
|
|
|
func (validator compositionValidator) Name() string {
|
|
return "generic-validator"
|
|
}
|
|
|
|
func (validator compositionValidator) ExecutionClass() contracts.ExecutionClass {
|
|
return contracts.ExecutionClassDeterministic
|
|
}
|
|
|
|
func (validator compositionValidator) Validate(ctx context.Context, req contracts.ValidationRequest) (contracts.ValidationResult, error) {
|
|
return contracts.ValidationResult{
|
|
Approved: true,
|
|
ReasonCode: "accepted",
|
|
Message: "output accepted",
|
|
}, nil
|
|
}
|
|
|
|
type compositionOutputEncoder struct{}
|
|
|
|
func (encoder compositionOutputEncoder) Key() string {
|
|
return "generic-output"
|
|
}
|
|
|
|
func (encoder compositionOutputEncoder) Encode(ctx context.Context, req contracts.OutputRequest) (contracts.OutputResult, error) {
|
|
payload := struct {
|
|
RunID string `json:"run_id"`
|
|
OutputCount int `json:"output_count"`
|
|
}{
|
|
RunID: req.Manifest.RunID,
|
|
OutputCount: len(req.NormalizeOutputs),
|
|
}
|
|
encoded, err := json.Marshal(payload)
|
|
if err != nil {
|
|
return contracts.OutputResult{}, err
|
|
}
|
|
|
|
return contracts.OutputResult{
|
|
Files: []contracts.OutputFile{
|
|
{
|
|
Name: "artifacts/generic.json",
|
|
ContentType: "application/json",
|
|
Bytes: encoded,
|
|
},
|
|
},
|
|
}, nil
|
|
}
|