Wire resolved validator chains into runner

This commit is contained in:
2026-07-07 21:27:28 +00:00
parent 666b4bf801
commit 5ef027b6f0
4 changed files with 320 additions and 138 deletions

View File

@@ -763,6 +763,110 @@ func TestRunPassesNormalizeReferencesToNormalizerRequest(t *testing.T) {
}
}
func TestRunPassesValidationRequestContextToValidators(t *testing.T) {
modules := defaultRunnerModules()
chunkValidator := &runnerChainValidator{name: "chain-chunk"}
extractValidator := &runnerChainValidator{name: "chain-extract", executionClass: contracts.ExecutionClassLLMBacked}
mergeValidator := &runnerChainValidator{name: "chain-merge"}
normalizeValidator := &runnerChainValidator{name: "chain-normalize"}
modules.validators[chunkValidator.name] = chunkValidator
modules.validators[extractValidator.name] = extractValidator
modules.validators[mergeValidator.name] = mergeValidator
modules.validators[normalizeValidator.name] = normalizeValidator
pipeline := resolvedPipeline()
pipeline.ChunkReferences.ReferenceSet = testReferenceSet("scene_guide", "chunk reference text")
pipeline.ArtifactLanes[0].ExtractReferences.ReferenceSet = testReferenceSet("roster", "extract reference text")
pipeline.ArtifactLanes[0].MergeReferences.ReferenceSet = testReferenceSet("merge_notes", "merge reference text")
pipeline.ArtifactLanes[0].NormalizeReferences.ReferenceSet = testReferenceSet("normalization_notes", "normalize reference text")
setResolvedValidatorChain(t, &pipeline, StageChunk, "", "chunk", resolvedValidatorForTest(chunkValidator))
setResolvedValidatorChain(t, &pipeline, StageExtract, "alpha", "extract-alpha", ResolvedValidator{
Binding: ModuleBinding{Module: extractValidator.name, LLMProfile: "validator-profile", Options: map[string]any{"strict": true}},
ExecutionClass: extractValidator.ExecutionClass(),
})
setResolvedValidatorChain(t, &pipeline, StageMerge, "alpha", "merge", resolvedValidatorForTest(mergeValidator))
setResolvedValidatorChain(t, &pipeline, StageNormalize, "alpha", "normalize", resolvedValidatorForTest(normalizeValidator))
rawInput := []byte("{\"source\":\"exact bytes\"}")
llmClient := fakeLLMClient{}
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{
Pipeline: pipeline,
Path: "session.json",
RawInput: rawInput,
LLMClient: llmClient,
SessionID: "session-123",
Metadata: map[string]any{"request": "test"},
})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
if len(chunkValidator.requests) != 2 {
t.Fatalf("chunk validator requests = %d, want one per chunk", len(chunkValidator.requests))
}
chunkReq := chunkValidator.requests[0]
if chunkReq.Stage != string(StageChunk) || chunkReq.ModuleKey != "chunk" || chunkReq.SourceID != "source-1" || chunkReq.SessionID != "session-123" {
t.Fatalf("chunk validation request = %#v, want stage/module/source/session provenance", chunkReq)
}
if chunkReq.LLMClient == nil || string(chunkReq.SourceInput.Content) != string(rawInput) {
t.Fatalf("chunk validation source/client = %#v, want full source input and LLM client", chunkReq.SourceInput)
}
if chunkReq.Chunk == nil || chunkReq.Chunk.ID != "chunk-0" || len(chunkReq.Chunks) != 2 || string(chunkReq.Payload.Content) != string(chunkReq.Chunk.Content) {
t.Fatalf("chunk validation chunk fields = %#v chunks=%#v payload=%s, want chunk payload and all chunks", chunkReq.Chunk, chunkReq.Chunks, chunkReq.Payload.Content)
}
if item := chunkReq.References.Slots["scene_guide"].Items[0]; string(item.Content) != "chunk reference text" {
t.Fatalf("chunk validation references = %#v, want chunk references", chunkReq.References)
}
if len(extractValidator.requests) != 2 {
t.Fatalf("extract validator requests = %d, want one per chunk", len(extractValidator.requests))
}
extractReq := extractValidator.requests[0]
if extractReq.Stage != string(StageExtract) || extractReq.LaneID != "alpha" || extractReq.ModuleKey != "extract-alpha" || extractReq.ChunkID != "chunk-0" || extractReq.ChunkIndex != 0 {
t.Fatalf("extract validation request = %#v, want extract provenance", extractReq)
}
if extractReq.LLMProfile != "validator-profile" || extractReq.Options["strict"] != true || extractReq.Metadata["request"] != "test" {
t.Fatalf("extract validator binding fields = profile %q options %#v metadata %#v", extractReq.LLMProfile, extractReq.Options, extractReq.Metadata)
}
if extractReq.Chunk == nil || string(extractReq.SourceInput.Content) != string(extractReq.Chunk.Content) {
t.Fatalf("extract source input = %#v chunk=%#v, want chunk material", extractReq.SourceInput, extractReq.Chunk)
}
if item := extractReq.References.Slots["roster"].Items[0]; string(item.Content) != "extract reference text" {
t.Fatalf("extract validation references = %#v, want extract references", extractReq.References)
}
if len(mergeValidator.requests) != 1 {
t.Fatalf("merge validator requests = %d, want one", len(mergeValidator.requests))
}
mergeReq := mergeValidator.requests[0]
if mergeReq.Stage != string(StageMerge) || mergeReq.LaneID != "alpha" || len(mergeReq.ExtractOutputs) != 2 {
t.Fatalf("merge validation request = %#v, want lane and extract outputs", mergeReq)
}
if mergeReq.ExtractOutputs[0].ChunkID != "chunk-0" || string(mergeReq.SourceInput.Content) != string(rawInput) {
t.Fatalf("merge validation upstream/source = %#v source=%#v, want ordered extracts and source input", mergeReq.ExtractOutputs, mergeReq.SourceInput)
}
if item := mergeReq.References.Slots["merge_notes"].Items[0]; string(item.Content) != "merge reference text" {
t.Fatalf("merge validation references = %#v, want merge references", mergeReq.References)
}
if len(normalizeValidator.requests) != 1 {
t.Fatalf("normalize validator requests = %d, want one", len(normalizeValidator.requests))
}
normalizeReq := normalizeValidator.requests[0]
if normalizeReq.Stage != string(StageNormalize) || normalizeReq.LaneID != "alpha" || string(normalizeReq.MergeOutput.Payload.Content) != `{"merged":true}` {
t.Fatalf("normalize validation request = %#v, want merge output context", normalizeReq)
}
if item := normalizeReq.References.Slots["normalization_notes"].Items[0]; string(item.Content) != "normalize reference text" {
t.Fatalf("normalize validation references = %#v, want normalize references", normalizeReq.References)
}
chunkReq.Payload.Content[0] = 'X'
chunkReq.Chunks[0].Content[0] = 'Y'
if got := string(modules.chunker.chunks[0].Content); got != `{"units":[{"id":1,"kind":"unit","text":"Source unit."}]}` {
t.Fatalf("validator request mutated original chunk content: %q", got)
}
}
func TestRunAllowsNilLLMClientWhenModulesDoNotUseIt(t *testing.T) {
_, err := New(newRunnerRegistries(t, defaultRunnerModules())).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
if err != nil {
@@ -889,9 +993,10 @@ func TestRunOmitsRejectedExtractOutputsFromMerge(t *testing.T) {
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-extract", approved: []bool{false, true}, reason: "bad_extract", message: "extract rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageExtract, "extract-alpha", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageExtract, "alpha", "extract-alpha", resolvedValidatorForTest(validator))
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
@@ -912,9 +1017,10 @@ func TestRunOmitsLaneWithNoAcceptedExtractOutputs(t *testing.T) {
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-extract", approved: []bool{false}, reason: "bad_extract", message: "extract rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageExtract, "extract-alpha", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageExtract, "alpha", "extract-alpha", resolvedValidatorForTest(validator))
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
@@ -937,9 +1043,10 @@ func TestRunRejectedMergePreventsNormalizeForLane(t *testing.T) {
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-merge", approved: []bool{false}, reason: "bad_merge", message: "merge rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageMerge, "merge", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageMerge, "alpha", "merge", resolvedValidatorForTest(validator))
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
@@ -959,9 +1066,10 @@ func TestRunRejectedNormalizePreventsOutputForLane(t *testing.T) {
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-normalize", approved: []bool{false}, reason: "bad_normalize", message: "normalize rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageNormalize, "normalize", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageNormalize, "alpha", "normalize", resolvedValidatorForTest(validator))
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
@@ -1005,8 +1113,8 @@ func TestRunRetriesSameModuleInputAfterValidatorRejection(t *testing.T) {
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-extract", approved: []bool{false, true, true}, reason: "bad_extract", message: "extract rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageExtract, "extract-alpha", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageExtract, "alpha", "extract-alpha", resolvedValidatorForTest(validator))
pipeline.ArtifactLanes[0].Extract.Retries = 1
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
@@ -1030,8 +1138,8 @@ func TestRunStopsRetryAfterConfiguredAttemptsAndRecordsAttemptCount(t *testing.T
modules := defaultRunnerModules()
validator := &runnerChainValidator{name: "chain-extract", approved: []bool{false}, reason: "bad_extract", message: "extract rejected"}
modules.validators[validator.name] = validator
modules.validatorChains = validatorChainRegistry(t, StageExtract, "extract-alpha", validator)
pipeline := resolvedPipeline()
setResolvedValidatorChain(t, &pipeline, StageExtract, "alpha", "extract-alpha", resolvedValidatorForTest(validator))
pipeline.ArtifactLanes[0].Extract.Retries = 1
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
@@ -1509,7 +1617,6 @@ type runnerModules struct {
mergers map[string]*runnerMerger
normalizers map[string]*runnerNormalizer
validators map[string]contracts.Validator
validatorChains *ValidatorChainRegistry
output *runnerOutputEncoder
inputBuildErr error
chunkerBuildErr error
@@ -1554,7 +1661,7 @@ func newRunnerRegistries(t *testing.T, modules *runnerModules) Registries {
Mergers: NewMergerRegistry(),
Normalizers: NewNormalizerRegistry(),
Validators: NewValidatorRegistry(),
ValidatorChains: modules.validatorChains,
ValidatorChains: NewValidatorChainRegistry(),
Outputs: NewOutputEncoderRegistry(),
}
if err := registries.Inputs.Register("input", func() (contracts.InputAdapter, error) {
@@ -1593,7 +1700,8 @@ func newRunnerRegistries(t *testing.T, modules *runnerModules) Registries {
}
for key, validator := range modules.validators {
validator := validator
if err := registries.Validators.Register(key, func() (contracts.Validator, error) { return validator, nil }); err != nil {
spec := ValidatorSpec{Key: key, ExecutionClass: validator.ExecutionClass()}
if err := registries.Validators.RegisterWithSpec(spec, func() (contracts.Validator, error) { return validator, nil }); err != nil {
t.Fatalf("register validator %q: %v", key, err)
}
}
@@ -1811,26 +1919,28 @@ func (normalizer *runnerNormalizer) Normalize(ctx context.Context, req contracts
}
type runnerValidator struct {
name string
approved []bool
reason string
message string
warnings []contracts.Warning
err error
order *[]string
calls int
requests []contracts.ValidationRequest
name string
executionClass contracts.ExecutionClass
approved []bool
reason string
message string
warnings []contracts.Warning
err error
order *[]string
calls int
requests []contracts.ValidationRequest
}
type runnerChainValidator struct {
name string
approved []bool
reason string
message string
warnings []contracts.Warning
err error
calls int
requests []contracts.ValidationRequest
name string
executionClass contracts.ExecutionClass
approved []bool
reason string
message string
warnings []contracts.Warning
err error
calls int
requests []contracts.ValidationRequest
}
func (validator *runnerChainValidator) Name() string {
@@ -1838,6 +1948,9 @@ func (validator *runnerChainValidator) Name() string {
}
func (validator *runnerChainValidator) ExecutionClass() contracts.ExecutionClass {
if validator.executionClass != "" {
return validator.executionClass
}
return contracts.ExecutionClassDeterministic
}
@@ -1868,6 +1981,9 @@ func (validator *runnerValidator) Name() string {
}
func (validator *runnerValidator) ExecutionClass() contracts.ExecutionClass {
if validator.executionClass != "" {
return validator.executionClass
}
return contracts.ExecutionClassDeterministic
}
@@ -2052,16 +2168,31 @@ func assertRunError(t *testing.T, err error, want string) {
}
}
func validatorChainRegistry(t *testing.T, stage ModuleStage, module string, validators ...contracts.Validator) *ValidatorChainRegistry {
func resolvedValidatorForTest(validator contracts.Validator) ResolvedValidator {
return ResolvedValidator{
Binding: Binding(validator.Name()),
ExecutionClass: validator.ExecutionClass(),
}
}
func setResolvedValidatorChain(t *testing.T, resolved *ResolvedPipeline, stage ModuleStage, laneID string, module string, validators ...ResolvedValidator) {
t.Helper()
registry := NewValidatorChainRegistry()
bindings := make([]ModuleBinding, 0, len(validators))
for _, validator := range validators {
bindings = append(bindings, Binding(validator.Name()))
if resolved == nil {
t.Fatal("resolved pipeline must not be nil")
}
if err := registry.Register(ValidatorChainMapping{Stage: stage, Module: module, Validators: bindings}); err != nil {
t.Fatalf("register validator chain: %v", err)
chain := ResolvedValidatorChain{
Stage: stage,
LaneID: laneID,
ModuleKey: module,
Validators: append([]ResolvedValidator(nil), validators...),
}
return registry
for index := range resolved.ValidatorChains {
existing := resolved.ValidatorChains[index]
if existing.Stage == stage && existing.LaneID == laneID && existing.ModuleKey == module {
resolved.ValidatorChains[index] = chain
return
}
}
resolved.ValidatorChains = append(resolved.ValidatorChains, chain)
}