59 lines
2.6 KiB
Python
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()
|