Implemented deterministic transcript normalization before the LLM stages

This commit is contained in:
2026-04-22 07:22:40 -05:00
parent 445329de81
commit c1672af9b4
12 changed files with 719 additions and 31 deletions

View File

@@ -1,6 +1,6 @@
# Audita # Audita
Audita takes raw audio transcripts and uses an LLM to identify and fix misheard words, jargon, domain-specific terms, and conservative readability issues. Audita takes raw audio transcripts, deterministically merges short same-speaker segments into speaking turns, and uses an LLM to identify and fix misheard words, jargon, domain-specific terms, and conservative readability issues.
## Development ## Development
@@ -47,6 +47,10 @@ Useful configuration can be supplied by CLI flag or environment variable:
- `AUDITA_MAX_RETRIES`, default `3` - `AUDITA_MAX_RETRIES`, default `3`
- `AUDITA_GLOSSARY_MAX_LLM_PASSES`, default `3`, for total glossary correction passes - `AUDITA_GLOSSARY_MAX_LLM_PASSES`, default `3`, for total glossary correction passes
- `AUDITA_GRAMMAR_MAX_LLM_PASSES`, default `3`, for total grammar/readability correction passes - `AUDITA_GRAMMAR_MAX_LLM_PASSES`, default `3`, for total grammar/readability correction passes
- `AUDITA_NORMALIZE_MAX_SEGMENT_GAP`, default `5.0`, for same-speaker gaps eligible for merging
- `AUDITA_NORMALIZE_ELLIPSIS_GAP`, default `2.0`, for same-speaker gaps that should be joined with ` ... `
- `AUDITA_NORMALIZE_MAX_SEGMENT_DURATION`, default `60.0`, for maximum merged segment duration
- `AUDITA_NORMALIZE_MAX_SEGMENT_TOKENS`, default `2048`, for maximum merged segment prompt payload size
- `AUDITA_WORK_DIR`, default `/tmp/audita` - `AUDITA_WORK_DIR`, default `/tmp/audita`
`AUDITA_WORK_DIR` stores per-run diagnostics while processing. Successful runs clean up their run directory unless corrections are skipped; failed runs and skipped-correction runs preserve diagnostics for debugging. `AUDITA_WORK_DIR` stores per-run diagnostics while processing. Successful runs clean up their run directory unless corrections are skipped; failed runs and skipped-correction runs preserve diagnostics for debugging.

View File

@@ -45,6 +45,26 @@ def _build_parser() -> argparse.ArgumentParser:
process.add_argument("--max-retries", type=int, help="maximum Instructor retries for structured response validation") 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("--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-max-llm-passes", type=int, help="maximum total LLM passes for grammar corrections")
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", type=Path, help="directory for per-run scratch diagnostics")
return parser return parser
@@ -61,6 +81,10 @@ def _process(args: argparse.Namespace) -> int:
max_retries=args.max_retries, max_retries=args.max_retries,
glossary_max_llm_passes=args.glossary_max_llm_passes, glossary_max_llm_passes=args.glossary_max_llm_passes,
grammar_max_llm_passes=args.grammar_max_llm_passes, grammar_max_llm_passes=args.grammar_max_llm_passes,
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=args.work_dir,
) )
) )

View File

