diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index ab69ea1f8c3..14fc77fa91f 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -152,6 +152,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport): # If we don't have a client or it's not a ClientSession, create one if not isinstance(self.client, ClientSession): + if hasattr(self, "_client_factory") and callable(self._client_factory): + self.client = self._client_factory() + else: + self.client = ClientSession() + # Don't return yet - check if the newly created session is valid + + # Check if the session itself is closed + if self.client.closed: + verbose_logger.debug("Session is closed, creating new session") + # Create a new session if hasattr(self, "_client_factory") and callable(self._client_factory): self.client = self._client_factory() else: @@ -209,28 +219,66 @@ class LiteLLMAiohttpTransport(AiohttpTransport): # Resolve proxy settings from environment variables proxy = await self._get_proxy_settings(request) - with map_aiohttp_exceptions(): - try: - data = request.content - except httpx.RequestNotRead: - data = request.stream # type: ignore - request.headers.pop("transfer-encoding", None) # handled by aiohttp + try: + with map_aiohttp_exceptions(): + try: + data = request.content + except httpx.RequestNotRead: + data = request.stream # type: ignore + request.headers.pop("transfer-encoding", None) # handled by aiohttp - response = await client_session.request( - method=request.method, - url=YarlURL(str(request.url), encoded=True), - headers=request.headers, - data=data, - allow_redirects=False, - auto_decompress=False, - timeout=ClientTimeout( - sock_connect=timeout.get("connect"), - sock_read=timeout.get("read"), - connect=timeout.get("pool"), - ), - proxy=proxy, - server_hostname=sni_hostname, - ).__aenter__() + response = await client_session.request( + method=request.method, + url=YarlURL(str(request.url), encoded=True), + headers=request.headers, + data=data, + allow_redirects=False, + auto_decompress=False, + timeout=ClientTimeout( + sock_connect=timeout.get("connect"), + sock_read=timeout.get("read"), + connect=timeout.get("pool"), + ), + proxy=proxy, + server_hostname=sni_hostname, + ).__aenter__() + except RuntimeError as e: + # Handle the case where session was closed between our check and actual use + if "Session is closed" in str(e): + verbose_logger.debug(f"Session closed during request, retrying with new session: {e}") + # Force creation of a new session + if hasattr(self, "_client_factory") and callable(self._client_factory): + self.client = self._client_factory() + else: + self.client = ClientSession() + client_session = self.client + + # Retry the request with the new session + with map_aiohttp_exceptions(): + try: + data = request.content + except httpx.RequestNotRead: + data = request.stream # type: ignore + request.headers.pop("transfer-encoding", None) # handled by aiohttp + + response = await client_session.request( + method=request.method, + url=YarlURL(str(request.url), encoded=True), + headers=request.headers, + data=data, + allow_redirects=False, + auto_decompress=False, + timeout=ClientTimeout( + sock_connect=timeout.get("connect"), + sock_read=timeout.get("read"), + connect=timeout.get("pool"), + ), + proxy=proxy, + server_hostname=sni_hostname, + ).__aenter__() + else: + # Re-raise if it's a different RuntimeError + raise return httpx.Response( status_code=response.status, diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index f296f076178..4059b5b2e74 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -181,6 +181,7 @@ 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) @@ -191,6 +192,13 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch): captured = {} 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): captured["proxy"] = kwargs.get("proxy") @@ -214,8 +222,97 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch): return Resp() - transport = LiteLLMAiohttpTransport(client=lambda: FakeSession()) + transport = LiteLLMAiohttpTransport(client=lambda: FakeSession()) # type: ignore request = httpx.Request("GET", "http://example.com") await transport.handle_async_request(request) assert captured["proxy"] == proxy_url + + +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() + + +def _make_mock_session(closed=False): + """Helper to create a mock aiohttp session""" + import asyncio + + class MockSession: + def __init__(self): + self.closed = closed + try: + self._loop = asyncio.get_running_loop() + except RuntimeError: + self._loop = None + + def request(self, *args, **kwargs): + return _make_mock_response() + + return MockSession() + + +@pytest.mark.asyncio +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 + + +@pytest.mark.asyncio +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 + try: + 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