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:
shin-bot-litellm 2026-02-01 01:40:20 +00:00
parent 8a57ee5efb
commit 3d3bb70e8e
2 changed files with 178 additions and 5 deletions

View file

@ -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

View file

@ -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()