返回 ViMax
test_agent_llm.py
根目录 / tests / test_agent_llm.py
1 import unittest
2 from types import SimpleNamespace
3 from unittest.mock import AsyncMock
4
5 from agent_runtime.llm import OpenAICompatibleLLM
6
7
8 class AgentLLMTests(unittest.IsolatedAsyncioTestCase):
9 async def test_string_response_retries_before_clear_error(self):
10 llm = OpenAICompatibleLLM(model="test", base_url="https://example.invalid/v1", api_key="test-key")
11 create = AsyncMock(side_effect=["data: [DONE]", "bad response"])
12 llm.client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=create)))
13 with self.assertRaisesRegex(RuntimeError, "returned a string"):
14 await llm.complete([], [])
15 self.assertEqual(create.await_count, 2)
16
17 async def test_string_response_retry_can_recover(self):
18 llm = OpenAICompatibleLLM(model="test", base_url="https://example.invalid/v1", api_key="test-key")
19 create = AsyncMock(side_effect=["data: [DONE]", {"choices": [{"message": {"content": "recovered", "tool_calls": []}}]}])
20 llm.client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=create)))
21 message = await llm.complete([], [])
22 self.assertEqual(message.text, "recovered")
23 self.assertEqual(create.await_count, 2)
24
25 async def test_tool_request_falls_back_to_plain_chat_after_bad_tool_responses(self):
26 llm = OpenAICompatibleLLM(model="test", base_url="https://example.invalid/v1", api_key="test-key")
27 create = AsyncMock(side_effect=[
28 "data: [DONE]",
29 "data: [DONE]",
30 {"choices": [{"message": {"content": "plain fallback", "tool_calls": []}}]},
31 ])
32 llm.client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=create)))
33 message = await llm.complete([], [{"type": "function", "function": {"name": "x", "parameters": {}}}])
34 self.assertEqual(message.text, "plain fallback")
35 self.assertEqual(create.await_count, 3)
36 self.assertIsNone(create.await_args_list[-1].kwargs.get("tools"))
37
38 async def test_dict_response_is_accepted(self):
39 llm = OpenAICompatibleLLM(model="test", base_url="https://example.invalid/v1", api_key="test-key")
40 llm.client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=AsyncMock(return_value={
41 "choices": [{"message": {"content": "hello", "tool_calls": []}}]
42 }))))
43 message = await llm.complete([], [])
44 self.assertEqual(message.text, "hello")
45
46
47 if __name__ == "__main__":
48 unittest.main()
49
49 lines PYTHON