From 4218c08fd8849c5467b97048e80c6a141635ea8c Mon Sep 17 00:00:00 2001 From: Yug Date: Wed, 29 Apr 2026 16:51:32 +0530 Subject: [PATCH] resolve --- .../mcp_server/mcp_server_manager.py | 2 + .../proxy/_experimental/mcp_server/server.py | 45 ++++----- tests/mcp_tests/test_coverage_boost.py | 6 +- tests/mcp_tests/test_server_coverage.py | 92 +++++++++++++++++++ .../_experimental/mcp_server/test_mcp_auth.py | 6 +- .../mcp_server/test_mcp_hook_extra_headers.py | 6 +- 6 files changed, 122 insertions(+), 35 deletions(-) create mode 100644 tests/mcp_tests/test_server_coverage.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c8712ad1646..a846c27f838 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1282,6 +1282,7 @@ class MCPServerManager: extra_headers: Optional[Dict[str, str]] = None, add_prefix: bool = True, raw_headers: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -1315,6 +1316,7 @@ class MCPServerManager: mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, + user_api_key_auth=user_api_key_auth, ) ## HANDLE OPENAPI TOOLS diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bd6b615259a..f1173e8371d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -102,6 +102,13 @@ try: Tool, ) from mcp.server.session import ServerSession as _McpServerSession + import weakref + + # Storage for session-specific auth context to avoid fragile monkey-patching + # and context-leakage between concurrent SSE sessions. + _session_auth_storage: "weakref.WeakKeyDictionary[Any, MCPAuthenticatedUser]" = ( + weakref.WeakKeyDictionary() + ) active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = ( contextvars.ContextVar("active_mcp_session", default=None) @@ -243,20 +250,12 @@ if MCP_AVAILABLE: json_response=False, # enables SSE streaming stateless=True, ) - # Create SSE session manager - sse_session_manager = StreamableHTTPSessionManager( - app=server, - event_store=None, - json_response=False, # Use SSE responses for this endpoint - 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_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm # Use async lock to prevent concurrent initialization async with _INITIALIZATION_LOCK: if _SESSION_MANAGERS_INITIALIZED: @@ -264,10 +263,8 @@ if MCP_AVAILABLE: verbose_logger.info("Initializing MCP session managers...") # 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_MANAGERS_INITIALIZED = True verbose_logger.info( "MCP Server started with StreamableHTTP session manager and SSE transport!" @@ -275,18 +272,15 @@ if MCP_AVAILABLE: async def shutdown_session_managers(): """Shutdown the session managers.""" - global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm if _SESSION_MANAGERS_INITIALIZED: verbose_logger.info("Shutting down MCP session managers...") 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 @@ -2688,10 +2682,11 @@ if MCP_AVAILABLE: "SSE connection established, running server loop..." ) - # ContextVars can sometimes be lost when the MCP SDK spawns internal tasks + # ContextVars are lost when the MCP SDK spawns internal tasks # (e.g. _receive_loop), so tool handlers can't read auth_context_var reliably. - # Storing it on the read_stream ensures it's isolated per SSE session. - streams[0]._litellm_auth_context = MCPAuthenticatedUser( # type: ignore + # Storing it in session-isolated storage ensures it's isolated per SSE session. + # We store it before calling server.run so it's available for the session's lifespan. + _session_auth_storage[streams[0]] = MCPAuthenticatedUser( # type: ignore user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2748,15 +2743,9 @@ if MCP_AVAILABLE: raw_headers, ) = await extract_mcp_auth_context(scope, path) _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - set_auth_context( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - client_ip=_sse_client_ip, - ) + # set_auth_context here is a no-op for actual tool execution since the SDK + # processes messages in background tasks that don't inherit this ContextVar. + # Authentication must be recovered from the session-auth-storage during execution. except HTTPException: raise except Exception as e: @@ -2923,7 +2912,7 @@ if MCP_AVAILABLE: session = request_ctx.get().session read_stream = getattr(session, "_read_stream", None) - stored = getattr(read_stream, "_litellm_auth_context", None) + stored = _session_auth_storage.get(read_stream) if read_stream else None except Exception as e: verbose_logger.debug( f"get_or_extract_auth_context FALLBACK failed: {e}" diff --git a/tests/mcp_tests/test_coverage_boost.py b/tests/mcp_tests/test_coverage_boost.py index 9955012f40b..76958e277d7 100644 --- a/tests/mcp_tests/test_coverage_boost.py +++ b/tests/mcp_tests/test_coverage_boost.py @@ -84,7 +84,10 @@ async def test_get_or_extract_auth_context_fallback(): mock_user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-1") from litellm.proxy._experimental.mcp_server.server import MCPAuthenticatedUser - mock_read_stream._litellm_auth_context = MCPAuthenticatedUser( + mock_session._read_stream = mock_read_stream + + from litellm.proxy._experimental.mcp_server.server import _session_auth_storage + _session_auth_storage[mock_read_stream] = MCPAuthenticatedUser( user_api_key_auth=mock_user_auth, mcp_auth_header=None, mcp_servers=None, @@ -93,7 +96,6 @@ async def test_get_or_extract_auth_context_fallback(): raw_headers=None, client_ip=None ) - mock_session._read_stream = mock_read_stream mock_request_ctx = MagicMock() mock_request_ctx.get.return_value.session = mock_session diff --git a/tests/mcp_tests/test_server_coverage.py b/tests/mcp_tests/test_server_coverage.py new file mode 100644 index 00000000000..3e360293759 --- /dev/null +++ b/tests/mcp_tests/test_server_coverage.py @@ -0,0 +1,92 @@ +import pytest +from unittest.mock import MagicMock, AsyncMock, patch +from litellm.proxy._types import MCPTransportType +from litellm.types.mcp_server.mcp_server_manager import MCPServer +from litellm.proxy._experimental.mcp_server.server import ( + _get_prompts_from_mcp_servers, + _get_resources_from_mcp_servers, + _get_resource_templates_from_mcp_servers, + _get_tools_from_mcp_servers, +) + +@pytest.mark.asyncio +async def test_get_prompts_from_mcp_servers_coverage(): + server1 = MCPServer( + server_id="test-1", name="test1", transport="stdio", url="http://test1" + ) + server2 = MCPServer( + server_id="test-2", name="test2", transport="stdio", url="http://test2" + ) + mock_prompt = MagicMock() + mock_prompt.name = "test_prompt" + + with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1, server2]): + with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_prompts_from_server", new_callable=AsyncMock) as mock_get: + mock_get.side_effect = [[mock_prompt], Exception("Server error")] + result = await _get_prompts_from_mcp_servers( + user_api_key_auth=None, + mcp_auth_header=None, + mcp_servers=["test1", "test2"] + ) + assert len(result) == 1 + assert result[0] == mock_prompt + +@pytest.mark.asyncio +async def test_get_resources_from_mcp_servers_coverage(): + server1 = MCPServer( + server_id="test-1", name="test1", transport="stdio", url="http://test1" + ) + mock_resource = MagicMock() + mock_resource.name = "test_resource" + + with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]): + with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resources_from_server", new_callable=AsyncMock) as mock_get: + mock_get.return_value = [mock_resource] + result = await _get_resources_from_mcp_servers( + user_api_key_auth=None, + mcp_auth_header=None, + mcp_servers=["test1"] + ) + assert len(result) == 1 + assert result[0] == mock_resource + +@pytest.mark.asyncio +async def test_get_resource_templates_from_mcp_servers_coverage(): + server1 = MCPServer( + server_id="test-1", name="test1", transport="stdio", url="http://test1" + ) + mock_template = MagicMock() + mock_template.name = "test_template" + + with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]): + with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resource_templates_from_server", new_callable=AsyncMock) as mock_get: + mock_get.return_value = [mock_template] + result = await _get_resource_templates_from_mcp_servers( + user_api_key_auth=None, + mcp_auth_header=None, + mcp_servers=["test1"] + ) + assert len(result) == 1 + assert result[0] == mock_template + +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_coverage(): + server1 = MCPServer( + server_id="test-1", name="test1", transport="stdio", url="http://test1" + ) + mock_tool = MagicMock() + mock_tool.name = "test_tool" + + with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]): + with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", new_callable=AsyncMock) as mock_get: + mock_get.return_value = [mock_tool] + # test with some tracking headers + result = await _get_tools_from_mcp_servers( + user_api_key_auth=None, + mcp_auth_header=None, + mcp_servers=["test1"], + log_list_tools_to_spendlogs=True, + litellm_trace_id="test-trace" + ) + assert len(result) == 1 + assert result[0] == mock_tool diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py index a16351aac80..a52180db1ff 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_auth.py @@ -40,12 +40,14 @@ async def test_get_or_extract_auth_context_fallback(): auth_data = UserAPIKeyAuth(api_key="fallback-key") auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data) - # Mock request_ctx.get().session._read_stream._litellm_auth_context + # Mock request_ctx.get().session._read_stream mock_session = MagicMock() mock_read_stream = MagicMock() - mock_read_stream._litellm_auth_context = auth_user mock_session._read_stream = mock_read_stream + from litellm.proxy._experimental.mcp_server.server import _session_auth_storage + _session_auth_storage[mock_read_stream] = auth_user + mock_request_ctx = MagicMock() mock_request_ctx.get.return_value.session = mock_session diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 084bc75c07d..c58ca053ab6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -548,7 +548,7 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None + server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs ): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() @@ -588,7 +588,7 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None + server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs ): captured_extra_headers["value"] = extra_headers mock_client = MagicMock() @@ -634,7 +634,7 @@ class TestHookHeaderMergePriority: captured_extra_headers: Dict[str, Any] = {} async def fake_create_mcp_client( - server, mcp_auth_header=None, extra_headers=None, stdio_env=None + server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs ): captured_extra_headers["value"] = extra_headers mock_client = MagicMock()