diff --git a/src/audita/modules/prompts.py b/src/audita/modules/prompts.py index a729555..7f62131 100644 --- a/src/audita/modules/prompts.py +++ b/src/audita/modules/prompts.py @@ -136,10 +136,13 @@ def build_grammar_proposal_messages(section: TranscriptSection, glossary: Glossa user = ( "Review this transcript section and return only grammar cleanup corrections that should be applied.\n\n" "Rules:\n" - "- Allowed changes are punctuation, capitalization, and spacing cleanup only.\n" + "- Allowed changes are punctuation, capitalization, spacing, and article cleanup only.\n" "- You may add, remove, or adjust commas, periods, quotation marks, apostrophes, dashes, ellipses, spacing, and capitalization when the underlying words stay the same.\n" + "- You may change the whole-word article \"a\" to \"an\" or \"an\" to \"a\" when the surrounding text otherwise stays the same.\n" + "- Homophone, spoken-form, and mistranscription corrections are handled during a later review stage; do not propose them here.\n" "- Do not make word substitutions, spelling fixes, homophone fixes, filler cleanup, repetition cleanup, paraphrases, or other semantic rewrites.\n" "- Do not change one written word into a different written word, except for capitalization changes to the same letters.\n" + "- If a possible correction depends on changing a content word into a different word, omit it here rather than bundling it together with formatting cleanup.\n" "- Treat glossary names and aliases as protected spellings and context.\n" "- Do not replace, Anglicize, normalize, lowercase, or otherwise alter protected glossary names or aliases away from their glossary spelling.\n" "- Preserve canonical glossary capitalization for protected names and aliases, even if they look unusual.\n" diff --git a/src/audita/validators/deterministic.py b/src/audita/validators/deterministic.py index 74a515a..69ee5a6 100644 --- a/src/audita/validators/deterministic.py +++ b/src/audita/validators/deterministic.py @@ -64,23 +64,45 @@ class GrammarOnlyValidator: validator_name=self.name, execution_kind=self.execution_kind, decisions=[ - ValidationDecision( - proposal_index=proposal.proposal_index, - approved=_grammar_semantic_key(proposal.original_text) == _grammar_semantic_key(proposal.corrected_text), - reason=( - None - if _grammar_semantic_key(proposal.original_text) == _grammar_semantic_key(proposal.corrected_text) - else "correction is not limited to punctuation, capitalization, and spacing" - ), - ) + _grammar_validation_decision(proposal.proposal_index, proposal.original_text, proposal.corrected_text) for proposal in context.proposals ], ) +def _grammar_validation_decision(proposal_index: int, original_text: str, corrected_text: str) -> ValidationDecision: + approved = _grammar_semantic_key(original_text) == _grammar_semantic_key(corrected_text) or ( + _grammar_article_semantic_key(original_text) == _grammar_article_semantic_key(corrected_text) + ) + return ValidationDecision( + proposal_index=proposal_index, + approved=approved, + reason=None if approved else "correction is not limited to punctuation, capitalization, and spacing", + ) + + def _grammar_semantic_key(text: str) -> str: return "".join( character.casefold() for character in text if not character.isspace() and character not in _GRAMMAR_PUNCTUATION ) + + +def _grammar_article_semantic_key(text: str) -> tuple[str, ...]: + return tuple("__article__" if token in {"a", "an"} else token for token in _grammar_word_tokens(text)) + + +def _grammar_word_tokens(text: str) -> tuple[str, ...]: + tokens: list[str] = [] + current: list[str] = [] + for character in text.casefold(): + if character.isalnum(): + current.append(character) + continue + if current: + tokens.append("".join(current)) + current = [] + if current: + tokens.append("".join(current)) + return tuple(tokens) diff --git a/tests/test_llm_validators.py b/tests/test_llm_validators.py index 8f2efeb..e7d4cae 100644 --- a/tests/test_llm_validators.py +++ b/tests/test_llm_validators.py @@ -361,6 +361,46 @@ def test_grammar_only_validator_allows_formatting_only_changes(tmp_path): ] +def test_grammar_only_validator_allows_indefinite_article_changes(tmp_path): + transcript = parse_transcript_json( + """ + [ + {"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "a intelligence saving throw"}, + {"id": 2, "speaker": "A", "start": 1.0, "end": 2.0, "text": "an owl"} + ] + """ + ) + proposals = [ + CorrectionProposal( + proposal_index=0, + module_instance="grammar", + module_key="grammar", + id=1, + original_text="a intelligence saving throw", + corrected_text="an intelligence saving throw", + confidence=0.95, + ), + CorrectionProposal( + proposal_index=1, + module_instance="grammar", + module_key="grammar", + id=2, + original_text="an owl", + corrected_text="a owl", + confidence=0.95, + ), + ] + + result = GrammarOnlyValidator("grammar_only_guard").validate( + _context(proposals=proposals, transcript=transcript, llm_client=None, tmp_path=tmp_path) + ) + + assert [(decision.proposal_index, decision.approved) for decision in result.decisions] == [ + (0, True), + (1, True), + ] + + def test_grammar_only_validator_rejects_word_level_changes(tmp_path): transcript = parse_transcript_json( """ @@ -368,7 +408,8 @@ def test_grammar_only_validator_rejects_word_level_changes(tmp_path): {"id": 1, "speaker": "A", "start": 0.0, "end": 1.0, "text": "their plan"}, {"id": 2, "speaker": "A", "start": 1.0, "end": 2.0, "text": "dam"}, {"id": 3, "speaker": "A", "start": 2.0, "end": 3.0, "text": "uh"}, - {"id": 4, "speaker": "A", "start": 3.0, "end": 4.0, "text": "I I agree"} + {"id": 4, "speaker": "A", "start": 3.0, "end": 4.0, "text": "I I agree"}, + {"id": 5, "speaker": "A", "start": 4.0, "end": 5.0, "text": "Eric"} ] """ ) @@ -409,13 +450,22 @@ def test_grammar_only_validator_rejects_word_level_changes(tmp_path): corrected_text="I", confidence=0.95, ), + CorrectionProposal( + proposal_index=4, + module_instance="grammar", + module_key="grammar", + id=5, + original_text="Eric", + corrected_text="Eric's", + confidence=0.95, + ), ] result = GrammarOnlyValidator("grammar_only_guard").validate( _context(proposals=proposals, transcript=transcript, llm_client=None, tmp_path=tmp_path) ) - assert [decision.approved for decision in result.decisions] == [False, False, False, False] + assert [decision.approved for decision in result.decisions] == [False, False, False, False, False] assert all( decision.reason == "correction is not limited to punctuation, capitalization, and spacing" for decision in result.decisions diff --git a/tests/test_module_proposals.py b/tests/test_module_proposals.py index 3544f80..2906ee9 100644 --- a/tests/test_module_proposals.py +++ b/tests/test_module_proposals.py @@ -220,9 +220,10 @@ def test_grammar_module_propose_writes_diagnostics_and_returns_proposals_without assert (tmp_path / "prompt-0000.json").exists() assert (tmp_path / "corrections-0000.json").exists() prompt_text = client.calls[0]["messages"][1]["content"] - assert "punctuation, capitalization, and spacing" in prompt_text + assert "punctuation, capitalization, spacing, and article cleanup" in prompt_text assert "exact text span" in prompt_text assert "word substitutions" in prompt_text + assert "later review stage" in prompt_text def test_grammar_prompt_is_explicitly_scoped_to_formatting_cleanup(): @@ -238,9 +239,11 @@ def test_grammar_prompt_is_explicitly_scoped_to_formatting_cleanup(): messages = build_grammar_proposal_messages(section, _glossary()) combined = messages[0]["content"] + messages[1]["content"] - assert "punctuation, capitalization, and spacing" in combined + assert "punctuation, capitalization, spacing, and article cleanup" in combined assert "word substitutions" in combined assert "homophone fixes" in combined + assert "later review stage" in combined + assert "mistranscription" in combined assert '"id": 1' in messages[1]["content"] @@ -629,6 +632,72 @@ def test_process_transcript_result_runs_grammar_module_with_full_validator_chain ] +def test_process_transcript_result_grammar_module_applies_indefinite_article_cleanup(tmp_path): + transcript = parse_source_transcript_json( + """ + [ + {"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Give me a intelligence saving throw."} + ] + """ + ) + base_config = AuditaConfig.from_sources(env={}) + config = AuditaConfig( + api_key=base_config.api_key, + model=base_config.model, + base_url=base_config.base_url, + max_retries=base_config.max_retries, + max_section_tokens=base_config.max_section_tokens, + glossary_confidence_threshold=base_config.glossary_confidence_threshold, + grammar_confidence_threshold=base_config.grammar_confidence_threshold, + homophones_confidence_threshold=base_config.homophones_confidence_threshold, + spoken_word_confidence_threshold=base_config.spoken_word_confidence_threshold, + normalize_max_segment_gap=base_config.normalize_max_segment_gap, + normalize_ellipsis_gap=base_config.normalize_ellipsis_gap, + normalize_max_segment_duration=base_config.normalize_max_segment_duration, + normalize_max_segment_tokens=base_config.normalize_max_segment_tokens, + work_dir=tmp_path / "work", + work_dir_retention="always", + ) + client = FakeStructuredLLMClient( + [ + { + "corrections": [ + { + "id": 1, + "original_text": "a intelligence saving throw", + "corrected_text": "an intelligence saving throw", + "confidence": 0.95, + } + ] + }, + { + "validations": [ + { + "correction_index": 0, + "approved": True, + "confidence": 0.99, + "reason": "Does not reverse the segment meaning.", + } + ] + }, + ] + ) + + result = process_transcript_result( + transcript, + _glossary(), + config, + module_keys=["grammar"], + llm_client=client, + ) + + assert result.transcript[0].text == "Give me an intelligence saving throw." + assert [call["stage_name"] for call in client.calls] == [ + "grammar:proposal", + "grammar:meaning_reversal_review", + ] + + def test_process_transcript_result_rejects_grammar_below_threshold_before_later_validators(tmp_path): transcript = parse_source_transcript_json( """ @@ -682,3 +751,58 @@ def test_process_transcript_result_rejects_grammar_below_threshold_before_later_ assert [call["stage_name"] for call in client.calls] == ["grammar:proposal"] assert result.report.skipped_corrections[0].reason == "proposal confidence below threshold" assert result.report.modules[0].validators[1].candidate_count == 0 + + +def test_process_transcript_result_grammar_module_still_rejects_homophone_style_proposals(tmp_path): + transcript = parse_source_transcript_json( + """ + [ + {"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "ChatGPT still can't really do that with a dam."} + ] + """ + ) + base_config = AuditaConfig.from_sources(env={}) + config = AuditaConfig( + api_key=base_config.api_key, + model=base_config.model, + base_url=base_config.base_url, + max_retries=base_config.max_retries, + max_section_tokens=base_config.max_section_tokens, + glossary_confidence_threshold=base_config.glossary_confidence_threshold, + grammar_confidence_threshold=base_config.grammar_confidence_threshold, + homophones_confidence_threshold=base_config.homophones_confidence_threshold, + spoken_word_confidence_threshold=base_config.spoken_word_confidence_threshold, + normalize_max_segment_gap=base_config.normalize_max_segment_gap, + normalize_ellipsis_gap=base_config.normalize_ellipsis_gap, + normalize_max_segment_duration=base_config.normalize_max_segment_duration, + normalize_max_segment_tokens=base_config.normalize_max_segment_tokens, + work_dir=tmp_path / "work", + work_dir_retention="always", + ) + client = FakeStructuredLLMClient( + [ + { + "corrections": [ + { + "id": 1, + "original_text": "dam", + "corrected_text": "damn", + "confidence": 0.95, + } + ] + } + ] + ) + + result = process_transcript_result( + transcript, + _glossary(), + config, + module_keys=["grammar"], + llm_client=client, + ) + + assert result.transcript[0].text == "ChatGPT still can't really do that with a dam." + assert [call["stage_name"] for call in client.calls] == ["grammar:proposal"] + assert result.report.skipped_corrections[0].source == "validator:grammar_only_guard" + assert result.report.skipped_corrections[0].reason == "correction is not limited to punctuation, capitalization, and spacing"