🐛 Fix proxy caching between requests in aiohttp transport (#17122)

* write a regression test

* impl fix

* add test for host case

* use the host as cache key
This commit is contained in:
Leslie Cheng 2025-12-02 14:37:45 -08:00 • committed by GitHub
parent db6c6eea89
commit de4ff120eb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 137 additions and 73 deletions

View file

@ -82,9 +82,7 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
async def __aiter__(self) -> typing.AsyncIterator[bytes]:
try:
async for chunk in self._aiohttp_response.content.iter_chunked(
self.CHUNK_SIZE
):
async for chunk in self._aiohttp_response.content.iter_chunked(self.CHUNK_SIZE):
yield chunk
except (
aiohttp.ClientPayloadError,
@ -120,16 +118,13 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
class AiohttpTransport(httpx.AsyncBaseTransport):
def __init__(
self, client: Union[ClientSession, Callable[[], ClientSession]]
) -> None:
def __init__(self, client: Union[ClientSession, Callable[[], ClientSession]]) -> None:
self.client = client
#########################################################
# Class variables for proxy settings
#########################################################
self.proxy: Optional[str] = None
self.checked_proxy_env_settings: bool = False
self.proxy_cache: Dict[str, Optional[str]] = {}
async def aclose(self) -> None:
if isinstance(self.client, ClientSession):
@ -184,11 +179,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
current_loop = asyncio.get_running_loop()
# If session is from a different or closed loop, recreate it
if (
session_loop is None
or session_loop != current_loop
or session_loop.is_closed()
):
if session_loop is None or session_loop != current_loop or session_loop.is_closed():
# Close old session to prevent leaks
old_session = self.client
try:
@ -215,7 +206,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
self.client = ClientSession()
return self.client
async def _make_aiohttp_request(
self,
client_session: ClientSession,
@ -226,20 +217,20 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
) -> ClientResponse:
"""
Helper function to make an aiohttp request with the given parameters.
Args:
client_session: The aiohttp ClientSession to use
request: The httpx Request to send
timeout: Timeout settings dict with 'connect', 'read', 'pool' keys
proxy: Optional proxy URL
sni_hostname: Optional SNI hostname for SSL
Returns:
ClientResponse from aiohttp
"""
from aiohttp import ClientTimeout
from yarl import URL as YarlURL
try:
data = request.content
except httpx.RequestNotRead:
@ -262,9 +253,9 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
proxy=proxy,
server_hostname=sni_hostname,
).__aenter__()
return response
async def handle_async_request(
self,
request: httpx.Request,
@ -297,7 +288,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
else:
self.client = ClientSession()
client_session = self.client
# Retry the request with the new session
with map_aiohttp_exceptions():
response = await self._make_aiohttp_request(
@ -317,45 +308,41 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
content=AiohttpResponseStream(response),
request=request,
)
async def _get_proxy_settings(self, request: httpx.Request):
proxy = None
if not (
litellm.disable_aiohttp_trust_env
or str_to_bool(os.getenv("DISABLE_AIOHTTP_TRUST_ENV", "False"))
):
if not (litellm.disable_aiohttp_trust_env or str_to_bool(os.getenv("DISABLE_AIOHTTP_TRUST_ENV", "False"))):
try:
proxy = self._proxy_from_env(request.url)
except Exception as e: # pragma: no cover - best effort
verbose_logger.debug(f"Error reading proxy env: {e}")
return proxy
def _proxy_from_env(self, url: httpx.URL) -> typing.Optional[str]:
"""
Return proxy URL from env for the given request URL
Only check the proxy env settings once, this is a costly operation for CPU % usage
."""
#########################################################
# Check if we've already checked the proxy env settings
#########################################################
if self.checked_proxy_env_settings is True:
return self.proxy
#########################################################
# set self.checked_proxy_env_settings to True
#########################################################
self.checked_proxy_env_settings = True
proxy_cache_key = url.host
if proxy_cache_key in self.proxy_cache:
return self.proxy_cache[proxy_cache_key]
proxies = urllib.request.getproxies()
if urllib.request.proxy_bypass(url.host):
return None
proxy_url = None
else:
proxy = proxies.get(url.scheme) or proxies.get("all")
if proxy and "://" not in proxy:
proxy = f"http://{proxy}"
proxy_url = proxy
proxy = proxies.get(url.scheme) or proxies.get("all")
if proxy and "://" not in proxy:
proxy = f"http://{proxy}"
self.proxy = proxy
return self.proxy
self.proxy_cache[proxy_cache_key] = proxy_url
return proxy_url

View file

@ -1,6 +1,6 @@
import asyncio
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import aiohttp
import aiohttp.client_exceptions
@ -8,14 +8,11 @@ import aiohttp.http_exceptions
import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path
from litellm.llms.custom_httpx.aiohttp_transport import (
AiohttpResponseStream,
LiteLLMAiohttpTransport,
map_aiohttp_exceptions,
)
@ -32,9 +29,7 @@ class MockAiohttpResponse:
):
self.status = status
self.headers = headers or {}
self.content = MockContent(
content_chunks, exception_to_raise, exception_at_chunk
)
self.content = MockContent(content_chunks, exception_to_raise, exception_at_chunk)
async def __aexit__(self, exc_type, exc_val, exc_tb):
pass
@ -74,7 +69,6 @@ async def test_aiohttp_response_stream_normal_flow():
@pytest.mark.asyncio
async def test_transfer_encoding_error_no_httpx_read_error():
"""Test that TransferEncodingError doesn't get converted to httpx.ReadError"""
import logging
# Create a TransferEncodingError wrapped in ClientPayloadError (like in real scenarios)
transfer_error = aiohttp.http_exceptions.TransferEncodingError(
@ -82,9 +76,7 @@ async def test_transfer_encoding_error_no_httpx_read_error():
)
# Wrap it in ClientPayloadError as aiohttp does
client_payload_error = aiohttp.ClientPayloadError(
"Response payload is not completed"
)
client_payload_error = aiohttp.ClientPayloadError("Response payload is not completed")
client_payload_error.__cause__ = transfer_error
mock_response = MockAiohttpResponse(
@ -111,9 +103,7 @@ async def test_transfer_encoding_error_no_httpx_read_error():
async def test_client_payload_error_graceful_handling():
"""Test that ClientPayloadError is handled gracefully without stacktrace"""
# Create a ClientPayloadError directly
client_error = aiohttp.client_exceptions.ClientPayloadError(
"Response payload is not completed"
)
client_error = aiohttp.client_exceptions.ClientPayloadError("Response payload is not completed")
mock_response = MockAiohttpResponse(
content_chunks=[b"data1", b"data2", b"data3"],
@ -181,7 +171,6 @@ async def test_timeout_exception_gets_mapped():
@pytest.mark.asyncio
async def test_handle_async_request_uses_env_proxy(monkeypatch):
"""Aiohttp transport should honor HTTP(S)_PROXY env vars"""
import asyncio
proxy_url = "http://proxy.local:3128"
monkeypatch.setenv("HTTP_PROXY", proxy_url)
monkeypatch.setenv("http_proxy", proxy_url)
@ -200,7 +189,7 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch):
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None
def request(self, *args, **kwargs):
captured["proxy"] = kwargs.get("proxy")
@ -231,30 +220,118 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch):
assert captured["proxy"] == proxy_url
@pytest.mark.asyncio
async def test_handle_async_request_uses_env_proxy_per_url(monkeypatch):
"""Aiohttp transport should honor HTTP(S)_PROXY env vars unless NO_PROXY matches"""
proxy_url = "http://proxy.local:3128"
monkeypatch.setenv("NO_PROXY", "example.com")
monkeypatch.setenv("HTTP_PROXY", proxy_url)
monkeypatch.setenv("http_proxy", proxy_url)
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
monkeypatch.setenv("https_proxy", proxy_url)
monkeypatch.delenv("DISABLE_AIOHTTP_TRUST_ENV", raising=False)
request_count = 0
proxied_count = 0
class FakeSession:
def __init__(self):
self.closed = False
try:
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None
def request(self, *args, **kwargs):
nonlocal request_count
nonlocal proxied_count
request_count += 1
if kwargs.get("proxy") is not None:
proxied_count += 1
class Resp:
status = 200
headers = {}
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
pass
@property
def content(self):
class C:
async def iter_chunked(self, size):
yield b""
return C()
return Resp()
transport = LiteLLMAiohttpTransport(client=lambda: FakeSession()) # type: ignore
request = httpx.Request("GET", "http://example.com")
await transport.handle_async_request(request)
request = httpx.Request("GET", "http://foo.com")
await transport.handle_async_request(request)
assert request_count == 2
assert proxied_count == 1
@pytest.mark.asyncio
async def test_handle_async_request_proxy_cache_per_host(monkeypatch):
"""Aiohttp transport should only cache a proxy per host rather than full URL"""
proxy_url = "http://proxy.local:3128"
monkeypatch.setenv("NO_PROXY", "example.com")
monkeypatch.setenv("HTTP_PROXY", proxy_url)
monkeypatch.setenv("http_proxy", proxy_url)
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
monkeypatch.setenv("https_proxy", proxy_url)
monkeypatch.delenv("DISABLE_AIOHTTP_TRUST_ENV", raising=False)
def factory():
return _make_mock_session()
transport = LiteLLMAiohttpTransport(client=factory) # type: ignore
request = httpx.Request("GET", "http://foo.com/path1")
await transport.handle_async_request(request)
request = httpx.Request("GET", "http://foo.com/path2")
await transport.handle_async_request(request)
assert len(transport.proxy_cache) == 1
def _make_mock_response(should_fail=False, fail_count={"count": 0}):
"""Helper to create a mock aiohttp response"""
class MockResp:
status = 200
headers = {}
async def __aenter__(self):
if should_fail and fail_count["count"] < 1:
fail_count["count"] += 1
raise RuntimeError("Session is closed")
return self
async def __aexit__(self, *args):
pass
@property
def content(self):
class C:
async def iter_chunked(self, size):
yield b"test"
return C()
return MockResp()
@pytest.mark.asyncio
async def test_handle_async_request_total_timeout_triggers():
"""
@ -298,10 +375,10 @@ async def test_handle_async_request_total_timeout_triggers():
await transport.aclose()
await runner.cleanup()
def _make_mock_session(closed=False):
"""Helper to create a mock aiohttp session"""
import asyncio
class MockSession:
def __init__(self):
self.closed = closed
@ -309,10 +386,10 @@ def _make_mock_session(closed=False):
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None
def request(self, *args, **kwargs):
return _make_mock_response()
return MockSession()
@ -320,14 +397,14 @@ def _make_mock_session(closed=False):
async def test_handle_closed_session_before_request():
"""Test that closed sessions are detected and recreated"""
counts = {"sessions": 0}
def factory():
counts["sessions"] += 1
return _make_mock_session(closed=counts["sessions"] == 1)
transport = LiteLLMAiohttpTransport(client=factory) # type: ignore
response = await transport.handle_async_request(httpx.Request("GET", "http://example.com"))
assert counts["sessions"] == 2 # Created 2 sessions: closed one, then open one
assert response.status_code == 200
@ -337,7 +414,7 @@ async def test_handle_session_closed_during_request():
"""Test that sessions closed during request are handled with retry"""
counts = {"sessions": 0, "requests": 0}
fail_count = {"count": 0}
class MockSession:
def __init__(self):
self.closed = False
@ -345,18 +422,18 @@ async def test_handle_session_closed_during_request():
self._loop = __import__("asyncio").get_running_loop()
except RuntimeError:
self._loop = None
def request(self, *args, **kwargs):
counts["requests"] += 1
return _make_mock_response(should_fail=True, fail_count=fail_count)
def factory():
counts["sessions"] += 1
return MockSession()
transport = LiteLLMAiohttpTransport(client=factory) # type: ignore
response = await transport.handle_async_request(httpx.Request("GET", "http://example.com"))
assert counts["requests"] == 2 # First request failed, second succeeded
assert counts["sessions"] == 2 # Created 2 sessions for retry
assert response.status_code == 200