Deliver target references during pipeline preparation

This commit is contained in:
2026-07-20 18:57:44 +00:00
parent ac53f83ac8
commit 7806dba509
5 changed files with 167 additions and 43 deletions

View File

@@ -19,6 +19,7 @@ type ModuleDependencies struct {
type BuildRequest struct {
Dependencies ModuleDependencies
Options map[string]any
References contracts.ReferenceSet
}
// OptionValidator validates one module binding without constructing it.
@@ -39,6 +40,7 @@ func cloneBuildRequest(request BuildRequest) BuildRequest {
return BuildRequest{
Dependencies: request.Dependencies,
Options: cloneOptions(request.Options),
References: CloneReferenceSet(request.References),
}
}

View File

@@ -72,6 +72,66 @@ func TestPrepareConstructsEverythingInStableOrder(t *testing.T) {
}
}
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
@@ -146,6 +206,15 @@ func constructionProfile() PipelineProfile {
}
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{}
@@ -153,7 +222,15 @@ func constructionRegistries(t *testing.T, built *[]string, failure *construction
if failure == nil {
failure = &constructionFailure{}
}
record := func(name string) { *built = append(*built, name) }
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{
@@ -164,14 +241,14 @@ func constructionRegistries(t *testing.T, built *[]string, failure *construction
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")
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(BuildRequest) (contracts.Chunker, error) {
record("chunk")
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)
@@ -179,45 +256,45 @@ func constructionRegistries(t *testing.T, built *[]string, failure *construction
extractSpec := defaultModuleSpec("extract", StageExtract)
extractSpec.ArtifactKind = "test/notes"
if err := RegisterExtractorBuilder(registries.Extractors, extractSpec, strict, func(request BuildRequest) (contracts.Extractor[codecNotes], error) {
record("extract")
record("extract", &request)
if failure.requireExtractorLLM && request.Dependencies.LLM == nil {
return nil, errors.New("structured LLM client is required")
}
return typedTestExtractor[codecNotes]{key: "extract"}, nil
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(BuildRequest) (contracts.Merger[codecNotes], error) {
record("merge")
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(BuildRequest) (contracts.Normalizer[codecNotes], error) {
record("normalize")
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(BuildRequest) (contracts.ChunkValidator, error) {
record("validator")
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(BuildRequest) (contracts.TypedValidator[codecNotes], error) {
record("validator")
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(BuildRequest) (contracts.OutputEncoder, error) {
record("output")
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
}
@@ -229,8 +306,41 @@ func constructionRegistries(t *testing.T, built *[]string, failure *construction
}
type constructionInput struct {
key string
requests []contracts.ParseRequest
extractRequests []contracts.TypedExtractionRequest
}
type constructionExtractor struct {
key string
requests []contracts.ParseRequest
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 }

View File

@@ -78,22 +78,22 @@ func Prepare(resolved ResolvedPipeline, registries Registries, deps ModuleDepend
resolved: stable,
dependencies: deps,
}
request := func(binding ModuleBinding) BuildRequest {
return BuildRequest{Dependencies: deps, Options: cloneOptions(binding.Options)}
request := func(binding ModuleBinding, references contracts.ReferenceSet) BuildRequest {
return BuildRequest{Dependencies: deps, Options: cloneOptions(binding.Options), References: references}
}
input, err := registries.Inputs.BuildWithRequest(stable.Input.Module, request(stable.Input))
input, err := registries.Inputs.BuildWithRequest(stable.Input.Module, request(stable.Input, contracts.ReferenceSet{}))
if err != nil {
return nil, constructionError(stable.ID, "", StageInput, stable.Input.Module, "", err)
}
prepared.input = input
chunker, err := registries.Chunkers.BuildWithRequest(stable.Chunk.Module, request(stable.Chunk))
chunker, err := registries.Chunkers.BuildWithRequest(stable.Chunk.Module, request(stable.Chunk, stable.ChunkReferences.ReferenceSet))
if err != nil {
return nil, constructionError(stable.ID, "", StageChunk, stable.Chunk.Module, "", err)
}
prepared.chunker = chunker
prepared.chunkValidators, err = prepareValidatorChain(stable, registries, deps, StageChunk, "", stable.Chunk.Module)
prepared.chunkValidators, err = prepareValidatorChain(stable, registries, deps, StageChunk, "", stable.Chunk.Module, stable.ChunkReferences.ReferenceSet)
if err != nil {
return nil, err
}
@@ -109,7 +109,7 @@ func Prepare(resolved ResolvedPipeline, registries Registries, deps ModuleDepend
prepared.lanes = append(prepared.lanes, executor)
}
output, err := registries.Outputs.BuildWithRequest(stable.Output.Module, request(stable.Output))
output, err := registries.Outputs.BuildWithRequest(stable.Output.Module, request(stable.Output, contracts.ReferenceSet{}))
if err != nil {
return nil, constructionError(stable.ID, "", StageOutput, stable.Output.Module, "", err)
}
@@ -119,14 +119,14 @@ func Prepare(resolved ResolvedPipeline, registries Registries, deps ModuleDepend
func prepareLane(pipeline ResolvedPipeline, lane ResolvedArtifactLane, registries Registries, deps ModuleDependencies) (preparedLaneExecutor, error) {
executor := preparedLaneExecutor{resolved: cloneResolvedArtifactLane(lane)}
request := func(binding ModuleBinding) BuildRequest {
return BuildRequest{Dependencies: deps, Options: cloneOptions(binding.Options)}
request := func(binding ModuleBinding, references contracts.ReferenceSet) BuildRequest {
return BuildRequest{Dependencies: deps, Options: cloneOptions(binding.Options), References: references}
}
extractEntry, ok := registries.Extractors.typedEntry(lane.Extract.Module)
if !ok {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageExtract, lane.Extract.Module, "", fmt.Errorf("typed construction entry is not registered"))
}
module, err := buildErasedModule(extractEntry.builder, request(lane.Extract), lane.Extract.Module, "extractor")
module, err := buildErasedModule(extractEntry.builder, request(lane.Extract, lane.ExtractReferences.ReferenceSet), lane.Extract.Module, "extractor")
if err != nil {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageExtract, lane.Extract.Module, "", err)
}
@@ -136,7 +136,7 @@ func prepareLane(pipeline ResolvedPipeline, lane ResolvedArtifactLane, registrie
}
executor.typed = &preparedTypedLane{extractor: module, extract: extractEntry.extract, codec: codec}
executor.extractValidators, err = prepareValidatorChain(pipeline, registries, deps, StageExtract, lane.ID, lane.Extract.Module)
executor.extractValidators, err = prepareValidatorChain(pipeline, registries, deps, StageExtract, lane.ID, lane.Extract.Module, lane.ExtractReferences.ReferenceSet)
if err != nil {
return preparedLaneExecutor{}, err
}
@@ -145,13 +145,13 @@ func prepareLane(pipeline ResolvedPipeline, lane ResolvedArtifactLane, registrie
if !ok {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageMerge, lane.Merge.Module, "", fmt.Errorf("typed construction entry is not registered"))
}
module, err = buildErasedModule(mergeEntry.builder, request(lane.Merge), lane.Merge.Module, "merger")
module, err = buildErasedModule(mergeEntry.builder, request(lane.Merge, lane.MergeReferences.ReferenceSet), lane.Merge.Module, "merger")
if err != nil {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageMerge, lane.Merge.Module, "", err)
}
executor.typed.merger = module
executor.typed.merge = mergeEntry.merge
executor.mergeValidators, err = prepareValidatorChain(pipeline, registries, deps, StageMerge, lane.ID, lane.Merge.Module)
executor.mergeValidators, err = prepareValidatorChain(pipeline, registries, deps, StageMerge, lane.ID, lane.Merge.Module, lane.MergeReferences.ReferenceSet)
if err != nil {
return preparedLaneExecutor{}, err
}
@@ -160,24 +160,24 @@ func prepareLane(pipeline ResolvedPipeline, lane ResolvedArtifactLane, registrie
if !ok {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageNormalize, lane.Normalize.Module, "", fmt.Errorf("typed construction entry is not registered"))
}
module, err = buildErasedModule(normalizeEntry.builder, request(lane.Normalize), lane.Normalize.Module, "normalizer")
module, err = buildErasedModule(normalizeEntry.builder, request(lane.Normalize, lane.NormalizeReferences.ReferenceSet), lane.Normalize.Module, "normalizer")
if err != nil {
return preparedLaneExecutor{}, constructionError(pipeline.ID, lane.ID, StageNormalize, lane.Normalize.Module, "", err)
}
executor.typed.normalizer = module
executor.typed.normalize = normalizeEntry.normalize
executor.normalizeValidators, err = prepareValidatorChain(pipeline, registries, deps, StageNormalize, lane.ID, lane.Normalize.Module)
executor.normalizeValidators, err = prepareValidatorChain(pipeline, registries, deps, StageNormalize, lane.ID, lane.Normalize.Module, lane.NormalizeReferences.ReferenceSet)
if err != nil {
return preparedLaneExecutor{}, err
}
return executor, nil
}
func prepareValidatorChain(pipeline ResolvedPipeline, registries Registries, deps ModuleDependencies, stage ModuleStage, laneID, moduleKey string) (preparedValidatorChain, error) {
func prepareValidatorChain(pipeline ResolvedPipeline, registries Registries, deps ModuleDependencies, stage ModuleStage, laneID, moduleKey string, references contracts.ReferenceSet) (preparedValidatorChain, error) {
resolved := resolvedValidatorChain(stage, laneID, moduleKey, pipeline.ValidatorChains)
prepared := preparedValidatorChain{resolved: resolved}
for _, validator := range resolved.Validators {
request := BuildRequest{Dependencies: deps, Options: cloneOptions(validator.Binding.Options)}
request := BuildRequest{Dependencies: deps, Options: cloneOptions(validator.Binding.Options), References: references}
built, err := buildPreparedValidator(registries.Validators, validator, request)
if err != nil {
return preparedValidatorChain{}, constructionError(pipeline.ID, laneID, stage, moduleKey, validator.Binding.Module, err)