275 lines
10 KiB
Go
275 lines
10 KiB
Go
package transcript
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"strings"
|
|
"testing"
|
|
|
|
"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"
|
|
)
|
|
|
|
func runPreparedPipeline(t *testing.T, registries pipeline.Registries, resolved pipeline.ResolvedPipeline, llmClient contracts.StructuredLLMClient, input pipeline.RunInput) (pipeline.RunOutput, error) {
|
|
t.Helper()
|
|
prepared, err := pipeline.Prepare(resolved, registries, pipeline.ModuleDependencies{LLM: llmClient})
|
|
if err != nil {
|
|
return pipeline.RunOutput{}, err
|
|
}
|
|
input.Prepared = prepared
|
|
return pipeline.New().Run(context.Background(), input)
|
|
}
|
|
|
|
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 := runPreparedPipeline(t, seriatimRunnerRegistries(t, extractor), resolved.ResolvedPipeline, nil, pipeline.RunInput{
|
|
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.NormalizeOutputs) != 1 {
|
|
t.Fatalf("len(NormalizeOutputs) = %d, want 1", len(output.NormalizeOutputs))
|
|
}
|
|
|
|
serializedOutput := output.NormalizeOutputs[0]
|
|
if serializedOutput.LaneID != "events" || serializedOutput.NormalizerKey != pipeline.DefaultNormalizeModule || serializedOutput.Artifact.Schema.ID != "fake.event" || serializedOutput.Artifact.Schema.Version != "v1" {
|
|
t.Fatalf("serialized output envelope = %#v, want fake extractor envelope", serializedOutput)
|
|
}
|
|
var payload struct {
|
|
Value string `json:"value"`
|
|
SourceRefs []source.SourceRef `json:"source_refs"`
|
|
}
|
|
if err := json.Unmarshal(serializedOutput.Artifact.Content, &payload); err != nil {
|
|
t.Fatalf("Unmarshal(serialized output) error = %v, want nil", err)
|
|
}
|
|
if len(payload.SourceRefs) != 1 {
|
|
t.Fatalf("len(SourceRefs) = %d, want 1", len(payload.SourceRefs))
|
|
}
|
|
if err := source.ValidateRef(expectedDoc, payload.SourceRefs[0]); err != nil {
|
|
t.Fatalf("ValidateRef() error = %v, want nil", err)
|
|
}
|
|
if payload.SourceRefs[0].StartUnitID != 1 || payload.SourceRefs[0].EndUnitID != 2 {
|
|
t.Fatalf("SourceRefs[0] = %#v, want Seriatim unit IDs", payload.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 := runPreparedPipeline(t, seriatimRunnerRegistries(t, &runnerSeriatimExtractor{}), resolved.ResolvedPipeline, nil, pipeline.RunInput{
|
|
RawInput: []byte(`{"metadata":{},"segments":[]}`),
|
|
SourceID: "invalid-source",
|
|
})
|
|
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[seriatimArtifact]) 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 := pipeline.RegisterExtractor[seriatimArtifact](extractors, pipeline.ModuleSpec{Key: "fake/extract", Stage: pipeline.StageExtract, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: seriatimArtifactKind, Requires: []string{"chunks", "transcript.speaker", "transcript.timestamps"}, Provides: []string{"fake.artifacts"}}, func() (contracts.Extractor[seriatimArtifact], error) {
|
|
return extractor, nil
|
|
}); err != nil {
|
|
t.Fatalf("register extractor: %v", err)
|
|
}
|
|
if err := pipeline.RegisterMerger[seriatimArtifact](mergers, pipeline.ModuleSpec{Key: pipeline.DefaultMergeModule, Stage: pipeline.StageMerge, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: seriatimArtifactKind}, func() (contracts.Merger[seriatimArtifact], error) {
|
|
return fakeMerger{}, nil
|
|
}); err != nil {
|
|
t.Fatalf("register merger: %v", err)
|
|
}
|
|
if err := pipeline.RegisterNormalizer[seriatimArtifact](normalizers, pipeline.ModuleSpec{Key: pipeline.DefaultNormalizeModule, Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: seriatimArtifactKind}, func() (contracts.Normalizer[seriatimArtifact], error) {
|
|
return fakeNormalizer{}, 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)
|
|
}
|
|
|
|
codecs := pipeline.NewArtifactCodecRegistry()
|
|
if err := pipeline.RegisterArtifactCodec(codecs, seriatimArtifactCodec{}); err != nil {
|
|
t.Fatalf("register artifact codec: %v", err)
|
|
}
|
|
return pipeline.Registries{
|
|
Inputs: inputs,
|
|
Chunkers: chunkers,
|
|
ArtifactCodecs: codecs,
|
|
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) Plan(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkPlanResult, error) {
|
|
return contracts.ChunkPlanResult{
|
|
Plan: source.ChunkPlan{SourceDigest: req.Source.Digest, Ranges: []source.ChunkRange{{StartUnitID: req.Source.Units[0].ID, EndUnitID: req.Source.Units[len(req.Source.Units)-1].ID}}},
|
|
}, nil
|
|
}
|
|
|
|
type runnerSeriatimExtractor struct {
|
|
calls int
|
|
}
|
|
|
|
func (e *runnerSeriatimExtractor) Key() string {
|
|
return "fake/extract"
|
|
}
|
|
|
|
func (e *runnerSeriatimExtractor) ReferenceSlots() []contracts.ReferenceSlot {
|
|
return nil
|
|
}
|
|
|
|
func (e *runnerSeriatimExtractor) Extract(ctx context.Context, req contracts.TypedExtractionRequest) (contracts.TypedExtractionResult[seriatimArtifact], error) {
|
|
e.calls++
|
|
if req.Source == nil {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("source must not be nil")
|
|
}
|
|
if req.Chunk == nil {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("chunk must not be nil")
|
|
}
|
|
if got := unitIDs(req.Source.Units); !equalInts(got, []int{1, 2}) {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("source unit IDs = %#v, want Seriatim segment IDs", got)
|
|
}
|
|
if got := unitIDs(req.Chunk.Units); !equalInts(got, []int{1, 2}) {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, 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.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("unit %d missing speaker metadata", unit.ID)
|
|
}
|
|
if _, ok := Start(unit); !ok {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("unit %d missing start metadata", unit.ID)
|
|
}
|
|
if _, ok := End(unit); !ok {
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{}, fmt.Errorf("unit %d missing end metadata", unit.ID)
|
|
}
|
|
}
|
|
|
|
return contracts.TypedExtractionResult[seriatimArtifact]{Value: seriatimArtifact{
|
|
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) []int {
|
|
ids := make([]int, 0, len(units))
|
|
for _, unit := range units {
|
|
ids = append(ids, unit.ID)
|
|
}
|
|
return ids
|
|
}
|
|
|
|
func equalInts(a, b []int) 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[seriatimArtifact] = (*runnerSeriatimExtractor)(nil)
|
|
_ contracts.OutputEncoder = runnerSeriatimOutput{}
|
|
)
|