322 lines
14 KiB
Python
322 lines
14 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_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},
|
|
})(),
|
|
) 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": {},
|
|
})()
|
|
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()
|