mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(aiohttp): keep keep-alive connector config when a session is rebuilt (#34962)
This commit is contained in:
parent
86ba228d92
commit
caede1c5a0
4 changed files with 171 additions and 39 deletions
|
|
@ -143,13 +143,26 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
|
||||
owns_session: bool = True,
|
||||
session_factory: Callable[[], ClientSession] | None = None,
|
||||
):
|
||||
self.client = client
|
||||
self._ssl_verify = ssl_verify # Store for per-request SSL override
|
||||
super().__init__(client=client, owns_session=owns_session)
|
||||
# Store the client factory for recreating sessions when needed
|
||||
if callable(client):
|
||||
self._client_factory = client
|
||||
default_factory: Callable[[], ClientSession] = client if callable(client) else ClientSession
|
||||
self._client_factory: Callable[[], ClientSession] = session_factory or default_factory
|
||||
|
||||
def _rebuild_session(self) -> ClientSession:
|
||||
"""
|
||||
Build a replacement session from the configured factory.
|
||||
|
||||
The replacement is reachable only from this transport, so the transport
|
||||
owns it from here on even when it was originally handed a session it did
|
||||
not own (the proxy's shared session).
|
||||
"""
|
||||
session = self._client_factory()
|
||||
self._owns_session = True
|
||||
return session
|
||||
|
||||
def _get_valid_client_session(self) -> ClientSession:
|
||||
"""
|
||||
|
|
@ -158,24 +171,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
This handles the case where the session was created in a different
|
||||
event loop that may have been closed (common in CI/CD environments).
|
||||
"""
|
||||
from aiohttp.client import ClientSession
|
||||
|
||||
# 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()
|
||||
self.client = self._rebuild_session()
|
||||
# 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:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
return self.client
|
||||
|
||||
# Check if the existing session is still valid for the current event loop
|
||||
|
|
@ -188,7 +193,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Close old session to prevent leaks
|
||||
old_session = self.client
|
||||
try:
|
||||
if not old_session.closed:
|
||||
if self._owns_session and not old_session.closed:
|
||||
try:
|
||||
asyncio.create_task(old_session.close())
|
||||
except RuntimeError:
|
||||
|
|
@ -198,17 +203,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
verbose_logger.debug(f"Error closing old session: {e}")
|
||||
|
||||
# Create a new session in the current event loop
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
|
||||
except (RuntimeError, AttributeError):
|
||||
# If we can't check the loop or session is invalid, recreate it
|
||||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._rebuild_session()
|
||||
|
||||
return self.client
|
||||
|
||||
|
|
@ -303,10 +302,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
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()
|
||||
self.client = self._rebuild_session()
|
||||
client_session = self.client
|
||||
|
||||
# Retry the request with the new session
|
||||
|
|
|
|||
|
|
@ -1013,17 +1013,6 @@ class AsyncHTTPHandler:
|
|||
|
||||
verbose_logger.debug("Creating AiohttpTransport...")
|
||||
|
||||
# Use shared session if provided and valid
|
||||
if shared_session is not None and not shared_session.closed:
|
||||
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
ssl_verify=ssl_for_transport,
|
||||
owns_session=False,
|
||||
)
|
||||
|
||||
# Create new session only if none provided or existing one is invalid
|
||||
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
|
||||
transport_connector_kwargs = {
|
||||
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
|
||||
|
|
@ -1041,11 +1030,26 @@ class AsyncHTTPHandler:
|
|||
if socket_factory is not None:
|
||||
transport_connector_kwargs["socket_factory"] = socket_factory
|
||||
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=lambda: ClientSession(
|
||||
def session_factory() -> ClientSession:
|
||||
return ClientSession(
|
||||
connector=TCPConnector(**transport_connector_kwargs),
|
||||
trust_env=trust_env,
|
||||
),
|
||||
)
|
||||
|
||||
# Use shared session if provided and valid
|
||||
if shared_session is not None and not shared_session.closed:
|
||||
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
ssl_verify=ssl_for_transport,
|
||||
owns_session=False,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
|
||||
# Create new session only if none provided or existing one is invalid
|
||||
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=session_factory,
|
||||
ssl_verify=ssl_for_transport,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import socket
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import aiohttp
|
||||
import pytest
|
||||
|
||||
|
||||
def _invoke_connector_factory(http_handler_module):
|
||||
"""
|
||||
|
|
@ -159,3 +162,32 @@ def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch
|
|||
setsockopt_calls[(socket.IPPROTO_TCP, fake_socket_module.TCP_KEEPALIVE)] == 60
|
||||
)
|
||||
assert (socket.IPPROTO_TCP, getattr(socket, "TCP_KEEPIDLE", -1)) not in setsockopt_calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shared_session_transport_rebuilds_with_socket_factory(monkeypatch):
|
||||
"""
|
||||
The proxy hands _create_aiohttp_transport an already-built shared session.
|
||||
When that session is rebuilt (closed session, or a session from another
|
||||
event loop) the replacement must still carry the keep-alive socket factory
|
||||
and the configured keepalive timeout, otherwise AIOHTTP_SO_KEEPALIVE stops
|
||||
protecting every later request served by that transport.
|
||||
"""
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True)
|
||||
|
||||
shared_session = aiohttp.ClientSession()
|
||||
transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(shared_session=shared_session)
|
||||
await shared_session.close()
|
||||
|
||||
rebuilt_session = MagicMock(name="rebuilt_session")
|
||||
|
||||
with patch.object(http_handler_module, "TCPConnector", return_value=MagicMock(name="connector")) as mock_tcp_connector:
|
||||
with patch.object(http_handler_module, "ClientSession", return_value=rebuilt_session):
|
||||
assert transport._get_valid_client_session() is rebuilt_session
|
||||
|
||||
assert mock_tcp_connector.call_count == 1
|
||||
assert callable(mock_tcp_connector.call_args.kwargs.get("socket_factory"))
|
||||
assert mock_tcp_connector.call_args.kwargs["keepalive_timeout"] == http_handler_module.AIOHTTP_KEEPALIVE_TIMEOUT
|
||||
|
|
|
|||
|
|
@ -727,3 +727,103 @@ async def test_response_stream_closes_response_on_generator_exit():
|
|||
await iterator.aclose()
|
||||
|
||||
assert mock_response.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_closed_shared_session_rebuild_uses_injected_session_factory():
|
||||
"""
|
||||
A transport handed an already-built session (the proxy's shared session)
|
||||
must rebuild through the injected factory. Rebuilding with a bare
|
||||
ClientSession drops the connector's keep-alive socket options, pool limits
|
||||
and DNS cache for every later request on that transport.
|
||||
"""
|
||||
shared_session = aiohttp.ClientSession()
|
||||
await shared_session.close()
|
||||
|
||||
rebuilt = []
|
||||
|
||||
def session_factory():
|
||||
session = _make_mock_session()
|
||||
rebuilt.append(session)
|
||||
return session
|
||||
|
||||
transport = LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
owns_session=False,
|
||||
session_factory=session_factory, # type: ignore
|
||||
)
|
||||
|
||||
assert transport._get_valid_client_session() in rebuilt
|
||||
|
||||
|
||||
def test_rebuild_without_running_loop_uses_injected_session_factory():
|
||||
"""
|
||||
The loop-validity fallback must also go through the injected factory, so a
|
||||
transport recovering outside a running event loop does not silently swap in
|
||||
an unconfigured session.
|
||||
"""
|
||||
rebuilt = []
|
||||
|
||||
def session_factory():
|
||||
session = _make_mock_session()
|
||||
rebuilt.append(session)
|
||||
return session
|
||||
|
||||
transport = LiteLLMAiohttpTransport(
|
||||
client=object(), # type: ignore
|
||||
session_factory=session_factory, # type: ignore
|
||||
)
|
||||
|
||||
assert transport._get_valid_client_session() in rebuilt
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rebuilt_session_becomes_transport_owned():
|
||||
"""
|
||||
A rebuilt session is reachable only from the transport, so aclose() must
|
||||
close it even when the transport was handed a session it did not own.
|
||||
"""
|
||||
shared_session = aiohttp.ClientSession()
|
||||
await shared_session.close()
|
||||
|
||||
replacement = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
owns_session=False,
|
||||
session_factory=lambda: replacement,
|
||||
)
|
||||
|
||||
assert transport._get_valid_client_session() is replacement
|
||||
|
||||
await transport.aclose()
|
||||
|
||||
assert replacement.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_loop_rebuild_does_not_close_unowned_session():
|
||||
"""
|
||||
A session the transport does not own (the proxy's shared session) is used by
|
||||
other transports too, so a rebuild must leave it open for them.
|
||||
"""
|
||||
shared_session = aiohttp.ClientSession()
|
||||
running_loop = asyncio.get_running_loop()
|
||||
other_loop = asyncio.new_event_loop()
|
||||
|
||||
replacement = _make_mock_session()
|
||||
transport = LiteLLMAiohttpTransport(
|
||||
client=shared_session,
|
||||
owns_session=False,
|
||||
session_factory=lambda: replacement, # type: ignore
|
||||
)
|
||||
|
||||
try:
|
||||
shared_session._loop = other_loop
|
||||
assert transport._get_valid_client_session() is replacement
|
||||
shared_session._loop = running_loop
|
||||
await asyncio.sleep(0.05)
|
||||
assert not shared_session.closed
|
||||
finally:
|
||||
shared_session._loop = running_loop
|
||||
other_loop.close()
|
||||
await shared_session.close()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue