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

241 lines
8.6 KiB
Go

package pipeline
import (
"context"
"errors"
"reflect"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
)
func TestResolvePipelineValidatesModuleAndValidatorOptions(t *testing.T) {
tests := []struct {
name string
mutate func(*PipelineProfile)
want []string
}{
{
name: "module",
mutate: func(profile *PipelineProfile) {
profile.Input.Options = map[string]any{"surprise": true}
},
want: []string{`pipeline "construction" input module "input" options`, `unknown option "surprise"`},
},
{
name: "validator",
mutate: func(profile *PipelineProfile) {
profile.Chunk.Validators = ValidatorOverride{Set: true, Validators: []ModuleBinding{{Module: "configured", Options: map[string]any{"surprise": true}}}}
},
want: []string{`pipeline "construction" chunk module "chunk" validator "configured" options`, `unknown option "surprise"`},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
registries, _ := constructionRegistries(t, nil, nil)
profile := constructionProfile()
test.mutate(&profile)
_, err := ResolvePipeline(profile, ResolveOptions{}, registries.catalog())
if err == nil {
t.Fatal("ResolvePipeline() error = nil, want option validation error")
}
for _, want := range test.want {
if !strings.Contains(err.Error(), want) {
t.Fatalf("ResolvePipeline() error = %q, want substring %q", err, want)
}
}
})
}
}
func TestPrepareConstructsEverythingInStableOrder(t *testing.T) {
var built []string
registries, _ := constructionRegistries(t, &built, nil)
resolved, err := ResolvePipeline(constructionProfile(), ResolveOptions{}, registries.catalog())
if err != nil {
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
}
prepared, err := Prepare(resolved, registries, ModuleDependencies{})
if err != nil {
t.Fatalf("Prepare() error = %v, want nil", err)
}
want := []string{"input", "chunk", "validator", "extract", "validator", "merge", "validator", "normalize", "validator", "output"}
if !reflect.DeepEqual(built, want) {
t.Fatalf("construction order = %#v, want %#v", built, want)
}
if prepared.Input.Module != "input" || prepared.Chunk.Module != "chunk" || prepared.Output.Module != "output" || len(prepared.ArtifactLanes) != 1 {
t.Fatalf("PreparedPipeline = %#v, want explicit resolved components", prepared)
}
}
func TestPrepareFailuresOccurBeforeInputParse(t *testing.T) {
tests := []struct {
name string
deps ModuleDependencies
configure func(*constructionFailure)
want string
wantBuilt []string
}{
{
name: "missing required llm dependency",
configure: func(failure *constructionFailure) {
failure.requireExtractorLLM = true
},
want: `lane "artifact" extract module "extract"`,
wantBuilt: []string{"input", "chunk", "validator", "extract"},
},
{
name: "late output construction",
configure: func(failure *constructionFailure) {
failure.output = errors.New("output unavailable")
},
want: `output module "output"`,
wantBuilt: []string{"input", "chunk", "validator", "extract", "validator", "merge", "validator", "normalize", "validator", "output"},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
failure := &constructionFailure{}
test.configure(failure)
var built []string
registries, input := constructionRegistries(t, &built, failure)
resolved, err := ResolvePipeline(constructionProfile(), ResolveOptions{}, registries.catalog())
if err != nil {
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
}
_, err = Prepare(resolved, registries, test.deps)
if err == nil || !strings.Contains(err.Error(), test.want) {
t.Fatalf("Prepare() error = %v, want substring %q", err, test.want)
}
if len(input.requests) != 0 {
t.Fatalf("input Parse calls = %d, want zero", len(input.requests))
}
if !reflect.DeepEqual(built, test.wantBuilt) {
t.Fatalf("construction order = %#v, want %#v", built, test.wantBuilt)
}
})
}
}
type constructionFailure struct {
requireExtractorLLM bool
output error
}
func constructionProfile() PipelineProfile {
validators := ValidatorOverride{Set: true, Validators: []ModuleBinding{{Module: "configured"}}}
return PipelineProfile{
ID: "construction",
Input: Binding("input"),
Chunk: ModuleBinding{Module: "chunk", Validators: validators},
Artifacts: map[string]ArtifactLaneProfile{
"artifact": {
Extract: ModuleBinding{Module: "extract", Validators: validators},
Merge: ModuleBinding{Module: "merge", Validators: validators},
Normalize: ModuleBinding{Module: "normalize", Validators: validators},
},
},
Output: Binding("output"),
}
}
func constructionRegistries(t *testing.T, built *[]string, failure *constructionFailure) (Registries, *constructionInput) {
t.Helper()
if built == nil {
built = &[]string{}
}
if failure == nil {
failure = &constructionFailure{}
}
record := func(name string) { *built = append(*built, name) }
strict := func(options map[string]any) error { return RejectUnknownOptions(options, "known") }
input := &constructionInput{key: "input"}
registries := Registries{
Inputs: NewInputAdapterRegistry(), Chunkers: NewChunkerRegistry(), ArtifactCodecs: NewArtifactCodecRegistry(),
Extractors: NewExtractorRegistry(), Mergers: NewMergerRegistry(), Normalizers: NewNormalizerRegistry(),
Validators: NewValidatorRegistry(), ValidatorChains: NewValidatorChainRegistry(), Outputs: NewOutputEncoderRegistry(),
}
if err := RegisterArtifactCodec(registries.ArtifactCodecs, notesCodec()); err != nil {
t.Fatal(err)
}
if err := registries.Inputs.RegisterBuilderWithSpec(defaultModuleSpec("input", StageInput), strict, func(BuildRequest) (contracts.InputAdapter, error) {
record("input")
return input, nil
}); err != nil {
t.Fatal(err)
}
if err := registries.Chunkers.RegisterBuilderWithSpec(defaultModuleSpec("chunk", StageChunk), strict, func(BuildRequest) (contracts.Chunker, error) {
record("chunk")
return &typedTestChunker{key: "chunk"}, nil
}); err != nil {
t.Fatal(err)
}
extractSpec := defaultModuleSpec("extract", StageExtract)
extractSpec.ArtifactKind = "test/notes"
if err := RegisterExtractorBuilder(registries.Extractors, extractSpec, strict, func(request BuildRequest) (contracts.Extractor[codecNotes], error) {
record("extract")
if failure.requireExtractorLLM && request.Dependencies.LLM == nil {
return nil, errors.New("structured LLM client is required")
}
return typedTestExtractor[codecNotes]{key: "extract"}, nil
}); err != nil {
t.Fatal(err)
}
mergeSpec := defaultModuleSpec("merge", StageMerge)
mergeSpec.ArtifactKind = "test/notes"
if err := RegisterMergerBuilder(registries.Mergers, mergeSpec, strict, func(BuildRequest) (contracts.Merger[codecNotes], error) {
record("merge")
return typedTestMerger[codecNotes]{key: "merge"}, nil
}); err != nil {
t.Fatal(err)
}
normalizeSpec := defaultModuleSpec("normalize", StageNormalize)
normalizeSpec.ArtifactKind = "test/notes"
if err := RegisterNormalizerBuilder(registries.Normalizers, normalizeSpec, strict, func(BuildRequest) (contracts.Normalizer[codecNotes], error) {
record("normalize")
return typedTestNormalizer[codecNotes]{key: "normalize"}, nil
}); err != nil {
t.Fatal(err)
}
validatorSpec := ValidatorSpec{Key: "configured", ExecutionClass: contracts.ExecutionClassDeterministic}
if err := RegisterChunkValidatorBuilder(registries.Validators, validatorSpec, strict, func(BuildRequest) (contracts.ChunkValidator, error) {
record("validator")
return typedTestChunkValidator{key: "configured"}, nil
}); err != nil {
t.Fatal(err)
}
if err := RegisterTypedValidatorBuilder(registries.Validators, "test/notes", validatorSpec, strict, func(BuildRequest) (contracts.TypedValidator[codecNotes], error) {
record("validator")
return typedTestValidator[codecNotes]{key: "configured"}, nil
}); err != nil {
t.Fatal(err)
}
if err := registries.Outputs.RegisterBuilderWithSpec(defaultModuleSpec("output", StageOutput), strict, func(BuildRequest) (contracts.OutputEncoder, error) {
record("output")
if failure.output != nil {
return nil, failure.output
}
return &typedTestOutput{key: "output"}, nil
}); err != nil {
t.Fatal(err)
}
return registries, input
}
type constructionInput struct {
key string
requests []contracts.ParseRequest
}
func (input *constructionInput) Key() string { return input.key }
func (input *constructionInput) Parse(_ context.Context, request contracts.ParseRequest) (*source.SourceDocument, error) {
input.requests = append(input.requests, request)
return typedTestDocument(), nil
}