mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(aiohttp): preserve trace_configs when session is recreated
Fixes #20174 When a user passes a shared_session with trace_configs to litellm, those trace_configs were lost when the session needed to be recreated (e.g., due to closed session or different event loop). This fix: 1. Stores the trace_configs from the original ClientSession when it is passed to LiteLLMAiohttpTransport 2. Uses the stored trace_configs when creating a new session via the new _create_session_with_trace_configs() helper method 3. Updates all session recreation points to use this helper method This allows aiohttp client tracing logs to work properly when using litellm's shared_session feature.
This commit is contained in:
parent
8a57ee5efb
commit
3d3bb70e8e
2 changed files with 178 additions and 5 deletions
|
|
@ -145,6 +145,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Store the client factory for recreating sessions when needed
|
||||
if callable(client):
|
||||
self._client_factory = client
|
||||
elif isinstance(client, ClientSession):
|
||||
# When a ClientSession is passed directly, preserve its trace_configs
|
||||
# so they can be reused when recreating the session
|
||||
# This fixes https://github.com/BerriAI/litellm/issues/20174
|
||||
self._trace_configs = getattr(client, "_trace_configs", None)
|
||||
|
||||
def _get_valid_client_session(self) -> ClientSession:
|
||||
"""
|
||||
|
|
@ -160,7 +165,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._create_session_with_trace_configs()
|
||||
# Don't return yet - check if the newly created session is valid
|
||||
|
||||
# Check if the session itself is closed
|
||||
|
|
@ -170,7 +175,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._create_session_with_trace_configs()
|
||||
return self.client
|
||||
|
||||
# Check if the existing session is still valid for the current event loop
|
||||
|
|
@ -196,17 +201,34 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._create_session_with_trace_configs()
|
||||
|
||||
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._create_session_with_trace_configs()
|
||||
|
||||
return self.client
|
||||
|
||||
def _create_session_with_trace_configs(self) -> ClientSession:
|
||||
"""
|
||||
Create a new ClientSession with trace_configs from the original session.
|
||||
|
||||
This preserves trace_configs when the session needs to be recreated
|
||||
(e.g., due to closed session or different event loop).
|
||||
|
||||
Fixes https://github.com/BerriAI/litellm/issues/20174
|
||||
"""
|
||||
trace_configs = getattr(self, "_trace_configs", None)
|
||||
if trace_configs:
|
||||
verbose_logger.debug(
|
||||
f"Creating new session with {len(trace_configs)} trace_configs from original session"
|
||||
)
|
||||
return ClientSession(trace_configs=trace_configs)
|
||||
return ClientSession()
|
||||
|
||||
async def _make_aiohttp_request(
|
||||
self,
|
||||
client_session: ClientSession,
|
||||
|
|
@ -285,7 +307,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if hasattr(self, "_client_factory") and callable(self._client_factory):
|
||||
self.client = self._client_factory()
|
||||
else:
|
||||
self.client = ClientSession()
|
||||
self.client = self._create_session_with_trace_configs()
|
||||
client_session = self.client
|
||||
|
||||
# Retry the request with the new session
|
||||
|
|
|
|||
|
|
@ -0,0 +1,151 @@
|
|||
"""
|
||||
Test that aiohttp trace_configs are preserved when session is recreated.
|
||||
|
||||
Fixes https://github.com/BerriAI/litellm/issues/20174
|
||||
|
||||
The issue was that when a user provides a shared_session with trace_configs,
|
||||
those trace_configs would be lost if the session needed to be recreated
|
||||
(e.g., due to closed session or different event loop).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import pytest
|
||||
import aiohttp
|
||||
|
||||
from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport
|
||||
|
||||
|
||||
class TestAiohttpTraceConfigsPreservation:
|
||||
"""Tests for trace_configs preservation in LiteLLMAiohttpTransport"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trace_configs_stored_from_client_session(self):
|
||||
"""Test that trace_configs are extracted and stored from a ClientSession"""
|
||||
# Create a trace config
|
||||
trace_config = aiohttp.TraceConfig()
|
||||
|
||||
# Track if our callback was invoked
|
||||
callback_invoked = False
|
||||
|
||||
async def on_request_start(session, trace_config_ctx, params):
|
||||
nonlocal callback_invoked
|
||||
callback_invoked = True
|
||||
|
||||
trace_config.on_request_start.append(on_request_start)
|
||||
|
||||
# Create a ClientSession with trace_configs
|
||||
async with aiohttp.ClientSession(trace_configs=[trace_config]) as session:
|
||||
# Create transport with the session
|
||||
transport = LiteLLMAiohttpTransport(client=session)
|
||||
|
||||
# Verify trace_configs were stored
|
||||
assert hasattr(transport, "_trace_configs")
|
||||
assert transport._trace_configs is not None
|
||||
assert len(transport._trace_configs) == 1
|
||||
assert transport._trace_configs[0] is trace_config
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trace_configs_used_when_creating_new_session(self):
|
||||
"""Test that stored trace_configs are used when creating a new session"""
|
||||
# Create a trace config
|
||||
trace_config = aiohttp.TraceConfig()
|
||||
|
||||
# Create a ClientSession with trace_configs
|
||||
session = aiohttp.ClientSession(trace_configs=[trace_config])
|
||||
|
||||
# Create transport with the session
|
||||
transport = LiteLLMAiohttpTransport(client=session)
|
||||
|
||||
# Close the session to force recreation
|
||||
await session.close()
|
||||
|
||||
# This should create a new session with the stored trace_configs
|
||||
new_session = transport._create_session_with_trace_configs()
|
||||
|
||||
try:
|
||||
# Verify the new session has trace_configs
|
||||
assert hasattr(new_session, "_trace_configs")
|
||||
assert new_session._trace_configs is not None
|
||||
assert len(new_session._trace_configs) == 1
|
||||
finally:
|
||||
await new_session.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_recreation_preserves_trace_configs(self):
|
||||
"""Test that _get_valid_client_session preserves trace_configs"""
|
||||
# Create a trace config
|
||||
trace_config = aiohttp.TraceConfig()
|
||||
|
||||
# Create a ClientSession with trace_configs
|
||||
session = aiohttp.ClientSession(trace_configs=[trace_config])
|
||||
|
||||
# Create transport with the session
|
||||
transport = LiteLLMAiohttpTransport(client=session)
|
||||
|
||||
# Close the session to force recreation
|
||||
await session.close()
|
||||
|
||||
# Get a valid session (should recreate with trace_configs)
|
||||
new_session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
# Verify the new session has trace_configs
|
||||
assert hasattr(new_session, "_trace_configs")
|
||||
assert new_session._trace_configs is not None
|
||||
assert len(new_session._trace_configs) == 1
|
||||
finally:
|
||||
await new_session.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_trace_configs_creates_plain_session(self):
|
||||
"""Test that a plain ClientSession without trace_configs still works"""
|
||||
# Create a plain ClientSession without trace_configs
|
||||
session = aiohttp.ClientSession()
|
||||
|
||||
# Create transport with the session
|
||||
transport = LiteLLMAiohttpTransport(client=session)
|
||||
|
||||
# The _trace_configs should be empty list or None
|
||||
stored_configs = getattr(transport, "_trace_configs", None)
|
||||
assert stored_configs is None or len(stored_configs) == 0
|
||||
|
||||
# Close the session to force recreation
|
||||
await session.close()
|
||||
|
||||
# This should create a plain session
|
||||
new_session = transport._create_session_with_trace_configs()
|
||||
|
||||
try:
|
||||
# Verify the new session was created (no error)
|
||||
assert new_session is not None
|
||||
assert not new_session.closed
|
||||
finally:
|
||||
await new_session.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callable_factory_takes_precedence(self):
|
||||
"""Test that a callable factory takes precedence over stored trace_configs"""
|
||||
# Create a trace config
|
||||
trace_config = aiohttp.TraceConfig()
|
||||
|
||||
# Track factory calls
|
||||
factory_calls = 0
|
||||
|
||||
def client_factory():
|
||||
nonlocal factory_calls
|
||||
factory_calls += 1
|
||||
return aiohttp.ClientSession()
|
||||
|
||||
# Create transport with a factory
|
||||
transport = LiteLLMAiohttpTransport(client=client_factory)
|
||||
|
||||
# Should not have _trace_configs when factory is used
|
||||
assert not hasattr(transport, "_trace_configs") or transport._trace_configs is None
|
||||
|
||||
# Get a valid session (should use factory)
|
||||
session = transport._get_valid_client_session()
|
||||
|
||||
try:
|
||||
assert factory_calls == 1
|
||||
finally:
|
||||
await session.close()
|
||||
Loading…
Add table
Reference in a new issue