diff --git a/PROJECT_KNOWLEDGE.md b/PROJECT_KNOWLEDGE.md index d67ce13..4921fa6 100644 --- a/PROJECT_KNOWLEDGE.md +++ b/PROJECT_KNOWLEDGE.md @@ -33,6 +33,12 @@ Implemented: default and may be revisited after empirical comparison without changing the orchestration API. - Interim Markdown protocol generation in `src/meeting_lab/protocol/`. +- Direct protocol prompt input protection: diarized transcripts are rendered as + compact adjacent-speaker blocks without per-segment timestamps. Every source + segment remains represented in order. A deterministic heuristic enforces a + configurable safe input budget, falls back to complete plain transcript text + when necessary, and fails before any Ollama request if even that input is too + large. Silent head/tail truncation is prohibited. - Non-LLM unit tests for chunking, extraction helpers, protocol rendering and gold-test runner validation. - Meeting Context V1 scaffold and documentation for manually maintained diff --git a/docs/diarization.md b/docs/diarization.md index 088cd3c..b2bf7c9 100644 --- a/docs/diarization.md +++ b/docs/diarization.md @@ -35,4 +35,28 @@ torchcodec file decoder. Anonymous `SPEAKER_XX` labels are aligned to Whisper segments by maximum temporal overlap with Community-1 exclusive diarization. The original Whisper transcript is preserved; the derived transcript under `diarization/` is used as -the unchanged direct-protocol generator's input. +the direct-protocol generator's source input. + +The full diarized JSON and timestamped text remain immutable audit artifacts, +but their per-segment formatting is too verbose for a full-meeting LLM prompt: +timestamps and repeated speaker labels can more than double input size. For +protocol generation, Meeting Lab deterministically groups only adjacent +segments assigned to the same anonymous speaker and omits timestamps. A later +return by the same speaker starts a new block, and unassigned segments remain +under `SPEAKER_UNASSIGNED`. `protocol/transcript_input.txt` preserves the exact +derived representation sent to prompt construction. + +Before contacting Ollama, Meeting Lab conservatively estimates prompt tokens +from UTF-8 byte count without adding a model tokenizer dependency. The default safe +budget is 16,200 estimated tokens, below the observed 16,386-token effective +boundary even when a larger `num_ctx` was requested. The estimate is calibrated +against the currently validated German BPD input and is configurable through +`MvpMeetingConfig.protocol_safe_input_token_budget` or +`--protocol-safe-input-token-budget`. + +If compact diarized input exceeds the budget, the generator deterministically +uses the complete plain segment transcript and records the fallback. If that +also exceeds the budget, generation fails before model lookup or generation; +it never truncates, chunks, summarizes, retries, or makes multiple protocol +calls implicitly. Full diarization artifacts are never overwritten by this +selection. diff --git a/scripts/run_direct_protocol.py b/scripts/run_direct_protocol.py index 9c21878..d3a4f4d 100644 --- a/scripts/run_direct_protocol.py +++ b/scripts/run_direct_protocol.py @@ -20,6 +20,7 @@ if str(REPO_ROOT) not in sys.path: from src.meeting_lab.llm.ollama import DEFAULT_ENDPOINT # noqa: E402 from src.meeting_lab.protocol.generate_direct_protocol import ( # noqa: E402 DEFAULT_MODEL, + DEFAULT_SAFE_INPUT_TOKEN_BUDGET, DirectProtocolResult, generate_direct_protocol, ) @@ -35,6 +36,11 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT) parser.add_argument("--model", default=DEFAULT_MODEL) parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT) + parser.add_argument( + "--safe-input-token-budget", + type=int, + default=DEFAULT_SAFE_INPUT_TOKEN_BUDGET, + ) return parser.parse_args(argv) @@ -64,6 +70,11 @@ def persist_result(run_dir: Path, result: DirectProtocolResult) -> Path: (protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8") write_json(protocol_dir / "raw_response.json", result.raw_response) write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata) + transcript_input = getattr(result, "transcript_input", None) + if transcript_input is not None: + (protocol_dir / "transcript_input.txt").write_text( + transcript_input, encoding="utf-8" + ) protocol_path = run_dir / "protocol.md" protocol_path.write_text(result.protocol_text, encoding="utf-8") return protocol_path @@ -116,6 +127,7 @@ def run(args: argparse.Namespace) -> tuple[int, Path, Path | None]: preserved_context, model=args.model, endpoint=args.ollama_endpoint, + safe_input_token_budget=args.safe_input_token_budget, ) protocol_path = persist_result(run_dir, result) metadata["status"] = "completed" diff --git a/scripts/run_mvp_meeting.py b/scripts/run_mvp_meeting.py index 10a97e9..9c36f7d 100644 --- a/scripts/run_mvp_meeting.py +++ b/scripts/run_mvp_meeting.py @@ -19,6 +19,7 @@ from src.meeting_lab.orchestration.mvp import ( # noqa: E402 DEFAULT_DIARIZATION_MODEL, DEFAULT_MODEL, DEFAULT_OUTPUT_ROOT, + DEFAULT_SAFE_INPUT_TOKEN_BUDGET, MvpMeetingConfig, create_unique_run_dir, run_mvp_meeting, @@ -53,6 +54,12 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: ) parser.add_argument("--model", default=DEFAULT_MODEL) parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT) + parser.add_argument( + "--protocol-safe-input-token-budget", + type=int, + default=DEFAULT_SAFE_INPUT_TOKEN_BUDGET, + help="Conservative estimated prompt-token limit before any Ollama request.", + ) parser.add_argument( "--diarization", choices=("auto", "gpu", "cpu", "off"), @@ -91,6 +98,7 @@ def config_from_args(args: argparse.Namespace) -> MvpMeetingConfig: threads=args.threads, model=args.model, ollama_endpoint=args.ollama_endpoint, + protocol_safe_input_token_budget=args.protocol_safe_input_token_budget, diarization=args.diarization, diarization_runtime=args.diarization_runtime, diarization_container_image=args.diarization_container_image, diff --git a/src/meeting_lab/orchestration/mvp.py b/src/meeting_lab/orchestration/mvp.py index 55e3653..5088c5b 100644 --- a/src/meeting_lab/orchestration/mvp.py +++ b/src/meeting_lab/orchestration/mvp.py @@ -30,6 +30,7 @@ from src.meeting_lab.models.meeting_context import ( from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus from src.meeting_lab.protocol.generate_direct_protocol import ( DEFAULT_MODEL, + DEFAULT_SAFE_INPUT_TOKEN_BUDGET, DirectProtocolResult, generate_direct_protocol, load_compact_transcript, @@ -54,6 +55,7 @@ class MvpMeetingConfig: threads: str | int = "auto" model: str = DEFAULT_MODEL ollama_endpoint: str = DEFAULT_ENDPOINT + protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET diarization: str = "off" diarization_runtime: str = "native" diarization_container_image: str | None = None @@ -127,6 +129,8 @@ def _validate_inputs( and not config.diarization_container_image ): raise ValueError("A diarization container image is required.") + if config.protocol_safe_input_token_budget <= 0: + raise ValueError("Protocol safe input token budget must be positive.") def _emit( @@ -154,6 +158,11 @@ def _persist_protocol(run_dir: Path, result: DirectProtocolResult) -> Path: (protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8") _write_json(protocol_dir / "raw_response.json", result.raw_response) _write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata) + transcript_input = getattr(result, "transcript_input", None) + if transcript_input is not None: + (protocol_dir / "transcript_input.txt").write_text( + transcript_input, encoding="utf-8" + ) protocol_path = run_dir / "protocol.md" protocol_path.write_text(result.protocol_text, encoding="utf-8") return protocol_path @@ -354,6 +363,7 @@ def run_mvp_meeting( preserved_context, model=config.model, endpoint=config.ollama_endpoint, + safe_input_token_budget=config.protocol_safe_input_token_budget, ) stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3) protocol_path = _persist_protocol(run_dir, result) diff --git a/src/meeting_lab/protocol/direct_protocol_prompt.py b/src/meeting_lab/protocol/direct_protocol_prompt.py index 5815adc..429df4a 100644 --- a/src/meeting_lab/protocol/direct_protocol_prompt.py +++ b/src/meeting_lab/protocol/direct_protocol_prompt.py @@ -9,11 +9,22 @@ Das Protokoll muss themenorientiert sein, nicht chronologisch und nicht nach tec Erzeuge keine reine Wiedergabe des Transkripts und verlängere das Protokoll nicht unnötig durch Wiederholungen. Synthetisiere zusammengehörige Aussagen, entferne Füllwörter und Gesprächsrauschen und erfinde keine Fakten, Entscheidungen, Zustimmungen, Verantwortlichen oder Fristen. Gib kein JSON, keine internen Labels und keine Analyse oder Denkprotokolle aus. Das Ergebnis soll als Markdown-Protokoll nach geringfügiger menschlicher Redaktion intern versendbar sein. Eine kompakte themenübergreifende Maßnahmenliste am Ende ist optional, wenn sie nützlich und vollständig belegt ist.""" +COMPACT_DIARIZED_PROTOCOL_INSTRUCTION = """Erstelle aus dem vollständigen Transkript und Meeting-Kontext ein vollständiges, professionelles internes Besprechungsprotokoll auf Deutsch. Das Transkript ist in aufeinanderfolgende anonyme Sprecherblöcke gegliedert. -def build_direct_protocol_prompt(transcript: str, meeting_context: str | None = None) -> str: +Beginne mit # Meeting Protocol. Gliedere themenorientiert mit ## und synthetisiere je Thema den relevanten Diskussionsverlauf, Kontext, unterschiedliche Positionen, Entscheidungsgrundlagen, Einschränkungen und ungelöste Meinungsverschiedenheiten so, dass Dritte ihn nachvollziehen können. Nenne Entscheidungen nur bei Beleg. Nenne Maßnahmen, Verantwortliche und Fristen nur bei expliziter Zuweisung, Annahme oder Bestätigung; Vorschläge sind keine Verpflichtungen. + +Entferne nur Wiederholungen, Füllwörter und Gesprächsrauschen. Erfinde keine Fakten oder Identitäten. Gib kein JSON, keine Sprecherlabels und kein Denkprotokoll aus. Eine belegte themenübergreifende Maßnahmenliste am Ende ist optional.""" + + +def build_direct_protocol_prompt( + transcript: str, + meeting_context: str | None = None, + *, + instruction: str = DIRECT_PROTOCOL_INSTRUCTION, +) -> str: context = meeting_context.strip() if meeting_context else "Kein Meeting-Kontext bereitgestellt." return ( - f"{DIRECT_PROTOCOL_INSTRUCTION}\n\n" + f"{instruction}\n\n" f"MEETING-KONTEXT:\n{context}\n\n" f"VOLLSTAENDIGES TRANSKRIPT:\n{transcript.strip()}\n" ) diff --git a/src/meeting_lab/protocol/generate_direct_protocol.py b/src/meeting_lab/protocol/generate_direct_protocol.py index 434a59d..4a0b046 100644 --- a/src/meeting_lab/protocol/generate_direct_protocol.py +++ b/src/meeting_lab/protocol/generate_direct_protocol.py @@ -18,13 +18,23 @@ from src.meeting_lab.models.meeting_context import ( load_meeting_context, render_meeting_context_for_prompt, ) -from src.meeting_lab.protocol.direct_protocol_prompt import build_direct_protocol_prompt +from src.meeting_lab.protocol.direct_protocol_prompt import ( + COMPACT_DIARIZED_PROTOCOL_INSTRUCTION, + build_direct_protocol_prompt, +) +from src.meeting_lab.protocol.transcript_input import ( + TranscriptInputError, + compact_diarized_transcript, + plain_segment_transcript, +) DEFAULT_MODEL = "qwen3.6:35B-A3B" DEFAULT_NUM_CTX = 32768 DEFAULT_NUM_PREDICT = 8192 DEFAULT_TIMEOUT = 1800 +DEFAULT_SAFE_INPUT_TOKEN_BUDGET = 16_200 +ESTIMATED_UTF8_BYTES_PER_TOKEN = 4.4 class DirectProtocolError(ValueError): @@ -38,9 +48,29 @@ class DirectProtocolResult: model_metadata: dict[str, Any] runtime_metadata: dict[str, Any] raw_response: dict[str, Any] + transcript_input: str | None = None + + +@dataclass(frozen=True) +class SelectedTranscriptInput: + text: str + prompt: str + representation: str + estimated_input_tokens: int + safe_input_token_budget: int + fallback_used: bool + diarization_enabled: bool def load_compact_transcript(path: Path) -> str: + data = _load_transcript_document(path) + text = data.get("text") + if not isinstance(text, str) or not text.strip(): + raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.") + return text + + +def _load_transcript_document(path: Path) -> dict[str, Any]: if not path.is_file(): raise DirectProtocolError(f"Transcript file does not exist: {path}") try: @@ -51,10 +81,73 @@ def load_compact_transcript(path: Path) -> str: raise DirectProtocolError("Transcript JSON must contain a top-level object.") if "text" not in data: raise DirectProtocolError("Transcript JSON must contain top-level 'text'.") - text = data["text"] - if not isinstance(text, str) or not text.strip(): - raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.") - return text + return data + + +def estimate_input_tokens(prompt: str) -> int: + """Estimate tokens without adding a model-specific tokenizer dependency.""" + byte_count = len(prompt.encode("utf-8")) + return max(1, int(byte_count / ESTIMATED_UTF8_BYTES_PER_TOKEN + 0.999999)) + + +def select_transcript_input( + transcript: dict[str, Any], + rendered_context: str | None, + *, + safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, +) -> SelectedTranscriptInput: + """Select complete prompt input without allowing silent tail truncation.""" + if safe_input_token_budget <= 0: + raise DirectProtocolError("Safe protocol input token budget must be positive.") + diarization_enabled = transcript.get("speaker_labels_anonymous") is True + if diarization_enabled: + try: + compact = compact_diarized_transcript(transcript.get("segments")) + plain_text = plain_segment_transcript(transcript.get("segments")) + except TranscriptInputError as exc: + raise DirectProtocolError(str(exc)) from exc + compact_prompt = build_direct_protocol_prompt( + compact.text, + rendered_context, + instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION, + ) + compact_estimate = estimate_input_tokens(compact_prompt) + if compact_estimate <= safe_input_token_budget: + return SelectedTranscriptInput( + text=compact.text, + prompt=compact_prompt, + representation="diarized_compact", + estimated_input_tokens=compact_estimate, + safe_input_token_budget=safe_input_token_budget, + fallback_used=False, + diarization_enabled=True, + ) + representation = "plain_transcript_fallback" + fallback_used = True + else: + plain_text = transcript.get("text") + if not isinstance(plain_text, str) or not plain_text.strip(): + raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.") + representation = "plain_transcript" + fallback_used = False + + plain_prompt = build_direct_protocol_prompt(plain_text, rendered_context) + plain_estimate = estimate_input_tokens(plain_prompt) + if plain_estimate > safe_input_token_budget: + raise DirectProtocolError( + "Protocol prompt/input is too large for the configured safe input budget " + f"({plain_estimate} estimated tokens > {safe_input_token_budget}). " + "No LLM request was made; silent truncation is not allowed." + ) + return SelectedTranscriptInput( + text=plain_text, + prompt=plain_prompt, + representation=representation, + estimated_input_tokens=plain_estimate, + safe_input_token_budget=safe_input_token_budget, + fallback_used=fallback_used, + diarization_enabled=diarization_enabled, + ) def generate_direct_protocol( @@ -66,21 +159,26 @@ def generate_direct_protocol( timeout: int = DEFAULT_TIMEOUT, num_ctx: int = DEFAULT_NUM_CTX, num_predict: int = DEFAULT_NUM_PREDICT, + safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, model_check: Callable[[str, str, int], dict[str, Any]] = require_model, generation_call: Callable[..., OllamaGeneration] = generate_once, ) -> DirectProtocolResult: - transcript = load_compact_transcript(transcript_path) + transcript = _load_transcript_document(transcript_path) context: MeetingContext | None = ( load_meeting_context(context_path) if context_path is not None else None ) rendered_context = render_meeting_context_for_prompt(context) if context else None - prompt = build_direct_protocol_prompt(transcript, rendered_context) + selected = select_transcript_input( + transcript, + rendered_context, + safe_input_token_budget=safe_input_token_budget, + ) model_metadata = model_check(endpoint, model, 10) generation = generation_call( endpoint, model, - prompt, + selected.prompt, timeout=timeout, num_ctx=num_ctx, num_predict=num_predict, @@ -101,11 +199,18 @@ def generate_direct_protocol( "think": False, "num_ctx": num_ctx, "num_predict": num_predict, + "selected_transcript_representation": selected.representation, + "estimated_input_tokens": selected.estimated_input_tokens, + "safe_input_token_budget": selected.safe_input_token_budget, + "input_token_estimation_method": "utf8_bytes_divided_by_4.4", + "fallback_used": selected.fallback_used, + "diarization_enabled": selected.diarization_enabled, } return DirectProtocolResult( protocol_text=generation.text, - exact_prompt=prompt, + exact_prompt=selected.prompt, model_metadata=model_metadata, runtime_metadata=runtime_metadata, raw_response=data, + transcript_input=selected.text, ) diff --git a/src/meeting_lab/protocol/transcript_input.py b/src/meeting_lab/protocol/transcript_input.py new file mode 100644 index 0000000..dd17aec --- /dev/null +++ b/src/meeting_lab/protocol/transcript_input.py @@ -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) diff --git a/tests/test_direct_protocol.py b/tests/test_direct_protocol.py index 1868e5a..cd01f96 100644 --- a/tests/test_direct_protocol.py +++ b/tests/test_direct_protocol.py @@ -232,6 +232,7 @@ class DirectProtocolCliTests(unittest.TestCase): return_value=type("Result", (), { "protocol_text": protocol_text, "exact_prompt": "exact prompt\n", + "transcript_input": "selected transcript\n", "raw_response": {"response": protocol_text}, "runtime_metadata": {"request_count": 1}, })(), @@ -245,6 +246,10 @@ class DirectProtocolCliTests(unittest.TestCase): (run_dir / "protocol/exact_prompt.txt").read_text(encoding="utf-8"), "exact prompt\n", ) + self.assertEqual( + (run_dir / "protocol/transcript_input.txt").read_text(encoding="utf-8"), + "selected transcript\n", + ) self.assertEqual( json.loads((run_dir / "protocol/raw_response.json").read_text())["response"], protocol_text, diff --git a/tests/test_mvp_api.py b/tests/test_mvp_api.py index 6091ac5..fefd7a1 100644 --- a/tests/test_mvp_api.py +++ b/tests/test_mvp_api.py @@ -180,6 +180,7 @@ class MvpApiTests(unittest.TestCase): self.assertEqual(delegated.whisper_executable, "whisper-cli") self.assertEqual(delegated.ffmpeg_executable, "ffmpeg") self.assertTrue(delegated.audio_normalization) + self.assertEqual(delegated.protocol_safe_input_token_budget, 16_200) self.assertEqual(api.call_args.kwargs["meeting_context"], context_data()) def test_cli_explicit_audio_normalization_values_are_propagated(self): diff --git a/tests/test_mvp_orchestrator.py b/tests/test_mvp_orchestrator.py index 86c4043..6d1aa6b 100644 --- a/tests/test_mvp_orchestrator.py +++ b/tests/test_mvp_orchestrator.py @@ -39,6 +39,7 @@ def protocol_result(model: str = "chosen:model") -> DirectProtocolResult: model_metadata={"model": model}, runtime_metadata={"model": model, "request_count": 1, "client_wall_time_seconds": 0.5}, raw_response={"response": text, "done": True}, + transcript_input="selected transcript\n", ) @@ -186,6 +187,7 @@ class MvpOrchestratorTests(unittest.TestCase): "transcript/runtime_metadata.json", "context/meeting_context.yaml", "protocol/exact_prompt.txt", + "protocol/transcript_input.txt", "protocol/raw_response.json", "protocol/runtime_metadata.json", "protocol.md", diff --git a/tests/test_protocol_transcript_input.py b/tests/test_protocol_transcript_input.py new file mode 100644 index 0000000..3fd2921 --- /dev/null +++ b/tests/test_protocol_transcript_input.py @@ -0,0 +1,231 @@ +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()