Files
meeting-lab/tests/test_protocol_threads.py
T

59 lines
2.6 KiB
Python

"""Protocol thread configuration through the public pipeline boundaries."""
import tempfile
import unittest
from dataclasses import replace
from pathlib import Path
from unittest.mock import Mock, patch
from src.meeting_lab.llm import ollama
from src.meeting_lab.orchestration import mvp
from tests import test_mvp_api as fixtures
class ProtocolThreadTests(unittest.TestCase):
def test_initial_and_regeneration_forward_default_and_explicit_threads(self):
for threads in (None, 4, 10, 16):
with self.subTest(threads=threads), tempfile.TemporaryDirectory() as directory:
config = replace(
fixtures.MvpApiTests().config(Path(directory)), protocol_num_thread=threads
)
with (
patch.object(mvp, "prepare_audio", side_effect=fixtures.fake_prepare),
patch.object(mvp, "transcribe_audio", side_effect=fixtures.fake_transcribe),
patch.object(
mvp, "generate_direct_protocol", side_effect=fixtures.fake_protocol
) as generate,
):
result = mvp.run_mvp_meeting(config, meeting_context=fixtures.context_data())
self.assertEqual(result.exit_code, 0)
self.assertEqual(generate.call_args.kwargs["num_thread"], threads)
mvp.regenerate_mvp_protocol(
result.run_dir,
meeting_context=fixtures.context_data(),
protocol_num_thread=threads,
)
self.assertEqual(generate.call_args.kwargs["num_thread"], threads)
def test_explicit_payload_and_invalid_values(self):
response = Mock()
response.json.return_value = {"response": "protocol"}
with patch.object(ollama.requests, "post", return_value=response) as post:
ollama.generate_once(
"url", "model", "prompt", timeout=1, num_ctx=32, num_predict=8, num_thread=10
)
self.assertEqual(post.call_args.kwargs["json"]["options"]["num_thread"], 10)
post.reset_mock()
for value in (0, -1, True, 1.5):
with self.assertRaises(ValueError):
ollama.generate_once(
"url",
"model",
"prompt",
timeout=1,
num_ctx=32,
num_predict=8,
num_thread=value,
)
post.assert_not_called()