Files
notarius/internal/framework/pipeline/runner_test.go

1728 lines
62 KiB
Go

package pipeline
import (
"context"
"errors"
"reflect"
"strings"
"testing"
"time"
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
validate "gitea.maximumdirect.net/eric/notarius/internal/framework/validate"
)
func TestNewAndDataTypes(t *testing.T) {
runner := New(Registries{})
if runner == nil {
t.Fatal("New() = nil, want runner")
}
input := RunInput{
Pipeline: resolvedPipeline(),
SourceID: "source-1",
Path: "input.txt",
RawInput: []byte("source text"),
LLMClient: fakeLLMClient{},
Metadata: map[string]any{"request": "test"},
}
output := RunOutput{
Manifest: artifacts.RunManifest{PipelineID: "pipeline-1"},
Approved: []artifacts.Artifact{{ExtractorKey: "extract-alpha"}},
Rejected: []artifacts.RejectedArtifact{{ValidatorName: "validator"}},
Warnings: []contracts.Warning{{ReasonCode: "note", Message: "message"}},
OutputFiles: []contracts.OutputFile{{Name: "artifacts/generic.json", ContentType: "application/json", Bytes: []byte(`{}`)}},
}
if input.Pipeline.ID != "pipeline-1" || input.SourceID != "source-1" {
t.Fatalf("RunInput = %#v, want constructed fields", input)
}
if output.Manifest.PipelineID != "pipeline-1" || len(output.Approved) != 1 || len(output.Rejected) != 1 || len(output.Warnings) != 1 || len(output.OutputFiles) != 1 {
t.Fatalf("RunOutput = %#v, want constructed fields", output)
}
}
func TestRunRejectsInvalidSetup(t *testing.T) {
tests := []struct {
name string
run func() (RunOutput, error)
error string
}{
{
name: "nil runner",
run: func() (RunOutput, error) { return (*Runner)(nil).Run(context.Background(), RunInput{}) },
error: "runner must not be nil",
},
{
name: "empty pipeline id",
run: func() (RunOutput, error) {
return New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: ResolvedPipeline{Digest: "sha256:pipeline"}})
},
error: "pipeline id",
},
{
name: "empty pipeline digest",
run: func() (RunOutput, error) {
pipeline := resolvedPipeline()
pipeline.Digest = ""
return New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: pipeline})
},
error: "pipeline digest",
},
{
name: "empty artifact lanes",
run: func() (RunOutput, error) {
pipeline := resolvedPipeline()
pipeline.ArtifactLanes = nil
return New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: pipeline})
},
error: "artifact lanes",
},
{
name: "missing input registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Inputs = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "input registry",
},
{
name: "missing chunker registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Chunkers = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "chunker registry",
},
{
name: "missing extractor registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Extractors = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "extractor registry",
},
{
name: "missing merger registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Mergers = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "merger registry",
},
{
name: "missing normalizer registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Normalizers = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "normalizer registry",
},
{
name: "missing validator registry only when configured validators are used",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Validators = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipelineWithValidators("configured")})
},
error: "validator registry",
},
{
name: "missing output registry",
run: func() (RunOutput, error) {
registries := newRunnerRegistries(t, nil)
registries.Outputs = nil
return New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
},
error: "output encoder registry",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
_, err := test.run()
assertRunError(t, err, test.error)
})
}
}
func TestRunAllowsNilValidatorRegistryWithoutConfiguredValidators(t *testing.T) {
registries := newRunnerRegistries(t, nil)
registries.Validators = nil
_, err := New(registries).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
}
func TestRunRejectsInputBuildParseAndInvalidSourceErrors(t *testing.T) {
buildErr := errors.New("build failed")
parseErr := errors.New("parse failed")
tests := []struct {
name string
configure func(*runnerModules)
want string
}{
{
name: "input build",
configure: func(modules *runnerModules) {
modules.inputBuildErr = buildErr
},
want: "build input adapter",
},
{
name: "input parse",
configure: func(modules *runnerModules) {
modules.input.err = parseErr
},
want: "parse input",
},
{
name: "invalid source",
configure: func(modules *runnerModules) {
modules.input.doc = &source.SourceDocument{}
},
want: "validate source document",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
test.configure(modules)
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, test.want)
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
})
}
}
func TestRunRejectsChunkerBuildChunkAndEmptyChunkErrors(t *testing.T) {
chunkErr := errors.New("chunk failed")
tests := []struct {
name string
configure func(*runnerModules)
want string
}{
{
name: "chunker build",
configure: func(modules *runnerModules) {
modules.chunkerBuildErr = errors.New("build failed")
},
want: "build chunker",
},
{
name: "chunker chunk",
configure: func(modules *runnerModules) {
modules.chunker.err = chunkErr
},
want: "chunk source",
},
{
name: "empty chunks",
configure: func(modules *runnerModules) {
modules.chunker.chunks = nil
},
want: "returned no chunks",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
test.configure(modules)
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, test.want)
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
})
}
}
func TestRunRejectsInvalidChunks(t *testing.T) {
tests := []struct {
name string
chunks []contracts.SourceChunk
want string
}{
{
name: "empty chunk id",
chunks: []contracts.SourceChunk{{ID: "", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u1")}}},
want: "id must not be empty",
},
{
name: "duplicate chunk id",
chunks: []contracts.SourceChunk{
{ID: "chunk-0", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u1")}},
{ID: "chunk-0", SourceID: "source-1", Index: 1, Units: []source.SourceUnit{unitWithID("u2")}},
},
want: "duplicated",
},
{
name: "wrong source id",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "other-source", Index: 0, Units: []source.SourceUnit{unitWithID("u1")}}},
want: "source_id",
},
{
name: "wrong index",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "source-1", Index: 1, Units: []source.SourceUnit{unitWithID("u1")}}},
want: "index",
},
{
name: "empty units",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "source-1", Index: 0}},
want: "units must not be empty",
},
{
name: "repeated unit inside chunk",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u1"), unitWithID("u1")}}},
want: "repeats source unit",
},
{
name: "unknown unit",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u9")}}},
want: "was not found",
},
{
name: "units out of source order",
chunks: []contracts.SourceChunk{{ID: "chunk-0", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u2"), unitWithID("u1")}}},
want: "source document order",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
modules.chunker.chunks = test.chunks
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, test.want)
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
if len(modules.extractors["extract-alpha"].requests) != 0 {
t.Fatalf("extractor calls = %d, want none after invalid chunks", len(modules.extractors["extract-alpha"].requests))
}
})
}
}
func TestRunAllowsPartialCoverageAndOverlappingChunks(t *testing.T) {
modules := defaultRunnerModules()
modules.chunker.chunks = []contracts.SourceChunk{
{ID: "chunk-0", SourceID: "source-1", Index: 0, Units: []source.SourceUnit{unitWithID("u1"), unitWithID("u2")}},
{ID: "chunk-1", SourceID: "source-1", Index: 1, Units: []source.SourceUnit{unitWithID("u2")}},
}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(output.Approved) != 2 {
t.Fatalf("len(Approved) = %d, want one candidate per accepted chunk", len(output.Approved))
}
}
func TestRunCanonicalizesChunkUnitsBeforeExtraction(t *testing.T) {
modules := defaultRunnerModules()
modules.input.doc = sourceDocumentWithUnitMetadata()
modules.chunker.chunks = []contracts.SourceChunk{
{
ID: "chunk-0",
SourceID: "source-1",
Index: 0,
Units: []source.SourceUnit{
{
ID: "u1",
Kind: "mutated-kind",
Text: "mutated text",
Metadata: map[string]any{
"speaker": "chunker-speaker",
"note": "chunker note",
},
},
},
},
}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(output.Approved) != 1 {
t.Fatalf("len(Approved) = %d, want 1", len(output.Approved))
}
extractor := modules.extractors["extract-alpha"]
if len(extractor.requests) != 1 {
t.Fatalf("len(extractor requests) = %d, want 1", len(extractor.requests))
}
chunk := extractor.requests[0].Chunk
if chunk == nil {
t.Fatal("extractor chunk = nil, want canonical chunk")
}
if chunk.Units[0].ID != "u1" || chunk.Units[0].Kind != "source-kind" || chunk.Units[0].Text != "source text" {
t.Fatalf("chunk unit = %#v, want source document unit values", chunk.Units[0])
}
if got := chunk.Units[0].Metadata["speaker"]; got != "source-speaker" {
t.Fatalf("chunk unit metadata = %#v, want source document metadata", chunk.Units[0].Metadata)
}
if got := chunk.Units[0].Metadata["topic"]; got != "source-topic" {
t.Fatalf("chunk unit metadata = %#v, want cloned source document metadata", chunk.Units[0].Metadata)
}
modules.input.doc.Units[0].Kind = "changed-kind"
modules.input.doc.Units[0].Text = "changed text"
modules.input.doc.Units[0].Metadata["speaker"] = "changed-speaker"
if chunk.Units[0].Kind != "source-kind" || chunk.Units[0].Text != "source text" || chunk.Units[0].Metadata["speaker"] != "source-speaker" {
t.Fatalf("chunk unit changed after source mutation: %#v", chunk.Units[0])
}
}
func TestRunPreservesChunkMetadataDuringCanonicalization(t *testing.T) {
modules := defaultRunnerModules()
modules.chunker.chunks = []contracts.SourceChunk{
{
ID: "chunk-0",
SourceID: "source-1",
Index: 0,
Units: []source.SourceUnit{
unitWithID("u1"),
},
Metadata: map[string]any{
"scene_title": "Original scene",
"boundary_note": "Chunker note",
},
},
}
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
extractor := modules.extractors["extract-alpha"]
if len(extractor.requests) != 1 || extractor.requests[0].Chunk == nil {
t.Fatalf("extractor requests = %#v, want one canonical chunk", extractor.requests)
}
if got := extractor.requests[0].Chunk.Metadata["scene_title"]; got != "Original scene" {
t.Fatalf("chunk metadata = %#v, want chunker metadata", extractor.requests[0].Chunk.Metadata)
}
if got := extractor.requests[0].Chunk.Metadata["boundary_note"]; got != "Chunker note" {
t.Fatalf("chunk metadata = %#v, want chunker metadata", extractor.requests[0].Chunk.Metadata)
}
modules.chunker.chunks[0].Metadata["scene_title"] = "changed"
modules.chunker.chunks[0].Metadata["boundary_note"] = "changed"
if got := extractor.requests[0].Chunk.Metadata["scene_title"]; got != "Original scene" {
t.Fatalf("chunk metadata aliased to chunker map: %#v", extractor.requests[0].Chunk.Metadata)
}
if got := extractor.requests[0].Chunk.Metadata["boundary_note"]; got != "Chunker note" {
t.Fatalf("chunk metadata aliased to chunker map: %#v", extractor.requests[0].Chunk.Metadata)
}
}
func TestRunExecutesChunksAndPassesChunkAndLLMClient(t *testing.T) {
modules := defaultRunnerModules()
llmClient := fakeLLMClient{}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{
Pipeline: resolvedPipeline(),
LLMClient: llmClient,
Metadata: map[string]any{"request": "test"},
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
extractor := modules.extractors["extract-alpha"]
if !reflect.DeepEqual(extractor.seenChunkIDs, []string{"chunk-0", "chunk-1"}) {
t.Fatalf("seen chunks = %#v, want both chunks", extractor.seenChunkIDs)
}
if len(modules.chunker.requests) != 1 || modules.chunker.requests[0].LLMClient == nil {
t.Fatalf("chunker LLM client = %#v, want client on chunk request", modules.chunker.requests)
}
if len(extractor.seenLLMClients) != 2 || extractor.seenLLMClients[0] == nil || extractor.seenLLMClients[1] == nil {
t.Fatalf("seen LLM clients = %#v, want client for each chunk", extractor.seenLLMClients)
}
normalizer := modules.normalizers["normalize"]
if len(normalizer.requests) != 1 || normalizer.requests[0].LLMClient == nil {
t.Fatalf("normalizer LLM client = %#v, want client on normalize request", normalizer.requests)
}
if extractor.seenMetadata[0]["request"] != "test" {
t.Fatalf("seen metadata = %#v, want request metadata", extractor.seenMetadata)
}
if len(output.Approved) != 2 {
t.Fatalf("len(Approved) = %d, want 2", len(output.Approved))
}
}
func TestRunPassesInputRequestFields(t *testing.T) {
modules := defaultRunnerModules()
metadata := map[string]any{"request": "test"}
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{
Pipeline: resolvedPipeline(),
SourceID: "source-1",
Path: "input.txt",
RawInput: []byte("source text"),
Metadata: metadata,
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(modules.input.requests) != 1 {
t.Fatalf("len(input requests) = %d, want 1", len(modules.input.requests))
}
req := modules.input.requests[0]
if req.SourceID != "source-1" || req.Path != "input.txt" || string(req.Raw) != "source text" {
t.Fatalf("ParseRequest = %#v, want source id, path, and raw input", req)
}
if req.Metadata["request"] != "test" {
t.Fatalf("ParseRequest.Metadata = %#v, want request metadata", req.Metadata)
}
}
func TestRunPassesModuleBindingConfigToStageRequests(t *testing.T) {
modules := defaultRunnerModules()
pipeline := resolvedPipelineWithValidators("configured")
pipeline.Input = ModuleBinding{Module: "input", LLMProfile: "input-profile", Options: map[string]any{"input_option": "input-value"}}
pipeline.Chunk = ModuleBinding{Module: "chunk", LLMProfile: "chunk-profile", Options: map[string]any{"chunk_option": "chunk-value"}}
pipeline.Output = ModuleBinding{Module: "output", LLMProfile: "output-profile", Options: map[string]any{"output_option": "output-value"}}
pipeline.ArtifactLanes[0].Extract = ModuleBinding{Module: "extract-alpha", LLMProfile: "extract-profile", Options: map[string]any{"extract_option": "extract-value"}}
pipeline.ArtifactLanes[0].Merge = ModuleBinding{Module: "merge", LLMProfile: "merge-profile", Options: map[string]any{"merge_option": "merge-value"}}
pipeline.ArtifactLanes[0].Normalize = ModuleBinding{Module: "normalize", LLMProfile: "normalize-profile", Options: map[string]any{"normalize_option": "normalize-value"}}
pipeline.ArtifactLanes[0].Validators[0] = ModuleBinding{Module: "configured", LLMProfile: "validator-profile", Options: map[string]any{"validator_option": "validator-value"}}
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if got := modules.input.requests[0].LLMProfile; got != "input-profile" {
t.Fatalf("input LLMProfile = %q, want input-profile", got)
}
if got := modules.input.requests[0].Options["input_option"]; got != "input-value" {
t.Fatalf("input Options = %#v, want input option", modules.input.requests[0].Options)
}
if got := modules.chunker.requests[0].LLMProfile; got != "chunk-profile" {
t.Fatalf("chunk LLMProfile = %q, want chunk-profile", got)
}
if got := modules.chunker.requests[0].Options["chunk_option"]; got != "chunk-value" {
t.Fatalf("chunk Options = %#v, want chunk option", modules.chunker.requests[0].Options)
}
if got := modules.extractors["extract-alpha"].requests[0].LLMProfile; got != "extract-profile" {
t.Fatalf("extract LLMProfile = %q, want extract-profile", got)
}
if got := modules.extractors["extract-alpha"].requests[0].Options["extract_option"]; got != "extract-value" {
t.Fatalf("extract Options = %#v, want extract option", modules.extractors["extract-alpha"].requests[0].Options)
}
if got := modules.mergers["merge"].requests[0].LLMProfile; got != "merge-profile" {
t.Fatalf("merge LLMProfile = %q, want merge-profile", got)
}
if got := modules.mergers["merge"].requests[0].Options["merge_option"]; got != "merge-value" {
t.Fatalf("merge Options = %#v, want merge option", modules.mergers["merge"].requests[0].Options)
}
if got := modules.normalizers["normalize"].requests[0].LLMProfile; got != "normalize-profile" {
t.Fatalf("normalize LLMProfile = %q, want normalize-profile", got)
}
if got := modules.normalizers["normalize"].requests[0].Options["normalize_option"]; got != "normalize-value" {
t.Fatalf("normalize Options = %#v, want normalize option", modules.normalizers["normalize"].requests[0].Options)
}
if got := modules.validators["configured"].requests[0].LLMProfile; got != "validator-profile" {
t.Fatalf("validator LLMProfile = %q, want validator-profile", got)
}
if got := modules.validators["configured"].requests[0].Options["validator_option"]; got != "validator-value" {
t.Fatalf("validator Options = %#v, want validator option", modules.validators["configured"].requests[0].Options)
}
if got := modules.output.requests[0].LLMProfile; got != "output-profile" {
t.Fatalf("output LLMProfile = %q, want output-profile", got)
}
if got := modules.output.requests[0].Options["output_option"]; got != "output-value" {
t.Fatalf("output Options = %#v, want output option", modules.output.requests[0].Options)
}
}
func TestRunPassesLaneReferencesToExtractorRequests(t *testing.T) {
modules := defaultRunnerModules()
pipeline := resolvedPipeline()
pipeline.ArtifactLanes[0].ExtractReferences.ReferenceSet = testReferenceSet("roster", "reference text")
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
req := modules.extractors["extract-alpha"].requests[0]
item := req.References.Slots["roster"].Items[0]
if string(item.Content) != "reference text" {
t.Fatalf("reference content = %q, want reference text", item.Content)
}
item.Content[0] = 'R'
if got := string(pipeline.ArtifactLanes[0].ExtractReferences.ReferenceSet.Slots["roster"].Items[0].Content); got != "reference text" {
t.Fatalf("runner mutated reference set content = %q", got)
}
}
func TestRunPassesChunkReferencesToChunkerRequest(t *testing.T) {
modules := defaultRunnerModules()
pipeline := resolvedPipeline()
pipeline.ChunkReferences.ReferenceSet = testReferenceSet("scene_guide", "chunk reference text")
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
req := modules.chunker.requests[0]
item := req.References.Slots["scene_guide"].Items[0]
if string(item.Content) != "chunk reference text" {
t.Fatalf("chunk reference content = %q, want chunk reference text", item.Content)
}
item.Content[0] = 'C'
if got := string(pipeline.ChunkReferences.ReferenceSet.Slots["scene_guide"].Items[0].Content); got != "chunk reference text" {
t.Fatalf("runner mutated chunk reference set content = %q", got)
}
}
func TestRunPassesNormalizeReferencesToNormalizerRequest(t *testing.T) {
modules := defaultRunnerModules()
pipeline := resolvedPipeline()
pipeline.ArtifactLanes[0].NormalizeReferences.ReferenceSet = testReferenceSet("normalization_notes", "normalize reference text")
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
req := modules.normalizers["normalize"].requests[0]
item := req.References.Slots["normalization_notes"].Items[0]
if string(item.Content) != "normalize reference text" {
t.Fatalf("normalize reference content = %q, want normalize reference text", item.Content)
}
item.Content[0] = 'N'
if got := string(pipeline.ArtifactLanes[0].NormalizeReferences.ReferenceSet.Slots["normalization_notes"].Items[0].Content); got != "normalize reference text" {
t.Fatalf("runner mutated normalize reference set content = %q", got)
}
}
func TestRunAllowsNilLLMClientWhenModulesDoNotUseIt(t *testing.T) {
_, err := New(newRunnerRegistries(t, defaultRunnerModules())).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil with nil LLM client when modules do not use it", err)
}
}
func TestRunIncludesInputWarnings(t *testing.T) {
modules := defaultRunnerModules()
warning := contracts.Warning{Scope: "reference", ReasonCode: "empty_reference", Message: "empty reference"}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{
Pipeline: resolvedPipeline(),
Warnings: []contracts.Warning{warning},
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(output.Warnings) != 1 || output.Warnings[0] != warning {
t.Fatalf("warnings = %#v, want input warning", output.Warnings)
}
}
func TestRunRecordsTopLevelModuleMetadataForSingletonModules(t *testing.T) {
modules := defaultRunnerModules()
modules.input.manifestMetadata = map[string]any{
"input_profile": "input-metadata",
}
modules.chunker.manifestMetadata = map[string]any{
"prompt_id": "dnd.scenes",
"prompt_version": "v1",
"prompt_sha256": "sha256:chunker-prompt",
"response_schema_key": "dnd_scenes",
"response_schema_id": "schema-dnd-scenes",
"response_schema_name": "dnd_scenes",
}
modules.output.manifestMetadata = map[string]any{
"output_profile": "output-metadata",
}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if output.Manifest.ModuleMetadata == nil {
t.Fatal("ModuleMetadata = nil, want module metadata map")
}
if got := output.Manifest.ModuleMetadata["input"]; !reflect.DeepEqual(got, modules.input.manifestMetadata) {
t.Fatalf("input module metadata = %#v, want %#v", got, modules.input.manifestMetadata)
}
if got := output.Manifest.ModuleMetadata["chunker"]; !reflect.DeepEqual(got, modules.chunker.manifestMetadata) {
t.Fatalf("chunker module metadata = %#v, want %#v", got, modules.chunker.manifestMetadata)
}
if got := output.Manifest.ModuleMetadata["output"]; !reflect.DeepEqual(got, modules.output.manifestMetadata) {
t.Fatalf("output module metadata = %#v, want %#v", got, modules.output.manifestMetadata)
}
modules.chunker.manifestMetadata["prompt_id"] = "changed"
if output.Manifest.ModuleMetadata["chunker"]["prompt_id"] != "dnd.scenes" {
t.Fatalf("chunker module metadata aliased to provider map: %#v", output.Manifest.ModuleMetadata["chunker"])
}
}
func TestRunPassesPerChunkCandidatesToMergeAndNormalize(t *testing.T) {
modules := defaultRunnerModules()
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
merger := modules.mergers["merge"]
if len(merger.requests) != 1 {
t.Fatalf("len(merge requests) = %d, want 1", len(merger.requests))
}
chunkArtifacts := merger.requests[0].ChunkArtifacts
if len(chunkArtifacts) != 2 {
t.Fatalf("len(ChunkArtifacts) = %d, want 2", len(chunkArtifacts))
}
if chunkArtifacts[0].Chunk.ID != "chunk-0" || chunkArtifacts[1].Chunk.ID != "chunk-1" {
t.Fatalf("merge chunks = %#v, want chunk order", chunkArtifacts)
}
if got := candidateIndices(chunkArtifacts[0].Candidates); !reflect.DeepEqual(got, []int{0}) {
t.Fatalf("first chunk candidate indices = %#v, want [0]", got)
}
if got := candidateIndices(chunkArtifacts[1].Candidates); !reflect.DeepEqual(got, []int{1}) {
t.Fatalf("second chunk candidate indices = %#v, want [1]", got)
}
normalizer := modules.normalizers["normalize"]
if len(normalizer.requests) != 1 {
t.Fatalf("len(normalize requests) = %d, want 1", len(normalizer.requests))
}
if got := candidateIndices(normalizer.requests[0].Candidates); !reflect.DeepEqual(got, []int{0, 1}) {
t.Fatalf("normalize candidate indices = %#v, want merged candidates", got)
}
}
func TestRunRejectsInvalidPostNormalizeCandidateEnvelope(t *testing.T) {
tests := []struct {
name string
candidates []artifacts.ArtifactCandidate
want string
}{
{
name: "duplicate index",
candidates: []artifacts.ArtifactCandidate{
runnerCandidate(0),
runnerCandidate(0),
},
want: "duplicated",
},
{
name: "missing extractor key",
candidates: []artifacts.ArtifactCandidate{
{Index: 0, ArtifactType: "artifact", SchemaVersion: "v1", Payload: []byte(`{"value":true}`)},
},
want: "extractor_key",
},
{
name: "mismatched schema version",
candidates: []artifacts.ArtifactCandidate{
{Index: 0, ExtractorKey: "extract-alpha", ArtifactType: "artifact", SchemaVersion: "other", Payload: []byte(`{"value":true}`)},
},
want: "schema_version",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
modules.normalizers["normalize"].result = test.candidates
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, test.want)
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
if len(output.Approved) != 0 {
t.Fatalf("len(Approved) = %d, want no approved artifacts", len(output.Approved))
}
})
}
}
func TestRunValidatorApprovalAndRejection(t *testing.T) {
modules := defaultRunnerModules()
rejectFirst := &runnerValidator{
name: "default-validator",
decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision {
return []contracts.ValidationDecision{
validate.Rejected(candidates[0].Index, "invalid", "not accepted"),
validate.Approved(candidates[1].Index),
}
},
}
modules.extractors["extract-alpha"].validators = []contracts.Validator{rejectFirst}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if output.Manifest.ValidationStatus != "rejected" {
t.Fatalf("ValidationStatus = %q, want rejected", output.Manifest.ValidationStatus)
}
if len(output.Approved) != 1 || len(output.Rejected) != 1 {
t.Fatalf("approved/rejected = %d/%d, want 1/1", len(output.Approved), len(output.Rejected))
}
if output.Rejected[0].ValidatorName != "default-validator" || output.Rejected[0].ReasonCode != "invalid" {
t.Fatalf("Rejected[0] = %#v, want rejection details", output.Rejected[0])
}
}
func TestRunUsesConfiguredValidatorsInLaneOrder(t *testing.T) {
modules := defaultRunnerModules()
var order []string
modules.validators["configured"] = &runnerValidator{name: "configured", decisions: approveAll, order: &order}
modules.validators["second-validator"] = &runnerValidator{name: "second-validator", decisions: approveAll, order: &order}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{
Pipeline: resolvedPipelineWithValidators("configured", "second-validator"),
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if !reflect.DeepEqual(order, []string{"configured", "second-validator"}) {
t.Fatalf("validator order = %#v, want configured order", order)
}
if got := output.Manifest.ArtifactLanes[0].Validators; !reflect.DeepEqual(got, []string{"configured", "second-validator"}) {
t.Fatalf("manifest validators = %#v, want configured validators", got)
}
}
func TestRunUsesDefaultValidatorsWhenLaneDoesNotConfigureValidators(t *testing.T) {
modules := defaultRunnerModules()
defaultValidator := &runnerValidator{name: "default-validator", decisions: approveAll}
configuredValidator := &runnerValidator{name: "configured", decisions: approveAll}
modules.extractors["extract-alpha"].validators = []contracts.Validator{defaultValidator}
modules.validators["configured"] = configuredValidator
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if defaultValidator.calls != 1 {
t.Fatalf("default validator calls = %d, want 1", defaultValidator.calls)
}
if configuredValidator.calls != 0 {
t.Fatalf("configured validator calls = %d, want 0", configuredValidator.calls)
}
}
func TestRunConfiguredValidatorsReplaceExtractorDefaults(t *testing.T) {
modules := defaultRunnerModules()
defaultValidator := &runnerValidator{name: "default-validator", decisions: approveAll}
configuredValidator := &runnerValidator{name: "configured", decisions: approveAll}
modules.extractors["extract-alpha"].validators = []contracts.Validator{defaultValidator}
modules.validators["configured"] = configuredValidator
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipelineWithValidators("configured")})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if defaultValidator.calls != 0 {
t.Fatalf("default validator calls = %d, want 0", defaultValidator.calls)
}
if configuredValidator.calls != 1 {
t.Fatalf("configured validator calls = %d, want 1", configuredValidator.calls)
}
}
func TestRunAssignsGlobalCandidateIndicesAcrossLanesAndChunks(t *testing.T) {
modules := defaultRunnerModules()
var seenIndices []int
recordIndices := func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision {
seenIndices = append(seenIndices, candidateIndices(candidates)...)
return approveAll(candidates)
}
modules.extractors["extract-alpha"].validators = []contracts.Validator{&runnerValidator{name: "alpha-validator", decisions: recordIndices}}
modules.extractors["extract-beta"] = &runnerExtractor{key: "extract-beta", artifactType: "artifact", schemaVersion: "v1", validators: []contracts.Validator{&runnerValidator{name: "beta-validator", decisions: recordIndices}}}
pipeline := resolvedPipeline()
pipeline.ArtifactLanes = append(pipeline.ArtifactLanes, ResolvedArtifactLane{
ID: "beta",
Extract: Binding("extract-beta"),
Merge: Binding("merge"),
Normalize: Binding("normalize"),
})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if !reflect.DeepEqual(seenIndices, []int{0, 1, 2, 3}) {
t.Fatalf("seen indices = %#v, want global indices", seenIndices)
}
if len(output.Approved) != 4 {
t.Fatalf("len(Approved) = %d, want 4", len(output.Approved))
}
}
func TestRunCollectsStageWarnings(t *testing.T) {
modules := defaultRunnerModules()
modules.chunker.warnings = []contracts.Warning{{ReasonCode: "chunk-warning", Message: "chunk warning"}}
modules.extractors["extract-alpha"].warnings = []contracts.Warning{{ReasonCode: "extract-warning", Message: "extract warning"}}
modules.mergers["merge"].warnings = []contracts.Warning{{ReasonCode: "merge-warning", Message: "merge warning"}}
modules.normalizers["normalize"].warnings = []contracts.Warning{{ReasonCode: "normalize-warning", Message: "normalize warning"}}
modules.extractors["extract-alpha"].validators = []contracts.Validator{&runnerValidator{
name: "default-validator",
decisions: approveAll,
warnings: []contracts.Warning{{ReasonCode: "validator-warning", Message: "validator warning"}},
}}
modules.output.warnings = []contracts.Warning{{ReasonCode: "output-warning", Message: "output warning"}}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
want := []string{"chunk-warning", "extract-warning", "extract-warning", "merge-warning", "normalize-warning", "validator-warning", "output-warning"}
if got := warningReasons(output.Warnings); !reflect.DeepEqual(got, want) {
t.Fatalf("warning reasons = %#v, want %#v", got, want)
}
}
func TestRunOutputEncoderReceivesManifestAndArtifacts(t *testing.T) {
modules := defaultRunnerModules()
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(output.OutputFiles) != 1 {
t.Fatalf("len(OutputFiles) = %d, want 1", len(output.OutputFiles))
}
file := output.OutputFiles[0]
if file.Name != "artifacts/generic.json" {
t.Fatalf("OutputFiles[0].Name = %q, want artifacts/generic.json", file.Name)
}
if file.ContentType != "application/json" {
t.Fatalf("ContentType = %q, want application/json", file.ContentType)
}
if string(file.Bytes) != `{"encoded":true}` {
t.Fatalf("OutputFiles[0].Bytes = %s, want encoded payload", file.Bytes)
}
if len(modules.output.requests) != 1 {
t.Fatalf("len(output requests) = %d, want 1", len(modules.output.requests))
}
req := modules.output.requests[0]
if req.Manifest.PipelineID != "pipeline-1" || req.Manifest.PipelineDigest != "sha256:pipeline" {
t.Fatalf("output manifest = %#v, want pipeline details", req.Manifest)
}
if len(req.Approved) != 2 {
t.Fatalf("len(output approved) = %d, want 2", len(req.Approved))
}
}
func TestRunRejectsUnsafeOutputFileNames(t *testing.T) {
tests := []struct {
name string
fileName string
}{
{name: "empty", fileName: ""},
{name: "absolute", fileName: "/tmp/output.json"},
{name: "parent", fileName: "artifacts/../manifest.json"},
{name: "backslash", fileName: `artifacts\manifest.json`},
{name: "unclean", fileName: "artifacts//manifest.json"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
modules.output.files = []contracts.OutputFile{
{Name: test.fileName, ContentType: "application/json", Bytes: []byte(`{}`)},
}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, "output file name")
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
})
}
}
func TestRunReturnsFailedManifestWhenOutputEncoderFails(t *testing.T) {
modules := defaultRunnerModules()
modules.output.err = errors.New("encode failed")
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, "encode failed")
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
if output.Manifest.CompletedAt == nil {
t.Fatal("CompletedAt = nil, want failed run completion timestamp")
}
if len(output.Approved) != 2 {
t.Fatalf("len(Approved) = %d, want partial approved output", len(output.Approved))
}
}
func TestRunManifestIncludesPipelineAndLaneDetails(t *testing.T) {
resolved := resolvedPipelineWithValidators("configured")
resolved.ChunkReferences.ReferenceSet = contracts.ReferenceSet{
Slots: map[string]contracts.ResolvedReferenceSlot{
"scene_guide": {
Slot: contracts.ReferenceSlot{Name: "scene_guide"},
Items: []contracts.ReferenceItem{
{
SlotName: "scene_guide",
MediaType: "text/plain; charset=utf-8",
Content: []byte("chunk reference content"),
Digest: "sha256:chunk-reference",
Origin: contracts.ReferenceOrigin{Type: "file", URI: "file:///tmp/scene-guide.txt"},
SizeBytes: int64(len("chunk reference content")),
BindingSource: contracts.ReferenceBindingSourceCLI,
},
},
},
},
}
resolved.ArtifactLanes[0].ExtractReferences.ReferenceSet = contracts.ReferenceSet{
Slots: map[string]contracts.ResolvedReferenceSlot{
"roster": {
Slot: contracts.ReferenceSlot{Name: "roster"},
Items: []contracts.ReferenceItem{
{
SlotName: "roster",
MediaType: "text/plain; charset=utf-8",
Content: []byte("reference content"),
Digest: "sha256:reference",
Origin: contracts.ReferenceOrigin{Type: "file", URI: "file:///tmp/roster.txt"},
SizeBytes: int64(len("reference content")),
BindingSource: contracts.ReferenceBindingSourceConfig,
},
},
},
},
}
resolved.ArtifactLanes[0].NormalizeReferences.ReferenceSet = contracts.ReferenceSet{
Slots: map[string]contracts.ResolvedReferenceSlot{
"normalization_notes": {
Slot: contracts.ReferenceSlot{Name: "normalization_notes"},
Items: []contracts.ReferenceItem{
{
SlotName: "normalization_notes",
MediaType: "text/plain; charset=utf-8",
Content: []byte("normalize reference content"),
Digest: "sha256:normalize-reference",
Origin: contracts.ReferenceOrigin{Type: "file", URI: "file:///tmp/normalize.txt"},
SizeBytes: int64(len("normalize reference content")),
BindingSource: contracts.ReferenceBindingSourceConfig,
},
},
},
},
}
output, err := New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: resolved})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
manifest := output.Manifest
if manifest.PipelineID != "pipeline-1" || manifest.PipelineDigest != "sha256:pipeline" {
t.Fatalf("manifest pipeline fields = %#v, want pipeline details", manifest)
}
if manifest.InputModule != "input" || manifest.Chunker != "chunk" || manifest.OutputEncoder != "output" {
t.Fatalf("manifest modules = %#v, want input/chunk/output modules", manifest)
}
if !reflect.DeepEqual(manifest.SourceDigests, []string{"sha256:source"}) {
t.Fatalf("SourceDigests = %#v, want source digest", manifest.SourceDigests)
}
if len(manifest.References) != 3 {
t.Fatalf("References = %#v, want three reference provenance entries", manifest.References)
}
chunkReference := manifest.References[0]
if chunkReference.Stage != string(StageChunk) || chunkReference.LaneID != "" || chunkReference.SlotName != "scene_guide" || chunkReference.Digest != "sha256:chunk-reference" {
t.Fatalf("chunk reference provenance = %#v, want chunk slot digest", chunkReference)
}
if chunkReference.OriginType != "file" || chunkReference.OriginURI != "file:///tmp/scene-guide.txt" || chunkReference.MediaType != "text/plain; charset=utf-8" || chunkReference.SizeBytes != int64(len("chunk reference content")) || chunkReference.BindingSource != contracts.ReferenceBindingSourceCLI {
t.Fatalf("chunk reference provenance = %#v, want origin/media/size/source", chunkReference)
}
extractReference := manifest.References[1]
if extractReference.Stage != string(StageExtract) || extractReference.LaneID != "alpha" || extractReference.SlotName != "roster" || extractReference.Digest != "sha256:reference" {
t.Fatalf("extract reference provenance = %#v, want lane slot digest", extractReference)
}
if extractReference.OriginType != "file" || extractReference.OriginURI != "file:///tmp/roster.txt" || extractReference.MediaType != "text/plain; charset=utf-8" || extractReference.SizeBytes != int64(len("reference content")) || extractReference.BindingSource != contracts.ReferenceBindingSourceConfig {
t.Fatalf("extract reference provenance = %#v, want origin/media/size/source", extractReference)
}
normalizeReference := manifest.References[2]
if normalizeReference.Stage != string(StageNormalize) || normalizeReference.LaneID != "alpha" || normalizeReference.SlotName != "normalization_notes" || normalizeReference.Digest != "sha256:normalize-reference" {
t.Fatalf("normalize reference provenance = %#v, want lane slot digest", normalizeReference)
}
if manifest.ValidationStatus != "approved" {
t.Fatalf("ValidationStatus = %q, want approved", manifest.ValidationStatus)
}
if len(manifest.ArtifactLanes) != 1 {
t.Fatalf("len(ArtifactLanes) = %d, want 1", len(manifest.ArtifactLanes))
}
lane := manifest.ArtifactLanes[0]
if lane.ID != "alpha" || lane.Extractor != "extract-alpha" || lane.Merger != "merge" || lane.Normalizer != "normalize" {
t.Fatalf("ArtifactLanes[0] = %#v, want lane details", lane)
}
if !reflect.DeepEqual(lane.Validators, []string{"configured"}) {
t.Fatalf("lane validators = %#v, want configured validator", lane.Validators)
}
}
func TestRunManifestIncludesRunTimingAndLLMProfiles(t *testing.T) {
startedAt := time.Now().Add(-time.Minute).UTC()
profiles := []artifacts.LLMProfileManifest{
{ID: "default", Provider: "openai-compatible", Model: "model-a"},
}
output, err := New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{
Pipeline: resolvedPipeline(),
RunID: "run-test",
StartedAt: startedAt,
LLMProfiles: profiles,
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
manifest := output.Manifest
if manifest.RunID != "run-test" {
t.Fatalf("RunID = %q, want run-test", manifest.RunID)
}
if manifest.StartedAt == nil || !manifest.StartedAt.Equal(startedAt) {
t.Fatalf("StartedAt = %v, want %s", manifest.StartedAt, startedAt)
}
if manifest.CompletedAt == nil || manifest.CompletedAt.Before(startedAt) {
t.Fatalf("CompletedAt = %v, want timestamp after start", manifest.CompletedAt)
}
if !reflect.DeepEqual(manifest.LLMProfiles, profiles) {
t.Fatalf("LLMProfiles = %#v, want %#v", manifest.LLMProfiles, profiles)
}
}
func TestRunManifestGeneratesRunIDAndTimestamps(t *testing.T) {
output, err := New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if !strings.HasPrefix(output.Manifest.RunID, "run-") {
t.Fatalf("RunID = %q, want generated run ID", output.Manifest.RunID)
}
if output.Manifest.StartedAt == nil {
t.Fatal("StartedAt = nil, want generated timestamp")
}
if output.Manifest.CompletedAt == nil {
t.Fatal("CompletedAt = nil, want generated timestamp")
}
}
func TestRunManifestIncludesExtractorMetadata(t *testing.T) {
modules := defaultRunnerModules()
modules.extractors["extract-alpha"].manifestMetadata = map[string]any{
"prompt_id": "test.prompt",
"response_schema_name": "test_schema",
}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
lane := output.Manifest.ArtifactLanes[0]
extractorMetadata, ok := lane.Metadata["extractor"].(map[string]any)
if !ok {
t.Fatalf("lane metadata = %#v, want extractor metadata", lane.Metadata)
}
if extractorMetadata["prompt_id"] != "test.prompt" || extractorMetadata["response_schema_name"] != "test_schema" {
t.Fatalf("extractor metadata = %#v, want prompt and schema metadata", extractorMetadata)
}
if output.Manifest.ModuleMetadata != nil {
if _, ok := output.Manifest.ModuleMetadata["extractor"]; ok {
t.Fatalf("top-level module metadata includes lane metadata key: %#v", output.Manifest.ModuleMetadata)
}
}
}
func TestRunReturnsPartialOutputWhenLaterLaneFails(t *testing.T) {
modules := defaultRunnerModules()
modules.extractors["extract-beta"] = &runnerExtractor{key: "extract-beta", artifactType: "artifact", schemaVersion: "v1", err: errors.New("extract failed")}
pipeline := resolvedPipeline()
pipeline.ArtifactLanes = append(pipeline.ArtifactLanes, ResolvedArtifactLane{
ID: "beta",
Extract: Binding("extract-beta"),
Merge: Binding("merge"),
Normalize: Binding("normalize"),
})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
assertRunError(t, err, "extract failed")
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
if len(output.Approved) != 2 {
t.Fatalf("len(Approved) = %d, want first lane approved output", len(output.Approved))
}
}
func TestRunSurfacesValidatorErrors(t *testing.T) {
tests := []struct {
name string
validator *runnerValidator
want string
}{
{name: "name mismatch", validator: &runnerValidator{name: "default-validator", resultName: "other", decisions: approveAll}, want: "returned result"},
{name: "cardinality", validator: &runnerValidator{name: "default-validator", decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { return nil }}, want: "0 decisions"},
{name: "error", validator: &runnerValidator{name: "default-validator", decisions: approveAll, err: errors.New("validator failed")}, want: "validator failed"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
modules := defaultRunnerModules()
modules.extractors["extract-alpha"].validators = []contracts.Validator{test.validator}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, test.want)
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
})
}
}
func TestRunRejectsNilDefaultValidator(t *testing.T) {
modules := defaultRunnerModules()
modules.extractors["extract-alpha"].validators = []contracts.Validator{nil}
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
assertRunError(t, err, "validator[0] must not be nil")
if output.Manifest.ValidationStatus != "failed" {
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
}
}
func resolvedPipeline() ResolvedPipeline {
return ResolvedPipeline{
ID: "pipeline-1",
Digest: "sha256:pipeline",
Input: Binding("input"),
Chunk: Binding("chunk"),
ChunkReferences: referenceTarget(StageChunk, "", "chunk", nil),
ArtifactLanes: []ResolvedArtifactLane{
{
ID: "alpha",
Extract: Binding("extract-alpha"),
Merge: Binding("merge"),
Normalize: Binding("normalize"),
ExtractReferences: referenceTarget(StageExtract, "alpha", "extract-alpha", nil),
NormalizeReferences: referenceTarget(StageNormalize, "alpha", "normalize", nil),
},
},
Output: Binding("output"),
}
}
func resolvedPipelineWithValidators(validators ...string) ResolvedPipeline {
pipeline := resolvedPipeline()
for _, validator := range validators {
pipeline.ArtifactLanes[0].Validators = append(pipeline.ArtifactLanes[0].Validators, Binding(validator))
}
return pipeline
}
func testReferenceSet(slotName string, content string) contracts.ReferenceSet {
return contracts.ReferenceSet{
Slots: map[string]contracts.ResolvedReferenceSlot{
slotName: {
Slot: contracts.ReferenceSlot{Name: slotName},
Items: []contracts.ReferenceItem{
{
SlotName: slotName,
MediaType: "text/plain; charset=utf-8",
Content: []byte(content),
Digest: "sha256:test",
Origin: contracts.ReferenceOrigin{Type: "file", URI: "file:///tmp/reference.txt"},
SizeBytes: int64(len(content)),
BindingSource: contracts.ReferenceBindingSourceConfig,
},
},
},
},
}
}
type runnerModules struct {
input *runnerInputAdapter
chunker *runnerChunker
extractors map[string]*runnerExtractor
mergers map[string]*runnerMerger
normalizers map[string]*runnerNormalizer
validators map[string]*runnerValidator
output *runnerOutputEncoder
inputBuildErr error
chunkerBuildErr error
}
func defaultRunnerModules() *runnerModules {
return &runnerModules{
input: &runnerInputAdapter{key: "input", doc: validSourceDocument()},
chunker: &runnerChunker{key: "chunk", chunks: []contracts.SourceChunk{sourceChunkWithID("chunk-0", 0), sourceChunkWithID("chunk-1", 1)}},
extractors: map[string]*runnerExtractor{
"extract-alpha": {key: "extract-alpha", artifactType: "artifact", schemaVersion: "v1"},
},
mergers: map[string]*runnerMerger{
"merge": {key: "merge"},
},
normalizers: map[string]*runnerNormalizer{
"normalize": {key: "normalize"},
},
validators: map[string]*runnerValidator{
"configured": {name: "configured", decisions: approveAll},
"second-validator": {name: "second-validator", decisions: approveAll},
},
output: &runnerOutputEncoder{
key: "output",
files: []contracts.OutputFile{
{Name: "artifacts/generic.json", ContentType: "application/json", Bytes: []byte(`{"encoded":true}`)},
},
},
}
}
func newRunnerRegistries(t *testing.T, modules *runnerModules) Registries {
t.Helper()
if modules == nil {
modules = defaultRunnerModules()
}
registries := Registries{
Inputs: NewInputAdapterRegistry(),
Chunkers: NewChunkerRegistry(),
Extractors: NewExtractorRegistry(),
Mergers: NewMergerRegistry(),
Normalizers: NewNormalizerRegistry(),
Validators: NewValidatorRegistry(),
Outputs: NewOutputEncoderRegistry(),
}
if err := registries.Inputs.Register("input", func() (contracts.InputAdapter, error) {
if modules.inputBuildErr != nil {
return nil, modules.inputBuildErr
}
return modules.input, nil
}); err != nil {
t.Fatalf("register input: %v", err)
}
if err := registries.Chunkers.Register("chunk", func() (contracts.Chunker, error) {
if modules.chunkerBuildErr != nil {
return nil, modules.chunkerBuildErr
}
return modules.chunker, nil
}); err != nil {
t.Fatalf("register chunker: %v", err)
}
for key, extractor := range modules.extractors {
extractor := extractor
if err := registries.Extractors.Register(key, func() (contracts.Extractor, error) { return extractor, nil }); err != nil {
t.Fatalf("register extractor %q: %v", key, err)
}
}
for key, merger := range modules.mergers {
merger := merger
if err := registries.Mergers.Register(key, func() (contracts.Merger, error) { return merger, nil }); err != nil {
t.Fatalf("register merger %q: %v", key, err)
}
}
for key, normalizer := range modules.normalizers {
normalizer := normalizer
if err := registries.Normalizers.Register(key, func() (contracts.Normalizer, error) { return normalizer, nil }); err != nil {
t.Fatalf("register normalizer %q: %v", key, err)
}
}
for key, validator := range modules.validators {
validator := validator
if err := registries.Validators.Register(key, func() (contracts.Validator, error) { return validator, nil }); err != nil {
t.Fatalf("register validator %q: %v", key, err)
}
}
if err := registries.Outputs.Register("output", func() (contracts.OutputEncoder, error) { return modules.output, nil }); err != nil {
t.Fatalf("register output: %v", err)
}
return registries
}
type runnerInputAdapter struct {
key string
doc *source.SourceDocument
err error
manifestMetadata map[string]any
requests []contracts.ParseRequest
}
func (adapter *runnerInputAdapter) Key() string {
return adapter.key
}
func (adapter *runnerInputAdapter) Parse(ctx context.Context, req contracts.ParseRequest) (*source.SourceDocument, error) {
adapter.requests = append(adapter.requests, req)
return adapter.doc, adapter.err
}
func (adapter *runnerInputAdapter) ManifestMetadata() map[string]any {
return adapter.manifestMetadata
}
type runnerChunker struct {
key string
chunks []contracts.SourceChunk
warnings []contracts.Warning
err error
manifestMetadata map[string]any
requests []contracts.ChunkRequest
}
func (chunker *runnerChunker) Key() string {
return chunker.key
}
func (chunker *runnerChunker) ReferenceSlots() []contracts.ReferenceSlot {
return nil
}
func (chunker *runnerChunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) {
chunker.requests = append(chunker.requests, req)
return contracts.ChunkResult{
Chunks: chunker.chunks,
Warnings: chunker.warnings,
}, chunker.err
}
func (chunker *runnerChunker) ManifestMetadata() map[string]any {
return chunker.manifestMetadata
}
type runnerExtractor struct {
key string
artifactType string
schemaVersion string
manifestMetadata map[string]any
candidates []artifacts.ArtifactCandidate
validators []contracts.Validator
warnings []contracts.Warning
err error
requests []contracts.ExtractionRequest
seenChunkIDs []string
seenLLMClients []contracts.StructuredLLMClient
seenMetadata []map[string]any
}
func (extractor *runnerExtractor) Key() string {
return extractor.key
}
func (extractor *runnerExtractor) ArtifactType() string {
return extractor.artifactType
}
func (extractor *runnerExtractor) SchemaVersion() string {
return extractor.schemaVersion
}
func (extractor *runnerExtractor) ReferenceSlots() []contracts.ReferenceSlot {
return nil
}
func (extractor *runnerExtractor) ManifestMetadata() map[string]any {
return extractor.manifestMetadata
}
func (extractor *runnerExtractor) Validators() []contracts.Validator {
return extractor.validators
}
func (extractor *runnerExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) {
extractor.requests = append(extractor.requests, req)
if req.Chunk != nil {
extractor.seenChunkIDs = append(extractor.seenChunkIDs, req.Chunk.ID)
}
extractor.seenLLMClients = append(extractor.seenLLMClients, req.LLMClient)
extractor.seenMetadata = append(extractor.seenMetadata, req.Metadata)
candidates := append([]artifacts.ArtifactCandidate(nil), extractor.candidates...)
if len(candidates) == 0 {
candidates = []artifacts.ArtifactCandidate{{Payload: []byte(`{"value":true}`)}}
}
return contracts.ExtractionResult{
Candidates: candidates,
Warnings: extractor.warnings,
}, extractor.err
}
type runnerMerger struct {
key string
result []artifacts.ArtifactCandidate
warnings []contracts.Warning
err error
requests []contracts.MergeRequest
}
func (merger *runnerMerger) Key() string {
return merger.key
}
func (merger *runnerMerger) Merge(ctx context.Context, req contracts.MergeRequest) (contracts.MergeResult, error) {
merger.requests = append(merger.requests, req)
candidates := append([]artifacts.ArtifactCandidate(nil), merger.result...)
if candidates == nil {
for _, chunkArtifacts := range req.ChunkArtifacts {
candidates = append(candidates, chunkArtifacts.Candidates...)
}
}
return contracts.MergeResult{
Candidates: candidates,
Warnings: merger.warnings,
}, merger.err
}
type runnerNormalizer struct {
key string
result []artifacts.ArtifactCandidate
warnings []contracts.Warning
err error
requests []contracts.NormalizeRequest
}
func (normalizer *runnerNormalizer) Key() string {
return normalizer.key
}
func (normalizer *runnerNormalizer) ReferenceSlots() []contracts.ReferenceSlot {
return nil
}
func (normalizer *runnerNormalizer) Normalize(ctx context.Context, req contracts.NormalizeRequest) (contracts.NormalizeResult, error) {
normalizer.requests = append(normalizer.requests, req)
candidates := append([]artifacts.ArtifactCandidate(nil), normalizer.result...)
if candidates == nil {
candidates = append(candidates, req.Candidates...)
}
return contracts.NormalizeResult{
Candidates: candidates,
Warnings: normalizer.warnings,
}, normalizer.err
}
type runnerValidator struct {
name string
resultName string
decisions func([]artifacts.ArtifactCandidate) []contracts.ValidationDecision
warnings []contracts.Warning
err error
order *[]string
calls int
requests []contracts.ValidationRequest
}
func (validator *runnerValidator) Name() string {
return validator.name
}
func (validator *runnerValidator) Validate(ctx context.Context, req contracts.ValidationRequest) (contracts.ValidationResult, error) {
validator.calls++
validator.requests = append(validator.requests, req)
if validator.order != nil {
*validator.order = append(*validator.order, validator.name)
}
resultName := validator.resultName
if resultName == "" {
resultName = validator.name
}
var decisions []contracts.ValidationDecision
if validator.decisions != nil {
decisions = validator.decisions(req.Candidates)
}
return contracts.ValidationResult{
ValidatorName: resultName,
Decisions: decisions,
Warnings: validator.warnings,
}, validator.err
}
type runnerOutputEncoder struct {
key string
files []contracts.OutputFile
warnings []contracts.Warning
err error
manifestMetadata map[string]any
requests []contracts.OutputRequest
}
func (encoder *runnerOutputEncoder) Key() string {
return encoder.key
}
func (encoder *runnerOutputEncoder) Encode(ctx context.Context, req contracts.OutputRequest) (contracts.OutputResult, error) {
encoder.requests = append(encoder.requests, req)
return contracts.OutputResult{
Files: encoder.files,
Warnings: encoder.warnings,
}, encoder.err
}
func (encoder *runnerOutputEncoder) ManifestMetadata() map[string]any {
return encoder.manifestMetadata
}
type fakeLLMClient struct{}
func (client fakeLLMClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
return contracts.StructuredCompletionResponse{}, nil
}
func approveAll(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision {
decisions := make([]contracts.ValidationDecision, 0, len(candidates))
for _, candidate := range candidates {
decisions = append(decisions, validate.Approved(candidate.Index))
}
return decisions
}
func validSourceDocument() *source.SourceDocument {
return &source.SourceDocument{
ID: "source-1",
Kind: "document",
Format: "text/plain",
Digest: "sha256:source",
Units: []source.SourceUnit{
{ID: "u1", Kind: "unit", Text: "Source unit."},
{ID: "u2", Kind: "unit", Text: "Second source unit."},
{ID: "u3", Kind: "unit", Text: "Third source unit."},
},
}
}
func sourceDocumentWithUnitMetadata() *source.SourceDocument {
return &source.SourceDocument{
ID: "source-1",
Kind: "document",
Format: "text/plain",
Digest: "sha256:source",
Units: []source.SourceUnit{
{
ID: "u1",
Kind: "source-kind",
Text: "source text",
Metadata: map[string]any{
"speaker": "source-speaker",
"topic": "source-topic",
},
},
{
ID: "u2",
Kind: "source-kind",
Text: "second source text",
Metadata: map[string]any{
"speaker": "source-speaker-2",
},
},
},
}
}
func sourceChunkWithID(id string, index int) contracts.SourceChunk {
return contracts.SourceChunk{
ID: id,
SourceID: "source-1",
Index: index,
Units: []source.SourceUnit{
{ID: "u1", Kind: "unit", Text: "Source unit."},
},
}
}
func unitWithID(id string) source.SourceUnit {
switch id {
case "u1":
return source.SourceUnit{ID: "u1", Kind: "unit", Text: "Source unit."}
case "u2":
return source.SourceUnit{ID: "u2", Kind: "unit", Text: "Second source unit."}
case "u3":
return source.SourceUnit{ID: "u3", Kind: "unit", Text: "Third source unit."}
default:
return source.SourceUnit{ID: id, Kind: "unit", Text: "Unknown source unit."}
}
}
func warningReasons(warnings []contracts.Warning) []string {
reasons := make([]string, 0, len(warnings))
for _, warning := range warnings {
reasons = append(reasons, warning.ReasonCode)
}
return reasons
}
func candidateIndices(candidates []artifacts.ArtifactCandidate) []int {
indices := make([]int, 0, len(candidates))
for _, candidate := range candidates {
indices = append(indices, candidate.Index)
}
return indices
}
func runnerCandidate(index int) artifacts.ArtifactCandidate {
return artifacts.ArtifactCandidate{
Index: index,
ExtractorKey: "extract-alpha",
ArtifactType: "artifact",
SchemaVersion: "v1",
Payload: []byte(`{"value":true}`),
}
}
func assertRunError(t *testing.T, err error, want string) {
t.Helper()
if err == nil {
t.Fatal("Run() error = nil, want error")
}
if !strings.Contains(err.Error(), want) {
t.Fatalf("Run() error = %q, want substring %q", err.Error(), want)
}
}