98 lines
3.6 KiB
Python
98 lines
3.6 KiB
Python
"""Deterministic transcript representations for one-call protocol prompts."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Any
|
|
|
|
|
|
class TranscriptInputError(ValueError):
|
|
"""Raised when a transcript cannot be represented without content loss."""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SpeakerBlock:
|
|
"""One contiguous run of transcript segments assigned to one speaker."""
|
|
|
|
speaker_id: str
|
|
segment_texts: tuple[str, ...]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class CompactDiarizedTranscript:
|
|
"""Compact prompt text plus structural evidence of segment preservation."""
|
|
|
|
text: str
|
|
blocks: tuple[SpeakerBlock, ...]
|
|
source_segment_count: int
|
|
|
|
@property
|
|
def represented_segment_count(self) -> int:
|
|
return sum(len(block.segment_texts) for block in self.blocks)
|
|
|
|
@property
|
|
def segment_texts(self) -> tuple[str, ...]:
|
|
return tuple(text for block in self.blocks for text in block.segment_texts)
|
|
|
|
|
|
def normalize_segment_text(value: Any, index: int) -> str:
|
|
"""Normalize formatting whitespace while retaining all semantic text."""
|
|
if not isinstance(value, str):
|
|
raise TranscriptInputError(f"Transcript segment {index} text must be a string.")
|
|
return " ".join(value.split())
|
|
|
|
|
|
def compact_diarized_transcript(segments: Any) -> CompactDiarizedTranscript:
|
|
"""Group only adjacent same-speaker segments and omit repeated timestamps."""
|
|
if not isinstance(segments, list) or not segments:
|
|
raise TranscriptInputError(
|
|
"Diarized transcript must contain a non-empty 'segments' list."
|
|
)
|
|
|
|
mutable_blocks: list[tuple[str, list[str]]] = []
|
|
source_texts: list[str] = []
|
|
for index, segment in enumerate(segments):
|
|
if not isinstance(segment, dict):
|
|
raise TranscriptInputError(f"Transcript segment {index} must be an object.")
|
|
speaker = segment.get("speaker_id") or "SPEAKER_UNASSIGNED"
|
|
if not isinstance(speaker, str) or not speaker.startswith("SPEAKER_"):
|
|
raise TranscriptInputError(
|
|
f"Transcript segment {index} must use an anonymous SPEAKER_ label."
|
|
)
|
|
text = normalize_segment_text(segment.get("text"), index)
|
|
source_texts.append(text)
|
|
if mutable_blocks and mutable_blocks[-1][0] == speaker:
|
|
mutable_blocks[-1][1].append(text)
|
|
else:
|
|
mutable_blocks.append((speaker, [text]))
|
|
|
|
blocks = tuple(
|
|
SpeakerBlock(speaker_id=speaker, segment_texts=tuple(texts))
|
|
for speaker, texts in mutable_blocks
|
|
)
|
|
rendered = "\n".join(
|
|
f"{block.speaker_id}: {' '.join(block.segment_texts)}" for block in blocks
|
|
)
|
|
result = CompactDiarizedTranscript(
|
|
text=rendered + "\n",
|
|
blocks=blocks,
|
|
source_segment_count=len(segments),
|
|
)
|
|
if result.represented_segment_count != len(segments):
|
|
raise TranscriptInputError("Compact diarized transcript lost source segments.")
|
|
if result.segment_texts != tuple(source_texts):
|
|
raise TranscriptInputError("Compact diarized transcript changed segment order or text.")
|
|
return result
|
|
|
|
|
|
def plain_segment_transcript(segments: Any) -> str:
|
|
"""Reconstruct plain transcript text from every segment in source order."""
|
|
if not isinstance(segments, list) or not segments:
|
|
raise TranscriptInputError("Transcript must contain a non-empty 'segments' list.")
|
|
texts = []
|
|
for index, segment in enumerate(segments):
|
|
if not isinstance(segment, dict):
|
|
raise TranscriptInputError(f"Transcript segment {index} must be an object.")
|
|
texts.append(normalize_segment_text(segment.get("text"), index))
|
|
return " ".join(texts)
|