Files
meeting-lab/tests/test_direct_protocol.py

342 lines
15 KiB
Python

import json
import tempfile
import unittest
from datetime import datetime
from pathlib import Path
from unittest.mock import Mock, patch
import requests
from scripts import run_direct_protocol
from src.meeting_lab.llm import ollama
from src.meeting_lab.llm.ollama import OllamaError, OllamaGeneration
from src.meeting_lab.protocol.generate_direct_protocol import (
DirectProtocolError,
generate_direct_protocol,
load_compact_transcript,
)
VALID_CONTEXT = """schema_version: "1"
meeting:
meeting_id: "test-meeting"
title: "Test Meeting"
language: "de"
participants: []
mentioned_people: []
organization:
departments: []
known_entities: {}
"""
def write_transcript(path: Path, text: str = "Wir besprechen den Projektstatus.") -> None:
path.write_text(json.dumps({"text": text, "segments": []}), encoding="utf-8")
def generation(text: str = "# Meeting Protocol\n\n## Status\nUnveraendert.") -> OllamaGeneration:
return OllamaGeneration(
raw_response={
"response": text,
"done": True,
"done_reason": "stop",
"prompt_eval_count": 123,
"eval_count": 17,
"prompt_eval_duration": 1000,
"eval_duration": 2000,
"total_duration": 4000,
},
text=text,
client_wall_time_seconds=0.25,
)
class TranscriptLoadingTests(unittest.TestCase):
def test_valid_transcript_is_accepted(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "transcript.json"
write_transcript(path)
self.assertEqual(load_compact_transcript(path), "Wir besprechen den Projektstatus.")
def test_missing_top_level_text_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "transcript.json"
path.write_text('{"segments": []}', encoding="utf-8")
with self.assertRaisesRegex(DirectProtocolError, "top-level 'text'"):
load_compact_transcript(path)
def test_empty_transcript_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "transcript.json"
write_transcript(path, " \n")
with self.assertRaisesRegex(DirectProtocolError, "non-empty string"):
load_compact_transcript(path)
def test_malformed_json_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "transcript.json"
path.write_text("{", encoding="utf-8")
with self.assertRaisesRegex(DirectProtocolError, "not valid JSON"):
load_compact_transcript(path)
class GeneratorTests(unittest.TestCase):
def test_prompt_requires_contextual_discussion_density_without_transcript_replay(self) -> None:
with tempfile.TemporaryDirectory() as directory:
transcript = Path(directory) / "transcript.json"
write_transcript(transcript)
result = generate_direct_protocol(
transcript,
model_check=Mock(return_value={}),
generation_call=Mock(return_value=generation()),
)
self.assertIn("vollständiges, strukturiertes", result.exact_prompt)
self.assertIn("relevante Diskussionsverläufe", result.exact_prompt)
self.assertIn("unterschiedliche Positionen", result.exact_prompt)
self.assertIn("Entscheidungsgrundlagen", result.exact_prompt)
self.assertIn("nicht am Meeting teilgenommen haben", result.exact_prompt)
self.assertIn("keine reine Wiedergabe des Transkripts", result.exact_prompt)
self.assertIn("nicht unnötig durch Wiederholungen", result.exact_prompt)
def test_optional_context_absent_and_generation_called_once(self) -> None:
with tempfile.TemporaryDirectory() as directory:
transcript = Path(directory) / "transcript.json"
write_transcript(transcript)
check = Mock(return_value={"model": "qwen3.6:35B-A3B"})
call = Mock(return_value=generation())
result = generate_direct_protocol(
transcript,
model_check=check,
generation_call=call,
)
self.assertIn("Kein Meeting-Kontext", result.exact_prompt)
self.assertEqual(check.call_count, 1)
self.assertEqual(call.call_count, 1)
self.assertEqual(call.call_args.args[1], "qwen3.6:35B-A3B")
self.assertEqual(call.call_args.kwargs["num_ctx"], 32768)
self.assertEqual(result.runtime_metadata["request_count"], 1)
self.assertEqual(result.runtime_metadata["prompt_token_count"], 123)
self.assertFalse(result.runtime_metadata["think"])
self.assertEqual(result.runtime_metadata["temperature"], 0.0)
def test_valid_context_is_loaded_and_rendered(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = root / "transcript.json"
context = root / "context.yaml"
write_transcript(transcript)
context.write_text(VALID_CONTEXT, encoding="utf-8")
result = generate_direct_protocol(
transcript,
context,
model_check=Mock(return_value={}),
generation_call=Mock(return_value=generation()),
)
self.assertIn("MEETING CONTEXT V1", result.exact_prompt)
self.assertIn("Test Meeting", result.exact_prompt)
def test_invalid_context_is_rejected_before_network_calls(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = root / "transcript.json"
context = root / "context.yaml"
write_transcript(transcript)
context.write_text("schema_version: wrong", encoding="utf-8")
check = Mock()
call = Mock()
with self.assertRaisesRegex(ValueError, "schema_version"):
generate_direct_protocol(
transcript,
context,
model_check=check,
generation_call=call,
)
check.assert_not_called()
call.assert_not_called()
class OllamaTests(unittest.TestCase):
def test_unavailable_endpoint_failure(self) -> None:
with patch.object(ollama.requests, "get", side_effect=requests.ConnectionError("down")):
with self.assertRaisesRegex(OllamaError, "not reachable"):
ollama.require_model("http://127.0.0.1:11434", "model")
def test_missing_model_failure(self) -> None:
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {"models": [{"name": "other:model"}]}
with patch.object(ollama.requests, "get", return_value=response):
with self.assertRaisesRegex(OllamaError, "not installed"):
ollama.require_model("http://127.0.0.1:11434", "model")
def test_request_settings_and_raw_response(self) -> None:
raw = {"response": "# Meeting Protocol", "done": True}
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = raw
with patch.object(ollama.requests, "post", return_value=response) as post:
result = ollama.generate_once(
"http://localhost:11434",
"qwen3.8:27b",
"prompt",
timeout=30,
num_ctx=32768,
num_predict=8192,
)
self.assertEqual(post.call_count, 1)
payload = post.call_args.kwargs["json"]
self.assertEqual(payload["model"], "qwen3.8:27b")
self.assertEqual(payload["options"]["temperature"], 0.0)
self.assertEqual(payload["options"]["num_ctx"], 32768)
self.assertFalse(payload["think"])
self.assertFalse(payload["stream"])
self.assertEqual(result.raw_response, raw)
def test_request_omits_num_thread_when_backend_selection_is_requested(self) -> None:
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {"response": "# Meeting Protocol"}
with patch.object(ollama.requests, "post", return_value=response) as post:
ollama.generate_once(
"http://localhost:11434",
"qwen3.8:27b",
"prompt",
timeout=30,
num_ctx=32768,
num_predict=8192,
)
payload = post.call_args.kwargs["json"]
self.assertNotIn("num_thread", payload["options"])
def test_malformed_response_failure_without_retry(self) -> None:
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {"message": "missing response"}
with patch.object(ollama.requests, "post", return_value=response) as post:
with self.assertRaisesRegex(OllamaError, "no string 'response'"):
ollama.generate_once("url", "model", "prompt", timeout=1, num_ctx=1, num_predict=1)
self.assertEqual(post.call_count, 1)
def test_empty_response_failure(self) -> None:
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {"response": " "}
with patch.object(ollama.requests, "post", return_value=response):
with self.assertRaisesRegex(OllamaError, "empty protocol"):
ollama.generate_once("url", "model", "prompt", timeout=1, num_ctx=1, num_predict=1)
class DirectProtocolCliTests(unittest.TestCase):
def test_artifacts_are_preserved_and_protocol_is_untouched(self) -> None:
protocol_text = "# Meeting Protocol\n\nExact output. \n"
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = root / "source.json"
context = root / "source.yaml"
write_transcript(transcript)
context.write_text(VALID_CONTEXT, encoding="utf-8")
args = run_direct_protocol.parse_args(
[str(transcript), "--context", str(context), "--output-root", str(root / "runs")]
)
with patch.object(
run_direct_protocol,
"generate_direct_protocol",
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},
"model_metadata": {},
})(),
) as generator:
code, run_dir, protocol_path = run_direct_protocol.run(args)
self.assertEqual(code, 0)
self.assertEqual(generator.call_count, 1)
self.assertEqual(protocol_path.read_text(encoding="utf-8"), protocol_text)
self.assertEqual(
(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,
)
self.assertEqual(
json.loads((run_dir / "protocol/runtime_metadata.json").read_text())["request_count"],
1,
)
self.assertTrue((run_dir / "transcript/transcript.json").is_file())
self.assertTrue((run_dir / "context/meeting_context.yaml").is_file())
self.assertTrue((run_dir / "input_manifest.json").is_file())
self.assertEqual(json.loads((run_dir / "run_metadata.json").read_text())["status"], "completed")
def test_unique_run_directories_do_not_overwrite(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
fixed = datetime(2026, 8, 20, 12, 0, 0)
first = run_direct_protocol.create_unique_run_dir(root, "meeting", lambda: fixed)
marker = first / "keep.txt"
marker.write_text("keep", encoding="utf-8")
second = run_direct_protocol.create_unique_run_dir(root, "meeting", lambda: fixed)
self.assertEqual(first.name, "meeting_20260820_120000")
self.assertEqual(second.name, "meeting_20260820_120000_01")
self.assertEqual(marker.read_text(encoding="utf-8"), "keep")
def test_failure_after_directory_creation_preserves_metadata(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = run_direct_protocol.parse_args(
[str(root / "missing.json"), "--output-root", str(root / "runs")]
)
code, run_dir, protocol_path = run_direct_protocol.run(args)
metadata = json.loads((run_dir / "run_metadata.json").read_text())
self.assertEqual(code, 2)
self.assertIsNone(protocol_path)
self.assertEqual(metadata["status"], "failed")
self.assertIn("does not exist", metadata["failure"])
def test_semantic_pipeline_functions_are_never_invoked(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = root / "source.json"
write_transcript(transcript)
args = run_direct_protocol.parse_args(
[str(transcript), "--output-root", str(root / "runs")]
)
fake_result = type("Result", (), {
"protocol_text": "# Meeting Protocol",
"exact_prompt": "prompt",
"raw_response": {"response": "# Meeting Protocol"},
"runtime_metadata": {},
"model_metadata": {},
"transcript_input": None,
})()
with (
patch("src.meeting_lab.extraction.extract_chunks.extract_input") as extraction,
patch("src.meeting_lab.consolidation.consolidate_facts.call_ollama") as consolidation,
patch.object(run_direct_protocol, "generate_direct_protocol", return_value=fake_result),
):
code, _run_dir, _protocol_path = run_direct_protocol.run(args)
self.assertEqual(code, 0)
extraction.assert_not_called()
consolidation.assert_not_called()
if __name__ == "__main__":
unittest.main()