Guard protocol generation against context truncation
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user