Add MCP client session caching for improved performance

Implement optional session caching in MCPClient to reuse connections
instead of creating new ones for every tool call.

Changes:
- Add use_session_cache and session_cache_ttl parameters to MCPClient
- Implement run_with_cached_session() for connection reuse
- Use idle timeout (resets on each use) so active connections stay alive
- Add close() method for explicit resource cleanup
- Extract shared transport creation into _create_transport_context()

Defaults to disabled (use_session_cache=False) for backwards compatibility.
This commit is contained in:
Ryan Crabbe 2026-01-20 22:07:56 -08:00
parent cbbd51a5ce
commit 6534a0c8ef
2 changed files with 464 additions and 8 deletions

View file

@ -4,6 +4,7 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers.
import asyncio
import base64
import time
from typing import Any, Awaitable, Callable, Dict, Generator, List, Optional, Tuple, TypeVar, Union
import httpx
@ -149,6 +150,8 @@ class MCPClient:
extra_headers: Optional[Dict[str, str]] = None,
ssl_verify: Optional[VerifyTypes] = None,
aws_auth: Optional[httpx.Auth] = None,
use_session_cache: bool = False,
session_cache_ttl: float = 300.0, # 5 minutes default
):
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
@ -159,6 +162,18 @@ class MCPClient:
self.extra_headers: Optional[Dict[str, str]] = extra_headers
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
self._aws_auth: Optional[httpx.Auth] = aws_auth
# Session caching configuration
self.use_session_cache: bool = use_session_cache
self.session_cache_ttl: float = session_cache_ttl
# Cached session state
self._cached_session: Optional[ClientSession] = None
self._cached_transport_ctx: Optional[Any] = None
self._cached_http_client: Optional[httpx.AsyncClient] = None
self._session_last_used_at: Optional[float] = None
self._session_lock: asyncio.Lock = asyncio.Lock()
# handle the basic auth value if provided
if auth_value:
self.update_auth_value(auth_value)
@ -349,6 +364,177 @@ class MCPClient:
return factory
def _create_transport_context(self) -> Tuple[Any, Optional[httpx.AsyncClient]]:
"""Create transport context based on transport type."""
http_client: Optional[httpx.AsyncClient] = None
if self.transport_type == MCPTransport.stdio:
if not self.stdio_config:
raise ValueError("stdio_config is required for stdio transport")
server_params = StdioServerParameters(
command=self.stdio_config.get("command", ""),
args=self.stdio_config.get("args", []),
env=self.stdio_config.get("env", {}),
)
return stdio_client(server_params), None
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
if self.transport_type == MCPTransport.sse:
return sse_client(
url=self.server_url,
timeout=self.timeout,
headers=headers,
httpx_client_factory=httpx_client_factory,
), None
verbose_logger.debug("litellm headers for streamable_http_client: %s", headers)
http_client = httpx_client_factory(
headers=headers,
timeout=httpx.Timeout(self.timeout),
)
return streamable_http_client(url=self.server_url, http_client=http_client), http_client
def _is_session_valid(self) -> bool:
"""Check if the cached session is still valid (not idle too long)."""
if self._cached_session is None:
return False
if self._session_last_used_at is None:
return False
# Check idle timeout
idle_time = time.time() - self._session_last_used_at
if idle_time > self.session_cache_ttl:
verbose_logger.debug(
f"MCP client cached session idle timeout (TTL={self.session_cache_ttl}s, idle={idle_time:.1f}s)"
)
return False
return True
async def _create_and_cache_session(self) -> ClientSession:
"""Create a new session and cache it."""
await self._cleanup_cached_session()
transport_ctx, http_client = self._create_transport_context()
session = None
try:
transport = await transport_ctx.__aenter__()
read_stream, write_stream = transport[0], transport[1]
session = ClientSession(read_stream, write_stream)
await session.__aenter__()
await session.initialize()
except Exception:
# Clean up on failure to prevent resource leak
if session is not None:
try:
await session.__aexit__(None, None, None)
except Exception:
pass
try:
await transport_ctx.__aexit__(None, None, None)
except Exception:
pass
if http_client is not None:
try:
await http_client.aclose()
except Exception:
pass
raise
self._cached_transport_ctx = transport_ctx
self._cached_session = session
self._cached_http_client = http_client
self._session_last_used_at = time.time()
verbose_logger.debug(
f"MCP client created and cached new session for {self.server_url or 'stdio'}"
)
return session
async def _cleanup_cached_session(self) -> None:
"""Clean up any cached session resources."""
if self._cached_session is not None:
try:
await self._cached_session.__aexit__(None, None, None)
except Exception as e:
verbose_logger.debug(f"Error closing cached session: {e}")
self._cached_session = None
if self._cached_transport_ctx is not None:
try:
await self._cached_transport_ctx.__aexit__(None, None, None)
except Exception as e:
verbose_logger.debug(f"Error closing cached transport: {e}")
self._cached_transport_ctx = None
if self._cached_http_client is not None:
try:
await self._cached_http_client.aclose()
except Exception as e:
verbose_logger.debug(f"Error closing cached http client: {e}")
self._cached_http_client = None
self._session_last_used_at = None
async def _get_or_create_session(self) -> ClientSession:
"""Get a cached session or create a new one."""
async with self._session_lock:
if self._is_session_valid():
verbose_logger.debug(
f"MCP client reusing cached session for {self.server_url or 'stdio'}"
)
return self._cached_session # type: ignore
return await self._create_and_cache_session()
def _is_connection_error(self, e: Exception) -> bool:
"""Check if exception indicates a broken/closed connection."""
if isinstance(e, (ConnectionError, ConnectionResetError, TimeoutError)):
return True
error_str = str(e).lower()
return "broken" in error_str or "closed" in error_str
async def run_with_cached_session(
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
) -> TSessionResult:
"""Run an operation using a cached session (connection pooling enabled)."""
try:
session = await self._get_or_create_session()
result = await operation(session)
self._session_last_used_at = time.time()
return result
except Exception as e:
if self._is_connection_error(e):
verbose_logger.warning(
f"MCP client cached session appears broken, retrying: {e}"
)
async with self._session_lock:
await self._cleanup_cached_session()
session = await self._create_and_cache_session()
result = await operation(session)
self._session_last_used_at = time.time()
return result
raise
async def close(self) -> None:
"""Close the client and clean up any cached sessions."""
async with self._session_lock:
await self._cleanup_cached_session()
verbose_logger.info(
f"MCP client closed for {self.server_url or 'stdio'}"
)
async def _run_operation(
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
) -> TSessionResult:
"""Run an operation, using cached session if enabled."""
if self.use_session_cache:
return await self.run_with_cached_session(operation)
return await self.run_with_session(operation)
async def list_tools(self) -> List[MCPTool]:
"""List available tools from the server."""
verbose_logger.debug(
@ -359,7 +545,7 @@ class MCPClient:
return await session.list_tools()
try:
result = await self.run_with_session(_list_tools_operation)
result = await self._run_operation(_list_tools_operation)
tool_count = len(result.tools)
tool_names = [tool.name for tool in result.tools]
verbose_logger.info(
@ -424,7 +610,7 @@ class MCPClient:
)
try:
tool_result = await self.run_with_session(_call_tool_operation)
tool_result = await self._run_operation(_call_tool_operation)
verbose_logger.info(
f"MCP client tool call '{call_tool_request_params.name}' completed successfully"
)
@ -474,7 +660,7 @@ class MCPClient:
return await session.list_prompts()
try:
result = await self.run_with_session(_list_prompts_operation)
result = await self._run_operation(_list_prompts_operation)
prompt_count = len(result.prompts)
prompt_names = [prompt.name for prompt in result.prompts]
verbose_logger.info(
@ -520,7 +706,7 @@ class MCPClient:
)
try:
get_prompt_result = await self.run_with_session(_get_prompt_operation)
get_prompt_result = await self._run_operation(_get_prompt_operation)
verbose_logger.info(
f"MCP client get_prompt '{get_prompt_request_params.name}' completed successfully"
)
@ -564,7 +750,7 @@ class MCPClient:
return await session.list_resources()
try:
result = await self.run_with_session(_list_resources_operation)
result = await self._run_operation(_list_resources_operation)
resource_count = len(result.resources)
resource_names = [resource.name for resource in result.resources]
verbose_logger.info(
@ -604,7 +790,7 @@ class MCPClient:
return await session.list_resource_templates()
try:
result = await self.run_with_session(_list_resource_templates_operation)
result = await self._run_operation(_list_resource_templates_operation)
resource_template_count = len(result.resourceTemplates)
resource_template_names = [
resourceTemplate.name for resourceTemplate in result.resourceTemplates
@ -645,7 +831,7 @@ class MCPClient:
return await session.read_resource(url)
try:
read_resource_result = await self.run_with_session(_read_resource_operation)
read_resource_result = await self._run_operation(_read_resource_operation)
verbose_logger.info(
f"MCP client read_resource '{url}' completed successfully"
)

View file

@ -1,5 +1,4 @@
import os
import ssl
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -312,5 +311,276 @@ class TestMCPClient:
assert MCPAuth.token.value == "token"
class TestMCPClientSessionCaching:
"""Test MCP Client session caching (connection pooling) functionality"""
def test_session_cache_defaults_to_disabled(self):
"""Test that session caching is disabled by default for backwards compatibility"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
)
assert client.use_session_cache is False
assert client.session_cache_ttl == 300.0
def test_session_cache_can_be_enabled(self):
"""Test that session caching can be enabled"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
session_cache_ttl=60.0,
)
assert client.use_session_cache is True
assert client.session_cache_ttl == 60.0
def test_is_session_valid_returns_false_when_no_session(self):
"""Test _is_session_valid returns False when no session is cached"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
assert client._is_session_valid() is False
def test_is_session_valid_returns_false_when_no_timestamp(self):
"""Test _is_session_valid returns False when session exists but no timestamp"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
client._cached_session = MagicMock()
client._session_last_used_at = None
assert client._is_session_valid() is False
def test_is_session_valid_returns_false_when_ttl_expired(self):
"""Test _is_session_valid returns False when TTL has expired"""
import time
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
session_cache_ttl=1.0, # 1 second TTL
)
client._cached_session = MagicMock()
client._session_last_used_at = time.time() - 2.0 # 2 seconds ago
assert client._is_session_valid() is False
def test_is_session_valid_returns_true_when_within_ttl(self):
"""Test _is_session_valid returns True when session is within TTL"""
import time
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
session_cache_ttl=300.0, # 5 minutes
)
client._cached_session = MagicMock()
client._session_last_used_at = time.time() # Just now
assert client._is_session_valid() is True
@pytest.mark.asyncio
async def test_cleanup_cached_session(self):
"""Test _cleanup_cached_session properly cleans up resources"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
# Setup mock cached resources
mock_session = MagicMock()
mock_session.__aexit__ = AsyncMock()
mock_transport_ctx = MagicMock()
mock_transport_ctx.__aexit__ = AsyncMock()
mock_http_client = MagicMock()
mock_http_client.aclose = AsyncMock()
client._cached_session = mock_session
client._cached_transport_ctx = mock_transport_ctx
client._cached_http_client = mock_http_client
client._session_last_used_at = 12345.0
await client._cleanup_cached_session()
# Verify cleanup was called
mock_session.__aexit__.assert_called_once()
mock_transport_ctx.__aexit__.assert_called_once()
mock_http_client.aclose.assert_called_once()
# Verify state was cleared
assert client._cached_session is None
assert client._cached_transport_ctx is None
assert client._cached_http_client is None
assert client._session_last_used_at is None
@pytest.mark.asyncio
async def test_close_cleans_up_session(self):
"""Test close() properly cleans up cached session"""
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
# Setup mock cached session
mock_session = MagicMock()
mock_session.__aexit__ = AsyncMock()
client._cached_session = mock_session
await client.close()
# Verify session was cleaned up
assert client._cached_session is None
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.sse_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_run_operation_uses_cached_session_when_enabled(
self, mock_session_class, mock_sse_client
):
"""Test _run_operation uses run_with_cached_session when caching is enabled"""
# Setup mocks
mock_transport = (MagicMock(), MagicMock())
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
mock_sse_client.return_value.__aexit__ = AsyncMock()
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance.__aexit__ = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_class.return_value = mock_session_instance
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
call_count = 0
async def _operation(session):
nonlocal call_count
call_count += 1
return f"result_{call_count}"
# First call should create session
result1 = await client._run_operation(_operation)
assert result1 == "result_1"
# Session should be cached
assert client._cached_session is not None
# Second call should reuse session (not create new one)
result2 = await client._run_operation(_operation)
assert result2 == "result_2"
# sse_client should only be called once (session reused)
assert mock_sse_client.call_count == 1
# Clean up
await client.close()
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.sse_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_run_operation_creates_new_session_when_disabled(
self, mock_session_class, mock_sse_client
):
"""Test _run_operation uses run_with_session (new connection) when caching is disabled"""
# Setup mocks
mock_transport = (MagicMock(), MagicMock())
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
mock_sse_client.return_value.__aexit__ = AsyncMock()
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance.__aexit__ = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_class.return_value = mock_session_instance
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=False, # Caching disabled
)
async def _operation(session):
return "result"
# First call
await client._run_operation(_operation)
# Second call
await client._run_operation(_operation)
# sse_client should be called twice (new session each time)
assert mock_sse_client.call_count == 2
# No cached session should exist
assert client._cached_session is None
@pytest.mark.asyncio
@patch("litellm.experimental_mcp_client.client.sse_client")
@patch("litellm.experimental_mcp_client.client.ClientSession")
async def test_get_or_create_session_creates_new_when_none_cached(
self, mock_session_class, mock_sse_client
):
"""Test _get_or_create_session creates new session when none is cached"""
# Setup mocks
mock_transport = (MagicMock(), MagicMock())
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
mock_sse_client.return_value.__aexit__ = AsyncMock()
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance.__aexit__ = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session_class.return_value = mock_session_instance
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
# No session cached initially
assert client._cached_session is None
# Get or create should create a new session
session = await client._get_or_create_session()
assert session is not None
assert client._cached_session is not None
assert client._session_last_used_at is not None
# Clean up
await client.close()
@pytest.mark.asyncio
async def test_get_or_create_session_reuses_valid_session(self):
"""Test _get_or_create_session reuses session when valid"""
import time
client = MCPClient(
server_url="http://localhost:8765/sse",
transport_type=MCPTransport.sse,
use_session_cache=True,
)
# Pre-cache a mock session
mock_session = MagicMock()
client._cached_session = mock_session
client._session_last_used_at = time.time()
# Get or create should return the cached session
session = await client._get_or_create_session()
assert session is mock_session
if __name__ == "__main__":
pytest.main([__file__])