import json import tempfile import unittest from pathlib import Path from unittest.mock import Mock from src.meeting_lab.diarization.alignment import diarized_transcript_text from src.meeting_lab.llm.ollama import OllamaGeneration from src.meeting_lab.protocol.direct_protocol_prompt import ( COMPACT_DIARIZED_PROTOCOL_INSTRUCTION, build_direct_protocol_prompt, ) from src.meeting_lab.protocol.generate_direct_protocol import ( DirectProtocolError, estimate_input_tokens, generate_direct_protocol, ) from src.meeting_lab.protocol.transcript_input import compact_diarized_transcript def segments() -> list[dict[str, object]]: return [ {"id": 0, "start": 0.0, "end": 1.0, "text": "First.", "speaker_id": "SPEAKER_01"}, {"id": 1, "start": 1.0, "end": 2.0, "text": "Second.", "speaker_id": "SPEAKER_01"}, {"id": 2, "start": 2.0, "end": 3.0, "text": "Third.", "speaker_id": "SPEAKER_04"}, {"id": 3, "start": 3.0, "end": 4.0, "text": "Unassigned.", "speaker_id": None}, {"id": 4, "start": 4.0, "end": 5.0, "text": "Last.", "speaker_id": "SPEAKER_01"}, ] def diarized_document(repetitions: int = 1) -> dict[str, object]: source = segments() * repetitions return { "text": diarized_transcript_text(source), "segments": source, "speaker_labels_anonymous": True, "alignment_source": "exclusive_diarization", } def completed_generation() -> OllamaGeneration: text = "# Meeting Protocol\n\nComplete." return OllamaGeneration( raw_response={ "response": text, "done": True, "done_reason": "stop", "prompt_eval_count": 100, "eval_count": 10, }, text=text, client_wall_time_seconds=0.1, ) class CompactDiarizedTranscriptTests(unittest.TestCase): def test_adjacent_segments_group_and_transitions_remain_separate(self) -> None: compact = compact_diarized_transcript(segments()) self.assertEqual( [block.speaker_id for block in compact.blocks], ["SPEAKER_01", "SPEAKER_04", "SPEAKER_UNASSIGNED", "SPEAKER_01"], ) self.assertEqual(compact.blocks[0].segment_texts, ("First.", "Second.")) self.assertEqual(compact.blocks[-1].segment_texts, ("Last.",)) self.assertEqual(compact.text.count("SPEAKER_01:"), 2) def test_every_segment_text_and_order_are_preserved(self) -> None: source = segments() compact = compact_diarized_transcript(source) self.assertEqual(compact.source_segment_count, len(source)) self.assertEqual(compact.represented_segment_count, len(source)) self.assertEqual( compact.segment_texts, tuple(str(segment["text"]) for segment in source), ) self.assertEqual(compact.segment_texts[0], "First.") self.assertEqual(compact.segment_texts[-1], "Last.") self.assertIn("SPEAKER_UNASSIGNED: Unassigned.", compact.text) def test_compact_form_is_materially_smaller_than_per_segment_format(self) -> None: source = [ { "start": index, "end": index + 1, "text": "Repeated transcript content.", "speaker_id": "SPEAKER_01", } for index in range(100) ] compact = compact_diarized_transcript(source).text verbose = diarized_transcript_text(source) self.assertLess(len(compact), len(verbose) * 0.6) class ProtocolInputBudgetTests(unittest.TestCase): def _write(self, root: Path, document: dict[str, object]) -> Path: path = root / "transcript.json" path.write_text(json.dumps(document), encoding="utf-8") return path def test_token_estimate_uses_utf8_bytes_for_non_ascii_safety(self) -> None: self.assertEqual(estimate_input_tokens("รค" * 44), 20) def test_compact_diarized_representation_selected_within_budget(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) document = diarized_document() transcript = self._write(root, document) compact = compact_diarized_transcript(document["segments"]).text budget = estimate_input_tokens( build_direct_protocol_prompt( compact, instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION, ) ) call = Mock(return_value=completed_generation()) result = generate_direct_protocol( transcript, safe_input_token_budget=budget, model_check=Mock(return_value={}), generation_call=call, ) self.assertEqual( result.runtime_metadata["selected_transcript_representation"], "diarized_compact", ) self.assertFalse(result.runtime_metadata["fallback_used"]) self.assertTrue(result.runtime_metadata["diarization_enabled"]) self.assertEqual(result.runtime_metadata["safe_input_token_budget"], budget) self.assertEqual(result.transcript_input, compact) self.assertEqual(call.call_count, 1) def test_plain_fallback_selected_when_diarized_compact_exceeds_budget(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) alternating_segments = [ { "id": index, "start": float(index), "end": float(index + 1), "text": "Word.", "speaker_id": f"SPEAKER_{index % 2:02d}", } for index in range(200) ] document = { "text": diarized_transcript_text(alternating_segments), "segments": alternating_segments, "speaker_labels_anonymous": True, "alignment_source": "exclusive_diarization", } transcript = self._write(root, document) compact = compact_diarized_transcript(document["segments"]).text plain = " ".join(str(segment["text"]) for segment in document["segments"]) compact_estimate = estimate_input_tokens( build_direct_protocol_prompt( compact, instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION, ) ) plain_estimate = estimate_input_tokens(build_direct_protocol_prompt(plain)) self.assertLess(plain_estimate, compact_estimate) result = generate_direct_protocol( transcript, safe_input_token_budget=plain_estimate, model_check=Mock(return_value={}), generation_call=Mock(return_value=completed_generation()), ) self.assertEqual( result.runtime_metadata["selected_transcript_representation"], "plain_transcript_fallback", ) self.assertTrue(result.runtime_metadata["fallback_used"]) self.assertEqual(result.transcript_input, plain) def test_oversized_plain_transcript_fails_before_any_network_call(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) transcript = self._write( root, {"text": "large input " * 1000, "segments": []}, ) model_check = Mock() generation_call = Mock() with self.assertRaisesRegex( DirectProtocolError, "No LLM request was made; silent truncation is not allowed", ): generate_direct_protocol( transcript, safe_input_token_budget=1, model_check=model_check, generation_call=generation_call, ) model_check.assert_not_called() generation_call.assert_not_called() def test_existing_plain_path_and_metadata_remain_direct(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) transcript = self._write( root, {"text": "Plain original transcript.", "segments": []}, ) result = generate_direct_protocol( transcript, model_check=Mock(return_value={}), generation_call=Mock(return_value=completed_generation()), ) self.assertEqual(result.transcript_input, "Plain original transcript.") self.assertEqual( result.runtime_metadata["selected_transcript_representation"], "plain_transcript", ) self.assertFalse(result.runtime_metadata["fallback_used"]) self.assertFalse(result.runtime_metadata["diarization_enabled"]) if __name__ == "__main__": unittest.main()