Complete Phase 12 grammar module
This commit is contained in:
194
internal/modules/grammar/module_test.go
Normal file
194
internal/modules/grammar/module_test.go
Normal file
@@ -0,0 +1,194 @@
|
||||
package grammar
|
||||
|
||||
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/proposal_generation"
|
||||
"gitea.maximumdirect.net/eric/audita/internal/framework/proposals"
|
||||
)
|
||||
|
||||
type fakeLLMClient struct {
|
||||
responses []map[string]any
|
||||
err error
|
||||
calls []contracts.StructuredCompletionRequest
|
||||
}
|
||||
|
||||
func (f *fakeLLMClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
|
||||
_ = ctx
|
||||
f.calls = append(f.calls, req)
|
||||
if f.err != nil {
|
||||
return contracts.StructuredCompletionResponse{}, f.err
|
||||
}
|
||||
if len(f.responses) == 0 {
|
||||
return contracts.StructuredCompletionResponse{}, errors.New("unexpected call")
|
||||
}
|
||||
payload := f.responses[0]
|
||||
f.responses = f.responses[1:]
|
||||
|
||||
target, ok := out.(*proposal_generation.StructuredCorrectionSet)
|
||||
if !ok {
|
||||
return contracts.StructuredCompletionResponse{}, errors.New("unexpected output type")
|
||||
}
|
||||
items, _ := payload["corrections"].([]map[string]any)
|
||||
for _, item := range items {
|
||||
target.Corrections = append(target.Corrections, proposal_generation.StructuredCorrectionProposal{
|
||||
TargetSegmentID: item["id"].(int),
|
||||
OriginalText: item["original_text"].(string),
|
||||
CorrectedText: item["corrected_text"].(string),
|
||||
Confidence: item["confidence"].(float64),
|
||||
})
|
||||
}
|
||||
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)
|
||||
}
|
||||
|
||||
func tinyTranscript() *schema.Transcript {
|
||||
return &schema.Transcript{Segments: []schema.Segment{
|
||||
{ID: 1, Speaker: "A", Start: 0, End: 1, Text: "hello ,world"},
|
||||
}}
|
||||
}
|
||||
|
||||
func tinyGlossary() *schema.Glossary {
|
||||
return &schema.Glossary{Entries: []schema.GlossaryEntry{
|
||||
{Name: "Jesters", Category: "faction", Summary: "Faction", Aliases: []string{"Jester"}},
|
||||
}}
|
||||
}
|
||||
|
||||
func TestBuildProposalMessagesConstraints(t *testing.T) {
|
||||
msgs, err := BuildProposalMessages(tinyTranscript(), tinyGlossary())
|
||||
if err != nil {
|
||||
t.Fatalf("BuildProposalMessages error: %v", err)
|
||||
}
|
||||
if len(msgs) != 2 {
|
||||
t.Fatalf("expected 2 messages, got %d", len(msgs))
|
||||
}
|
||||
combined := msgs[0].Content + "\n" + msgs[1].Content
|
||||
for _, want := range []string{
|
||||
"Transcript section:",
|
||||
"Protected glossary/context:",
|
||||
"punctuation, capitalization, spacing, and article cleanup only",
|
||||
"Do not make word substitutions",
|
||||
"Do not return speaker, start, or end fields",
|
||||
`"id": 1`,
|
||||
} {
|
||||
if !strings.Contains(combined, want) {
|
||||
t.Fatalf("expected prompt to contain %q", want)
|
||||
}
|
||||
}
|
||||
if strings.Contains(strings.ToLower(combined), "summarize") {
|
||||
t.Fatalf("prompt should not invite summarization")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrammarModuleReplacementPolicy(t *testing.T) {
|
||||
m, err := New()
|
||||
if err != nil {
|
||||
t.Fatalf("New error: %v", err)
|
||||
}
|
||||
if m.ReplacementPolicy() != proposals.ReplacementPolicyRequireUnique {
|
||||
t.Fatalf("unexpected replacement policy: %q", m.ReplacementPolicy())
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrammarModuleValidatorChain(t *testing.T) {
|
||||
m, err := New()
|
||||
if err != nil {
|
||||
t.Fatalf("New error: %v", err)
|
||||
}
|
||||
got := make([]string, 0, len(m.Validators()))
|
||||
for _, v := range m.Validators() {
|
||||
got = append(got, v.Name())
|
||||
}
|
||||
want := []string{
|
||||
"no_effect",
|
||||
"original_text_presence",
|
||||
"confidence_threshold",
|
||||
"protected_glossary_terms",
|
||||
"non_empty_correction",
|
||||
"grammar_only_guard",
|
||||
"meaning_reversal_review",
|
||||
}
|
||||
if strings.Join(got, ",") != strings.Join(want, ",") {
|
||||
t.Fatalf("unexpected validator chain\n got: %v\nwant: %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrammarModuleProposeMapsCorrectionsAndWritesDiagnostics(t *testing.T) {
|
||||
secret := "super-secret-key"
|
||||
client := &fakeLLMClient{
|
||||
responses: []map[string]any{
|
||||
{
|
||||
"corrections": []map[string]any{
|
||||
{"id": 1, "original_text": "hello ,world", "corrected_text": "Hello, world", "confidence": 0.93},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
scheduler := &countingScheduler{}
|
||||
cfg := config.Default()
|
||||
cfg.PrimaryLLM.APIKey = secret
|
||||
module, err := New()
|
||||
if err != nil {
|
||||
t.Fatalf("New error: %v", err)
|
||||
}
|
||||
diagDir := t.TempDir()
|
||||
out, err := module.Propose(context.Background(), contracts.ProposalRequest{
|
||||
ExecutionContext: contracts.ExecutionContext{
|
||||
Config: &cfg,
|
||||
WorkingTranscript: tinyTranscript(),
|
||||
Glossary: tinyGlossary(),
|
||||
DiagnosticsDir: diagDir,
|
||||
},
|
||||
RunSpec: contracts.ModuleRunSpec{
|
||||
ModuleKey: "grammar",
|
||||
InstanceName: "grammar",
|
||||
ReplacementPolicy: proposals.ReplacementPolicyRequireUnique,
|
||||
},
|
||||
LLMClient: client,
|
||||
LLMScheduler: scheduler,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Propose error: %v", err)
|
||||
}
|
||||
if scheduler.runs != 1 {
|
||||
t.Fatalf("expected scheduler run count 1, got %d", scheduler.runs)
|
||||
}
|
||||
if len(client.calls) != 1 || client.calls[0].StageName != "grammar:proposal" {
|
||||
t.Fatalf("expected one grammar:proposal call, got %+v", client.calls)
|
||||
}
|
||||
if len(out) != 1 || out[0].CorrectedText != "Hello, world" {
|
||||
t.Fatalf("unexpected proposals: %+v", out)
|
||||
}
|
||||
diagFiles, globErr := filepath.Glob(filepath.Join(diagDir, "grammar", "*proposal*response-payload.json"))
|
||||
if globErr != nil {
|
||||
t.Fatalf("glob error: %v", globErr)
|
||||
}
|
||||
if len(diagFiles) == 0 {
|
||||
t.Fatalf("expected diagnostics response payload under %s", filepath.Join(diagDir, "grammar"))
|
||||
}
|
||||
for _, f := range diagFiles {
|
||||
raw, readErr := os.ReadFile(f)
|
||||
if readErr != nil {
|
||||
t.Fatalf("read diag %q: %v", f, readErr)
|
||||
}
|
||||
if strings.Contains(string(raw), secret) {
|
||||
t.Fatalf("secret leaked in diagnostics: %s", string(raw))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user