Wire resolved validator chains into runner
This commit is contained in:
@@ -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