Guard protocol generation against context truncation

This commit is contained in:
2026-08-24 16:48:06 +02:00
parent d94436af43
commit d77bfedb6e
12 changed files with 524 additions and 12 deletions
@@ -0,0 +1,97 @@
"""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)