Deliver target references during pipeline preparation
This commit is contained in:
@@ -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 }
|
||||
|
||||
Reference in New Issue
Block a user