Files
audita/internal/framework/runner/runner_test.go

689 lines
33 KiB
Go

package runner
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"gitea.maximumdirect.net/eric/audita/internal/core/config"
"gitea.maximumdirect.net/eric/audita/internal/core/schema"
"gitea.maximumdirect.net/eric/audita/internal/framework/contracts"
"gitea.maximumdirect.net/eric/audita/internal/framework/llm"
"gitea.maximumdirect.net/eric/audita/internal/framework/modules"
"gitea.maximumdirect.net/eric/audita/internal/framework/proposal_generation"
"gitea.maximumdirect.net/eric/audita/internal/framework/proposals"
"gitea.maximumdirect.net/eric/audita/internal/framework/validators"
)
type fakeFactory struct {
modules map[string]contracts.TranscriptModule
}
func (f fakeFactory) ModuleForSpec(spec contracts.ModuleRunSpec) (contracts.TranscriptModule, error) {
m, ok := f.modules[spec.InstanceName]
if !ok {
return nil, errors.New("module not registered")
}
return m, nil
}
type fakeModule struct {
key string
policy proposals.ReplacementPolicy
validators []contracts.Validator
proposeF func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error)
}
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) {
if m.proposeF == nil {
return nil, nil
}
return m.proposeF(req)
}
type fakeValidator struct {
name string
validateF func(req contracts.ValidationRequest) (validators.Result, error)
}
func (v fakeValidator) Name() string { return v.name }
func (v fakeValidator) Validate(ctx context.Context, req contracts.ValidationRequest) (validators.Result, error) {
_ = ctx
return v.validateF(req)
}
func TestRunnerOneModuleAppliesProposal(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "teh cat"}}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"grammar": fakeModule{key: "grammar", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 0.9}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "grammar", InstanceName: "grammar"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected corrected text, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults) != 1 || len(out.ModuleResults[0].AppliedChanges) != 1 {
t.Fatalf("expected one module with one applied change, got %+v", out.ModuleResults)
}
}
func TestRunnerModulesRunSequentially(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "teh cat"}}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m1": fakeModule{key: "m1", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}},
"m2": fakeModule{key: "m2", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
if req.WorkingTranscript.Segments[0].Text != "the cat" {
t.Fatalf("second module did not see first module changes: %q", req.WorkingTranscript.Segments[0].Text)
}
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "cat", CorrectedText: "dog", Confidence: 1}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m1", InstanceName: "m1"}, {ModuleKey: "m2", InstanceName: "m2"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "the dog" {
t.Fatalf("expected sequential updates, got %q", out.FinalTranscript.Segments[0].Text)
}
}
func TestRunnerSkippedRecorded(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "word word"}}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "word", CorrectedText: "term", Confidence: 1}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if len(out.ModuleResults[0].SkippedChanges) != 1 {
t.Fatalf("expected one skipped change, got %+v", out.ModuleResults[0].SkippedChanges)
}
}
func TestRunnerFailureReturnsPartialProgress(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "teh cat"}}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m1": fakeModule{key: "m1", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}},
"m2": fakeModule{key: "m2", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return nil, errors.New("boom")
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m1", InstanceName: "m1"}, {ModuleKey: "m2", InstanceName: "m2"}}})
if err == nil {
t.Fatal("expected error")
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected partial transcript progress preserved, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults) != 2 || out.ModuleResults[1].Status != ModuleStatusFailed {
t.Fatalf("expected second module failed in results, got %+v", out.ModuleResults)
}
}
func TestRunnerDoesNotMutateInputTranscript(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "teh"}}}
before := transcript.Segments[0].Text
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}},
}})
_, err := r.Run(context.Background(), RunInput{Config: ptrConfig(config.Default()), Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if transcript.Segments[0].Text != before {
t.Fatalf("expected input transcript unchanged, got %q", transcript.Segments[0].Text)
}
}
func TestRunnerRepeatedModuleInstanceNames(t *testing.T) {
specs, err := contracts.ResolveModuleRunSpecs([]string{"glossary", "glossary"})
if err != nil {
t.Fatalf("ResolveModuleRunSpecs error: %v", err)
}
if specs[0].InstanceName != "glossary_1" || specs[1].InstanceName != "glossary_2" {
t.Fatalf("unexpected instance names: %+v", specs)
}
}
func TestRunnerValidatorApprovedProposalApplied(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "teh cat"}}}
allowAll := fakeValidator{name: "allow", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
decisions := make([]validators.Decision, len(req.CandidateProposal))
for i, p := range req.CandidateProposal {
decisions[i] = validators.Decision{ProposalIndex: p.ProposalIndex, Approved: true, ReasonCode: validators.ReasonApproved, Message: "approved"}
}
return validators.Result{ValidatorName: "allow", Decisions: decisions}, nil
}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{
key: "m",
policy: proposals.ReplacementPolicyRequireUnique,
validators: []contracts.Validator{allowAll},
proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
},
},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected proposal applied, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults[0].ValidatorDecisions) != 1 {
t.Fatalf("expected validator decisions recorded")
}
}
func TestRunnerValidatorRejectedProposalNotApplied(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "teh cat"}}}
rejectAll := fakeValidator{name: "reject", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
decisions := make([]validators.Decision, len(req.CandidateProposal))
for i, p := range req.CandidateProposal {
decisions[i] = validators.Decision{ProposalIndex: p.ProposalIndex, Approved: false, ReasonCode: validators.ReasonNoEffect, Message: "rejected"}
}
return validators.Result{ValidatorName: "reject", Decisions: decisions}, nil
}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{rejectAll}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "teh cat" {
t.Fatalf("expected rejected proposal not applied, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults[0].ValidatorRejected) != 1 {
t.Fatalf("expected validator rejection recorded, got %+v", out.ModuleResults[0].ValidatorRejected)
}
}
func TestRunnerMultipleValidatorsRunInOrderAndFilterSurvivors(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "one two"}}}
first := fakeValidator{name: "first", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
if len(req.CandidateProposal) != 2 {
t.Fatalf("expected first validator to see 2 candidates, got %d", len(req.CandidateProposal))
}
return validators.Result{
ValidatorName: "first",
Decisions: []validators.Decision{
{ProposalIndex: req.CandidateProposal[0].ProposalIndex, Approved: true, ReasonCode: validators.ReasonApproved, Message: "ok"},
{ProposalIndex: req.CandidateProposal[1].ProposalIndex, Approved: false, ReasonCode: validators.ReasonNoEffect, Message: "reject"},
},
}, nil
}}
second := fakeValidator{name: "second", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
if len(req.CandidateProposal) != 1 {
t.Fatalf("expected second validator to see only survivors, got %d", len(req.CandidateProposal))
}
return validators.Result{ValidatorName: "second", Decisions: []validators.Decision{{ProposalIndex: req.CandidateProposal[0].ProposalIndex, Approved: true, ReasonCode: validators.ReasonApproved, Message: "ok"}}}, nil
}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{first, second}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{
{TargetSegmentID: 1, OriginalText: "one", CorrectedText: "ONE", Confidence: 1},
{TargetSegmentID: 1, OriginalText: "two", CorrectedText: "TWO", Confidence: 1},
}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "ONE two" {
t.Fatalf("expected only survivor applied, got %q", out.FinalTranscript.Segments[0].Text)
}
}
func TestRunnerValidatorCardinalityErrorStopsPipelineWithPartialProgress(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "teh cat"}}}
good := fakeModule{key: "m1", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}}
badValidator := fakeValidator{name: "bad", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
// Missing one decision triggers cardinality error.
return validators.Result{ValidatorName: "bad", Decisions: nil}, nil
}}
bad := fakeModule{key: "m2", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{badValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "cat", CorrectedText: "dog", Confidence: 1}}, nil
}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{"m1": good, "m2": bad}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m1", InstanceName: "m1"}, {ModuleKey: "m2", InstanceName: "m2"}}})
if err == nil {
t.Fatal("expected cardinality error")
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected partial progress preserved, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults) != 2 || out.ModuleResults[1].Status != ModuleStatusFailed {
t.Fatalf("expected second module failed")
}
}
func TestRunnerApplicationSkipAfterValidatorApprovalReported(t *testing.T) {
transcript := &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "word word"}}}
allow := fakeValidator{name: "allow", validateF: func(req contracts.ValidationRequest) (validators.Result, error) {
return validators.Result{ValidatorName: "allow", Decisions: []validators.Decision{{ProposalIndex: req.CandidateProposal[0].ProposalIndex, Approved: true, ReasonCode: validators.ReasonApproved, Message: "ok"}}}, nil
}}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{allow}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "word", CorrectedText: "term", Confidence: 1}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{Transcript: transcript, ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}}})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if len(out.ModuleResults[0].SkippedChanges) != 1 {
t.Fatalf("expected application skip recorded")
}
if len(out.ModuleResults[0].ValidatorRejected) != 0 {
t.Fatalf("expected no validator rejection")
}
}
func ptrConfig(c config.Config) *config.Config { return &c }
type fakeStructuredClient struct {
responses []validators.LLMValidationResponse
err error
calls int
}
func (f *fakeStructuredClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
_ = ctx
_ = req
f.calls++
if f.err != nil {
return contracts.StructuredCompletionResponse{}, f.err
}
if len(f.responses) == 0 {
return contracts.StructuredCompletionResponse{}, errors.New("unexpected call")
}
resp := f.responses[0]
f.responses = f.responses[1:]
target := out.(*validators.LLMValidationResponse)
*target = resp
return contracts.StructuredCompletionResponse{}, nil
}
type countingScheduler struct{ runs int }
func (s *countingScheduler) Run(ctx context.Context, fn func(context.Context) error) error {
s.runs++
return fn(ctx)
}
type lenEstimator struct{}
func (lenEstimator) EstimateTokens(text string) int { return len(text) }
func TestRunnerLLMValidatorApprovalApplied(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 0, Approved: true, Confidence: 0.9, Reason: "ok"}}}}}
llmValidator, err := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
if err != nil {
t.Fatalf("NewLLMBackedValidator: %v", err)
}
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "gestures", CorrectedText: "Jesters", Confidence: 1}}, nil
}},
}})
cfg := config.Default()
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "There were gestures at the temple."}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "There were Jesters at the temple." {
t.Fatalf("expected applied LLM-approved proposal, got %q", out.FinalTranscript.Segments[0].Text)
}
}
func TestRunnerLLMValidatorRejectionPreventsApplication(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 0, Approved: false, Confidence: 0.95, Reason: "reject"}}}}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "gestures", CorrectedText: "Jesters", Confidence: 1}}, nil
}},
}})
cfg := config.Default()
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "There were gestures at the temple."}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "There were gestures at the temple." {
t.Fatalf("expected rejected proposal not applied")
}
if len(out.ModuleResults[0].ValidatorRejected) != 1 {
t.Fatalf("expected validator rejection record")
}
}
func TestRunnerLLMValidatorMalformedResponseFailsWithPartialProgress(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{
"m1": fakeModule{key: "m1", policy: proposals.ReplacementPolicyRequireUnique, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 1}}, nil
}},
"m2": fakeModule{key: "m2", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "cat", CorrectedText: "dog", Confidence: 1}}, nil
}},
}})
cfg := config.Default()
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "teh cat"}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m1", InstanceName: "m1"}, {ModuleKey: "m2", InstanceName: "m2"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected llm validator failure")
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected partial progress retained")
}
}
func TestRunnerLLMValidatorMissingDecisionFails(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{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "gestures", CorrectedText: "Jesters", Confidence: 1}}, nil
}},
}})
cfg := config.Default()
_, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "There were gestures at the temple."}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected missing decision failure")
}
}
func TestRunnerLLMValidatorDuplicateDecisionFails(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"},
}}}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "gestures", CorrectedText: "Jesters", Confidence: 1}}, nil
}},
}})
cfg := config.Default()
_, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "There were gestures at the temple."}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
})
if err == nil {
t.Fatal("expected duplicate decision failure")
}
}
func TestRunnerLLMValidatorBatchingAndSchedulerUsage(t *testing.T) {
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{
{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 0, Approved: true, Confidence: 0.9, Reason: "ok"}}},
{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 1, Approved: true, Confidence: 0.9, Reason: "ok"}}},
}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
llmValidator.SetTokenEstimator(lenEstimator{})
scheduler := &countingScheduler{}
cfg := config.Default()
cfg.ValidationMaxPromptTokens = 260
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyReplaceAll, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{
{TargetSegmentID: 1, OriginalText: "alpha", CorrectedText: strings.Repeat("B", 40), Confidence: 1},
{TargetSegmentID: 1, OriginalText: "gamma", CorrectedText: strings.Repeat("D", 40), Confidence: 1},
}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "alpha gamma"}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
ValidationLLMScheduler: scheduler,
})
if err != nil {
t.Fatalf("unexpected run error: %v", err)
}
if scheduler.runs < 2 {
t.Fatalf("expected scheduler to run per batch, got %d", scheduler.runs)
}
if client.calls < 2 {
t.Fatalf("expected multiple llm calls for batching, got %d", client.calls)
}
if len(out.ModuleResults[0].AppliedChanges) != 2 {
t.Fatalf("expected both approved proposals applied")
}
}
func TestRunnerLLMValidatorDiagnosticsWrittenAndRedacted(t *testing.T) {
secret := "super-secret-key"
client := &fakeStructuredClient{responses: []validators.LLMValidationResponse{{Validations: []validators.LLMValidationDecision{{CorrectionIndex: 0, Approved: true, Confidence: 0.9, Reason: secret}}}}}
llmValidator, _ := validators.NewLLMBackedValidator("spoken_form_plausibility_review", validators.LLMValidatorTypeSpokenFormPlausibility, "")
cfg := config.Default()
cfg.PrimaryLLM.APIKey = secret
cfg.ValidationLLM.APIKey = secret
diagDir := t.TempDir()
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": fakeModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique, validators: []contracts.Validator{llmValidator}, proposeF: func(req contracts.ProposalRequest) ([]proposals.CorrectionProposal, error) {
return []proposals.CorrectionProposal{{TargetSegmentID: 1, OriginalText: "gestures", CorrectedText: "Jesters", Confidence: 1}}, nil
}},
}})
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "There were gestures at the temple."}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ValidationLLMClient: client,
ValidationDiagnosticsDir: diagDir,
})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if len(out.ModuleResults[0].ValidatorDecisions) == 0 || out.ModuleResults[0].ValidatorDecisions[0].DiagnosticArtifactPath == "" {
t.Fatalf("expected diagnostic artifact path on decision")
}
raw, readErr := os.ReadFile(out.ModuleResults[0].ValidatorDecisions[0].DiagnosticArtifactPath)
if readErr != nil {
t.Fatalf("read diagnostic: %v", readErr)
}
if strings.Contains(string(raw), secret) {
t.Fatalf("secret leaked in diagnostics: %s", string(raw))
}
if !strings.Contains(string(raw), "[REDACTED]") {
t.Fatalf("expected redaction marker in diagnostics")
}
}
func TestRunnerAcceptsLLMSchedulerType(t *testing.T) {
s, err := llm.NewScheduler(1)
if err != nil {
t.Fatalf("NewScheduler: %v", err)
}
if s == nil {
t.Fatal("expected scheduler instance")
}
}
type proposalGenerationModule struct {
key string
policy proposals.ReplacementPolicy
}
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) {
result, err := proposal_generation.GenerateCandidates(ctx, proposal_generation.Request{
ModuleKey: req.RunSpec.ModuleKey,
ModuleInstance: req.RunSpec.InstanceName,
ReplacementPolicy: req.RunSpec.ReplacementPolicy,
WorkingTranscript: req.WorkingTranscript,
Config: req.Config,
Glossary: req.Glossary,
Messages: []contracts.LLMMessage{
{Role: "system", Content: "return transcript corrections"},
{Role: "user", Content: "produce one safe correction"},
},
LLMClient: req.LLMClient,
Scheduler: req.LLMScheduler,
DiagnosticsDir: req.DiagnosticsDir,
})
if err != nil {
return nil, err
}
return result.Corrections, nil
}
type fakeProposalStructuredClient struct {
responses []proposal_generation.StructuredCorrectionSet
}
func (f *fakeProposalStructuredClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
_ = ctx
_ = req
if len(f.responses) == 0 {
return contracts.StructuredCompletionResponse{}, errors.New("unexpected call")
}
target, ok := out.(*proposal_generation.StructuredCorrectionSet)
if !ok {
return contracts.StructuredCompletionResponse{}, errors.New("unexpected output type")
}
*target = f.responses[0]
f.responses = f.responses[1:]
return contracts.StructuredCompletionResponse{}, nil
}
func TestRunnerProposalGenerationHelperFlowsThroughPipeline(t *testing.T) {
client := &fakeProposalStructuredClient{
responses: []proposal_generation.StructuredCorrectionSet{
{
Corrections: []proposal_generation.StructuredCorrectionProposal{
{TargetSegmentID: 1, OriginalText: "teh", CorrectedText: "the", Confidence: 0.9},
},
},
},
}
scheduler := &countingScheduler{}
cfg := config.Default()
diagDir := t.TempDir()
r := New(fakeFactory{modules: map[string]contracts.TranscriptModule{
"m": proposalGenerationModule{key: "m", policy: proposals.ReplacementPolicyRequireUnique},
}})
out, err := r.Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "teh cat"}}},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "m", InstanceName: "m"}},
ProposalLLMClient: client,
ProposalLLMScheduler: scheduler,
ProposalDiagnosticsDir: diagDir,
})
if err != nil {
t.Fatalf("unexpected run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "the cat" {
t.Fatalf("expected proposal-generated correction to apply, got %q", out.FinalTranscript.Segments[0].Text)
}
if scheduler.runs != 1 {
t.Fatalf("expected proposal scheduler use, got %d runs", scheduler.runs)
}
if len(out.ModuleResults) != 1 || len(out.ModuleResults[0].AppliedChanges) != 1 {
t.Fatalf("expected one applied change, got %+v", out.ModuleResults)
}
matches, globErr := filepath.Glob(filepath.Join(diagDir, "m", "*proposal-generation*response-payload.json"))
if globErr != nil {
t.Fatalf("glob diagnostics: %v", globErr)
}
if len(matches) == 0 {
t.Fatalf("expected proposal-generation diagnostics artifacts in %s", filepath.Join(diagDir, "m"))
}
}
func TestRunnerGrammarModuleUsesConfidenceThreshold(t *testing.T) {
client := &fakeProposalStructuredClient{
responses: []proposal_generation.StructuredCorrectionSet{
{
Corrections: []proposal_generation.StructuredCorrectionProposal{
{TargetSegmentID: 1, OriginalText: "hello ,world", CorrectedText: "Hello, world", Confidence: 0.5},
},
},
},
}
cfg := config.Default()
cfg.Thresholds.Grammar = 0.9
factory := modules.NewFactory(modules.Dependencies{})
out, err := New(factory).Run(context.Background(), RunInput{
Config: &cfg,
Transcript: &schema.Transcript{Segments: []schema.Segment{{ID: 1, Text: "hello ,world"}}},
Glossary: &schema.Glossary{},
ModuleSpecs: []contracts.ModuleRunSpec{{ModuleKey: "grammar", InstanceName: "grammar"}},
ProposalLLMClient: client,
})
if err != nil {
t.Fatalf("Run error: %v", err)
}
if out.FinalTranscript.Segments[0].Text != "hello ,world" {
t.Fatalf("expected no changes due to confidence threshold, got %q", out.FinalTranscript.Segments[0].Text)
}
if len(out.ModuleResults) != 1 || len(out.ModuleResults[0].ValidatorRejected) == 0 {
t.Fatalf("expected validator rejection, got %+v", out.ModuleResults)
}
if out.ModuleResults[0].ValidatorRejected[0].ReasonCode != validators.ReasonLowConfidence {
t.Fatalf("expected low confidence reason, got %+v", out.ModuleResults[0].ValidatorRejected[0])
}
}