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
|
||||
|
||||
Reference in New Issue
Block a user