fix(aiohttp): reliably close recycled client sessions

This commit is contained in:
Alexey 2026-03-20 20:56:15 +03:00
parent d7c419bfee
commit 2627bd037f
3 changed files with 247 additions and 28 deletions

View file

@ -4,7 +4,7 @@ import os
import ssl
import typing
import urllib.request
from typing import Any, Callable, Dict, Optional, Union
from typing import Any, Callable, ClassVar, Dict, Optional, Union
import aiohttp
import aiohttp.client_exceptions
@ -149,6 +149,8 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
Credit to: https://github.com/karpetrosyan/httpx-aiohttp for this implementation
"""
_background_close_tasks: ClassVar[set[asyncio.Task[Any]]] = set()
def __init__(
self,
client: Union[ClientSession, Callable[[], ClientSession]],
@ -162,6 +164,41 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
if callable(client):
self._client_factory = client
@classmethod
def _discard_background_close_task(cls, task: asyncio.Task[Any]) -> None:
cls._background_close_tasks.discard(task)
with contextlib.suppress(asyncio.CancelledError):
close_error = task.exception()
if close_error is not None:
verbose_logger.debug(
"Error closing old session in background task: %s",
close_error,
)
@classmethod
def _prune_background_close_tasks(cls) -> None:
for task in tuple(cls._background_close_tasks):
if task.done():
cls._discard_background_close_task(task)
def _schedule_session_close(self, session: ClientSession) -> None:
if session.closed:
return
cls = type(self)
cls._prune_background_close_tasks()
close_coro = session.close()
try:
task = asyncio.create_task(close_coro)
except RuntimeError:
close_coro.close()
verbose_logger.debug("Old session from different loop, relying on GC")
return
cls._background_close_tasks.add(task)
task.add_done_callback(cls._discard_background_close_task)
def _get_valid_client_session(self) -> ClientSession:
"""
Helper to get a valid ClientSession for the current event loop.
@ -203,14 +240,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
# Close old session to prevent leaks
old_session = self.client
try:
if not old_session.closed:
try:
asyncio.create_task(old_session.close())
except RuntimeError:
# Different event loop - can't schedule task, rely on GC
verbose_logger.debug(
"Old session from different loop, relying on GC"
)
self._schedule_session_close(old_session)
except Exception as e:
verbose_logger.debug(f"Error closing old session: {e}")
@ -220,12 +250,21 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
else:
self.client = ClientSession()
except (RuntimeError, AttributeError):
except (RuntimeError, AttributeError) as e:
# If we can't check the loop or session is invalid, recreate it
old_session = self.client
try:
self._schedule_session_close(old_session)
except Exception as close_error:
verbose_logger.debug(f"Error closing old session: {close_error}")
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
verbose_logger.debug(
"Error checking session loop, created new session: %s",
e,
)
return self.client
@ -319,6 +358,9 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
f"Session closed during request, retrying with new session: {e}"
)
# Force creation of a new session
old_session = self.client
if isinstance(old_session, ClientSession):
self._schedule_session_close(old_session)
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:

View file

