diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 582d044023c..5ad2dd54853 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -11,15 +11,15 @@ from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParamete from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client from mcp.client.streamable_http import streamable_http_client +from mcp.types import CallToolRequestParams as MCPCallToolRequestParams +from mcp.types import CallToolResult as MCPCallToolResult from mcp.types import ( - CallToolRequestParams as MCPCallToolRequestParams, GetPromptRequestParams, GetPromptResult, Prompt, ResourceTemplate, + TextContent, ) -from mcp.types import CallToolResult as MCPCallToolResult -from mcp.types import TextContent from mcp.types import Tool as MCPTool from pydantic import AnyUrl @@ -80,6 +80,8 @@ class MCPClient: """Open a session, run the provided coroutine, and clean up.""" transport_ctx = None http_client: Optional[httpx.AsyncClient] = None + transport = None + session_ctx = None try: if self.transport_type == MCPTransport.stdio: @@ -119,20 +121,50 @@ class MCPClient: if transport_ctx is None: raise RuntimeError("Failed to create transport context") - async with transport_ctx as transport: + # Enter transport context + transport = await transport_ctx.__aenter__() + try: read_stream, write_stream = transport[0], transport[1] session_ctx = ClientSession(read_stream, write_stream) - async with session_ctx as session: + + # Enter session context + session = await session_ctx.__aenter__() + try: await session.initialize() - return await operation(session) + result = await operation(session) + return result + finally: + # Ensure session context is properly exited + if session_ctx is not None: + try: + await session_ctx.__aexit__(None, None, None) + except Exception as e: + verbose_logger.debug( + f"Error during session context exit: {e}" + ) + finally: + # Ensure transport context is properly exited + if transport_ctx is not None: + try: + await transport_ctx.__aexit__(None, None, None) + except Exception as e: + verbose_logger.debug( + f"Error during transport context exit: {e}" + ) except Exception: verbose_logger.warning( "MCP client run_with_session failed for %s", self.server_url or "stdio" ) raise finally: + # Always clean up http_client if it was created if http_client is not None: - await http_client.aclose() + try: + await http_client.aclose() + except Exception as e: + verbose_logger.debug( + f"Error during http_client cleanup: {e}" + ) def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): """