Support protocol thread configuration
This commit is contained in:
@@ -198,6 +198,23 @@ class OllamaTests(unittest.TestCase):
|
||||
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
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user