diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 7362ec31dbe..3ee9817524d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -41,7 +41,6 @@ except ImportError as e: # Global variables to track initialization _SESSION_MANAGERS_INITIALIZED = False -_SESSION_MANAGER_TASK = None if MCP_AVAILABLE: from mcp.server import Server @@ -109,48 +108,48 @@ if MCP_AVAILABLE: stateless=True, ) + # Context managers for proper lifecycle management + _session_manager_cm = None + _sse_session_manager_cm = None + async def initialize_session_managers(): """Initialize the session managers. Can be called from main app lifespan.""" - global _SESSION_MANAGERS_INITIALIZED, _SESSION_MANAGER_TASK + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm if _SESSION_MANAGERS_INITIALIZED: return verbose_logger.info("Initializing MCP session managers...") - # Create a task to run the session managers - async def run_session_managers(): - async with session_manager.run(): - async with sse_session_manager.run(): - verbose_logger.info( - "MCP Server started with StreamableHTTP and SSE session managers!" - ) - try: - # Keep running until cancelled - while True: - await asyncio.sleep(1) - except asyncio.CancelledError: - verbose_logger.info("MCP session managers shutting down...") - raise + # Start the session managers with context managers + _session_manager_cm = session_manager.run() + _sse_session_manager_cm = sse_session_manager.run() + + # Enter the context managers + await _session_manager_cm.__aenter__() + await _sse_session_manager_cm.__aenter__() - _SESSION_MANAGER_TASK = asyncio.create_task(run_session_managers()) _SESSION_MANAGERS_INITIALIZED = True - verbose_logger.info("MCP session managers initialization completed!") + verbose_logger.info("MCP Server started with StreamableHTTP and SSE session managers!") async def shutdown_session_managers(): """Shutdown the session managers.""" - global _SESSION_MANAGERS_INITIALIZED, _SESSION_MANAGER_TASK + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm - if _SESSION_MANAGER_TASK and not _SESSION_MANAGER_TASK.done(): + if _SESSION_MANAGERS_INITIALIZED: verbose_logger.info("Shutting down MCP session managers...") - _SESSION_MANAGER_TASK.cancel() - try: - await _SESSION_MANAGER_TASK - except asyncio.CancelledError: - pass - _SESSION_MANAGERS_INITIALIZED = False - _SESSION_MANAGER_TASK = None + try: + if _session_manager_cm: + await _session_manager_cm.__aexit__(None, None, None) + if _sse_session_manager_cm: + await _sse_session_manager_cm.__aexit__(None, None, None) + except Exception as e: + verbose_logger.exception(f"Error during session manager shutdown: {e}") + + _session_manager_cm = None + _sse_session_manager_cm = None + _SESSION_MANAGERS_INITIALIZED = False @contextlib.asynccontextmanager async def lifespan(app) -> AsyncIterator[None]: