From 445329de81b2a7c6d89e88ed2a61b46e73c897fd Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 21 Apr 2026 15:42:09 -0500 Subject: [PATCH] Implemented a second LLM stage for grammatical review --- README.md | 8 +- pyproject.toml | 2 +- src/audita/cli.py | 16 ++- src/audita/config.py | 46 +++++-- src/audita/corrections.py | 12 +- src/audita/passes.py | 24 +++- src/audita/pipeline.py | 250 +++++++++++++++++++++++++------------- src/audita/prompts.py | 46 +++++++ tests/test_cli.py | 7 +- tests/test_config.py | 61 ++++++++-- tests/test_corrections.py | 32 +++++ tests/test_pipeline.py | 184 ++++++++++++++++++++++++++-- tests/test_prompts.py | 61 +++++++++- 13 files changed, 620 insertions(+), 129 deletions(-) diff --git a/README.md b/README.md index f311531..29340fe 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # Audita -Audita takes raw audio transcripts and uses an LLM to identify and fix misheard words, jargon, and domain-specific terms. +Audita takes raw audio transcripts and uses an LLM to identify and fix misheard words, jargon, domain-specific terms, and conservative readability issues. ## Development @@ -42,9 +42,11 @@ Useful configuration can be supplied by CLI flag or environment variable: - `AUDITA_MODEL`, default `openrouter/mistralai/mistral-small-3.2-24b-instruct` - `AUDITA_BASE_URL`, default `https://openrouter.ai/api/v1` - `AUDITA_MAX_SECTION_TOKENS`, default `16000` -- `AUDITA_CONFIDENCE_THRESHOLD`, default `0.60` +- `AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD`, default `0.60` +- `AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD`, default `0.60` - `AUDITA_MAX_RETRIES`, default `3` - `AUDITA_GLOSSARY_MAX_LLM_PASSES`, default `3`, for total glossary correction passes +- `AUDITA_GRAMMAR_MAX_LLM_PASSES`, default `3`, for total grammar/readability correction passes - `AUDITA_WORK_DIR`, default `/tmp/audita` -`AUDITA_WORK_DIR` stores per-run diagnostics while processing. Successful runs clean up their run directory unless corrections are skipped for target mismatches; failed runs and skipped-correction runs preserve diagnostics for debugging. +`AUDITA_WORK_DIR` stores per-run diagnostics while processing. Successful runs clean up their run directory unless corrections are skipped; failed runs and skipped-correction runs preserve diagnostics for debugging. diff --git a/pyproject.toml b/pyproject.toml index 1489bb4..0132a6a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "audita" version = "0.1.0" -description = "Correct audio transcripts with glossary-guided LLM passes." +description = "Correct audio transcripts with staged LLM review passes." readme = "README.md" requires-python = ">=3.9" license = { text = "BSD-3-Clause" } diff --git a/src/audita/cli.py b/src/audita/cli.py index f78367f..1842518 100644 --- a/src/audita/cli.py +++ b/src/audita/cli.py @@ -32,9 +32,19 @@ def _build_parser() -> argparse.ArgumentParser: process.add_argument("--model", help="OpenRouter model to use") process.add_argument("--base-url", help="OpenAI-compatible API base URL") process.add_argument("--max-section-tokens", type=int, help="maximum estimated tokens per transcript section") - process.add_argument("--confidence-threshold", type=float, help="minimum confidence required to apply a correction") + process.add_argument( + "--glossary-confidence-threshold", + type=float, + help="minimum confidence required to apply a glossary correction", + ) + process.add_argument( + "--grammar-confidence-threshold", + type=float, + help="minimum confidence required to apply a grammar correction", + ) process.add_argument("--max-retries", type=int, help="maximum Instructor retries for structured response validation") process.add_argument("--glossary-max-llm-passes", type=int, help="maximum total LLM passes for glossary corrections") + process.add_argument("--grammar-max-llm-passes", type=int, help="maximum total LLM passes for grammar corrections") process.add_argument("--work-dir", type=Path, help="directory for per-run scratch diagnostics") return parser @@ -46,9 +56,11 @@ def _process(args: argparse.Namespace) -> int: model=args.model, base_url=args.base_url, max_section_tokens=args.max_section_tokens, - confidence_threshold=args.confidence_threshold, + glossary_confidence_threshold=args.glossary_confidence_threshold, + grammar_confidence_threshold=args.grammar_confidence_threshold, max_retries=args.max_retries, glossary_max_llm_passes=args.glossary_max_llm_passes, + grammar_max_llm_passes=args.grammar_max_llm_passes, work_dir=args.work_dir, ) ) diff --git a/src/audita/config.py b/src/audita/config.py index cb8311a..2704503 100644 --- a/src/audita/config.py +++ b/src/audita/config.py @@ -9,10 +9,12 @@ from .errors import AuditaConfigError DEFAULT_MODEL = "openrouter/mistralai/mistral-small-3.2-24b-instruct" DEFAULT_BASE_URL = "https://openrouter.ai/api/v1" DEFAULT_MAX_SECTION_TOKENS = 16000 -DEFAULT_CONFIDENCE_THRESHOLD = 0.60 +DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD = 0.60 +DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD = 0.60 DEFAULT_MAX_RETRIES = 3 DEFAULT_WORK_DIR = "/tmp/audita" DEFAULT_GLOSSARY_MAX_LLM_PASSES = 3 +DEFAULT_GRAMMAR_MAX_LLM_PASSES = 3 @dataclass(frozen=True) @@ -20,9 +22,11 @@ class ConfigOverrides: model: Optional[str] = None base_url: Optional[str] = None max_section_tokens: Optional[int] = None - confidence_threshold: Optional[float] = None + glossary_confidence_threshold: Optional[float] = None + grammar_confidence_threshold: Optional[float] = None max_retries: Optional[int] = None glossary_max_llm_passes: Optional[int] = None + grammar_max_llm_passes: Optional[int] = None work_dir: Optional[Path] = None @@ -32,9 +36,11 @@ class AuditaConfig: model: str = DEFAULT_MODEL base_url: str = DEFAULT_BASE_URL max_section_tokens: int = DEFAULT_MAX_SECTION_TOKENS - confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD + glossary_confidence_threshold: float = DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD + grammar_confidence_threshold: float = DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD max_retries: int = DEFAULT_MAX_RETRIES glossary_max_llm_passes: int = DEFAULT_GLOSSARY_MAX_LLM_PASSES + grammar_max_llm_passes: int = DEFAULT_GRAMMAR_MAX_LLM_PASSES work_dir: Path = Path(DEFAULT_WORK_DIR) @classmethod @@ -55,11 +61,17 @@ class AuditaConfig: DEFAULT_MAX_SECTION_TOKENS, "AUDITA_MAX_SECTION_TOKENS", ) - confidence_threshold = _select_float( - selected.confidence_threshold, - source.get("AUDITA_CONFIDENCE_THRESHOLD"), - DEFAULT_CONFIDENCE_THRESHOLD, - "AUDITA_CONFIDENCE_THRESHOLD", + glossary_confidence_threshold = _select_float( + selected.glossary_confidence_threshold, + source.get("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD"), + DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD, + "AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD", + ) + grammar_confidence_threshold = _select_float( + selected.grammar_confidence_threshold, + source.get("AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD"), + DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD, + "AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD", ) max_retries = _select_int( selected.max_retries, @@ -73,6 +85,12 @@ class AuditaConfig: DEFAULT_GLOSSARY_MAX_LLM_PASSES, "AUDITA_GLOSSARY_MAX_LLM_PASSES", ) + grammar_max_llm_passes = _select_int( + selected.grammar_max_llm_passes, + source.get("AUDITA_GRAMMAR_MAX_LLM_PASSES"), + DEFAULT_GRAMMAR_MAX_LLM_PASSES, + "AUDITA_GRAMMAR_MAX_LLM_PASSES", + ) work_dir_value = selected.work_dir or Path(source.get("AUDITA_WORK_DIR") or DEFAULT_WORK_DIR) config = cls( @@ -80,9 +98,11 @@ class AuditaConfig: model=model, base_url=base_url, max_section_tokens=max_section_tokens, - confidence_threshold=confidence_threshold, + glossary_confidence_threshold=glossary_confidence_threshold, + grammar_confidence_threshold=grammar_confidence_threshold, max_retries=max_retries, glossary_max_llm_passes=glossary_max_llm_passes, + grammar_max_llm_passes=grammar_max_llm_passes, work_dir=Path(work_dir_value), ) config.validate() @@ -97,12 +117,16 @@ class AuditaConfig: raise AuditaConfigError("AUDITA_BASE_URL must not be empty.") if self.max_section_tokens <= 0: raise AuditaConfigError("AUDITA_MAX_SECTION_TOKENS must be greater than zero.") - if not 0.0 <= self.confidence_threshold <= 1.0: - raise AuditaConfigError("AUDITA_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0.") + if not 0.0 <= self.glossary_confidence_threshold <= 1.0: + raise AuditaConfigError("AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0.") + if not 0.0 <= self.grammar_confidence_threshold <= 1.0: + raise AuditaConfigError("AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0.") if self.max_retries < 0: raise AuditaConfigError("AUDITA_MAX_RETRIES must be greater than or equal to zero.") if self.glossary_max_llm_passes < 1: raise AuditaConfigError("AUDITA_GLOSSARY_MAX_LLM_PASSES must be greater than or equal to one.") + if self.grammar_max_llm_passes < 1: + raise AuditaConfigError("AUDITA_GRAMMAR_MAX_LLM_PASSES must be greater than or equal to one.") def _get_required_env(env: Mapping[str, str], name: str) -> str: diff --git a/src/audita/corrections.py b/src/audita/corrections.py index 3dd90e5..bdf6e8c 100644 --- a/src/audita/corrections.py +++ b/src/audita/corrections.py @@ -1,9 +1,11 @@ from dataclasses import asdict, dataclass -from typing import Dict, Iterable, List, Optional, Tuple +from typing import Dict, Iterable, List, Literal, Optional, Tuple from .errors import AuditaValidationError from .schemas import CorrectionCandidate, TranscriptSegment +ReplacementMode = Literal["replace_all", "require_unique"] + @dataclass(frozen=True) class SkippedCorrection: @@ -30,9 +32,12 @@ def apply_corrections( transcript: List[TranscriptSegment], corrections: Iterable[CorrectionCandidate], confidence_threshold: float, + replacement_mode: ReplacementMode = "replace_all", ) -> CorrectionApplicationResult: if not 0.0 <= confidence_threshold <= 1.0: raise AuditaValidationError("Confidence threshold must be between 0.0 and 1.0.") + if replacement_mode not in ("replace_all", "require_unique"): + raise AuditaValidationError("Replacement mode must be replace_all or require_unique.") revised = list(transcript) id_to_position = _id_to_position(revised) @@ -44,7 +49,7 @@ def apply_corrections( ignored_ids.append(correction.id) continue - reason, actual_text = _target_error(revised, id_to_position, correction) + reason, actual_text = _target_error(revised, id_to_position, correction, replacement_mode) if reason is not None: skipped.append(_skip(correction, reason, actual_text=actual_text)) continue @@ -67,6 +72,7 @@ def _target_error( transcript: List[TranscriptSegment], id_to_position: Dict[int, int], correction: CorrectionCandidate, + replacement_mode: ReplacementMode, ) -> Tuple[Optional[str], Optional[str]]: if correction.id not in id_to_position: return "id does not exist in transcript", None @@ -80,6 +86,8 @@ def _target_error( match_count = segment.text.count(correction.original_text) if match_count == 0: return "original_text does not match any substring in segment text", segment.text + if replacement_mode == "require_unique" and match_count > 1: + return "original_text appears more than once in segment text", segment.text return None, None diff --git a/src/audita/passes.py b/src/audita/passes.py index 97e07b9..780916f 100644 --- a/src/audita/passes.py +++ b/src/audita/passes.py @@ -4,7 +4,7 @@ from typing import List, Protocol from .chunking import TranscriptSection from .config import AuditaConfig -from .prompts import build_glossary_correction_messages +from .prompts import build_glossary_correction_messages, build_grammar_correction_messages from .schemas import CorrectionCandidate, CorrectionSet, Glossary @@ -45,3 +45,25 @@ class GlossaryCorrectionPass: response_path = run_dir / f"corrections-{section.section_index:04d}.json" response_path.write_text(response.model_dump_json(indent=2) + "\n", encoding="utf-8") return list(response.corrections) + + +class GrammarCorrectionPass: + def __init__(self, llm_client: LLMClient) -> None: + self._llm_client = llm_client + + def run( + self, + section: TranscriptSection, + glossary: Glossary, + config: AuditaConfig, + run_dir: Path, + retry_pass: bool = False, + ) -> List[CorrectionCandidate]: + messages = build_grammar_correction_messages(section, glossary, retry_pass=retry_pass) + prompt_path = run_dir / f"prompt-{section.section_index:04d}.json" + prompt_path.write_text(json.dumps(messages, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + + response = self._llm_client.create_corrections(messages, config) + response_path = run_dir / f"corrections-{section.section_index:04d}.json" + response_path.write_text(response.model_dump_json(indent=2) + "\n", encoding="utf-8") + return list(response.corrections) diff --git a/src/audita/pipeline.py b/src/audita/pipeline.py index 51f8a74..a4b1e86 100644 --- a/src/audita/pipeline.py +++ b/src/audita/pipeline.py @@ -1,21 +1,31 @@ import json import shutil +from dataclasses import dataclass from datetime import datetime from pathlib import Path -from typing import Callable, Dict, List, Optional +from typing import Callable, Dict, List, Optional, Tuple from uuid import uuid4 from .chunking import IndexedSegment, TranscriptSection, chunk_indexed_segments from .config import AuditaConfig -from .corrections import SkippedCorrection, apply_corrections +from .corrections import ReplacementMode, SkippedCorrection, apply_corrections from .errors import AuditaError -from .passes import GlossaryCorrectionPass, LLMClient +from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient from .schemas import Glossary, TranscriptSegment, parse_transcript_json ProgressCallback = Callable[[str], None] +@dataclass(frozen=True) +class StageSpec: + name: str + correction_pass: CorrectionPass + max_llm_passes: int + confidence_threshold: float + replacement_mode: ReplacementMode + + def process_transcript( transcript: List[TranscriptSegment], glossary: Glossary, @@ -26,8 +36,8 @@ def process_transcript( run_dir = _create_run_dir(config.work_dir) try: _log(progress, f"Created work directory {run_dir}") - pass_summaries: List[dict] = [] - _write_run_metadata(run_dir, config, pass_summaries) + stage_summaries: List[dict] = [] + _write_run_metadata(run_dir, config, stage_summaries) if llm_client is None: from .llm import InstructorLLMClient @@ -35,87 +45,56 @@ def process_transcript( llm_client = InstructorLLMClient(config) working = list(transcript) - correction_pass = GlossaryCorrectionPass(llm_client) - unresolved_retry_skips: Dict[int, SkippedCorrection] = {} - final_nonretry_skips: List[SkippedCorrection] = [] - - for pass_number in range(1, config.glossary_max_llm_passes + 1): - if pass_number == 1: - indexed_segments = _indexed_segments_for_ids(working, [segment.id for segment in working]) - else: - retry_ids = sorted(unresolved_retry_skips) - if not retry_ids: - break - indexed_segments = _indexed_segments_for_ids(working, retry_ids) - if not indexed_segments: - break - - pass_dir = run_dir / f"pass-{pass_number:04d}" - pass_dir.mkdir() - sections = chunk_indexed_segments(indexed_segments, config.max_section_tokens) - corrections = [] - - for section in sections: - _write_and_validate_section(pass_dir, section) - _log( - progress, - f"Processing glossary pass {pass_number}/{config.glossary_max_llm_passes} " - f"section {section.section_index + 1}/{len(sections)} " - f"({len(section.segments)} segments, estimated {section.token_count} tokens)", - ) - corrections.extend( - correction_pass.run( - section, - glossary, - config, - pass_dir, - retry_pass=pass_number > 1, - ) - ) - - application_result = apply_corrections( - working, - corrections, - config.confidence_threshold, - ) - working = application_result.transcript - - for correction_id in application_result.applied_ids: - unresolved_retry_skips.pop(correction_id, None) - for correction_id in application_result.ignored_ids: - unresolved_retry_skips.pop(correction_id, None) - for skipped in application_result.skipped: - if _is_retryable_skip(skipped, working): - unresolved_retry_skips[skipped.id] = skipped - else: - final_nonretry_skips.append(skipped) - - pass_summaries.append( - { - "pass_number": pass_number, - "retry_pass": pass_number > 1, - "section_count": len(sections), - "segment_count": len(indexed_segments), - "corrections_returned": len(corrections), - "applied_count": len(application_result.applied_ids), - "ignored_below_threshold_count": len(application_result.ignored_ids), - "skipped_count": len(application_result.skipped), - "retry_segment_count": len(unresolved_retry_skips), - } - ) - _write_run_metadata(run_dir, config, pass_summaries) - - if not unresolved_retry_skips: - break - - final_skipped = final_nonretry_skips + [ - unresolved_retry_skips[correction_id] for correction_id in sorted(unresolved_retry_skips) + stages = [ + StageSpec( + name="glossary", + correction_pass=GlossaryCorrectionPass(llm_client), + max_llm_passes=config.glossary_max_llm_passes, + confidence_threshold=config.glossary_confidence_threshold, + replacement_mode="replace_all", + ), + StageSpec( + name="grammar", + correction_pass=GrammarCorrectionPass(llm_client), + max_llm_passes=config.grammar_max_llm_passes, + confidence_threshold=config.grammar_confidence_threshold, + replacement_mode="require_unique", + ), ] + final_skipped: List[Tuple[str, SkippedCorrection]] = [] + + for stage in stages: + stage_dir = run_dir / stage.name + stage_dir.mkdir() + stage_summary = { + "stage": stage.name, + "max_llm_passes": stage.max_llm_passes, + "confidence_threshold": stage.confidence_threshold, + "replacement_mode": stage.replacement_mode, + "passes": [], + } + stage_summaries.append(stage_summary) + _write_run_metadata(run_dir, config, stage_summaries) + + working, stage_skipped = _run_correction_stage( + working, + glossary, + config, + stage, + run_dir, + stage_dir, + stage_summaries, + stage_summary["passes"], + progress, + ) + final_skipped.extend((stage.name, skipped) for skipped in stage_skipped) + _write_run_metadata(run_dir, config, stage_summaries) + _write_skipped_corrections(run_dir, final_skipped) - for skipped in final_skipped: + for stage_name, skipped in final_skipped: _log( progress, - f"Skipping correction for id {skipped.id}: {skipped.reason}", + f"Skipping {stage_name} correction for id {skipped.id}: {skipped.reason}", ) revised = _sort_transcript_chronologically(working) except Exception as exc: @@ -132,6 +111,97 @@ def process_transcript( return revised +def _run_correction_stage( + transcript: List[TranscriptSegment], + glossary: Glossary, + config: AuditaConfig, + stage: StageSpec, + run_dir: Path, + stage_dir: Path, + stage_summaries: List[dict], + pass_summaries: List[dict], + progress: Optional[ProgressCallback], +) -> Tuple[List[TranscriptSegment], List[SkippedCorrection]]: + working = list(transcript) + unresolved_retry_skips: Dict[int, SkippedCorrection] = {} + final_nonretry_skips: List[SkippedCorrection] = [] + + for pass_number in range(1, stage.max_llm_passes + 1): + if pass_number == 1: + indexed_segments = _indexed_segments_for_ids(working, [segment.id for segment in working]) + else: + retry_ids = sorted(unresolved_retry_skips) + if not retry_ids: + break + indexed_segments = _indexed_segments_for_ids(working, retry_ids) + if not indexed_segments: + break + + pass_dir = stage_dir / f"pass-{pass_number:04d}" + pass_dir.mkdir() + sections = chunk_indexed_segments(indexed_segments, config.max_section_tokens) + corrections = [] + + for section in sections: + _write_and_validate_section(pass_dir, section) + _log( + progress, + f"Processing {stage.name} pass {pass_number}/{stage.max_llm_passes} " + f"section {section.section_index + 1}/{len(sections)} " + f"({len(section.segments)} segments, estimated {section.token_count} tokens)", + ) + corrections.extend( + stage.correction_pass.run( + section, + glossary, + config, + pass_dir, + retry_pass=pass_number > 1, + ) + ) + + application_result = apply_corrections( + working, + corrections, + stage.confidence_threshold, + replacement_mode=stage.replacement_mode, + ) + working = application_result.transcript + + for correction_id in application_result.applied_ids: + unresolved_retry_skips.pop(correction_id, None) + for correction_id in application_result.ignored_ids: + unresolved_retry_skips.pop(correction_id, None) + for skipped in application_result.skipped: + if _is_retryable_skip(skipped, working): + unresolved_retry_skips[skipped.id] = skipped + else: + final_nonretry_skips.append(skipped) + + pass_summaries.append( + { + "pass_number": pass_number, + "retry_pass": pass_number > 1, + "section_count": len(sections), + "segment_count": len(indexed_segments), + "corrections_returned": len(corrections), + "applied_count": len(application_result.applied_ids), + "ignored_below_threshold_count": len(application_result.ignored_ids), + "skipped_count": len(application_result.skipped), + "retry_segment_count": len(unresolved_retry_skips), + } + ) + _write_run_metadata(run_dir, config, stage_summaries) + + if not unresolved_retry_skips: + break + + final_skipped = final_nonretry_skips + [ + unresolved_retry_skips[correction_id] for correction_id in sorted(unresolved_retry_skips) + ] + return working, final_skipped + + def _create_run_dir(work_dir: Path) -> Path: work_dir.mkdir(parents=True, exist_ok=True) timestamp = datetime.utcnow().strftime("%Y%m%dT%H%M%SZ") @@ -143,16 +213,18 @@ def _create_run_dir(work_dir: Path) -> Path: def _write_run_metadata( run_dir: Path, config: AuditaConfig, - pass_summaries: List[dict], + stage_summaries: List[dict], ) -> None: metadata = { "model": config.model, "base_url": config.base_url, "max_section_tokens": config.max_section_tokens, - "confidence_threshold": config.confidence_threshold, + "glossary_confidence_threshold": config.glossary_confidence_threshold, + "grammar_confidence_threshold": config.grammar_confidence_threshold, "max_retries": config.max_retries, "glossary_max_llm_passes": config.glossary_max_llm_passes, - "passes": pass_summaries, + "grammar_max_llm_passes": config.grammar_max_llm_passes, + "stages": stage_summaries, } (run_dir / "metadata.json").write_text( json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", @@ -167,11 +239,15 @@ def _write_and_validate_section(run_dir: Path, section: TranscriptSection) -> No parse_transcript_json(section_json, require_sequential_ids=False) -def _write_skipped_corrections(run_dir: Path, skipped: List[SkippedCorrection]) -> None: +def _write_skipped_corrections(run_dir: Path, skipped: List[Tuple[str, SkippedCorrection]]) -> None: skipped_path = run_dir / "skipped-corrections.json" skipped_path.write_text( json.dumps( - {"skipped_corrections": [item.to_dict() for item in skipped]}, + { + "skipped_corrections": [ + {"stage": stage_name, **item.to_dict()} for stage_name, item in skipped + ] + }, ensure_ascii=False, indent=2, ) diff --git a/src/audita/prompts.py b/src/audita/prompts.py index d327c2a..bfff834 100644 --- a/src/audita/prompts.py +++ b/src/audita/prompts.py @@ -57,3 +57,49 @@ def build_glossary_correction_messages( f"Transcript section:\n{section_json}" ) return [{"role": "system", "content": system}, {"role": "user", "content": user}] + + +def build_grammar_correction_messages( + section: TranscriptSection, + glossary: Glossary, + retry_pass: bool = False, +) -> List[Message]: + glossary_json = json.dumps(glossary.model_dump(mode="json"), ensure_ascii=False, indent=2) + section_json = json.dumps(section.prompt_payload(), ensure_ascii=False, indent=2) + + system = ( + "You are Audita, a conservative transcript readability assistant. " + "Improve readability only where the change preserves the speaker's words and meaning. " + "Allowed changes are capitalization, commas, periods, em dashes, ellipses, homophone fixes, " + "and spelling fixes. Do not paraphrase, summarize, reorder words, or change content." + ) + retry_guidance = "" + if retry_pass: + retry_guidance = ( + "Retry guidance:\n" + "These segments are being retried because previous correction spans did not apply cleanly. " + "Copy original_text exactly from the current segment text, using only a span that appears exactly once.\n\n" + ) + + user = ( + "Review this transcript section and return only readability corrections that should be applied.\n\n" + f"{retry_guidance}" + "Rules:\n" + "- Allowed corrections are only capitalization changes, punctuation changes involving commas, periods, em dashes, and ellipses, homophone fixes, and spelling fixes.\n" + "- Do not add, remove, reorder, or replace words except for clear homophone or spelling corrections that preserve the spoken content.\n" + "- Do not paraphrase, summarize, clarify, smooth style, or change the speaker's intent or meaning.\n" + "- Treat the glossary as protected vocabulary and context; do not introduce new glossary substitutions during this grammar pass.\n" + "- Use the exact id from the input segment.\n" + "- For returned corrections, original_text must be only the exact text span that needs replacement, not the full segment text unless the whole segment is the replacement span.\n" + "- Choose an original_text span that appears exactly once in the current segment text.\n" + "- corrected_text must be only the replacement text for that span, not the full corrected segment text unless the whole segment is the replacement span.\n" + "- Each returned correction must contain only id, original_text, corrected_text, and confidence.\n" + "- Do not return corrections where original_text and corrected_text are identical.\n" + "- Do not return speaker, start, or end fields.\n" + "- Return only changed segments; do not return entries for unchanged segments.\n" + "- confidence must be between 0.0 and 1.0.\n" + "- If no corrections are needed, return an empty corrections list.\n\n" + f"Protected glossary/context:\n{glossary_json}\n\n" + f"Transcript section:\n{section_json}" + ) + return [{"role": "system", "content": system}, {"role": "user", "content": user}] diff --git a/tests/test_cli.py b/tests/test_cli.py index e237716..8a1ae50 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -16,4 +16,9 @@ def test_process_help_includes_glossary_pass_flag(capsys): main(["process", "--help"]) assert exc.value.code == 0 - assert "--glossary-max-llm-passes" in capsys.readouterr().out + output = capsys.readouterr().out + assert "--glossary-max-llm-passes" in output + assert "--grammar-max-llm-passes" in output + assert "--glossary-confidence-threshold" in output + assert "--grammar-confidence-threshold" in output + assert "--confidence-threshold" not in output diff --git a/tests/test_config.py b/tests/test_config.py index e04f6e4..b7a9e32 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -3,18 +3,29 @@ from pathlib import Path import pytest from audita.config import AuditaConfig, ConfigOverrides -from audita.config import DEFAULT_CONFIDENCE_THRESHOLD, DEFAULT_GLOSSARY_MAX_LLM_PASSES, DEFAULT_MAX_RETRIES, DEFAULT_MAX_SECTION_TOKENS, DEFAULT_WORK_DIR +from audita.config import ( + DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD, + DEFAULT_GLOSSARY_MAX_LLM_PASSES, + DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD, + DEFAULT_GRAMMAR_MAX_LLM_PASSES, + DEFAULT_MAX_RETRIES, + DEFAULT_MAX_SECTION_TOKENS, + DEFAULT_WORK_DIR, +) from audita.errors import AuditaConfigError def test_config_uses_defaults_with_api_key(): config = AuditaConfig.from_sources(env={"OPENROUTER_API_KEY": "key"}) - assert config.confidence_threshold == DEFAULT_CONFIDENCE_THRESHOLD - assert config.confidence_threshold == 0.6 + assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD + assert config.glossary_confidence_threshold == 0.6 + assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD + assert config.grammar_confidence_threshold == 0.6 assert config.max_section_tokens == DEFAULT_MAX_SECTION_TOKENS assert config.max_retries == DEFAULT_MAX_RETRIES assert config.glossary_max_llm_passes == DEFAULT_GLOSSARY_MAX_LLM_PASSES + assert config.grammar_max_llm_passes == DEFAULT_GRAMMAR_MAX_LLM_PASSES assert config.work_dir == Path(DEFAULT_WORK_DIR) @@ -23,17 +34,21 @@ def test_config_env_overrides_defaults(): env={ "OPENROUTER_API_KEY": "key", "AUDITA_MAX_SECTION_TOKENS": "42", - "AUDITA_CONFIDENCE_THRESHOLD": "0.9", + "AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD": "0.9", + "AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD": "0.7", "AUDITA_MAX_RETRIES": "5", "AUDITA_GLOSSARY_MAX_LLM_PASSES": "7", + "AUDITA_GRAMMAR_MAX_LLM_PASSES": "4", "AUDITA_WORK_DIR": "/tmp/custom-audita", } ) assert config.max_section_tokens == 42 - assert config.confidence_threshold == 0.9 + assert config.glossary_confidence_threshold == 0.9 + assert config.grammar_confidence_threshold == 0.7 assert config.max_retries == 5 assert config.glossary_max_llm_passes == 7 + assert config.grammar_max_llm_passes == 4 assert config.work_dir == Path("/tmp/custom-audita") @@ -44,21 +59,26 @@ def test_config_cli_overrides_env(): "AUDITA_MAX_SECTION_TOKENS": "42", "AUDITA_MAX_RETRIES": "5", "AUDITA_GLOSSARY_MAX_LLM_PASSES": "7", + "AUDITA_GRAMMAR_MAX_LLM_PASSES": "6", "AUDITA_WORK_DIR": "/tmp/env-audita", }, overrides=ConfigOverrides( max_section_tokens=100, - confidence_threshold=0.7, + glossary_confidence_threshold=0.7, + grammar_confidence_threshold=0.65, max_retries=3, glossary_max_llm_passes=2, + grammar_max_llm_passes=3, work_dir=Path("/tmp/cli-audita"), ), ) assert config.max_section_tokens == 100 - assert config.confidence_threshold == 0.7 + assert config.glossary_confidence_threshold == 0.7 + assert config.grammar_confidence_threshold == 0.65 assert config.max_retries == 3 assert config.glossary_max_llm_passes == 2 + assert config.grammar_max_llm_passes == 3 assert config.work_dir == Path("/tmp/cli-audita") @@ -79,3 +99,30 @@ def test_config_rejects_invalid_glossary_pass_count(): AuditaConfig.from_sources( env={"OPENROUTER_API_KEY": "key", "AUDITA_GLOSSARY_MAX_LLM_PASSES": "0"} ) + + +def test_config_rejects_invalid_grammar_pass_count(): + with pytest.raises(AuditaConfigError): + AuditaConfig.from_sources( + env={"OPENROUTER_API_KEY": "key", "AUDITA_GRAMMAR_MAX_LLM_PASSES": "0"} + ) + + +def test_config_rejects_invalid_stage_thresholds(): + with pytest.raises(AuditaConfigError): + AuditaConfig.from_sources( + env={"OPENROUTER_API_KEY": "key", "AUDITA_GLOSSARY_CONFIDENCE_THRESHOLD": "1.1"} + ) + with pytest.raises(AuditaConfigError): + AuditaConfig.from_sources( + env={"OPENROUTER_API_KEY": "key", "AUDITA_GRAMMAR_CONFIDENCE_THRESHOLD": "-0.1"} + ) + + +def test_legacy_confidence_threshold_env_is_ignored(): + config = AuditaConfig.from_sources( + env={"OPENROUTER_API_KEY": "key", "AUDITA_CONFIDENCE_THRESHOLD": "0.9"} + ) + + assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD + assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD diff --git a/tests/test_corrections.py b/tests/test_corrections.py index 71e4715..ee05c23 100644 --- a/tests/test_corrections.py +++ b/tests/test_corrections.py @@ -144,6 +144,33 @@ def test_apply_corrections_replaces_all_repeated_substrings(): assert result.skipped == [] +def test_apply_corrections_requires_unique_match_when_configured(): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "there and there"} + ] + """ + ) + correction = CorrectionCandidate( + id=1, + original_text="there", + corrected_text="their", + confidence=0.8, + ) + + result = apply_corrections( + transcript, + [correction], + confidence_threshold=0.8, + replacement_mode="require_unique", + ) + + assert result.transcript[0].text == "there and there" + assert len(result.skipped) == 1 + assert "more than once" in result.skipped[0].reason + + def test_apply_corrections_skips_empty_original_text(): transcript = _transcript() correction = CorrectionCandidate( @@ -163,3 +190,8 @@ def test_apply_corrections_skips_empty_original_text(): def test_apply_corrections_rejects_invalid_threshold(): with pytest.raises(AuditaValidationError): apply_corrections(_transcript(), [], confidence_threshold=1.1) + + +def test_apply_corrections_rejects_invalid_replacement_mode(): + with pytest.raises(AuditaValidationError): + apply_corrections(_transcript(), [], confidence_threshold=0.8, replacement_mode="unknown") diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index 6da3329..c1254e3 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -17,13 +17,15 @@ class FakeLLMClient: return self.responses.pop(0) -def _config(tmp_path, glossary_max_llm_passes=3): +def _config(tmp_path, glossary_max_llm_passes=3, grammar_max_llm_passes=3): return AuditaConfig( api_key="key", max_section_tokens=16000, - confidence_threshold=0.8, + glossary_confidence_threshold=0.8, + grammar_confidence_threshold=0.8, max_retries=3, glossary_max_llm_passes=glossary_max_llm_passes, + grammar_max_llm_passes=grammar_max_llm_passes, work_dir=tmp_path / "work", ) @@ -57,7 +59,12 @@ def test_pipeline_processes_with_fake_llm_and_cleans_work_dir(tmp_path): corrected_text="Chauntea", confidence=0.95, ) - fake_client = FakeLLMClient([CorrectionSet(corrections=[correction])]) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[correction]), + CorrectionSet(corrections=[]), + ] + ) revised = process_transcript( _transcript(), @@ -68,18 +75,23 @@ def test_pipeline_processes_with_fake_llm_and_cleans_work_dir(tmp_path): assert [segment.speaker for segment in revised] == ["Mike", "Eric"] assert revised[1].text == "I ask Chauntea." - assert fake_client.calls == 1 + assert fake_client.calls == 2 assert list((tmp_path / "work").iterdir()) == [] -def test_pipeline_skips_bad_correction_and_preserves_diagnostics(tmp_path): +def test_pipeline_skips_bad_glossary_correction_and_preserves_diagnostics(tmp_path): correction = CorrectionCandidate( id=1, original_text="Different text.", corrected_text="Chauntea", confidence=0.95, ) - fake_client = FakeLLMClient([CorrectionSet(corrections=[correction])]) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[correction]), + CorrectionSet(corrections=[]), + ] + ) progress = [] revised = process_transcript( @@ -91,12 +103,13 @@ def test_pipeline_skips_bad_correction_and_preserves_diagnostics(tmp_path): ) assert revised[1].text == "I ask Chontia." - assert any("Skipping correction for id 1" in message for message in progress) + assert any("Skipping glossary correction for id 1" in message for message in progress) preserved = list((tmp_path / "work").iterdir()) assert len(preserved) == 1 skipped_path = preserved[0] / "skipped-corrections.json" assert skipped_path.exists() diagnostics = json.loads(skipped_path.read_text(encoding="utf-8")) + assert diagnostics["skipped_corrections"][0]["stage"] == "glossary" assert diagnostics["skipped_corrections"][0]["id"] == 1 assert "does not match any substring" in diagnostics["skipped_corrections"][0]["reason"] @@ -118,6 +131,7 @@ def test_pipeline_retries_skipped_segment_and_cleans_work_dir_when_fixed(tmp_pat [ CorrectionSet(corrections=[first_pass]), CorrectionSet(corrections=[second_pass]), + CorrectionSet(corrections=[]), ] ) @@ -128,7 +142,7 @@ def test_pipeline_retries_skipped_segment_and_cleans_work_dir_when_fixed(tmp_pat llm_client=fake_client, ) - assert fake_client.calls == 2 + assert fake_client.calls == 3 assert [segment.speaker for segment in revised] == ["Mike", "Eric"] assert revised[1].text == "I ask Chauntea." assert list((tmp_path / "work").iterdir()) == [] @@ -157,6 +171,7 @@ def test_pipeline_retry_prompt_contains_only_valid_deduped_ids(tmp_path): [ CorrectionSet(corrections=[first_bad, second_bad_same_segment, invalid_segment]), CorrectionSet(corrections=[]), + CorrectionSet(corrections=[]), ] ) @@ -167,14 +182,14 @@ def test_pipeline_retry_prompt_contains_only_valid_deduped_ids(tmp_path): llm_client=fake_client, ) - assert fake_client.calls == 2 + assert fake_client.calls == 3 retry_prompt = fake_client.messages[1][1]["content"] retry_payload = json.loads(retry_prompt.split("Transcript section:\n", maxsplit=1)[1]) assert retry_payload == [{"id": 1, "original_text": "I ask Chontia."}] assert "Retry guidance" in retry_prompt -def test_pipeline_writes_pass_metadata_for_unresolved_retries(tmp_path): +def test_pipeline_writes_stage_metadata_for_unresolved_retries(tmp_path): first_pass = CorrectionCandidate( id=1, original_text="Contia", @@ -185,6 +200,7 @@ def test_pipeline_writes_pass_metadata_for_unresolved_retries(tmp_path): [ CorrectionSet(corrections=[first_pass]), CorrectionSet(corrections=[]), + CorrectionSet(corrections=[]), ] ) @@ -197,8 +213,150 @@ def test_pipeline_writes_pass_metadata_for_unresolved_retries(tmp_path): run_dirs = list((tmp_path / "work").iterdir()) assert len(run_dirs) == 1 + assert (run_dirs[0] / "glossary" / "pass-0001").exists() + assert (run_dirs[0] / "grammar" / "pass-0001").exists() metadata = json.loads((run_dirs[0] / "metadata.json").read_text(encoding="utf-8")) assert metadata["glossary_max_llm_passes"] == 2 - assert [item["pass_number"] for item in metadata["passes"]] == [1, 2] - assert metadata["passes"][0]["retry_segment_count"] == 1 - assert metadata["passes"][1]["retry_pass"] is True + assert metadata["grammar_max_llm_passes"] == 3 + assert metadata["glossary_confidence_threshold"] == 0.8 + assert metadata["grammar_confidence_threshold"] == 0.8 + assert [item["stage"] for item in metadata["stages"]] == ["glossary", "grammar"] + assert [item["pass_number"] for item in metadata["stages"][0]["passes"]] == [1, 2] + assert metadata["stages"][0]["passes"][0]["retry_segment_count"] == 1 + assert metadata["stages"][0]["passes"][1]["retry_pass"] is True + + +def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "Eric", "start": 10.0, "end": 11.0, "text": "i ask Chontia."}, + {"id": 2, "speaker": "Mike", "start": 0.0, "end": 1.0, "text": "Then Lyra."} + ] + """ + ) + glossary_correction = CorrectionCandidate( + id=1, + original_text="Chontia", + corrected_text="Chauntea", + confidence=0.95, + ) + grammar_correction = CorrectionCandidate( + id=1, + original_text="i", + corrected_text="I", + confidence=0.95, + ) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[glossary_correction]), + CorrectionSet(corrections=[grammar_correction]), + ] + ) + + revised = process_transcript( + transcript, + _glossary(), + _config(tmp_path), + llm_client=fake_client, + ) + + grammar_prompt = fake_client.messages[1][1]["content"] + grammar_payload = json.loads(grammar_prompt.split("Transcript section:\n", maxsplit=1)[1]) + assert grammar_payload[0]["original_text"] == "i ask Chauntea." + assert revised[1].text == "I ask Chauntea." + + +def test_grammar_stage_retries_repeated_span_and_applies_unique_retry(tmp_path): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "there and there."} + ] + """ + ) + repeated_span = CorrectionCandidate( + id=1, + original_text="there", + corrected_text="their", + confidence=0.95, + ) + unique_retry = CorrectionCandidate( + id=1, + original_text="there and there", + corrected_text="their and there", + confidence=0.95, + ) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[]), + CorrectionSet(corrections=[repeated_span]), + CorrectionSet(corrections=[unique_retry]), + ] + ) + + revised = process_transcript( + transcript, + _glossary(), + _config(tmp_path, grammar_max_llm_passes=2), + llm_client=fake_client, + ) + + assert fake_client.calls == 3 + assert revised[0].text == "their and there." + retry_prompt = fake_client.messages[2][1]["content"] + retry_payload = json.loads(retry_prompt.split("Transcript section:\n", maxsplit=1)[1]) + assert retry_payload == [{"id": 1, "original_text": "there and there."}] + + +def test_below_threshold_grammar_corrections_are_not_retried(tmp_path): + correction = CorrectionCandidate( + id=1, + original_text="I", + corrected_text="i", + confidence=0.7, + ) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[]), + CorrectionSet(corrections=[correction]), + ] + ) + + revised = process_transcript( + _transcript(), + _glossary(), + _config(tmp_path, grammar_max_llm_passes=3), + llm_client=fake_client, + ) + + assert fake_client.calls == 2 + assert revised[1].text == "I ask Chontia." + + +def test_unresolved_grammar_skip_preserves_diagnostics(tmp_path): + correction = CorrectionCandidate( + id=1, + original_text="a", + corrected_text="A", + confidence=0.95, + ) + fake_client = FakeLLMClient( + [ + CorrectionSet(corrections=[]), + CorrectionSet(corrections=[correction]), + ] + ) + + process_transcript( + _transcript(), + _glossary(), + _config(tmp_path, grammar_max_llm_passes=1), + llm_client=fake_client, + ) + + run_dirs = list((tmp_path / "work").iterdir()) + assert len(run_dirs) == 1 + diagnostics = json.loads((run_dirs[0] / "skipped-corrections.json").read_text(encoding="utf-8")) + assert diagnostics["skipped_corrections"][0]["stage"] == "grammar" + assert "more than once" in diagnostics["skipped_corrections"][0]["reason"] diff --git a/tests/test_prompts.py b/tests/test_prompts.py index 172f94c..1e2bd4e 100644 --- a/tests/test_prompts.py +++ b/tests/test_prompts.py @@ -1,7 +1,7 @@ import json from audita.chunking import chunk_transcript -from audita.prompts import build_glossary_correction_messages +from audita.prompts import build_glossary_correction_messages, build_grammar_correction_messages from audita.schemas import parse_glossary_yaml, parse_transcript_json @@ -65,3 +65,62 @@ def test_prompt_uses_simplified_segment_payload(): assert "speaker" not in prompt_segments[0] assert "start" not in prompt_segments[0] assert "end" not in prompt_segments[0] + + +def test_grammar_prompt_limits_readability_corrections_and_protects_glossary(): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "then lyra went their"} + ] + """ + ) + glossary = parse_glossary_yaml( + """ + glossary: + - name: "Lyra" + category: npc + summary: "Lyra is a hostile NPC." + """ + ) + section = chunk_transcript(transcript, max_section_tokens=16000)[0] + + messages = build_grammar_correction_messages(section, glossary) + prompt_text = "\n".join(message["content"] for message in messages) + + assert "capitalization" in prompt_text + assert "commas, periods, em dashes, and ellipses" in prompt_text + assert "homophone fixes" in prompt_text + assert "spelling fixes" in prompt_text + assert "Do not paraphrase" in prompt_text + assert "protected vocabulary" in prompt_text + assert "appears exactly once" in prompt_text + assert "Do not return speaker, start, or end fields" in prompt_text + + +def test_grammar_prompt_uses_simplified_segment_payload(): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "then lyra went their"} + ] + """ + ) + glossary = parse_glossary_yaml( + """ + glossary: + - name: "Lyra" + category: npc + summary: "Lyra is a hostile NPC." + """ + ) + section = chunk_transcript(transcript, max_section_tokens=16000)[0] + + messages = build_grammar_correction_messages(section, glossary) + transcript_json = messages[1]["content"].split("Transcript section:\n", maxsplit=1)[1] + prompt_segments = json.loads(transcript_json) + + assert prompt_segments == [{"id": 1, "original_text": "then lyra went their"}] + assert "speaker" not in prompt_segments[0] + assert "start" not in prompt_segments[0] + assert "end" not in prompt_segments[0]