diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index 132191c946c..7ad5b5bd41a 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -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: 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 6e2e60ba0dd..8c9adba05ae 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -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() diff --git a/tests/test_litellm/test_streaming_connection_cleanup.py b/tests/test_litellm/test_streaming_connection_cleanup.py index 677046bc66c..f4e5197e1d6 100644 --- a/tests/test_litellm/test_streaming_connection_cleanup.py +++ b/tests/test_litellm/test_streaming_connection_cleanup.py @@ -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."""