Files
notarius/internal/modules/input/seriatim/runner_test.go

282 lines
8.7 KiB
Go

package seriatim
import (
"context"
"encoding/json"
"fmt"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
"gitea.maximumdirect.net/eric/notarius/internal/modules/merge/appendorder"
"gitea.maximumdirect.net/eric/notarius/internal/modules/normalize/noop"
)
func TestRunnerProcessesSeriatimInputWithFakeModules(t *testing.T) {
raw := readFixture(t, "testdata/valid_minimal.json")
expectedDoc, err := New().Parse(context.Background(), contracts.ParseRequest{Raw: raw})
if err != nil {
t.Fatalf("Parse() error = %v, want nil", err)
}
resolved, err := loadPipelineConfig(t).Resolve(configResolveInput(t))
if err != nil {
t.Fatalf("Resolve() error = %v, want nil", err)
}
extractor := &runnerSeriatimExtractor{}
output, err := pipeline.New(seriatimRunnerRegistries(t, extractor)).Run(context.Background(), pipeline.RunInput{
Pipeline: resolved.ResolvedPipeline,
RawInput: raw,
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if output.Manifest.InputModule != Key {
t.Fatalf("manifest input module = %q, want %q", output.Manifest.InputModule, Key)
}
if got := output.Manifest.SourceDigests; len(got) != 1 || got[0] != expectedDoc.Digest {
t.Fatalf("manifest source digests = %#v, want %q", got, expectedDoc.Digest)
}
if len(output.Approved) != 1 {
t.Fatalf("len(Approved) = %d, want 1", len(output.Approved))
}
artifact := output.Approved[0]
if artifact.ExtractorKey != "fake/extract" || artifact.ArtifactType != "fake.event" || artifact.SchemaVersion != "v1" {
t.Fatalf("approved artifact envelope = %#v, want fake extractor envelope", artifact)
}
if len(artifact.SourceRefs) != 1 {
t.Fatalf("len(SourceRefs) = %d, want 1", len(artifact.SourceRefs))
}
if err := source.ValidateRef(expectedDoc, artifact.SourceRefs[0]); err != nil {
t.Fatalf("ValidateRef() error = %v, want nil", err)
}
if artifact.SourceRefs[0].StartUnitID != "seg-001" || artifact.SourceRefs[0].EndUnitID != "seg-002" {
t.Fatalf("SourceRefs[0] = %#v, want Seriatim unit IDs", artifact.SourceRefs[0])
}
if extractor.calls != 1 {
t.Fatalf("extractor calls = %d, want 1", extractor.calls)
}
if len(output.OutputFiles) != 1 {
t.Fatalf("len(OutputFiles) = %d, want 1", len(output.OutputFiles))
}
if output.OutputFiles[0].ContentType != "application/json" {
t.Fatalf("ContentType = %q, want application/json", output.OutputFiles[0].ContentType)
}
}
func TestRunnerFailsOnInvalidSeriatimInput(t *testing.T) {
resolved, err := loadPipelineConfig(t).Resolve(configResolveInput(t))
if err != nil {
t.Fatalf("Resolve() error = %v, want nil", err)
}
output, err := pipeline.New(seriatimRunnerRegistries(t, &runnerSeriatimExtractor{})).Run(context.Background(), pipeline.RunInput{
Pipeline: resolved.ResolvedPipeline,
RawInput: []byte(`{"metadata":{},"segments":[]}`),
SourceID: "invalid-source",
LLMClient: nil,
})
if err == nil {
t.Fatal("Run() error = nil, want invalid input error")
}
if !strings.Contains(err.Error(), "parse input with adapter") || !strings.Contains(err.Error(), "seriatim input") {
t.Fatalf("Run() error = %q, want Seriatim parse context", err.Error())
}
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
}
func configResolveInput(t *testing.T) config.ResolveInput {
t.Helper()
return config.ResolveInput{
PipelineID: "seriatim-fixture",
Catalog: seriatimTestCatalog(t, ModuleSpec()),
}
}
func seriatimRunnerRegistries(t *testing.T, extractor contracts.Extractor) pipeline.Registries {
t.Helper()
inputs := pipeline.NewInputAdapterRegistry()
chunkers := pipeline.NewChunkerRegistry()
extractors := pipeline.NewExtractorRegistry()
mergers := pipeline.NewMergerRegistry()
normalizers := pipeline.NewNormalizerRegistry()
outputs := pipeline.NewOutputEncoderRegistry()
if err := Register(inputs); err != nil {
t.Fatalf("register seriatim input: %v", err)
}
if err := chunkers.Register("fake/chunk", func() (contracts.Chunker, error) {
return runnerSeriatimChunker{}, nil
}); err != nil {
t.Fatalf("register chunker: %v", err)
}
if err := extractors.Register("fake/extract", func() (contracts.Extractor, error) {
return extractor, nil
}); err != nil {
t.Fatalf("register extractor: %v", err)
}
if err := mergers.Register(pipeline.DefaultMergeModule, func() (contracts.Merger, error) {
return appendorder.New(), nil
}); err != nil {
t.Fatalf("register merger: %v", err)
}
if err := normalizers.Register(pipeline.DefaultNormalizeModule, func() (contracts.Normalizer, error) {
return noop.New(), nil
}); err != nil {
t.Fatalf("register normalizer: %v", err)
}
if err := outputs.Register(pipeline.DefaultOutputModule, func() (contracts.OutputEncoder, error) {
return runnerSeriatimOutput{}, nil
}); err != nil {
t.Fatalf("register output: %v", err)
}
return pipeline.Registries{
Inputs: inputs,
Chunkers: chunkers,
Extractors: extractors,
Mergers: mergers,
Normalizers: normalizers,
Outputs: outputs,
}
}
type runnerSeriatimChunker struct{}
func (runnerSeriatimChunker) Key() string {
return "fake/chunk"
}
func (runnerSeriatimChunker) ReferenceSlots() []contracts.ReferenceSlot {
return nil
}
func (runnerSeriatimChunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) {
return contracts.ChunkResult{
Chunks: []contracts.SourceChunk{
{
ID: req.Source.ID + ":chunk:0",
SourceID: req.Source.ID,
Index: 0,
Units: append([]source.SourceUnit(nil), req.Source.Units...),
},
},
}, nil
}
type runnerSeriatimExtractor struct {
calls int
}
func (e *runnerSeriatimExtractor) Key() string {
return "fake/extract"
}
func (e *runnerSeriatimExtractor) ArtifactType() string {
return "fake.event"
}
func (e *runnerSeriatimExtractor) SchemaVersion() string {
return "v1"
}
func (e *runnerSeriatimExtractor) ReferenceSlots() []contracts.ReferenceSlot {
return nil
}
func (e *runnerSeriatimExtractor) Validators() []contracts.Validator {
return nil
}
func (e *runnerSeriatimExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) {
e.calls++
if req.Source == nil {
return contracts.ExtractionResult{}, fmt.Errorf("source must not be nil")
}
if req.Chunk == nil {
return contracts.ExtractionResult{}, fmt.Errorf("chunk must not be nil")
}
if got := unitIDs(req.Source.Units); !equalStrings(got, []string{"seg-001", "seg-002"}) {
return contracts.ExtractionResult{}, fmt.Errorf("source unit IDs = %#v, want Seriatim segment IDs", got)
}
if got := unitIDs(req.Chunk.Units); !equalStrings(got, []string{"seg-001", "seg-002"}) {
return contracts.ExtractionResult{}, fmt.Errorf("chunk unit IDs = %#v, want Seriatim segment IDs", got)
}
for _, unit := range req.Chunk.Units {
if speaker, ok := Speaker(unit); !ok || speaker == "" {
return contracts.ExtractionResult{}, fmt.Errorf("unit %q missing speaker metadata", unit.ID)
}
if _, ok := Start(unit); !ok {
return contracts.ExtractionResult{}, fmt.Errorf("unit %q missing start metadata", unit.ID)
}
if _, ok := End(unit); !ok {
return contracts.ExtractionResult{}, fmt.Errorf("unit %q missing end metadata", unit.ID)
}
}
return contracts.ExtractionResult{
Candidates: []artifacts.ArtifactCandidate{
{
Payload: json.RawMessage(`{"value":"seriatim-source-ref"}`),
SourceRefs: []source.SourceRef{
{
SourceID: req.Source.ID,
StartUnitID: req.Chunk.Units[0].ID,
EndUnitID: req.Chunk.Units[len(req.Chunk.Units)-1].ID,
},
},
},
},
}, nil
}
type runnerSeriatimOutput struct{}
func (runnerSeriatimOutput) Key() string {
return pipeline.DefaultOutputModule
}
func (runnerSeriatimOutput) Encode(ctx context.Context, req contracts.OutputRequest) (contracts.OutputResult, error) {
return contracts.OutputResult{
Files: []contracts.OutputFile{
{Name: "output.json", ContentType: "application/json", Bytes: []byte(`{"encoded":true}`)},
},
}, nil
}
func unitIDs(units []source.SourceUnit) []string {
ids := make([]string, 0, len(units))
for _, unit := range units {
ids = append(ids, unit.ID)
}
return ids
}
func equalStrings(a, b []string) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
var (
_ contracts.Chunker = runnerSeriatimChunker{}
_ contracts.Extractor = (*runnerSeriatimExtractor)(nil)
_ contracts.OutputEncoder = runnerSeriatimOutput{}
)