@@ -1,3 +1,4 @@
import math
import os import os
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
@@ -15,6 +16,10 @@ DEFAULT_MAX_RETRIES = 3
DEFAULT_WORK_DIR = "/tmp/audita" DEFAULT_WORK_DIR = "/tmp/audita"
DEFAULT_GLOSSARY_MAX_LLM_PASSES = 3 DEFAULT_GLOSSARY_MAX_LLM_PASSES = 3
DEFAULT_GRAMMAR_MAX_LLM_PASSES = 3 DEFAULT_GRAMMAR_MAX_LLM_PASSES = 3
DEFAULT_NORMALIZE_MAX_SEGMENT_GAP = 5.0
DEFAULT_NORMALIZE_ELLIPSIS_GAP = 2.0
DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION = 60.0
DEFAULT_NORMALIZE_MAX_SEGMENT_TOKENS = 2048
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -27,6 +32,10 @@ class ConfigOverrides:
max_retries: Optional[int] = None max_retries: Optional[int] = None
glossary_max_llm_passes: Optional[int] = None glossary_max_llm_passes: Optional[int] = None
grammar_max_llm_passes: Optional[int] = None grammar_max_llm_passes: Optional[int] = 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: Optional[Path] = None
@@ -41,6 +50,10 @@ class AuditaConfig:
max_retries: int = DEFAULT_MAX_RETRIES max_retries: int = DEFAULT_MAX_RETRIES
glossary_max_llm_passes: int = DEFAULT_GLOSSARY_MAX_LLM_PASSES glossary_max_llm_passes: int = DEFAULT_GLOSSARY_MAX_LLM_PASSES
grammar_max_llm_passes: int = DEFAULT_GRAMMAR_MAX_LLM_PASSES grammar_max_llm_passes: int = DEFAULT_GRAMMAR_MAX_LLM_PASSES
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: Path = Path(DEFAULT_WORK_DIR)
@classmethod @classmethod
@@ -91,6 +104,30 @@ class AuditaConfig:
DEFAULT_GRAMMAR_MAX_LLM_PASSES, DEFAULT_GRAMMAR_MAX_LLM_PASSES,
"AUDITA_GRAMMAR_MAX_LLM_PASSES", "AUDITA_GRAMMAR_MAX_LLM_PASSES",
) )
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_value = selected.work_dir or Path(source.get("AUDITA_WORK_DIR") or DEFAULT_WORK_DIR)
config = cls( config = cls(
@@ -103,6 +140,10 @@ class AuditaConfig:
max_retries=max_retries, max_retries=max_retries,
glossary_max_llm_passes=glossary_max_llm_passes, glossary_max_llm_passes=glossary_max_llm_passes,
grammar_max_llm_passes=grammar_max_llm_passes, grammar_max_llm_passes=grammar_max_llm_passes,
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=Path(work_dir_value),
) )
config.validate() config.validate()
@@ -127,6 +168,24 @@ class AuditaConfig:
raise AuditaConfigError("AUDITA_GLOSSARY_MAX_LLM_PASSES must be greater than or equal to one.") raise AuditaConfigError("AUDITA_GLOSSARY_MAX_LLM_PASSES must be greater than or equal to one.")
if self.grammar_max_llm_passes < 1: if self.grammar_max_llm_passes < 1:
raise AuditaConfigError("AUDITA_GRAMMAR_MAX_LLM_PASSES must be greater than or equal to one.") raise AuditaConfigError("AUDITA_GRAMMAR_MAX_LLM_PASSES must be greater than or equal to one.")
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.")
def _get_required_env(env: Mapping[str, str], name: str) -> str: def _get_required_env(env: Mapping[str, str], name: str) -> str:

View File

@@ -1,12 +1,12 @@
from pathlib import Path from pathlib import Path
from typing import List from typing import List
from .schemas import Glossary, TranscriptSegment from .schemas import Glossary, SourceTranscriptSegment, TranscriptSegment
from .schemas import parse_glossary_yaml, parse_transcript_json, transcript_to_json from .schemas import parse_glossary_yaml, parse_source_transcript_json, transcript_to_json
def load_transcript(path: Path) -> List[TranscriptSegment]: def load_transcript(path: Path) -> List[SourceTranscriptSegment]:
return parse_transcript_json(path.read_text(encoding="utf-8")) return parse_source_transcript_json(path.read_text(encoding="utf-8"))
def load_glossary(path: Path) -> Glossary: def load_glossary(path: Path) -> Glossary:
@@ -15,4 +15,3 @@ def load_glossary(path: Path) -> Glossary:
def write_transcript(path: Path, segments: List[TranscriptSegment]) -> None: def write_transcript(path: Path, segments: List[TranscriptSegment]) -> None:
path.write_text(transcript_to_json(segments), encoding="utf-8") path.write_text(transcript_to_json(segments), encoding="utf-8")

176
src/audita/normalization.py Normal file
View File

@@ -0,0 +1,176 @@
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)
]