@ -10,6 +10,7 @@ import pytest
sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path
from litellm.llms.custom_httpx import aiohttp_transport as aiohttp_transport_module
from litellm.llms.custom_httpx.aiohttp_transport import (
AiohttpResponseStream,
AiohttpTransport,
@ -517,30 +518,172 @@ async def test_handle_closed_session_before_request():
@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}
async def test_schedule_session_close_handles_missing_running_loop(monkeypatch):
"""If create_task cannot run, the close coroutine should still be explicitly closed."""
coro_closed = {"value": False}
class MockCloseCoro:
def close(self):
coro_closed["value"] = True
class MockSession:
def __init__(self):
self.closed = False
try:
self._loop = __import__("asyncio").get_running_loop()
except RuntimeError:
self._loop = None
closed = False
def request(self, *args, **kwargs):
counts["requests"] += 1
return _make_mock_response(should_fail=True, fail_count=fail_count)
def close(self):
return MockCloseCoro()
session = MockSession()
transport = LiteLLMAiohttpTransport(client=session)
LiteLLMAiohttpTransport._background_close_tasks.clear()
def raise_runtime_error(_coro):
raise RuntimeError("no running event loop")
monkeypatch.setattr(asyncio, "create_task", raise_runtime_error)
try:
transport._schedule_session_close(session)
assert coro_closed["value"] is True
finally:
LiteLLMAiohttpTransport._background_close_tasks.clear()
@pytest.mark.asyncio
async def test_discard_background_close_task_logs_failed_close(monkeypatch):
"""Failed background close tasks should be removed and logged."""
debug_messages = []
def capture_debug(message, *args, **kwargs):
if args:
debug_messages.append(message % args)
else:
debug_messages.append(message)
async def fail_close():
raise RuntimeError("close failed")
monkeypatch.setattr(aiohttp_transport_module.verbose_logger, "debug", capture_debug)
task = asyncio.create_task(fail_close())
LiteLLMAiohttpTransport._background_close_tasks.clear()
LiteLLMAiohttpTransport._background_close_tasks.add(task)
try:
await asyncio.gather(task, return_exceptions=True)
LiteLLMAiohttpTransport._discard_background_close_task(task)
assert task not in LiteLLMAiohttpTransport._background_close_tasks
assert any("Error closing old session in background task: close failed" in msg for msg in debug_messages)
finally:
LiteLLMAiohttpTransport._background_close_tasks.clear()
@pytest.mark.asyncio
async def test_schedule_session_close_prunes_stale_done_tasks():
"""Scheduling a new close should prune completed stale tasks from the tracking set."""
async def completed_task():
return None
stale_task = asyncio.create_task(completed_task())
await stale_task
session = aiohttp.ClientSession()
transport = LiteLLMAiohttpTransport(client=session)
LiteLLMAiohttpTransport._background_close_tasks.clear()
LiteLLMAiohttpTransport._background_close_tasks.add(stale_task)
try:
transport._schedule_session_close(session)
pending_tasks = list(LiteLLMAiohttpTransport._background_close_tasks)
if pending_tasks:
await asyncio.gather(*pending_tasks)
assert stale_task not in LiteLLMAiohttpTransport._background_close_tasks
assert LiteLLMAiohttpTransport._background_close_tasks == set()
assert session.closed is True
finally:
LiteLLMAiohttpTransport._background_close_tasks.clear()
if not session.closed:
await session.close()
@pytest.mark.asyncio
async def test_get_valid_client_session_closes_old_session_on_loop_check_runtime_error(monkeypatch):
"""Loop-check exceptions should still attempt to close the old session before recreation."""
old_session = aiohttp.ClientSession()
new_session = aiohttp.ClientSession()
scheduled_sessions = []
transport = LiteLLMAiohttpTransport(client=lambda: new_session)
transport.client = old_session
def mock_schedule(session):
scheduled_sessions.append(session)
monkeypatch.setattr(
transport,
"_schedule_session_close",
mock_schedule,
)
def raise_runtime_error():
raise RuntimeError("loop unavailable")
monkeypatch.setattr(asyncio, "get_running_loop", raise_runtime_error)
try:
session = transport._get_valid_client_session()
assert session is new_session
assert scheduled_sessions == [old_session]
finally:
await old_session.close()
await new_session.close()
@pytest.mark.asyncio
async def test_handle_session_closed_during_request():
"""Test that sessions closed during request are handled with retry"""
counts = {"sessions": 0}
scheduled_sessions = []
def factory():
counts["sessions"] += 1
return MockSession()
return aiohttp.ClientSession()
transport = LiteLLMAiohttpTransport(client=factory) # type: ignore
response = await transport.handle_async_request(httpx.Request("GET", "http://example.com"))
original_schedule = transport._schedule_session_close
call_count = {"value": 0}
assert counts["requests"] == 2 # First request failed, second succeeded
assert counts["sessions"] == 2 # Created 2 sessions for retry
assert response.status_code == 200
def tracked_schedule(session):
scheduled_sessions.append(session)
return original_schedule(session)
transport._schedule_session_close = tracked_schedule # type: ignore[method-assign]
async def tracked_make_request(*args, **kwargs):
call_count["value"] += 1
if call_count["value"] == 1:
raise RuntimeError("Session is closed")
return await _make_mock_response().__aenter__()
transport._make_aiohttp_request = tracked_make_request # type: ignore[method-assign]
try:
response = await transport.handle_async_request(httpx.Request("GET", "http://example.com"))
assert call_count["value"] == 2 # First request failed, second succeeded
assert counts["sessions"] == 2 # Created 2 sessions for retry
assert response.status_code == 200
assert len(scheduled_sessions) == 1
assert scheduled_sessions[0].closed is False
finally:
if isinstance(transport.client, aiohttp.ClientSession) and not transport.client.closed:
await transport.client.close()

View file

@ -7,6 +7,7 @@ import os
import sys
from unittest.mock import MagicMock, patch
import aiohttp
import anyio
import httpx
import pytest
@ -65,6 +66,39 @@ async def test_aiohttp_transport_response_uses_stream_not_content():
assert isinstance(response.stream, AiohttpResponseStream)
@pytest.mark.asyncio
async def test_get_valid_client_session_closes_old_session_on_loop_mismatch():
"""Loop-mismatch recreation should close the old session in the background."""
class FakeLoop:
def __init__(self, closed: bool = False):
self._closed = closed
def is_closed(self):
return self._closed
old_session = aiohttp.ClientSession()
old_session._loop = FakeLoop() # type: ignore[attr-defined]
new_session = aiohttp.ClientSession()
transport = LiteLLMAiohttpTransport(client=lambda: new_session)
transport.client = old_session
LiteLLMAiohttpTransport._background_close_tasks.clear()
try:
session = transport._get_valid_client_session()
pending_tasks = list(LiteLLMAiohttpTransport._background_close_tasks)
if pending_tasks:
await asyncio.gather(*pending_tasks)
assert session is new_session
assert old_session.closed is True
assert LiteLLMAiohttpTransport._background_close_tasks == set()
finally:
LiteLLMAiohttpTransport._background_close_tasks.clear()
if not new_session.closed:
await new_session.close()
@pytest.mark.asyncio
async def test_aiohttp_response_stream_aclose_releases_connection():
"""AiohttpResponseStream.aclose() must call __aexit__ on the aiohttp response."""