Add validator chain provenance
This commit is contained in:
@@ -18,14 +18,14 @@ import (
|
||||
)
|
||||
|
||||
type Registries struct {
|
||||
Inputs *InputAdapterRegistry
|
||||
Chunkers *ChunkerRegistry
|
||||
Extractors *ExtractorRegistry
|
||||
Mergers *MergerRegistry
|
||||
Normalizers *NormalizerRegistry
|
||||
Validators *ValidatorRegistry
|
||||
RawValidators *RawValidationRegistry
|
||||
Outputs *OutputEncoderRegistry
|
||||
Inputs *InputAdapterRegistry
|
||||
Chunkers *ChunkerRegistry
|
||||
Extractors *ExtractorRegistry
|
||||
Mergers *MergerRegistry
|
||||
Normalizers *NormalizerRegistry
|
||||
Validators *ValidatorRegistry
|
||||
ValidatorChains *ValidatorChainRegistry
|
||||
Outputs *OutputEncoderRegistry
|
||||
}
|
||||
|
||||
type Runner struct {
|
||||
@@ -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, attempt)
|
||||
rejection, err := r.validateChunksRaw(ctx, doc, chunker.Key(), chunks, input.Metadata, input.Pipeline.ValidatorChains, attempt)
|
||||
if err != nil || rejection != nil {
|
||||
return false, rejection, err
|
||||
}
|
||||
@@ -240,6 +240,7 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
|
||||
schema: extractOutput.Schema,
|
||||
payload: extractOutput.Payload,
|
||||
metadata: input.Metadata,
|
||||
chains: input.Pipeline.ValidatorChains,
|
||||
attempt: attempt,
|
||||
})
|
||||
if err != nil || rejection != nil {
|
||||
@@ -296,6 +297,7 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
|
||||
schema: mergeOutput.Schema,
|
||||
payload: mergeOutput.Payload,
|
||||
metadata: input.Metadata,
|
||||
chains: input.Pipeline.ValidatorChains,
|
||||
attempt: attempt,
|
||||
})
|
||||
if err != nil || rejection != nil {
|
||||
@@ -346,6 +348,7 @@ func (r *Runner) runLane(ctx context.Context, input RunInput, doc *source.Source
|
||||
schema: normalizeOutput.Schema,
|
||||
payload: normalizeOutput.Payload,
|
||||
metadata: input.Metadata,
|
||||
chains: input.Pipeline.ValidatorChains,
|
||||
attempt: attempt,
|
||||
})
|
||||
if err != nil || rejection != nil {
|
||||
@@ -378,6 +381,7 @@ type rawValidationTarget struct {
|
||||
schema contracts.ResponseSchema
|
||||
payload contracts.RawPayload
|
||||
metadata map[string]any
|
||||
chains []ResolvedValidatorChain
|
||||
attempt int
|
||||
}
|
||||
|
||||
@@ -428,7 +432,7 @@ 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, attempt int) (*contracts.RejectedOutput, error) {
|
||||
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 {
|
||||
_, rejection, err := r.validateRaw(ctx, rawValidationTarget{
|
||||
stage: StageChunk,
|
||||
@@ -443,6 +447,7 @@ func (r *Runner) validateChunksRaw(ctx context.Context, doc *source.SourceDocume
|
||||
Metadata: cloneMetadata(chunk.Metadata),
|
||||
},
|
||||
metadata: metadata,
|
||||
chains: chains,
|
||||
attempt: attempt,
|
||||
})
|
||||
if err != nil {
|
||||
@@ -456,10 +461,16 @@ func (r *Runner) validateChunksRaw(ctx context.Context, doc *source.SourceDocume
|
||||
}
|
||||
|
||||
func (r *Runner) validateRaw(ctx context.Context, target rawValidationTarget) ([]contracts.Warning, *contracts.RejectedOutput, error) {
|
||||
validators := r.registries.RawValidators.Validators(target.stage, target.moduleKey)
|
||||
if len(validators) == 0 {
|
||||
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
|
||||
}
|
||||
if r.registries.Validators == nil {
|
||||
return nil, nil, fmt.Errorf("validator registry must not be nil")
|
||||
}
|
||||
|
||||
request := contracts.ValidationRequest{
|
||||
Stage: string(target.stage),
|
||||
@@ -475,7 +486,11 @@ func (r *Runner) validateRaw(ctx context.Context, target rawValidationTarget) ([
|
||||
}
|
||||
|
||||
var warnings []contracts.Warning
|
||||
for _, validator := range validators {
|
||||
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)
|
||||
}
|
||||
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)
|
||||
@@ -507,6 +522,60 @@ 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),
|
||||
}
|
||||
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 {
|
||||
for _, chain := range chains {
|
||||
if chain.Stage != stage {
|
||||
continue
|
||||
}
|
||||
if chain.ModuleKey != moduleKey {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(chain.LaneID) != strings.TrimSpace(laneID) {
|
||||
continue
|
||||
}
|
||||
return ResolvedValidatorChain{
|
||||
Stage: chain.Stage,
|
||||
LaneID: chain.LaneID,
|
||||
ModuleKey: chain.ModuleKey,
|
||||
Validators: cloneResolvedValidators(chain.Validators),
|
||||
}
|
||||
}
|
||||
return ResolvedValidatorChain{
|
||||
Stage: stage,
|
||||
LaneID: strings.TrimSpace(laneID),
|
||||
ModuleKey: strings.TrimSpace(moduleKey),
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) validateRegistries(pipeline ResolvedPipeline) error {
|
||||
if r.registries.Inputs == nil {
|
||||
return fmt.Errorf("input registry must not be nil")
|
||||
@@ -580,16 +649,17 @@ func manifestFromPipeline(input RunInput) artifacts.RunManifest {
|
||||
|
||||
pipeline := input.Pipeline
|
||||
manifest := artifacts.RunManifest{
|
||||
PipelineID: pipeline.ID,
|
||||
PipelineDigest: pipeline.Digest,
|
||||
InputModule: pipeline.Input.Module,
|
||||
Chunker: pipeline.Chunk.Module,
|
||||
OutputEncoder: pipeline.Output.Module,
|
||||
ArtifactLanes: make([]artifacts.ArtifactLaneManifest, 0, len(pipeline.ArtifactLanes)),
|
||||
RunID: runID,
|
||||
StartedAt: timePtr(startedAt),
|
||||
References: ReferenceProvenance(pipeline),
|
||||
LLMProfiles: cloneLLMProfiles(input.LLMProfiles),
|
||||
PipelineID: pipeline.ID,
|
||||
PipelineDigest: pipeline.Digest,
|
||||
InputModule: pipeline.Input.Module,
|
||||
Chunker: pipeline.Chunk.Module,
|
||||
OutputEncoder: pipeline.Output.Module,
|
||||
ArtifactLanes: make([]artifacts.ArtifactLaneManifest, 0, len(pipeline.ArtifactLanes)),
|
||||
ValidatorChains: validatorChainManifests(pipeline.ValidatorChains),
|
||||
RunID: runID,
|
||||
StartedAt: timePtr(startedAt),
|
||||
References: ReferenceProvenance(pipeline),
|
||||
LLMProfiles: cloneLLMProfiles(input.LLMProfiles),
|
||||
}
|
||||
// The runner does not currently maintain a cache or idempotency key. Reference
|
||||
// digests are recorded in manifest provenance and intentionally kept separate
|
||||
@@ -607,6 +677,29 @@ func manifestFromPipeline(input RunInput) artifacts.RunManifest {
|
||||
return manifest
|
||||
}
|
||||
|
||||
func validatorChainManifests(chains []ResolvedValidatorChain) []artifacts.ValidatorChainManifest {
|
||||
if len(chains) == 0 {
|
||||
return nil
|
||||
}
|
||||
manifests := make([]artifacts.ValidatorChainManifest, 0, len(chains))
|
||||
for _, chain := range chains {
|
||||
manifest := artifacts.ValidatorChainManifest{
|
||||
Stage: string(chain.Stage),
|
||||
LaneID: chain.LaneID,
|
||||
ModuleKey: chain.ModuleKey,
|
||||
Validators: make([]artifacts.ValidatorManifest, 0, len(chain.Validators)),
|
||||
}
|
||||
for _, validator := range chain.Validators {
|
||||
manifest.Validators = append(manifest.Validators, artifacts.ValidatorManifest{
|
||||
Key: validator.Binding.Module,
|
||||
ExecutionClass: string(validator.ExecutionClass),
|
||||
})
|
||||
}
|
||||
manifests = append(manifests, manifest)
|
||||
}
|
||||
return manifests
|
||||
}
|
||||
|
||||
func failOutput(output RunOutput) RunOutput {
|
||||
if output.Manifest.PipelineID != "" {
|
||||
populateRawOutputManifest(&output)
|
||||
|
||||
Reference in New Issue
Block a user