Files
agentos/tests/test_model.py
T
2026-07-31 17:54:01 +08:00

86 lines
3.2 KiB
Python

import io
import sys
import threading
import time
import unittest
from concurrent.futures import ThreadPoolExecutor
from agentos.model import SubprocessModelClient, _extract_json
class ModelParsingTests(unittest.TestCase):
def test_extracts_json_after_loading_banner(self) -> None:
value = _extract_json('loaded\n[{"option_id":"1"}]\ntrailing')
self.assertEqual(value, [{"option_id": "1"}])
def test_falls_back_when_option_json_is_invalid(self) -> None:
client = SubprocessModelClient(("false",), timeout_s=1)
options = client._fallback_options("整理目录")
self.assertEqual(len(options), 2)
self.assertEqual(options[0].risk, "low")
def test_generation_streams_and_captures_the_same_output(self) -> None:
script = (
"import json, sys; request=json.load(sys.stdin); "
"assert sys.argv[1:] == ['--request-json-stdin']; "
"assert request['system_prompt'] == 'dynamic system'; "
"assert request['system_prompt_candidates'] == ['dynamic system']; "
"assert request['prompt'] == 'prompt'; "
"assert request['max_new_tokens'] == 10; "
"print('Qwen3-14B 已加载到 NPU, 用时 1s', flush=True); "
'print(\'[{"option_id": "1"}]\', flush=True)'
)
stream = io.StringIO()
client = SubprocessModelClient((sys.executable, "-c", script), timeout_s=5, stream=stream)
generated = client._generate("prompt", "dynamic system", max_tokens=10)
self.assertEqual(generated, '[{"option_id": "1"}]')
self.assertIn("Qwen3-14B 已加载到 NPU", stream.getvalue())
self.assertIn(generated, stream.getvalue())
def test_generation_can_be_captured_without_background_streaming(self) -> None:
script = "import json, sys; json.load(sys.stdin); print('background result', flush=True)"
stream = io.StringIO()
client = SubprocessModelClient(
(sys.executable, "-c", script),
timeout_s=5,
stream=stream,
)
with client.suppress_streaming():
generated = client._generate("prompt", "system")
self.assertEqual(generated, "background result")
self.assertEqual(stream.getvalue(), "")
def test_generation_is_serialized_for_one_npu(self) -> None:
active = 0
max_active = 0
state_lock = threading.Lock()
class RecordingClient(SubprocessModelClient):
def _generate_locked(self, prompt, system_prompt, max_tokens):
nonlocal active, max_active
del system_prompt, max_tokens
with state_lock:
active += 1
max_active = max(max_active, active)
try:
time.sleep(0.03)
return prompt
finally:
with state_lock:
active -= 1
client = RecordingClient(("unused",))
with ThreadPoolExecutor(max_workers=4) as executor:
outputs = list(executor.map(lambda value: client._generate(value, "system"), "abcd"))
self.assertEqual(outputs, list("abcd"))
self.assertEqual(max_active, 1)
if __name__ == "__main__":
unittest.main()