Removed historical prototype code
This commit is contained in:
@@ -1,6 +0,0 @@
|
||||
"""Audita transcript correction package."""
|
||||
|
||||
__all__ = ["__version__"]
|
||||
|
||||
__version__ = "0.1.0"
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
from .cli import main
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
||||
@@ -1,134 +0,0 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from math import ceil
|
||||
from typing import Any, List, Optional, Protocol
|
||||
|
||||
from .errors import AuditaValidationError
|
||||
from .schemas import TranscriptSegment, parse_transcript_json
|
||||
|
||||
|
||||
class TokenEstimatorProtocol(Protocol):
|
||||
def estimate_json(self, value: Any) -> int:
|
||||
...
|
||||
|
||||
|
||||
class TokenEstimator:
|
||||
def __init__(self, fallback_chars_per_token: int = 4) -> None:
|
||||
self._fallback_chars_per_token = fallback_chars_per_token
|
||||
self._encoding = None
|
||||
try:
|
||||
import tiktoken
|
||||
|
||||
self._encoding = tiktoken.get_encoding("cl100k_base")
|
||||
except Exception:
|
||||
self._encoding = None
|
||||
|
||||
def estimate_text(self, text: str) -> int:
|
||||
if self._encoding is not None:
|
||||
return len(self._encoding.encode(text))
|
||||
return max(1, ceil(len(text) / self._fallback_chars_per_token))
|
||||
|
||||
def estimate_json(self, value: Any) -> int:
|
||||
return self.estimate_text(json.dumps(value, ensure_ascii=False, separators=(",", ":")))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class IndexedSegment:
|
||||
index: int
|
||||
segment: TranscriptSegment
|
||||
|
||||
def transcript_payload(self) -> dict:
|
||||
return self.segment.model_dump(mode="json")
|
||||
|
||||
def prompt_payload(self) -> dict:
|
||||
return {"id": self.segment.id, "original_text": self.segment.text}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TranscriptSection:
|
||||
section_index: int
|
||||
start_index: int
|
||||
segments: List[IndexedSegment]
|
||||
token_count: int
|
||||
|
||||
def transcript_payload(self) -> List[dict]:
|
||||
return [item.transcript_payload() for item in self.segments]
|
||||
|
||||
def prompt_payload(self) -> List[dict]:
|
||||
return [item.prompt_payload() for item in self.segments]
|
||||
|
||||
def transcript_json(self) -> str:
|
||||
return json.dumps(self.transcript_payload(), ensure_ascii=False, indent=2) + "\n"
|
||||
|
||||
|
||||
def chunk_transcript(
|
||||
segments: List[TranscriptSegment],
|
||||
max_section_tokens: int,
|
||||
estimator: Optional[TokenEstimatorProtocol] = None,
|
||||
) -> List[TranscriptSection]:
|
||||
if max_section_tokens <= 0:
|
||||
raise AuditaValidationError("Maximum section token count must be greater than zero.")
|
||||
if not segments:
|
||||
raise AuditaValidationError("Transcript must contain at least one segment.")
|
||||
|
||||
indexed = [IndexedSegment(index=index, segment=segment) for index, segment in enumerate(segments)]
|
||||
return chunk_indexed_segments(indexed, max_section_tokens, estimator=estimator)
|
||||
|
||||
|
||||
def chunk_indexed_segments(
|
||||
indexed_segments: List[IndexedSegment],
|
||||
max_section_tokens: int,
|
||||
estimator: Optional[TokenEstimatorProtocol] = None,
|
||||
) -> List[TranscriptSection]:
|
||||
if max_section_tokens <= 0:
|
||||
raise AuditaValidationError("Maximum section token count must be greater than zero.")
|
||||
if not indexed_segments:
|
||||
raise AuditaValidationError("Transcript must contain at least one segment.")
|
||||
|
||||
token_estimator = TokenEstimator() if estimator is None else estimator
|
||||
indexed = list(indexed_segments)
|
||||
sections: List[TranscriptSection] = []
|
||||
current: List[IndexedSegment] = []
|
||||
current_tokens = 0
|
||||
|
||||
for item in indexed:
|
||||
single_payload = [item.prompt_payload()]
|
||||
single_tokens = token_estimator.estimate_json(single_payload)
|
||||
if single_tokens > max_section_tokens:
|
||||
raise AuditaValidationError(
|
||||
"A single transcript segment exceeds the maximum section token limit. "
|
||||
"Raise the limit or pre-split the transcript."
|
||||
)
|
||||
|
||||
candidate = current + [item]
|
||||
candidate_tokens = token_estimator.estimate_json(
|
||||
[candidate_item.prompt_payload() for candidate_item in candidate]
|
||||
)
|
||||
if current and candidate_tokens > max_section_tokens:
|
||||
sections.append(_make_section(len(sections), current, current_tokens))
|
||||
current = [item]
|
||||
current_tokens = single_tokens
|
||||
else:
|
||||
current = candidate
|
||||
current_tokens = candidate_tokens
|
||||
|
||||
if current:
|
||||
sections.append(_make_section(len(sections), current, current_tokens))
|
||||
|
||||
for section in sections:
|
||||
parse_transcript_json(section.transcript_json(), require_sequential_ids=False)
|
||||
|
||||
return sections
|
||||
|
||||
|
||||
def _make_section(
|
||||
section_index: int,
|
||||
segments: List[IndexedSegment],
|
||||
token_count: int,
|
||||
) -> TranscriptSection:
|
||||
return TranscriptSection(
|
||||
section_index=section_index,
|
||||
start_index=segments[0].index,
|
||||
segments=list(segments),
|
||||
token_count=token_count,
|
||||
)
|
||||
@@ -1,137 +0,0 @@
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional, Sequence
|
||||
|
||||
from .config import AuditaConfig, ConfigOverrides
|
||||
from .errors import AuditaError
|
||||
from .io import load_glossary, load_transcript, write_report, write_transcript
|
||||
from .pipeline import process_transcript_result
|
||||
from .schemas import transcript_to_json
|
||||
|
||||
|
||||
def main(argv: Optional[Sequence[str]] = None) -> int:
|
||||
parser = _build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
if args.command == "process":
|
||||
return _process(args)
|
||||
|
||||
parser.print_help(sys.stderr)
|
||||
return 2
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(prog="audita")
|
||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
process = subparsers.add_parser("process", help="correct a transcript using a glossary")
|
||||
process.add_argument("transcript", type=Path, help="path to the input transcript JSON")
|
||||
process.add_argument("--glossary", type=Path, required=True, help="path to the glossary YAML")
|
||||
process.add_argument("--output", type=Path, help="write corrected transcript JSON to this path")
|
||||
process.add_argument("--report-json", type=Path, help="write structured run report JSON to this path")
|
||||
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(
|
||||
"--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(
|
||||
"--grammar-validation-enabled",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=None,
|
||||
help="enable semantic validation for grammar corrections",
|
||||
)
|
||||
process.add_argument(
|
||||
"--grammar-validation-confidence-threshold",
|
||||
type=float,
|
||||
help="minimum validator confidence required to apply a validated grammar correction",
|
||||
)
|
||||
process.add_argument(
|
||||
"--grammar-spoken-form-validation-confidence-threshold",
|
||||
type=float,
|
||||
help="minimum validator confidence required to apply a spoken-form rescue correction",
|
||||
)
|
||||
process.add_argument(
|
||||
"--normalize-max-segment-gap",
|
||||
type=float,
|
||||
help="maximum same-speaker gap in seconds eligible for deterministic merging",
|
||||
)
|
||||
process.add_argument(
|
||||
"--normalize-ellipsis-gap",
|
||||
type=float,
|
||||
help="minimum same-speaker gap in seconds that uses an ellipsis joiner",
|
||||
)
|
||||
process.add_argument(
|
||||
"--normalize-max-segment-duration",
|
||||
type=float,
|
||||
help="maximum merged segment duration in seconds",
|
||||
)
|
||||
process.add_argument(
|
||||
"--normalize-max-segment-tokens",
|
||||
type=int,
|
||||
help="maximum estimated tokens for a merged segment prompt payload",
|
||||
)
|
||||
process.add_argument("--work-dir", type=Path, help="directory for per-run scratch diagnostics")
|
||||
process.add_argument(
|
||||
"--work-dir-retention",
|
||||
choices=("auto", "always", "never"),
|
||||
help="whether to retain the per-run work directory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def _process(args: argparse.Namespace) -> int:
|
||||
try:
|
||||
config = AuditaConfig.from_sources(
|
||||
overrides=ConfigOverrides(
|
||||
model=args.model,
|
||||
base_url=args.base_url,
|
||||
max_section_tokens=args.max_section_tokens,
|
||||
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,
|
||||
grammar_validation_enabled=args.grammar_validation_enabled,
|
||||
grammar_validation_confidence_threshold=args.grammar_validation_confidence_threshold,
|
||||
grammar_spoken_form_validation_confidence_threshold=(
|
||||
args.grammar_spoken_form_validation_confidence_threshold
|
||||
),
|
||||
normalize_max_segment_gap=args.normalize_max_segment_gap,
|
||||
normalize_ellipsis_gap=args.normalize_ellipsis_gap,
|
||||
normalize_max_segment_duration=args.normalize_max_segment_duration,
|
||||
normalize_max_segment_tokens=args.normalize_max_segment_tokens,
|
||||
work_dir=args.work_dir,
|
||||
work_dir_retention=args.work_dir_retention,
|
||||
)
|
||||
)
|
||||
transcript = load_transcript(args.transcript)
|
||||
glossary = load_glossary(args.glossary)
|
||||
result = process_transcript_result(
|
||||
transcript,
|
||||
glossary,
|
||||
config,
|
||||
progress=lambda message: print(message, file=sys.stderr),
|
||||
)
|
||||
|
||||
if args.output is not None:
|
||||
write_transcript(args.output, result.transcript)
|
||||
else:
|
||||
sys.stdout.write(transcript_to_json(result.transcript))
|
||||
if args.report_json is not None:
|
||||
write_report(args.report_json, result.report)
|
||||
return 0
|
||||
except AuditaError as exc:
|
||||
print(f"audita: error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
@@ -1,313 +0,0 @@
|
||||
import math
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Mapping, Optional
|
||||
|
||||
from .errors import AuditaConfigError
|
||||
|
||||
|
||||
DEFAULT_MODEL = "openrouter/google/gemma-4-31b-it"
|
||||
DEFAULT_BASE_URL = "https://openrouter.ai/api/v1"
|
||||
DEFAULT_MAX_SECTION_TOKENS = 6144
|
||||
DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD = 0.80
|
||||
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD = 0.80
|
||||
DEFAULT_MAX_RETRIES = 3
|
||||
DEFAULT_WORK_DIR = "/tmp/audita"
|
||||
DEFAULT_WORK_DIR_RETENTION = "auto"
|
||||
DEFAULT_GLOSSARY_MAX_LLM_PASSES = 3
|
||||
DEFAULT_GRAMMAR_MAX_LLM_PASSES = 3
|
||||
DEFAULT_GRAMMAR_VALIDATION_ENABLED = True
|
||||
DEFAULT_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD = 0.80
|
||||
DEFAULT_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD = 0.80
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP = 4.0
|
||||
DEFAULT_NORMALIZE_ELLIPSIS_GAP = 3.5
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION = 60.0
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS = 2048
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConfigOverrides:
|
||||
model: Optional[str] = None
|
||||
base_url: Optional[str] = None
|
||||
max_section_tokens: Optional[int] = 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
|
||||
grammar_validation_enabled: Optional[bool] = None
|
||||
grammar_validation_confidence_threshold: Optional[float] = None
|
||||
grammar_spoken_form_validation_confidence_threshold: Optional[float] = None
|
||||
normalize_max_segment_gap: Optional[float] = None
|
||||
normalize_ellipsis_gap: Optional[float] = None
|
||||
normalize_max_segment_duration: Optional[float] = None
|
||||
normalize_max_segment_tokens: Optional[int] = None
|
||||
work_dir: Optional[Path] = None
|
||||
work_dir_retention: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AuditaConfig:
|
||||
api_key: str
|
||||
model: str = DEFAULT_MODEL
|
||||
base_url: str = DEFAULT_BASE_URL
|
||||
max_section_tokens: int = DEFAULT_MAX_SECTION_TOKENS
|
||||
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
|
||||
grammar_validation_enabled: bool = DEFAULT_GRAMMAR_VALIDATION_ENABLED
|
||||
grammar_validation_confidence_threshold: float = DEFAULT_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD
|
||||
grammar_spoken_form_validation_confidence_threshold: float = (
|
||||
DEFAULT_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD
|
||||
)
|
||||
normalize_max_segment_gap: float = DEFAULT_NORMALIZE_MAX_SEGMENT_GAP
|
||||
normalize_ellipsis_gap: float = DEFAULT_NORMALIZE_ELLIPSIS_GAP
|
||||
normalize_max_segment_duration: float = DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION
|
||||
normalize_max_segment_tokens: int = DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS
|
||||
work_dir: Path = Path(DEFAULT_WORK_DIR)
|
||||
work_dir_retention: str = DEFAULT_WORK_DIR_RETENTION
|
||||
|
||||
@classmethod
|
||||
def from_sources(
|
||||
cls,
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
overrides: Optional[ConfigOverrides] = None,
|
||||
) -> "AuditaConfig":
|
||||
source = os.environ if env is None else env
|
||||
selected = ConfigOverrides() if overrides is None else overrides
|
||||
|
||||
api_key = _get_required_env(source, "OPENROUTER_API_KEY")
|
||||
model = selected.model or source.get("AUDITA_MODEL") or DEFAULT_MODEL
|
||||
base_url = selected.base_url or source.get("AUDITA_BASE_URL") or DEFAULT_BASE_URL
|
||||
max_section_tokens = _select_int(
|
||||
selected.max_section_tokens,
|
||||
source.get("AUDITA_MAX_SECTION_TOKENS"),
|
||||
DEFAULT_MAX_SECTION_TOKENS,
|
||||
"AUDITA_MAX_SECTION_TOKENS",
|
||||
)
|
||||
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,
|
||||
source.get("AUDITA_MAX_RETRIES"),
|
||||
DEFAULT_MAX_RETRIES,
|
||||
"AUDITA_MAX_RETRIES",
|
||||
)
|
||||
glossary_max_llm_passes = _select_int(
|
||||
selected.glossary_max_llm_passes,
|
||||
source.get("AUDITA_GLOSSARY_MAX_LLM_PASSES"),
|
||||
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",
|
||||
)
|
||||
grammar_validation_enabled = _select_bool(
|
||||
selected.grammar_validation_enabled,
|
||||
source.get("AUDITA_GRAMMAR_VALIDATION_ENABLED"),
|
||||
DEFAULT_GRAMMAR_VALIDATION_ENABLED,
|
||||
"AUDITA_GRAMMAR_VALIDATION_ENABLED",
|
||||
)
|
||||
grammar_validation_confidence_threshold = _select_float(
|
||||
selected.grammar_validation_confidence_threshold,
|
||||
source.get("AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD"),
|
||||
DEFAULT_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD,
|
||||
"AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD",
|
||||
)
|
||||
grammar_spoken_form_validation_confidence_threshold = _select_float(
|
||||
selected.grammar_spoken_form_validation_confidence_threshold,
|
||||
source.get("AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD"),
|
||||
DEFAULT_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD,
|
||||
"AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD",
|
||||
)
|
||||
normalize_max_segment_gap = _select_float(
|
||||
selected.normalize_max_segment_gap,
|
||||
source.get("AUDITA_NORMALIZE_MAX_SEGMENT_GAP"),
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP,
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_GAP",
|
||||
)
|
||||
normalize_ellipsis_gap = _select_float(
|
||||
selected.normalize_ellipsis_gap,
|
||||
source.get("AUDITA_NORMALIZE_ELLIPSIS_GAP"),
|
||||
DEFAULT_NORMALIZE_ELLIPSIS_GAP,
|
||||
"AUDITA_NORMALIZE_ELLIPSIS_GAP",
|
||||
)
|
||||
normalize_max_segment_duration = _select_float(
|
||||
selected.normalize_max_segment_duration,
|
||||
source.get("AUDITA_NORMALIZE_MAX_SEGMENT_DURATION"),
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION,
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_DURATION",
|
||||
)
|
||||
normalize_max_segment_tokens = _select_int(
|
||||
selected.normalize_max_segment_tokens,
|
||||
source.get("AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS"),
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS,
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS",
|
||||
)
|
||||
work_dir_value = selected.work_dir or Path(source.get("AUDITA_WORK_DIR") or DEFAULT_WORK_DIR)
|
||||
work_dir_retention = _select_choice(
|
||||
selected.work_dir_retention,
|
||||
source.get("AUDITA_WORK_DIR_RETENTION"),
|
||||
DEFAULT_WORK_DIR_RETENTION,
|
||||
"AUDITA_WORK_DIR_RETENTION",
|
||||
("auto", "always", "never"),
|
||||
)
|
||||
|
||||
config = cls(
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
base_url=base_url,
|
||||
max_section_tokens=max_section_tokens,
|
||||
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,
|
||||
grammar_validation_enabled=grammar_validation_enabled,
|
||||
grammar_validation_confidence_threshold=grammar_validation_confidence_threshold,
|
||||
grammar_spoken_form_validation_confidence_threshold=grammar_spoken_form_validation_confidence_threshold,
|
||||
normalize_max_segment_gap=normalize_max_segment_gap,
|
||||
normalize_ellipsis_gap=normalize_ellipsis_gap,
|
||||
normalize_max_segment_duration=normalize_max_segment_duration,
|
||||
normalize_max_segment_tokens=normalize_max_segment_tokens,
|
||||
work_dir=Path(work_dir_value),
|
||||
work_dir_retention=work_dir_retention,
|
||||
)
|
||||
config.validate()
|
||||
return config
|
||||
|
||||
def validate(self) -> None:
|
||||
if not self.api_key.strip():
|
||||
raise AuditaConfigError("OPENROUTER_API_KEY must not be empty.")
|
||||
if not self.model.strip():
|
||||
raise AuditaConfigError("AUDITA_MODEL must not be empty.")
|
||||
if not self.base_url.strip():
|
||||
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.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.")
|
||||
if not 0.0 <= self.grammar_validation_confidence_threshold <= 1.0:
|
||||
raise AuditaConfigError(
|
||||
"AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0."
|
||||
)
|
||||
if not 0.0 <= self.grammar_spoken_form_validation_confidence_threshold <= 1.0:
|
||||
raise AuditaConfigError(
|
||||
"AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD must be between 0.0 and 1.0."
|
||||
)
|
||||
if not math.isfinite(self.normalize_max_segment_gap):
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_MAX_SEGMENT_GAP must be finite.")
|
||||
if not math.isfinite(self.normalize_ellipsis_gap):
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_ELLIPSIS_GAP must be finite.")
|
||||
if not math.isfinite(self.normalize_max_segment_duration):
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_MAX_SEGMENT_DURATION must be finite.")
|
||||
if self.normalize_max_segment_gap < 0:
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_MAX_SEGMENT_GAP must be greater than or equal to zero.")
|
||||
if self.normalize_ellipsis_gap < 0:
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_ELLIPSIS_GAP must be greater than or equal to zero.")
|
||||
if self.normalize_ellipsis_gap > self.normalize_max_segment_gap:
|
||||
raise AuditaConfigError(
|
||||
"AUDITA_NORMALIZE_ELLIPSIS_GAP must be less than or equal to AUDITA_NORMALIZE_MAX_SEGMENT_GAP."
|
||||
)
|
||||
if self.normalize_max_segment_duration <= 0:
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_MAX_SEGMENT_DURATION must be greater than zero.")
|
||||
if self.normalize_max_segment_tokens <= 0:
|
||||
raise AuditaConfigError("AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS must be greater than zero.")
|
||||
if self.work_dir_retention not in ("auto", "always", "never"):
|
||||
raise AuditaConfigError("AUDITA_WORK_DIR_RETENTION must be one of auto, always, or never.")
|
||||
|
||||
|
||||
def _get_required_env(env: Mapping[str, str], name: str) -> str:
|
||||
value = env.get(name)
|
||||
if value is None or not value.strip():
|
||||
raise AuditaConfigError(f"{name} is required.")
|
||||
return value
|
||||
|
||||
|
||||
def _select_int(cli_value: Optional[int], env_value: Optional[str], default: int, name: str) -> int:
|
||||
if cli_value is not None:
|
||||
return cli_value
|
||||
if env_value is None:
|
||||
return default
|
||||
try:
|
||||
return int(env_value)
|
||||
except ValueError as exc:
|
||||
raise AuditaConfigError(f"{name} must be an integer.") from exc
|
||||
|
||||
|
||||
def _select_bool(
|
||||
cli_value: Optional[bool],
|
||||
env_value: Optional[str],
|
||||
default: bool,
|
||||
name: str,
|
||||
) -> bool:
|
||||
if cli_value is not None:
|
||||
return cli_value
|
||||
if env_value is None:
|
||||
return default
|
||||
normalized = env_value.strip().casefold()
|
||||
if normalized in ("1", "true", "yes", "on"):
|
||||
return True
|
||||
if normalized in ("0", "false", "no", "off"):
|
||||
return False
|
||||
raise AuditaConfigError(f"{name} must be a boolean.")
|
||||
|
||||
|
||||
def _select_float(
|
||||
cli_value: Optional[float],
|
||||
env_value: Optional[str],
|
||||
default: float,
|
||||
name: str,
|
||||
) -> float:
|
||||
if cli_value is not None:
|
||||
return cli_value
|
||||
if env_value is None:
|
||||
return default
|
||||
try:
|
||||
return float(env_value)
|
||||
except ValueError as exc:
|
||||
raise AuditaConfigError(f"{name} must be a number.") from exc
|
||||
|
||||
|
||||
def _select_choice(
|
||||
cli_value: Optional[str],
|
||||
env_value: Optional[str],
|
||||
default: str,
|
||||
name: str,
|
||||
choices: tuple[str, ...],
|
||||
) -> str:
|
||||
if cli_value is not None:
|
||||
value = cli_value
|
||||
elif env_value is not None:
|
||||
value = env_value
|
||||
else:
|
||||
return default
|
||||
normalized = value.strip().casefold()
|
||||
if normalized not in choices:
|
||||
allowed = ", ".join(choices)
|
||||
raise AuditaConfigError(f"{name} must be one of {allowed}.")
|
||||
return normalized
|
||||
@@ -1,150 +0,0 @@
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Callable, Dict, Iterable, List, Literal, Optional, Tuple
|
||||
|
||||
from .errors import AuditaValidationError
|
||||
from .schemas import CorrectionCandidate, TranscriptSegment
|
||||
|
||||
ReplacementMode = Literal["replace_all", "require_unique"]
|
||||
CorrectionGuard = Callable[[str, str], Optional[str]]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SkippedCorrection:
|
||||
id: int
|
||||
reason: str
|
||||
original_text: str
|
||||
corrected_text: str
|
||||
confidence: float
|
||||
actual_text: Optional[str] = None
|
||||
validation_confidence: Optional[float] = None
|
||||
validation_reason: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AppliedCorrection:
|
||||
id: int
|
||||
original_text: str
|
||||
corrected_text: str
|
||||
confidence: float
|
||||
segment_text_before: str
|
||||
segment_text_after: str
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CorrectionApplicationResult:
|
||||
transcript: List[TranscriptSegment]
|
||||
skipped: List[SkippedCorrection]
|
||||
applied_corrections: List[AppliedCorrection]
|
||||
applied_ids: List[int]
|
||||
ignored_ids: List[int]
|
||||
ignored: List[SkippedCorrection]
|
||||
|
||||
|
||||
def apply_corrections(
|
||||
transcript: List[TranscriptSegment],
|
||||
corrections: Iterable[CorrectionCandidate],
|
||||
confidence_threshold: float,
|
||||
replacement_mode: ReplacementMode = "replace_all",
|
||||
correction_guard: Optional[CorrectionGuard] = None,
|
||||
) -> 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)
|
||||
skipped: List[SkippedCorrection] = []
|
||||
applied_corrections: List[AppliedCorrection] = []
|
||||
applied_ids: List[int] = []
|
||||
ignored_ids: List[int] = []
|
||||
ignored: List[SkippedCorrection] = []
|
||||
for correction in corrections:
|
||||
if correction.confidence < confidence_threshold:
|
||||
ignored_ids.append(correction.id)
|
||||
ignored.append(_skip(correction, "correction confidence below threshold"))
|
||||
continue
|
||||
|
||||
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
|
||||
|
||||
position = id_to_position[correction.id]
|
||||
segment = revised[position]
|
||||
revised_text = segment.text.replace(correction.original_text, correction.corrected_text)
|
||||
if correction_guard is not None:
|
||||
reason = correction_guard(correction.original_text, correction.corrected_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})
|
||||
applied_corrections.append(
|
||||
AppliedCorrection(
|
||||
id=correction.id,
|
||||
original_text=correction.original_text,
|
||||
corrected_text=correction.corrected_text,
|
||||
confidence=correction.confidence,
|
||||
segment_text_before=segment.text,
|
||||
segment_text_after=revised_text,
|
||||
)
|
||||
)
|
||||
applied_ids.append(correction.id)
|
||||
|
||||
return CorrectionApplicationResult(
|
||||
transcript=revised,
|
||||
skipped=skipped,
|
||||
applied_corrections=applied_corrections,
|
||||
applied_ids=applied_ids,
|
||||
ignored_ids=ignored_ids,
|
||||
ignored=ignored,
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
segment = transcript[id_to_position[correction.id]]
|
||||
if correction.original_text == "":
|
||||
return "original_text is empty", segment.text
|
||||
if correction.original_text == correction.corrected_text:
|
||||
return "original_text and corrected_text are identical", segment.text
|
||||
|
||||
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
|
||||
|
||||
|
||||
def _id_to_position(transcript: List[TranscriptSegment]) -> Dict[int, int]:
|
||||
return {segment.id: position for position, segment in enumerate(transcript)}
|
||||
|
||||
|
||||
def _skip(
|
||||
correction: CorrectionCandidate,
|
||||
reason: str,
|
||||
actual_text: Optional[str] = None,
|
||||
) -> SkippedCorrection:
|
||||
return SkippedCorrection(
|
||||
id=correction.id,
|
||||
reason=reason,
|
||||
original_text=correction.original_text,
|
||||
corrected_text=correction.corrected_text,
|
||||
confidence=correction.confidence,
|
||||
actual_text=actual_text,
|
||||
)
|
||||
@@ -1,15 +0,0 @@
|
||||
class AuditaError(Exception):
|
||||
"""Base exception for user-facing Audita failures."""
|
||||
|
||||
|
||||
class AuditaValidationError(AuditaError):
|
||||
"""Raised when input data does not match Audita's expected schema."""
|
||||
|
||||
|
||||
class AuditaConfigError(AuditaError):
|
||||
"""Raised when runtime configuration is invalid or incomplete."""
|
||||
|
||||
|
||||
class AuditaLLMError(AuditaError):
|
||||
"""Raised when an LLM request or structured response fails."""
|
||||
|
||||
@@ -1,22 +0,0 @@
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
from .reporting import RunReport
|
||||
from .schemas import Glossary, SourceTranscriptSegment, TranscriptSegment
|
||||
from .schemas import parse_glossary_yaml, parse_source_transcript_json, transcript_to_json
|
||||
|
||||
|
||||
def load_transcript(path: Path) -> List[SourceTranscriptSegment]:
|
||||
return parse_source_transcript_json(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def load_glossary(path: Path) -> Glossary:
|
||||
return parse_glossary_yaml(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def write_transcript(path: Path, segments: List[TranscriptSegment]) -> None:
|
||||
path.write_text(transcript_to_json(segments), encoding="utf-8")
|
||||
|
||||
|
||||
def write_report(path: Path, report: RunReport) -> None:
|
||||
path.write_text(report.to_json(), encoding="utf-8")
|
||||
@@ -1,80 +0,0 @@
|
||||
from typing import List
|
||||
|
||||
from .config import AuditaConfig
|
||||
from .errors import AuditaLLMError
|
||||
from .prompts import Message
|
||||
from .schemas import CorrectionSet, GrammarSpokenFormValidationSet, GrammarValidationSet
|
||||
|
||||
|
||||
class InstructorLLMClient:
|
||||
def __init__(self, config: AuditaConfig) -> None:
|
||||
try:
|
||||
import instructor
|
||||
from openai import OpenAI
|
||||
except ImportError as exc:
|
||||
raise AuditaLLMError(
|
||||
"The LLM dependencies are not installed. Run `uv sync` before using audita."
|
||||
) from exc
|
||||
|
||||
self._instructor = instructor
|
||||
openai_client = OpenAI(api_key=config.api_key, base_url=config.base_url)
|
||||
self._client = instructor.patch(openai_client, mode=instructor.Mode.TOOLS)
|
||||
|
||||
def create_corrections(self, messages: List[Message], config: AuditaConfig) -> CorrectionSet:
|
||||
model = _normalize_openrouter_model(config.model)
|
||||
try:
|
||||
return self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
response_model=CorrectionSet,
|
||||
max_retries=config.max_retries,
|
||||
extra_body={"provider": {"require_parameters": True}},
|
||||
)
|
||||
except Exception as exc:
|
||||
raise AuditaLLMError(
|
||||
"LLM correction request failed. Confirm the configured OpenRouter model "
|
||||
"supports tool calling or structured outputs."
|
||||
) from exc
|
||||
|
||||
def create_grammar_validations(self, messages: List[Message], config: AuditaConfig) -> GrammarValidationSet:
|
||||
model = _normalize_openrouter_model(config.model)
|
||||
try:
|
||||
return self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
response_model=GrammarValidationSet,
|
||||
max_retries=config.max_retries,
|
||||
extra_body={"provider": {"require_parameters": True}},
|
||||
)
|
||||
except Exception as exc:
|
||||
raise AuditaLLMError(
|
||||
"LLM grammar validation request failed. Confirm the configured OpenRouter model "
|
||||
"supports tool calling or structured outputs."
|
||||
) from exc
|
||||
|
||||
def create_grammar_spoken_form_validations(
|
||||
self,
|
||||
messages: List[Message],
|
||||
config: AuditaConfig,
|
||||
) -> GrammarSpokenFormValidationSet:
|
||||
model = _normalize_openrouter_model(config.model)
|
||||
try:
|
||||
return self._client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
response_model=GrammarSpokenFormValidationSet,
|
||||
max_retries=config.max_retries,
|
||||
extra_body={"provider": {"require_parameters": True}},
|
||||
)
|
||||
except Exception as exc:
|
||||
raise AuditaLLMError(
|
||||
"LLM spoken-form grammar validation request failed. Confirm the configured OpenRouter model "
|
||||
"supports tool calling or structured outputs."
|
||||
) from exc
|
||||
|
||||
|
||||
def _normalize_openrouter_model(model: str) -> str:
|
||||
prefix = "openrouter/"
|
||||
if model.startswith(prefix):
|
||||
return model[len(prefix) :]
|
||||
return model
|
||||
@@ -1,176 +0,0 @@
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from .chunking import TokenEstimator, TokenEstimatorProtocol
|
||||
from .schemas import SourceTranscriptSegment, TranscriptSegment
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NormalizationSummary:
|
||||
source_segment_count: int
|
||||
normalized_segment_count: int
|
||||
merge_count: int
|
||||
max_segment_gap: float
|
||||
ellipsis_gap: float
|
||||
max_segment_duration: float
|
||||
max_segment_tokens: int
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NormalizationResult:
|
||||
transcript: List[TranscriptSegment]
|
||||
summary: NormalizationSummary
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _WorkingSegment:
|
||||
speaker: str
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
order: int
|
||||
|
||||
|
||||
def normalize_transcript(
|
||||
segments: List[SourceTranscriptSegment],
|
||||
max_segment_gap: float,
|
||||
ellipsis_gap: float,
|
||||
max_segment_duration: float,
|
||||
max_segment_tokens: int,
|
||||
estimator: Optional[TokenEstimatorProtocol] = None,
|
||||
) -> NormalizationResult:
|
||||
token_estimator = TokenEstimator() if estimator is None else estimator
|
||||
working = [
|
||||
_WorkingSegment(
|
||||
speaker=segment.speaker,
|
||||
start=segment.start,
|
||||
end=segment.end,
|
||||
text=segment.text,
|
||||
order=index,
|
||||
)
|
||||
for index, segment in enumerate(segments)
|
||||
]
|
||||
working.sort(key=lambda segment: (segment.start, segment.end, segment.order))
|
||||
|
||||
merge_count = 0
|
||||
while True:
|
||||
candidate_index = _shortest_mergeable_gap_index(
|
||||
working,
|
||||
max_segment_gap,
|
||||
ellipsis_gap,
|
||||
max_segment_duration,
|
||||
max_segment_tokens,
|
||||
token_estimator,
|
||||
)
|
||||
if candidate_index is None:
|
||||
break
|
||||
left = working[candidate_index]
|
||||
right = working[candidate_index + 1]
|
||||
working[candidate_index : candidate_index + 2] = [
|
||||
_merge_segments(left, right, ellipsis_gap)
|
||||
]
|
||||
merge_count += 1
|
||||
|
||||
normalized = _assign_ids(working)
|
||||
summary = NormalizationSummary(
|
||||
source_segment_count=len(segments),
|
||||
normalized_segment_count=len(normalized),
|
||||
merge_count=merge_count,
|
||||
max_segment_gap=max_segment_gap,
|
||||
ellipsis_gap=ellipsis_gap,
|
||||
max_segment_duration=max_segment_duration,
|
||||
max_segment_tokens=max_segment_tokens,
|
||||
)
|
||||
return NormalizationResult(transcript=normalized, summary=summary)
|
||||
|
||||
|
||||
def _shortest_mergeable_gap_index(
|
||||
segments: List[_WorkingSegment],
|
||||
max_segment_gap: float,
|
||||
ellipsis_gap: float,
|
||||
max_segment_duration: float,
|
||||
max_segment_tokens: int,
|
||||
estimator: TokenEstimatorProtocol,
|
||||
) -> Optional[int]:
|
||||
best_index = None
|
||||
best_gap = None
|
||||
for index in range(len(segments) - 1):
|
||||
left = segments[index]
|
||||
right = segments[index + 1]
|
||||
gap = right.start - left.end
|
||||
if not _can_merge(
|
||||
left,
|
||||
right,
|
||||
gap,
|
||||
max_segment_gap,
|
||||
ellipsis_gap,
|
||||
max_segment_duration,
|
||||
max_segment_tokens,
|
||||
estimator,
|
||||
):
|
||||
continue
|
||||
if best_gap is None or gap < best_gap:
|
||||
best_index = index
|
||||
best_gap = gap
|
||||
return best_index
|
||||
|
||||
|
||||
def _can_merge(
|
||||
left: _WorkingSegment,
|
||||
right: _WorkingSegment,
|
||||
gap: float,
|
||||
max_segment_gap: float,
|
||||
ellipsis_gap: float,
|
||||
max_segment_duration: float,
|
||||
max_segment_tokens: int,
|
||||
estimator: TokenEstimatorProtocol,
|
||||
) -> bool:
|
||||
if left.speaker != right.speaker:
|
||||
return False
|
||||
if gap < 0 or gap > max_segment_gap:
|
||||
return False
|
||||
if right.end - left.start > max_segment_duration:
|
||||
return False
|
||||
merged_text = _joined_text(left.text, right.text, gap, ellipsis_gap)
|
||||
return _estimate_prompt_tokens(merged_text, estimator) <= max_segment_tokens
|
||||
|
||||
|
||||
def _merge_segments(
|
||||
left: _WorkingSegment,
|
||||
right: _WorkingSegment,
|
||||
ellipsis_gap: float,
|
||||
) -> _WorkingSegment:
|
||||
gap = right.start - left.end
|
||||
return _WorkingSegment(
|
||||
speaker=left.speaker,
|
||||
start=left.start,
|
||||
end=right.end,
|
||||
text=_joined_text(left.text, right.text, gap, ellipsis_gap),
|
||||
order=left.order,
|
||||
)
|
||||
|
||||
|
||||
def _joined_text(left_text: str, right_text: str, gap: float, ellipsis_gap: float) -> str:
|
||||
joiner = " " if gap <= ellipsis_gap else " ... "
|
||||
return f"{left_text.rstrip()}{joiner}{right_text.lstrip()}"
|
||||
|
||||
|
||||
def _estimate_prompt_tokens(text: str, estimator: TokenEstimatorProtocol) -> int:
|
||||
return estimator.estimate_json([{"id": 1, "original_text": text}])
|
||||
|
||||
|
||||
def _assign_ids(segments: List[_WorkingSegment]) -> List[TranscriptSegment]:
|
||||
ordered = sorted(segments, key=lambda segment: (segment.start, segment.end, segment.order))
|
||||
return [
|
||||
TranscriptSegment(
|
||||
id=index + 1,
|
||||
speaker=segment.speaker,
|
||||
start=segment.start,
|
||||
end=segment.end,
|
||||
text=segment.text,
|
||||
)
|
||||
for index, segment in enumerate(ordered)
|
||||
]
|
||||
@@ -1,85 +0,0 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import List, Protocol
|
||||
|
||||
from .chunking import TranscriptSection
|
||||
from .config import AuditaConfig
|
||||
from .prompts import build_glossary_correction_messages, build_grammar_correction_messages
|
||||
from .schemas import (
|
||||
CorrectionCandidate,
|
||||
CorrectionSet,
|
||||
Glossary,
|
||||
GrammarSpokenFormValidationSet,
|
||||
GrammarValidationSet,
|
||||
)
|
||||
|
||||
|
||||
class LLMClient(Protocol):
|
||||
def create_corrections(self, messages: List[dict], config: AuditaConfig) -> CorrectionSet:
|
||||
...
|
||||
|
||||
def create_grammar_validations(self, messages: List[dict], config: AuditaConfig) -> GrammarValidationSet:
|
||||
...
|
||||
|
||||
def create_grammar_spoken_form_validations(
|
||||
self,
|
||||
messages: List[dict],
|
||||
config: AuditaConfig,
|
||||
) -> GrammarSpokenFormValidationSet:
|
||||
...
|
||||
|
||||
|
||||
class CorrectionPass(Protocol):
|
||||
def run(
|
||||
self,
|
||||
section: TranscriptSection,
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
run_dir: Path,
|
||||
retry_pass: bool = False,
|
||||
) -> List[CorrectionCandidate]:
|
||||
...
|
||||
|
||||
|
||||
class GlossaryCorrectionPass:
|
||||
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_glossary_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)
|
||||
|
||||
|
||||
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)
|
||||
@@ -1,624 +0,0 @@
|
||||
import json
|
||||
import shutil
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
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 CorrectionGuard, ReplacementMode, SkippedCorrection, apply_corrections
|
||||
from .errors import AuditaError
|
||||
from .normalization import NormalizationResult, normalize_transcript
|
||||
from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient
|
||||
from .prompts import build_grammar_spoken_form_validation_messages, build_grammar_validation_messages
|
||||
from .protection import ProtectedVocabulary
|
||||
from .reporting import AppliedChange, ProcessResult, ReportedSkippedCorrection, RunReport
|
||||
from .semantic_validation import (
|
||||
filter_with_meaning_preserving_validations,
|
||||
filter_with_spoken_form_validations,
|
||||
keep_corrections_with_indexes,
|
||||
select_grammar_validation_candidates,
|
||||
)
|
||||
from .schemas import CorrectionCandidate, Glossary, SourceTranscriptSegment, 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
|
||||
correction_guard: Optional[CorrectionGuard] = None
|
||||
protected_vocabulary: Optional[ProtectedVocabulary] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StageRunResult:
|
||||
transcript: List[TranscriptSegment]
|
||||
applied_changes: List[AppliedChange]
|
||||
skipped_corrections: List[ReportedSkippedCorrection]
|
||||
|
||||
|
||||
def process_transcript(
|
||||
transcript: List[SourceTranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client: Optional[LLMClient] = None,
|
||||
progress: Optional[ProgressCallback] = None,
|
||||
) -> List[TranscriptSegment]:
|
||||
return process_transcript_result(
|
||||
transcript,
|
||||
glossary,
|
||||
config,
|
||||
llm_client=llm_client,
|
||||
progress=progress,
|
||||
).transcript
|
||||
|
||||
|
||||
def process_transcript_result(
|
||||
transcript: List[SourceTranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
llm_client: Optional[LLMClient] = None,
|
||||
progress: Optional[ProgressCallback] = None,
|
||||
) -> ProcessResult:
|
||||
run_dir = _create_run_dir(config.work_dir)
|
||||
stage_summaries: List[dict] = []
|
||||
normalization_summary: Optional[dict] = None
|
||||
applied_changes: List[AppliedChange] = []
|
||||
final_skipped: List[ReportedSkippedCorrection] = []
|
||||
working: List[TranscriptSegment] = []
|
||||
try:
|
||||
_log(progress, f"Created work directory {run_dir}")
|
||||
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries, work_dir_retained=True)
|
||||
|
||||
normalization_result = normalize_transcript(
|
||||
transcript,
|
||||
max_segment_gap=config.normalize_max_segment_gap,
|
||||
ellipsis_gap=config.normalize_ellipsis_gap,
|
||||
max_segment_duration=config.normalize_max_segment_duration,
|
||||
max_segment_tokens=config.normalize_max_segment_tokens,
|
||||
)
|
||||
normalization_summary = normalization_result.summary.to_dict()
|
||||
_write_normalization_diagnostics(run_dir, transcript, normalization_result)
|
||||
_log(
|
||||
progress,
|
||||
"Normalized transcript from "
|
||||
f"{normalization_result.summary.source_segment_count} to "
|
||||
f"{normalization_result.summary.normalized_segment_count} segments",
|
||||
)
|
||||
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries, work_dir_retained=True)
|
||||
|
||||
if llm_client is None:
|
||||
from .llm import InstructorLLMClient
|
||||
|
||||
llm_client = InstructorLLMClient(config)
|
||||
|
||||
working = list(normalization_result.transcript)
|
||||
protected_vocabulary = ProtectedVocabulary.from_glossary(glossary)
|
||||
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",
|
||||
correction_guard=protected_vocabulary.violation_reason,
|
||||
protected_vocabulary=protected_vocabulary,
|
||||
),
|
||||
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",
|
||||
correction_guard=protected_vocabulary.violation_reason,
|
||||
protected_vocabulary=protected_vocabulary,
|
||||
),
|
||||
]
|
||||
|
||||
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, normalization_summary, stage_summaries, work_dir_retained=True)
|
||||
|
||||
stage_result = _run_correction_stage(
|
||||
working,
|
||||
glossary,
|
||||
config,
|
||||
stage,
|
||||
run_dir,
|
||||
normalization_summary,
|
||||
stage_dir,
|
||||
stage_summaries,
|
||||
stage_summary["passes"],
|
||||
llm_client,
|
||||
progress,
|
||||
)
|
||||
working = stage_result.transcript
|
||||
applied_changes.extend(stage_result.applied_changes)
|
||||
final_skipped.extend(stage_result.skipped_corrections)
|
||||
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries, work_dir_retained=True)
|
||||
|
||||
_write_skipped_corrections(run_dir, final_skipped)
|
||||
for skipped in final_skipped:
|
||||
_log(
|
||||
progress,
|
||||
f"Skipping {skipped.stage} correction for id {skipped.id}: {skipped.reason}",
|
||||
)
|
||||
revised = _sort_transcript_chronologically(working)
|
||||
except Exception as exc:
|
||||
_write_skipped_corrections(run_dir, final_skipped)
|
||||
report = _build_run_report(
|
||||
config=config,
|
||||
normalization_summary=normalization_summary,
|
||||
stage_summaries=stage_summaries,
|
||||
applied_changes=applied_changes,
|
||||
skipped_corrections=final_skipped,
|
||||
work_dir_retention=config.work_dir_retention,
|
||||
work_dir_retained=True,
|
||||
run_dir=run_dir,
|
||||
transcript=working,
|
||||
status="failed",
|
||||
error=str(exc),
|
||||
)
|
||||
_write_run_report(run_dir / "report.json", report)
|
||||
message = f"{exc} Diagnostics preserved at {run_dir}"
|
||||
if isinstance(exc, AuditaError):
|
||||
raise type(exc)(message) from exc
|
||||
raise AuditaError(message) from exc
|
||||
|
||||
work_dir_retained = _should_retain_run_dir(config.work_dir_retention, bool(final_skipped))
|
||||
report = _build_run_report(
|
||||
config=config,
|
||||
normalization_summary=normalization_summary,
|
||||
stage_summaries=stage_summaries,
|
||||
applied_changes=applied_changes,
|
||||
skipped_corrections=final_skipped,
|
||||
work_dir_retention=config.work_dir_retention,
|
||||
work_dir_retained=work_dir_retained,
|
||||
run_dir=run_dir,
|
||||
transcript=revised,
|
||||
status="success",
|
||||
error=None,
|
||||
)
|
||||
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries, work_dir_retained=work_dir_retained)
|
||||
if work_dir_retained:
|
||||
_write_run_report(run_dir / "report.json", report)
|
||||
if final_skipped:
|
||||
_log(progress, f"Skipped correction diagnostics preserved at {run_dir}")
|
||||
else:
|
||||
_log(progress, f"Work directory preserved at {run_dir}")
|
||||
else:
|
||||
shutil.rmtree(run_dir)
|
||||
_log(progress, "Removed work directory after successful run")
|
||||
return ProcessResult(
|
||||
transcript=revised,
|
||||
report=report,
|
||||
run_dir=run_dir,
|
||||
work_dir_retained=work_dir_retained,
|
||||
)
|
||||
|
||||
|
||||
def _run_correction_stage(
|
||||
transcript: List[TranscriptSegment],
|
||||
glossary: Glossary,
|
||||
config: AuditaConfig,
|
||||
stage: StageSpec,
|
||||
run_dir: Path,
|
||||
normalization_summary: Optional[dict],
|
||||
stage_dir: Path,
|
||||
stage_summaries: List[dict],
|
||||
pass_summaries: List[dict],
|
||||
llm_client: LLMClient,
|
||||
progress: Optional[ProgressCallback],
|
||||
) -> StageRunResult:
|
||||
working = list(transcript)
|
||||
stage_applied_changes: List[AppliedChange] = []
|
||||
unresolved_retry_skips: Dict[int, ReportedSkippedCorrection] = {}
|
||||
final_nonretry_skips: List[ReportedSkippedCorrection] = []
|
||||
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
corrections_for_application, validation_skips, validation_summary = _validate_grammar_corrections(
|
||||
working,
|
||||
corrections,
|
||||
config,
|
||||
stage,
|
||||
pass_dir,
|
||||
llm_client,
|
||||
)
|
||||
final_nonretry_skips.extend(
|
||||
_reported_skip(stage.name, pass_number, skipped) for skipped in validation_skips
|
||||
)
|
||||
|
||||
application_result = apply_corrections(
|
||||
working,
|
||||
corrections_for_application,
|
||||
stage.confidence_threshold,
|
||||
replacement_mode=stage.replacement_mode,
|
||||
correction_guard=stage.correction_guard,
|
||||
)
|
||||
working = application_result.transcript
|
||||
stage_applied_changes.extend(
|
||||
AppliedChange(
|
||||
stage=stage.name,
|
||||
pass_number=pass_number,
|
||||
id=applied.id,
|
||||
original_text=applied.original_text,
|
||||
corrected_text=applied.corrected_text,
|
||||
confidence=applied.confidence,
|
||||
segment_text_before=applied.segment_text_before,
|
||||
segment_text_after=applied.segment_text_after,
|
||||
)
|
||||
for applied in application_result.applied_corrections
|
||||
)
|
||||
|
||||
next_retry_skips: Dict[int, ReportedSkippedCorrection] = {}
|
||||
for ignored in application_result.ignored:
|
||||
reported_ignored = _reported_skip(stage.name, pass_number, ignored)
|
||||
if _is_retryable_skip(reported_ignored, working):
|
||||
next_retry_skips[reported_ignored.id] = reported_ignored
|
||||
else:
|
||||
final_nonretry_skips.append(reported_ignored)
|
||||
for skipped in application_result.skipped:
|
||||
reported_skip = _reported_skip(stage.name, pass_number, skipped)
|
||||
if _is_retryable_skip(reported_skip, working):
|
||||
next_retry_skips[reported_skip.id] = reported_skip
|
||||
else:
|
||||
final_nonretry_skips.append(reported_skip)
|
||||
unresolved_retry_skips = next_retry_skips
|
||||
|
||||
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_corrections),
|
||||
"ignored_below_threshold_count": len(application_result.ignored_ids),
|
||||
"skipped_count": len(application_result.skipped),
|
||||
"retry_segment_count": len(unresolved_retry_skips),
|
||||
**validation_summary,
|
||||
}
|
||||
)
|
||||
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries, work_dir_retained=True)
|
||||
|
||||
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 StageRunResult(
|
||||
transcript=working,
|
||||
applied_changes=stage_applied_changes,
|
||||
skipped_corrections=final_skipped,
|
||||
)
|
||||
|
||||
|
||||
def _validate_grammar_corrections(
|
||||
transcript: List[TranscriptSegment],
|
||||
corrections: List[CorrectionCandidate],
|
||||
config: AuditaConfig,
|
||||
stage: StageSpec,
|
||||
pass_dir: Path,
|
||||
llm_client: LLMClient,
|
||||
) -> Tuple[List[CorrectionCandidate], List[SkippedCorrection], dict]:
|
||||
validation_summary = {
|
||||
"validation_candidate_count": 0,
|
||||
"validation_approved_count": 0,
|
||||
"validation_rejected_count": 0,
|
||||
"validation_bypassed_count": 0,
|
||||
"spoken_form_validation_candidate_count": 0,
|
||||
"spoken_form_validation_approved_count": 0,
|
||||
"spoken_form_validation_rejected_count": 0,
|
||||
}
|
||||
if stage.name != "grammar":
|
||||
return corrections, [], validation_summary
|
||||
if not config.grammar_validation_enabled:
|
||||
validation_summary["validation_bypassed_count"] = len(corrections)
|
||||
return corrections, [], validation_summary
|
||||
if stage.protected_vocabulary is None:
|
||||
validation_summary["validation_bypassed_count"] = len(corrections)
|
||||
return corrections, [], validation_summary
|
||||
|
||||
candidates, bypassed_count = select_grammar_validation_candidates(
|
||||
transcript,
|
||||
corrections,
|
||||
stage.confidence_threshold,
|
||||
stage.replacement_mode,
|
||||
stage.protected_vocabulary,
|
||||
)
|
||||
validation_summary["validation_candidate_count"] = len(candidates)
|
||||
validation_summary["validation_bypassed_count"] = bypassed_count
|
||||
if not candidates:
|
||||
return corrections, [], validation_summary
|
||||
|
||||
payload = [candidate.to_prompt_payload() for candidate in candidates]
|
||||
messages = build_grammar_validation_messages(payload)
|
||||
prompt_path = pass_dir / "validation-prompt-0000.json"
|
||||
prompt_path.write_text(json.dumps(messages, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
response = llm_client.create_grammar_validations(messages, config)
|
||||
response_path = pass_dir / "validation-response-0000.json"
|
||||
response_path.write_text(response.model_dump_json(indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
meaning_result = filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
response,
|
||||
config.grammar_validation_confidence_threshold,
|
||||
)
|
||||
validation_summary["validation_approved_count"] = meaning_result.approved_count
|
||||
validation_summary["validation_rejected_count"] = meaning_result.rejected_count
|
||||
validation_summary["spoken_form_validation_candidate_count"] = len(meaning_result.rescue_candidates)
|
||||
|
||||
candidate_indexes = {candidate.correction_index for candidate in candidates}
|
||||
allowed_indexes = set(range(len(corrections))) - candidate_indexes
|
||||
allowed_indexes.update(meaning_result.approved_correction_indexes)
|
||||
|
||||
if not meaning_result.rescue_candidates:
|
||||
return keep_corrections_with_indexes(corrections, allowed_indexes), [], validation_summary
|
||||
|
||||
spoken_form_payload = [candidate.to_prompt_payload() for candidate in meaning_result.rescue_candidates]
|
||||
spoken_form_messages = build_grammar_spoken_form_validation_messages(spoken_form_payload)
|
||||
spoken_form_prompt_path = pass_dir / "spoken-form-validation-prompt-0000.json"
|
||||
spoken_form_prompt_path.write_text(
|
||||
json.dumps(spoken_form_messages, ensure_ascii=False, indent=2) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
spoken_form_response = llm_client.create_grammar_spoken_form_validations(spoken_form_messages, config)
|
||||
spoken_form_response_path = pass_dir / "spoken-form-validation-response-0000.json"
|
||||
spoken_form_response_path.write_text(spoken_form_response.model_dump_json(indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
spoken_form_result = filter_with_spoken_form_validations(
|
||||
meaning_result.rescue_candidates,
|
||||
spoken_form_response,
|
||||
config.grammar_spoken_form_validation_confidence_threshold,
|
||||
)
|
||||
validation_summary["spoken_form_validation_approved_count"] = spoken_form_result.approved_count
|
||||
validation_summary["spoken_form_validation_rejected_count"] = spoken_form_result.rejected_count
|
||||
allowed_indexes.update(spoken_form_result.approved_correction_indexes)
|
||||
return (
|
||||
keep_corrections_with_indexes(corrections, allowed_indexes),
|
||||
spoken_form_result.skipped,
|
||||
validation_summary,
|
||||
)
|
||||
|
||||
|
||||
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")
|
||||
run_dir = work_dir / f"run-{timestamp}-{uuid4().hex[:8]}"
|
||||
run_dir.mkdir()
|
||||
return run_dir
|
||||
|
||||
|
||||
def _write_run_metadata(
|
||||
run_dir: Path,
|
||||
config: AuditaConfig,
|
||||
normalization_summary: Optional[dict],
|
||||
stage_summaries: List[dict],
|
||||
work_dir_retained: bool,
|
||||
) -> None:
|
||||
metadata = {
|
||||
**_config_summary(config),
|
||||
"normalization": normalization_summary,
|
||||
"stages": stage_summaries,
|
||||
"work_dir_retained": work_dir_retained,
|
||||
}
|
||||
(run_dir / "metadata.json").write_text(
|
||||
json.dumps(metadata, ensure_ascii=False, indent=2) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def _write_normalization_diagnostics(
|
||||
run_dir: Path,
|
||||
source: List[SourceTranscriptSegment],
|
||||
result: NormalizationResult,
|
||||
) -> None:
|
||||
normalization_dir = run_dir / "normalization"
|
||||
normalization_dir.mkdir()
|
||||
(normalization_dir / "source-transcript.json").write_text(
|
||||
_segments_to_json(source),
|
||||
encoding="utf-8",
|
||||
)
|
||||
(normalization_dir / "normalized-transcript.json").write_text(
|
||||
_segments_to_json(result.transcript),
|
||||
encoding="utf-8",
|
||||
)
|
||||
(normalization_dir / "summary.json").write_text(
|
||||
json.dumps(result.summary.to_dict(), ensure_ascii=False, indent=2) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def _segments_to_json(segments: List[SourceTranscriptSegment]) -> str:
|
||||
payload = [segment.model_dump(mode="json", exclude_none=True) for segment in segments]
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2) + "\n"
|
||||
|
||||
|
||||
def _write_and_validate_section(run_dir: Path, section: TranscriptSection) -> None:
|
||||
section_path = run_dir / f"section-{section.section_index:04d}.json"
|
||||
section_json = section.transcript_json()
|
||||
section_path.write_text(section_json, encoding="utf-8")
|
||||
parse_transcript_json(section_json, require_sequential_ids=False)
|
||||
|
||||
|
||||
def _write_skipped_corrections(run_dir: Path, skipped: List[ReportedSkippedCorrection]) -> None:
|
||||
skipped_path = run_dir / "skipped-corrections.json"
|
||||
skipped_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"skipped_corrections": [item.to_dict() for item in skipped]
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def _indexed_segments_for_ids(
|
||||
transcript: List[TranscriptSegment],
|
||||
ids: List[int],
|
||||
) -> List[IndexedSegment]:
|
||||
id_to_position = {segment.id: position for position, segment in enumerate(transcript)}
|
||||
return [
|
||||
IndexedSegment(index=id_to_position[segment_id], segment=transcript[id_to_position[segment_id]])
|
||||
for segment_id in ids
|
||||
if segment_id in id_to_position
|
||||
]
|
||||
|
||||
|
||||
def _is_retryable_skip(skipped: ReportedSkippedCorrection, transcript: List[TranscriptSegment]) -> bool:
|
||||
return any(segment.id == skipped.id for segment in transcript)
|
||||
|
||||
|
||||
def _sort_transcript_chronologically(
|
||||
transcript: List[TranscriptSegment],
|
||||
) -> List[TranscriptSegment]:
|
||||
indexed = list(enumerate(transcript))
|
||||
indexed.sort(key=lambda item: (item[1].start, item[1].end, item[0]))
|
||||
return [segment for _, segment in indexed]
|
||||
|
||||
|
||||
def _log(progress: Optional[ProgressCallback], message: str) -> None:
|
||||
if progress is not None:
|
||||
progress(message)
|
||||
|
||||
|
||||
def _reported_skip(stage: str, pass_number: int, skipped: SkippedCorrection) -> ReportedSkippedCorrection:
|
||||
return ReportedSkippedCorrection(
|
||||
stage=stage,
|
||||
pass_number=pass_number,
|
||||
id=skipped.id,
|
||||
reason=skipped.reason,
|
||||
original_text=skipped.original_text,
|
||||
corrected_text=skipped.corrected_text,
|
||||
confidence=skipped.confidence,
|
||||
actual_text=skipped.actual_text,
|
||||
validation_confidence=skipped.validation_confidence,
|
||||
validation_reason=skipped.validation_reason,
|
||||
)
|
||||
|
||||
|
||||
def _should_retain_run_dir(work_dir_retention: str, has_final_skipped: bool) -> bool:
|
||||
if work_dir_retention == "always":
|
||||
return True
|
||||
if work_dir_retention == "never":
|
||||
return False
|
||||
return has_final_skipped
|
||||
|
||||
|
||||
def _config_summary(config: AuditaConfig) -> dict:
|
||||
return {
|
||||
"model": config.model,
|
||||
"base_url": config.base_url,
|
||||
"max_section_tokens": config.max_section_tokens,
|
||||
"glossary_confidence_threshold": config.glossary_confidence_threshold,
|
||||
"grammar_confidence_threshold": config.grammar_confidence_threshold,
|
||||
"grammar_validation_enabled": config.grammar_validation_enabled,
|
||||
"grammar_validation_confidence_threshold": config.grammar_validation_confidence_threshold,
|
||||
"grammar_spoken_form_validation_confidence_threshold": (
|
||||
config.grammar_spoken_form_validation_confidence_threshold
|
||||
),
|
||||
"max_retries": config.max_retries,
|
||||
"glossary_max_llm_passes": config.glossary_max_llm_passes,
|
||||
"grammar_max_llm_passes": config.grammar_max_llm_passes,
|
||||
"normalize_max_segment_gap": config.normalize_max_segment_gap,
|
||||
"normalize_ellipsis_gap": config.normalize_ellipsis_gap,
|
||||
"normalize_max_segment_duration": config.normalize_max_segment_duration,
|
||||
"normalize_max_segment_tokens": config.normalize_max_segment_tokens,
|
||||
"work_dir_retention": config.work_dir_retention,
|
||||
}
|
||||
|
||||
|
||||
def _build_run_report(
|
||||
config: AuditaConfig,
|
||||
normalization_summary: Optional[dict],
|
||||
stage_summaries: List[dict],
|
||||
applied_changes: List[AppliedChange],
|
||||
skipped_corrections: List[ReportedSkippedCorrection],
|
||||
work_dir_retention: str,
|
||||
work_dir_retained: bool,
|
||||
run_dir: Path,
|
||||
transcript: List[TranscriptSegment],
|
||||
status: str,
|
||||
error: Optional[str],
|
||||
) -> RunReport:
|
||||
totals = {
|
||||
"output_segment_count": len(transcript),
|
||||
"applied_change_count": len(applied_changes),
|
||||
"skipped_correction_count": len(skipped_corrections),
|
||||
}
|
||||
if normalization_summary is not None:
|
||||
totals["source_segment_count"] = normalization_summary["source_segment_count"]
|
||||
totals["normalized_segment_count"] = normalization_summary["normalized_segment_count"]
|
||||
return RunReport(
|
||||
status=status,
|
||||
config=_config_summary(config),
|
||||
normalization=normalization_summary,
|
||||
stages=stage_summaries,
|
||||
applied_changes=applied_changes,
|
||||
skipped_corrections=skipped_corrections,
|
||||
totals=totals,
|
||||
work_dir_retention=work_dir_retention,
|
||||
work_dir_retained=work_dir_retained,
|
||||
work_dir=str(run_dir) if work_dir_retained else None,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
def _write_run_report(path: Path, report: RunReport) -> None:
|
||||
path.write_text(report.to_json(), encoding="utf-8")
|
||||
@@ -1,159 +0,0 @@
|
||||
import json
|
||||
from typing import Dict, List
|
||||
|
||||
from .chunking import TranscriptSection
|
||||
from .schemas import Glossary
|
||||
|
||||
|
||||
Message = Dict[str, str]
|
||||
|
||||
|
||||
def build_glossary_correction_messages(
|
||||
section: TranscriptSection,
|
||||
glossary: Glossary,
|
||||
retry_pass: bool = False,
|
||||
) -> List[Message]:
|
||||
glossary_json = json.dumps(glossary.model_dump(mode="json", exclude_none=True), ensure_ascii=False, indent=2)
|
||||
section_json = json.dumps(section.prompt_payload(), ensure_ascii=False, indent=2)
|
||||
|
||||
system = (
|
||||
"You are Audita, a careful transcript correction assistant. "
|
||||
"Identify only transcription errors that are strongly supported by the glossary. "
|
||||
"A valid correction must be acoustically plausible: the original transcript text "
|
||||
"should sound similar to the proposed correction when spoken aloud. "
|
||||
"Do not make generic grammar, spelling, capitalization, or style edits. "
|
||||
"Do not substitute an unrelated glossary term just because it could fit the topic. "
|
||||
"Do not rewrite unchanged transcript segments. "
|
||||
"Preserve speaker names, timestamps, and meaning."
|
||||
)
|
||||
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 the span that needs replacement.\n\n"
|
||||
)
|
||||
|
||||
user = (
|
||||
"Review this transcript section and return only corrections that should be applied.\n\n"
|
||||
f"{retry_guidance}"
|
||||
"Rules:\n"
|
||||
"- Correct domain-specific names, aliases, jargon, deities, locations, NPCs, players, and similar terms only when both the glossary and surrounding transcript context support the correction.\n"
|
||||
"- The correction must be likely to fix a transcription error: the original words should be phonetically or acoustically similar to the corrected words in spoken English.\n"
|
||||
"- Appropriate example: correcting \"gestures\" to \"Jesters\" can be valid if \"Jesters\" appears in the glossary and nearby context supports that inference.\n"
|
||||
"- Inappropriate example: correcting \"Lyra\" to \"Jesters\" should be omitted because those words are not similar in spoken English, even if \"Jesters\" appears in the glossary.\n"
|
||||
"- Do not replace one clear glossary term, character name, location, or ordinary word with a different glossary term unless it is a plausible mishearing.\n"
|
||||
"- Treat glossary names and aliases already present in the transcript as protected spellings.\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"
|
||||
"- Assign high confidence only when the correction is supported by glossary evidence, local context, and spoken-word similarity; otherwise omit the correction.\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"
|
||||
"- 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"Glossary:\n{glossary_json}\n\n"
|
||||
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", exclude_none=True), 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 glossary names and aliases as protected spellings and context.\n"
|
||||
"- You may correct clear transcription or spelling errors toward glossary names or aliases when the correction preserves the spoken content.\n"
|
||||
"- Do not autocorrect, Anglicize, replace, normalize, lowercase, or otherwise alter protected glossary names or aliases that already appear correctly in the transcript.\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"
|
||||
"- 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}]
|
||||
|
||||
|
||||
def build_grammar_validation_messages(validation_payload: List[dict]) -> List[Message]:
|
||||
payload_json = json.dumps(validation_payload, ensure_ascii=False, indent=2)
|
||||
system = (
|
||||
"You are Audita, a conservative semantic validation assistant. "
|
||||
"Evaluate whether each proposed grammar correction preserves the same spoken content and meaning. "
|
||||
"Do not judge whether the correction is more polished, and do not try to rescue likely homophone or transcription fixes. "
|
||||
"Judge only whether the corrected text preserves the same written meaning."
|
||||
)
|
||||
user = (
|
||||
"Review these proposed grammar corrections and decide whether each correction preserves meaning.\n\n"
|
||||
"Rules:\n"
|
||||
"- Return one validation decision for every correction_index in the input.\n"
|
||||
"- Reject corrections that add or remove negation, reverse meaning, introduce antonyms, change names, change quantities, change actions, change who did what, or otherwise substantively alter the speaker's meaning.\n"
|
||||
"- Reject corrections like changing \"became visible\" to \"became invisible\" because that reverses the meaning.\n"
|
||||
"- Allow capitalization and punctuation changes when they preserve meaning.\n"
|
||||
"- If a correction changes meaning because the original transcript may have used the wrong homophone or a phonetic misspelling, reject it here; that question is handled in a separate spoken-form validation step.\n"
|
||||
"- Do not use domain knowledge to second-guess protected glossary terms; glossary-protected corrections are excluded from this validation step.\n"
|
||||
"- Each returned validation must contain only correction_index, is_meaning_preserving, confidence, and reason.\n"
|
||||
"- confidence must be between 0.0 and 1.0.\n\n"
|
||||
f"Corrections to validate:\n{payload_json}"
|
||||
)
|
||||
return [{"role": "system", "content": system}, {"role": "user", "content": user}]
|
||||
|
||||
|
||||
def build_grammar_spoken_form_validation_messages(validation_payload: List[dict]) -> List[Message]:
|
||||
payload_json = json.dumps(validation_payload, ensure_ascii=False, indent=2)
|
||||
system = (
|
||||
"You are Audita, a conservative spoken-form validation assistant. "
|
||||
"Evaluate whether each rejected grammar correction is still a likely homophone, spoken-form, or transcription fix "
|
||||
"supported by the local segment context. "
|
||||
"Approve only corrections that plausibly recover the intended spoken words from a mistaken transcript rendering."
|
||||
)
|
||||
user = (
|
||||
"Review these rejected grammar corrections and decide whether each one is a likely spoken-form correction.\n\n"
|
||||
"Rules:\n"
|
||||
"- Return one validation decision for every correction_index in the input.\n"
|
||||
"- Approve a correction only when the original transcript text is plausibly a mistaken homophone, phonetic rendering, or transcription error, and the corrected text better matches the likely spoken words in context.\n"
|
||||
"- Allow examples like changing \"dam\" to \"damn\" when the surrounding phrase strongly supports the intended spoken phrase.\n"
|
||||
"- Reject examples like changing \"became visible\" to \"became invisible\" because that is a semantic reversal, not a likely spoken-form correction.\n"
|
||||
"- Reject paraphrases, stylistic rewrites, content additions or removals, and any meaning-changing edit that is not clearly explained by a transcription or spoken-form mistake.\n"
|
||||
"- Do not use domain knowledge to second-guess protected glossary terms; glossary-protected corrections are excluded from this validation step.\n"
|
||||
"- Each returned validation must contain only correction_index, is_likely_spoken_form_correction, confidence, and reason.\n"
|
||||
"- confidence must be between 0.0 and 1.0.\n\n"
|
||||
f"Corrections to validate:\n{payload_json}"
|
||||
)
|
||||
return [{"role": "system", "content": system}, {"role": "user", "content": user}]
|
||||
@@ -1,121 +0,0 @@
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Pattern
|
||||
|
||||
from .schemas import Glossary
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProtectedVocabulary:
|
||||
terms_by_folded: Dict[str, "_ProtectedTermDefinition"]
|
||||
pattern: Optional[Pattern[str]]
|
||||
|
||||
@classmethod
|
||||
def from_glossary(cls, glossary: Glossary) -> "ProtectedVocabulary":
|
||||
terms_by_folded: Dict[str, _ProtectedTermDefinition] = {}
|
||||
for identity, entry in enumerate(glossary.glossary):
|
||||
entry_terms = [entry.name, *entry.aliases]
|
||||
for term in entry_terms:
|
||||
_add_term(terms_by_folded, term, identity)
|
||||
_add_term(terms_by_folded, f"{term}s", identity)
|
||||
if entry.plural is not None:
|
||||
_add_term(terms_by_folded, entry.plural, identity)
|
||||
|
||||
terms = [definition.canonical for definition in terms_by_folded.values()]
|
||||
if not terms:
|
||||
return cls(terms_by_folded=terms_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(terms_by_folded=terms_by_folded, pattern=pattern)
|
||||
|
||||
def violation_reason(self, before: str, after: str) -> Optional[str]:
|
||||
before_occurrences = self._occurrences_by_identity(before)
|
||||
after_occurrences = self._occurrences_by_identity(after)
|
||||
|
||||
reason = self._validate_identity_preservation(before_occurrences, after_occurrences)
|
||||
if reason is not None:
|
||||
return reason
|
||||
return self._validate_capitalization_transitions(before_occurrences, after_occurrences)
|
||||
|
||||
def _occurrences(self, text: str) -> List["_ProtectedOccurrence"]:
|
||||
if self.pattern is None:
|
||||
return []
|
||||
occurrences = []
|
||||
for match in self.pattern.finditer(text):
|
||||
matched_text = match.group(0)
|
||||
definition = self.terms_by_folded[matched_text.casefold()]
|
||||
occurrences.append(
|
||||
_ProtectedOccurrence(
|
||||
text=matched_text,
|
||||
identity=definition.identity,
|
||||
canonical=definition.canonical,
|
||||
)
|
||||
)
|
||||
return occurrences
|
||||
|
||||
def _occurrences_by_identity(self, text: str) -> Dict[int, List["_ProtectedOccurrence"]]:
|
||||
occurrences_by_identity: Dict[int, List["_ProtectedOccurrence"]] = {}
|
||||
for occurrence in self._occurrences(text):
|
||||
occurrences_by_identity.setdefault(occurrence.identity, []).append(occurrence)
|
||||
return occurrences_by_identity
|
||||
|
||||
def _validate_identity_preservation(
|
||||
self,
|
||||
before_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
after_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
) -> Optional[str]:
|
||||
for identity, before_items in before_occurrences.items():
|
||||
if len(after_occurrences.get(identity, [])) < len(before_items):
|
||||
return "correction changes protected glossary term usage"
|
||||
return None
|
||||
|
||||
def _validate_capitalization_transitions(
|
||||
self,
|
||||
before_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
after_occurrences: Dict[int, List["_ProtectedOccurrence"]],
|
||||
) -> Optional[str]:
|
||||
for identity, after_items in after_occurrences.items():
|
||||
before_items = before_occurrences.get(identity, [])
|
||||
before_count = len(before_items)
|
||||
for index, after_item in enumerate(after_items):
|
||||
if index < before_count:
|
||||
before_item = before_items[index]
|
||||
if after_item.text == before_item.text:
|
||||
continue
|
||||
if after_item.text == after_item.canonical:
|
||||
continue
|
||||
return "correction changes protected glossary term capitalization"
|
||||
if after_item.text != after_item.canonical:
|
||||
return "correction changes protected glossary term capitalization"
|
||||
return None
|
||||
|
||||
def contains_term(self, text: str) -> bool:
|
||||
return bool(self._occurrences(text))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ProtectedOccurrence:
|
||||
text: str
|
||||
identity: int
|
||||
canonical: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ProtectedTermDefinition:
|
||||
identity: int
|
||||
canonical: str
|
||||
|
||||
|
||||
def _add_term(
|
||||
terms_by_folded: Dict[str, _ProtectedTermDefinition],
|
||||
term: str,
|
||||
identity: int,
|
||||
) -> None:
|
||||
stripped = term.strip()
|
||||
if not stripped:
|
||||
return
|
||||
terms_by_folded.setdefault(
|
||||
stripped.casefold(),
|
||||
_ProtectedTermDefinition(identity=identity, canonical=stripped),
|
||||
)
|
||||
@@ -1 +0,0 @@
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
import json
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
from .schemas import TranscriptSegment
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AppliedChange:
|
||||
stage: str
|
||||
pass_number: int
|
||||
id: int
|
||||
original_text: str
|
||||
corrected_text: str
|
||||
confidence: float
|
||||
segment_text_before: str
|
||||
segment_text_after: str
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ReportedSkippedCorrection:
|
||||
stage: str
|
||||
pass_number: int
|
||||
id: int
|
||||
reason: str
|
||||
original_text: str
|
||||
corrected_text: str
|
||||
confidence: float
|
||||
actual_text: Optional[str] = None
|
||||
validation_confidence: Optional[float] = None
|
||||
validation_reason: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RunReport:
|
||||
status: str
|
||||
config: dict
|
||||
normalization: Optional[dict]
|
||||
stages: List[dict]
|
||||
applied_changes: List[AppliedChange]
|
||||
skipped_corrections: List[ReportedSkippedCorrection]
|
||||
totals: dict
|
||||
work_dir_retention: str
|
||||
work_dir_retained: bool
|
||||
work_dir: Optional[str]
|
||||
error: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"status": self.status,
|
||||
"config": self.config,
|
||||
"normalization": self.normalization,
|
||||
"stages": self.stages,
|
||||
"applied_changes": [item.to_dict() for item in self.applied_changes],
|
||||
"skipped_corrections": [item.to_dict() for item in self.skipped_corrections],
|
||||
"totals": self.totals,
|
||||
"work_dir_retention": self.work_dir_retention,
|
||||
"work_dir_retained": self.work_dir_retained,
|
||||
"work_dir": self.work_dir,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
def to_json(self) -> str:
|
||||
return json.dumps(self.to_dict(), ensure_ascii=False, indent=2) + "\n"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProcessResult:
|
||||
transcript: List[TranscriptSegment]
|
||||
report: RunReport
|
||||
run_dir: Path
|
||||
work_dir_retained: bool
|
||||
@@ -1,350 +0,0 @@
|
||||
import json
|
||||
import math
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictStr, TypeAdapter
|
||||
from pydantic import ValidationError, field_validator, model_validator
|
||||
|
||||
from .errors import AuditaValidationError
|
||||
|
||||
|
||||
class TranscriptSegment(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: int = Field(ge=1)
|
||||
speaker: StrictStr
|
||||
start: float
|
||||
end: float
|
||||
text: StrictStr
|
||||
|
||||
@field_validator("id", mode="before")
|
||||
@classmethod
|
||||
def require_integer_id(cls, value: Any) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("must be an integer")
|
||||
return value
|
||||
|
||||
@field_validator("speaker", "text")
|
||||
@classmethod
|
||||
def require_non_empty_text(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("must not be empty")
|
||||
return value
|
||||
|
||||
@field_validator("start", "end", mode="before")
|
||||
@classmethod
|
||||
def require_number(cls, value: Any) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("must be a JSON number")
|
||||
number = float(value)
|
||||
if not math.isfinite(number):
|
||||
raise ValueError("must be finite")
|
||||
if number < 0:
|
||||
raise ValueError("must be non-negative")
|
||||
return number
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_ordered_times(self) -> "TranscriptSegment":
|
||||
if self.end < self.start:
|
||||
raise ValueError("end must be greater than or equal to start")
|
||||
return self
|
||||
|
||||
|
||||
class SourceTranscriptSegment(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: Optional[int] = Field(default=None, ge=1)
|
||||
speaker: StrictStr
|
||||
start: float
|
||||
end: float
|
||||
text: StrictStr
|
||||
|
||||
@field_validator("id", mode="before")
|
||||
@classmethod
|
||||
def require_optional_integer_id(cls, value: Any) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("must be an integer")
|
||||
return value
|
||||
|
||||
@field_validator("speaker", "text")
|
||||
@classmethod
|
||||
def require_non_empty_text(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("must not be empty")
|
||||
return value
|
||||
|
||||
@field_validator("start", "end", mode="before")
|
||||
@classmethod
|
||||
def require_number(cls, value: Any) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("must be a JSON number")
|
||||
number = float(value)
|
||||
if not math.isfinite(number):
|
||||
raise ValueError("must be finite")
|
||||
if number < 0:
|
||||
raise ValueError("must be non-negative")
|
||||
return number
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_ordered_times(self) -> "SourceTranscriptSegment":
|
||||
if self.end < self.start:
|
||||
raise ValueError("end must be greater than or equal to start")
|
||||
return self
|
||||
|
||||
|
||||
class GlossaryEntry(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
name: StrictStr
|
||||
aliases: List[StrictStr] = Field(default_factory=list)
|
||||
plural: Optional[StrictStr] = None
|
||||
category: StrictStr
|
||||
summary: StrictStr
|
||||
|
||||
@field_validator("name", "category", "summary")
|
||||
@classmethod
|
||||
def require_non_empty_text(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("must not be empty")
|
||||
return value
|
||||
|
||||
@field_validator("aliases")
|
||||
@classmethod
|
||||
def require_non_empty_aliases(cls, aliases: List[str]) -> List[str]:
|
||||
for alias in aliases:
|
||||
if not alias.strip():
|
||||
raise ValueError("aliases must not contain empty strings")
|
||||
return aliases
|
||||
|
||||
@field_validator("plural")
|
||||
@classmethod
|
||||
def require_non_empty_plural(cls, plural: Optional[str]) -> Optional[str]:
|
||||
if plural is not None and not plural.strip():
|
||||
raise ValueError("plural must not be empty")
|
||||
return plural
|
||||
|
||||
|
||||
class Glossary(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
glossary: List[GlossaryEntry]
|
||||
|
||||
@field_validator("glossary")
|
||||
@classmethod
|
||||
def require_entries(cls, entries: List[GlossaryEntry]) -> List[GlossaryEntry]:
|
||||
if not entries:
|
||||
raise ValueError("glossary must contain at least one entry")
|
||||
return entries
|
||||
|
||||
|
||||
class CorrectionCandidate(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
id: int = Field(ge=1)
|
||||
original_text: StrictStr
|
||||
corrected_text: StrictStr
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
|
||||
@field_validator("id", mode="before")
|
||||
@classmethod
|
||||
def require_integer_id(cls, value: Any) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("must be an integer")
|
||||
return value
|
||||
|
||||
@field_validator("confidence", mode="before")
|
||||
@classmethod
|
||||
def require_number(cls, value: Any) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("must be a JSON number")
|
||||
number = float(value)
|
||||
if not math.isfinite(number):
|
||||
raise ValueError("must be finite")
|
||||
return number
|
||||
|
||||
|
||||
class CorrectionSet(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
corrections: List[CorrectionCandidate] = Field(default_factory=list)
|
||||
|
||||
|
||||
class GrammarValidationDecision(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
correction_index: int = Field(ge=0)
|
||||
is_meaning_preserving: bool
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
reason: StrictStr
|
||||
|
||||
@field_validator("correction_index", mode="before")
|
||||
@classmethod
|
||||
def require_integer_index(cls, value: Any) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("must be an integer")
|
||||
return value
|
||||
|
||||
@field_validator("is_meaning_preserving", mode="before")
|
||||
@classmethod
|
||||
def require_boolean(cls, value: Any) -> bool:
|
||||
if not isinstance(value, bool):
|
||||
raise ValueError("must be a boolean")
|
||||
return value
|
||||
|
||||
@field_validator("confidence", mode="before")
|
||||
@classmethod
|
||||
def require_number(cls, value: Any) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("must be a JSON number")
|
||||
number = float(value)
|
||||
if not math.isfinite(number):
|
||||
raise ValueError("must be finite")
|
||||
return number
|
||||
|
||||
@field_validator("reason")
|
||||
@classmethod
|
||||
def require_non_empty_reason(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("must not be empty")
|
||||
return value
|
||||
|
||||
|
||||
class GrammarValidationSet(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
validations: List[GrammarValidationDecision] = Field(default_factory=list)
|
||||
|
||||
|
||||
class GrammarSpokenFormValidationDecision(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
correction_index: int = Field(ge=0)
|
||||
is_likely_spoken_form_correction: bool
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
reason: StrictStr
|
||||
|
||||
@field_validator("correction_index", mode="before")
|
||||
@classmethod
|
||||
def require_integer_index(cls, value: Any) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError("must be an integer")
|
||||
return value
|
||||
|
||||
@field_validator("is_likely_spoken_form_correction", mode="before")
|
||||
@classmethod
|
||||
def require_boolean(cls, value: Any) -> bool:
|
||||
if not isinstance(value, bool):
|
||||
raise ValueError("must be a boolean")
|
||||
return value
|
||||
|
||||
@field_validator("confidence", mode="before")
|
||||
@classmethod
|
||||
def require_number(cls, value: Any) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("must be a JSON number")
|
||||
number = float(value)
|
||||
if not math.isfinite(number):
|
||||
raise ValueError("must be finite")
|
||||
return number
|
||||
|
||||
@field_validator("reason")
|
||||
@classmethod
|
||||
def require_non_empty_reason(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("must not be empty")
|
||||
return value
|
||||
|
||||
|
||||
class GrammarSpokenFormValidationSet(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
validations: List[GrammarSpokenFormValidationDecision] = Field(default_factory=list)
|
||||
|
||||
|
||||
_TRANSCRIPT_ADAPTER = TypeAdapter(List[TranscriptSegment])
|
||||
_SOURCE_TRANSCRIPT_ADAPTER = TypeAdapter(List[SourceTranscriptSegment])
|
||||
|
||||
|
||||
def validate_transcript_data(
|
||||
data: Any,
|
||||
require_sequential_ids: bool = True,
|
||||
) -> List[TranscriptSegment]:
|
||||
if not isinstance(data, list):
|
||||
raise AuditaValidationError("Transcript must be a JSON array.")
|
||||
if not data:
|
||||
raise AuditaValidationError("Transcript must contain at least one segment.")
|
||||
try:
|
||||
transcript = _TRANSCRIPT_ADAPTER.validate_python(data)
|
||||
except ValidationError as exc:
|
||||
raise AuditaValidationError(f"Transcript schema validation failed: {exc}") from exc
|
||||
if require_sequential_ids:
|
||||
_validate_sequential_ids(transcript)
|
||||
return transcript
|
||||
|
||||
|
||||
def validate_source_transcript_data(data: Any) -> List[SourceTranscriptSegment]:
|
||||
if not isinstance(data, list):
|
||||
raise AuditaValidationError("Transcript must be a JSON array.")
|
||||
if not data:
|
||||
raise AuditaValidationError("Transcript must contain at least one segment.")
|
||||
try:
|
||||
return _SOURCE_TRANSCRIPT_ADAPTER.validate_python(data)
|
||||
except ValidationError as exc:
|
||||
raise AuditaValidationError(f"Transcript schema validation failed: {exc}") from exc
|
||||
|
||||
|
||||
def _validate_sequential_ids(transcript: List[TranscriptSegment]) -> None:
|
||||
ids = [segment.id for segment in transcript]
|
||||
expected = list(range(1, len(transcript) + 1))
|
||||
if ids != expected:
|
||||
raise AuditaValidationError("Transcript segment ids must be sequential starting at 1.")
|
||||
|
||||
|
||||
def parse_transcript_json(
|
||||
raw: str,
|
||||
require_sequential_ids: bool = True,
|
||||
) -> List[TranscriptSegment]:
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AuditaValidationError(f"Transcript is not valid JSON: {exc}") from exc
|
||||
return validate_transcript_data(data, require_sequential_ids=require_sequential_ids)
|
||||
|
||||
|
||||
def parse_source_transcript_json(raw: str) -> List[SourceTranscriptSegment]:
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise AuditaValidationError(f"Transcript is not valid JSON: {exc}") from exc
|
||||
return validate_source_transcript_data(data)
|
||||
|
||||
|
||||
def parse_glossary_yaml(raw: str) -> Glossary:
|
||||
try:
|
||||
import yaml
|
||||
except ImportError as exc:
|
||||
raise AuditaValidationError("PyYAML is required to read glossary files.") from exc
|
||||
|
||||
try:
|
||||
data = yaml.safe_load(raw)
|
||||
except yaml.YAMLError as exc:
|
||||
raise AuditaValidationError(f"Glossary is not valid YAML: {exc}") from exc
|
||||
|
||||
if data is None:
|
||||
data = {}
|
||||
try:
|
||||
return Glossary.model_validate(data)
|
||||
except ValidationError as exc:
|
||||
raise AuditaValidationError(f"Glossary schema validation failed: {exc}") from exc
|
||||
|
||||
|
||||
def transcript_to_json(segments: List[TranscriptSegment]) -> str:
|
||||
payload = [segment.model_dump(mode="json") for segment in segments]
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2) + "\n"
|
||||
|
||||
|
||||
def source_transcript_to_json(segments: List[SourceTranscriptSegment]) -> str:
|
||||
payload = [segment.model_dump(mode="json", exclude_none=True) for segment in segments]
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2) + "\n"
|
||||
@@ -1,223 +0,0 @@
|
||||
import string
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
from .corrections import ReplacementMode, SkippedCorrection
|
||||
from .errors import AuditaLLMError
|
||||
from .protection import ProtectedVocabulary
|
||||
from .schemas import (
|
||||
CorrectionCandidate,
|
||||
GrammarSpokenFormValidationSet,
|
||||
GrammarValidationSet,
|
||||
TranscriptSegment,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GrammarValidationCandidate:
|
||||
correction_index: int
|
||||
correction: CorrectionCandidate
|
||||
original_segment_text: str
|
||||
corrected_segment_text: str
|
||||
|
||||
def to_prompt_payload(self) -> dict:
|
||||
return {
|
||||
"correction_index": self.correction_index,
|
||||
"id": self.correction.id,
|
||||
"original_segment_text": self.original_segment_text,
|
||||
"corrected_segment_text": self.corrected_segment_text,
|
||||
"original_text": self.correction.original_text,
|
||||
"corrected_text": self.correction.corrected_text,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GrammarMeaningValidationResult:
|
||||
approved_correction_indexes: List[int]
|
||||
rescue_candidates: List[GrammarValidationCandidate]
|
||||
approved_count: int
|
||||
rejected_count: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GrammarSpokenFormValidationResult:
|
||||
approved_correction_indexes: List[int]
|
||||
skipped: List[SkippedCorrection]
|
||||
approved_count: int
|
||||
rejected_count: int
|
||||
|
||||
|
||||
def select_grammar_validation_candidates(
|
||||
transcript: List[TranscriptSegment],
|
||||
corrections: List[CorrectionCandidate],
|
||||
confidence_threshold: float,
|
||||
replacement_mode: ReplacementMode,
|
||||
protected_vocabulary: ProtectedVocabulary,
|
||||
) -> Tuple[List[GrammarValidationCandidate], int]:
|
||||
id_to_segment = {segment.id: segment for segment in transcript}
|
||||
candidates: List[GrammarValidationCandidate] = []
|
||||
bypassed_count = 0
|
||||
for index, correction in enumerate(corrections):
|
||||
if correction.confidence < confidence_threshold:
|
||||
bypassed_count += 1
|
||||
continue
|
||||
if _contains_protected_term(protected_vocabulary, correction):
|
||||
bypassed_count += 1
|
||||
continue
|
||||
if is_capitalization_or_punctuation_only(correction.original_text, correction.corrected_text):
|
||||
bypassed_count += 1
|
||||
continue
|
||||
|
||||
segment = id_to_segment.get(correction.id)
|
||||
corrected_segment_text = _preview_corrected_segment_text(segment, correction, replacement_mode)
|
||||
if corrected_segment_text is None:
|
||||
bypassed_count += 1
|
||||
continue
|
||||
|
||||
candidates.append(
|
||||
GrammarValidationCandidate(
|
||||
correction_index=index,
|
||||
correction=correction,
|
||||
original_segment_text=segment.text,
|
||||
corrected_segment_text=corrected_segment_text,
|
||||
)
|
||||
)
|
||||
return candidates, bypassed_count
|
||||
|
||||
|
||||
def filter_with_meaning_preserving_validations(
|
||||
candidates: List[GrammarValidationCandidate],
|
||||
validation_set: GrammarValidationSet,
|
||||
confidence_threshold: float,
|
||||
) -> GrammarMeaningValidationResult:
|
||||
decisions_by_index = _index_validation_decisions(
|
||||
validation_set.validations,
|
||||
candidates,
|
||||
duplicate_error="LLM grammar validation response included duplicate correction_index values.",
|
||||
unknown_error="LLM grammar validation response included an unknown correction_index.",
|
||||
missing_error="LLM grammar validation response omitted correction_index values.",
|
||||
)
|
||||
|
||||
approved_correction_indexes: List[int] = []
|
||||
rescue_candidates: List[GrammarValidationCandidate] = []
|
||||
for candidate in candidates:
|
||||
decision = decisions_by_index[candidate.correction_index]
|
||||
if decision.is_meaning_preserving and decision.confidence >= confidence_threshold:
|
||||
approved_correction_indexes.append(candidate.correction_index)
|
||||
else:
|
||||
rescue_candidates.append(candidate)
|
||||
|
||||
return GrammarMeaningValidationResult(
|
||||
approved_correction_indexes=approved_correction_indexes,
|
||||
rescue_candidates=rescue_candidates,
|
||||
approved_count=len(approved_correction_indexes),
|
||||
rejected_count=len(rescue_candidates),
|
||||
)
|
||||
|
||||
|
||||
def filter_with_spoken_form_validations(
|
||||
candidates: List[GrammarValidationCandidate],
|
||||
validation_set: GrammarSpokenFormValidationSet,
|
||||
confidence_threshold: float,
|
||||
) -> GrammarSpokenFormValidationResult:
|
||||
decisions_by_index = _index_validation_decisions(
|
||||
validation_set.validations,
|
||||
candidates,
|
||||
duplicate_error="LLM spoken-form validation response included duplicate correction_index values.",
|
||||
unknown_error="LLM spoken-form validation response included an unknown correction_index.",
|
||||
missing_error="LLM spoken-form validation response omitted correction_index values.",
|
||||
)
|
||||
|
||||
approved_correction_indexes: List[int] = []
|
||||
skipped: List[SkippedCorrection] = []
|
||||
for candidate in candidates:
|
||||
decision = decisions_by_index[candidate.correction_index]
|
||||
if decision.is_likely_spoken_form_correction and decision.confidence >= confidence_threshold:
|
||||
approved_correction_indexes.append(candidate.correction_index)
|
||||
continue
|
||||
skipped.append(
|
||||
SkippedCorrection(
|
||||
id=candidate.correction.id,
|
||||
reason="grammar validation rejected semantic change",
|
||||
original_text=candidate.correction.original_text,
|
||||
corrected_text=candidate.correction.corrected_text,
|
||||
confidence=candidate.correction.confidence,
|
||||
actual_text=candidate.original_segment_text,
|
||||
validation_confidence=decision.confidence,
|
||||
validation_reason=decision.reason,
|
||||
)
|
||||
)
|
||||
|
||||
return GrammarSpokenFormValidationResult(
|
||||
approved_correction_indexes=approved_correction_indexes,
|
||||
skipped=skipped,
|
||||
approved_count=len(approved_correction_indexes),
|
||||
rejected_count=len(skipped),
|
||||
)
|
||||
|
||||
|
||||
def keep_corrections_with_indexes(
|
||||
corrections: List[CorrectionCandidate],
|
||||
allowed_indexes: Iterable[int],
|
||||
) -> List[CorrectionCandidate]:
|
||||
allowed_index_set = set(allowed_indexes)
|
||||
return [correction for index, correction in enumerate(corrections) if index in allowed_index_set]
|
||||
|
||||
|
||||
def is_capitalization_or_punctuation_only(original_text: str, corrected_text: str) -> bool:
|
||||
return _semantic_key(original_text) == _semantic_key(corrected_text)
|
||||
|
||||
|
||||
def _contains_protected_term(
|
||||
protected_vocabulary: ProtectedVocabulary,
|
||||
correction: CorrectionCandidate,
|
||||
) -> bool:
|
||||
return protected_vocabulary.contains_term(correction.original_text) or protected_vocabulary.contains_term(
|
||||
correction.corrected_text
|
||||
)
|
||||
|
||||
|
||||
def _preview_corrected_segment_text(
|
||||
segment: Optional[TranscriptSegment],
|
||||
correction: CorrectionCandidate,
|
||||
replacement_mode: ReplacementMode,
|
||||
) -> Optional[str]:
|
||||
if segment is None:
|
||||
return None
|
||||
if correction.original_text == "" or correction.original_text == correction.corrected_text:
|
||||
return None
|
||||
match_count = segment.text.count(correction.original_text)
|
||||
if match_count == 0:
|
||||
return None
|
||||
if replacement_mode == "require_unique" and match_count > 1:
|
||||
return None
|
||||
return segment.text.replace(correction.original_text, correction.corrected_text)
|
||||
|
||||
|
||||
def _index_validation_decisions(
|
||||
decisions: Iterable[Any],
|
||||
candidates: List[GrammarValidationCandidate],
|
||||
duplicate_error: str,
|
||||
unknown_error: str,
|
||||
missing_error: str,
|
||||
) -> Dict[int, Any]:
|
||||
candidate_indexes = {candidate.correction_index for candidate in candidates}
|
||||
decisions_by_index: Dict[int, Any] = {}
|
||||
for decision in decisions:
|
||||
if decision.correction_index in decisions_by_index:
|
||||
raise AuditaLLMError(duplicate_error)
|
||||
if decision.correction_index not in candidate_indexes:
|
||||
raise AuditaLLMError(unknown_error)
|
||||
decisions_by_index[decision.correction_index] = decision
|
||||
|
||||
missing_indexes = sorted(candidate_indexes - set(decisions_by_index))
|
||||
if missing_indexes:
|
||||
raise AuditaLLMError(missing_error)
|
||||
return decisions_by_index
|
||||
|
||||
|
||||
_PUNCTUATION = set(string.punctuation) | {"—", "–", "…", "“", "”", "‘", "’"}
|
||||
|
||||
|
||||
def _semantic_key(text: str) -> str:
|
||||
return "".join(character.casefold() for character in text if not character.isspace() and character not in _PUNCTUATION)
|
||||
@@ -1 +0,0 @@
|
||||
"""Archived prototype regression suite."""
|
||||
@@ -1,39 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from audita_prototype.chunking import chunk_transcript
|
||||
from audita_prototype.errors import AuditaValidationError
|
||||
from audita_prototype.schemas import parse_transcript_json
|
||||
|
||||
|
||||
class CountEstimator:
|
||||
def estimate_json(self, value):
|
||||
return len(value) * 10
|
||||
|
||||
|
||||
def _segments(count):
|
||||
payload = [
|
||||
{"id": i + 1, "speaker": "Eric", "start": float(i), "end": float(i + 1), "text": f"Segment {i}"}
|
||||
for i in range(count)
|
||||
]
|
||||
import json
|
||||
|
||||
return parse_transcript_json(json.dumps(payload))
|
||||
|
||||
|
||||
def test_chunk_transcript_splits_on_segment_boundaries():
|
||||
sections = chunk_transcript(_segments(5), max_section_tokens=20, estimator=CountEstimator())
|
||||
|
||||
assert [len(section.segments) for section in sections] == [2, 2, 1]
|
||||
assert [section.start_index for section in sections] == [0, 2, 4]
|
||||
|
||||
|
||||
def test_chunk_transcript_allows_exact_limit():
|
||||
sections = chunk_transcript(_segments(2), max_section_tokens=20, estimator=CountEstimator())
|
||||
|
||||
assert len(sections) == 1
|
||||
assert len(sections[0].segments) == 2
|
||||
|
||||
|
||||
def test_chunk_transcript_rejects_oversized_single_segment():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
chunk_transcript(_segments(1), max_section_tokens=9, estimator=CountEstimator())
|
||||
@@ -1,87 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from audita_prototype.cli import main
|
||||
from audita_prototype.reporting import ProcessResult, RunReport
|
||||
from audita_prototype.schemas import parse_transcript_json
|
||||
|
||||
|
||||
def test_cli_help_uses_audita_program_name(capsys):
|
||||
with pytest.raises(SystemExit) as exc:
|
||||
main(["--help"])
|
||||
|
||||
assert exc.value.code == 0
|
||||
assert capsys.readouterr().out.startswith("usage: audita ")
|
||||
|
||||
|
||||
def test_process_help_includes_glossary_pass_flag(capsys):
|
||||
with pytest.raises(SystemExit) as exc:
|
||||
main(["process", "--help"])
|
||||
|
||||
assert exc.value.code == 0
|
||||
output = capsys.readouterr().out
|
||||
assert "--report-json" in output
|
||||
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 "--grammar-validation-enabled" in output
|
||||
assert "--grammar-validation-confidence-threshold" in output
|
||||
assert "--grammar-spoken-form-validation-confidence-threshold" in output
|
||||
assert "--work-dir-retention" in output
|
||||
assert "--normalize-max-segment-gap" in output
|
||||
assert "--normalize-ellipsis-gap" in output
|
||||
assert "--normalize-max-segment-duration" in output
|
||||
assert "--normalize-max-segment-tokens" in output
|
||||
assert "--confidence-threshold" not in output
|
||||
|
||||
|
||||
def test_cli_process_writes_report_json(monkeypatch, tmp_path):
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Fixed."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
report = RunReport(
|
||||
status="success",
|
||||
config={"model": "m", "base_url": "b"},
|
||||
normalization={"source_segment_count": 1, "normalized_segment_count": 1, "merge_count": 0},
|
||||
stages=[],
|
||||
applied_changes=[],
|
||||
skipped_corrections=[],
|
||||
totals={"output_segment_count": 1, "applied_change_count": 0, "skipped_correction_count": 0},
|
||||
work_dir_retention="auto",
|
||||
work_dir_retained=False,
|
||||
work_dir=None,
|
||||
error=None,
|
||||
)
|
||||
result = ProcessResult(
|
||||
transcript=transcript,
|
||||
report=report,
|
||||
run_dir=tmp_path / "run",
|
||||
work_dir_retained=False,
|
||||
)
|
||||
|
||||
monkeypatch.setattr("audita_prototype.cli.AuditaConfig.from_sources", lambda overrides=None: object())
|
||||
monkeypatch.setattr("audita_prototype.cli.load_transcript", lambda path: [])
|
||||
monkeypatch.setattr("audita_prototype.cli.load_glossary", lambda path: object())
|
||||
monkeypatch.setattr("audita_prototype.cli.process_transcript_result", lambda *args, **kwargs: result)
|
||||
|
||||
output_path = tmp_path / "out.json"
|
||||
report_path = tmp_path / "report.json"
|
||||
exit_code = main(
|
||||
[
|
||||
"process",
|
||||
"transcript.json",
|
||||
"--glossary",
|
||||
"glossary.yaml",
|
||||
"--output",
|
||||
str(output_path),
|
||||
"--report-json",
|
||||
str(report_path),
|
||||
]
|
||||
)
|
||||
|
||||
assert exit_code == 0
|
||||
assert report_path.exists()
|
||||
@@ -1,236 +0,0 @@
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from audita_prototype.config import AuditaConfig, ConfigOverrides
|
||||
from audita_prototype.config import (
|
||||
DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD,
|
||||
DEFAULT_GLOSSARY_MAX_LLM_PASSES,
|
||||
DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD,
|
||||
DEFAULT_GRAMMAR_MAX_LLM_PASSES,
|
||||
DEFAULT_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD,
|
||||
DEFAULT_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD,
|
||||
DEFAULT_GRAMMAR_VALIDATION_ENABLED,
|
||||
DEFAULT_MAX_RETRIES,
|
||||
DEFAULT_MAX_SECTION_TOKENS,
|
||||
DEFAULT_NORMALIZE_ELLIPSIS_GAP,
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION,
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP,
|
||||
DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS,
|
||||
DEFAULT_WORK_DIR,
|
||||
DEFAULT_WORK_DIR_RETENTION,
|
||||
)
|
||||
from audita_prototype.errors import AuditaConfigError
|
||||
|
||||
|
||||
def test_config_uses_defaults_with_api_key():
|
||||
config = AuditaConfig.from_sources(env={"OPENROUTER_API_KEY": "key"})
|
||||
|
||||
assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
|
||||
assert config.glossary_confidence_threshold == 0.8
|
||||
assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_CONFIDENCE_THRESHOLD
|
||||
assert config.grammar_confidence_threshold == 0.8
|
||||
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.grammar_validation_enabled == DEFAULT_GRAMMAR_VALIDATION_ENABLED
|
||||
assert config.grammar_validation_enabled is True
|
||||
assert config.grammar_validation_confidence_threshold == DEFAULT_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD
|
||||
assert config.grammar_validation_confidence_threshold == 0.8
|
||||
assert (
|
||||
config.grammar_spoken_form_validation_confidence_threshold
|
||||
== DEFAULT_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD
|
||||
)
|
||||
assert config.grammar_spoken_form_validation_confidence_threshold == 0.8
|
||||
assert config.normalize_max_segment_gap == DEFAULT_NORMALIZE_MAX_SEGMENT_GAP
|
||||
assert config.normalize_max_segment_gap == 4.0
|
||||
assert config.normalize_ellipsis_gap == DEFAULT_NORMALIZE_ELLIPSIS_GAP
|
||||
assert config.normalize_max_segment_duration == DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION
|
||||
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 == 2048
|
||||
assert config.work_dir == Path(DEFAULT_WORK_DIR)
|
||||
assert config.work_dir_retention == DEFAULT_WORK_DIR_RETENTION
|
||||
assert config.work_dir_retention == "auto"
|
||||
|
||||
|
||||
def test_config_env_overrides_defaults():
|
||||
config = AuditaConfig.from_sources(
|
||||
env={
|
||||
"OPENROUTER_API_KEY": "key",
|
||||
"AUDITA_MAX_SECTION_TOKENS": "42",
|
||||
"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_GRAMMAR_VALIDATION_ENABLED": "false",
|
||||
"AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD": "0.91",
|
||||
"AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD": "0.87",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_GAP": "4.5",
|
||||
"AUDITA_NORMALIZE_ELLIPSIS_GAP": "1.5",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_DURATION": "45.0",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS": "512",
|
||||
"AUDITA_WORK_DIR": "/tmp/custom-audita",
|
||||
"AUDITA_WORK_DIR_RETENTION": "always",
|
||||
}
|
||||
)
|
||||
|
||||
assert config.max_section_tokens == 42
|
||||
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.grammar_validation_enabled is False
|
||||
assert config.grammar_validation_confidence_threshold == 0.91
|
||||
assert config.grammar_spoken_form_validation_confidence_threshold == 0.87
|
||||
assert config.normalize_max_segment_gap == 4.5
|
||||
assert config.normalize_ellipsis_gap == 1.5
|
||||
assert config.normalize_max_segment_duration == 45.0
|
||||
assert config.normalize_max_segment_tokens == 512
|
||||
assert config.work_dir == Path("/tmp/custom-audita")
|
||||
assert config.work_dir_retention == "always"
|
||||
|
||||
|
||||
def test_config_cli_overrides_env():
|
||||
config = AuditaConfig.from_sources(
|
||||
env={
|
||||
"OPENROUTER_API_KEY": "key",
|
||||
"AUDITA_MAX_SECTION_TOKENS": "42",
|
||||
"AUDITA_MAX_RETRIES": "5",
|
||||
"AUDITA_GLOSSARY_MAX_LLM_PASSES": "7",
|
||||
"AUDITA_GRAMMAR_MAX_LLM_PASSES": "6",
|
||||
"AUDITA_GRAMMAR_VALIDATION_ENABLED": "false",
|
||||
"AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD": "0.91",
|
||||
"AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD": "0.87",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_GAP": "4.5",
|
||||
"AUDITA_NORMALIZE_ELLIPSIS_GAP": "1.5",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_DURATION": "45.0",
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS": "512",
|
||||
"AUDITA_WORK_DIR": "/tmp/env-audita",
|
||||
"AUDITA_WORK_DIR_RETENTION": "always",
|
||||
},
|
||||
overrides=ConfigOverrides(
|
||||
max_section_tokens=100,
|
||||
glossary_confidence_threshold=0.7,
|
||||
grammar_confidence_threshold=0.65,
|
||||
max_retries=3,
|
||||
glossary_max_llm_passes=2,
|
||||
grammar_max_llm_passes=3,
|
||||
grammar_validation_enabled=True,
|
||||
grammar_validation_confidence_threshold=0.75,
|
||||
grammar_spoken_form_validation_confidence_threshold=0.72,
|
||||
normalize_max_segment_gap=3.0,
|
||||
normalize_ellipsis_gap=1.0,
|
||||
normalize_max_segment_duration=30.0,
|
||||
normalize_max_segment_tokens=256,
|
||||
work_dir=Path("/tmp/cli-audita"),
|
||||
work_dir_retention="never",
|
||||
),
|
||||
)
|
||||
|
||||
assert config.max_section_tokens == 100
|
||||
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.grammar_validation_enabled is True
|
||||
assert config.grammar_validation_confidence_threshold == 0.75
|
||||
assert config.grammar_spoken_form_validation_confidence_threshold == 0.72
|
||||
assert config.normalize_max_segment_gap == 3.0
|
||||
assert config.normalize_ellipsis_gap == 1.0
|
||||
assert config.normalize_max_segment_duration == 30.0
|
||||
assert config.normalize_max_segment_tokens == 256
|
||||
assert config.work_dir == Path("/tmp/cli-audita")
|
||||
assert config.work_dir_retention == "never"
|
||||
|
||||
|
||||
def test_config_requires_api_key():
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(env={})
|
||||
|
||||
|
||||
def test_config_rejects_bad_env_int():
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(
|
||||
env={"OPENROUTER_API_KEY": "key", "AUDITA_MAX_SECTION_TOKENS": "many"}
|
||||
)
|
||||
|
||||
|
||||
def test_config_rejects_invalid_glossary_pass_count():
|
||||
with pytest.raises(AuditaConfigError):
|
||||
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"}
|
||||
)
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(
|
||||
env={"OPENROUTER_API_KEY": "key", "AUDITA_GRAMMAR_VALIDATION_CONFIDENCE_THRESHOLD": "1.1"}
|
||||
)
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(
|
||||
env={
|
||||
"OPENROUTER_API_KEY": "key",
|
||||
"AUDITA_GRAMMAR_SPOKEN_FORM_VALIDATION_CONFIDENCE_THRESHOLD": "-0.1",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_config_rejects_invalid_grammar_validation_enabled():
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(
|
||||
env={"OPENROUTER_API_KEY": "key", "AUDITA_GRAMMAR_VALIDATION_ENABLED": "maybe"}
|
||||
)
|
||||
|
||||
|
||||
def test_config_rejects_invalid_work_dir_retention():
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(
|
||||
env={"OPENROUTER_API_KEY": "key", "AUDITA_WORK_DIR_RETENTION": "sometimes"}
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def test_config_rejects_invalid_normalization_values():
|
||||
invalid_envs = [
|
||||
{"AUDITA_NORMALIZE_MAX_SEGMENT_GAP": "-1"},
|
||||
{"AUDITA_NORMALIZE_ELLIPSIS_GAP": "-1"},
|
||||
{
|
||||
"AUDITA_NORMALIZE_MAX_SEGMENT_GAP": "1",
|
||||
"AUDITA_NORMALIZE_ELLIPSIS_GAP": "2",
|
||||
},
|
||||
{"AUDITA_NORMALIZE_MAX_SEGMENT_DURATION": "0"},
|
||||
{"AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS": "0"},
|
||||
{"AUDITA_NORMALIZE_MAX_SEGMENT_GAP": "nan"},
|
||||
]
|
||||
for env in invalid_envs:
|
||||
with pytest.raises(AuditaConfigError):
|
||||
AuditaConfig.from_sources(env={"OPENROUTER_API_KEY": "key", **env})
|
||||
@@ -1,294 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from audita_prototype.corrections import apply_corrections
|
||||
from audita_prototype.errors import AuditaValidationError
|
||||
from audita_prototype.protection import ProtectedVocabulary
|
||||
from audita_prototype.schemas import CorrectionCandidate, parse_glossary_yaml, parse_transcript_json
|
||||
|
||||
|
||||
def _transcript():
|
||||
return parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 10.0, "end": 11.0, "text": "I ask Chontia for help."},
|
||||
{"id": 2, "speaker": "Mike", "start": 0.0, "end": 1.0, "text": "Then Lyra."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_apply_corrections_uses_threshold_and_preserves_id_order():
|
||||
transcript = _transcript()
|
||||
corrections = [
|
||||
CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Chontia",
|
||||
corrected_text="Chauntea",
|
||||
confidence=0.8,
|
||||
)
|
||||
]
|
||||
|
||||
result = apply_corrections(transcript, corrections, confidence_threshold=0.8)
|
||||
|
||||
assert [segment.speaker for segment in result.transcript] == ["Eric", "Mike"]
|
||||
assert result.transcript[0].text == "I ask Chauntea for help."
|
||||
assert result.skipped == []
|
||||
assert len(result.applied_corrections) == 1
|
||||
assert result.applied_corrections[0].segment_text_before == "I ask Chontia for help."
|
||||
assert result.applied_corrections[0].segment_text_after == "I ask Chauntea for help."
|
||||
|
||||
|
||||
def test_apply_corrections_ignores_below_threshold():
|
||||
transcript = _transcript()
|
||||
corrections = [
|
||||
CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Chontia",
|
||||
corrected_text="Chauntea",
|
||||
confidence=0.79,
|
||||
)
|
||||
]
|
||||
|
||||
result = apply_corrections(transcript, corrections, confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "I ask Chontia for help."
|
||||
assert result.skipped == []
|
||||
assert result.ignored_ids == [1]
|
||||
assert len(result.ignored) == 1
|
||||
assert result.ignored[0].id == 1
|
||||
assert result.ignored[0].reason == "correction confidence below threshold"
|
||||
|
||||
|
||||
def test_apply_corrections_allows_multiple_distinct_spans_in_one_segment():
|
||||
transcript = _transcript()
|
||||
first = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Chontia",
|
||||
corrected_text="Chauntea",
|
||||
confidence=0.8,
|
||||
)
|
||||
second = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="help",
|
||||
corrected_text="guidance",
|
||||
confidence=0.9,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [first, second], confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "I ask Chauntea for guidance."
|
||||
assert result.skipped == []
|
||||
assert len(result.applied_corrections) == 2
|
||||
|
||||
|
||||
def test_apply_corrections_skips_missing_substring():
|
||||
transcript = _transcript()
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Different text.",
|
||||
corrected_text="Chauntea",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "I ask Chontia for help."
|
||||
assert len(result.skipped) == 1
|
||||
assert result.skipped[0].id == 1
|
||||
assert result.skipped[0].actual_text == "I ask Chontia for help."
|
||||
assert "does not match any substring" in result.skipped[0].reason
|
||||
|
||||
|
||||
def test_apply_corrections_skips_missing_id():
|
||||
transcript = _transcript()
|
||||
correction = CorrectionCandidate(
|
||||
id=99,
|
||||
original_text="Missing.",
|
||||
corrected_text="Still missing.",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["I ask Chontia for help.", "Then Lyra."]
|
||||
assert len(result.skipped) == 1
|
||||
assert result.skipped[0].id == 99
|
||||
assert "does not exist" in result.skipped[0].reason
|
||||
|
||||
|
||||
def test_apply_corrections_skips_no_op():
|
||||
transcript = _transcript()
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Chontia",
|
||||
corrected_text="Chontia",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "I ask Chontia for help."
|
||||
assert len(result.skipped) == 1
|
||||
assert "identical" in result.skipped[0].reason
|
||||
|
||||
|
||||
def test_apply_corrections_replaces_all_repeated_substrings():
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Bane met Bane."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="Bane",
|
||||
corrected_text="Bain",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "Bain met Bain."
|
||||
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_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_guards_only_replacement_span():
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "You have to keep it bind. Svend sees the jesters."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Svend"
|
||||
category: pc
|
||||
summary: "Svend is a player character."
|
||||
- name: "Jesters"
|
||||
category: faction
|
||||
summary: "The Jesters are a faction."
|
||||
"""
|
||||
)
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="keep it bind",
|
||||
corrected_text="keep in mind",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(
|
||||
transcript,
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
correction_guard=ProtectedVocabulary.from_glossary(glossary).violation_reason,
|
||||
)
|
||||
|
||||
assert result.transcript[0].text == "You have to keep in mind. Svend sees the jesters."
|
||||
assert result.skipped == []
|
||||
|
||||
|
||||
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():
|
||||
transcript = _transcript()
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="",
|
||||
corrected_text="Chauntea",
|
||||
confidence=0.8,
|
||||
)
|
||||
|
||||
result = apply_corrections(transcript, [correction], confidence_threshold=0.8)
|
||||
|
||||
assert result.transcript[0].text == "I ask Chontia for help."
|
||||
assert len(result.skipped) == 1
|
||||
assert "empty" in result.skipped[0].reason
|
||||
|
||||
|
||||
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")
|
||||
@@ -1,33 +0,0 @@
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
|
||||
|
||||
def test_prototype_package_is_importable():
|
||||
package_root = ROOT / "src" / "audita_prototype"
|
||||
assert package_root.is_dir()
|
||||
assert (package_root / "__main__.py").is_file()
|
||||
|
||||
|
||||
def test_prototype_module_help_smoke():
|
||||
env = os.environ.copy()
|
||||
env["PYTHONPATH"] = str(ROOT / "src")
|
||||
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-m", "audita_prototype", "--help"],
|
||||
cwd=ROOT,
|
||||
text=True,
|
||||
capture_output=True,
|
||||
env=env,
|
||||
check=False,
|
||||
)
|
||||
|
||||
assert result.returncode == 0
|
||||
assert result.stdout.startswith("usage: audita ")
|
||||
@@ -1,153 +0,0 @@
|
||||
from audita_prototype.normalization import normalize_transcript
|
||||
from audita_prototype.schemas import parse_source_transcript_json
|
||||
|
||||
|
||||
class WordEstimator:
|
||||
def estimate_json(self, value):
|
||||
return len(value[0]["original_text"].split())
|
||||
|
||||
|
||||
def _normalize(raw, **overrides):
|
||||
defaults = {
|
||||
"max_segment_gap": 5.0,
|
||||
"ellipsis_gap": 2.0,
|
||||
"max_segment_duration": 60.0,
|
||||
"max_segment_tokens": 2048,
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return normalize_transcript(parse_source_transcript_json(raw), **defaults)
|
||||
|
||||
|
||||
def test_same_speaker_short_gap_merges_with_space():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hello"},
|
||||
{"speaker": "Eric", "start": 2.0, "end": 3.0, "text": "there"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert len(result.transcript) == 1
|
||||
assert result.transcript[0].id == 1
|
||||
assert result.transcript[0].text == "Hello there"
|
||||
assert result.transcript[0].start == 0.0
|
||||
assert result.transcript[0].end == 3.0
|
||||
assert result.summary.merge_count == 1
|
||||
|
||||
|
||||
def test_same_speaker_larger_allowed_gap_merges_with_ellipsis():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hello"},
|
||||
{"speaker": "Eric", "start": 4.0, "end": 5.0, "text": "there"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert result.transcript[0].text == "Hello ... there"
|
||||
|
||||
|
||||
def test_different_speakers_do_not_merge():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hello"},
|
||||
{"speaker": "Mike", "start": 1.5, "end": 2.0, "text": "there"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["Hello", "there"]
|
||||
assert result.summary.merge_count == 0
|
||||
|
||||
|
||||
def test_gap_above_max_does_not_merge():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hello"},
|
||||
{"speaker": "Eric", "start": 7.0, "end": 8.0, "text": "there"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["Hello", "there"]
|
||||
|
||||
|
||||
def test_overlapping_segments_do_not_merge():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 2.0, "text": "Hello"},
|
||||
{"speaker": "Eric", "start": 1.5, "end": 3.0, "text": "there"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["Hello", "there"]
|
||||
|
||||
|
||||
def test_max_duration_prevents_merge():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 40.0, "text": "Hello"},
|
||||
{"speaker": "Eric", "start": 45.0, "end": 50.0, "text": "there"}
|
||||
]
|
||||
""",
|
||||
max_segment_duration=45.0,
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["Hello", "there"]
|
||||
|
||||
|
||||
def test_max_token_limit_prevents_merge():
|
||||
segments = parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "one two"},
|
||||
{"speaker": "Eric", "start": 2.0, "end": 3.0, "text": "three four"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
result = normalize_transcript(
|
||||
segments,
|
||||
max_segment_gap=5.0,
|
||||
ellipsis_gap=2.0,
|
||||
max_segment_duration=60.0,
|
||||
max_segment_tokens=3,
|
||||
estimator=WordEstimator(),
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["one two", "three four"]
|
||||
|
||||
|
||||
def test_shortest_gap_merges_first():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "A"},
|
||||
{"speaker": "Eric", "start": 3.0, "end": 4.0, "text": "B"},
|
||||
{"speaker": "Eric", "start": 4.5, "end": 5.0, "text": "C"}
|
||||
]
|
||||
""",
|
||||
max_segment_duration=4.0,
|
||||
)
|
||||
|
||||
assert [segment.text for segment in result.transcript] == ["A", "B C"]
|
||||
|
||||
|
||||
def test_fresh_ids_are_assigned_chronologically_and_source_ids_are_discarded():
|
||||
result = _normalize(
|
||||
"""
|
||||
[
|
||||
{"id": 99, "speaker": "Mike", "start": 10.0, "end": 11.0, "text": "Second"},
|
||||
{"id": 42, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "First"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert [(segment.id, segment.text) for segment in result.transcript] == [(1, "First"), (2, "Second")]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,222 +0,0 @@
|
||||
import json
|
||||
|
||||
from audita_prototype.chunking import chunk_transcript
|
||||
from audita_prototype.prompts import (
|
||||
build_glossary_correction_messages,
|
||||
build_grammar_correction_messages,
|
||||
build_grammar_spoken_form_validation_messages,
|
||||
build_grammar_validation_messages,
|
||||
)
|
||||
from audita_prototype.schemas import parse_glossary_yaml, parse_transcript_json
|
||||
|
||||
|
||||
def test_prompt_requires_acoustically_plausible_transcription_errors():
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "The gestures are nearby."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Jesters"
|
||||
category: faction
|
||||
summary: "The Jesters are a local faction."
|
||||
- name: "Lyra"
|
||||
category: npc
|
||||
summary: "Lyra is a hostile NPC."
|
||||
"""
|
||||
)
|
||||
section = chunk_transcript(transcript, max_section_tokens=16000)[0]
|
||||
|
||||
messages = build_glossary_correction_messages(section, glossary)
|
||||
prompt_text = "\n".join(message["content"] for message in messages)
|
||||
|
||||
assert "acoustically plausible" in prompt_text
|
||||
assert "phonetically or acoustically similar" in prompt_text
|
||||
assert '"gestures" to "Jesters"' in prompt_text
|
||||
assert '"Lyra" to "Jesters"' in prompt_text
|
||||
assert "should be omitted" in prompt_text
|
||||
assert "glossary names and aliases already present in the transcript as protected spellings" in prompt_text
|
||||
assert "Do not replace, Anglicize, normalize, lowercase" in prompt_text
|
||||
assert "Preserve canonical glossary capitalization" in prompt_text
|
||||
assert "exact text span that needs replacement" in prompt_text
|
||||
assert "replacement text for that span" in prompt_text
|
||||
assert "Do not return corrections where original_text and corrected_text are identical" in prompt_text
|
||||
|
||||
|
||||
def test_prompt_uses_simplified_segment_payload():
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "The gestures are nearby."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Jesters"
|
||||
category: faction
|
||||
summary: "The Jesters are a local faction."
|
||||
"""
|
||||
)
|
||||
section = chunk_transcript(transcript, max_section_tokens=16000)[0]
|
||||
|
||||
messages = build_glossary_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": "The gestures are nearby."}]
|
||||
assert "speaker" not in prompt_segments[0]
|
||||
assert "start" not in prompt_segments[0]
|
||||
assert "end" not in prompt_segments[0]
|
||||
|
||||
|
||||
def test_prompts_do_not_include_inferred_plurals():
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "The gestures are nearby."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Godfrey"
|
||||
aliases:
|
||||
- "Jester"
|
||||
category: npc
|
||||
summary: "Godfrey is an NPC."
|
||||
"""
|
||||
)
|
||||
section = chunk_transcript(transcript, max_section_tokens=16000)[0]
|
||||
|
||||
glossary_messages = build_glossary_correction_messages(section, glossary)
|
||||
glossary_json = glossary_messages[1]["content"].split("Glossary:\n", maxsplit=1)[1].split(
|
||||
"\n\nTranscript section:",
|
||||
maxsplit=1,
|
||||
)[0]
|
||||
grammar_messages = build_grammar_correction_messages(section, glossary)
|
||||
grammar_json = grammar_messages[1]["content"].split("Protected glossary/context:\n", maxsplit=1)[1].split(
|
||||
"\n\nTranscript section:",
|
||||
maxsplit=1,
|
||||
)[0]
|
||||
|
||||
for prompt_glossary in (json.loads(glossary_json), json.loads(grammar_json)):
|
||||
entry = prompt_glossary["glossary"][0]
|
||||
assert "plural" not in entry
|
||||
assert "Godfreys" not in json.dumps(prompt_glossary)
|
||||
assert "Jesters" not in json.dumps(prompt_glossary)
|
||||
|
||||
|
||||
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 "glossary names and aliases as protected spellings" in prompt_text
|
||||
assert "correct clear transcription or spelling errors toward glossary names or aliases" in prompt_text
|
||||
assert "Do not autocorrect, Anglicize, replace, normalize, lowercase" in prompt_text
|
||||
assert "that already appear correctly in the transcript" in prompt_text
|
||||
assert "Preserve canonical glossary capitalization" 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]
|
||||
|
||||
|
||||
def test_grammar_validation_prompt_rejects_semantic_changes():
|
||||
messages = build_grammar_validation_messages(
|
||||
[
|
||||
{
|
||||
"correction_index": 0,
|
||||
"id": 1,
|
||||
"original_segment_text": "He became visible.",
|
||||
"corrected_segment_text": "He became invisible.",
|
||||
"original_text": "visible",
|
||||
"corrected_text": "invisible",
|
||||
}
|
||||
]
|
||||
)
|
||||
prompt_text = "\n".join(message["content"] for message in messages)
|
||||
|
||||
assert "preserves meaning" in prompt_text
|
||||
assert "became visible" in prompt_text
|
||||
assert "became invisible" in prompt_text
|
||||
assert "reverses the meaning" in prompt_text
|
||||
assert "do not try to rescue likely homophone or transcription fixes" in prompt_text
|
||||
assert "handled in a separate spoken-form validation step" in prompt_text
|
||||
assert "correction_index" in prompt_text
|
||||
assert "is_meaning_preserving" in prompt_text
|
||||
|
||||
|
||||
def test_grammar_spoken_form_validation_prompt_allows_homophone_rescue():
|
||||
messages = build_grammar_spoken_form_validation_messages(
|
||||
[
|
||||
{
|
||||
"correction_index": 0,
|
||||
"id": 1,
|
||||
"original_segment_text": "ChatGPT still can't really do that with a dam.",
|
||||
"corrected_segment_text": "ChatGPT still can't really do that with a damn.",
|
||||
"original_text": "dam",
|
||||
"corrected_text": "damn",
|
||||
}
|
||||
]
|
||||
)
|
||||
prompt_text = "\n".join(message["content"] for message in messages)
|
||||
|
||||
assert "likely homophone, spoken-form, or transcription fix" in prompt_text
|
||||
assert '"dam" to "damn"' in prompt_text
|
||||
assert '"became visible" to "became invisible"' in prompt_text
|
||||
assert "is_likely_spoken_form_correction" in prompt_text
|
||||
@@ -1,191 +0,0 @@
|
||||
from audita_prototype.protection import ProtectedVocabulary
|
||||
from audita_prototype.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."
|
||||
- name: "Popov"
|
||||
category: npc
|
||||
summary: "Popov is an allied NPC."
|
||||
- name: "Jesters"
|
||||
aliases:
|
||||
- "Jester"
|
||||
category: faction
|
||||
summary: "The Jesters are a faction."
|
||||
- name: "Svend"
|
||||
category: pc
|
||||
summary: "Svend is a player character."
|
||||
- name: "Godfrey"
|
||||
category: npc
|
||||
summary: "Godfrey is an NPC."
|
||||
- name: "Lyra"
|
||||
category: npc
|
||||
summary: "Lyra is an NPC."
|
||||
- name: "Loviator"
|
||||
category: deity
|
||||
summary: "Loviator is a deity."
|
||||
"""
|
||||
)
|
||||
return ProtectedVocabulary.from_glossary(glossary)
|
||||
|
||||
|
||||
def test_protection_blocks_replacing_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("Hrank moves.", "Frank moves.")
|
||||
== "correction changes protected glossary term usage"
|
||||
)
|
||||
|
||||
|
||||
def test_protection_blocks_replacing_possessive_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Popov's exhausted.", "Pawpaw's exhausted.") is not None
|
||||
|
||||
|
||||
def test_protection_blocks_lowercasing_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("Hrank moves.", "hrank moves.")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
|
||||
|
||||
def test_protection_blocks_noncanonical_uppercase_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Popov moves.", "POPOV 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_unchanged_noncanonical_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("jesters advance.", "jesters advance.") is None
|
||||
|
||||
|
||||
def test_protection_allows_correction_toward_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Pawpaw moves.", "Popov moves.") is None
|
||||
assert vocabulary.violation_reason("gestures", "Jesters") is None
|
||||
assert vocabulary.violation_reason("rank", "Hrank") is None
|
||||
assert vocabulary.violation_reason("spend", "Svend") is None
|
||||
|
||||
|
||||
def test_protection_allows_possessive_correction_toward_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Pawpaw's exhausted.", "Popov's exhausted.") is None
|
||||
|
||||
|
||||
def test_protection_blocks_noncanonical_introduced_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("gestures", "jesters")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
assert (
|
||||
vocabulary.violation_reason("rank", "hrank")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
assert (
|
||||
vocabulary.violation_reason("spend", "svend")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
|
||||
|
||||
def test_protection_blocks_changed_noncanonical_variant():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("jesters advance.", "JESTERS advance.")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
|
||||
|
||||
def test_protection_allows_inferred_name_plural():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Godfrey's", "Godfreys") is None
|
||||
|
||||
|
||||
def test_protection_allows_inferred_alias_plural():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("gesture", "Jesters") is None
|
||||
|
||||
|
||||
def test_protection_allows_explicit_plural():
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Mox"
|
||||
plural: "Moxen"
|
||||
category: faction
|
||||
summary: "The Mox are a faction."
|
||||
"""
|
||||
)
|
||||
vocabulary = ProtectedVocabulary.from_glossary(glossary)
|
||||
|
||||
assert vocabulary.violation_reason("Mox's", "Moxen") is None
|
||||
assert (
|
||||
vocabulary.violation_reason("Mox's", "moxen")
|
||||
== "correction changes protected glossary term capitalization"
|
||||
)
|
||||
|
||||
|
||||
def test_protection_allows_punctuation_around_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert vocabulary.violation_reason("Popov, moves.", "Popov. Moves.") is None
|
||||
|
||||
|
||||
def test_protection_allows_quote_wrapping_sentence_with_unchanged_lowercase_protected_term():
|
||||
vocabulary = _vocabulary()
|
||||
before = (
|
||||
"When you say that, Popov will say, when I was in that room with the jesters, "
|
||||
"I just knew that Godfrey and Lyra came directly from Loviator herself."
|
||||
)
|
||||
after = (
|
||||
'When you say that, Popov will say, "When I was in that room with the jesters, '
|
||||
'I just knew that Godfrey and Lyra came directly from Loviator herself."'
|
||||
)
|
||||
|
||||
assert vocabulary.violation_reason(before, after) 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
|
||||
|
||||
|
||||
def test_protection_blocks_removing_preexisting_protected_occurrence():
|
||||
vocabulary = _vocabulary()
|
||||
|
||||
assert (
|
||||
vocabulary.violation_reason("Jesters flank the Jesters.", "Jesters flank the gestures.")
|
||||
== "correction changes protected glossary term usage"
|
||||
)
|
||||
@@ -1,336 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from audita_prototype.errors import AuditaLLMError
|
||||
from audita_prototype.protection import ProtectedVocabulary
|
||||
from audita_prototype.schemas import (
|
||||
CorrectionCandidate,
|
||||
GrammarSpokenFormValidationDecision,
|
||||
GrammarSpokenFormValidationSet,
|
||||
GrammarValidationDecision,
|
||||
GrammarValidationSet,
|
||||
parse_glossary_yaml,
|
||||
parse_transcript_json,
|
||||
)
|
||||
from audita_prototype.semantic_validation import (
|
||||
filter_with_meaning_preserving_validations,
|
||||
filter_with_spoken_form_validations,
|
||||
keep_corrections_with_indexes,
|
||||
select_grammar_validation_candidates,
|
||||
)
|
||||
|
||||
|
||||
def _transcript():
|
||||
return parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "he became visible and then gestures arrived"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _vocabulary():
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Jesters"
|
||||
category: faction
|
||||
summary: "The Jesters are a faction."
|
||||
"""
|
||||
)
|
||||
return ProtectedVocabulary.from_glossary(glossary)
|
||||
|
||||
|
||||
def test_protected_vocabulary_correction_bypasses_validation():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="gestures",
|
||||
corrected_text="Jesters",
|
||||
confidence=0.95,
|
||||
)
|
||||
|
||||
candidates, bypassed_count = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert bypassed_count == 1
|
||||
|
||||
|
||||
def test_capitalization_only_correction_bypasses_validation():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="he",
|
||||
corrected_text="He",
|
||||
confidence=0.95,
|
||||
)
|
||||
|
||||
candidates, bypassed_count = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert bypassed_count == 1
|
||||
|
||||
|
||||
def test_punctuation_only_correction_bypasses_validation():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="visible",
|
||||
corrected_text="visible.",
|
||||
confidence=0.95,
|
||||
)
|
||||
|
||||
candidates, bypassed_count = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert bypassed_count == 1
|
||||
|
||||
|
||||
def test_meaning_sensitive_substitution_requires_validation():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="visible",
|
||||
corrected_text="invisible",
|
||||
confidence=0.95,
|
||||
)
|
||||
|
||||
candidates, bypassed_count = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
|
||||
assert len(candidates) == 1
|
||||
assert candidates[0].correction_index == 0
|
||||
assert candidates[0].original_segment_text == "he became visible and then gestures arrived"
|
||||
assert candidates[0].corrected_segment_text == "he became invisible and then gestures arrived"
|
||||
assert bypassed_count == 0
|
||||
|
||||
|
||||
def test_meaning_preserving_validation_approves_without_rescue():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="bind",
|
||||
corrected_text="mind",
|
||||
confidence=0.95,
|
||||
)
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Keep in bind."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
candidates, _ = select_grammar_validation_candidates(
|
||||
transcript,
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
|
||||
result = filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
GrammarValidationSet(
|
||||
validations=[
|
||||
GrammarValidationDecision(
|
||||
correction_index=0,
|
||||
is_meaning_preserving=True,
|
||||
confidence=0.95,
|
||||
reason="This preserves the intended meaning.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
|
||||
assert result.approved_correction_indexes == [0]
|
||||
assert result.rescue_candidates == []
|
||||
assert result.approved_count == 1
|
||||
assert result.rejected_count == 0
|
||||
|
||||
|
||||
def test_spoken_form_validation_can_rescue_homophone_fix():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="dam",
|
||||
corrected_text="damn",
|
||||
confidence=0.95,
|
||||
)
|
||||
transcript = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "ChatGPT still can't really do that with a dam."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
candidates, _ = select_grammar_validation_candidates(
|
||||
transcript,
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
meaning_result = filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
GrammarValidationSet(
|
||||
validations=[
|
||||
GrammarValidationDecision(
|
||||
correction_index=0,
|
||||
is_meaning_preserving=False,
|
||||
confidence=0.99,
|
||||
reason="Written meaning changes from a barrier to a curse word.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
spoken_form_result = filter_with_spoken_form_validations(
|
||||
meaning_result.rescue_candidates,
|
||||
GrammarSpokenFormValidationSet(
|
||||
validations=[
|
||||
GrammarSpokenFormValidationDecision(
|
||||
correction_index=0,
|
||||
is_likely_spoken_form_correction=True,
|
||||
confidence=0.95,
|
||||
reason="The surrounding phrase strongly supports the intended spoken phrase with a curse word.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
|
||||
kept = keep_corrections_with_indexes([correction], spoken_form_result.approved_correction_indexes)
|
||||
assert meaning_result.approved_correction_indexes == []
|
||||
assert kept == [correction]
|
||||
assert spoken_form_result.skipped == []
|
||||
|
||||
|
||||
def test_spoken_form_validation_rejects_non_homophone_semantic_change():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="visible",
|
||||
corrected_text="invisible",
|
||||
confidence=0.95,
|
||||
)
|
||||
candidates, _ = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
meaning_result = filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
GrammarValidationSet(
|
||||
validations=[
|
||||
GrammarValidationDecision(
|
||||
correction_index=0,
|
||||
is_meaning_preserving=False,
|
||||
confidence=0.99,
|
||||
reason="This reverses visible to invisible.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
spoken_form_result = filter_with_spoken_form_validations(
|
||||
meaning_result.rescue_candidates,
|
||||
GrammarSpokenFormValidationSet(
|
||||
validations=[
|
||||
GrammarSpokenFormValidationDecision(
|
||||
correction_index=0,
|
||||
is_likely_spoken_form_correction=False,
|
||||
confidence=0.99,
|
||||
reason="This is a semantic reversal, not a likely spoken-form transcription error.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
|
||||
assert spoken_form_result.approved_correction_indexes == []
|
||||
assert spoken_form_result.rejected_count == 1
|
||||
assert spoken_form_result.skipped[0].reason == "grammar validation rejected semantic change"
|
||||
assert spoken_form_result.skipped[0].validation_reason == (
|
||||
"This is a semantic reversal, not a likely spoken-form transcription error."
|
||||
)
|
||||
|
||||
|
||||
def test_validation_rejects_duplicate_unknown_and_missing_decisions():
|
||||
correction = CorrectionCandidate(
|
||||
id=1,
|
||||
original_text="visible",
|
||||
corrected_text="invisible",
|
||||
confidence=0.95,
|
||||
)
|
||||
candidates, _ = select_grammar_validation_candidates(
|
||||
_transcript(),
|
||||
[correction],
|
||||
confidence_threshold=0.8,
|
||||
replacement_mode="require_unique",
|
||||
protected_vocabulary=_vocabulary(),
|
||||
)
|
||||
meaning_decision = GrammarValidationDecision(
|
||||
correction_index=0,
|
||||
is_meaning_preserving=True,
|
||||
confidence=0.95,
|
||||
reason="Preserves meaning.",
|
||||
)
|
||||
spoken_form_decision = GrammarSpokenFormValidationDecision(
|
||||
correction_index=0,
|
||||
is_likely_spoken_form_correction=True,
|
||||
confidence=0.95,
|
||||
reason="Likely spoken-form correction.",
|
||||
)
|
||||
|
||||
with pytest.raises(AuditaLLMError):
|
||||
filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
GrammarValidationSet(validations=[meaning_decision, meaning_decision]),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
with pytest.raises(AuditaLLMError):
|
||||
filter_with_meaning_preserving_validations(
|
||||
candidates,
|
||||
GrammarValidationSet(validations=[]),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
with pytest.raises(AuditaLLMError):
|
||||
filter_with_spoken_form_validations(
|
||||
candidates,
|
||||
GrammarSpokenFormValidationSet(
|
||||
validations=[
|
||||
GrammarSpokenFormValidationDecision(
|
||||
correction_index=99,
|
||||
is_likely_spoken_form_correction=True,
|
||||
confidence=0.95,
|
||||
reason="Unknown.",
|
||||
)
|
||||
]
|
||||
),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
with pytest.raises(AuditaLLMError):
|
||||
filter_with_spoken_form_validations(
|
||||
candidates,
|
||||
GrammarSpokenFormValidationSet(validations=[spoken_form_decision, spoken_form_decision]),
|
||||
confidence_threshold=0.8,
|
||||
)
|
||||
@@ -1,225 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from audita_prototype.errors import AuditaValidationError
|
||||
from audita_prototype.schemas import parse_glossary_yaml, parse_source_transcript_json, parse_transcript_json
|
||||
|
||||
|
||||
def test_valid_transcript_parses():
|
||||
segments = parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.25, "text": "Then Lyra."}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert len(segments) == 1
|
||||
assert segments[0].id == 1
|
||||
assert segments[0].speaker == "Eric"
|
||||
|
||||
|
||||
def test_transcript_rejects_extra_fields():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.25, "text": "Hi", "extra": true}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_bad_timestamps():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 2.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_missing_id():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_duplicate_ids():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"},
|
||||
{"id": 1, "speaker": "Mike", "start": 1.0, "end": 2.0, "text": "There"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_nonsequential_ids():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"},
|
||||
{"id": 3, "speaker": "Mike", "start": 1.0, "end": 2.0, "text": "There"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_zero_or_negative_id():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 0, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_noninteger_id():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 1.5, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_transcript_rejects_empty_input():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_transcript_json("[]")
|
||||
|
||||
|
||||
def test_source_transcript_accepts_missing_ids():
|
||||
segments = parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert segments[0].id is None
|
||||
assert segments[0].speaker == "Eric"
|
||||
|
||||
|
||||
def test_source_transcript_accepts_present_nonsequential_ids():
|
||||
segments = parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"id": 10, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi"},
|
||||
{"id": 4, "speaker": "Mike", "start": 1.0, "end": 2.0, "text": "There"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
assert [segment.id for segment in segments] == [10, 4]
|
||||
|
||||
|
||||
def test_source_transcript_rejects_extra_fields():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "Hi", "extra": true}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_source_transcript_rejects_bad_timestamps():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "Eric", "start": 2.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_source_transcript_rejects_empty_values():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_source_transcript_json(
|
||||
"""
|
||||
[
|
||||
{"speaker": "", "start": 0.0, "end": 1.0, "text": "Hi"}
|
||||
]
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_source_transcript_rejects_invalid_json():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_source_transcript_json("{")
|
||||
|
||||
|
||||
def test_valid_glossary_parses():
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Lyra"
|
||||
category: npc
|
||||
summary: "Lyra is a hostile NPC."
|
||||
"""
|
||||
)
|
||||
|
||||
assert glossary.glossary[0].name == "Lyra"
|
||||
assert glossary.glossary[0].plural is None
|
||||
|
||||
|
||||
def test_glossary_accepts_optional_plural():
|
||||
glossary = parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Godfrey"
|
||||
plural: "Godfreys"
|
||||
category: npc
|
||||
summary: "Godfrey is an NPC."
|
||||
"""
|
||||
)
|
||||
|
||||
assert glossary.glossary[0].plural == "Godfreys"
|
||||
|
||||
|
||||
def test_glossary_rejects_empty_plural():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Godfrey"
|
||||
plural: ""
|
||||
category: npc
|
||||
summary: "Godfrey is an NPC."
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def test_glossary_rejects_empty_entries():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_glossary_yaml("glossary: []")
|
||||
|
||||
|
||||
def test_glossary_rejects_extra_fields():
|
||||
with pytest.raises(AuditaValidationError):
|
||||
parse_glossary_yaml(
|
||||
"""
|
||||
glossary:
|
||||
- name: "Lyra"
|
||||
category: npc
|
||||
summary: "Lyra is a hostile NPC."
|
||||
extra: "nope"
|
||||
"""
|
||||
)
|
||||
Reference in New Issue
Block a user