mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
a94fefe580
commit
ce36a2a9f1
2 changed files with 167 additions and 22 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue