mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
cbbd51a5ce
commit
6534a0c8ef
2 changed files with 464 additions and 8 deletions
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue