fix: handle closed aiohttp sessions with detection and retry

Fixes RuntimeError "Session is closed" by:
- Checking session.closed before use and recreating if needed
- Catching RuntimeError during requests and retrying with new session
- Validating newly created sessions aren't already closed

Adds tests for both proactive detection and reactive retry scenarios.
This commit is contained in:
AlexsanderHamir 2025-10-10 15:12:43 -07:00
parent a94fefe580
commit ce36a2a9f1
2 changed files with 167 additions and 22 deletions

View file

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

View file

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