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

@@ -208,8 +208,9 @@ The production CLI currently registers these module keys:
- normalize: `noop`
- output: `json`
Configured validator module lists are reserved for a future validator-chain
feature and are rejected by current configuration validation.
Validator chain overrides are configured on `chunk`, lane `extract`, lane
`merge`, and lane `normalize` bindings. Validator keys are resolved against the
registered validator catalog.
For YAML structure, Scriptorium profile sources, environment overrides, and
module binding syntax, see [Configuration](config.md).

View File

@@ -183,9 +183,11 @@ Within an artifact lane, the runner:
The current runner handoff is raw-output based. Extractors, mergers, and
normalizers do not advertise validator chains through their module interfaces.
Runner-side raw validation chains receive the raw module output plus
stage, lane, module, source, and chunk provenance. Empty raw validation chains
approve output by default.
Runner-side raw validation chains receive the raw module output plus stage,
lane, module, source, chunk, schema, session, reference, LLM client/profile,
binding option, and run metadata context. Merge validators also receive the
ordered extract outputs used by the merge, and normalize validators receive the
accepted merge output. Empty raw validation chains approve output by default.
Resolved validator chains come from central default mappings unless a
stage-local config override is set on `chunk`, lane `extract`, lane `merge`, or

View File

@@ -127,7 +127,7 @@ func (r *Runner) Run(ctx context.Context, input RunInput) (output RunOutput, err
if err != nil {
return false, nil, fmt.Errorf("validate chunks from chunker %q: %w", chunker.Key(), err)
}
rejection, err := r.validateChunksRaw(ctx, doc, chunker.Key(), chunks, input.Metadata, input.Pipeline.ValidatorChains, attempt)
rejection, err := r.validateChunksRaw(ctx, doc, chunker.Key(), chunks, sourceInput, sessionID, input.Pipeline.ChunkReferences.ReferenceSet, input.LLMClient, input.Metadata, input.Pipeline.ValidatorChains, attempt)
if err != nil || rejection != nil {
return false, rejection, err
}
@@ -230,18 +230,23 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
extractOutput.ChunkIndex = chunk.Index
extractOutput.Payload.Warnings = append(extractOutput.Payload.Warnings, result.Warnings...)
validationWarnings, rejection, err := r.validateRaw(ctx, rawValidationTarget{
stage: StageExtract,
laneID: lane.ID,
moduleKey: extractor.Key(),
source: doc,
sourceID: doc.ID,
chunkID: chunk.ID,
chunkIndex: chunk.Index,
schema: extractOutput.Schema,
payload: extractOutput.Payload,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
stage: StageExtract,
laneID: lane.ID,
moduleKey: extractor.Key(),
source: doc,
sourceID: doc.ID,
chunkID: chunk.ID,
chunkIndex: chunk.Index,
chunk: &chunk,
sourceInput: chunkInputMaterial(sourceInput, chunk),
sessionID: sessionID,
references: lane.ExtractReferences.ReferenceSet,
llmClient: input.LLMClient,
schema: extractOutput.Schema,
payload: extractOutput.Payload,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
})
if err != nil || rejection != nil {
return false, rejection, err
@@ -289,16 +294,21 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
mergeOutput.SourceID = doc.ID
mergeOutput.Payload.Warnings = append(mergeOutput.Payload.Warnings, mergeResult.Warnings...)
validationWarnings, rejection, err := r.validateRaw(ctx, rawValidationTarget{
stage: StageMerge,
laneID: lane.ID,
moduleKey: merger.Key(),
source: doc,
sourceID: doc.ID,
schema: mergeOutput.Schema,
payload: mergeOutput.Payload,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
stage: StageMerge,
laneID: lane.ID,
moduleKey: merger.Key(),
source: doc,
sourceID: doc.ID,
sourceInput: sourceInput.Clone(),
sessionID: sessionID,
references: lane.MergeReferences.ReferenceSet,
llmClient: input.LLMClient,
schema: mergeOutput.Schema,
payload: mergeOutput.Payload,
extractOutputs: extractOutputs,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
})
if err != nil || rejection != nil {
return false, rejection, err
@@ -340,16 +350,21 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
normalizeOutput.SourceID = doc.ID
normalizeOutput.Payload.Warnings = append(normalizeOutput.Payload.Warnings, normalizeResult.Warnings...)
validationWarnings, rejection, err := r.validateRaw(ctx, rawValidationTarget{
stage: StageNormalize,
laneID: lane.ID,
moduleKey: normalizer.Key(),
source: doc,
sourceID: doc.ID,
schema: normalizeOutput.Schema,
payload: normalizeOutput.Payload,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
stage: StageNormalize,
laneID: lane.ID,
moduleKey: normalizer.Key(),
source: doc,
sourceID: doc.ID,
sourceInput: sourceInput.Clone(),
sessionID: sessionID,
references: lane.NormalizeReferences.ReferenceSet,
llmClient: input.LLMClient,
schema: normalizeOutput.Schema,
payload: normalizeOutput.Payload,
mergeOutput: acceptedMerge,
metadata: input.Metadata,
chains: input.Pipeline.ValidatorChains,
attempt: attempt,
})
if err != nil || rejection != nil {
return false, rejection, err
@@ -371,18 +386,26 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
}
type rawValidationTarget struct {
stage ModuleStage
laneID string
moduleKey string
source *source.SourceDocument
sourceID string
chunkID string
chunkIndex int
schema contracts.ResponseSchema
payload contracts.RawPayload
metadata map[string]any
chains []ResolvedValidatorChain
attempt int
stage ModuleStage
laneID string
moduleKey string
source *source.SourceDocument
sourceID string
sourceInput contracts.LLMInputMaterial
sessionID string
references contracts.ReferenceSet
llmClient contracts.StructuredLLMClient
chunkID string
chunkIndex int
chunk *contracts.SourceChunk
chunks []contracts.SourceChunk
schema contracts.ResponseSchema
payload contracts.RawPayload
extractOutputs []contracts.ExtractOutput
mergeOutput contracts.MergeOutput
metadata map[string]any
chains []ResolvedValidatorChain
attempt int
}
func runWithRetry(ctx context.Context, retries int, run func(attempt int) (bool, *contracts.RejectedOutput, error)) (bool, *contracts.RejectedOutput, error) {
@@ -432,15 +455,22 @@ func runWithRetry(ctx context.Context, retries int, run func(attempt int) (bool,
return false, lastRejection, nil
}
func (r *Runner) validateChunksRaw(ctx context.Context, doc *source.SourceDocument, moduleKey string, chunks []contracts.SourceChunk, metadata map[string]any, chains []ResolvedValidatorChain, attempt int) (*contracts.RejectedOutput, error) {
for _, chunk := range chunks {
func (r *Runner) validateChunksRaw(ctx context.Context, doc *source.SourceDocument, moduleKey string, chunks []contracts.SourceChunk, sourceInput contracts.LLMInputMaterial, sessionID string, references contracts.ReferenceSet, llmClient contracts.StructuredLLMClient, metadata map[string]any, chains []ResolvedValidatorChain, attempt int) (*contracts.RejectedOutput, error) {
for index := range chunks {
chunk := chunks[index]
_, rejection, err := r.validateRaw(ctx, rawValidationTarget{
stage: StageChunk,
moduleKey: moduleKey,
source: doc,
sourceID: doc.ID,
chunkID: chunk.ID,
chunkIndex: chunk.Index,
stage: StageChunk,
moduleKey: moduleKey,
source: doc,
sourceID: doc.ID,
sourceInput: sourceInput.Clone(),
sessionID: sessionID,
references: references,
llmClient: llmClient,
chunkID: chunk.ID,
chunkIndex: chunk.Index,
chunk: &chunk,
chunks: chunks,
payload: contracts.RawPayload{
Content: append([]byte(nil), chunk.Content...),
MediaType: chunk.MediaType,
@@ -462,9 +492,6 @@ func (r *Runner) validateChunksRaw(ctx context.Context, doc *source.SourceDocume
func (r *Runner) validateRaw(ctx context.Context, target rawValidationTarget) ([]contracts.Warning, *contracts.RejectedOutput, error) {
chain := resolvedValidatorChain(target.stage, target.laneID, target.moduleKey, target.chains)
if len(chain.Validators) == 0 {
chain = r.registryValidatorChain(target.stage, target.laneID, target.moduleKey)
}
if len(chain.Validators) == 0 {
return nil, nil, nil
}
@@ -472,25 +499,13 @@ func (r *Runner) validateRaw(ctx context.Context, target rawValidationTarget) ([
return nil, nil, fmt.Errorf("validator registry must not be nil")
}
request := contracts.ValidationRequest{
Stage: string(target.stage),
LaneID: target.laneID,
ModuleKey: target.moduleKey,
Source: target.source,
SourceID: target.sourceID,
ChunkID: target.chunkID,
ChunkIndex: target.chunkIndex,
Schema: target.schema,
Payload: cloneRawPayload(target.payload),
Metadata: cloneMetadata(target.metadata),
}
var warnings []contracts.Warning
for _, validatorBinding := range chain.Validators {
validator, err := r.registries.Validators.Build(validatorBinding.Binding.Module)
if err != nil {
return nil, nil, fmt.Errorf("build validator %q: %w", validatorBinding.Binding.Module, err)
}
request := target.validationRequest(validatorBinding.Binding)
result, err := validator.Validate(ctx, request)
if err != nil {
return nil, nil, fmt.Errorf("validate raw %s output with validator %q: %w", target.stage, validator.Name(), err)
@@ -522,33 +537,29 @@ func (r *Runner) validateRaw(ctx context.Context, target rawValidationTarget) ([
return warnings, nil, nil
}
func (r *Runner) registryValidatorChain(stage ModuleStage, laneID string, moduleKey string) ResolvedValidatorChain {
chain := ResolvedValidatorChain{
Stage: stage,
LaneID: strings.TrimSpace(laneID),
ModuleKey: strings.TrimSpace(moduleKey),
func (target rawValidationTarget) validationRequest(binding ModuleBinding) contracts.ValidationRequest {
return contracts.ValidationRequest{
Stage: string(target.stage),
LaneID: target.laneID,
ModuleKey: target.moduleKey,
Source: target.source,
SourceID: target.sourceID,
SourceInput: target.sourceInput.Clone(),
SessionID: target.sessionID,
References: CloneReferenceSet(target.references),
LLMClient: target.llmClient,
LLMProfile: binding.LLMProfile,
Options: cloneOptions(binding.Options),
Metadata: cloneMetadata(target.metadata),
Schema: target.schema,
Payload: cloneRawPayload(target.payload),
ChunkID: target.chunkID,
ChunkIndex: target.chunkIndex,
Chunk: cloneSourceChunkPtr(target.chunk),
Chunks: cloneSourceChunks(target.chunks),
ExtractOutputs: cloneExtractOutputs(target.extractOutputs),
MergeOutput: cloneMergeOutput(target.mergeOutput),
}
if r == nil || r.registries.ValidatorChains == nil {
return chain
}
bindings := r.registries.ValidatorChains.Validators(stage, moduleKey)
if len(bindings) == 0 {
return chain
}
chain.Validators = make([]ResolvedValidator, 0, len(bindings))
for _, binding := range bindings {
executionClass := contracts.ExecutionClass("")
if r.registries.Validators != nil {
if spec, ok := r.registries.Validators.Spec(binding.Module); ok {
executionClass = spec.ExecutionClass
}
}
chain.Validators = append(chain.Validators, ResolvedValidator{
Binding: cloneModuleBinding(binding),
ExecutionClass: executionClass,
})
}
return chain
}
func resolvedValidatorChain(stage ModuleStage, laneID string, moduleKey string, chains []ResolvedValidatorChain) ResolvedValidatorChain {
@@ -1001,6 +1012,43 @@ func cloneRawPayload(payload contracts.RawPayload) contracts.RawPayload {
}
}
func cloneSourceChunkPtr(chunk *contracts.SourceChunk) *contracts.SourceChunk {
if chunk == nil {
return nil
}
cloned := cloneSourceChunk(*chunk)
return &cloned
}
func cloneSourceChunk(chunk contracts.SourceChunk) contracts.SourceChunk {
chunk.Content = append([]byte(nil), chunk.Content...)
chunk.Units = cloneSourceUnits(chunk.Units)
chunk.Metadata = cloneMetadata(chunk.Metadata)
return chunk
}
func cloneSourceChunks(chunks []contracts.SourceChunk) []contracts.SourceChunk {
if len(chunks) == 0 {
return nil
}
out := make([]contracts.SourceChunk, 0, len(chunks))
for _, chunk := range chunks {
out = append(out, cloneSourceChunk(chunk))
}
return out
}
func cloneSourceUnits(units []source.SourceUnit) []source.SourceUnit {
if len(units) == 0 {
return nil
}
out := make([]source.SourceUnit, 0, len(units))
for _, unit := range units {
out = append(out, cloneSourceUnit(unit))
}
return out
}
func cloneExtractOutput(output contracts.ExtractOutput) contracts.ExtractOutput {
output.Payload = cloneRawPayload(output.Payload)
return output

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)
}