返回 ViMax
test_robustness.py
根目录 / tests / test_robustness.py
1 """Regression tests for error-boundary and durability fixes.
2
3 Covers: LLM client retry/empty-choices handling, agent-loop turn error
4 boundary, session index corruption/atomicity/concurrency, and bounded
5 retry policies with backoff across agents and API clients.
6 """
7
8 import tempfile
9 import threading
10 import unittest
11 from pathlib import Path
12 from unittest.mock import AsyncMock, MagicMock, patch
13
14 from tenacity.stop import stop_never
15 from tenacity.wait import wait_none
16
17 from agent_runtime.llm import OpenAICompatibleLLM
18 from agent_runtime.loop import AgentLoop
19 from agent_runtime.prompts import PromptBuilder
20 from agent_runtime.session_index import SessionIndex
21 from agent_runtime.tool_executor import ToolExecutor
22 from agent_runtime.tools import ToolRegistry
23 from agents.screenwriter import Screenwriter
24 from agents.script_planner import ScriptPlanner
25 from tools.image_generator_doubao_seedream_yunwu_api import ImageGeneratorDoubaoSeedreamYunwuAPI
26 from tools.image_generator_nanobanana_google_api import ImageGeneratorNanobananaGoogleAPI
27 from tools.image_generator_nanobanana_yunwu_api import ImageGeneratorNanobananaYunwuAPI
28 from tools.reranker_bge_silicon_api import RerankerBgeSiliconapi
29
30
31 class FakeStatusError(Exception):
32 def __init__(self, status_code):
33 self.status_code = status_code
34 super().__init__(f"http status {status_code}")
35
36
37 def _fake_completion(text="ok"):
38 message = MagicMock()
39 message.content = text
40 message.tool_calls = None
41 message.model_dump.return_value = {}
42 return MagicMock(choices=[MagicMock(message=message)])
43
44
45 class TestLLMClient(unittest.IsolatedAsyncioTestCase):
46 def _llm(self, create):
47 llm = OpenAICompatibleLLM(model="m", base_url="http://localhost:1", api_key="k")
48 llm.client = MagicMock(chat=MagicMock(completions=MagicMock(create=create)))
49 return llm
50
51 async def test_retries_rate_limit_then_succeeds(self):
52 create = AsyncMock(side_effect=[FakeStatusError(429), _fake_completion("recovered")])
53 llm = self._llm(create)
54 result = await llm.complete([{"role": "user", "content": "x"}], tools=[])
55 self.assertEqual(result.text, "recovered")
56 self.assertEqual(create.await_count, 2)
57
58 async def test_does_not_retry_auth_errors(self):
59 create = AsyncMock(side_effect=FakeStatusError(401))
60 llm = self._llm(create)
61 with self.assertRaises(FakeStatusError):
62 await llm.complete([{"role": "user", "content": "x"}], tools=[])
63 self.assertEqual(create.await_count, 1)
64
65 async def test_gives_up_after_bounded_attempts(self):
66 create = AsyncMock(side_effect=FakeStatusError(500))
67 llm = self._llm(create)
68 with self.assertRaises(FakeStatusError):
69 await llm.complete([{"role": "user", "content": "x"}], tools=[])
70 self.assertLessEqual(create.await_count, 4)
71 self.assertGreater(create.await_count, 1)
72
73 async def test_empty_choices_raises_clear_error(self):
74 create = AsyncMock(return_value=MagicMock(choices=[]))
75 llm = self._llm(create)
76 with self.assertRaisesRegex(RuntimeError, "choice"):
77 await llm.complete([{"role": "user", "content": "x"}], tools=[])
78
79
80 class BoomLLM:
81 async def complete(self, messages, tools):
82 raise RuntimeError("boom-llm")
83
84
85 class TestLoopErrorBoundary(unittest.IsolatedAsyncioTestCase):
86 async def test_llm_failure_emits_error_and_persists_failed_turn(self):
87 with tempfile.TemporaryDirectory() as tmp:
88 index = SessionIndex(tmp)
89 registry = ToolRegistry([])
90 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), BoomLLM())
91 events = [event async for event in loop.stream_events("hi")]
92 kinds = [event["type"] for event in events]
93 self.assertIn("error", kinds)
94 error_event = next(event for event in events if event["type"] == "error")
95 self.assertIn("boom-llm", error_event["message"])
96 self.assertEqual(events[-2]["type"], "done")
97 self.assertEqual(events[-1]["type"], "session")
98 active = index.active()
99 records = index.get(active["session_id"])["recent_turn_records"]
100 self.assertEqual(records[-1]["status"], "failed")
101
102
103 class TestSessionIndexDurability(unittest.TestCase):
104 def test_corrupt_sessions_file_is_backed_up_not_silently_replaced(self):
105 with tempfile.TemporaryDirectory() as tmp:
106 index = SessionIndex(tmp)
107 index.create(idea="precious work", session_id="keep-me")
108 index.sessions_path.write_text("{ definitely not json", encoding="utf-8")
109 data = index.load()
110 self.assertEqual(data["sessions"], {})
111 backups = list(index.vimax_dir.glob("sessions.json.corrupt-*"))
112 self.assertEqual(len(backups), 1, "corrupt state must be preserved for recovery, not discarded")
113 self.assertIn("definitely not json", backups[0].read_text(encoding="utf-8"))
114
115 def test_save_is_atomic_and_leaves_no_temp_files(self):
116 with tempfile.TemporaryDirectory() as tmp:
117 index = SessionIndex(tmp)
118 index.create(session_id="roundtrip")
119 self.assertEqual(list(index.vimax_dir.glob("*.tmp")), [])
120 self.assertIn("roundtrip", index.load()["sessions"])
121
122 def test_concurrent_creates_do_not_lose_sessions(self):
123 with tempfile.TemporaryDirectory() as tmp:
124 index_a = SessionIndex(tmp)
125 index_b = SessionIndex(tmp)
126
127 def worker(index, tag):
128 for i in range(40):
129 index.create(session_id=f"s-{tag}-{i}")
130
131 threads = [
132 threading.Thread(target=worker, args=(index_a, "a")),
133 threading.Thread(target=worker, args=(index_b, "b")),
134 ]
135 for thread in threads:
136 thread.start()
137 for thread in threads:
138 thread.join()
139 sessions = index_a.load()["sessions"]
140 self.assertEqual(len(sessions), 80, "concurrent read-modify-write must not lose sessions")
141
142
143 class TestBoundedRetryPolicies(unittest.TestCase):
144 CASES = [
145 ("Screenwriter.write_script_based_on_story", Screenwriter.write_script_based_on_story),
146 ("ScriptPlanner.plan_script", ScriptPlanner.plan_script),
147 ("RerankerBgeSiliconapi.__call__", RerankerBgeSiliconapi.__call__),
148 ("ImageGeneratorDoubaoSeedreamYunwuAPI.generate_single_image", ImageGeneratorDoubaoSeedreamYunwuAPI.generate_single_image),
149 ("ImageGeneratorNanobananaGoogleAPI.generate_single_image", ImageGeneratorNanobananaGoogleAPI.generate_single_image),
150 ("ImageGeneratorNanobananaYunwuAPI.generate_single_image", ImageGeneratorNanobananaYunwuAPI.generate_single_image),
151 ]
152
153 def test_every_retry_is_bounded_with_backoff(self):
154 for name, fn in self.CASES:
155 with self.subTest(name=name):
156 retrying = getattr(fn, "retry", None)
157 self.assertIsNotNone(retrying, f"{name} must have a retry policy")
158 self.assertIsNot(retrying.stop, stop_never, f"{name} must not retry forever")
159 self.assertNotIsInstance(retrying.wait, wait_none, f"{name} must back off between attempts")
160
161
162 class _FakeResponse:
163 def __init__(self, payload, status=200):
164 self.payload = payload
165 self.status = status
166
167 async def __aenter__(self):
168 return self
169
170 async def __aexit__(self, exc_type, exc, tb):
171 return False
172
173 async def json(self):
174 return self.payload
175
176
177 class _FakeSession:
178 def __init__(self, scripted):
179 self.scripted = list(scripted)
180 self.calls = 0
181
182 async def __aenter__(self):
183 return self
184
185 async def __aexit__(self, exc_type, exc, tb):
186 return False
187
188 def _next(self):
189 response = self.scripted[min(self.calls, len(self.scripted) - 1)]
190 self.calls += 1
191 return _FakeResponse(*response)
192
193 def post(self, url, **kwargs):
194 return self._next()
195
196 def get(self, url, **kwargs):
197 return self._next()
198
199
200 class TestClientHttpErrors(unittest.IsolatedAsyncioTestCase):
201 async def test_reranker_surfaces_http_error_without_retry(self):
202 session = _FakeSession([
203 ({"message": "invalid api key"}, 401),
204 ({"results": []}, 200),
205 ])
206 reranker = RerankerBgeSiliconapi(api_key="bad", base_url="http://x")
207 with patch("tools.reranker_bge_silicon_api.aiohttp.ClientSession", return_value=session):
208 with self.assertRaisesRegex(RuntimeError, "401"):
209 await reranker(documents=["doc"], query="q", top_n=1)
210 self.assertEqual(session.calls, 1, "4xx must fail fast with the real error, not retry into KeyError")
211
212 async def test_seedream_surfaces_http_error_without_retry(self):
213 session = _FakeSession([
214 ({"error": {"message": "invalid api key"}}, 401),
215 ({"data": [{"url": "http://img"}]}, 200),
216 ])
217 generator = ImageGeneratorDoubaoSeedreamYunwuAPI(api_key="bad")
218 with patch("tools.image_generator_doubao_seedream_yunwu_api.aiohttp.ClientSession", return_value=session):
219 with self.assertRaisesRegex(RuntimeError, "401"):
220 await generator.generate_single_image(prompt="p")
221 self.assertEqual(session.calls, 1)
222
223
224 if __name__ == "__main__":
225 unittest.main()
226
226 lines PYTHON