Wire resolved validator chains into runner
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user