Make module-stage LLM handling resilient and report warnings

This commit is contained in:
2026-05-23 10:07:06 -05:00
parent a84941d681
commit a3655f5540
43 changed files with 856 additions and 217 deletions

View File

@@ -15,6 +15,7 @@ import (
"gitea.maximumdirect.net/eric/audita/internal/framework/llm"
"gitea.maximumdirect.net/eric/audita/internal/framework/proposals"
"gitea.maximumdirect.net/eric/audita/internal/framework/validators"
stagewarnings "gitea.maximumdirect.net/eric/audita/internal/framework/warnings"
validatormetadata "gitea.maximumdirect.net/eric/audita/internal/validators/metadata"
)
@@ -37,18 +38,19 @@ type ValidationScheduler = contracts.LLMScheduler
// ModuleResult captures deterministic per-module execution output.
type ModuleResult struct {
ModuleKey string `json:"module_key"`
ModuleInstance string `json:"module_instance"`
ReplacementPolicy proposals.ReplacementPolicy `json:"replacement_policy"`
Status string `json:"status"`
ProposalCount int `json:"proposal_count"`
ValidatorDecisions []ValidatorDecisionRecord `json:"validator_decisions,omitempty"`
ValidatorRejected []ValidatorRejectedChange `json:"validator_rejected,omitempty"`
AppliedChanges []proposals.AppliedChange `json:"applied_changes,omitempty"`
SkippedChanges []proposals.SkippedChange `json:"skipped_changes,omitempty"`
ErrorMessage string `json:"error_message,omitempty"`
StartedAt time.Time `json:"started_at"`
CompletedAt time.Time `json:"completed_at"`
ModuleKey string `json:"module_key"`
ModuleInstance string `json:"module_instance"`
ReplacementPolicy proposals.ReplacementPolicy `json:"replacement_policy"`
Status string `json:"status"`
ProposalCount int `json:"proposal_count"`
Warnings []stagewarnings.StageWarning `json:"warnings,omitempty"`
ValidatorDecisions []ValidatorDecisionRecord `json:"validator_decisions,omitempty"`
ValidatorRejected []ValidatorRejectedChange `json:"validator_rejected,omitempty"`
AppliedChanges []proposals.AppliedChange `json:"applied_changes,omitempty"`
SkippedChanges []proposals.SkippedChange `json:"skipped_changes,omitempty"`
ErrorMessage string `json:"error_message,omitempty"`
StartedAt time.Time `json:"started_at"`
CompletedAt time.Time `json:"completed_at"`
}
type ValidatorDecisionRecord struct {
@@ -176,6 +178,7 @@ func (r *Runner) Run(ctx context.Context, input RunInput) (RunOutput, error) {
ReplacementPolicy: policy,
Status: ModuleStatusFailed,
ProposalCount: pipelineResult.ProposalCount,
Warnings: pipelineResult.Warnings,
ValidatorDecisions: pipelineResult.ValidatorDecisions,
ValidatorRejected: pipelineResult.ValidatorRejected,
ErrorMessage: pipelineErr.Error(),
@@ -199,6 +202,7 @@ func (r *Runner) Run(ctx context.Context, input RunInput) (RunOutput, error) {
ReplacementPolicy: policy,
Status: ModuleStatusSuccess,
ProposalCount: pipelineResult.ProposalCount,
Warnings: pipelineResult.Warnings,
ValidatorDecisions: pipelineResult.ValidatorDecisions,
ValidatorRejected: pipelineResult.ValidatorRejected,
AppliedChanges: applyResult.Applied,
@@ -232,12 +236,14 @@ type collectSectionProposalsInput struct {
type sectionProposals struct {
meta contracts.SectionMetadata
corrected []proposals.CorrectionProposal
warnings []stagewarnings.StageWarning
}
type sectionProposalResult struct {
sectionPos int
section chunking.Section
corrected []proposals.CorrectionProposal
warnings []stagewarnings.StageWarning
err error
}
@@ -247,12 +253,14 @@ type sectionValidationResult struct {
approved []proposals.EnrichedCorrectionProposal
decisions []ValidatorDecisionRecord
rejected []ValidatorRejectedChange
warnings []stagewarnings.StageWarning
err error
}
type modulePipelineResult struct {
ProposalCount int
Approved []proposals.EnrichedCorrectionProposal
Warnings []stagewarnings.StageWarning
ValidatorDecisions []ValidatorDecisionRecord
ValidatorRejected []ValidatorRejectedChange
}
@@ -271,7 +279,7 @@ func collectSectionProposals(ctx context.Context, input collectSectionProposalsI
go func() {
defer wg.Done()
meta := contracts.SectionMetadataFromSection(section)
corrected, err := input.Module.Propose(withModuleInstanceContext(runCtx, input.Spec.InstanceName), contracts.ProposalRequest{
proposalResult, err := input.Module.Propose(withModuleInstanceContext(runCtx, input.Spec.InstanceName), contracts.ProposalRequest{
ExecutionContext: contracts.ExecutionContext{
Config: input.Config,
WorkingTranscript: transcriptFromSection(section),
@@ -291,7 +299,8 @@ func collectSectionProposals(ctx context.Context, input collectSectionProposalsI
case results <- sectionProposalResult{
sectionPos: sectionPos,
section: section,
corrected: corrected,
corrected: proposalResult.Proposals,
warnings: proposalResult.Warnings,
err: err,
}:
case <-runCtx.Done():
@@ -308,6 +317,7 @@ func collectSectionProposals(ctx context.Context, input collectSectionProposalsI
func runModulePipeline(ctx context.Context, input collectSectionProposalsInput) (modulePipelineResult, error) {
out := modulePipelineResult{
Approved: make([]proposals.EnrichedCorrectionProposal, 0),
Warnings: make([]stagewarnings.StageWarning, 0),
ValidatorDecisions: make([]ValidatorDecisionRecord, 0),
ValidatorRejected: make([]ValidatorRejectedChange, 0),
}
@@ -352,6 +362,7 @@ func runModulePipeline(ctx context.Context, input collectSectionProposalsInput)
if firstErr != nil {
continue
}
out.Warnings = append(out.Warnings, result.warnings...)
pending[result.sectionPos] = result
for {
@@ -401,6 +412,7 @@ func runModulePipeline(ctx context.Context, input collectSectionProposalsInput)
approved: validated.approved,
decisions: validated.decisions,
rejected: validated.rejected,
warnings: validated.warnings,
err: err,
}
}(nextSectionToProcess, sectionEnriched, sectionMeta)
@@ -424,6 +436,7 @@ func runModulePipeline(ctx context.Context, input collectSectionProposalsInput)
break
}
out.Approved = append(out.Approved, res.approved...)
out.Warnings = append(out.Warnings, res.warnings...)
out.ValidatorDecisions = append(out.ValidatorDecisions, res.decisions...)
out.ValidatorRejected = append(out.ValidatorRejected, res.rejected...)
}
@@ -440,6 +453,38 @@ func runModulePipeline(ctx context.Context, input collectSectionProposalsInput)
}
return validatorOrder[out.ValidatorRejected[i].ValidatorName] < validatorOrder[out.ValidatorRejected[j].ValidatorName]
})
sort.SliceStable(out.Warnings, func(i, j int) bool {
leftSection, rightSection := -1, -1
if out.Warnings[i].SectionIndex != nil {
leftSection = *out.Warnings[i].SectionIndex
}
if out.Warnings[j].SectionIndex != nil {
rightSection = *out.Warnings[j].SectionIndex
}
if leftSection != rightSection {
return leftSection < rightSection
}
leftBatch, rightBatch := -1, -1
if out.Warnings[i].BatchIndex != nil {
leftBatch = *out.Warnings[i].BatchIndex
}
if out.Warnings[j].BatchIndex != nil {
rightBatch = *out.Warnings[j].BatchIndex
}
if leftBatch != rightBatch {
return leftBatch < rightBatch
}
if out.Warnings[i].ValidatorName != out.Warnings[j].ValidatorName {
return validatorOrder[out.Warnings[i].ValidatorName] < validatorOrder[out.Warnings[j].ValidatorName]
}
if out.Warnings[i].Scope != out.Warnings[j].Scope {
return out.Warnings[i].Scope < out.Warnings[j].Scope
}
if out.Warnings[i].ReasonCode != out.Warnings[j].ReasonCode {
return out.Warnings[i].ReasonCode < out.Warnings[j].ReasonCode
}
return out.Warnings[i].Message < out.Warnings[j].Message
})
if firstErr != nil {
return out, firstErr
@@ -467,11 +512,13 @@ type validateSectionCandidatesResult struct {
approved []proposals.EnrichedCorrectionProposal
decisions []ValidatorDecisionRecord
rejected []ValidatorRejectedChange
warnings []stagewarnings.StageWarning
}
func validateSectionCandidates(ctx context.Context, input validateSectionCandidatesInput) (validateSectionCandidatesResult, error) {
decisions := make([]ValidatorDecisionRecord, 0)
rejected := make([]ValidatorRejectedChange, 0)
warnings := make([]stagewarnings.StageWarning, 0)
eligible := append([]proposals.EnrichedCorrectionProposal(nil), input.SectionEnriched...)
for _, validator := range input.Validators {
@@ -512,6 +559,7 @@ func validateSectionCandidates(ctx context.Context, input validateSectionCandida
approved: eligible,
decisions: decisions,
rejected: rejected,
warnings: warnings,
}, fmt.Errorf("validator %q failed: %w", validator.Name(), err)
}
if err := validators.EnforceDecisionCardinality(eligible, vResult.Decisions); err != nil {
@@ -519,8 +567,10 @@ func validateSectionCandidates(ctx context.Context, input validateSectionCandida
approved: eligible,
decisions: decisions,
rejected: rejected,
warnings: warnings,
}, fmt.Errorf("validator %q cardinality failed: %w", validator.Name(), err)
}
warnings = append(warnings, vResult.Warnings...)
nextEligible := make([]proposals.EnrichedCorrectionProposal, 0, len(eligible))
byIndex := make(map[int]proposals.EnrichedCorrectionProposal, len(eligible))
@@ -562,6 +612,7 @@ func validateSectionCandidates(ctx context.Context, input validateSectionCandida
approved: eligible,
decisions: decisions,
rejected: rejected,
warnings: warnings,
}, nil
}

View File

@@ -47,11 +47,12 @@ type fakeModule struct {
func (m fakeModule) Key() string { return m.key }
func (m fakeModule) ReplacementPolicy() proposals.ReplacementPolicy { return m.policy }
func (m fakeModule) Validators() []contracts.Validator { return m.validators }
func (m fakeModule) Propose(ctx context.Context, req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
func (m fakeModule) Propose(ctx context.Context, req contracts.ProposalRequest) (contracts.ProposalResult, error) {
if m.proposeF == nil {
return nil, nil
return contracts.ProposalResult{}, nil
}
return m.proposeF(req)
proposalsOut, err := m.proposeF(req)
return contracts.ProposalResult{Proposals: proposalsOut}, err
}
type fakeValidator struct {
@@ -1191,7 +1192,7 @@ func TestRunnerLLMValidatorRejectionPreventsApplication(t *testing.T) {
}
}
func TestRunnerLLMValidatorMalformedResponseFailsWithPartialProgress(t *testing.T) {
func TestRunnerLLMValidatorMalformedResponseRejectsBatchAndKeepsPartialProgress(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 99, Approved: true, Confidence: 0.9, Reason: "bad index"}}}}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
@@ -1209,15 +1210,21 @@ func TestRunnerLLMValidatorMalformedResponseFailsWithPartialProgress(t *testing.
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m1", InstanceName: "m1"}, {ModuleKey: "m2", InstanceName: "m2"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected llm validator failure")
if err != nil {
t.Fatalf("expected malformed validator response to downgrade, got %v", err)
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected partial progress retained")
}
if len(out.ModuleResults) != 2 || len(out.ModuleResults[1].ValidatorRejected) != 1 {
t.Fatalf("expected second module rejection, got %+v", out.ModuleResults)
}
if len(out.ModuleResults[1].Warnings) != 1 || out.ModuleResults[1].Warnings[0].ReasonCode != validators.ReasonValidatorMalformed {
t.Fatalf("expected malformed warning, got %+v", out.ModuleResults[1].Warnings)
}
}
func TestRunnerLLMValidatorMissingDecisionFails(t *testing.T) {
func TestRunnerLLMValidatorMissingDecisionRejectsBatch(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{}}}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
@@ -1232,12 +1239,12 @@ func TestRunnerLLMValidatorMissingDecisionFails(t *testing.T) {
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected missing decision failure")
if err != nil {
t.Fatalf("expected missing decision downgrade, got %v", err)
}
}
func TestRunnerLLMValidatorDuplicateDecisionFails(t *testing.T) {
func TestRunnerLLMValidatorDuplicateDecisionRejectsBatch(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{
{CorrectionIndex: 0, Approved: true, Confidence: 0.9, Reason: "ok"},
{CorrectionIndex: 0, Approved: false, Confidence: 0.9, Reason: "dup"},
@@ -1255,8 +1262,8 @@ func TestRunnerLLMValidatorDuplicateDecisionFails(t *testing.T) {
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected duplicate decision failure")
if err != nil {
t.Fatalf("expected duplicate decision downgrade, got %v", err)
}
}
@@ -1355,7 +1362,7 @@ type proposalGenerationModule struct {
func (m proposalGenerationModule) Key() string { return m.key }
func (m proposalGenerationModule) ReplacementPolicy() proposals.ReplacementPolicy { return m.policy }
func (m proposalGenerationModule) Validators() []contracts.Validator { return nil }
func (m proposalGenerationModule) Propose(ctx context.Context, req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
func (m proposalGenerationModule) Propose(ctx context.Context, req contracts.ProposalRequest) (contracts.ProposalResult, error) {
result, err := proposal_generation.GenerateCandidates(ctx, proposal_generation.Request{
ModuleKey: req.RunSpec.ModuleKey,
ModuleInstance: req.RunSpec.InstanceName,
@@ -1372,9 +1379,9 @@ func (m proposalGenerationModule) Propose(ctx context.Context, req contracts.Pro
DiagnosticsDir: req.DiagnosticsDir,
})
if err != nil {
return nil, err
return contracts.ProposalResult{}, err
}
return result.Corrections, nil
return contracts.ProposalResult{Proposals: result.Corrections, Warnings: result.Warnings}, nil
}
type fakeProposalStructuredClient struct {