Skip to content

Commit 2f7012f

Browse files
test: add retry behavior tests for AI provider
1 parent 9b60265 commit 2f7012f

1 file changed

Lines changed: 77 additions & 9 deletions

File tree

backend/tests/test_ai_provider.py

Lines changed: 77 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121
import pytest
2222

2323

24-
2524
def _make_llm_response(text: str) -> MagicMock:
2625
"""Return a fake httpx.Response with an OpenAI-compatible JSON body."""
2726
resp = MagicMock()
@@ -35,15 +34,15 @@ def _make_llm_response(text: str) -> MagicMock:
3534
def _make_error_response(status_code: int = 500) -> MagicMock:
3635
"""Return a fake httpx.Response whose raise_for_status() raises."""
3736
resp = MagicMock()
38-
resp.status_code = status_code
39-
37+
resp.status_code = status_code
38+
4039
mock_response = MagicMock()
41-
mock_response.status_code = status_code
42-
40+
mock_response.status_code = status_code
41+
4342
resp.raise_for_status.side_effect = httpx.HTTPStatusError(
4443
message=f"HTTP {status_code}",
4544
request=MagicMock(),
46-
response=mock_response,
45+
response=mock_response,
4746
)
4847
return resp
4948

@@ -64,12 +63,11 @@ def _reload_module(env: dict):
6463
"""Reload ai_provider so module-level env vars are re-evaluated."""
6564
with patch.dict(os.environ, env, clear=False):
6665
import app.services.ai_provider as mod
66+
6767
importlib.reload(mod)
6868
return mod
6969

7070

71-
72-
7371
@pytest.fixture()
7472
def enabled_env():
7573
"""Env vars that enable the LLM provider."""
@@ -199,6 +197,7 @@ async def test_whitespace_only_response_stripped_to_empty(self, enabled_env):
199197
patcher.stop()
200198
assert result == ""
201199

200+
202201
class TestCallLlmPayload:
203202

204203
@pytest.mark.asyncio
@@ -267,6 +266,7 @@ async def test_calls_correct_ollama_endpoint(self, ollama_env):
267266
url_called = mock_client.post.call_args[0][0]
268267
assert url_called == "http://localhost:11434/v1/chat/completions"
269268

269+
270270
class TestCallLlmErrors:
271271

272272
@pytest.mark.asyncio
@@ -316,7 +316,9 @@ async def test_returns_none_on_timeout(self, enabled_env):
316316
async def test_returns_none_on_network_error(self, enabled_env):
317317
mod = _reload_module(enabled_env)
318318
mock_client = AsyncMock()
319-
mock_client.post = AsyncMock(side_effect=httpx.ConnectError("connection refused"))
319+
mock_client.post = AsyncMock(
320+
side_effect=httpx.ConnectError("connection refused")
321+
)
320322

321323
with patch("app.services.ai_provider.httpx.AsyncClient") as MockCls:
322324
MockCls.return_value.__aenter__ = AsyncMock(return_value=mock_client)
@@ -356,3 +358,69 @@ async def test_returns_none_on_empty_choices(self, enabled_env):
356358
finally:
357359
patcher.stop()
358360
assert result is None
361+
362+
@pytest.mark.asyncio
363+
async def test_retries_on_500_error(self,enabled_env):
364+
mod = _reload_module(enabled_env)
365+
366+
mock_client = AsyncMock()
367+
mock_client.post = AsyncMock(
368+
side_effect=httpx.HTTPStatusError(
369+
"HTTP 500",
370+
request=MagicMock(),
371+
response=MagicMock(status_code=500),
372+
)
373+
)
374+
375+
with patch("app.services.ai_provider.httpx.AsyncClient") as MockCls:
376+
MockCls.return_value.__aenter__ = AsyncMock(return_value=mock_client)
377+
MockCls.return_value.__aexit__ = AsyncMock(return_value=False)
378+
379+
with patch("app.services.ai_provider.asyncio.sleep", new=AsyncMock()):
380+
result = await mod.call_llm("sys", "usr")
381+
382+
assert result is None
383+
assert mock_client.post.call_count == mod.LLM_MAX_RETRIES + 1
384+
@pytest.mark.asyncio
385+
async def test_retries_on_429_error(self,enabled_env):
386+
mod=_reload_module(enabled_env)
387+
388+
mock_client=AsyncMock()
389+
mock_client.post=AsyncMock(
390+
side_effect=httpx.HTTPStatusError(
391+
"HTTP 429",
392+
request=MagicMock(),
393+
response=MagicMock(status_code=429),
394+
)
395+
)
396+
397+
with patch("app.services.ai_provider.httpx.AsyncClient") as MockCls:
398+
MockCls.return_value.__aenter__=AsyncMock(return_value=mock_client)
399+
MockCls.return_value.__aexit__=AsyncMock(return_value=False)
400+
401+
with patch("app.services.ai_provider.asyncio.sleep", new=AsyncMock()):
402+
result=await mod.call_llm("sys","usr")
403+
404+
assert result is None
405+
assert mock_client.post.call_count == mod.LLM_MAX_RETRIES + 1
406+
407+
@pytest.mark.asyncio
408+
async def test_retries_on_401_error(self,enabled_env):
409+
mod=_reload_module(enabled_env)
410+
411+
mock_client=AsyncMock()
412+
mock_client.post=AsyncMock(
413+
side_effect=httpx.HTTPStatusError(
414+
"HTTP 401",
415+
request=MagicMock(),
416+
response=MagicMock(status_code=401),
417+
)
418+
)
419+
420+
with patch("app.services.ai_provider.httpx.AsyncClient") as MockCls:
421+
MockCls.return_value.__aenter__=AsyncMock(return_value=mock_client)
422+
MockCls.return_value.__aexit__=AsyncMock(return_value=False)
423+
424+
result = await mod.call_llm("sys", "usr")
425+
assert result is None
426+
assert mock_client.post.call_count == 1

0 commit comments

Comments
 (0)