mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(aiohttp): prevent closing shared ClientSession in AiohttpTransport (#21117)
When a shared ClientSession is passed to LiteLLMAiohttpTransport, calling aclose() on the transport would close the shared session, breaking other clients still using it. Add owns_session parameter (default True for backwards compatibility) to AiohttpTransport and LiteLLMAiohttpTransport. When a shared session is provided in http_handler.py, owns_session=False is set to prevent the transport from closing a session it does not own. This aligns AiohttpTransport with the ownership pattern already used in AiohttpHandler (aiohttp_handler.py).
This commit is contained in:
parent
5df06b484b
commit
e19dcea159
3 changed files with 42 additions and 3 deletions
|
|
@ -119,8 +119,13 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
|
|||
|
||||
|
||||
class AiohttpTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self, client: Union[ClientSession, Callable[[], ClientSession]]) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
owns_session: bool = True,
|
||||
) -> None:
|
||||
self.client = client
|
||||
self._owns_session = owns_session
|
||||
|
||||
#########################################################
|
||||
# Class variables for proxy settings
|
||||
|
|
@ -128,7 +133,7 @@ class AiohttpTransport(httpx.AsyncBaseTransport):
|
|||
self.proxy_cache: Dict[str, Optional[str]] = {}
|
||||
|
||||
async def aclose(self) -> None:
|
||||
if isinstance(self.client, ClientSession):
|
||||
if self._owns_session and isinstance(self.client, ClientSession):
|
||||
await self.client.close()
|
||||
|
||||
|
||||
|
|
@ -144,10 +149,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
self,
|
||||
client: Union[ClientSession, Callable[[], ClientSession]],
|
||||
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
|
||||
owns_session: bool = True,
|
||||
):
|
||||
self.client = client
|
||||
self._ssl_verify = ssl_verify # Store for per-request SSL override
|
||||
super().__init__(client=client)
|
||||
super().__init__(client=client, owns_session=owns_session)
|
||||
# Store the client factory for recreating sessions when needed
|
||||
if callable(client):
|
||||
self._client_factory = client
|
||||
|
|
|
|||
|
|
@ -866,6 +866,7 @@ class AsyncHTTPHandler:
|
|||
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
|
||||
|
|
|
|||
|
|
@ -12,10 +12,42 @@ sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory
|
|||
|
||||
from litellm.llms.custom_httpx.aiohttp_transport import (
|
||||
AiohttpResponseStream,
|
||||
AiohttpTransport,
|
||||
LiteLLMAiohttpTransport,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclose_does_not_close_shared_session():
|
||||
"""Test that aclose() does not close a session it does not own (shared session)."""
|
||||
session = aiohttp.ClientSession()
|
||||
try:
|
||||
transport = LiteLLMAiohttpTransport(client=session, owns_session=False)
|
||||
await transport.aclose()
|
||||
assert not session.closed, "Shared session should not be closed by transport"
|
||||
finally:
|
||||
await session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclose_closes_owned_session():
|
||||
"""Test that aclose() closes a session it owns."""
|
||||
session = aiohttp.ClientSession()
|
||||
transport = LiteLLMAiohttpTransport(client=session, owns_session=True)
|
||||
await transport.aclose()
|
||||
assert session.closed, "Owned session should be closed by transport"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_owns_session_defaults_to_true():
|
||||
"""Test that owns_session defaults to True for backwards compatibility."""
|
||||
session = aiohttp.ClientSession()
|
||||
transport = AiohttpTransport(client=session)
|
||||
assert transport._owns_session is True
|
||||
await transport.aclose()
|
||||
assert session.closed
|
||||
|
||||
|
||||
class MockAiohttpResponse:
|
||||
"""Mock aiohttp ClientResponse for testing"""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue