Declare producer correction capability
This commit is contained in:
@@ -21,25 +21,28 @@ const (
|
||||
)
|
||||
|
||||
type ModuleSpec struct {
|
||||
Key string
|
||||
Stage ModuleStage
|
||||
ExecutionClass contracts.ExecutionClass
|
||||
ArtifactKind contracts.ArtifactKind
|
||||
Provides []string
|
||||
Requires []string
|
||||
ReferenceSlots []contracts.ReferenceSlot
|
||||
Key string
|
||||
Stage ModuleStage
|
||||
ExecutionClass contracts.ExecutionClass
|
||||
CorrectionProtocol contracts.CorrectionProtocol
|
||||
ArtifactKind contracts.ArtifactKind
|
||||
Provides []string
|
||||
Requires []string
|
||||
ReferenceSlots []contracts.ReferenceSlot
|
||||
}
|
||||
|
||||
func normalizeModuleSpec(spec ModuleSpec) ModuleSpec {
|
||||
executionClass := contracts.ExecutionClass(strings.TrimSpace(string(spec.ExecutionClass)))
|
||||
correctionProtocol := contracts.CorrectionProtocol(strings.TrimSpace(string(spec.CorrectionProtocol)))
|
||||
return ModuleSpec{
|
||||
Key: strings.TrimSpace(spec.Key),
|
||||
Stage: spec.Stage,
|
||||
ExecutionClass: executionClass,
|
||||
ArtifactKind: normalizeArtifactKind(spec.ArtifactKind),
|
||||
Provides: normalizeCapabilities(spec.Provides),
|
||||
Requires: normalizeCapabilities(spec.Requires),
|
||||
ReferenceSlots: normalizeReferenceSlots(spec.ReferenceSlots),
|
||||
Key: strings.TrimSpace(spec.Key),
|
||||
Stage: spec.Stage,
|
||||
ExecutionClass: executionClass,
|
||||
CorrectionProtocol: correctionProtocol,
|
||||
ArtifactKind: normalizeArtifactKind(spec.ArtifactKind),
|
||||
Provides: normalizeCapabilities(spec.Provides),
|
||||
Requires: normalizeCapabilities(spec.Requires),
|
||||
ReferenceSlots: normalizeReferenceSlots(spec.ReferenceSlots),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -70,13 +73,14 @@ func normalizeCapabilities(values []string) []string {
|
||||
|
||||
func cloneModuleSpec(spec ModuleSpec) ModuleSpec {
|
||||
return ModuleSpec{
|
||||
Key: spec.Key,
|
||||
Stage: spec.Stage,
|
||||
ExecutionClass: spec.ExecutionClass,
|
||||
ArtifactKind: spec.ArtifactKind,
|
||||
Provides: append([]string(nil), spec.Provides...),
|
||||
Requires: append([]string(nil), spec.Requires...),
|
||||
ReferenceSlots: contracts.CloneReferenceSlots(spec.ReferenceSlots),
|
||||
Key: spec.Key,
|
||||
Stage: spec.Stage,
|
||||
ExecutionClass: spec.ExecutionClass,
|
||||
CorrectionProtocol: spec.CorrectionProtocol,
|
||||
ArtifactKind: spec.ArtifactKind,
|
||||
Provides: append([]string(nil), spec.Provides...),
|
||||
Requires: append([]string(nil), spec.Requires...),
|
||||
ReferenceSlots: contracts.CloneReferenceSlots(spec.ReferenceSlots),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -93,6 +97,19 @@ func validateModuleSpec(kind string, expectedStage ModuleStage, spec ModuleSpec)
|
||||
if spec.ExecutionClass != contracts.ExecutionClassDeterministic && spec.ExecutionClass != contracts.ExecutionClassLLMBacked {
|
||||
return fmt.Errorf("%s %q has unsupported execution class %q", kind, spec.Key, spec.ExecutionClass)
|
||||
}
|
||||
if spec.CorrectionProtocol != "" {
|
||||
if err := spec.CorrectionProtocol.Validate(); err != nil {
|
||||
return fmt.Errorf("%s %q correction protocol: %w", kind, spec.Key, err)
|
||||
}
|
||||
if spec.ExecutionClass != contracts.ExecutionClassLLMBacked {
|
||||
return fmt.Errorf("%s %q correction protocol requires an LLM-backed execution class", kind, spec.Key)
|
||||
}
|
||||
switch spec.Stage {
|
||||
case StageChunk, StageExtract, StageMerge, StageNormalize:
|
||||
default:
|
||||
return fmt.Errorf("%s %q correction protocol is not supported for %q stage", kind, spec.Key, spec.Stage)
|
||||
}
|
||||
}
|
||||
if spec.ArtifactKind != "" && spec.Stage != StageExtract && spec.Stage != StageMerge && spec.Stage != StageNormalize {
|
||||
return fmt.Errorf("%s %q must not declare an artifact kind", kind, spec.Key)
|
||||
}
|
||||
|
||||
@@ -32,14 +32,70 @@ func TestValidateModuleSpecRequiresSupportedExecutionClass(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloneModuleSpecPreservesExecutionClass(t *testing.T) {
|
||||
spec := normalizeModuleSpec(ModuleSpec{Key: " module ", Stage: StageChunk, ExecutionClass: contracts.ExecutionClassLLMBacked})
|
||||
func TestCloneModuleSpecPreservesCorrectionProtocol(t *testing.T) {
|
||||
spec := normalizeModuleSpec(ModuleSpec{
|
||||
Key: " module ",
|
||||
Stage: StageChunk,
|
||||
ExecutionClass: contracts.ExecutionClassLLMBacked,
|
||||
CorrectionProtocol: " single_response_v1 ",
|
||||
})
|
||||
if spec.CorrectionProtocol != contracts.CorrectionProtocolSingleResponseV1 {
|
||||
t.Fatalf("normalized CorrectionProtocol = %q, want %q", spec.CorrectionProtocol, contracts.CorrectionProtocolSingleResponseV1)
|
||||
}
|
||||
cloned := cloneModuleSpec(spec)
|
||||
if !reflect.DeepEqual(cloned, spec) {
|
||||
t.Fatalf("cloneModuleSpec() = %#v, want %#v", cloned, spec)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateModuleSpecCorrectionProtocolEligibility(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
stage ModuleStage
|
||||
class contracts.ExecutionClass
|
||||
protocol contracts.CorrectionProtocol
|
||||
want string
|
||||
}{
|
||||
{name: "LLM chunk", stage: StageChunk, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1},
|
||||
{name: "LLM extract", stage: StageExtract, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1},
|
||||
{name: "LLM merge", stage: StageMerge, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1},
|
||||
{name: "LLM normalize", stage: StageNormalize, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1},
|
||||
{name: "deterministic chunk", stage: StageChunk, class: contracts.ExecutionClassDeterministic, protocol: contracts.CorrectionProtocolSingleResponseV1, want: "LLM-backed"},
|
||||
{name: "input", stage: StageInput, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1, want: "not supported"},
|
||||
{name: "output", stage: StageOutput, class: contracts.ExecutionClassLLMBacked, protocol: contracts.CorrectionProtocolSingleResponseV1, want: "not supported"},
|
||||
{name: "unknown protocol", stage: StageChunk, class: contracts.ExecutionClassLLMBacked, protocol: "unsupported", want: "unsupported correction protocol"},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
spec := normalizeModuleSpec(ModuleSpec{
|
||||
Key: "module",
|
||||
Stage: test.stage,
|
||||
ExecutionClass: test.class,
|
||||
CorrectionProtocol: test.protocol,
|
||||
})
|
||||
err := validateModuleSpec("module", test.stage, spec)
|
||||
if test.want == "" && err != nil {
|
||||
t.Fatalf("validateModuleSpec() error = %v, want nil", err)
|
||||
}
|
||||
if test.want != "" && (err == nil || !strings.Contains(err.Error(), test.want)) {
|
||||
t.Fatalf("validateModuleSpec() error = %v, want %q", err, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeValidatorSpecRejectsCorrectionProtocol(t *testing.T) {
|
||||
_, err := normalizeValidatorSpec(ValidatorSpec{
|
||||
Key: "validator",
|
||||
ExecutionClass: contracts.ExecutionClassLLMBacked,
|
||||
CorrectionProtocol: contracts.CorrectionProtocolSingleResponseV1,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "correction protocol") {
|
||||
t.Fatalf("normalizeValidatorSpec() error = %v, want correction protocol error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateModuleSpecAllowsReferenceSlotsForEligibleStages(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -13,10 +13,11 @@ import (
|
||||
// resolved pipeline. Its implementation values are private so execution cannot
|
||||
// replace or reconfigure them after preparation.
|
||||
type PreparedPipeline struct {
|
||||
Input ModuleBinding
|
||||
Chunk ModuleBinding
|
||||
Steps []PreparedPipelineStep
|
||||
Output ModuleBinding
|
||||
Input ModuleBinding
|
||||
Chunk ModuleBinding
|
||||
ChunkCorrectionProtocol contracts.CorrectionProtocol
|
||||
Steps []PreparedPipelineStep
|
||||
Output ModuleBinding
|
||||
|
||||
resolved ResolvedPipeline
|
||||
dependencies ModuleDependencies
|
||||
@@ -36,7 +37,10 @@ type PreparedPipelineStep struct {
|
||||
}
|
||||
|
||||
type PreparedArtifactLane struct {
|
||||
Resolved ResolvedArtifactLane
|
||||
Resolved ResolvedArtifactLane
|
||||
ExtractCorrectionProtocol contracts.CorrectionProtocol
|
||||
MergeCorrectionProtocol contracts.CorrectionProtocol
|
||||
NormalizeCorrectionProtocol contracts.CorrectionProtocol
|
||||
}
|
||||
|
||||
type preparedLaneExecutor struct {
|
||||
@@ -92,12 +96,13 @@ func Prepare(resolved ResolvedPipeline, registries Registries, deps ModuleDepend
|
||||
}
|
||||
stable := cloneResolvedPipeline(resolved)
|
||||
prepared := &PreparedPipeline{
|
||||
Input: cloneModuleBinding(stable.Input),
|
||||
Chunk: cloneModuleBinding(stable.Chunk),
|
||||
Output: cloneModuleBinding(stable.Output),
|
||||
resolved: stable,
|
||||
dependencies: deps,
|
||||
artifactCodecs: registries.ArtifactCodecs,
|
||||
Input: cloneModuleBinding(stable.Input),
|
||||
Chunk: cloneModuleBinding(stable.Chunk),
|
||||
ChunkCorrectionProtocol: stable.ChunkCorrectionProtocol,
|
||||
Output: cloneModuleBinding(stable.Output),
|
||||
resolved: stable,
|
||||
dependencies: deps,
|
||||
artifactCodecs: registries.ArtifactCodecs,
|
||||
}
|
||||
request := func(binding ModuleBinding, references contracts.ReferenceSet) BuildRequest {
|
||||
return BuildRequest{Dependencies: deps, Options: binding.Options, References: references}
|
||||
@@ -131,7 +136,12 @@ func Prepare(resolved ResolvedPipeline, registries Registries, deps ModuleDepend
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
preparedStep.ArtifactLanes = append(preparedStep.ArtifactLanes, PreparedArtifactLane{Resolved: cloneResolvedArtifactLane(lane)})
|
||||
preparedStep.ArtifactLanes = append(preparedStep.ArtifactLanes, PreparedArtifactLane{
|
||||
Resolved: cloneResolvedArtifactLane(lane),
|
||||
ExtractCorrectionProtocol: lane.ExtractCorrectionProtocol,
|
||||
MergeCorrectionProtocol: lane.MergeCorrectionProtocol,
|
||||
NormalizeCorrectionProtocol: lane.NormalizeCorrectionProtocol,
|
||||
})
|
||||
preparedStep.lanes = append(preparedStep.lanes, executor)
|
||||
}
|
||||
prepared.Steps[stepIndex] = preparedStep
|
||||
|
||||
@@ -204,26 +204,29 @@ type ResolvedReferenceTarget struct {
|
||||
}
|
||||
|
||||
type ResolvedArtifactLane struct {
|
||||
StepID string
|
||||
ID string
|
||||
ArtifactKind contracts.ArtifactKind `json:"artifact_kind,omitempty"`
|
||||
ArtifactSchemaID string `json:"artifact_schema_id,omitempty"`
|
||||
ArtifactSchemaName string `json:"artifact_schema_name,omitempty"`
|
||||
ArtifactSchemaVersion string `json:"artifact_schema_version,omitempty"`
|
||||
ArtifactSchemaDigest string `json:"artifact_schema_digest,omitempty"`
|
||||
Extract ModuleBinding
|
||||
ExtractExecutionClass contracts.ExecutionClass `json:"extract_execution_class"`
|
||||
ExtractValidationPolicy ValidationPolicy `json:"extract_validation_policy"`
|
||||
Merge ModuleBinding
|
||||
MergeExecutionClass contracts.ExecutionClass `json:"merge_execution_class"`
|
||||
MergeValidationPolicy ValidationPolicy `json:"merge_validation_policy"`
|
||||
Normalize ModuleBinding
|
||||
NormalizeExecutionClass contracts.ExecutionClass `json:"normalize_execution_class"`
|
||||
NormalizeValidationPolicy ValidationPolicy `json:"normalize_validation_policy"`
|
||||
Validators []ModuleBinding
|
||||
ExtractReferences ResolvedReferenceTarget `json:"extract_references"`
|
||||
MergeReferences ResolvedReferenceTarget `json:"merge_references"`
|
||||
NormalizeReferences ResolvedReferenceTarget `json:"normalize_references"`
|
||||
StepID string
|
||||
ID string
|
||||
ArtifactKind contracts.ArtifactKind `json:"artifact_kind,omitempty"`
|
||||
ArtifactSchemaID string `json:"artifact_schema_id,omitempty"`
|
||||
ArtifactSchemaName string `json:"artifact_schema_name,omitempty"`
|
||||
ArtifactSchemaVersion string `json:"artifact_schema_version,omitempty"`
|
||||
ArtifactSchemaDigest string `json:"artifact_schema_digest,omitempty"`
|
||||
Extract ModuleBinding
|
||||
ExtractExecutionClass contracts.ExecutionClass `json:"extract_execution_class"`
|
||||
ExtractCorrectionProtocol contracts.CorrectionProtocol `json:"extract_correction_protocol,omitempty"`
|
||||
ExtractValidationPolicy ValidationPolicy `json:"extract_validation_policy"`
|
||||
Merge ModuleBinding
|
||||
MergeExecutionClass contracts.ExecutionClass `json:"merge_execution_class"`
|
||||
MergeCorrectionProtocol contracts.CorrectionProtocol `json:"merge_correction_protocol,omitempty"`
|
||||
MergeValidationPolicy ValidationPolicy `json:"merge_validation_policy"`
|
||||
Normalize ModuleBinding
|
||||
NormalizeExecutionClass contracts.ExecutionClass `json:"normalize_execution_class"`
|
||||
NormalizeCorrectionProtocol contracts.CorrectionProtocol `json:"normalize_correction_protocol,omitempty"`
|
||||
NormalizeValidationPolicy ValidationPolicy `json:"normalize_validation_policy"`
|
||||
Validators []ModuleBinding
|
||||
ExtractReferences ResolvedReferenceTarget `json:"extract_references"`
|
||||
MergeReferences ResolvedReferenceTarget `json:"merge_references"`
|
||||
NormalizeReferences ResolvedReferenceTarget `json:"normalize_references"`
|
||||
}
|
||||
|
||||
type ResolvedPipelineStep struct {
|
||||
@@ -252,9 +255,10 @@ type ResolvedPipeline struct {
|
||||
Input ModuleBinding
|
||||
InputExecutionClass contracts.ExecutionClass `json:"input_execution_class"`
|
||||
Chunk ModuleBinding
|
||||
ChunkExecutionClass contracts.ExecutionClass `json:"chunk_execution_class"`
|
||||
ChunkValidationPolicy ValidationPolicy `json:"chunk_validation_policy"`
|
||||
ChunkReferences ResolvedReferenceTarget `json:"chunk_references"`
|
||||
ChunkExecutionClass contracts.ExecutionClass `json:"chunk_execution_class"`
|
||||
ChunkCorrectionProtocol contracts.CorrectionProtocol `json:"chunk_correction_protocol,omitempty"`
|
||||
ChunkValidationPolicy ValidationPolicy `json:"chunk_validation_policy"`
|
||||
ChunkReferences ResolvedReferenceTarget `json:"chunk_references"`
|
||||
Steps []ResolvedPipelineStep
|
||||
ValidatorChains []ResolvedValidatorChain `json:"validator_chains"`
|
||||
Output ModuleBinding
|
||||
@@ -431,6 +435,7 @@ func ResolvePipeline(profile PipelineProfile, options ResolveOptions, catalog Mo
|
||||
InputExecutionClass: inputModuleSpec.ExecutionClass,
|
||||
Chunk: chunk,
|
||||
ChunkExecutionClass: chunkSpec.ExecutionClass,
|
||||
ChunkCorrectionProtocol: chunkSpec.CorrectionProtocol,
|
||||
ChunkReferences: referenceTarget(StageChunk, "", chunk.Module, chunkReferences),
|
||||
Output: output,
|
||||
}
|
||||
@@ -615,6 +620,7 @@ func resolveArtifactLane(
|
||||
lane.ExtractReferences = referenceTarget(StageExtract, laneID, lane.Extract.Module, references)
|
||||
lane.ExtractReferences.StepID = strings.TrimSpace(stepID)
|
||||
lane.ExtractExecutionClass = extractSpec.ExecutionClass
|
||||
lane.ExtractCorrectionProtocol = extractSpec.CorrectionProtocol
|
||||
capabilities.add(extractSpec.Provides...)
|
||||
|
||||
mergeSpec, err := mergerSpecForArtifact(catalog, lane.Merge.Module, lane.ArtifactKind, artifactType)
|
||||
@@ -641,6 +647,7 @@ func resolveArtifactLane(
|
||||
lane.MergeReferences = referenceTarget(StageMerge, laneID, lane.Merge.Module, mergeReferences)
|
||||
lane.MergeReferences.StepID = strings.TrimSpace(stepID)
|
||||
lane.MergeExecutionClass = mergeSpec.ExecutionClass
|
||||
lane.MergeCorrectionProtocol = mergeSpec.CorrectionProtocol
|
||||
capabilities.add(mergeSpec.Provides...)
|
||||
|
||||
normalizeSpec, err := normalizerSpecForArtifact(catalog, lane.Normalize.Module, lane.ArtifactKind, artifactType)
|
||||
@@ -667,6 +674,7 @@ func resolveArtifactLane(
|
||||
lane.NormalizeReferences = referenceTarget(StageNormalize, laneID, lane.Normalize.Module, normalizeReferences)
|
||||
lane.NormalizeReferences.StepID = strings.TrimSpace(stepID)
|
||||
lane.NormalizeExecutionClass = normalizeSpec.ExecutionClass
|
||||
lane.NormalizeCorrectionProtocol = normalizeSpec.CorrectionProtocol
|
||||
capabilities.add(normalizeSpec.Provides...)
|
||||
|
||||
if len(lane.Validators) > 0 {
|
||||
@@ -928,6 +936,9 @@ func resolveValidatorChain(pipelineID string, laneID string, stage ModuleStage,
|
||||
if err != nil {
|
||||
return ResolvedValidatorChain{}, fmt.Errorf("pipeline %q %s validator chain for module %q: %w", pipelineID, stage, chain.ModuleKey, err)
|
||||
}
|
||||
if validator.Retries > 0 && spec.ExecutionClass != contracts.ExecutionClassLLMBacked {
|
||||
return ResolvedValidatorChain{}, fmt.Errorf("pipeline %q %s validator %q retries require an LLM-backed execution class", pipelineID, stage, validator.Module)
|
||||
}
|
||||
chain.Validators = append(chain.Validators, ResolvedValidator{
|
||||
Binding: cloneModuleBinding(validator),
|
||||
ExecutionClass: spec.ExecutionClass,
|
||||
@@ -1637,6 +1648,7 @@ func resolvedPipelineDigest(resolved ResolvedPipeline) (string, error) {
|
||||
InputExecutionClass contracts.ExecutionClass
|
||||
Chunk ModuleBinding
|
||||
ChunkExecutionClass contracts.ExecutionClass
|
||||
ChunkCorrectionProtocol contracts.CorrectionProtocol
|
||||
ChunkValidationPolicy ValidationPolicy
|
||||
ChunkReferences ResolvedReferenceTarget
|
||||
Steps []ResolvedPipelineStep
|
||||
@@ -1651,6 +1663,7 @@ func resolvedPipelineDigest(resolved ResolvedPipeline) (string, error) {
|
||||
InputExecutionClass: resolved.InputExecutionClass,
|
||||
Chunk: resolved.Chunk,
|
||||
ChunkExecutionClass: resolved.ChunkExecutionClass,
|
||||
ChunkCorrectionProtocol: resolved.ChunkCorrectionProtocol,
|
||||
ChunkReferences: resolved.ChunkReferences,
|
||||
Steps: resolved.Steps,
|
||||
ValidatorChains: resolved.ValidatorChains,
|
||||
|
||||
@@ -126,6 +126,109 @@ func TestModuleCatalogExecutionClassLooksUpRegisteredMetadata(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePipelineCarriesCorrectionProtocolsIntoPreparedMetadata(t *testing.T) {
|
||||
correction := contracts.CorrectionProtocolSingleResponseV1
|
||||
catalog := newProfileCatalogWithOverrides(t,
|
||||
ModuleSpec{Key: "llm-input", Stage: StageInput, ExecutionClass: contracts.ExecutionClassLLMBacked, Provides: []string{"source"}},
|
||||
ModuleSpec{Key: "llm-chunk", Stage: StageChunk, ExecutionClass: contracts.ExecutionClassLLMBacked, CorrectionProtocol: correction, Requires: []string{"source"}, Provides: []string{"chunk"}},
|
||||
ModuleSpec{Key: "llm-extractor", Stage: StageExtract, ExecutionClass: contracts.ExecutionClassLLMBacked, CorrectionProtocol: correction, ArtifactKind: "test/notes", Requires: []string{"chunk"}, Provides: []string{"candidate"}},
|
||||
ModuleSpec{Key: "llm-merge", Stage: StageMerge, ExecutionClass: contracts.ExecutionClassLLMBacked, CorrectionProtocol: correction, ArtifactKind: "test/notes", Requires: []string{"candidate"}, Provides: []string{"merged"}},
|
||||
ModuleSpec{Key: "llm-normalize", Stage: StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, CorrectionProtocol: correction, ArtifactKind: "test/notes", Requires: []string{"merged"}, Provides: []string{"normalized"}},
|
||||
ModuleSpec{Key: "llm-output", Stage: StageOutput, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: []string{"normalized"}, Provides: []string{"encoded"}},
|
||||
)
|
||||
resolved, err := ResolvePipeline(llmProfilePipeline(), ResolveOptions{}, catalog)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||
}
|
||||
if resolved.ChunkCorrectionProtocol != correction {
|
||||
t.Fatalf("ChunkCorrectionProtocol = %q, want %q", resolved.ChunkCorrectionProtocol, correction)
|
||||
}
|
||||
lane := resolved.Steps[0].ArtifactLanes[0]
|
||||
if lane.ExtractCorrectionProtocol != correction || lane.MergeCorrectionProtocol != correction || lane.NormalizeCorrectionProtocol != correction {
|
||||
t.Fatalf("lane correction protocols = %q/%q/%q, want %q", lane.ExtractCorrectionProtocol, lane.MergeCorrectionProtocol, lane.NormalizeCorrectionProtocol, correction)
|
||||
}
|
||||
withoutCapability := cloneResolvedPipeline(resolved)
|
||||
withoutCapability.ChunkCorrectionProtocol = ""
|
||||
withoutCapability.Steps[0].ArtifactLanes[0].ExtractCorrectionProtocol = ""
|
||||
withoutCapability.Steps[0].ArtifactLanes[0].MergeCorrectionProtocol = ""
|
||||
withoutCapability.Steps[0].ArtifactLanes[0].NormalizeCorrectionProtocol = ""
|
||||
withoutCapabilityDigest, err := resolvedPipelineDigest(withoutCapability)
|
||||
if err != nil {
|
||||
t.Fatalf("resolvedPipelineDigest() error = %v", err)
|
||||
}
|
||||
if withoutCapabilityDigest == resolved.Digest {
|
||||
t.Fatalf("resolved digest = %q with and without correction capability, want changed", resolved.Digest)
|
||||
}
|
||||
|
||||
prepared, err := Prepare(resolved, registriesFromModuleCatalog(catalog), ModuleDependencies{})
|
||||
if err != nil {
|
||||
t.Fatalf("Prepare() error = %v, want nil", err)
|
||||
}
|
||||
if prepared.ChunkCorrectionProtocol != correction {
|
||||
t.Fatalf("prepared ChunkCorrectionProtocol = %q, want %q", prepared.ChunkCorrectionProtocol, correction)
|
||||
}
|
||||
preparedLane := prepared.Steps[0].ArtifactLanes[0]
|
||||
if preparedLane.ExtractCorrectionProtocol != correction || preparedLane.MergeCorrectionProtocol != correction || preparedLane.NormalizeCorrectionProtocol != correction {
|
||||
t.Fatalf("prepared lane correction protocols = %q/%q/%q, want %q", preparedLane.ExtractCorrectionProtocol, preparedLane.MergeCorrectionProtocol, preparedLane.NormalizeCorrectionProtocol, correction)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareAllowsLLMProducerWithoutCorrectionCapability(t *testing.T) {
|
||||
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"}},
|
||||
)
|
||||
profile := llmProfilePipeline()
|
||||
profile.Chunk.Retries = 1
|
||||
resolved, err := ResolvePipeline(profile, ResolveOptions{}, catalog)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||
}
|
||||
if resolved.ChunkCorrectionProtocol != "" {
|
||||
t.Fatalf("ChunkCorrectionProtocol = %q, want empty unsupported value", resolved.ChunkCorrectionProtocol)
|
||||
}
|
||||
if _, err := Prepare(resolved, registriesFromModuleCatalog(catalog), ModuleDependencies{}); err != nil {
|
||||
t.Fatalf("Prepare() error = %v, want nil", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePipelineValidatorRetriesRequireLLMBackedValidator(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
class contracts.ExecutionClass
|
||||
want string
|
||||
}{
|
||||
{name: "deterministic validator", class: contracts.ExecutionClassDeterministic, want: "retries require an LLM-backed"},
|
||||
{name: "LLM validator", class: contracts.ExecutionClassLLMBacked},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
catalog := newProfileCatalog(t)
|
||||
registerProfileValidatorSpec(t, catalog, ValidatorSpec{Key: "retrying-validator", ExecutionClass: test.class})
|
||||
profile := baselineProfile()
|
||||
lane := profile.Artifacts["events"]
|
||||
lane.Extract.Validators = ValidatorOverride{Set: true, Validators: []ModuleBinding{{Module: "retrying-validator", Retries: 1}}}
|
||||
profile.Artifacts["events"] = lane
|
||||
resolved, err := ResolvePipeline(profile, ResolveOptions{}, catalog)
|
||||
if test.want != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), test.want) {
|
||||
t.Fatalf("ResolvePipeline() error = %v, want %q", err, test.want)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("ResolvePipeline() error = %v, want nil", err)
|
||||
}
|
||||
if got := resolved.ValidatorChains[1].Validators[0].Binding.Retries; got != 1 {
|
||||
t.Fatalf("validator retries = %d, want 1", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePipelineAppliesDefaults(t *testing.T) {
|
||||
resolved, err := ResolvePipeline(PipelineProfile{
|
||||
ID: "defaulted",
|
||||
|
||||
@@ -11,8 +11,9 @@ import (
|
||||
)
|
||||
|
||||
type ValidatorSpec struct {
|
||||
Key string `json:"key"`
|
||||
ExecutionClass contracts.ExecutionClass `json:"execution_class"`
|
||||
Key string `json:"key"`
|
||||
ExecutionClass contracts.ExecutionClass `json:"execution_class"`
|
||||
CorrectionProtocol contracts.CorrectionProtocol `json:"correction_protocol,omitempty"`
|
||||
}
|
||||
|
||||
type SerializedValidatorSpec struct {
|
||||
@@ -321,7 +322,11 @@ func (r *ValidatorRegistry) RegisteredKeys() []string {
|
||||
}
|
||||
|
||||
func normalizeValidatorSpec(spec ValidatorSpec) (ValidatorSpec, error) {
|
||||
normalized := ValidatorSpec{Key: strings.TrimSpace(spec.Key), ExecutionClass: spec.ExecutionClass}
|
||||
normalized := ValidatorSpec{
|
||||
Key: strings.TrimSpace(spec.Key),
|
||||
ExecutionClass: spec.ExecutionClass,
|
||||
CorrectionProtocol: contracts.CorrectionProtocol(strings.TrimSpace(string(spec.CorrectionProtocol))),
|
||||
}
|
||||
if normalized.Key == "" {
|
||||
return ValidatorSpec{}, fmt.Errorf("validator key must not be empty")
|
||||
}
|
||||
@@ -330,6 +335,9 @@ func normalizeValidatorSpec(spec ValidatorSpec) (ValidatorSpec, error) {
|
||||
default:
|
||||
return ValidatorSpec{}, fmt.Errorf("validator %q execution class %q is not supported", normalized.Key, normalized.ExecutionClass)
|
||||
}
|
||||
if normalized.CorrectionProtocol != "" {
|
||||
return ValidatorSpec{}, fmt.Errorf("validator %q correction protocol is not supported", normalized.Key)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user