Initial commit
This commit is contained in:
@@ -0,0 +1,85 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user