Files
meeting-lab/tests/test_protocol_transcript_input.py

316 lines
12 KiB
Python

import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import Mock, patch
from src.meeting_lab.diarization.alignment import diarized_transcript_text
from src.meeting_lab.llm import ollama
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_mapped_speakers_and_statements_reach_final_ollama_payload(self) -> None:
diarized = {
"text": "",
"segments": [
{
"start": 0.0,
"end": 1.0,
"speaker_id": "SPEAKER_00",
"text": "We will run the trial on Wednesday.",
},
{
"start": 1.0,
"end": 2.0,
"speaker_id": "SPEAKER_01",
"text": "I will prepare the raw materials before then.",
},
{
"start": 2.0,
"end": 3.0,
"speaker_id": "SPEAKER_00",
"text": "Good. Anna owns the material preparation.",
},
],
"speaker_labels_anonymous": True,
"alignment_source": "exclusive_diarization",
}
context = {
"schema_version": "1",
"meeting": {
"meeting_id": "speaker-test",
"title": "Speaker test",
"language": "en",
},
"participants": [
{"participant_id": "martin", "display_name": "Martin"},
{"participant_id": "anna", "display_name": "Anna"},
],
"speaker_mappings": {"SPEAKER_00": "martin", "SPEAKER_01": "anna"},
"mentioned_people": [],
"organization": {"departments": []},
"known_entities": {},
}
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {"response": "# Meeting Protocol\n", "done": True}
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = self._write(root, diarized)
context_path = root / "context.yaml"
context_path.write_text(json.dumps(context), encoding="utf-8")
with patch.object(ollama.requests, "post", return_value=response) as post:
result = generate_direct_protocol(
transcript,
context_path,
model="qwen3.8:27b",
model_check=Mock(return_value={}),
)
prompt = post.call_args.kwargs["json"]["prompt"]
self.assertEqual(result.exact_prompt, prompt)
self.assertIn("- SPEAKER_00: Martin (participant_id: martin)", prompt)
self.assertIn("- SPEAKER_01: Anna (participant_id: anna)", prompt)
self.assertIn("SPEAKER_00: We will run the trial on Wednesday.", prompt)
self.assertIn(
"SPEAKER_01: I will prepare the raw materials before then.", prompt
)
self.assertIn("SPEAKER_00: Good. Anna owns the material preparation.", prompt)
self.assertNotIn("Martin: We will run the trial on Wednesday.", prompt)
self.assertNotIn("Anna: I will prepare the raw materials before then.", prompt)
self.assertIn("autoritativen SPEAKER_XX-zu-Teilnehmer-Zuordnungen", prompt)
self.assertIn("Ich-Zusage", prompt)
self.assertIn("nur erwähnten Personen", prompt)
self.assertIn("keine persönliche Verantwortung", prompt)
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)
self.assertNotIn("SPEAKER_00", result.exact_prompt)
self.assertNotIn("SPEAKER_01", result.exact_prompt)
self.assertFalse(result.runtime_metadata["speaker_attribution_available"])
self.assertEqual(
result.runtime_metadata["speaker_attribution_loss_reason"],
"plain_transcript_fallback",
)
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"])
self.assertIsNone(result.runtime_metadata["speaker_attribution_available"])
if __name__ == "__main__":
unittest.main()