| 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 |