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

351 lines
13 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 TestPrepareDeliversTargetReferencesAsIndependentBuildInputs(t *testing.T) {
var built []string
var observations []constructionBuildObservation
registries, input := constructionRegistriesWithHooks(t, &built, nil,
func(name string, request BuildRequest) {
observations = append(observations, constructionBuildObservation{Name: name, Request: request})
},
func(name string, request *BuildRequest) {
if name != "extract" {
return
}
slot := request.References.Slots["extract"]
slot.Items[0].Content = []byte("mutated by extractor builder")
request.References.Slots["extract"] = slot
},
)
resolved, err := ResolvePipeline(constructionProfile(), ResolveOptions{}, registries.catalog())
if err != nil {
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
}
resolved.ChunkReferences.ReferenceSet = constructionReferenceSet("chunk", "chunk reference")
resolved.ArtifactLanes[0].ExtractReferences.ReferenceSet = constructionReferenceSet("extract", "extract reference")
resolved.ArtifactLanes[0].MergeReferences.ReferenceSet = constructionReferenceSet("merge", "merge reference")
resolved.ArtifactLanes[0].NormalizeReferences.ReferenceSet = constructionReferenceSet("normalize", "normalize reference")
prepared, err := Prepare(resolved, registries, ModuleDependencies{})
if err != nil {
t.Fatalf("Prepare() error = %v, want nil", err)
}
wantNames := []string{"input", "chunk", "validator", "extract", "validator", "merge", "validator", "normalize", "validator", "output"}
if !reflect.DeepEqual(built, wantNames) {
t.Fatalf("construction order = %#v, want %#v", built, wantNames)
}
wantContents := []string{"", "chunk reference", "chunk reference", "extract reference", "extract reference", "merge reference", "merge reference", "normalize reference", "normalize reference", ""}
if len(observations) != len(wantContents) {
t.Fatalf("observed %d build requests, want %d", len(observations), len(wantContents))
}
for i, want := range wantContents {
if got := constructionReferenceContent(observations[i].Request.References); got != want {
t.Errorf("build request %d (%s) reference content = %q, want %q", i, observations[i].Name, got, want)
}
}
if got := constructionReferenceContent(resolved.ArtifactLanes[0].ExtractReferences.ReferenceSet); got != "extract reference" {
t.Fatalf("resolved extract references = %q, want original content", got)
}
_, err = prepared.lanes[0].typed.extract(context.Background(), prepared.lanes[0].typed.extractor, contracts.TypedExtractionRequest{
References: CloneReferenceSet(resolved.ArtifactLanes[0].ExtractReferences.ReferenceSet),
})
if err != nil {
t.Fatalf("prepared extractor operation error = %v, want nil", err)
}
if len(input.extractRequests) != 1 {
t.Fatalf("runtime extraction requests = %d, want one", len(input.extractRequests))
}
if got := constructionReferenceContent(input.extractRequests[0].References); got != "extract reference" {
t.Fatalf("runtime extraction references = %q, want original content", got)
}
}
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) {
return constructionRegistriesWithHooks(t, built, failure, nil, nil)
}
type constructionBuildObservation struct {
Name string
Request BuildRequest
}
func constructionRegistriesWithHooks(t *testing.T, built *[]string, failure *constructionFailure, observe func(string, BuildRequest), mutate func(string, *BuildRequest)) (Registries, *constructionInput) {
t.Helper()
if built == nil {
built = &[]string{}
}
if failure == nil {
failure = &constructionFailure{}
}
record := func(name string, request *BuildRequest) {
*built = append(*built, name)
if observe != nil {
observe(name, cloneBuildRequest(*request))
}
if mutate != nil {
mutate(name, request)
}
}
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(request BuildRequest) (contracts.InputAdapter, error) {
record("input", &request)
return input, nil
}); err != nil {
t.Fatal(err)
}
if err := registries.Chunkers.RegisterBuilderWithSpec(defaultModuleSpec("chunk", StageChunk), strict, func(request BuildRequest) (contracts.Chunker, error) {
record("chunk", &request)
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", &request)
if failure.requireExtractorLLM && request.Dependencies.LLM == nil {
return nil, errors.New("structured LLM client is required")
}
return &constructionExtractor{key: "extract", requests: &input.extractRequests}, nil
}); err != nil {
t.Fatal(err)
}
mergeSpec := defaultModuleSpec("merge", StageMerge)
mergeSpec.ArtifactKind = "test/notes"
if err := RegisterMergerBuilder(registries.Mergers, mergeSpec, strict, func(request BuildRequest) (contracts.Merger[codecNotes], error) {
record("merge", &request)
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(request BuildRequest) (contracts.Normalizer[codecNotes], error) {
record("normalize", &request)
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(request BuildRequest) (contracts.ChunkValidator, error) {
record("validator", &request)
return typedTestChunkValidator{key: "configured"}, nil
}); err != nil {
t.Fatal(err)
}
if err := RegisterTypedValidatorBuilder(registries.Validators, "test/notes", validatorSpec, strict, func(request BuildRequest) (contracts.TypedValidator[codecNotes], error) {
record("validator", &request)
return typedTestValidator[codecNotes]{key: "configured"}, nil
}); err != nil {
t.Fatal(err)
}
if err := registries.Outputs.RegisterBuilderWithSpec(defaultModuleSpec("output", StageOutput), strict, func(request BuildRequest) (contracts.OutputEncoder, error) {
record("output", &request)
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
extractRequests []contracts.TypedExtractionRequest
}
type constructionExtractor struct {
key string
requests *[]contracts.TypedExtractionRequest
}
func (extractor *constructionExtractor) Key() string { return extractor.key }
func (*constructionExtractor) ReferenceSlots() []contracts.ReferenceSlot { return nil }
func (extractor *constructionExtractor) Extract(_ context.Context, request contracts.TypedExtractionRequest) (contracts.TypedExtractionResult[codecNotes], error) {
if extractor.requests != nil {
*extractor.requests = append(*extractor.requests, request)
}
return contracts.TypedExtractionResult[codecNotes]{}, nil
}
func constructionReferenceSet(slotName, content string) contracts.ReferenceSet {
return contracts.ReferenceSet{Slots: map[string]contracts.ResolvedReferenceSlot{
slotName: {
Slot: contracts.ReferenceSlot{Name: slotName},
Items: []contracts.ReferenceItem{{SlotName: slotName, Content: []byte(content)}},
},
}}
}
func constructionReferenceContent(references contracts.ReferenceSet) string {
for _, slot := range references.Slots {
if len(slot.Items) > 0 {
return string(slot.Items[0].Content)
}
}
return ""
}
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
}