diff --git a/internal/framework/pipeline/registry_integration_test.go b/internal/framework/pipeline/registry_integration_test.go index 494f546..86e25e1 100644 --- a/internal/framework/pipeline/registry_integration_test.go +++ b/internal/framework/pipeline/registry_integration_test.go @@ -12,55 +12,129 @@ import ( validate "gitea.maximumdirect.net/eric/notarius/internal/framework/validate" ) -func TestRunnerUsesExtractorRegistry(t *testing.T) { - var builtKeys []string - var executedKeys []string - registry := NewExtractorRegistry() +func TestRunnerUsesRegistries(t *testing.T) { + var built []string + var executed []string + registries := integrationRegistries(t, &built, &executed) - registerIntegrationExtractor(t, registry, "second", &builtKeys, &executedKeys, []contracts.Validator{ - integrationValidator{name: "reject-second", approve: false}, - }) - registerIntegrationExtractor(t, registry, "first", &builtKeys, &executedKeys, []contracts.Validator{ - integrationValidator{name: "approve-first", approve: true}, - }) - - output, err := New(registry).Run(context.Background(), RunInput{ - Source: integrationSourceDocument(), - ExtractorKeys: []string{"second", "first"}, + output, err := New(registries).Run(context.Background(), RunInput{ + Pipeline: integrationPipeline(), + SourceID: "source-1", + RawInput: []byte("source text"), }) if err != nil { t.Fatalf("Run() error = %v, want nil", err) } - if !reflect.DeepEqual(builtKeys, []string{"second", "first"}) { - t.Fatalf("built keys = %#v, want configured order", builtKeys) + wantBuilt := []string{"input", "chunk", "extract-first", "merge", "normalize", "extract-second", "merge", "normalize", "output"} + if !reflect.DeepEqual(built, wantBuilt) { + t.Fatalf("built = %#v, want %#v", built, wantBuilt) } - if !reflect.DeepEqual(executedKeys, []string{"second", "first"}) { - t.Fatalf("executed keys = %#v, want configured order", executedKeys) + if !reflect.DeepEqual(executed, []string{"extract-first:chunk-0", "extract-second:chunk-0"}) { + t.Fatalf("executed = %#v, want extractor chunk execution", executed) } - if got := artifactKeys(output.Approved); !reflect.DeepEqual(got, []string{"first"}) { - t.Fatalf("approved keys = %#v, want [first]", got) + if got := artifactKeys(output.Approved); !reflect.DeepEqual(got, []string{"extract-first"}) { + t.Fatalf("approved keys = %#v, want [extract-first]", got) } - if got := rejectedKeys(output.Rejected); !reflect.DeepEqual(got, []string{"second"}) { - t.Fatalf("rejected keys = %#v, want [second]", got) + if got := rejectedKeys(output.Rejected); !reflect.DeepEqual(got, []string{"extract-second"}) { + t.Fatalf("rejected keys = %#v, want [extract-second]", got) } } -func registerIntegrationExtractor(t *testing.T, registry *ExtractorRegistry, key string, builtKeys *[]string, executedKeys *[]string, validators []contracts.Validator) { +func integrationRegistries(t *testing.T, built, executed *[]string) Registries { + t.Helper() + + registries := Registries{ + Inputs: NewInputAdapterRegistry(), + Chunkers: NewChunkerRegistry(), + Extractors: NewExtractorRegistry(), + Mergers: NewMergerRegistry(), + Normalizers: NewNormalizerRegistry(), + Outputs: NewOutputEncoderRegistry(), + } + if err := registries.Inputs.Register("input", func() (contracts.InputAdapter, error) { + *built = append(*built, "input") + return integrationInput{}, nil + }); err != nil { + t.Fatalf("register input: %v", err) + } + if err := registries.Chunkers.Register("chunk", func() (contracts.Chunker, error) { + *built = append(*built, "chunk") + return integrationChunker{}, nil + }); err != nil { + t.Fatalf("register chunker: %v", err) + } + registerIntegrationExtractor(t, registries.Extractors, "extract-first", built, executed, []contracts.Validator{ + integrationValidator{name: "approve-first", approve: true}, + }) + registerIntegrationExtractor(t, registries.Extractors, "extract-second", built, executed, []contracts.Validator{ + integrationValidator{name: "reject-second", approve: false}, + }) + if err := registries.Mergers.Register("merge", func() (contracts.Merger, error) { + *built = append(*built, "merge") + return integrationMerger{}, nil + }); err != nil { + t.Fatalf("register merger: %v", err) + } + if err := registries.Normalizers.Register("normalize", func() (contracts.Normalizer, error) { + *built = append(*built, "normalize") + return integrationNormalizer{}, nil + }); err != nil { + t.Fatalf("register normalizer: %v", err) + } + if err := registries.Outputs.Register("output", func() (contracts.OutputEncoder, error) { + *built = append(*built, "output") + return integrationOutput{}, nil + }); err != nil { + t.Fatalf("register output: %v", err) + } + return registries +} + +func registerIntegrationExtractor(t *testing.T, registry *ExtractorRegistry, key string, built, executed *[]string, validators []contracts.Validator) { t.Helper() if err := registry.Register(key, func() (contracts.Extractor, error) { - *builtKeys = append(*builtKeys, key) - return integrationExtractor{key: key, executedKeys: executedKeys, validators: validators}, nil + *built = append(*built, key) + return integrationExtractor{key: key, executed: executed, validators: validators}, nil }); err != nil { t.Fatalf("Register(%q) error = %v, want nil", key, err) } } +type integrationInput struct{} + +func (input integrationInput) Key() string { + return "input" +} + +func (input integrationInput) Parse(ctx context.Context, req contracts.ParseRequest) (*source.SourceDocument, error) { + return integrationSourceDocument(), nil +} + +type integrationChunker struct{} + +func (chunker integrationChunker) Key() string { + return "chunk" +} + +func (chunker integrationChunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) { + return contracts.ChunkResult{ + Chunks: []contracts.SourceChunk{ + { + ID: "chunk-0", + SourceID: req.Source.ID, + Index: 0, + Units: req.Source.Units, + }, + }, + }, nil +} + type integrationExtractor struct { - key string - executedKeys *[]string - validators []contracts.Validator + key string + executed *[]string + validators []contracts.Validator } func (extractor integrationExtractor) Key() string { @@ -80,7 +154,7 @@ func (extractor integrationExtractor) Validators() []contracts.Validator { } func (extractor integrationExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) { - *extractor.executedKeys = append(*extractor.executedKeys, extractor.key) + *extractor.executed = append(*extractor.executed, extractor.key+":"+req.Chunk.ID) return contracts.ExtractionResult{ Candidates: []artifacts.ArtifactCandidate{ {Payload: []byte(`{"value":true}`)}, @@ -88,6 +162,40 @@ func (extractor integrationExtractor) Extract(ctx context.Context, req contracts }, nil } +type integrationNormalizer struct{} + +type integrationMerger struct{} + +func (merger integrationMerger) Key() string { + return "merge" +} + +func (merger integrationMerger) Merge(ctx context.Context, req contracts.MergeRequest) (contracts.MergeResult, error) { + var candidates []artifacts.ArtifactCandidate + for _, chunkArtifacts := range req.ChunkArtifacts { + candidates = append(candidates, chunkArtifacts.Candidates...) + } + return contracts.MergeResult{Candidates: candidates}, nil +} + +func (normalizer integrationNormalizer) Key() string { + return "normalize" +} + +func (normalizer integrationNormalizer) Normalize(ctx context.Context, req contracts.NormalizeRequest) (contracts.NormalizeResult, error) { + return contracts.NormalizeResult{Candidates: req.Candidates}, nil +} + +type integrationOutput struct{} + +func (output integrationOutput) Key() string { + return "output" +} + +func (output integrationOutput) Encode(ctx context.Context, req contracts.OutputRequest) (contracts.OutputResult, error) { + return contracts.OutputResult{Bytes: []byte(`{}`), ContentType: "application/json"}, nil +} + type integrationValidator struct { name string approve bool @@ -112,6 +220,30 @@ func (validator integrationValidator) Validate(ctx context.Context, req contract }, nil } +func integrationPipeline() ResolvedPipeline { + return ResolvedPipeline{ + ID: "pipeline-1", + Digest: "sha256:pipeline", + Input: Binding("input"), + Chunk: Binding("chunk"), + ArtifactLanes: []ResolvedArtifactLane{ + { + ID: "first", + Extract: Binding("extract-first"), + Merge: Binding("merge"), + Normalize: Binding("normalize"), + }, + { + ID: "second", + Extract: Binding("extract-second"), + Merge: Binding("merge"), + Normalize: Binding("normalize"), + }, + }, + Output: Binding("output"), + } +} + func integrationSourceDocument() *source.SourceDocument { return &source.SourceDocument{ ID: "source-1", diff --git a/internal/framework/pipeline/runner.go b/internal/framework/pipeline/runner.go index 559871c..7e51dd7 100644 --- a/internal/framework/pipeline/runner.go +++ b/internal/framework/pipeline/runner.go @@ -10,29 +10,40 @@ import ( "gitea.maximumdirect.net/eric/notarius/internal/framework/validate" ) -type ExtractorFactory interface { - Build(key string) (contracts.Extractor, error) +type Registries struct { + Inputs *InputAdapterRegistry + Chunkers *ChunkerRegistry + Extractors *ExtractorRegistry + Mergers *MergerRegistry + Normalizers *NormalizerRegistry + Validators *ValidatorRegistry + Outputs *OutputEncoderRegistry } type Runner struct { - extractors ExtractorFactory + registries Registries } -func New(extractors ExtractorFactory) *Runner { - return &Runner{extractors: extractors} +func New(registries Registries) *Runner { + return &Runner{registries: registries} } type RunInput struct { - Source *source.SourceDocument - ExtractorKeys []string - LLMClient contracts.StructuredLLMClient - Metadata map[string]any + Pipeline ResolvedPipeline + SourceID string + Path string + RawInput []byte + LLMClient contracts.StructuredLLMClient + Metadata map[string]any } type RunOutput struct { - Approved []artifacts.Artifact `json:"approved,omitempty"` - Rejected []artifacts.RejectedArtifact `json:"rejected,omitempty"` - Warnings []contracts.Warning `json:"warnings,omitempty"` + Manifest artifacts.RunManifest `json:"manifest"` + Approved []artifacts.Artifact `json:"approved,omitempty"` + Rejected []artifacts.RejectedArtifact `json:"rejected,omitempty"` + Warnings []contracts.Warning `json:"warnings,omitempty"` + EncodedOutput []byte `json:"-"` + ContentType string `json:"content_type,omitempty"` } func (r *Runner) Run(ctx context.Context, input RunInput) (RunOutput, error) { @@ -40,54 +51,283 @@ func (r *Runner) Run(ctx context.Context, input RunInput) (RunOutput, error) { if r == nil { return output, fmt.Errorf("runner must not be nil") } - if r.extractors == nil { - return output, fmt.Errorf("runner extractor factory must not be nil") + if err := validateRunInput(input); err != nil { + return output, err } - if err := source.ValidateDocument(input.Source); err != nil { - return output, fmt.Errorf("validate source document: %w", err) + if err := r.validateRegistries(input.Pipeline); err != nil { + return output, err } - if len(input.ExtractorKeys) == 0 { - return output, fmt.Errorf("extractor keys must not be empty") + + output.Manifest = manifestFromPipeline(input.Pipeline) + + adapter, err := r.registries.Inputs.Build(input.Pipeline.Input.Module) + if err != nil { + return failOutput(output), fmt.Errorf("build input adapter %q: %w", input.Pipeline.Input.Module, err) + } + doc, err := adapter.Parse(ctx, contracts.ParseRequest{ + SourceID: input.SourceID, + Path: input.Path, + Raw: input.RawInput, + Metadata: input.Metadata, + }) + if err != nil { + return failOutput(output), fmt.Errorf("parse input with adapter %q: %w", adapter.Key(), err) + } + if err := source.ValidateDocument(doc); err != nil { + return failOutput(output), fmt.Errorf("validate source document: %w", err) + } + output.Manifest.SourceDigests = []string{doc.Digest} + + chunker, err := r.registries.Chunkers.Build(input.Pipeline.Chunk.Module) + if err != nil { + return failOutput(output), fmt.Errorf("build chunker %q: %w", input.Pipeline.Chunk.Module, err) + } + chunkResult, err := chunker.Chunk(ctx, contracts.ChunkRequest{ + Source: doc, + Metadata: input.Metadata, + }) + output.Warnings = append(output.Warnings, chunkResult.Warnings...) + if err != nil { + return failOutput(output), fmt.Errorf("chunk source with chunker %q: %w", chunker.Key(), err) + } + if len(chunkResult.Chunks) == 0 { + return failOutput(output), fmt.Errorf("chunker %q returned no chunks", chunker.Key()) } nextCandidateIndex := 0 - for _, extractorKey := range input.ExtractorKeys { - extractor, err := r.extractors.Build(extractorKey) - if err != nil { - return output, fmt.Errorf("build extractor %q: %w", extractorKey, err) - } - if extractor == nil { - return output, fmt.Errorf("build extractor %q: returned nil extractor", extractorKey) + for _, lane := range input.Pipeline.ArtifactLanes { + if err := r.runLane(ctx, input, doc, chunkResult.Chunks, lane, &output, &nextCandidateIndex); err != nil { + return failOutput(output), err } + } + if len(output.Rejected) > 0 { + output.Manifest.ValidationStatus = "rejected" + } else { + output.Manifest.ValidationStatus = "approved" + } + + encoder, err := r.registries.Outputs.Build(input.Pipeline.Output.Module) + if err != nil { + return failOutput(output), fmt.Errorf("build output encoder %q: %w", input.Pipeline.Output.Module, err) + } + encoded, err := encoder.Encode(ctx, contracts.OutputRequest{ + Manifest: output.Manifest, + Approved: output.Approved, + Rejected: output.Rejected, + Warnings: output.Warnings, + Metadata: input.Metadata, + }) + output.Warnings = append(output.Warnings, encoded.Warnings...) + if err != nil { + return failOutput(output), fmt.Errorf("encode output with encoder %q: %w", encoder.Key(), err) + } + output.EncodedOutput = encoded.Bytes + output.ContentType = encoded.ContentType + + return output, nil +} + +func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.SourceDocument, chunks []contracts.SourceChunk, lane ResolvedArtifactLane, output *RunOutput, nextCandidateIndex *int) error { + extractor, err := r.registries.Extractors.Build(lane.Extract.Module) + if err != nil { + return fmt.Errorf("build extractor %q for lane %q: %w", lane.Extract.Module, lane.ID, err) + } + merger, err := r.registries.Mergers.Build(lane.Merge.Module) + if err != nil { + return fmt.Errorf("build merger %q for lane %q: %w", lane.Merge.Module, lane.ID, err) + } + normalizer, err := r.registries.Normalizers.Build(lane.Normalize.Module) + if err != nil { + return fmt.Errorf("build normalizer %q for lane %q: %w", lane.Normalize.Module, lane.ID, err) + } + + var validators []contracts.Validator + if len(lane.Validators) > 0 { + validators, err = r.buildConfiguredValidators(lane) + if err != nil { + return err + } + } else { + validators = extractor.Validators() + } + + chunkArtifacts := make([]contracts.ChunkArtifacts, 0, len(chunks)) + for index := range chunks { + chunk := chunks[index] result, err := extractor.Extract(ctx, contracts.ExtractionRequest{ - Source: input.Source, + Source: doc, + Chunk: &chunk, LLMClient: input.LLMClient, Metadata: input.Metadata, }) output.Warnings = append(output.Warnings, result.Warnings...) if err != nil { - return output, fmt.Errorf("extract with extractor %q: %w", extractor.Key(), err) + return fmt.Errorf("extract lane %q chunk %q with extractor %q: %w", lane.ID, chunk.ID, extractor.Key(), err) } - candidates, err := normalizeCandidates(extractor, result.Candidates, &nextCandidateIndex) + candidates, err := normalizeCandidates(extractor, result.Candidates, nextCandidateIndex) if err != nil { - return output, err - } - - approved, rejected, warnings, err := runValidators(ctx, extractor, input.Source, candidates, input.Metadata) - output.Warnings = append(output.Warnings, warnings...) - output.Rejected = append(output.Rejected, rejected...) - if err != nil { - return output, err - } - - for _, candidate := range approved { - output.Approved = append(output.Approved, artifacts.ArtifactFromCandidate(candidate)) + return err } + chunkArtifacts = append(chunkArtifacts, contracts.ChunkArtifacts{ + Chunk: chunk, + Candidates: candidates, + }) } - return output, nil + mergeResult, err := merger.Merge(ctx, contracts.MergeRequest{ + Source: doc, + LaneID: lane.ID, + ChunkArtifacts: chunkArtifacts, + Metadata: input.Metadata, + }) + output.Warnings = append(output.Warnings, mergeResult.Warnings...) + if err != nil { + return fmt.Errorf("merge lane %q with merger %q: %w", lane.ID, merger.Key(), err) + } + + normalizeResult, err := normalizer.Normalize(ctx, contracts.NormalizeRequest{ + Source: doc, + LaneID: lane.ID, + Candidates: mergeResult.Candidates, + Metadata: input.Metadata, + }) + output.Warnings = append(output.Warnings, normalizeResult.Warnings...) + if err != nil { + return fmt.Errorf("normalize lane %q with normalizer %q: %w", lane.ID, normalizer.Key(), err) + } + + approved, rejected, warnings, err := runValidators(ctx, extractor.Key(), validators, doc, normalizeResult.Candidates, input.Metadata) + output.Warnings = append(output.Warnings, warnings...) + output.Rejected = append(output.Rejected, rejected...) + if err != nil { + return err + } + + for _, candidate := range approved { + output.Approved = append(output.Approved, artifacts.ArtifactFromCandidate(candidate)) + } + return nil +} + +func (r *Runner) buildConfiguredValidators(lane ResolvedArtifactLane) ([]contracts.Validator, error) { + validators := make([]contracts.Validator, 0, len(lane.Validators)) + for _, binding := range lane.Validators { + validator, err := r.registries.Validators.Build(binding.Module) + if err != nil { + return nil, fmt.Errorf("build validator %q for lane %q: %w", binding.Module, lane.ID, err) + } + validators = append(validators, validator) + } + return validators, nil +} + +func (r *Runner) validateRegistries(pipeline ResolvedPipeline) error { + if r.registries.Inputs == nil { + return fmt.Errorf("input registry must not be nil") + } + if r.registries.Chunkers == nil { + return fmt.Errorf("chunker registry must not be nil") + } + if r.registries.Extractors == nil { + return fmt.Errorf("extractor registry must not be nil") + } + if r.registries.Mergers == nil { + return fmt.Errorf("merger registry must not be nil") + } + if r.registries.Normalizers == nil { + return fmt.Errorf("normalizer registry must not be nil") + } + if r.registries.Outputs == nil { + return fmt.Errorf("output encoder registry must not be nil") + } + if pipelineUsesConfiguredValidators(pipeline) && r.registries.Validators == nil { + return fmt.Errorf("validator registry must not be nil") + } + return nil +} + +func validateRunInput(input RunInput) error { + if input.Pipeline.ID == "" { + return fmt.Errorf("resolved pipeline id must not be empty") + } + if input.Pipeline.Digest == "" { + return fmt.Errorf("resolved pipeline digest must not be empty") + } + if input.Pipeline.Input.Module == "" { + return fmt.Errorf("resolved pipeline input module must not be empty") + } + if input.Pipeline.Chunk.Module == "" { + return fmt.Errorf("resolved pipeline chunk module must not be empty") + } + if input.Pipeline.Output.Module == "" { + return fmt.Errorf("resolved pipeline output module must not be empty") + } + if len(input.Pipeline.ArtifactLanes) == 0 { + return fmt.Errorf("resolved pipeline artifact lanes must not be empty") + } + for _, lane := range input.Pipeline.ArtifactLanes { + if lane.ID == "" { + return fmt.Errorf("resolved pipeline artifact lane id must not be empty") + } + if lane.Extract.Module == "" { + return fmt.Errorf("resolved pipeline lane %q extract module must not be empty", lane.ID) + } + if lane.Merge.Module == "" { + return fmt.Errorf("resolved pipeline lane %q merge module must not be empty", lane.ID) + } + if lane.Normalize.Module == "" { + return fmt.Errorf("resolved pipeline lane %q normalize module must not be empty", lane.ID) + } + for _, validator := range lane.Validators { + if validator.Module == "" { + return fmt.Errorf("resolved pipeline lane %q validator module must not be empty", lane.ID) + } + } + } + return nil +} + +func manifestFromPipeline(pipeline ResolvedPipeline) artifacts.RunManifest { + manifest := artifacts.RunManifest{ + PipelineID: pipeline.ID, + PipelineDigest: pipeline.Digest, + InputModule: pipeline.Input.Module, + Chunker: pipeline.Chunk.Module, + OutputEncoder: pipeline.Output.Module, + ArtifactLanes: make([]artifacts.ArtifactLaneManifest, 0, len(pipeline.ArtifactLanes)), + } + + for _, lane := range pipeline.ArtifactLanes { + laneManifest := artifacts.ArtifactLaneManifest{ + ID: lane.ID, + Extractor: lane.Extract.Module, + Merger: lane.Merge.Module, + Normalizer: lane.Normalize.Module, + } + for _, validator := range lane.Validators { + laneManifest.Validators = append(laneManifest.Validators, validator.Module) + } + manifest.ArtifactLanes = append(manifest.ArtifactLanes, laneManifest) + } + return manifest +} + +func failOutput(output RunOutput) RunOutput { + if output.Manifest.PipelineID != "" { + output.Manifest.ValidationStatus = "failed" + } + return output +} + +func pipelineUsesConfiguredValidators(pipeline ResolvedPipeline) bool { + for _, lane := range pipeline.ArtifactLanes { + if len(lane.Validators) > 0 { + return true + } + } + return false } func normalizeCandidates(extractor contracts.Extractor, candidates []artifacts.ArtifactCandidate, nextIndex *int) ([]artifacts.ArtifactCandidate, error) { @@ -119,14 +359,14 @@ func normalizeCandidates(extractor contracts.Extractor, candidates []artifacts.A return normalized, nil } -func runValidators(ctx context.Context, extractor contracts.Extractor, doc *source.SourceDocument, candidates []artifacts.ArtifactCandidate, metadata map[string]any) ([]artifacts.ArtifactCandidate, []artifacts.RejectedArtifact, []contracts.Warning, error) { +func runValidators(ctx context.Context, extractorKey string, validators []contracts.Validator, doc *source.SourceDocument, candidates []artifacts.ArtifactCandidate, metadata map[string]any) ([]artifacts.ArtifactCandidate, []artifacts.RejectedArtifact, []contracts.Warning, error) { eligible := candidates var rejected []artifacts.RejectedArtifact var warnings []contracts.Warning - for validatorIndex, validator := range extractor.Validators() { + for validatorIndex, validator := range validators { if validator == nil { - return nil, rejected, warnings, fmt.Errorf("extractor %q validator[%d] must not be nil", extractor.Key(), validatorIndex) + return nil, rejected, warnings, fmt.Errorf("extractor %q validator[%d] must not be nil", extractorKey, validatorIndex) } result, err := validator.Validate(ctx, contracts.ValidationRequest{ Source: doc, @@ -135,13 +375,13 @@ func runValidators(ctx context.Context, extractor contracts.Extractor, doc *sour }) warnings = append(warnings, result.Warnings...) if err != nil { - return nil, rejected, warnings, fmt.Errorf("validate extractor %q with validator %q: %w", extractor.Key(), validator.Name(), err) + return nil, rejected, warnings, fmt.Errorf("validate extractor %q with validator %q: %w", extractorKey, validator.Name(), err) } if result.ValidatorName != validator.Name() { return nil, rejected, warnings, fmt.Errorf("validator %q returned result for %q", validator.Name(), result.ValidatorName) } if err := validate.EnforceDecisionCardinality(eligible, result.Decisions); err != nil { - return nil, rejected, warnings, fmt.Errorf("validate extractor %q with validator %q: %w", extractor.Key(), validator.Name(), err) + return nil, rejected, warnings, fmt.Errorf("validate extractor %q with validator %q: %w", extractorKey, validator.Name(), err) } decisions := make(map[int]contracts.ValidationDecision, len(result.Decisions)) diff --git a/internal/framework/pipeline/runner_test.go b/internal/framework/pipeline/runner_test.go index da938b9..69d67ac 100644 --- a/internal/framework/pipeline/runner_test.go +++ b/internal/framework/pipeline/runner_test.go @@ -14,26 +14,32 @@ import ( ) func TestNewAndDataTypes(t *testing.T) { - r := New(fakeFactory{}) - if r == nil { + runner := New(Registries{}) + if runner == nil { t.Fatal("New() = nil, want runner") } input := RunInput{ - Source: validSourceDocument(), - ExtractorKeys: []string{"generic-extractor"}, - Metadata: map[string]any{"request": "test"}, + Pipeline: resolvedPipeline(), + SourceID: "source-1", + Path: "input.txt", + RawInput: []byte("source text"), + LLMClient: fakeLLMClient{}, + Metadata: map[string]any{"request": "test"}, } output := RunOutput{ - Approved: []artifacts.Artifact{{ExtractorKey: "generic-extractor"}}, - Rejected: []artifacts.RejectedArtifact{{ValidatorName: "generic-validator"}}, - Warnings: []contracts.Warning{{ReasonCode: "note", Message: "message"}}, + Manifest: artifacts.RunManifest{PipelineID: "pipeline-1"}, + Approved: []artifacts.Artifact{{ExtractorKey: "extract-alpha"}}, + Rejected: []artifacts.RejectedArtifact{{ValidatorName: "validator"}}, + Warnings: []contracts.Warning{{ReasonCode: "note", Message: "message"}}, + EncodedOutput: []byte(`{}`), + ContentType: "application/json", } - if input.Source.ID != "source-1" { - t.Fatalf("RunInput.Source.ID = %q, want source-1", input.Source.ID) + if input.Pipeline.ID != "pipeline-1" || input.SourceID != "source-1" { + t.Fatalf("RunInput = %#v, want constructed fields", input) } - if len(output.Approved) != 1 || len(output.Rejected) != 1 || len(output.Warnings) != 1 { + if output.Manifest.PipelineID != "pipeline-1" || len(output.Approved) != 1 || len(output.Rejected) != 1 || len(output.Warnings) != 1 { t.Fatalf("RunOutput = %#v, want constructed fields", output) } } @@ -50,345 +56,765 @@ func TestRunRejectsInvalidSetup(t *testing.T) { error: "runner must not be nil", }, { - name: "nil factory", - run: func() (RunOutput, error) { return New(nil).Run(context.Background(), RunInput{}) }, - error: "factory", + 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", - run: func() (RunOutput, error) { - return New(fakeFactory{}).Run(context.Background(), RunInput{Source: &source.SourceDocument{}, ExtractorKeys: []string{"generic-extractor"}}) + configure: func(modules *runnerModules) { + modules.input.doc = &source.SourceDocument{} }, - error: "validate source document", - }, - { - name: "empty extractors", - run: func() (RunOutput, error) { - return New(fakeFactory{}).Run(context.Background(), RunInput{Source: validSourceDocument()}) - }, - error: "extractor keys must not be empty", + want: "validate source document", }, } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - _, err := tt.run() - if err == nil { - t.Fatal("Run() error = nil, want error") - } - if !strings.Contains(err.Error(), tt.error) { - t.Fatalf("Run() error = %q, want substring %q", err.Error(), tt.error) + 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 TestRunRejectsNilExtractorFromFactory(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": nil, - }} +func TestRunRejectsChunkerBuildChunkAndEmptyChunkErrors(t *testing.T) { + chunkErr := errors.New("chunk failed") - _, err := New(factory).Run(context.Background(), RunInput{ - Source: validSourceDocument(), - ExtractorKeys: []string{"generic-extractor"}, - }) - - assertRunError(t, err, "returned nil extractor") -} - -func TestRunUsesConfiguredExtractorOrderAndAssignsGlobalIndices(t *testing.T) { - var order []string - var seenIndices []int - recordIndices := func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { - decisions := make([]contracts.ValidationDecision, 0, len(candidates)) - for _, candidate := range candidates { - seenIndices = append(seenIndices, candidate.Index) - decisions = append(decisions, validate.Approved(candidate.Index)) - } - return decisions - } - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "second": fakeExtractor{key: "second", artifactType: "artifact", schemaVersion: "v1", candidateCount: 1, validators: []contracts.Validator{fakeValidator{name: "recorder-second", decisions: recordIndices}}, order: &order}, - "first": fakeExtractor{key: "first", artifactType: "artifact", schemaVersion: "v1", candidateCount: 2, validators: []contracts.Validator{fakeValidator{name: "recorder-first", decisions: recordIndices}}, order: &order}, - }} - - output, err := New(factory).Run(context.Background(), RunInput{ - Source: validSourceDocument(), - ExtractorKeys: []string{"second", "first"}, - }) - if err != nil { - t.Fatalf("Run() error = %v, want nil", err) - } - - if !reflect.DeepEqual(order, []string{"second", "first"}) { - t.Fatalf("order = %#v, want configured order", order) - } - if !reflect.DeepEqual(seenIndices, []int{0, 1, 2}) { - t.Fatalf("seen indices = %#v, want [0 1 2]", seenIndices) - } - if len(output.Approved) != 3 { - t.Fatalf("len(Approved) = %d, want 3", len(output.Approved)) - } -} - -func TestRunFillsEmptyCandidateExtractorMetadata(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{key: "generic-extractor", artifactType: "generic-artifact", schemaVersion: "v1", candidates: []artifacts.ArtifactCandidate{{Payload: []byte(`{"value":true}`)}}}, - }} - - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) - if err != nil { - t.Fatalf("Run() error = %v, want nil", err) - } - - artifact := output.Approved[0] - if artifact.ExtractorKey != "generic-extractor" || artifact.ArtifactType != "generic-artifact" || artifact.SchemaVersion != "v1" { - t.Fatalf("approved artifact metadata = %#v, want extractor metadata", artifact) - } -} - -func TestRunRejectsCandidateMetadataMismatches(t *testing.T) { tests := []struct { name string - candidate artifacts.ArtifactCandidate - error string + configure func(*runnerModules) + want string }{ - {name: "extractor key", candidate: artifacts.ArtifactCandidate{ExtractorKey: "other"}, error: "extractor_key"}, - {name: "artifact type", candidate: artifacts.ArtifactCandidate{ArtifactType: "other"}, error: "artifact_type"}, - {name: "schema version", candidate: artifacts.ArtifactCandidate{SchemaVersion: "other"}, error: "schema_version"}, + { + 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 _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{key: "generic-extractor", artifactType: "generic-artifact", schemaVersion: "v1", candidates: []artifacts.ArtifactCandidate{tt.candidate}}, - }} + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + modules := defaultRunnerModules() + test.configure(modules) - _, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) + output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()}) - if err == nil { - t.Fatal("Run() error = nil, want error") - } - if !strings.Contains(err.Error(), tt.error) { - t.Fatalf("Run() error = %q, want substring %q", err.Error(), tt.error) + assertRunError(t, err, test.want) + if output.Manifest.ValidationStatus != "failed" { + t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus) } }) } } -func TestRunApprovesCandidatesWithoutValidators(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{key: "generic-extractor", artifactType: "generic-artifact", schemaVersion: "v1", candidateCount: 2}, - }} +func TestRunExecutesChunksAndPassesChunkAndLLMClient(t *testing.T) { + modules := defaultRunnerModules() + llmClient := fakeLLMClient{} - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) + 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(extractor.seenLLMClients) != 2 || extractor.seenLLMClients[0] == nil || extractor.seenLLMClients[1] == nil { + t.Fatalf("seen LLM clients = %#v, want client for each chunk", extractor.seenLLMClients) + } + 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)) } - if len(output.Rejected) != 0 { - t.Fatalf("len(Rejected) = %d, want 0", len(output.Rejected)) - } } -func TestRunValidatorApprovalProducesApprovedArtifacts(t *testing.T) { - validator := fakeValidator{name: "generic-validator", decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { - return []contracts.ValidationDecision{validate.Approved(candidates[0].Index)} - }} - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{key: "generic-extractor", artifactType: "generic-artifact", schemaVersion: "v1", candidateCount: 1, validators: []contracts.Validator{validator}}, - }} +func TestRunPassesInputRequestFields(t *testing.T) { + modules := defaultRunnerModules() + metadata := map[string]any{"request": "test"} - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) + _, 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(output.Approved) != 1 { - t.Fatalf("len(Approved) = %d, want 1", len(output.Approved)) + + 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 TestRunValidatorRejectionRemovesCandidateFromLaterValidators(t *testing.T) { - var laterSeen int - rejectFirst := fakeValidator{name: "reject-first", decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { - return []contracts.ValidationDecision{ - validate.Rejected(candidates[0].Index, "invalid", "not accepted"), - validate.Approved(candidates[1].Index), - } - }} - approveRemaining := fakeValidator{name: "approve-remaining", decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { - laterSeen = len(candidates) - return []contracts.ValidationDecision{validate.Approved(candidates[0].Index)} - }} - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{key: "generic-extractor", artifactType: "generic-artifact", schemaVersion: "v1", candidateCount: 2, validators: []contracts.Validator{rejectFirst, approveRemaining}}, - }} +func TestRunPassesPerChunkCandidatesToMergeAndNormalize(t *testing.T) { + modules := defaultRunnerModules() - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) + _, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()}) if err != nil { t.Fatalf("Run() error = %v, want nil", err) } - if laterSeen != 1 { - t.Fatalf("later validator saw %d candidates, want 1", laterSeen) + + merger := modules.mergers["merge"] + if len(merger.requests) != 1 { + t.Fatalf("len(merge requests) = %d, want 1", len(merger.requests)) } - if len(output.Rejected) != 1 { - t.Fatalf("len(Rejected) = %d, want 1", len(output.Rejected)) + chunkArtifacts := merger.requests[0].ChunkArtifacts + if len(chunkArtifacts) != 2 { + t.Fatalf("len(ChunkArtifacts) = %d, want 2", len(chunkArtifacts)) } - if output.Rejected[0].ValidatorName != "reject-first" || output.Rejected[0].ReasonCode != "invalid" { + 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 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]) } - if len(output.Approved) != 1 { - t.Fatalf("len(Approved) = %d, want 1", len(output.Approved)) - } } -func TestRunSurfacesValidatorNameMismatch(t *testing.T) { - validator := fakeValidator{name: "generic-validator", resultName: "other-validator", decisions: approveAll} - factory := factoryWithValidator(validator) +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} - _, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) - - assertRunError(t, err, "returned result") -} - -func TestRunSurfacesValidatorCardinalityError(t *testing.T) { - validator := fakeValidator{name: "generic-validator", decisions: func(candidates []artifacts.ArtifactCandidate) []contracts.ValidationDecision { - return nil - }} - factory := factoryWithValidator(validator) - - _, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) - - assertRunError(t, err, "0 decisions for 1 candidates") -} - -func TestRunRejectsNilValidator(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{ - key: "generic-extractor", - artifactType: "generic-artifact", - schemaVersion: "v1", - candidateCount: 1, - validators: []contracts.Validator{nil}, - }, - }} - - _, err := New(factory).Run(context.Background(), RunInput{ - Source: validSourceDocument(), - ExtractorKeys: []string{"generic-extractor"}, + output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{ + Pipeline: resolvedPipelineWithValidators("configured", "second-validator"), }) - - assertRunError(t, err, "validator[0] must not be nil") -} - -func TestRunCollectsExtractorAndValidatorWarnings(t *testing.T) { - validator := fakeValidator{ - name: "generic-validator", - decisions: approveAll, - warnings: []contracts.Warning{{ReasonCode: "validator-warning", Message: "validator warning"}}, - } - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{ - key: "generic-extractor", - artifactType: "generic-artifact", - schemaVersion: "v1", - candidateCount: 1, - validators: []contracts.Validator{validator}, - warnings: []contracts.Warning{{ReasonCode: "extractor-warning", Message: "extractor warning"}}, - }, - }} - - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"generic-extractor"}}) if err != nil { t.Fatalf("Run() error = %v, want nil", err) } - if got := warningReasons(output.Warnings); !reflect.DeepEqual(got, []string{"extractor-warning", "validator-warning"}) { - t.Fatalf("warning reasons = %#v, want extractor and validator warnings", got) + + 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 TestRunReturnsPartialOutputWhenLaterExtractorFails(t *testing.T) { - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "ok": fakeExtractor{key: "ok", artifactType: "artifact", schemaVersion: "v1", candidateCount: 1}, - "fail": fakeExtractor{key: "fail", artifactType: "artifact", schemaVersion: "v1", err: errors.New("extract failed")}, +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(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"ok", "fail"}}) + output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()}) + if err != nil { + t.Fatalf("Run() error = %v, want nil", err) + } - assertRunError(t, err, "extract with extractor") - if len(output.Approved) != 1 { + 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 output.ContentType != "application/json" { + t.Fatalf("ContentType = %q, want application/json", output.ContentType) + } + if string(output.EncodedOutput) != `{"encoded":true}` { + t.Fatalf("EncodedOutput = %s, want encoded payload", output.EncodedOutput) + } + 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 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 len(output.Approved) != 2 { t.Fatalf("len(Approved) = %d, want partial approved output", len(output.Approved)) } } -func TestRunReturnsPartialOutputWhenLaterValidatorFails(t *testing.T) { - validatorErr := errors.New("validator failed") - factory := fakeFactory{extractors: map[string]contracts.Extractor{ - "ok": fakeExtractor{key: "ok", artifactType: "artifact", schemaVersion: "v1", candidateCount: 1}, - "fail": fakeExtractor{key: "fail", artifactType: "artifact", schemaVersion: "v1", candidateCount: 1, validators: []contracts.Validator{fakeValidator{name: "failing-validator", err: validatorErr}}}, - }} +func TestRunManifestIncludesPipelineAndLaneDetails(t *testing.T) { + output, err := New(newRunnerRegistries(t, nil)).Run(context.Background(), RunInput{Pipeline: resolvedPipelineWithValidators("configured")}) + if err != nil { + t.Fatalf("Run() error = %v, want nil", err) + } - output, err := New(factory).Run(context.Background(), RunInput{Source: validSourceDocument(), ExtractorKeys: []string{"ok", "fail"}}) - - assertRunError(t, err, "failing-validator") - if len(output.Approved) != 1 { - t.Fatalf("len(Approved) = %d, want partial approved output", len(output.Approved)) + 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 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) } } -type fakeFactory struct { - extractors map[string]contracts.Extractor - err error +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 (factory fakeFactory) Build(key string) (contracts.Extractor, error) { - if factory.err != nil { - return nil, factory.err +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"}, } - extractor, ok := factory.extractors[key] - if !ok { - return nil, errors.New("missing extractor") + + 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) + } + }) } - return extractor, nil } -type fakeExtractor struct { +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"), + ArtifactLanes: []ResolvedArtifactLane{ + { + ID: "alpha", + Extract: Binding("extract-alpha"), + Merge: Binding("merge"), + Normalize: Binding("normalize"), + }, + }, + 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 +} + +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", bytes: []byte(`{"encoded":true}`), contentType: "application/json"}, + } +} + +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 + 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 +} + +type runnerChunker struct { + key string + chunks []contracts.SourceChunk + warnings []contracts.Warning + err error + requests []contracts.ChunkRequest +} + +func (chunker *runnerChunker) Key() string { + return chunker.key +} + +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 +} + +type runnerExtractor struct { key string artifactType string schemaVersion string - candidateCount int candidates []artifacts.ArtifactCandidate validators []contracts.Validator warnings []contracts.Warning err error - order *[]string + seenChunkIDs []string + seenLLMClients []contracts.StructuredLLMClient + seenMetadata []map[string]any } -func (extractor fakeExtractor) Key() string { +func (extractor *runnerExtractor) Key() string { return extractor.key } -func (extractor fakeExtractor) ArtifactType() string { +func (extractor *runnerExtractor) ArtifactType() string { return extractor.artifactType } -func (extractor fakeExtractor) SchemaVersion() string { +func (extractor *runnerExtractor) SchemaVersion() string { return extractor.schemaVersion } -func (extractor fakeExtractor) Validators() []contracts.Validator { +func (extractor *runnerExtractor) Validators() []contracts.Validator { return extractor.validators } -func (extractor fakeExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) { - if extractor.order != nil { - *extractor.order = append(*extractor.order, extractor.key) +func (extractor *runnerExtractor) Extract(ctx context.Context, req contracts.ExtractionRequest) (contracts.ExtractionResult, error) { + 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...) - for len(candidates) < extractor.candidateCount { - candidates = append(candidates, artifacts.ArtifactCandidate{Payload: []byte(`{"value":true}`)}) + if len(candidates) == 0 { + candidates = []artifacts.ArtifactCandidate{{Payload: []byte(`{"value":true}`)}} } return contracts.ExtractionResult{ Candidates: candidates, @@ -396,19 +822,75 @@ func (extractor fakeExtractor) Extract(ctx context.Context, req contracts.Extrac }, extractor.err } -type fakeValidator struct { +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) 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 } -func (validator fakeValidator) Name() string { +func (validator *runnerValidator) Name() string { return validator.name } -func (validator fakeValidator) Validate(ctx context.Context, req contracts.ValidationRequest) (contracts.ValidationResult, error) { +func (validator *runnerValidator) Validate(ctx context.Context, req contracts.ValidationRequest) (contracts.ValidationResult, error) { + validator.calls++ + if validator.order != nil { + *validator.order = append(*validator.order, validator.name) + } resultName := validator.resultName if resultName == "" { resultName = validator.name @@ -424,16 +906,32 @@ func (validator fakeValidator) Validate(ctx context.Context, req contracts.Valid }, validator.err } -func factoryWithValidator(validator contracts.Validator) fakeFactory { - return fakeFactory{extractors: map[string]contracts.Extractor{ - "generic-extractor": fakeExtractor{ - key: "generic-extractor", - artifactType: "generic-artifact", - schemaVersion: "v1", - candidateCount: 1, - validators: []contracts.Validator{validator}, - }, - }} +type runnerOutputEncoder struct { + key string + bytes []byte + contentType string + warnings []contracts.Warning + err error + 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{ + Bytes: encoder.bytes, + ContentType: encoder.contentType, + Warnings: encoder.warnings, + }, encoder.err +} + +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 { @@ -449,7 +947,18 @@ func validSourceDocument() *source.SourceDocument { ID: "source-1", Kind: "document", Format: "text/plain", - Digest: "sha256:abc123", + Digest: "sha256:source", + Units: []source.SourceUnit{ + {ID: "u1", Kind: "unit", Text: "Source unit."}, + }, + } +} + +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."}, }, @@ -464,6 +973,14 @@ func warningReasons(warnings []contracts.Warning) []string { 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 assertRunError(t *testing.T, err error, want string) { t.Helper()