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 categories: Optional[List[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, categories=None if segment.categories is None else list(segment.categories), 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) return NormalizationResult( transcript=normalized, 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, ), ) 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), categories=_merged_categories(left.categories, right.categories), 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 _merged_categories(left: Optional[List[str]], right: Optional[List[str]]) -> Optional[List[str]]: merged: List[str] = [] for category in (left or []) + (right or []): if category not in merged: merged.append(category) return merged or None 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, categories=None if segment.categories is None else list(segment.categories), ) for index, segment in enumerate(ordered) ]