View File

@@ -10,8 +10,9 @@ from .chunking import IndexedSegment, TranscriptSection, chunk_indexed_segments
from .config import AuditaConfig from .config import AuditaConfig
from .corrections import ReplacementMode, SkippedCorrection, apply_corrections from .corrections import ReplacementMode, SkippedCorrection, apply_corrections
from .errors import AuditaError from .errors import AuditaError
from .normalization import NormalizationResult, normalize_transcript
from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient from .passes import CorrectionPass, GlossaryCorrectionPass, GrammarCorrectionPass, LLMClient
from .schemas import Glossary, TranscriptSegment, parse_transcript_json from .schemas import Glossary, SourceTranscriptSegment, TranscriptSegment, parse_transcript_json
ProgressCallback = Callable[[str], None] ProgressCallback = Callable[[str], None]
@@ -27,7 +28,7 @@ class StageSpec:
def process_transcript( def process_transcript(
transcript: List[TranscriptSegment], transcript: List[SourceTranscriptSegment],
glossary: Glossary, glossary: Glossary,
config: AuditaConfig, config: AuditaConfig,
llm_client: Optional[LLMClient] = None, llm_client: Optional[LLMClient] = None,
@@ -37,14 +38,32 @@ def process_transcript(
try: try:
_log(progress, f"Created work directory {run_dir}") _log(progress, f"Created work directory {run_dir}")
stage_summaries: List[dict] = [] stage_summaries: List[dict] = []
_write_run_metadata(run_dir, config, stage_summaries) normalization_summary: Optional[dict] = None
_write_run_metadata(run_dir, config, normalization_summary, stage_summaries)
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)
if llm_client is None: if llm_client is None:
from .llm import InstructorLLMClient from .llm import InstructorLLMClient
llm_client = InstructorLLMClient(config) llm_client = InstructorLLMClient(config)
working = list(transcript) working = list(normalization_result.transcript)
stages = [ stages = [
StageSpec( StageSpec(
name="glossary", name="glossary",
@@ -74,7 +93,7 @@ def process_transcript(
"passes": [], "passes": [],
} }
stage_summaries.append(stage_summary) stage_summaries.append(stage_summary)
_write_run_metadata(run_dir, config, stage_summaries) _write_run_metadata(run_dir, config, normalization_summary, stage_summaries)
working, stage_skipped = _run_correction_stage( working, stage_skipped = _run_correction_stage(
working, working,
@@ -82,13 +101,14 @@ def process_transcript(
config, config,
stage, stage,
run_dir, run_dir,
normalization_summary,
stage_dir, stage_dir,
stage_summaries, stage_summaries,
stage_summary["passes"], stage_summary["passes"],
progress, progress,
) )
final_skipped.extend((stage.name, skipped) for skipped in stage_skipped) final_skipped.extend((stage.name, skipped) for skipped in stage_skipped)
_write_run_metadata(run_dir, config, stage_summaries) _write_run_metadata(run_dir, config, normalization_summary, stage_summaries)
_write_skipped_corrections(run_dir, final_skipped) _write_skipped_corrections(run_dir, final_skipped)
for stage_name, skipped in final_skipped: for stage_name, skipped in final_skipped:
@@ -117,6 +137,7 @@ def _run_correction_stage(
config: AuditaConfig, config: AuditaConfig,
stage: StageSpec, stage: StageSpec,
run_dir: Path, run_dir: Path,
normalization_summary: Optional[dict],
stage_dir: Path, stage_dir: Path,
stage_summaries: List[dict], stage_summaries: List[dict],
pass_summaries: List[dict], pass_summaries: List[dict],
@@ -191,7 +212,7 @@ def _run_correction_stage(
"retry_segment_count": len(unresolved_retry_skips), "retry_segment_count": len(unresolved_retry_skips),
} }
) )
_write_run_metadata(run_dir, config, stage_summaries) _write_run_metadata(run_dir, config, normalization_summary, stage_summaries)
if not unresolved_retry_skips: if not unresolved_retry_skips:
break break
@@ -213,6 +234,7 @@ def _create_run_dir(work_dir: Path) -> Path:
def _write_run_metadata( def _write_run_metadata(
run_dir: Path, run_dir: Path,
config: AuditaConfig, config: AuditaConfig,
normalization_summary: Optional[dict],
stage_summaries: List[dict], stage_summaries: List[dict],
) -> None: ) -> None:
metadata = { metadata = {
@@ -224,6 +246,7 @@ def _write_run_metadata(
"max_retries": config.max_retries, "max_retries": config.max_retries,
"glossary_max_llm_passes": config.glossary_max_llm_passes, "glossary_max_llm_passes": config.glossary_max_llm_passes,
"grammar_max_llm_passes": config.grammar_max_llm_passes, "grammar_max_llm_passes": config.grammar_max_llm_passes,
"normalization": normalization_summary,
"stages": stage_summaries, "stages": stage_summaries,
} }
(run_dir / "metadata.json").write_text( (run_dir / "metadata.json").write_text(
@@ -232,6 +255,32 @@ def _write_run_metadata(
) )
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: def _write_and_validate_section(run_dir: Path, section: TranscriptSection) -> None:
section_path = run_dir / f"section-{section.section_index:04d}.json" section_path = run_dir / f"section-{section.section_index:04d}.json"
section_json = section.transcript_json() section_json = section.transcript_json()

View File

@@ -1,6 +1,6 @@
import json import json
import math import math
from typing import Any, List from typing import Any, List, Optional
from pydantic import BaseModel, ConfigDict, Field, StrictStr, TypeAdapter from pydantic import BaseModel, ConfigDict, Field, StrictStr, TypeAdapter
from pydantic import ValidationError, field_validator, model_validator from pydantic import ValidationError, field_validator, model_validator
@@ -50,6 +50,50 @@ class TranscriptSegment(BaseModel):
return self 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): class GlossaryEntry(BaseModel):
model_config = ConfigDict(extra="forbid") model_config = ConfigDict(extra="forbid")
@@ -120,6 +164,7 @@ class CorrectionSet(BaseModel):
_TRANSCRIPT_ADAPTER = TypeAdapter(List[TranscriptSegment]) _TRANSCRIPT_ADAPTER = TypeAdapter(List[TranscriptSegment])
_SOURCE_TRANSCRIPT_ADAPTER = TypeAdapter(List[SourceTranscriptSegment])
def validate_transcript_data( def validate_transcript_data(
@@ -139,6 +184,17 @@ def validate_transcript_data(
return 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: def _validate_sequential_ids(transcript: List[TranscriptSegment]) -> None:
ids = [segment.id for segment in transcript] ids = [segment.id for segment in transcript]
expected = list(range(1, len(transcript) + 1)) expected = list(range(1, len(transcript) + 1))
@@ -157,6 +213,14 @@ def parse_transcript_json(
return validate_transcript_data(data, require_sequential_ids=require_sequential_ids) 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: def parse_glossary_yaml(raw: str) -> Glossary:
try: try:
import yaml import yaml
@@ -179,3 +243,8 @@ def parse_glossary_yaml(raw: str) -> Glossary:
def transcript_to_json(segments: List[TranscriptSegment]) -> str: def transcript_to_json(segments: List[TranscriptSegment]) -> str:
payload = [segment.model_dump(mode="json") for segment in segments] payload = [segment.model_dump(mode="json") for segment in segments]
return json.dumps(payload, ensure_ascii=False, indent=2) + "\n" 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"

View File

@@ -21,4 +21,8 @@ def test_process_help_includes_glossary_pass_flag(capsys):
assert "--grammar-max-llm-passes" in output assert "--grammar-max-llm-passes" in output
assert "--glossary-confidence-threshold" in output assert "--glossary-confidence-threshold" in output
assert "--grammar-confidence-threshold" in output assert "--grammar-confidence-threshold" 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 assert "--confidence-threshold" not in output

View File

@@ -10,6 +10,10 @@ from audita.config import (
DEFAULT_GRAMMAR_MAX_LLM_PASSES, DEFAULT_GRAMMAR_MAX_LLM_PASSES,
DEFAULT_MAX_RETRIES, DEFAULT_MAX_RETRIES,
DEFAULT_MAX_SECTION_TOKENS, 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,
) )
from audita.errors import AuditaConfigError from audita.errors import AuditaConfigError
@@ -26,6 +30,14 @@ def test_config_uses_defaults_with_api_key():
assert config.max_retries == DEFAULT_MAX_RETRIES assert config.max_retries == DEFAULT_MAX_RETRIES
assert config.glossary_max_llm_passes == DEFAULT_GLOSSARY_MAX_LLM_PASSES 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_max_llm_passes == DEFAULT_GRAMMAR_MAX_LLM_PASSES
assert config.normalize_max_segment_gap == DEFAULT_NORMALIZE_MAX_SEGMENT_GAP
assert config.normalize_max_segment_gap == 5.0
assert config.normalize_ellipsis_gap == DEFAULT_NORMALIZE_ELLIPSIS_GAP
assert config.normalize_ellipsis_gap == 2.0
assert config.normalize_max_segment_duration == DEFAULT_NORMALIZE_MAX_SEGMENT_DURATION
assert config.normalize_max_segment_duration == 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 == Path(DEFAULT_WORK_DIR)
@@ -39,6 +51,10 @@ def test_config_env_overrides_defaults():
"AUDITA_MAX_RETRIES": "5", "AUDITA_MAX_RETRIES": "5",
"AUDITA_GLOSSARY_MAX_LLM_PASSES": "7", "AUDITA_GLOSSARY_MAX_LLM_PASSES": "7",
"AUDITA_GRAMMAR_MAX_LLM_PASSES": "4", "AUDITA_GRAMMAR_MAX_LLM_PASSES": "4",
"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": "/tmp/custom-audita",
} }
) )
@@ -49,6 +65,10 @@ def test_config_env_overrides_defaults():
assert config.max_retries == 5 assert config.max_retries == 5
assert config.glossary_max_llm_passes == 7 assert config.glossary_max_llm_passes == 7
assert config.grammar_max_llm_passes == 4 assert config.grammar_max_llm_passes == 4
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 == Path("/tmp/custom-audita")
@@ -60,6 +80,10 @@ def test_config_cli_overrides_env():
"AUDITA_MAX_RETRIES": "5", "AUDITA_MAX_RETRIES": "5",
"AUDITA_GLOSSARY_MAX_LLM_PASSES": "7", "AUDITA_GLOSSARY_MAX_LLM_PASSES": "7",
"AUDITA_GRAMMAR_MAX_LLM_PASSES": "6", "AUDITA_GRAMMAR_MAX_LLM_PASSES": "6",
"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": "/tmp/env-audita",
}, },
overrides=ConfigOverrides( overrides=ConfigOverrides(
@@ -69,6 +93,10 @@ def test_config_cli_overrides_env():
max_retries=3, max_retries=3,
glossary_max_llm_passes=2, glossary_max_llm_passes=2,
grammar_max_llm_passes=3, grammar_max_llm_passes=3,
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=Path("/tmp/cli-audita"),
), ),
) )
@@ -79,6 +107,10 @@ def test_config_cli_overrides_env():
assert config.max_retries == 3 assert config.max_retries == 3
assert config.glossary_max_llm_passes == 2 assert config.glossary_max_llm_passes == 2
assert config.grammar_max_llm_passes == 3 assert config.grammar_max_llm_passes == 3
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 == Path("/tmp/cli-audita")
@@ -126,3 +158,20 @@ def test_legacy_confidence_threshold_env_is_ignored():
assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD assert config.glossary_confidence_threshold == DEFAULT_GLOSSARY_CONFIDENCE_THRESHOLD
assert config.grammar_confidence_threshold == DEFAULT_GRAMMAR_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})

