Implement runner retries and raw validation
This commit is contained in:
@@ -854,6 +854,209 @@ func TestRunPassesPerChunkRawOutputsToMergeAndNormalize(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunPassesChunkContentAndMediaTypeToExtractors(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
modules.chunker.chunks = []contracts.SourceChunk{
|
||||
sourceChunkWithContent("chunk-0", 0, []byte(`{"chunk":0}`), "application/vnd.test+json"),
|
||||
}
|
||||
|
||||
_, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
req := modules.extractors["extract-alpha"].requests[0]
|
||||
if req.Chunk == nil {
|
||||
t.Fatal("extractor chunk = nil, want chunk")
|
||||
}
|
||||
if got := string(req.Chunk.Content); got != `{"chunk":0}` {
|
||||
t.Fatalf("chunk content = %q, want raw chunk content", got)
|
||||
}
|
||||
if req.Chunk.MediaType != "application/vnd.test+json" {
|
||||
t.Fatalf("chunk media type = %q, want application/vnd.test+json", req.Chunk.MediaType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunOmitsRejectedExtractOutputsFromMerge(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-extract", approved: []bool{false, true}, reason: "bad_extract", message: "extract rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageExtract, "extract-alpha", validator)
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
if len(output.Rejected) != 1 || output.Rejected[0].Stage != string(StageExtract) || output.Rejected[0].ChunkID != "chunk-0" {
|
||||
t.Fatalf("rejected outputs = %#v, want rejected first extract", output.Rejected)
|
||||
}
|
||||
extractOutputs := modules.mergers["merge"].requests[0].ExtractOutputs
|
||||
if len(extractOutputs) != 1 || extractOutputs[0].ChunkID != "chunk-1" {
|
||||
t.Fatalf("merge extract outputs = %#v, want only accepted second chunk", extractOutputs)
|
||||
}
|
||||
if output.Manifest.ValidationStatus != "rejected" {
|
||||
t.Fatalf("ValidationStatus = %q, want rejected", output.Manifest.ValidationStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunOmitsLaneWithNoAcceptedExtractOutputs(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-extract", approved: []bool{false}, reason: "bad_extract", message: "extract rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageExtract, "extract-alpha", validator)
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
if len(output.Rejected) != 2 {
|
||||
t.Fatalf("len(Rejected) = %d, want one rejected record per chunk", len(output.Rejected))
|
||||
}
|
||||
if len(modules.mergers["merge"].requests) != 0 {
|
||||
t.Fatalf("merge requests = %d, want none", len(modules.mergers["merge"].requests))
|
||||
}
|
||||
if len(modules.normalizers["normalize"].requests) != 0 {
|
||||
t.Fatalf("normalize requests = %d, want none", len(modules.normalizers["normalize"].requests))
|
||||
}
|
||||
if len(modules.output.requests) != 1 || len(modules.output.requests[0].NormalizeOutputs) != 0 {
|
||||
t.Fatalf("output normalize outputs = %#v, want none", modules.output.requests)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRejectedMergePreventsNormalizeForLane(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-merge", approved: []bool{false}, reason: "bad_merge", message: "merge rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageMerge, "merge", validator)
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
if len(output.Rejected) != 1 || output.Rejected[0].Stage != string(StageMerge) {
|
||||
t.Fatalf("rejected outputs = %#v, want rejected merge", output.Rejected)
|
||||
}
|
||||
if len(modules.normalizers["normalize"].requests) != 0 {
|
||||
t.Fatalf("normalize requests = %d, want none", len(modules.normalizers["normalize"].requests))
|
||||
}
|
||||
if len(modules.output.requests[0].NormalizeOutputs) != 0 {
|
||||
t.Fatalf("output normalize outputs = %#v, want none", modules.output.requests[0].NormalizeOutputs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRejectedNormalizePreventsOutputForLane(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-normalize", approved: []bool{false}, reason: "bad_normalize", message: "normalize rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageNormalize, "normalize", validator)
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: resolvedPipeline()})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
if len(output.Rejected) != 1 || output.Rejected[0].Stage != string(StageNormalize) {
|
||||
t.Fatalf("rejected outputs = %#v, want rejected normalize", output.Rejected)
|
||||
}
|
||||
if len(output.NormalizeOutputs) != 0 {
|
||||
t.Fatalf("NormalizeOutputs = %#v, want none", output.NormalizeOutputs)
|
||||
}
|
||||
if len(modules.output.requests[0].NormalizeOutputs) != 0 {
|
||||
t.Fatalf("output normalize outputs = %#v, want none", modules.output.requests[0].NormalizeOutputs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRetriesSameModuleInputAfterFrameworkError(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
modules.extractors["extract-alpha"].failuresBeforeSuccess = 1
|
||||
modules.extractors["extract-alpha"].failureErr = errors.New("transient extract failure")
|
||||
pipeline := resolvedPipeline()
|
||||
pipeline.ArtifactLanes[0].Extract.Retries = 1
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
requests := modules.extractors["extract-alpha"].requests
|
||||
if len(requests) != 3 {
|
||||
t.Fatalf("extract requests = %d, want retry plus remaining chunk", len(requests))
|
||||
}
|
||||
if requests[0].Chunk.ID != "chunk-0" || requests[1].Chunk.ID != "chunk-0" {
|
||||
t.Fatalf("retried chunks = %q, %q; want same first chunk input", requests[0].Chunk.ID, requests[1].Chunk.ID)
|
||||
}
|
||||
if output.Manifest.ValidationStatus != "approved" {
|
||||
t.Fatalf("ValidationStatus = %q, want approved", output.Manifest.ValidationStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRetriesSameModuleInputAfterValidatorRejection(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-extract", approved: []bool{false, true, true}, reason: "bad_extract", message: "extract rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageExtract, "extract-alpha", validator)
|
||||
pipeline := resolvedPipeline()
|
||||
pipeline.ArtifactLanes[0].Extract.Retries = 1
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
requests := modules.extractors["extract-alpha"].requests
|
||||
if len(requests) != 3 {
|
||||
t.Fatalf("extract requests = %d, want retry plus remaining chunk", len(requests))
|
||||
}
|
||||
if requests[0].Chunk.ID != "chunk-0" || requests[1].Chunk.ID != "chunk-0" {
|
||||
t.Fatalf("retried chunks = %q, %q; want same first chunk input", requests[0].Chunk.ID, requests[1].Chunk.ID)
|
||||
}
|
||||
if len(output.Rejected) != 0 {
|
||||
t.Fatalf("Rejected = %#v, want transient rejection omitted after retry approval", output.Rejected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunStopsRetryAfterConfiguredAttemptsAndRecordsAttemptCount(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
validator := &runnerRawValidator{name: "raw-extract", approved: []bool{false}, reason: "bad_extract", message: "extract rejected"}
|
||||
modules.rawValidators = rawValidationRegistry(t, StageExtract, "extract-alpha", validator)
|
||||
pipeline := resolvedPipeline()
|
||||
pipeline.ArtifactLanes[0].Extract.Retries = 1
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(context.Background(), RunInput{Pipeline: pipeline})
|
||||
if err != nil {
|
||||
t.Fatalf("Run() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
if len(output.Rejected) != 2 {
|
||||
t.Fatalf("len(Rejected) = %d, want rejected record per chunk", len(output.Rejected))
|
||||
}
|
||||
if output.Rejected[0].AttemptCount != 2 || output.Rejected[1].AttemptCount != 2 {
|
||||
t.Fatalf("attempt counts = %#v, want final attempt count 2", output.Rejected)
|
||||
}
|
||||
if len(modules.extractors["extract-alpha"].requests) != 4 {
|
||||
t.Fatalf("extract requests = %d, want two attempts per chunk", len(modules.extractors["extract-alpha"].requests))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunContextCancellationStopsRetries(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
modules.extractors["extract-alpha"].err = errors.New("extract failed")
|
||||
pipeline := resolvedPipeline()
|
||||
pipeline.ArtifactLanes[0].Extract.Retries = 2
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
output, err := New(newRunnerRegistries(t, modules)).Run(ctx, RunInput{Pipeline: pipeline})
|
||||
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Run() error = %v, want context.Canceled", err)
|
||||
}
|
||||
if len(modules.extractors["extract-alpha"].requests) != 0 {
|
||||
t.Fatalf("extract requests = %d, want none after cancellation", len(modules.extractors["extract-alpha"].requests))
|
||||
}
|
||||
if output.Manifest.ValidationStatus != "failed" {
|
||||
t.Fatalf("ValidationStatus = %q, want failed", output.Manifest.ValidationStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRecordsConfiguredValidatorsInManifest(t *testing.T) {
|
||||
modules := defaultRunnerModules()
|
||||
|
||||
@@ -1278,6 +1481,7 @@ type runnerModules struct {
|
||||
mergers map[string]*runnerMerger
|
||||
normalizers map[string]*runnerNormalizer
|
||||
validators map[string]*runnerValidator
|
||||
rawValidators *RawValidationRegistry
|
||||
output *runnerOutputEncoder
|
||||
inputBuildErr error
|
||||
chunkerBuildErr error
|
||||
@@ -1316,13 +1520,14 @@ func newRunnerRegistries(t *testing.T, modules *runnerModules) Registries {
|
||||
}
|
||||
|
||||
registries := Registries{
|
||||
Inputs: NewInputAdapterRegistry(),
|
||||
Chunkers: NewChunkerRegistry(),
|
||||
Extractors: NewExtractorRegistry(),
|
||||
Mergers: NewMergerRegistry(),
|
||||
Normalizers: NewNormalizerRegistry(),
|
||||
Validators: NewValidatorRegistry(),
|
||||
Outputs: NewOutputEncoderRegistry(),
|
||||
Inputs: NewInputAdapterRegistry(),
|
||||
Chunkers: NewChunkerRegistry(),
|
||||
Extractors: NewExtractorRegistry(),
|
||||
Mergers: NewMergerRegistry(),
|
||||
Normalizers: NewNormalizerRegistry(),
|
||||
Validators: NewValidatorRegistry(),
|
||||
RawValidators: modules.rawValidators,
|
||||
Outputs: NewOutputEncoderRegistry(),
|
||||
}
|
||||
if err := registries.Inputs.Register("input", func() (contracts.InputAdapter, error) {
|
||||
if modules.inputBuildErr != nil {
|
||||
@@ -1392,12 +1597,14 @@ func (adapter *runnerInputAdapter) ManifestMetadata() map[string]any {
|
||||
}
|
||||
|
||||
type runnerChunker struct {
|
||||
key string
|
||||
chunks []contracts.SourceChunk
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
manifestMetadata map[string]any
|
||||
requests []contracts.ChunkRequest
|
||||
key string
|
||||
chunks []contracts.SourceChunk
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
failureErr error
|
||||
failuresBeforeSuccess int
|
||||
manifestMetadata map[string]any
|
||||
requests []contracts.ChunkRequest
|
||||
}
|
||||
|
||||
func (chunker *runnerChunker) Key() string {
|
||||
@@ -1410,6 +1617,14 @@ func (chunker *runnerChunker) ReferenceSlots() []contracts.ReferenceSlot {
|
||||
|
||||
func (chunker *runnerChunker) Chunk(ctx context.Context, req contracts.ChunkRequest) (contracts.ChunkResult, error) {
|
||||
chunker.requests = append(chunker.requests, req)
|
||||
if chunker.failuresBeforeSuccess > 0 {
|
||||
chunker.failuresBeforeSuccess--
|
||||
err := chunker.failureErr
|
||||
if err == nil {
|
||||
err = errors.New("transient chunk failure")
|
||||
}
|
||||
return contracts.ChunkResult{}, err
|
||||
}
|
||||
return contracts.ChunkResult{
|
||||
Chunks: chunker.chunks,
|
||||
Warnings: chunker.warnings,
|
||||
@@ -1421,15 +1636,17 @@ func (chunker *runnerChunker) ManifestMetadata() map[string]any {
|
||||
}
|
||||
|
||||
type runnerExtractor struct {
|
||||
key string
|
||||
manifestMetadata map[string]any
|
||||
output *contracts.ExtractOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
requests []contracts.ExtractionRequest
|
||||
seenChunkIDs []string
|
||||
seenLLMClients []contracts.StructuredLLMClient
|
||||
seenMetadata []map[string]any
|
||||
key string
|
||||
manifestMetadata map[string]any
|
||||
output *contracts.ExtractOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
failureErr error
|
||||
failuresBeforeSuccess int
|
||||
requests []contracts.ExtractionRequest
|
||||
seenChunkIDs []string
|
||||
seenLLMClients []contracts.StructuredLLMClient
|
||||
seenMetadata []map[string]any
|
||||
}
|
||||
|
||||
func (extractor *runnerExtractor) Key() string {
|
||||
@@ -1452,6 +1669,15 @@ func (extractor *runnerExtractor) Extract(ctx context.Context, req contracts.Ext
|
||||
extractor.seenLLMClients = append(extractor.seenLLMClients, req.LLMClient)
|
||||
extractor.seenMetadata = append(extractor.seenMetadata, req.Metadata)
|
||||
|
||||
if extractor.failuresBeforeSuccess > 0 {
|
||||
extractor.failuresBeforeSuccess--
|
||||
err := extractor.failureErr
|
||||
if err == nil {
|
||||
err = errors.New("transient extract failure")
|
||||
}
|
||||
return contracts.ExtractionResult{}, err
|
||||
}
|
||||
|
||||
output := contracts.ExtractOutput{
|
||||
Schema: contracts.ResponseSchema{ID: "runner.raw", Name: "runner_raw", Version: "v1"},
|
||||
Payload: contracts.RawPayload{
|
||||
@@ -1472,11 +1698,13 @@ func (extractor *runnerExtractor) Extract(ctx context.Context, req contracts.Ext
|
||||
}
|
||||
|
||||
type runnerMerger struct {
|
||||
key string
|
||||
result *contracts.MergeOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
requests []contracts.MergeRequest
|
||||
key string
|
||||
result *contracts.MergeOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
failureErr error
|
||||
failuresBeforeSuccess int
|
||||
requests []contracts.MergeRequest
|
||||
}
|
||||
|
||||
func (merger *runnerMerger) Key() string {
|
||||
@@ -1485,6 +1713,14 @@ func (merger *runnerMerger) Key() string {
|
||||
|
||||
func (merger *runnerMerger) Merge(ctx context.Context, req contracts.MergeRequest) (contracts.MergeResult, error) {
|
||||
merger.requests = append(merger.requests, req)
|
||||
if merger.failuresBeforeSuccess > 0 {
|
||||
merger.failuresBeforeSuccess--
|
||||
err := merger.failureErr
|
||||
if err == nil {
|
||||
err = errors.New("transient merge failure")
|
||||
}
|
||||
return contracts.MergeResult{}, err
|
||||
}
|
||||
output := contracts.MergeOutput{
|
||||
LaneID: req.LaneID,
|
||||
SourceID: req.Source.ID,
|
||||
@@ -1504,11 +1740,13 @@ func (merger *runnerMerger) Merge(ctx context.Context, req contracts.MergeReques
|
||||
}
|
||||
|
||||
type runnerNormalizer struct {
|
||||
key string
|
||||
result *contracts.NormalizeOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
requests []contracts.NormalizeRequest
|
||||
key string
|
||||
result *contracts.NormalizeOutput
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
failureErr error
|
||||
failuresBeforeSuccess int
|
||||
requests []contracts.NormalizeRequest
|
||||
}
|
||||
|
||||
func (normalizer *runnerNormalizer) Key() string {
|
||||
@@ -1521,6 +1759,14 @@ func (normalizer *runnerNormalizer) ReferenceSlots() []contracts.ReferenceSlot {
|
||||
|
||||
func (normalizer *runnerNormalizer) Normalize(ctx context.Context, req contracts.NormalizeRequest) (contracts.NormalizeResult, error) {
|
||||
normalizer.requests = append(normalizer.requests, req)
|
||||
if normalizer.failuresBeforeSuccess > 0 {
|
||||
normalizer.failuresBeforeSuccess--
|
||||
err := normalizer.failureErr
|
||||
if err == nil {
|
||||
err = errors.New("transient normalize failure")
|
||||
}
|
||||
return contracts.NormalizeResult{}, err
|
||||
}
|
||||
output := contracts.NormalizeOutput{
|
||||
LaneID: req.LaneID,
|
||||
SourceID: req.MergeOutput.SourceID,
|
||||
@@ -1547,6 +1793,43 @@ type runnerValidator struct {
|
||||
requests []contracts.ValidationRequest
|
||||
}
|
||||
|
||||
type runnerRawValidator struct {
|
||||
name string
|
||||
approved []bool
|
||||
reason string
|
||||
message string
|
||||
warnings []contracts.Warning
|
||||
err error
|
||||
calls int
|
||||
requests []contracts.RawValidationRequest
|
||||
}
|
||||
|
||||
func (validator *runnerRawValidator) Name() string {
|
||||
return validator.name
|
||||
}
|
||||
|
||||
func (validator *runnerRawValidator) ValidateRaw(ctx context.Context, req contracts.RawValidationRequest) (contracts.RawValidationResult, error) {
|
||||
validator.calls++
|
||||
validator.requests = append(validator.requests, req)
|
||||
if validator.err != nil {
|
||||
return contracts.RawValidationResult{}, validator.err
|
||||
}
|
||||
approved := true
|
||||
if len(validator.approved) > 0 {
|
||||
index := validator.calls - 1
|
||||
if index >= len(validator.approved) {
|
||||
index = len(validator.approved) - 1
|
||||
}
|
||||
approved = validator.approved[index]
|
||||
}
|
||||
return contracts.RawValidationResult{
|
||||
Approved: approved,
|
||||
ReasonCode: validator.reason,
|
||||
Message: validator.message,
|
||||
Warnings: validator.warnings,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (validator *runnerValidator) Name() string {
|
||||
return validator.name
|
||||
}
|
||||
@@ -1668,6 +1951,13 @@ func sourceChunkWithID(id string, index int) contracts.SourceChunk {
|
||||
}
|
||||
}
|
||||
|
||||
func sourceChunkWithContent(id string, index int, content []byte, mediaType string) contracts.SourceChunk {
|
||||
chunk := sourceChunkWithID(id, index)
|
||||
chunk.Content = append([]byte(nil), content...)
|
||||
chunk.MediaType = mediaType
|
||||
return chunk
|
||||
}
|
||||
|
||||
func unitWithID(id string) source.SourceUnit {
|
||||
switch id {
|
||||
case "u1":
|
||||
@@ -1723,3 +2013,13 @@ func assertRunError(t *testing.T, err error, want string) {
|
||||
t.Fatalf("Run() error = %q, want substring %q", err.Error(), want)
|
||||
}
|
||||
}
|
||||
|
||||
func rawValidationRegistry(t *testing.T, stage ModuleStage, module string, validators ...contracts.RawValidator) *RawValidationRegistry {
|
||||
t.Helper()
|
||||
|
||||
registry := NewRawValidationRegistry()
|
||||
if err := registry.Register(stage, module, validators...); err != nil {
|
||||
t.Fatalf("register raw validators: %v", err)
|
||||
}
|
||||
return registry
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user