Resolve pipeline LLM profile defaults
This commit is contained in:
@@ -276,7 +276,7 @@ func TestRunLLMProfileOverrideAndValidationUseInjectedBoundaries(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("validator profile remains distinct", func(t *testing.T) {
|
t.Run("runtime override applies to validators", func(t *testing.T) {
|
||||||
roots := newStateTestRoots(t)
|
roots := newStateTestRoots(t)
|
||||||
profileDir := writeRunContractProfiles(t, "override-profile", "validator-profile")
|
profileDir := writeRunContractProfiles(t, "override-profile", "validator-profile")
|
||||||
prependRunContractConfig(t, roots, fmt.Sprintf("promptkit:\n profile_dir: %q\n", profileDir))
|
prependRunContractConfig(t, roots, fmt.Sprintf("promptkit:\n profile_dir: %q\n", profileDir))
|
||||||
@@ -294,11 +294,11 @@ func TestRunLLMProfileOverrideAndValidationUseInjectedBoundaries(t *testing.T) {
|
|||||||
if code != 0 || stderr.Len() != 0 {
|
if code != 0 || stderr.Len() != 0 {
|
||||||
t.Fatalf("code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String())
|
t.Fatalf("code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String())
|
||||||
}
|
}
|
||||||
if len(factoryProfiles) != 1 || factoryProfiles[0] != "" {
|
if len(factoryProfiles) != 1 || factoryProfiles[0] != "override-profile" {
|
||||||
t.Fatalf("factory profiles = %#v, want one call without a unique profile", factoryProfiles)
|
t.Fatalf("factory profiles = %#v, want one override profile", factoryProfiles)
|
||||||
}
|
}
|
||||||
if len(validatorProfiles) != 1 || validatorProfiles[0] != "validator-profile" {
|
if len(validatorProfiles) != 1 || validatorProfiles[0] != "override-profile" {
|
||||||
t.Fatalf("validator profiles = %#v, want configured validator profile", validatorProfiles)
|
t.Fatalf("validator profiles = %#v, want runtime override", validatorProfiles)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
@@ -826,16 +826,16 @@ func (h *stateTestHarness) options() Options {
|
|||||||
if err := registries.Inputs.RegisterBuilderWithSpec(pipeline.ModuleSpec{Key: "test/input", Stage: pipeline.StageInput, ExecutionClass: contracts.ExecutionClassDeterministic, Provides: []string{"source"}}, func(map[string]any) error { return nil }, func(pipeline.BuildRequest) (contracts.InputAdapter, error) { return stateTestInput{}, nil }); err != nil {
|
if err := registries.Inputs.RegisterBuilderWithSpec(pipeline.ModuleSpec{Key: "test/input", Stage: pipeline.StageInput, ExecutionClass: contracts.ExecutionClassDeterministic, Provides: []string{"source"}}, func(map[string]any) error { return nil }, func(pipeline.BuildRequest) (contracts.InputAdapter, error) { return stateTestInput{}, nil }); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if err := registries.Chunkers.RegisterBuilderWithSpec(pipeline.ModuleSpec{Key: "test/chunk", Stage: pipeline.StageChunk, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"source"}, Provides: []string{"chunks"}, ReferenceSlots: []contracts.ReferenceSlot{{Name: "cache-reference"}}}, func(map[string]any) error { return nil }, func(pipeline.BuildRequest) (contracts.Chunker, error) { return stateTestChunker{h}, nil }); err != nil {
|
if err := registries.Chunkers.RegisterBuilderWithSpec(pipeline.ModuleSpec{Key: "test/chunk", Stage: pipeline.StageChunk, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"source"}, Provides: []string{"chunks"}, ReferenceSlots: []contracts.ReferenceSlot{{Name: "cache-reference"}}}, func(map[string]any) error { return nil }, func(pipeline.BuildRequest) (contracts.Chunker, error) { return stateTestChunker{h}, nil }); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterExtractor(registries.Extractors, pipeline.ModuleSpec{Key: "test/extract", Stage: pipeline.StageExtract, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"chunks"}, Provides: []string{"artifact"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Extractor[stateTestArtifact], error) { return stateTestExtractor{h}, nil }); err != nil {
|
if err := pipeline.RegisterExtractor(registries.Extractors, pipeline.ModuleSpec{Key: "test/extract", Stage: pipeline.StageExtract, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"chunks"}, Provides: []string{"artifact"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Extractor[stateTestArtifact], error) { return stateTestExtractor{h}, nil }); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterMerger(registries.Mergers, pipeline.ModuleSpec{Key: "test/merge", Stage: pipeline.StageMerge, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"artifact"}, Provides: []string{"merged"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Merger[stateTestArtifact], error) { return stateTestMerger{harness: h}, nil }); err != nil {
|
if err := pipeline.RegisterMerger(registries.Mergers, pipeline.ModuleSpec{Key: "test/merge", Stage: pipeline.StageMerge, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"artifact"}, Provides: []string{"merged"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Merger[stateTestArtifact], error) { return stateTestMerger{harness: h}, nil }); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterNormalizer(registries.Normalizers, pipeline.ModuleSpec{Key: "test/normalize", Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"merged"}, Provides: []string{"normalized"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Normalizer[stateTestArtifact], error) { return stateTestNormalizer{harness: h}, nil }); err != nil {
|
if err := pipeline.RegisterNormalizer(registries.Normalizers, pipeline.ModuleSpec{Key: "test/normalize", Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"merged"}, Provides: []string{"normalized"}, ArtifactKind: stateTestArtifactKind}, func() (contracts.Normalizer[stateTestArtifact], error) { return stateTestNormalizer{harness: h}, nil }); err != nil {
|
||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
if err := registries.Outputs.RegisterWithSpec(pipeline.ModuleSpec{Key: "test/output", Stage: pipeline.StageOutput, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"normalized"}, Provides: []string{"output"}}, func() (contracts.OutputEncoder, error) {
|
if err := registries.Outputs.RegisterWithSpec(pipeline.ModuleSpec{Key: "test/output", Stage: pipeline.StageOutput, ExecutionClass: contracts.ExecutionClassDeterministic, Requires: []string{"normalized"}, Provides: []string{"output"}}, func() (contracts.OutputEncoder, error) {
|
||||||
|
|||||||
@@ -42,12 +42,10 @@ func (c Config) Resolve(input ResolveInput) (EffectiveConfig, error) {
|
|||||||
}
|
}
|
||||||
profile = clonePipelineProfile(profile)
|
profile = clonePipelineProfile(profile)
|
||||||
profile.ID = pipelineID
|
profile.ID = pipelineID
|
||||||
if override := strings.TrimSpace(input.LLMProfileOverride); override != "" {
|
|
||||||
applyLLMProfileOverride(&profile, override)
|
|
||||||
}
|
|
||||||
|
|
||||||
resolved, err := pipeline.ResolvePipeline(profile, pipeline.ResolveOptions{
|
resolved, err := pipeline.ResolvePipeline(profile, pipeline.ResolveOptions{
|
||||||
Only: input.Only,
|
Only: input.Only,
|
||||||
|
LLMProfileOverride: input.LLMProfileOverride,
|
||||||
ReferenceOverrides: append([]pipeline.ReferenceBinding(nil), input.ReferenceOverrides...),
|
ReferenceOverrides: append([]pipeline.ReferenceBinding(nil), input.ReferenceOverrides...),
|
||||||
ReferenceUnbinds: append([]pipeline.ReferenceUnbind(nil), input.ReferenceUnbinds...),
|
ReferenceUnbinds: append([]pipeline.ReferenceUnbind(nil), input.ReferenceUnbinds...),
|
||||||
}, input.Catalog)
|
}, input.Catalog)
|
||||||
@@ -65,22 +63,6 @@ func (c Config) Resolve(input ResolveInput) (EffectiveConfig, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func applyLLMProfileOverride(profile *pipeline.PipelineProfile, profileID string) {
|
|
||||||
profile.Chunk.LLMProfile = profileID
|
|
||||||
apply := func(artifacts map[string]pipeline.ArtifactLaneProfile) {
|
|
||||||
for laneID, lane := range artifacts {
|
|
||||||
lane.Extract.LLMProfile = profileID
|
|
||||||
lane.Merge.LLMProfile = profileID
|
|
||||||
lane.Normalize.LLMProfile = profileID
|
|
||||||
artifacts[laneID] = lane
|
|
||||||
}
|
|
||||||
}
|
|
||||||
apply(profile.Artifacts)
|
|
||||||
for index := range profile.Steps {
|
|
||||||
apply(profile.Steps[index].Artifacts)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func lookupPipelineProfile(profiles map[string]pipeline.PipelineProfile, pipelineID string) (pipeline.PipelineProfile, bool) {
|
func lookupPipelineProfile(profiles map[string]pipeline.PipelineProfile, pipelineID string) (pipeline.PipelineProfile, bool) {
|
||||||
pipelineID = strings.TrimSpace(pipelineID)
|
pipelineID = strings.TrimSpace(pipelineID)
|
||||||
for rawID, profile := range profiles {
|
for rawID, profile := range profiles {
|
||||||
|
|||||||
@@ -204,7 +204,7 @@ func TestEffectiveConfigResolutionFailuresRetainContext(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEffectiveConfigLLMProfileOverrideChangesDigestWithoutOverridingValidators(t *testing.T) {
|
func TestEffectiveConfigLLMProfileOverrideChangesDigestAndOverridesValidators(t *testing.T) {
|
||||||
profile := effectiveProfile()
|
profile := effectiveProfile()
|
||||||
profile.Chunk.LLMProfile = "chunk-profile"
|
profile.Chunk.LLMProfile = "chunk-profile"
|
||||||
lane := profile.Artifacts["lane"]
|
lane := profile.Artifacts["lane"]
|
||||||
@@ -237,8 +237,8 @@ func TestEffectiveConfigLLMProfileOverrideChangesDigestWithoutOverridingValidato
|
|||||||
t.Fatalf("pipeline profile override was not applied: %#v", resolved)
|
t.Fatalf("pipeline profile override was not applied: %#v", resolved)
|
||||||
}
|
}
|
||||||
validators := findEffectiveValidatorChain(resolved, pipeline.StageExtract, "lane")
|
validators := findEffectiveValidatorChain(resolved, pipeline.StageExtract, "lane")
|
||||||
if len(validators.Validators) != 1 || validators.Validators[0].Binding.LLMProfile != "validator-profile" {
|
if len(validators.Validators) != 1 || validators.Validators[0].Binding.LLMProfile != "override-profile" {
|
||||||
t.Fatalf("validator profile was overridden: %#v", validators)
|
t.Fatalf("validator profile = %#v, want runtime override", validators)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -493,7 +493,7 @@ func effectiveCatalog(t *testing.T) pipeline.ModuleCatalog {
|
|||||||
chunkSpec := pipeline.ModuleSpec{
|
chunkSpec := pipeline.ModuleSpec{
|
||||||
Key: pipeline.DefaultChunkModule,
|
Key: pipeline.DefaultChunkModule,
|
||||||
Stage: pipeline.StageChunk,
|
Stage: pipeline.StageChunk,
|
||||||
ExecutionClass: contracts.ExecutionClassDeterministic,
|
ExecutionClass: contracts.ExecutionClassLLMBacked,
|
||||||
Requires: []string{"source"},
|
Requires: []string{"source"},
|
||||||
Provides: []string{"chunk"},
|
Provides: []string{"chunk"},
|
||||||
ReferenceSlots: []contracts.ReferenceSlot{{Name: "chunk-ref"}},
|
ReferenceSlots: []contracts.ReferenceSlot{{Name: "chunk-ref"}},
|
||||||
@@ -509,12 +509,12 @@ func effectiveCatalog(t *testing.T) pipeline.ModuleCatalog {
|
|||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterExtractor(catalog.Extractors, pipeline.ModuleSpec{Key: "extract", Stage: pipeline.StageExtract, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: effectiveArtifactKind, Requires: []string{"chunk"}, Provides: []string{"candidate"}}, func() (contracts.Extractor[effectiveArtifact], error) {
|
if err := pipeline.RegisterExtractor(catalog.Extractors, pipeline.ModuleSpec{Key: "extract", Stage: pipeline.StageExtract, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: effectiveArtifactKind, Requires: []string{"chunk"}, Provides: []string{"candidate"}}, func() (contracts.Extractor[effectiveArtifact], error) {
|
||||||
return effectiveExtractor{key: "extract"}, nil
|
return effectiveExtractor{key: "extract"}, nil
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterMerger(catalog.Mergers, pipeline.ModuleSpec{Key: pipeline.DefaultMergeModule, Stage: pipeline.StageMerge, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: effectiveArtifactKind, Requires: []string{"candidate"}, Provides: []string{"merged"}}, func() (contracts.Merger[effectiveArtifact], error) {
|
if err := pipeline.RegisterMerger(catalog.Mergers, pipeline.ModuleSpec{Key: pipeline.DefaultMergeModule, Stage: pipeline.StageMerge, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: effectiveArtifactKind, Requires: []string{"candidate"}, Provides: []string{"merged"}}, func() (contracts.Merger[effectiveArtifact], error) {
|
||||||
return effectiveMerger{key: pipeline.DefaultMergeModule}, nil
|
return effectiveMerger{key: pipeline.DefaultMergeModule}, nil
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -524,7 +524,7 @@ func effectiveCatalog(t *testing.T) pipeline.ModuleCatalog {
|
|||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := pipeline.RegisterNormalizer(catalog.Normalizers, pipeline.ModuleSpec{Key: pipeline.DefaultNormalizeModule, Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassDeterministic, ArtifactKind: effectiveArtifactKind, Requires: []string{"merged"}, Provides: []string{"normalized"}}, func() (contracts.Normalizer[effectiveArtifact], error) {
|
if err := pipeline.RegisterNormalizer(catalog.Normalizers, pipeline.ModuleSpec{Key: pipeline.DefaultNormalizeModule, Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: effectiveArtifactKind, Requires: []string{"merged"}, Provides: []string{"normalized"}}, func() (contracts.Normalizer[effectiveArtifact], error) {
|
||||||
return effectiveNormalizer{key: pipeline.DefaultNormalizeModule}, nil
|
return effectiveNormalizer{key: pipeline.DefaultNormalizeModule}, nil
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
|
|||||||
@@ -117,6 +117,7 @@ type PipelineStepProfile struct {
|
|||||||
|
|
||||||
type PipelineProfile struct {
|
type PipelineProfile struct {
|
||||||
ID string `json:"id"`
|
ID string `json:"id"`
|
||||||
|
LLMProfile string `json:"llm_profile,omitempty"`
|
||||||
Input ModuleBinding `json:"input"`
|
Input ModuleBinding `json:"input"`
|
||||||
Chunk ModuleBinding `json:"chunk,omitempty"`
|
Chunk ModuleBinding `json:"chunk,omitempty"`
|
||||||
Artifacts map[string]ArtifactLaneProfile `json:"artifacts"`
|
Artifacts map[string]ArtifactLaneProfile `json:"artifacts"`
|
||||||
@@ -127,6 +128,7 @@ type PipelineProfile struct {
|
|||||||
|
|
||||||
type ResolveOptions struct {
|
type ResolveOptions struct {
|
||||||
Only []string
|
Only []string
|
||||||
|
LLMProfileOverride string
|
||||||
ReferenceOverrides []ReferenceBinding
|
ReferenceOverrides []ReferenceBinding
|
||||||
ReferenceUnbinds []ReferenceUnbind
|
ReferenceUnbinds []ReferenceUnbind
|
||||||
}
|
}
|
||||||
@@ -435,6 +437,9 @@ func ResolvePipeline(profile PipelineProfile, options ResolveOptions, catalog Mo
|
|||||||
return ResolvedPipeline{}, capabilityError(pipelineID, "", StageOutput, resolved.Output.Module, missing)
|
return ResolvedPipeline{}, capabilityError(pipelineID, "", StageOutput, resolved.Output.Module, missing)
|
||||||
}
|
}
|
||||||
resolved.OutputExecutionClass = outputSpec.ExecutionClass
|
resolved.OutputExecutionClass = outputSpec.ExecutionClass
|
||||||
|
if err := applyEffectiveLLMProfiles(&resolved, profile.LLMProfile, options.LLMProfileOverride); err != nil {
|
||||||
|
return ResolvedPipeline{}, err
|
||||||
|
}
|
||||||
if err := validateResolvedOptions(resolved, catalog, configuredLaneIDs); err != nil {
|
if err := validateResolvedOptions(resolved, catalog, configuredLaneIDs); err != nil {
|
||||||
return ResolvedPipeline{}, err
|
return ResolvedPipeline{}, err
|
||||||
}
|
}
|
||||||
@@ -836,9 +841,6 @@ func resolveValidatorChain(pipelineID string, laneID string, stage ModuleStage,
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return ResolvedValidatorChain{}, fmt.Errorf("pipeline %q %s validator chain for module %q: %w", pipelineID, stage, chain.ModuleKey, err)
|
return ResolvedValidatorChain{}, fmt.Errorf("pipeline %q %s validator chain for module %q: %w", pipelineID, stage, chain.ModuleKey, err)
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(validator.LLMProfile) != "" && spec.ExecutionClass != contracts.ExecutionClassLLMBacked {
|
|
||||||
return ResolvedValidatorChain{}, fmt.Errorf("pipeline %q %s validator chain for module %q assigns llm_profile to deterministic validator %q", pipelineID, stage, chain.ModuleKey, validator.Module)
|
|
||||||
}
|
|
||||||
chain.Validators = append(chain.Validators, ResolvedValidator{
|
chain.Validators = append(chain.Validators, ResolvedValidator{
|
||||||
Binding: cloneModuleBinding(validator),
|
Binding: cloneModuleBinding(validator),
|
||||||
ExecutionClass: spec.ExecutionClass,
|
ExecutionClass: spec.ExecutionClass,
|
||||||
@@ -1250,6 +1252,66 @@ func resolveBinding(binding ModuleBinding, defaultModule string) ModuleBinding {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func applyEffectiveLLMProfiles(resolved *ResolvedPipeline, pipelineProfile, overrideProfile string) error {
|
||||||
|
pipelineProfile = strings.TrimSpace(pipelineProfile)
|
||||||
|
overrideProfile = strings.TrimSpace(overrideProfile)
|
||||||
|
|
||||||
|
apply := func(stage ModuleStage, laneID, module string, binding *ModuleBinding, executionClass contracts.ExecutionClass, kind string) error {
|
||||||
|
binding.LLMProfile = strings.TrimSpace(binding.LLMProfile)
|
||||||
|
if executionClass != contracts.ExecutionClassLLMBacked {
|
||||||
|
if binding.LLMProfile != "" {
|
||||||
|
if laneID == "" {
|
||||||
|
return fmt.Errorf("pipeline %q %s %q assigns llm_profile to deterministic %s %q", resolved.ID, stage, module, kind, binding.Module)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("pipeline %q lane %q %s %q assigns llm_profile to deterministic %s %q", resolved.ID, laneID, stage, module, kind, binding.Module)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if overrideProfile != "" {
|
||||||
|
binding.LLMProfile = overrideProfile
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if binding.LLMProfile == "" {
|
||||||
|
binding.LLMProfile = pipelineProfile
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := apply(StageInput, "", resolved.Input.Module, &resolved.Input, resolved.InputExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := apply(StageChunk, "", resolved.Chunk.Module, &resolved.Chunk, resolved.ChunkExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for stepIndex := range resolved.Steps {
|
||||||
|
for laneIndex := range resolved.Steps[stepIndex].ArtifactLanes {
|
||||||
|
lane := &resolved.Steps[stepIndex].ArtifactLanes[laneIndex]
|
||||||
|
if err := apply(StageExtract, lane.ID, lane.Extract.Module, &lane.Extract, lane.ExtractExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := apply(StageMerge, lane.ID, lane.Merge.Module, &lane.Merge, lane.MergeExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := apply(StageNormalize, lane.ID, lane.Normalize.Module, &lane.Normalize, lane.NormalizeExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := apply(StageOutput, "", resolved.Output.Module, &resolved.Output, resolved.OutputExecutionClass, "module"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for chainIndex := range resolved.ValidatorChains {
|
||||||
|
chain := &resolved.ValidatorChains[chainIndex]
|
||||||
|
for validatorIndex := range chain.Validators {
|
||||||
|
validator := &chain.Validators[validatorIndex]
|
||||||
|
if err := apply(chain.Stage, chain.LaneID, chain.ModuleKey, &validator.Binding, validator.ExecutionClass, "validator"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func resolveBindings(bindings []ModuleBinding, defaultModule string) []ModuleBinding {
|
func resolveBindings(bindings []ModuleBinding, defaultModule string) []ModuleBinding {
|
||||||
if len(bindings) == 0 {
|
if len(bindings) == 0 {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -24,13 +24,13 @@ func TestResolvePipelineWithExplicitModules(t *testing.T) {
|
|||||||
|
|
||||||
resolved, err := ResolvePipeline(PipelineProfile{
|
resolved, err := ResolvePipeline(PipelineProfile{
|
||||||
ID: " campaign ",
|
ID: " campaign ",
|
||||||
Input: ModuleBinding{Module: " text ", LLMProfile: " fast "},
|
Input: Binding(" text "),
|
||||||
Chunk: ModuleBinding{Module: " window ", Options: map[string]any{
|
Chunk: ModuleBinding{Module: " window ", Options: map[string]any{
|
||||||
"size": 10,
|
"size": 10,
|
||||||
}},
|
}},
|
||||||
Artifacts: map[string]ArtifactLaneProfile{
|
Artifacts: map[string]ArtifactLaneProfile{
|
||||||
" records ": {
|
" records ": {
|
||||||
Extract: ModuleBinding{Module: " record-extractor ", LLMProfile: " careful "},
|
Extract: Binding(" record-extractor "),
|
||||||
Merge: Binding(" dedupe "),
|
Merge: Binding(" dedupe "),
|
||||||
Normalize: Binding(" canonical "),
|
Normalize: Binding(" canonical "),
|
||||||
},
|
},
|
||||||
@@ -44,7 +44,7 @@ func TestResolvePipelineWithExplicitModules(t *testing.T) {
|
|||||||
if resolved.ID != "campaign" {
|
if resolved.ID != "campaign" {
|
||||||
t.Fatalf("ID = %q, want campaign", resolved.ID)
|
t.Fatalf("ID = %q, want campaign", resolved.ID)
|
||||||
}
|
}
|
||||||
if !reflect.DeepEqual(resolved.Input, ModuleBinding{Module: "text", LLMProfile: "fast"}) {
|
if !reflect.DeepEqual(resolved.Input, ModuleBinding{Module: "text"}) {
|
||||||
t.Fatalf("Input = %#v, want trimmed explicit input", resolved.Input)
|
t.Fatalf("Input = %#v, want trimmed explicit input", resolved.Input)
|
||||||
}
|
}
|
||||||
if resolved.InputExecutionClass != contracts.ExecutionClassDeterministic {
|
if resolved.InputExecutionClass != contracts.ExecutionClassDeterministic {
|
||||||
@@ -66,7 +66,7 @@ func TestResolvePipelineWithExplicitModules(t *testing.T) {
|
|||||||
if lane.ID != "records" {
|
if lane.ID != "records" {
|
||||||
t.Fatalf("lane.ID = %q, want records", lane.ID)
|
t.Fatalf("lane.ID = %q, want records", lane.ID)
|
||||||
}
|
}
|
||||||
if !reflect.DeepEqual(lane.Extract, ModuleBinding{Module: "record-extractor", LLMProfile: "careful"}) {
|
if !reflect.DeepEqual(lane.Extract, ModuleBinding{Module: "record-extractor"}) {
|
||||||
t.Fatalf("lane.Extract = %#v, want explicit extractor", lane.Extract)
|
t.Fatalf("lane.Extract = %#v, want explicit extractor", lane.Extract)
|
||||||
}
|
}
|
||||||
if lane.ExtractExecutionClass != contracts.ExecutionClassDeterministic || lane.MergeExecutionClass != contracts.ExecutionClassDeterministic || lane.NormalizeExecutionClass != contracts.ExecutionClassDeterministic {
|
if lane.ExtractExecutionClass != contracts.ExecutionClassDeterministic || lane.MergeExecutionClass != contracts.ExecutionClassDeterministic || lane.NormalizeExecutionClass != contracts.ExecutionClassDeterministic {
|
||||||
@@ -159,6 +159,176 @@ func TestResolvePipelineAppliesDefaults(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestResolvePipelineAppliesEffectiveLLMProfiles(t *testing.T) {
|
||||||
|
for _, test := range []struct {
|
||||||
|
name string
|
||||||
|
profile PipelineProfile
|
||||||
|
options ResolveOptions
|
||||||
|
want map[string]string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "runtime override",
|
||||||
|
profile: func() PipelineProfile {
|
||||||
|
profile := llmProfilePipeline()
|
||||||
|
profile.LLMProfile = " pipeline "
|
||||||
|
profile.Chunk.LLMProfile = "binding"
|
||||||
|
return profile
|
||||||
|
}(),
|
||||||
|
options: ResolveOptions{LLMProfileOverride: " runtime "},
|
||||||
|
want: llmProfileValues("runtime"),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "binding exception",
|
||||||
|
profile: func() PipelineProfile {
|
||||||
|
profile := llmProfilePipeline()
|
||||||
|
profile.LLMProfile = "pipeline"
|
||||||
|
lane := profile.Artifacts["events"]
|
||||||
|
lane.Extract.LLMProfile = "extract"
|
||||||
|
lane.Extract.Validators = ValidatorOverride{
|
||||||
|
Set: true,
|
||||||
|
Validators: []ModuleBinding{{Module: "llm-validator", LLMProfile: "validator"}},
|
||||||
|
}
|
||||||
|
profile.Artifacts["events"] = lane
|
||||||
|
return profile
|
||||||
|
}(),
|
||||||
|
want: func() map[string]string {
|
||||||
|
values := llmProfileValues("pipeline")
|
||||||
|
values["extract"] = "extract"
|
||||||
|
values["validator:extract:events"] = "validator"
|
||||||
|
return values
|
||||||
|
}(),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "pipeline default",
|
||||||
|
profile: func() PipelineProfile {
|
||||||
|
profile := llmProfilePipeline()
|
||||||
|
profile.LLMProfile = "pipeline"
|
||||||
|
return profile
|
||||||
|
}(),
|
||||||
|
want: llmProfileValues("pipeline"),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "prompt fallback",
|
||||||
|
profile: llmProfilePipeline(),
|
||||||
|
want: llmProfileValues(""),
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
resolved, err := ResolvePipeline(test.profile, test.options, llmProfileCatalog(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if got := resolvedLLMProfileValues(resolved); !reflect.DeepEqual(got, test.want) {
|
||||||
|
t.Fatalf("resolved profiles = %#v, want %#v", got, test.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolvePipelineAppliesProfilesOnlyToSelectedLLMBackedBindings(t *testing.T) {
|
||||||
|
profile := llmProfilePipeline()
|
||||||
|
profile.LLMProfile = "pipeline"
|
||||||
|
profile.Artifacts["notes"] = ArtifactLaneProfile{Extract: Binding("llm-extractor")}
|
||||||
|
|
||||||
|
resolved, err := ResolvePipeline(profile, ResolveOptions{Only: []string{"events"}}, llmProfileCatalog(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if got := laneIDs(resolved.Steps[0].ArtifactLanes); !reflect.DeepEqual(got, []string{"events"}) {
|
||||||
|
t.Fatalf("selected lanes = %#v, want events only", got)
|
||||||
|
}
|
||||||
|
if got := resolvedLLMProfileValues(resolved); !reflect.DeepEqual(got, llmProfileValues("pipeline")) {
|
||||||
|
t.Fatalf("resolved profiles = %#v, want selected LLM bindings only", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolvePipelineLeavesUnusedProfilesOffDeterministicBindings(t *testing.T) {
|
||||||
|
profile := baselineProfile()
|
||||||
|
profile.LLMProfile = "unused"
|
||||||
|
|
||||||
|
resolved, err := ResolvePipeline(profile, ResolveOptions{LLMProfileOverride: "also-unused"}, newProfileCatalog(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if resolved.Input.LLMProfile != "" || resolved.Chunk.LLMProfile != "" || resolved.Output.LLMProfile != "" {
|
||||||
|
t.Fatalf("deterministic pipeline profiles = input %q chunk %q output %q, want empty", resolved.Input.LLMProfile, resolved.Chunk.LLMProfile, resolved.Output.LLMProfile)
|
||||||
|
}
|
||||||
|
lane := resolved.Steps[0].ArtifactLanes[0]
|
||||||
|
if lane.Extract.LLMProfile != "" || lane.Merge.LLMProfile != "" || lane.Normalize.LLMProfile != "" {
|
||||||
|
t.Fatalf("deterministic lane profiles = extract %q merge %q normalize %q, want empty", lane.Extract.LLMProfile, lane.Merge.LLMProfile, lane.Normalize.LLMProfile)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolvePipelineRejectsLLMProfileForDeterministicModule(t *testing.T) {
|
||||||
|
for _, test := range []struct {
|
||||||
|
name string
|
||||||
|
mutate func(*PipelineProfile)
|
||||||
|
}{
|
||||||
|
{name: "input", mutate: func(profile *PipelineProfile) { profile.Input.LLMProfile = "invalid" }},
|
||||||
|
{name: "chunk", mutate: func(profile *PipelineProfile) { profile.Chunk.LLMProfile = "invalid" }},
|
||||||
|
{name: "extract", mutate: func(profile *PipelineProfile) {
|
||||||
|
lane := profile.Artifacts["events"]
|
||||||
|
lane.Extract.LLMProfile = "invalid"
|
||||||
|
profile.Artifacts["events"] = lane
|
||||||
|
}},
|
||||||
|
{name: "merge", mutate: func(profile *PipelineProfile) {
|
||||||
|
lane := profile.Artifacts["events"]
|
||||||
|
lane.Merge.LLMProfile = "invalid"
|
||||||
|
profile.Artifacts["events"] = lane
|
||||||
|
}},
|
||||||
|
{name: "normalize", mutate: func(profile *PipelineProfile) {
|
||||||
|
lane := profile.Artifacts["events"]
|
||||||
|
lane.Normalize.LLMProfile = "invalid"
|
||||||
|
profile.Artifacts["events"] = lane
|
||||||
|
}},
|
||||||
|
{name: "output", mutate: func(profile *PipelineProfile) { profile.Output.LLMProfile = "invalid" }},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
profile := baselineProfile()
|
||||||
|
test.mutate(&profile)
|
||||||
|
_, err := ResolvePipeline(profile, ResolveOptions{}, newProfileCatalog(t))
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "llm_profile") || !strings.Contains(err.Error(), "deterministic") {
|
||||||
|
t.Fatalf("ResolvePipeline() error = %v, want deterministic profile rejection", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolvePipelineDigestUsesEffectiveLLMProfiles(t *testing.T) {
|
||||||
|
inherited := llmProfilePipeline()
|
||||||
|
inherited.LLMProfile = "shared"
|
||||||
|
explicit := llmProfilePipeline()
|
||||||
|
explicit.Input.LLMProfile = "shared"
|
||||||
|
explicit.Chunk.LLMProfile = "shared"
|
||||||
|
explicit.Output.LLMProfile = "shared"
|
||||||
|
lane := explicit.Artifacts["events"]
|
||||||
|
lane.Extract.LLMProfile = "shared"
|
||||||
|
lane.Merge.LLMProfile = "shared"
|
||||||
|
lane.Normalize.LLMProfile = "shared"
|
||||||
|
explicit.Artifacts["events"] = lane
|
||||||
|
|
||||||
|
inheritedResolved, err := ResolvePipeline(inherited, ResolveOptions{}, llmProfileCatalogWithoutValidatorChains(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline(inherited) error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
explicitResolved, err := ResolvePipeline(explicit, ResolveOptions{}, llmProfileCatalogWithoutValidatorChains(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline(explicit) error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if inheritedResolved.Digest != explicitResolved.Digest {
|
||||||
|
t.Fatalf("effective profile digests differ: %q != %q", inheritedResolved.Digest, explicitResolved.Digest)
|
||||||
|
}
|
||||||
|
|
||||||
|
explicit.Chunk.LLMProfile = "different"
|
||||||
|
changedResolved, err := ResolvePipeline(explicit, ResolveOptions{}, llmProfileCatalogWithoutValidatorChains(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ResolvePipeline(changed) error = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if explicitResolved.Digest == changedResolved.Digest {
|
||||||
|
t.Fatalf("digest = %q after effective profile change, want different", explicitResolved.Digest)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestResolvePipelineRecordsValidatorChains(t *testing.T) {
|
func TestResolvePipelineRecordsValidatorChains(t *testing.T) {
|
||||||
catalog := newProfileCatalog(t)
|
catalog := newProfileCatalog(t)
|
||||||
if err := catalog.ValidatorChains.Register(ValidatorChainMapping{
|
if err := catalog.ValidatorChains.Register(ValidatorChainMapping{
|
||||||
@@ -1543,6 +1713,92 @@ func defaultProfileSpecs() []ModuleSpec {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func llmProfileCatalog(t *testing.T) ModuleCatalog {
|
||||||
|
t.Helper()
|
||||||
|
catalog := newProfileCatalogWithOverrides(t,
|
||||||
|
ModuleSpec{Key: "llm-input", Stage: StageInput, ExecutionClass: contracts.ExecutionClassLLMBacked, Provides: []string{"source"}},
|
||||||
|
ModuleSpec{Key: "llm-chunk", Stage: StageChunk, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"source"}, Provides: []string{"chunk"}},
|
||||||
|
ModuleSpec{Key: "llm-extractor", Stage: StageExtract, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: "test/notes", Requires: []string{"chunk"}, Provides: []string{"candidate"}},
|
||||||
|
ModuleSpec{Key: "llm-merge", Stage: StageMerge, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: "test/notes", Requires: []string{"candidate"}, Provides: []string{"merged"}},
|
||||||
|
ModuleSpec{Key: "llm-normalize", Stage: StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, ArtifactKind: "test/notes", Requires: []string{"merged"}, Provides: []string{"normalized"}},
|
||||||
|
ModuleSpec{Key: "llm-output", Stage: StageOutput, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"normalized"}, Provides: []string{"encoded"}},
|
||||||
|
)
|
||||||
|
if err := RegisterChunkValidator(catalog.Validators, ValidatorSpec{Key: "llm-chunk-validator", ExecutionClass: contracts.ExecutionClassLLMBacked}, func() (contracts.ChunkValidator, error) {
|
||||||
|
return typedTestChunkValidator{key: "llm-chunk-validator"}, nil
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("register chunk validator: %v", err)
|
||||||
|
}
|
||||||
|
registerProfileValidatorSpec(t, catalog, ValidatorSpec{Key: "llm-validator", ExecutionClass: contracts.ExecutionClassLLMBacked})
|
||||||
|
for _, mapping := range []ValidatorChainMapping{
|
||||||
|
{Stage: StageChunk, Module: "llm-chunk", Validators: []ModuleBinding{Binding("llm-chunk-validator")}},
|
||||||
|
{Stage: StageExtract, Module: "llm-extractor", Validators: []ModuleBinding{Binding("llm-validator")}},
|
||||||
|
{Stage: StageMerge, Module: "llm-merge", Validators: []ModuleBinding{Binding("llm-validator")}},
|
||||||
|
{Stage: StageNormalize, Module: "llm-normalize", Validators: []ModuleBinding{Binding("llm-validator")}},
|
||||||
|
} {
|
||||||
|
if err := catalog.ValidatorChains.Register(mapping); err != nil {
|
||||||
|
t.Fatalf("register validator chain %#v: %v", mapping, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return catalog
|
||||||
|
}
|
||||||
|
|
||||||
|
func llmProfileCatalogWithoutValidatorChains(t *testing.T) ModuleCatalog {
|
||||||
|
catalog := llmProfileCatalog(t)
|
||||||
|
catalog.ValidatorChains = NewValidatorChainRegistry()
|
||||||
|
return catalog
|
||||||
|
}
|
||||||
|
|
||||||
|
func llmProfilePipeline() PipelineProfile {
|
||||||
|
return PipelineProfile{
|
||||||
|
ID: "llm-profile",
|
||||||
|
Input: Binding("llm-input"),
|
||||||
|
Chunk: Binding("llm-chunk"),
|
||||||
|
Output: Binding("llm-output"),
|
||||||
|
Artifacts: map[string]ArtifactLaneProfile{
|
||||||
|
"events": {
|
||||||
|
Extract: Binding("llm-extractor"),
|
||||||
|
Merge: Binding("llm-merge"),
|
||||||
|
Normalize: Binding("llm-normalize"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func llmProfileValues(profile string) map[string]string {
|
||||||
|
return map[string]string{
|
||||||
|
"input": profile,
|
||||||
|
"chunk": profile,
|
||||||
|
"extract": profile,
|
||||||
|
"merge": profile,
|
||||||
|
"normalize": profile,
|
||||||
|
"output": profile,
|
||||||
|
"validator:chunk:": profile,
|
||||||
|
"validator:extract:events": profile,
|
||||||
|
"validator:merge:events": profile,
|
||||||
|
"validator:normalize:events": profile,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolvedLLMProfileValues(resolved ResolvedPipeline) map[string]string {
|
||||||
|
values := map[string]string{
|
||||||
|
"input": resolved.Input.LLMProfile,
|
||||||
|
"chunk": resolved.Chunk.LLMProfile,
|
||||||
|
"output": resolved.Output.LLMProfile,
|
||||||
|
}
|
||||||
|
lane := resolved.Steps[0].ArtifactLanes[0]
|
||||||
|
values["extract"] = lane.Extract.LLMProfile
|
||||||
|
values["merge"] = lane.Merge.LLMProfile
|
||||||
|
values["normalize"] = lane.Normalize.LLMProfile
|
||||||
|
for _, chain := range resolved.ValidatorChains {
|
||||||
|
for _, validator := range chain.Validators {
|
||||||
|
if validator.ExecutionClass == contracts.ExecutionClassLLMBacked {
|
||||||
|
values["validator:"+string(chain.Stage)+":"+chain.LaneID] = validator.Binding.LLMProfile
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return values
|
||||||
|
}
|
||||||
|
|
||||||
func registerProfileSpecs(t *testing.T, catalog ModuleCatalog, specs ...ModuleSpec) {
|
func registerProfileSpecs(t *testing.T, catalog ModuleCatalog, specs ...ModuleSpec) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user