153
tests/test_normalization.py Normal file
View File

@@ -0,0 +1,153 @@
from audita.normalization import normalize_transcript
from audita.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")]

View File

@@ -2,7 +2,7 @@ import json
from audita.config import AuditaConfig from audita.config import AuditaConfig
from audita.pipeline import process_transcript from audita.pipeline import process_transcript
from audita.schemas import CorrectionCandidate, CorrectionSet, parse_glossary_yaml, parse_transcript_json from audita.schemas import CorrectionCandidate, CorrectionSet, parse_glossary_yaml, parse_source_transcript_json
class FakeLLMClient: class FakeLLMClient:
@@ -42,11 +42,11 @@ def _glossary():
def _transcript(): def _transcript():
return parse_transcript_json( return parse_source_transcript_json(
""" """
[ [
{"id": 1, "speaker": "Eric", "start": 10.0, "end": 11.0, "text": "I ask Chontia."}, {"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "I ask Chontia."},
{"id": 2, "speaker": "Mike", "start": 0.0, "end": 1.0, "text": "Then Lyra."} {"speaker": "Mike", "start": 10.0, "end": 11.0, "text": "Then Lyra."}
] ]
""" """
) )
@@ -73,12 +73,44 @@ def test_pipeline_processes_with_fake_llm_and_cleans_work_dir(tmp_path):
llm_client=fake_client, llm_client=fake_client,
) )
assert [segment.speaker for segment in revised] == ["Mike", "Eric"] assert [segment.speaker for segment in revised] == ["Eric", "Mike"]
assert revised[1].text == "I ask Chauntea." assert revised[0].text == "I ask Chauntea."
assert fake_client.calls == 2 assert fake_client.calls == 2
assert list((tmp_path / "work").iterdir()) == [] assert list((tmp_path / "work").iterdir()) == []
def test_pipeline_normalizes_before_llm_prompts(tmp_path):
transcript = parse_source_transcript_json(
"""
[
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "I ask"},
{"speaker": "Eric", "start": 2.0, "end": 3.0, "text": "Chontia."},
{"speaker": "Mike", "start": 10.0, "end": 11.0, "text": "Then Lyra."}
]
"""
)
fake_client = FakeLLMClient(
[
CorrectionSet(corrections=[]),
CorrectionSet(corrections=[]),
]
)
progress = []
process_transcript(
transcript,
_glossary(),
_config(tmp_path),
llm_client=fake_client,
progress=progress.append,
)
glossary_prompt = fake_client.messages[0][1]["content"]
glossary_payload = json.loads(glossary_prompt.split("Transcript section:\n", maxsplit=1)[1])
assert glossary_payload[0] == {"id": 1, "original_text": "I ask Chontia."}
assert any("Normalized transcript from 3 to 2 segments" in message for message in progress)
def test_pipeline_skips_bad_glossary_correction_and_preserves_diagnostics(tmp_path): def test_pipeline_skips_bad_glossary_correction_and_preserves_diagnostics(tmp_path):
correction = CorrectionCandidate( correction = CorrectionCandidate(
id=1, id=1,
@@ -102,7 +134,7 @@ def test_pipeline_skips_bad_glossary_correction_and_preserves_diagnostics(tmp_pa
progress=progress.append, progress=progress.append,
) )
assert revised[1].text == "I ask Chontia." assert revised[0].text == "I ask Chontia."
assert any("Skipping glossary correction for id 1" in message for message in progress) assert any("Skipping glossary correction for id 1" in message for message in progress)
preserved = list((tmp_path / "work").iterdir()) preserved = list((tmp_path / "work").iterdir())
assert len(preserved) == 1 assert len(preserved) == 1
@@ -143,8 +175,8 @@ def test_pipeline_retries_skipped_segment_and_cleans_work_dir_when_fixed(tmp_pat
) )
assert fake_client.calls == 3 assert fake_client.calls == 3
assert [segment.speaker for segment in revised] == ["Mike", "Eric"] assert [segment.speaker for segment in revised] == ["Eric", "Mike"]
assert revised[1].text == "I ask Chauntea." assert revised[0].text == "I ask Chauntea."
assert list((tmp_path / "work").iterdir()) == [] assert list((tmp_path / "work").iterdir()) == []
@@ -215,7 +247,13 @@ def test_pipeline_writes_stage_metadata_for_unresolved_retries(tmp_path):
assert len(run_dirs) == 1 assert len(run_dirs) == 1
assert (run_dirs[0] / "glossary" / "pass-0001").exists() assert (run_dirs[0] / "glossary" / "pass-0001").exists()
assert (run_dirs[0] / "grammar" / "pass-0001").exists() assert (run_dirs[0] / "grammar" / "pass-0001").exists()
assert (run_dirs[0] / "normalization" / "source-transcript.json").exists()
assert (run_dirs[0] / "normalization" / "normalized-transcript.json").exists()
assert (run_dirs[0] / "normalization" / "summary.json").exists()
metadata = json.loads((run_dirs[0] / "metadata.json").read_text(encoding="utf-8")) metadata = json.loads((run_dirs[0] / "metadata.json").read_text(encoding="utf-8"))
assert metadata["normalization"]["source_segment_count"] == 2
assert metadata["normalization"]["normalized_segment_count"] == 2
assert metadata["normalization"]["merge_count"] == 0
assert metadata["glossary_max_llm_passes"] == 2 assert metadata["glossary_max_llm_passes"] == 2
assert metadata["grammar_max_llm_passes"] == 3 assert metadata["grammar_max_llm_passes"] == 3
assert metadata["glossary_confidence_threshold"] == 0.8 assert metadata["glossary_confidence_threshold"] == 0.8
@@ -227,11 +265,11 @@ def test_pipeline_writes_stage_metadata_for_unresolved_retries(tmp_path):
def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path): def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path):
transcript = parse_transcript_json( transcript = parse_source_transcript_json(
""" """
[ [
{"id": 1, "speaker": "Eric", "start": 10.0, "end": 11.0, "text": "i ask Chontia."}, {"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "i ask Chontia."},
{"id": 2, "speaker": "Mike", "start": 0.0, "end": 1.0, "text": "Then Lyra."} {"speaker": "Mike", "start": 10.0, "end": 11.0, "text": "Then Lyra."}
] ]
""" """
) )
@@ -264,14 +302,14 @@ def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path):
grammar_prompt = fake_client.messages[1][1]["content"] grammar_prompt = fake_client.messages[1][1]["content"]
grammar_payload = json.loads(grammar_prompt.split("Transcript section:\n", maxsplit=1)[1]) grammar_payload = json.loads(grammar_prompt.split("Transcript section:\n", maxsplit=1)[1])
assert grammar_payload[0]["original_text"] == "i ask Chauntea." assert grammar_payload[0]["original_text"] == "i ask Chauntea."
assert revised[1].text == "I ask Chauntea." assert revised[0].text == "I ask Chauntea."
def test_grammar_stage_retries_repeated_span_and_applies_unique_retry(tmp_path): def test_grammar_stage_retries_repeated_span_and_applies_unique_retry(tmp_path):
transcript = parse_transcript_json( transcript = parse_source_transcript_json(
""" """
[ [
{"id": 1, "speaker": "Eric", "start": 0.0, "end": 1.0, "text": "there and there."} {"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "there and there."}
] ]
""" """
) )
@@ -331,7 +369,7 @@ def test_below_threshold_grammar_corrections_are_not_retried(tmp_path):
) )
assert fake_client.calls == 2 assert fake_client.calls == 2
assert revised[1].text == "I ask Chontia." assert revised[0].text == "I ask Chontia."
def test_unresolved_grammar_skip_preserves_diagnostics(tmp_path): def test_unresolved_grammar_skip_preserves_diagnostics(tmp_path):

View File

@@ -1,7 +1,7 @@
import pytest import pytest
from audita.errors import AuditaValidationError from audita.errors import AuditaValidationError
from audita.schemas import parse_glossary_yaml, parse_transcript_json from audita.schemas import parse_glossary_yaml, parse_source_transcript_json, parse_transcript_json
def test_valid_transcript_parses(): def test_valid_transcript_parses():
@@ -102,6 +102,70 @@ def test_transcript_rejects_empty_input():
parse_transcript_json("[]") 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(): def test_valid_glossary_parses():
glossary = parse_glossary_yaml( glossary = parse_glossary_yaml(
""" """