mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
fix(mavvrik): use get_async_httpx_client instead of httpx.AsyncClient per request
BerriAI enforce_async_clients check requires using the shared cached client from get_async_httpx_client() instead of creating httpx.AsyncClient() per request. Creating per-request clients adds +500ms latency overhead. _http.py now calls get_async_httpx_client(LoggingCallback).client to get the shared httpx.AsyncClient from LiteLLM's in-memory client cache, then calls .request() on it directly (AsyncHTTPHandler only has verb-specific wrappers, not a generic request() method). Test mocks updated to patch get_async_httpx_client instead of httpx.AsyncClient. Co-Authored-By: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
4e0b1907d4
commit
376d9a59ee
2 changed files with 123 additions and 80 deletions
|
|
@ -10,6 +10,9 @@ Retry behaviour:
|
|||
- 5xx responses and network errors: retry up to MAX_RETRIES times
|
||||
- 4xx responses: returned immediately (client-side error, no retry)
|
||||
- Backoff: RETRY_BACKOFF_BASE * 2^attempt seconds between retries
|
||||
|
||||
Uses get_async_httpx_client() from litellm's shared client cache — avoids
|
||||
creating a new AsyncClient per request which adds +500ms latency overhead.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
|
@ -18,6 +21,8 @@ from typing import Any, Dict, Optional
|
|||
import httpx
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
_MAX_RETRIES = 3
|
||||
_RETRY_BACKOFF_BASE = 1.0 # seconds; doubles each retry
|
||||
|
|
@ -55,41 +60,46 @@ async def http_request(
|
|||
"""
|
||||
tag = label or method
|
||||
last_exc: Exception = RuntimeError("unknown error")
|
||||
# Use the shared cached client — avoids creating a new AsyncClient per request
|
||||
# which adds +500ms latency overhead. Access .client for the generic request() method
|
||||
# since AsyncHTTPHandler only exposes verb-specific wrappers (get/post/put/patch).
|
||||
http = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
).client
|
||||
|
||||
async with httpx.AsyncClient() as http:
|
||||
for attempt in range(_MAX_RETRIES):
|
||||
try:
|
||||
resp = await http.request(
|
||||
method,
|
||||
url,
|
||||
headers=headers,
|
||||
json=json,
|
||||
params=params,
|
||||
content=content,
|
||||
timeout=timeout,
|
||||
)
|
||||
for attempt in range(_MAX_RETRIES):
|
||||
try:
|
||||
resp = await http.request(
|
||||
method=method,
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=json,
|
||||
params=params,
|
||||
content=content,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
if resp.status_code < 500:
|
||||
return resp # success or 4xx — return immediately, no retry
|
||||
if resp.status_code < 500:
|
||||
return resp # success or 4xx — return immediately, no retry
|
||||
|
||||
last_exc = RuntimeError(
|
||||
f"{tag} failed: {resp.status_code} {resp.text[:200]}"
|
||||
)
|
||||
last_exc = RuntimeError(
|
||||
f"{tag} failed: {resp.status_code} {resp.text[:200]}"
|
||||
)
|
||||
|
||||
except httpx.RequestError as exc:
|
||||
last_exc = exc
|
||||
except httpx.RequestError as exc:
|
||||
last_exc = exc
|
||||
|
||||
if attempt < _MAX_RETRIES - 1:
|
||||
wait = _RETRY_BACKOFF_BASE * (2**attempt)
|
||||
verbose_proxy_logger.warning(
|
||||
"mavvrik: %s attempt %d/%d failed, retrying in %.1fs: %s",
|
||||
tag,
|
||||
attempt + 1,
|
||||
_MAX_RETRIES,
|
||||
wait,
|
||||
last_exc,
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
if attempt < _MAX_RETRIES - 1:
|
||||
wait = _RETRY_BACKOFF_BASE * (2**attempt)
|
||||
verbose_proxy_logger.warning(
|
||||
"mavvrik: %s attempt %d/%d failed, retrying in %.1fs: %s",
|
||||
tag,
|
||||
attempt + 1,
|
||||
_MAX_RETRIES,
|
||||
wait,
|
||||
last_exc,
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
|
||||
raise RuntimeError(
|
||||
f"mavvrik: {tag} failed after {_MAX_RETRIES} attempts: {last_exc}"
|
||||
|
|
|
|||
|
|
@ -19,17 +19,17 @@ def _mock_response(status_code: int, text: str = "") -> MagicMock:
|
|||
return resp
|
||||
|
||||
|
||||
def _mock_http(return_value=None, side_effect=None):
|
||||
"""Return a patched httpx.AsyncClient context manager."""
|
||||
http = MagicMock()
|
||||
def _mock_shared_client(return_value=None, side_effect=None):
|
||||
"""Mock get_async_httpx_client returning a handler whose .client.request is stubbed."""
|
||||
inner_client = MagicMock()
|
||||
if side_effect:
|
||||
http.request = AsyncMock(side_effect=side_effect)
|
||||
inner_client.request = AsyncMock(side_effect=side_effect)
|
||||
else:
|
||||
http.request = AsyncMock(return_value=return_value)
|
||||
ctx = MagicMock()
|
||||
ctx.__aenter__ = AsyncMock(return_value=http)
|
||||
ctx.__aexit__ = AsyncMock(return_value=False)
|
||||
return ctx, http
|
||||
inner_client.request = AsyncMock(return_value=return_value)
|
||||
|
||||
handler = MagicMock()
|
||||
handler.client = inner_client
|
||||
return handler, inner_client
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -40,17 +40,26 @@ def _mock_http(return_value=None, side_effect=None):
|
|||
class TestHttpRequestSuccess:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_response_on_2xx(self):
|
||||
ctx, _ = _mock_http(return_value=_mock_response(200))
|
||||
with patch("httpx.AsyncClient", return_value=ctx):
|
||||
handler, _ = _mock_shared_client(return_value=_mock_response(200))
|
||||
with patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
):
|
||||
resp = await http_request("GET", "https://example.com")
|
||||
assert resp.status_code == 200
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_4xx_without_retry(self):
|
||||
ctx, http = _mock_http(return_value=_mock_response(401, "Unauthorized"))
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep", new_callable=AsyncMock
|
||||
) as mock_sleep:
|
||||
handler, http = _mock_shared_client(
|
||||
return_value=_mock_response(401, "Unauthorized")
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep,
|
||||
):
|
||||
resp = await http_request("GET", "https://example.com")
|
||||
assert resp.status_code == 401
|
||||
assert http.request.call_count == 1
|
||||
|
|
@ -64,13 +73,13 @@ class TestHttpRequestSuccess:
|
|||
captured.append(headers)
|
||||
return _mock_response(200)
|
||||
|
||||
ctx = MagicMock()
|
||||
http = MagicMock()
|
||||
http.request = fake_request
|
||||
ctx.__aenter__ = AsyncMock(return_value=http)
|
||||
ctx.__aexit__ = AsyncMock(return_value=False)
|
||||
handler = MagicMock()
|
||||
handler.client.request = fake_request
|
||||
|
||||
with patch("httpx.AsyncClient", return_value=ctx):
|
||||
with patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
):
|
||||
await http_request(
|
||||
"POST", "https://example.com", headers={"x-api-key": "secret"}
|
||||
)
|
||||
|
|
@ -85,13 +94,13 @@ class TestHttpRequestSuccess:
|
|||
captured.append({"json": json, "params": params})
|
||||
return _mock_response(200)
|
||||
|
||||
ctx = MagicMock()
|
||||
http = MagicMock()
|
||||
http.request = fake_request
|
||||
ctx.__aenter__ = AsyncMock(return_value=http)
|
||||
ctx.__aexit__ = AsyncMock(return_value=False)
|
||||
handler = MagicMock()
|
||||
handler.client.request = fake_request
|
||||
|
||||
with patch("httpx.AsyncClient", return_value=ctx):
|
||||
with patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
):
|
||||
await http_request(
|
||||
"GET", "https://example.com", json={"key": "val"}, params={"q": "1"}
|
||||
)
|
||||
|
|
@ -107,13 +116,13 @@ class TestHttpRequestSuccess:
|
|||
captured.append(content)
|
||||
return _mock_response(201)
|
||||
|
||||
ctx = MagicMock()
|
||||
http = MagicMock()
|
||||
http.request = fake_request
|
||||
ctx.__aenter__ = AsyncMock(return_value=http)
|
||||
ctx.__aexit__ = AsyncMock(return_value=False)
|
||||
handler = MagicMock()
|
||||
handler.client.request = fake_request
|
||||
|
||||
with patch("httpx.AsyncClient", return_value=ctx):
|
||||
with patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
):
|
||||
await http_request("PUT", "https://example.com", content=b"gzip-data")
|
||||
|
||||
assert captured[0] == b"gzip-data"
|
||||
|
|
@ -127,9 +136,15 @@ class TestHttpRequestSuccess:
|
|||
class TestHttpRequestRetry:
|
||||
@pytest.mark.asyncio
|
||||
async def test_retries_on_5xx_then_raises(self):
|
||||
ctx, http = _mock_http(return_value=_mock_response(503, "unavailable"))
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep", new_callable=AsyncMock
|
||||
handler, http = _mock_shared_client(
|
||||
return_value=_mock_response(503, "unavailable")
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch("asyncio.sleep", new_callable=AsyncMock),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="failed after"):
|
||||
await http_request("GET", "https://example.com")
|
||||
|
|
@ -137,9 +152,13 @@ class TestHttpRequestRetry:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retries_on_network_error_then_raises(self):
|
||||
ctx, http = _mock_http(side_effect=httpx.ConnectError("timeout"))
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep", new_callable=AsyncMock
|
||||
handler, http = _mock_shared_client(side_effect=httpx.ConnectError("timeout"))
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch("asyncio.sleep", new_callable=AsyncMock),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="failed after"):
|
||||
await http_request("GET", "https://example.com")
|
||||
|
|
@ -149,10 +168,14 @@ class TestHttpRequestRetry:
|
|||
async def test_succeeds_on_second_attempt(self):
|
||||
fail = _mock_response(503, "err")
|
||||
ok = _mock_response(200)
|
||||
ctx, http = _mock_http()
|
||||
handler, http = _mock_shared_client()
|
||||
http.request = AsyncMock(side_effect=[fail, ok])
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep", new_callable=AsyncMock
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch("asyncio.sleep", new_callable=AsyncMock),
|
||||
):
|
||||
resp = await http_request("GET", "https://example.com")
|
||||
assert resp.status_code == 200
|
||||
|
|
@ -160,12 +183,18 @@ class TestHttpRequestRetry:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_uses_exponential_backoff(self):
|
||||
ctx, http = _mock_http(return_value=_mock_response(503, "err"))
|
||||
handler, http = _mock_shared_client(return_value=_mock_response(503, "err"))
|
||||
sleep_calls = []
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=lambda s: sleep_calls.append(s),
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch(
|
||||
"asyncio.sleep",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=lambda s: sleep_calls.append(s),
|
||||
),
|
||||
):
|
||||
with pytest.raises(RuntimeError):
|
||||
await http_request("GET", "https://example.com")
|
||||
|
|
@ -174,9 +203,13 @@ class TestHttpRequestRetry:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_message_contains_label(self):
|
||||
ctx, _ = _mock_http(return_value=_mock_response(503, "err"))
|
||||
with patch("httpx.AsyncClient", return_value=ctx), patch(
|
||||
"asyncio.sleep", new_callable=AsyncMock
|
||||
handler, _ = _mock_shared_client(return_value=_mock_response(503, "err"))
|
||||
with (
|
||||
patch(
|
||||
"litellm.integrations.mavvrik._http.get_async_httpx_client",
|
||||
return_value=handler,
|
||||
),
|
||||
patch("asyncio.sleep", new_callable=AsyncMock),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="initiate"):
|
||||
await http_request("POST", "https://example.com", label="initiate")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue