Grammar corrections now run through a deterministic guard that blocks protected glossary names/aliases from being changed
This commit is contained in:
@@ -1,10 +1,11 @@
|
|||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from typing import Dict, Iterable, List, Literal, Optional, Tuple
|
from typing import Callable, Dict, Iterable, List, Literal, Optional, Tuple
|
||||||
|
|
||||||
from .errors import AuditaValidationError
|
from .errors import AuditaValidationError
|
||||||
from .schemas import CorrectionCandidate, TranscriptSegment
|
from .schemas import CorrectionCandidate, TranscriptSegment
|
||||||
|
|
||||||
ReplacementMode = Literal["replace_all", "require_unique"]
|
ReplacementMode = Literal["replace_all", "require_unique"]
|
||||||
|
CorrectionGuard = Callable[[str, str], Optional[str]]
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -33,6 +34,7 @@ def apply_corrections(
|
|||||||
corrections: Iterable[CorrectionCandidate],
|
corrections: Iterable[CorrectionCandidate],
|
||||||
confidence_threshold: float,
|
confidence_threshold: float,
|
||||||
replacement_mode: ReplacementMode = "replace_all",
|
replacement_mode: ReplacementMode = "replace_all",
|
||||||
|
correction_guard: Optional[CorrectionGuard] = None,
|
||||||
) -> CorrectionApplicationResult:
|
) -> CorrectionApplicationResult:
|
||||||
if not 0.0 <= confidence_threshold <= 1.0:
|
if not 0.0 <= confidence_threshold <= 1.0:
|
||||||
raise AuditaValidationError("Confidence threshold must be between 0.0 and 1.0.")
|
raise AuditaValidationError("Confidence threshold must be between 0.0 and 1.0.")
|
||||||
@@ -57,6 +59,11 @@ def apply_corrections(
|
|||||||
position = id_to_position[correction.id]
|
position = id_to_position[correction.id]
|
||||||
segment = revised[position]
|
segment = revised[position]
|
||||||
revised_text = segment.text.replace(correction.original_text, correction.corrected_text)
|
revised_text = segment.text.replace(correction.original_text, correction.corrected_text)
|
||||||
|
if correction_guard is not None:
|
||||||
|
reason = correction_guard(segment.text, revised_text)
|
||||||
|
if reason is not None:
|
||||||
|
skipped.append(_skip(correction, reason, actual_text=segment.text))
|
||||||
|
continue
|
||||||
revised[position] = segment.model_copy(update={"text": revised_text})
|
revised[position] = segment.model_copy(update={"text": revised_text})
|
||||||
applied_ids.append(correction.id)
|
applied_ids.append(correction.id)
|
||||||
|
|
||||||
|
|||||||
@@ -8,10 +8,11 @@ from uuid import uuid4
|
|||||||
|
|
||||||
from .chunking import IndexedSegment, TranscriptSection, chunk_indexed_segments
|
from .chunking import IndexedSegment, TranscriptSection, chunk_indexed_segments
|
||||||
from .config import AuditaConfig
|
from .config import AuditaConfig
|
||||||
from .corrections import ReplacementMode, SkippedCorrection, apply_corrections
|
from .corrections import CorrectionGuard, ReplacementMode, SkippedCorrection, apply_corrections
|
||||||
from .errors import AuditaError
|
from .errors import AuditaError
|
||||||
from .normalization import NormalizationResult, normalize_transcript
|
from .normalization import NormalizationResult, normalize_transcript
|
||||||
from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient
|
from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient
|
||||||
|
from .protection import ProtectedVocabulary
|
||||||
from .schemas import Glossary, SourceTranscriptSegment, TranscriptSegment, parse_transcript_json
|
from .schemas import Glossary, SourceTranscriptSegment, TranscriptSegment, parse_transcript_json
|
||||||
|
|
||||||
|
|
||||||
@@ -25,6 +26,7 @@ class StageSpec:
|
|||||||
max_llm_passes: int
|
max_llm_passes: int
|
||||||
confidence_threshold: float
|
confidence_threshold: float
|
||||||
replacement_mode: ReplacementMode
|
replacement_mode: ReplacementMode
|
||||||
|
correction_guard: Optional[CorrectionGuard] = None
|
||||||
|
|
||||||
|
|
||||||
def process_transcript(
|
def process_transcript(
|
||||||
@@ -64,6 +66,7 @@ def process_transcript(
|
|||||||
llm_client = InstructorLLMClient(config)
|
llm_client = InstructorLLMClient(config)
|
||||||
|
|
||||||
working = list(normalization_result.transcript)
|
working = list(normalization_result.transcript)
|
||||||
|
protected_vocabulary = ProtectedVocabulary.from_glossary(glossary)
|
||||||
stages = [
|
stages = [
|
||||||
StageSpec(
|
StageSpec(
|
||||||
name="glossary",
|
name="glossary",
|
||||||
@@ -78,6 +81,7 @@ def process_transcript(
|
|||||||
max_llm_passes=config.grammar_max_llm_passes,
|
max_llm_passes=config.grammar_max_llm_passes,
|
||||||
confidence_threshold=config.grammar_confidence_threshold,
|
confidence_threshold=config.grammar_confidence_threshold,
|
||||||
replacement_mode="require_unique",
|
replacement_mode="require_unique",
|
||||||
|
correction_guard=protected_vocabulary.violation_reason,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
final_skipped: List[Tuple[str, SkippedCorrection]] = []
|
final_skipped: List[Tuple[str, SkippedCorrection]] = []
|
||||||
@@ -186,6 +190,7 @@ def _run_correction_stage(
|
|||||||
corrections,
|
corrections,
|
||||||
stage.confidence_threshold,
|
stage.confidence_threshold,
|
||||||
replacement_mode=stage.replacement_mode,
|
replacement_mode=stage.replacement_mode,
|
||||||
|
correction_guard=stage.correction_guard,
|
||||||
)
|
)
|
||||||
working = application_result.transcript
|
working = application_result.transcript
|
||||||
|
|
||||||
|
|||||||
@@ -88,7 +88,9 @@ def build_grammar_correction_messages(
|
|||||||
"- Allowed corrections are only capitalization changes, punctuation changes involving commas, periods, em dashes, and ellipses, homophone fixes, and spelling fixes.\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 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"
|
"- 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"
|
"- Treat glossary names and aliases as protected spellings and context; do not introduce new glossary substitutions during this grammar pass.\n"
|
||||||
|
"- Do not autocorrect, Anglicize, replace, normalize, lowercase, or otherwise alter protected glossary names or aliases.\n"
|
||||||
|
"- Preserve canonical glossary capitalization for protected names and aliases, even if they look unusual.\n"
|
||||||
"- Use the exact id from the input segment.\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"
|
"- 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"
|
"- Choose an original_text span that appears exactly once in the current segment text.\n"
|
||||||
|
|||||||
63
src/audita/protection.py
Normal file
63
src/audita/protection.py
Normal file
@@ -0,0 +1,63 @@
|
|||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Dict, List, Optional, Pattern
|
||||||
|
|
||||||
|
from .schemas import Glossary
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ProtectedVocabulary:
|
||||||
|
canonical_by_folded: Dict[str, str]
|
||||||
|
pattern: Optional[Pattern[str]]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_glossary(cls, glossary: Glossary) -> "ProtectedVocabulary":
|
||||||
|
canonical_by_folded: Dict[str, str] = {}
|
||||||
|
for entry in glossary.glossary:
|
||||||
|
_add_term(canonical_by_folded, entry.name)
|
||||||
|
for alias in entry.aliases:
|
||||||
|
_add_term(canonical_by_folded, alias)
|
||||||
|
|
||||||
|
terms = list(canonical_by_folded.values())
|
||||||
|
if not terms:
|
||||||
|
return cls(canonical_by_folded=canonical_by_folded, pattern=None)
|
||||||
|
|
||||||
|
alternatives = sorted((re.escape(term) for term in terms), key=len, reverse=True)
|
||||||
|
pattern = re.compile(r"(?<!\w)(" + "|".join(alternatives) + r")(?!\w)", flags=re.IGNORECASE)
|
||||||
|
return cls(canonical_by_folded=canonical_by_folded, pattern=pattern)
|
||||||
|
|
||||||
|
def violation_reason(self, before: str, after: str) -> Optional[str]:
|
||||||
|
before_terms = self._terms(before)
|
||||||
|
after_terms = self._terms(after)
|
||||||
|
if len(before_terms) != len(after_terms):
|
||||||
|
return "grammar correction changes protected glossary term usage"
|
||||||
|
|
||||||
|
for before_term, after_term in zip(before_terms, after_terms):
|
||||||
|
if before_term.folded != after_term.folded:
|
||||||
|
return "grammar correction changes protected glossary term usage"
|
||||||
|
canonical = self.canonical_by_folded[after_term.folded]
|
||||||
|
if after_term.text != canonical:
|
||||||
|
return "grammar correction changes protected glossary term capitalization"
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _terms(self, text: str) -> List["_ProtectedTerm"]:
|
||||||
|
if self.pattern is None:
|
||||||
|
return []
|
||||||
|
return [
|
||||||
|
_ProtectedTerm(text=match.group(0), folded=match.group(0).casefold())
|
||||||
|
for match in self.pattern.finditer(text)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _ProtectedTerm:
|
||||||
|
text: str
|
||||||
|
folded: str
|
||||||
|
|
||||||
|
|
||||||
|
def _add_term(canonical_by_folded: Dict[str, str], term: str) -> None:
|
||||||
|
stripped = term.strip()
|
||||||
|
if not stripped:
|
||||||
|
return
|
||||||
|
canonical_by_folded.setdefault(stripped.casefold(), stripped)
|
||||||
@@ -33,7 +33,6 @@ def test_config_uses_defaults_with_api_key():
|
|||||||
assert config.normalize_max_segment_gap == DEFAULT_NORMALIZE_MAX_SEGMENT_GAP
|
assert config.normalize_max_segment_gap == DEFAULT_NORMALIZE_MAX_SEGMENT_GAP
|
||||||
assert config.normalize_max_segment_gap == 5.0
|
assert config.normalize_max_segment_gap == 5.0
|
||||||
assert config.normalize_ellipsis_gap == DEFAULT_NORMALIZE_ELLIPSIS_GAP
|
assert config.normalize_ellipsis_gap == DEFAULT_NORMALIZE_ELLIPSIS_GAP
|
||||||
assert config.normalize_ellipsis_gap == 2.0
|
|
||||||
assert config.normalize_max_segment_duration == DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION
|
assert config.normalize_max_segment_duration == DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION
|
||||||
assert config.normalize_max_segment_duration == 60.0
|
assert config.normalize_max_segment_duration == 60.0
|
||||||
assert config.normalize_max_segment_tokens == DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS
|
assert config.normalize_max_segment_tokens == DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS
|
||||||
|
|||||||
@@ -171,6 +171,56 @@ def test_apply_corrections_requires_unique_match_when_configured():
|
|||||||
assert "more than once" in result.skipped[0].reason
|
assert "more than once" in result.skipped[0].reason
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_corrections_skips_when_guard_rejects_replacement():
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hrank moves."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
correction = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Hrank",
|
||||||
|
corrected_text="Frank",
|
||||||
|
confidence=0.8,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = apply_corrections(
|
||||||
|
transcript,
|
||||||
|
[correction],
|
||||||
|
confidence_threshold=0.8,
|
||||||
|
replacement_mode="require_unique",
|
||||||
|
correction_guard=lambda before, after: "protected term changed" if before != after else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.transcript[0].text == "Hrank moves."
|
||||||
|
assert len(result.skipped) == 1
|
||||||
|
assert result.skipped[0].reason == "protected term changed"
|
||||||
|
assert result.skipped[0].actual_text == "Hrank moves."
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_corrections_without_guard_remains_permissive():
|
||||||
|
transcript = parse_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Frank moves."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
correction = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Frank",
|
||||||
|
corrected_text="Hrank",
|
||||||
|
confidence=0.8,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||||
|
|
||||||
|
assert result.transcript[0].text == "Hrank moves."
|
||||||
|
assert result.skipped == []
|
||||||
|
|
||||||
|
|
||||||
def test_apply_corrections_skips_empty_original_text():
|
def test_apply_corrections_skips_empty_original_text():
|
||||||
transcript = _transcript()
|
transcript = _transcript()
|
||||||
correction = CorrectionCandidate(
|
correction = CorrectionCandidate(
|
||||||
|
|||||||
@@ -305,6 +305,59 @@ def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path):
|
|||||||
assert revised[0].text == "I ask Chauntea."
|
assert revised[0].text == "I ask Chauntea."
|
||||||
|
|
||||||
|
|
||||||
|
def test_grammar_stage_cannot_reverse_glossary_protected_term(tmp_path):
|
||||||
|
transcript = parse_source_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Frank moves."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
glossary = parse_glossary_yaml(
|
||||||
|
"""
|
||||||
|
glossary:
|
||||||
|
- name: "Hrank"
|
||||||
|
category: pc
|
||||||
|
summary: "Hrank is a player character."
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
glossary_correction = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Frank",
|
||||||
|
corrected_text="Hrank",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
grammar_reversal = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Hrank",
|
||||||
|
corrected_text="Frank",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
fake_client = FakeLLMClient(
|
||||||
|
[
|
||||||
|
CorrectionSet(corrections=[glossary_correction]),
|
||||||
|
CorrectionSet(corrections=[grammar_reversal]),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
progress = []
|
||||||
|
|
||||||
|
revised = process_transcript(
|
||||||
|
transcript,
|
||||||
|
glossary,
|
||||||
|
_config(tmp_path, grammar_max_llm_passes=1),
|
||||||
|
llm_client=fake_client,
|
||||||
|
progress=progress.append,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert revised[0].text == "Hrank moves."
|
||||||
|
assert any("Skipping grammar correction for id 1" in message for message in progress)
|
||||||
|
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 "protected glossary term" in diagnostics["skipped_corrections"][0]["reason"]
|
||||||
|
|
||||||
|
|
||||||
def test_grammar_stage_retries_repeated_span_and_applies_unique_retry(tmp_path):
|
def test_grammar_stage_retries_repeated_span_and_applies_unique_retry(tmp_path):
|
||||||
transcript = parse_source_transcript_json(
|
transcript = parse_source_transcript_json(
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -93,7 +93,9 @@ def test_grammar_prompt_limits_readability_corrections_and_protects_glossary():
|
|||||||
assert "homophone fixes" in prompt_text
|
assert "homophone fixes" in prompt_text
|
||||||
assert "spelling fixes" in prompt_text
|
assert "spelling fixes" in prompt_text
|
||||||
assert "Do not paraphrase" in prompt_text
|
assert "Do not paraphrase" in prompt_text
|
||||||
assert "protected vocabulary" in prompt_text
|
assert "glossary names and aliases as protected spellings" in prompt_text
|
||||||
|
assert "Do not autocorrect, Anglicize, replace, normalize, lowercase" in prompt_text
|
||||||
|
assert "Preserve canonical glossary capitalization" in prompt_text
|
||||||
assert "appears exactly once" in prompt_text
|
assert "appears exactly once" in prompt_text
|
||||||
assert "Do not return speaker, start, or end fields" in prompt_text
|
assert "Do not return speaker, start, or end fields" in prompt_text
|
||||||
|
|
||||||
|
|||||||
52
tests/test_protection.py
Normal file
52
tests/test_protection.py
Normal file
@@ -0,0 +1,52 @@
|
|||||||
|
from audita.protection import ProtectedVocabulary
|
||||||
|
from audita.schemas import parse_glossary_yaml
|
||||||
|
|
||||||
|
|
||||||
|
def _vocabulary():
|
||||||
|
glossary = parse_glossary_yaml(
|
||||||
|
"""
|
||||||
|
glossary:
|
||||||
|
- name: "Hrank"
|
||||||
|
aliases:
|
||||||
|
- "Greenfield"
|
||||||
|
category: pc
|
||||||
|
summary: "Hrank Greenfield is a player character."
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
return ProtectedVocabulary.from_glossary(glossary)
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_blocks_replacing_protected_term():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("Hrank moves.", "Frank moves.") is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_blocks_lowercasing_protected_term():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("Hrank moves.", "hrank moves.") is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_allows_canonical_capitalization():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("hrank moves.", "Hrank moves.") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_allows_punctuation_around_protected_term():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("Hrank, moves.", "Hrank. Moves.") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_does_not_match_terms_inside_larger_words():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("The shrank spell worked.", "The shrank spell works.") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_protection_applies_to_aliases():
|
||||||
|
vocabulary = _vocabulary()
|
||||||
|
|
||||||
|
assert vocabulary.violation_reason("Greenfield waits.", "greenfield waits.") is not None
|
||||||
Reference in New Issue
Block a user