Files
notarius/internal/framework/pipeline/runner_rejection_warnings_test.go

149 lines
7.5 KiB
Go

package pipeline
import (
"context"
"fmt"
"reflect"
"strings"
"testing"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
)
type diagnosticChunker struct {
key string
plan source.ChunkPlan
calls int
}
func (c *diagnosticChunker) Key() string { return c.key }
func (*diagnosticChunker) ReferenceSlots() []contracts.ReferenceSlot { return nil }
func (c *diagnosticChunker) Plan(context.Context, contracts.ChunkRequest) (contracts.ChunkPlanResult, error) {
c.calls++
return contracts.ChunkPlanResult{
Plan: source.CloneChunkPlan(c.plan),
Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("operation-%d", c.calls), "operation diagnostic")},
}, nil
}
type chunkValidationFunc struct {
name string
validate func(context.Context, contracts.ChunkValidationRequest) (contracts.ValidationResult, error)
}
func (v chunkValidationFunc) Name() string { return v.name }
func (chunkValidationFunc) ExecutionClass() contracts.ExecutionClass {
return contracts.ExecutionClassDeterministic
}
func (v chunkValidationFunc) Validate(ctx context.Context, request contracts.ChunkValidationRequest) (contracts.ValidationResult, error) {
return v.validate(ctx, request)
}
func TestRunnerPromotesOnlyTerminalRejectionDiagnostics(t *testing.T) {
for _, target := range []ModuleStage{StageChunk, StageExtract, StageMerge, StageNormalize} {
t.Run(string(target), func(t *testing.T) {
prepared := preparedAttemptDebugPipeline(t)
lane := &prepared.Steps[0].lanes[0]
attempts := 0
first := func() contracts.ValidationResult {
return contracts.ValidationResult{Approved: true, Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("validator-%d", attempts), "validator diagnostic")}}
}
reject := func() contracts.ValidationResult {
return contracts.ValidationResult{Approved: false, ReasonCode: "rejected", Message: "rejected", CorrectionGuidance: "return an acceptable candidate"}
}
debug := newCapturedDebugRecorder()
recorder := &extractCaptureRecorder{CheckpointRecorder: NoopCheckpointRecorder()}
switch target {
case StageChunk:
prepared.resolved.ChunkValidationPolicy.SemanticRejection = SemanticRejectionRejectOutput
chunker := prepared.chunker.(*typedTestChunker)
prepared.chunker = &diagnosticChunker{key: prepared.resolved.Chunk.Module, plan: source.CloneChunkPlan(chunker.plan)}
prepared.resolved.Chunk.Retries = 1
prepared.chunkValidators.validators = []preparedValidator{
{resolved: ResolvedValidator{Binding: Binding("warning-approval"), Target: ValidatorTargetChunk}, chunk: chunkValidationFunc{name: "warning-approval", validate: func(context.Context, contracts.ChunkValidationRequest) (contracts.ValidationResult, error) {
return first(), nil
}}},
{resolved: ResolvedValidator{Binding: Binding("warning-rejection"), Target: ValidatorTargetChunk}, chunk: chunkValidationFunc{name: "warning-rejection", validate: func(context.Context, contracts.ChunkValidationRequest) (contracts.ValidationResult, error) {
return reject(), nil
}}},
}
chunkerWithDiagnostics := prepared.chunker.(*diagnosticChunker)
first = func() contracts.ValidationResult {
return contracts.ValidationResult{Approved: true, Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("validator-%d", chunkerWithDiagnostics.calls), "validator diagnostic")}}
}
case StageExtract:
lane.resolved.Extract.Retries = 1
lane.resolved.ExtractValidationPolicy.SemanticRejection = SemanticRejectionRejectOutput
installExtractOperation(prepared, 0, func(context.Context, contracts.TypedExtractionRequest) (erasedTypedResult, error) {
attempts++
return erasedTypedResult{Value: codecNotes{Items: []string{"extract"}}, Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("operation-%d", attempts), "operation diagnostic")}, ModelCandidate: attemptCandidate(t, fmt.Sprintf(`{"items":["%d"]}`, attempts))}, nil
})
lane.extractValidators.validators = rejectionDiagnosticTypedValidators(first, reject)
case StageMerge:
lane.resolved.Merge.Retries = 1
lane.resolved.MergeValidationPolicy.SemanticRejection = SemanticRejectionRejectOutput
lane.typed.merge = func(context.Context, any, contracts.TypedMergeRequest[any]) (erasedTypedResult, error) {
attempts++
return erasedTypedResult{Value: codecNotes{Items: []string{"merge"}}, Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("operation-%d", attempts), "operation diagnostic")}, ModelCandidate: attemptCandidate(t, fmt.Sprintf(`{"items":["%d"]}`, attempts))}, nil
}
lane.mergeValidators.validators = rejectionDiagnosticTypedValidators(first, reject)
case StageNormalize:
lane.resolved.Normalize.Retries = 1
lane.resolved.NormalizeValidationPolicy.SemanticRejection = SemanticRejectionRejectOutput
lane.typed.normalize = func(context.Context, any, contracts.TypedNormalizeRequest[any]) (erasedTypedResult, error) {
attempts++
return erasedTypedResult{Value: codecNotes{Items: []string{"normalize"}}, Diagnostics: []contracts.ProducerDiagnostic{producerDiagnostic(fmt.Sprintf("operation-%d", attempts), "operation diagnostic")}, ModelCandidate: attemptCandidate(t, fmt.Sprintf(`{"items":["%d"]}`, attempts))}, nil
}
lane.normalizeValidators.validators = rejectionDiagnosticTypedValidators(first, reject)
}
output, err := New().Run(context.Background(), RunInput{Prepared: prepared, RawInput: []byte("input"), Checkpoints: recorder, Debug: debug})
if err != nil {
t.Fatalf("Run() error = %v, want nil", err)
}
wantScopes := []string{"operation-2", "validator-2"}
wantAttempts := 2
if target == StageChunk {
wantScopes = []string{"operation-1", "validator-1"}
wantAttempts = 1
}
if got := rejectionDiagnosticScopes(output.Diagnostics.Groups); !reflect.DeepEqual(got, wantScopes) {
t.Fatalf("published diagnostic scopes = %#v, want %#v", got, wantScopes)
}
if len(output.Rejected) != 1 || output.Rejected[0].AttemptCount != wantAttempts {
t.Fatalf("rejections = %#v, want terminal rejection after %d attempt(s)", output.Rejected, wantAttempts)
}
attemptPath := fmt.Sprintf("%s/notes/attempt-01.json", target)
if target == StageChunk {
attemptPath = "chunk/attempt-01.json"
} else if target == StageExtract {
attemptPath = "extract/notes/chunk-000001/attempt-01.json"
}
if !strings.Contains(string(debug.json[attemptPath]), "rejection") {
t.Fatalf("first attempt debug = %s, want rejection", debug.json[attemptPath])
}
})
}
}
func rejectionDiagnosticTypedValidators(first func() contracts.ValidationResult, reject func() contracts.ValidationResult) []preparedValidator {
return []preparedValidator{
{resolved: ResolvedValidator{Binding: Binding("warning-approval"), Target: ValidatorTargetTyped, ArtifactKind: "test/notes"}, typedValidate: func(context.Context, any, typedValidationTarget) (contracts.ValidationResult, error) {
return first(), nil
}},
{resolved: ResolvedValidator{Binding: Binding("warning-rejection"), Target: ValidatorTargetTyped, ArtifactKind: "test/notes"}, typedValidate: func(context.Context, any, typedValidationTarget) (contracts.ValidationResult, error) {
return reject(), nil
}},
}
}
func rejectionDiagnosticScopes(groups []contracts.DiagnosticGroup) []string {
scopes := make([]string, len(groups))
for index := range groups {
scopes[index] = groups[index].Samples[0].Scope
}
return scopes
}