diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f1173e8371d..2ec17fac476 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -242,7 +242,7 @@ if MCP_AVAILABLE: # client-visible POST path becomes /mcp/messages. from mcp.server.sse import SseServerTransport as _McpSseServerTransport - sse = _McpSseServerTransport("/messages") + sse = _McpSseServerTransport("/mcp/messages") # Create session managers (StreamableHTTP — stateless by default) session_manager = StreamableHTTPSessionManager( app=server, @@ -2743,7 +2743,7 @@ if MCP_AVAILABLE: raw_headers, ) = await extract_mcp_auth_context(scope, path) _sse_client_ip = IPAddressUtils.get_mcp_client_ip(request) - # set_auth_context here is a no-op for actual tool execution since the SDK + # 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: @@ -2788,7 +2788,7 @@ if MCP_AVAILABLE: session = request_ctx.get().session read_stream = getattr(session, "_read_stream", None) - return getattr(read_stream, "_litellm_auth_context", None) + return _session_auth_storage.get(read_stream) if read_stream else None except Exception: return None diff --git a/tests/mcp_tests/test_coverage_boost.py b/tests/mcp_tests/test_coverage_boost.py index 76958e277d7..5f106f2c67e 100644 --- a/tests/mcp_tests/test_coverage_boost.py +++ b/tests/mcp_tests/test_coverage_boost.py @@ -1,5 +1,5 @@ import pytest -from unittest.mock import MagicMock, AsyncMock, patch +from unittest.mock import MagicMock, patch from litellm.proxy._experimental.mcp_server.sampling_handler import ( _convert_single_content, _convert_openai_response_to_mcp_result, @@ -11,25 +11,49 @@ from litellm.proxy._types import UserAPIKeyAuth # Mock MCP types try: from mcp.types import ( - TextContent, ImageContent, SamplingMessage, - CreateMessageRequestParams, ToolUseContent, ToolResultContent + TextContent, + ImageContent, + SamplingMessage, + CreateMessageRequestParams, + ToolUseContent, + ToolResultContent, ) except ImportError: + class TextContent: - def __init__(self, type="text", text=""): self.type = type; self.text = text + def __init__(self, type="text", text=""): + self.type = type + self.text = text + class ImageContent: def __init__(self, type="image", data="", mimeType="image/png"): - self.type = type; self.data = data; self.mimeType = mimeType + self.type = type + self.data = data + self.mimeType = mimeType + class SamplingMessage: - def __init__(self, role, content): self.role = role; self.content = content + def __init__(self, role, content): + self.role = role + self.content = content + class CreateMessageRequestParams: - def __init__(self, messages, maxTokens=100): self.messages = messages; self.maxTokens = maxTokens + def __init__(self, messages, maxTokens=100): + self.messages = messages + self.maxTokens = maxTokens + class ToolUseContent: def __init__(self, type="tool_use", id=None, name=None, input=None): - self.type = type; self.id = id; self.name = name; self.input = input + self.type = type + self.id = id + self.name = name + self.input = input + class ToolResultContent: def __init__(self, type="tool_result", toolUseId=None, content=None): - self.type = type; self.toolUseId = toolUseId; self.content = content + self.type = type + self.toolUseId = toolUseId + self.content = content + class MockAudioContent: def __init__(self, data="audio_data", mimeType="audio/wav"): @@ -37,6 +61,7 @@ class MockAudioContent: self.data = data self.mimeType = mimeType + def test_convert_audio_content(): audio = MockAudioContent() result = _convert_single_content(audio) @@ -44,6 +69,7 @@ def test_convert_audio_content(): assert result["input_audio"]["data"] == "audio_data" assert result["input_audio"]["format"] == "wav" + def test_convert_openai_response_to_mcp_result_with_tool_calls(): mock_choice = MagicMock() mock_choice.message.content = "I will search now" @@ -51,42 +77,51 @@ def test_convert_openai_response_to_mcp_result_with_tool_calls(): mock_tool_call.id = "call_1" mock_tool_call.function.name = "search" mock_tool_call.function.arguments = '{"q": "test"}' - + mock_choice.message.tool_calls = [mock_tool_call] mock_choice.finish_reason = "tool_calls" - + mock_response = MagicMock() mock_response.choices = [mock_choice] mock_response.model = "gpt-4" - + result = _convert_openai_response_to_mcp_result(mock_response, model_name="gpt-4") assert result.role == "assistant" # It should have both text and tool use content # Depending on implementation it might return CreateMessageResultWithTools assert hasattr(result, "content") + @pytest.mark.asyncio async def test_handle_sampling_no_package_error(): params = CreateMessageRequestParams( - messages=[SamplingMessage(role="user", content=TextContent(type="text", text="hi"))], - maxTokens=100 + messages=[ + SamplingMessage(role="user", content=TextContent(type="text", text="hi")) + ], + maxTokens=100, ) - with patch("litellm.proxy._experimental.mcp_server.sampling_handler.MCP_SAMPLING_AVAILABLE", False): + with patch( + "litellm.proxy._experimental.mcp_server.sampling_handler.MCP_SAMPLING_AVAILABLE", + False, + ): result = await handle_sampling_create_message(context=None, params=params) assert hasattr(result, "message") assert "MCP sampling is not available" in result.message + @pytest.mark.asyncio async def test_get_or_extract_auth_context_fallback(): # Test fallback to session read_stream when ContextVar is empty mock_session = MagicMock() mock_read_stream = MagicMock() mock_user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-1") - + from litellm.proxy._experimental.mcp_server.server import 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, @@ -94,22 +129,32 @@ async def test_get_or_extract_auth_context_fallback(): mcp_server_auth_headers=None, oauth2_headers=None, raw_headers=None, - client_ip=None + client_ip=None, ) - + mock_request_ctx = MagicMock() mock_request_ctx.get.return_value.session = mock_session - - with patch("litellm.proxy._experimental.mcp_server.server.get_auth_context", return_value=(None, None, None, None, None, {}, None)): + + with patch( + "litellm.proxy._experimental.mcp_server.server.get_auth_context", + return_value=(None, None, None, None, None, {}, None), + ): with patch("mcp.server.lowlevel.server.request_ctx", mock_request_ctx): result = await get_or_extract_auth_context() assert result[0] == mock_user_auth assert result[0].api_key is not None + @pytest.mark.asyncio async def test_get_or_extract_auth_context_exception_handling(): # Test that it handles exceptions in fallback gracefully - with patch("litellm.proxy._experimental.mcp_server.server.get_auth_context", return_value=(None, None, None, None, None, {}, None)): - with patch("mcp.server.lowlevel.server.request_ctx", side_effect=Exception("Context error")): + with patch( + "litellm.proxy._experimental.mcp_server.server.get_auth_context", + return_value=(None, None, None, None, None, {}, None), + ): + with patch( + "mcp.server.lowlevel.server.request_ctx", + side_effect=Exception("Context error"), + ): result = await get_or_extract_auth_context() assert result[0] is None diff --git a/tests/mcp_tests/test_server_coverage.py b/tests/mcp_tests/test_server_coverage.py index 3e360293759..6e11e5458e0 100644 --- a/tests/mcp_tests/test_server_coverage.py +++ b/tests/mcp_tests/test_server_coverage.py @@ -1,6 +1,5 @@ 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, @@ -9,6 +8,7 @@ from litellm.proxy._experimental.mcp_server.server import ( _get_tools_from_mcp_servers, ) + @pytest.mark.asyncio async def test_get_prompts_from_mcp_servers_coverage(): server1 = MCPServer( @@ -19,18 +19,25 @@ async def test_get_prompts_from_mcp_servers_coverage(): ) 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: + + 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"] + 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( @@ -38,18 +45,23 @@ async def test_get_resources_from_mcp_servers_coverage(): ) 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: + + 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"] + 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( @@ -57,18 +69,23 @@ async def test_get_resource_templates_from_mcp_servers_coverage(): ) 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: + + 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"] + 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( @@ -76,9 +93,15 @@ async def test_get_tools_from_mcp_servers_coverage(): ) 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: + + 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( @@ -86,7 +109,61 @@ async def test_get_tools_from_mcp_servers_coverage(): mcp_auth_header=None, mcp_servers=["test1"], log_list_tools_to_spendlogs=True, - litellm_trace_id="test-trace" + litellm_trace_id="test-trace", ) assert len(result) == 1 assert result[0] == mock_tool + + +@pytest.mark.asyncio +async def test_handle_stale_mcp_session(): + """Test the handle_stale_mcp_session logic for handling missing session IDs on multiple workers.""" + from litellm.proxy._experimental.mcp_server.server import _handle_stale_mcp_session + + # 1. DELETE request for non-existent session + scope_delete = { + "headers": [(b"mcp-session-id", b"stale-session-123")], + "method": "DELETE", + "type": "http", + } + mock_mgr = MagicMock() + mock_mgr._server_instances = {} + mock_receive = AsyncMock() + mock_send = AsyncMock() + + result_delete = await _handle_stale_mcp_session( + scope_delete, mock_receive, mock_send, mock_mgr + ) + assert result_delete is True + # The JSONResponse success should have been sent + assert mock_send.call_count >= 1 + + # 2. POST request for non-existent session -> header should be stripped + scope_post = { + "headers": [ + (b"mcp-session-id", b"stale-session-123"), + (b"content-type", b"application/json"), + ], + "method": "POST", + } + result_post = await _handle_stale_mcp_session( + scope_post, mock_receive, mock_send, mock_mgr + ) + assert result_post is False + headers = dict(scope_post["headers"]) + assert b"mcp-session-id" not in headers + assert b"content-type" in headers + + # 3. Request with valid session -> should return False immediately + mock_mgr._server_instances = {"valid-session-123": MagicMock()} + scope_valid = { + "headers": [(b"mcp-session-id", b"valid-session-123")], + "method": "POST", + } + result_valid = await _handle_stale_mcp_session( + scope_valid, mock_receive, mock_send, mock_mgr + ) + assert result_valid is False + # Header should not be stripped + headers_valid = dict(scope_valid["headers"]) + assert b"mcp-session-id" in headers_valid 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 a52180db1ff..a1fcff0604a 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 @@ -46,6 +46,7 @@ async def test_get_or_extract_auth_context_fallback(): 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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 9df6408b0d7..a709af811b2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -853,21 +853,16 @@ async def test_concurrent_initialize_session_managers(): # Reset state before test original_initialized = mcp_server._SESSION_MANAGERS_INITIALIZED original_session_cm = mcp_server._session_manager_cm - original_sse_session_cm = mcp_server._sse_session_manager_cm try: mcp_server._SESSION_MANAGERS_INITIALIZED = False mcp_server._session_manager_cm = None - mcp_server._sse_session_manager_cm = None # Mock the session managers to avoid actual MCP initialization with ( patch( "litellm.proxy._experimental.mcp_server.server.session_manager" ) as mock_session_manager, - patch( - "litellm.proxy._experimental.mcp_server.server.sse_session_manager" - ) as mock_sse_session_manager, patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"), ): # Mock the run() method to return a mock context manager @@ -876,7 +871,6 @@ async def test_concurrent_initialize_session_managers(): mock_cm.__aexit__ = AsyncMock() mock_session_manager.run.return_value = mock_cm - mock_sse_session_manager.run.return_value = mock_cm # Create multiple concurrent tasks that call initialize_session_managers async def init_task(): @@ -896,14 +890,11 @@ async def test_concurrent_initialize_session_managers(): assert ( mock_session_manager.run.call_count == 1 ), f"Expected 1 call to session_manager.run(), got {mock_session_manager.run.call_count}" - assert ( - mock_sse_session_manager.run.call_count == 1 - ), f"Expected 1 call to sse_session_manager.run(), got {mock_sse_session_manager.run.call_count}" # The context managers should only be entered once each assert ( - mock_cm.__aenter__.call_count == 2 - ), f"Expected 2 calls to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}" + mock_cm.__aenter__.call_count == 1 + ), f"Expected 1 call to __aenter__ (one for each session manager), got {mock_cm.__aenter__.call_count}" # State should be properly set assert mcp_server._SESSION_MANAGERS_INITIALIZED is True @@ -912,7 +903,6 @@ async def test_concurrent_initialize_session_managers(): # Restore original state mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized mcp_server._session_manager_cm = original_session_cm - mcp_server._sse_session_manager_cm = original_sse_session_cm @pytest.mark.asyncio @@ -1077,20 +1067,13 @@ async def test_oauth2_headers_passed_to_mcp_client(): captured_client_args = {} async def mock_create_mcp_client( - server, - mcp_auth_header=None, - extra_headers=None, - stdio_env=None, + *args, + **kwargs, ): # Capture the arguments for verification - captured_client_args.update( - { - "server": server, - "mcp_auth_header": mcp_auth_header, - "extra_headers": extra_headers, - "stdio_env": stdio_env, - } - ) + captured_client_args.update(kwargs) + if args and len(args) > 0: + captured_client_args["server"] = args[0] # Return a mock client that doesn't actually connect mock_client = MagicMock() return mock_client