mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Guard against stdio transport with session caching and fix review issues
- Raise ValueError when use_session_cache=True with stdio transport (stdio spawns child processes that can become orphans if cached) - Fix duplicate _create_transport_context method (second shadowed first) - Fix concurrent retry stampede (identity check before replacing session) - Add _closed guard to prevent operations after close() - Improve _is_connection_error to use type-name matching (no false positives) - Remove unused orig variable in tests - Add tests for stdio caching guard
This commit is contained in:
parent
62121806d9
commit
4afcbc9d66
2 changed files with 386 additions and 353 deletions
|
|
@ -163,7 +163,13 @@ class MCPClient:
|
|||
self.ssl_verify: Optional[VerifyTypes] = ssl_verify
|
||||
self._aws_auth: Optional[httpx.Auth] = aws_auth
|
||||
|
||||
# Session caching configuration
|
||||
# Session caching configuration — stdio transport spawns child processes
|
||||
# that can become orphans if not explicitly cleaned up, so disallow caching.
|
||||
if use_session_cache and transport_type == MCPTransport.stdio:
|
||||
raise ValueError(
|
||||
"Session caching is not supported for stdio transport. "
|
||||
"stdio spawns child processes that cannot be safely cached."
|
||||
)
|
||||
self.use_session_cache: bool = use_session_cache
|
||||
self.session_cache_ttl: float = session_cache_ttl
|
||||
|
||||
|
|
@ -173,6 +179,7 @@ class MCPClient:
|
|||
self._cached_http_client: Optional[httpx.AsyncClient] = None
|
||||
self._session_last_used_at: Optional[float] = None
|
||||
self._session_lock: asyncio.Lock = asyncio.Lock()
|
||||
self._closed: bool = False
|
||||
|
||||
# handle the basic auth value if provided
|
||||
if auth_value:
|
||||
|
|
@ -364,38 +371,6 @@ 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:
|
||||
|
|
@ -455,7 +430,7 @@ class MCPClient:
|
|||
return session
|
||||
|
||||
async def _cleanup_cached_session(self) -> None:
|
||||
"""Clean up any cached session resources."""
|
||||
"""Close and clean up any cached session resources."""
|
||||
if self._cached_session is not None:
|
||||
try:
|
||||
await self._cached_session.__aexit__(None, None, None)
|
||||
|
|
@ -482,6 +457,8 @@ class MCPClient:
|
|||
async def _get_or_create_session(self) -> ClientSession:
|
||||
"""Get a cached session or create a new one."""
|
||||
async with self._session_lock:
|
||||
if self._closed:
|
||||
raise RuntimeError("MCPClient is closed")
|
||||
if self._is_session_valid():
|
||||
verbose_logger.debug(
|
||||
f"MCP client reusing cached session for {self.server_url or 'stdio'}"
|
||||
|
|
@ -492,10 +469,13 @@ class MCPClient:
|
|||
|
||||
def _is_connection_error(self, e: Exception) -> bool:
|
||||
"""Check if exception indicates a broken/closed connection."""
|
||||
if isinstance(e, (ConnectionError, ConnectionResetError, TimeoutError)):
|
||||
# anyio / MCP-specific exceptions (not subclasses of ConnectionError)
|
||||
type_name = type(e).__name__
|
||||
if type_name in ("BrokenResourceError", "ClosedResourceError", "EndOfStream"):
|
||||
return True
|
||||
error_str = str(e).lower()
|
||||
return "broken" in error_str or "closed" in error_str
|
||||
if isinstance(e, ConnectionError):
|
||||
return True
|
||||
return False
|
||||
|
||||
async def run_with_cached_session(
|
||||
self, operation: Callable[[ClientSession], Awaitable[TSessionResult]]
|
||||
|
|
@ -512,8 +492,15 @@ class MCPClient:
|
|||
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()
|
||||
if self._cached_session is session:
|
||||
# First caller to detect failure — replace session
|
||||
session = await self._create_and_cache_session()
|
||||
elif self._cached_session is not None:
|
||||
# Another caller already replaced it — reuse
|
||||
session = self._cached_session
|
||||
else:
|
||||
# Session was closed, create new
|
||||
session = await self._create_and_cache_session()
|
||||
result = await operation(session)
|
||||
self._session_last_used_at = time.time()
|
||||
return result
|
||||
|
|
@ -522,6 +509,7 @@ class MCPClient:
|
|||
async def close(self) -> None:
|
||||
"""Close the client and clean up any cached sessions."""
|
||||
async with self._session_lock:
|
||||
self._closed = True
|
||||
await self._cleanup_cached_session()
|
||||
verbose_logger.info(
|
||||
f"MCP client closed for {self.server_url or 'stdio'}"
|
||||
|
|
|
|||
|
|
@ -312,354 +312,399 @@ class TestMCPClient:
|
|||
assert MCPAuth.token.value == "token"
|
||||
|
||||
|
||||
class TestMCPClientSessionCaching:
|
||||
"""Test MCP Client session caching (connection pooling) functionality"""
|
||||
class TestMCPClientSessionCachingE2E:
|
||||
"""E2E tests for session caching — exercise full flows through _run_operation."""
|
||||
|
||||
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
|
||||
@staticmethod
|
||||
def _mock_transport():
|
||||
ctx = MagicMock()
|
||||
ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
|
||||
ctx.__aexit__ = AsyncMock(return_value=False)
|
||||
return ctx
|
||||
|
||||
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
|
||||
@staticmethod
|
||||
def _mock_http_client():
|
||||
c = MagicMock()
|
||||
c.aclose = AsyncMock()
|
||||
return c
|
||||
|
||||
@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()
|
||||
async def test_caching_reuses_session_across_calls(self, mock_session_cls):
|
||||
"""With caching ON, multiple operations share one transport/session."""
|
||||
sessions = []
|
||||
|
||||
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
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
sessions.append(s)
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://localhost:8765/sse",
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
)
|
||||
transports = []
|
||||
|
||||
call_count = 0
|
||||
def tracked():
|
||||
t = self._mock_transport()
|
||||
h = self._mock_http_client()
|
||||
transports.append(t)
|
||||
return t, h
|
||||
|
||||
async def _operation(session):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return f"result_{call_count}"
|
||||
client._create_transport_context = tracked
|
||||
|
||||
# First call should create session
|
||||
result1 = await client._run_operation(_operation)
|
||||
assert result1 == "result_1"
|
||||
seen = []
|
||||
|
||||
# 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
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_with_cached_session_retries_on_connection_error(self, monkeypatch):
|
||||
"""Test run_with_cached_session retries once on connection error and succeeds"""
|
||||
client = MCPClient(
|
||||
server_url="http://localhost:8765/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
session_cache_ttl=300.0,
|
||||
)
|
||||
|
||||
mock_session = AsyncMock()
|
||||
|
||||
create_calls = {"count": 0}
|
||||
|
||||
async def fake_create_and_cache_session():
|
||||
create_calls["count"] += 1
|
||||
return mock_session
|
||||
|
||||
def fake_is_connection_error(exc: BaseException) -> bool:
|
||||
return True
|
||||
|
||||
async def fake_cleanup_cached_session() -> None:
|
||||
return None
|
||||
|
||||
op_calls = {"count": 0}
|
||||
|
||||
async def _operation(session):
|
||||
op_calls["count"] += 1
|
||||
if op_calls["count"] == 1:
|
||||
raise RuntimeError("Simulated connection error")
|
||||
async def op(session):
|
||||
seen.append(id(session))
|
||||
return "ok"
|
||||
|
||||
monkeypatch.setattr(client, "_create_and_cache_session", fake_create_and_cache_session)
|
||||
monkeypatch.setattr(client, "_is_connection_error", fake_is_connection_error)
|
||||
monkeypatch.setattr(client, "_cleanup_cached_session", fake_cleanup_cached_session)
|
||||
await client._run_operation(op)
|
||||
await client._run_operation(op)
|
||||
await client._run_operation(op)
|
||||
|
||||
result = await client.run_with_cached_session(_operation)
|
||||
|
||||
assert result == "ok"
|
||||
assert op_calls["count"] == 2
|
||||
assert create_calls["count"] == 2
|
||||
assert len(transports) == 1, "should create transport only once"
|
||||
assert len(sessions) == 1, "should create session only once"
|
||||
assert len(set(seen)) == 1, "all ops should see the same session"
|
||||
await client.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_access_uses_single_cached_session(self, monkeypatch):
|
||||
"""Test concurrent operations share a single cached session"""
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_caching_disabled_creates_fresh_session_each_call(self, mock_session_cls):
|
||||
"""With caching OFF, each operation gets a new transport/session."""
|
||||
sessions = []
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
sessions.append(s)
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://localhost:8765/sse",
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=False,
|
||||
)
|
||||
transports = []
|
||||
|
||||
def tracked():
|
||||
t = self._mock_transport()
|
||||
h = self._mock_http_client()
|
||||
transports.append(t)
|
||||
return t, h
|
||||
|
||||
client._create_transport_context = tracked
|
||||
|
||||
async def op(session):
|
||||
return "ok"
|
||||
|
||||
await client._run_operation(op)
|
||||
await client._run_operation(op)
|
||||
|
||||
assert len(transports) == 2, "should create a new transport each time"
|
||||
assert len(sessions) == 2, "should create a new session each time"
|
||||
assert client._cached_session is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_ttl_expiry_creates_new_session_and_closes_old(self, mock_session_cls):
|
||||
"""After TTL expires, next operation creates a new session and closes old resources."""
|
||||
import time as time_mod
|
||||
|
||||
sessions = []
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
sessions.append(s)
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
session_cache_ttl=300.0,
|
||||
session_cache_ttl=1.0,
|
||||
)
|
||||
transports = []
|
||||
http_clients = []
|
||||
|
||||
mock_session = AsyncMock()
|
||||
def tracked():
|
||||
t = self._mock_transport()
|
||||
h = self._mock_http_client()
|
||||
transports.append(t)
|
||||
http_clients.append(h)
|
||||
return t, h
|
||||
|
||||
create_calls = {"count": 0}
|
||||
client._create_transport_context = tracked
|
||||
|
||||
async def fake_create_and_cache_session():
|
||||
create_calls["count"] += 1
|
||||
return mock_session
|
||||
async def op(session):
|
||||
return "ok"
|
||||
|
||||
monkeypatch.setattr(client, "_create_and_cache_session", fake_create_and_cache_session)
|
||||
await client._run_operation(op)
|
||||
assert len(transports) == 1
|
||||
|
||||
async def _operation(session):
|
||||
await asyncio.sleep(0)
|
||||
return id(session)
|
||||
# Expire TTL
|
||||
client._session_last_used_at = time_mod.time() - 2.0
|
||||
|
||||
async def run_one():
|
||||
return await client._run_operation(_operation)
|
||||
await client._run_operation(op)
|
||||
assert len(transports) == 2, "should create a second session after TTL expiry"
|
||||
|
||||
# Old resources should have been closed
|
||||
sessions[0].__aexit__.assert_called()
|
||||
transports[0].__aexit__.assert_called()
|
||||
http_clients[0].aclose.assert_called()
|
||||
|
||||
await client.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_connection_error_retries_with_new_session(self, mock_session_cls):
|
||||
"""On connection error, old session is closed and retry uses a new session."""
|
||||
sessions = []
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
sessions.append(s)
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
)
|
||||
transports = []
|
||||
http_clients = []
|
||||
|
||||
def tracked():
|
||||
t = self._mock_transport()
|
||||
h = self._mock_http_client()
|
||||
transports.append(t)
|
||||
http_clients.append(h)
|
||||
return t, h
|
||||
|
||||
client._create_transport_context = tracked
|
||||
|
||||
# Setup initial session
|
||||
async def setup(s):
|
||||
return "ok"
|
||||
|
||||
await client._run_operation(setup)
|
||||
broken = sessions[0]
|
||||
|
||||
# Operation that fails on the broken session, succeeds on the new one
|
||||
async def retry_op(session):
|
||||
if session is broken:
|
||||
raise ConnectionError("pipe broken")
|
||||
return "recovered"
|
||||
|
||||
result = await client._run_operation(retry_op)
|
||||
assert result == "recovered"
|
||||
assert len(transports) == 2
|
||||
|
||||
# Old resources closed
|
||||
broken.__aexit__.assert_called()
|
||||
transports[0].__aexit__.assert_called()
|
||||
http_clients[0].aclose.assert_called()
|
||||
|
||||
await client.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_concurrent_errors_create_single_recovery_session(self, mock_session_cls):
|
||||
"""Multiple concurrent callers hitting a broken session share one recovery."""
|
||||
sessions = []
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
sessions.append(s)
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
)
|
||||
transport_count = [0]
|
||||
|
||||
def tracked():
|
||||
transport_count[0] += 1
|
||||
return self._mock_transport(), self._mock_http_client()
|
||||
|
||||
client._create_transport_context = tracked
|
||||
|
||||
# Create initial session
|
||||
async def setup(s):
|
||||
return "ok"
|
||||
|
||||
await client._run_operation(setup)
|
||||
assert transport_count[0] == 1
|
||||
broken = sessions[0]
|
||||
|
||||
async def failing_op(session):
|
||||
await asyncio.sleep(0) # yield so all coroutines start concurrently
|
||||
if session is broken:
|
||||
raise ConnectionError("connection lost")
|
||||
return "recovered"
|
||||
|
||||
results = await asyncio.gather(
|
||||
run_one(),
|
||||
run_one(),
|
||||
run_one(),
|
||||
run_one(),
|
||||
client.run_with_cached_session(failing_op),
|
||||
client.run_with_cached_session(failing_op),
|
||||
client.run_with_cached_session(failing_op),
|
||||
)
|
||||
|
||||
assert len(set(results)) == 1
|
||||
assert create_calls["count"] >= 1
|
||||
assert all(r == "recovered" for r in results)
|
||||
# Only 1 recovery session created (2 total: original + recovery)
|
||||
assert transport_count[0] == 2
|
||||
assert len(sessions) == 2
|
||||
|
||||
await client.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_close_cleans_up_and_prevents_reuse(self, mock_session_cls):
|
||||
"""close() cleans up resources; further operations raise RuntimeError."""
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
)
|
||||
transports = []
|
||||
http_clients = []
|
||||
|
||||
def tracked():
|
||||
t = self._mock_transport()
|
||||
h = self._mock_http_client()
|
||||
transports.append(t)
|
||||
http_clients.append(h)
|
||||
return t, h
|
||||
|
||||
client._create_transport_context = tracked
|
||||
|
||||
async def op(s):
|
||||
return "ok"
|
||||
|
||||
await client._run_operation(op)
|
||||
await client.close()
|
||||
|
||||
# Resources cleaned up
|
||||
transports[0].__aexit__.assert_called()
|
||||
http_clients[0].aclose.assert_called()
|
||||
|
||||
# Further operations should fail
|
||||
with pytest.raises(RuntimeError, match="MCPClient is closed"):
|
||||
await client._run_operation(op)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_non_connection_error_propagates_without_retry(self, mock_session_cls):
|
||||
"""Non-connection errors propagate immediately; no retry, no session churn."""
|
||||
|
||||
def make_session(*a, **kw):
|
||||
s = MagicMock()
|
||||
s.__aenter__ = AsyncMock(return_value=s)
|
||||
s.__aexit__ = AsyncMock()
|
||||
s.initialize = AsyncMock()
|
||||
return s
|
||||
|
||||
mock_session_cls.side_effect = make_session
|
||||
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
use_session_cache=True,
|
||||
)
|
||||
transport_count = [0]
|
||||
|
||||
def tracked():
|
||||
transport_count[0] += 1
|
||||
return self._mock_transport(), self._mock_http_client()
|
||||
|
||||
client._create_transport_context = tracked
|
||||
|
||||
async def bad_op(session):
|
||||
raise ValueError("application error")
|
||||
|
||||
with pytest.raises(ValueError, match="application error"):
|
||||
await client.run_with_cached_session(bad_op)
|
||||
|
||||
# No second session created (no retry)
|
||||
assert transport_count[0] == 1
|
||||
await client.close()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"exc,expected",
|
||||
[
|
||||
(type("BrokenResourceError", (Exception,), {})("broken"), True),
|
||||
(type("ClosedResourceError", (Exception,), {})("closed"), True),
|
||||
(type("EndOfStream", (Exception,), {})("eof"), True),
|
||||
(ConnectionError("conn reset"), True),
|
||||
(ConnectionResetError("reset"), True), # subclass of ConnectionError
|
||||
(ValueError("bad input"), False),
|
||||
(RuntimeError("something failed"), False),
|
||||
(TimeoutError("timed out"), False), # not treated as connection error
|
||||
],
|
||||
)
|
||||
def test_is_connection_error_classification(self, exc, expected):
|
||||
"""Connection error classification is precise — no false positives."""
|
||||
client = MCPClient(
|
||||
server_url="http://test/sse",
|
||||
transport_type=MCPTransport.sse,
|
||||
)
|
||||
assert client._is_connection_error(exc) is expected
|
||||
|
||||
|
||||
class TestMCPClientStdioCachingGuard:
|
||||
"""Session caching must be rejected for stdio transport."""
|
||||
|
||||
def test_stdio_with_session_cache_raises(self):
|
||||
with pytest.raises(ValueError, match="not supported for stdio transport"):
|
||||
MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config={"command": "echo", "args": [], "env": {}},
|
||||
use_session_cache=True,
|
||||
)
|
||||
|
||||
def test_stdio_without_session_cache_ok(self):
|
||||
client = MCPClient(
|
||||
transport_type=MCPTransport.stdio,
|
||||
stdio_config={"command": "echo", "args": [], "env": {}},
|
||||
use_session_cache=False,
|
||||
)
|
||||
assert client.use_session_cache is False
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue