From 12f1b360896342d2caddc776720a14d5156c4214 Mon Sep 17 00:00:00 2001 From: Yug Date: Fri, 1 May 2026 07:41:59 +0530 Subject: [PATCH] resolve --- litellm/experimental_mcp_client/client.py | 49 +++++- .../proxy/_experimental/mcp_server/server.py | 24 ++- tests/mcp_tests/test_sampling_handler.py | 154 +++++++++++++----- .../test_mcp_elicitation_handler.py | 35 +++- .../mcp_server/test_mcp_sampling_handler.py | 4 +- .../mcp_server/test_mcp_server_manager.py | 4 + .../mcp_server/test_mcp_stale_session.py | 133 ++++++++------- .../mcp_server/test_semantic_tool_filter.py | 28 +--- 8 files changed, 311 insertions(+), 120 deletions(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index fe2fb19e10b..2c55ae56400 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -239,7 +239,7 @@ class MCPClient: server_params = StdioServerParameters( command=self.stdio_config.get("command", ""), args=self.stdio_config.get("args", []), - env=self.stdio_config.get("env", None), + env=self._get_safe_stdio_env(self.stdio_config.get("env")), ) return stdio_client(server_params), None if self.transport_type == MCPTransport.sse: @@ -273,6 +273,53 @@ class MCPClient: ) return transport_ctx, http_client + def _get_safe_stdio_env( + self, provided_env: Optional[Dict[str, str]] + ) -> Optional[Dict[str, str]]: + """ + Return a safe environment for the stdio subprocess. + + If provided_env is set, we use it as-is. + If provided_env is None, we return a minimal allowlist from the parent environment + to avoid leaking sensitive LiteLLM keys (OPENAI_API_KEY, etc.) to sub-processes. + """ + if provided_env is not None: + return provided_env + + import os + + # Minimal allowlist of safe/standard environment variables + safe_keys = { + "PATH", + "HOME", + "USER", + "LOGNAME", + "TMPDIR", + "TMP", + "TEMP", + "SHELL", + "LANG", + "LC_ALL", + # Node/Package manager caches + "NPM_CONFIG_CACHE", + "PNPM_HOME", + "XDG_CACHE_HOME", + "XDG_CONFIG_HOME", + "XDG_DATA_HOME", + # System info + "SYSTEMROOT", + "COMSPEC", + "PATHEXT", + "WINDIR", + } + + safe_env = {} + for key in safe_keys: + if key in os.environ: + safe_env[key] = os.environ[key] + + return safe_env + async def _execute_session_operation( self, transport_ctx: Any, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9e8eea04140..4ad63648768 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2559,17 +2559,24 @@ if MCP_AVAILABLE: return user_api_key_auth.model_copy(update={"object_permission": updated_op}) async def _render_mcp_error( - e: Exception, scope: Scope, receive: Receive, send: Send + e: Exception, + scope: Scope, + receive: Receive, + send: Send, + error_prefix: Optional[str] = None, ) -> None: """Render an exception as a JSON response for ASGI handlers.""" from starlette.responses import JSONResponse status_code = 500 detail = str(e) + headers = {} if isinstance(e, HTTPException): status_code = e.status_code detail = e.detail + if e.headers: + headers.update(e.headers) elif hasattr(e, "status_code"): status_code = getattr(e, "status_code") elif hasattr(e, "code"): @@ -2579,9 +2586,14 @@ if MCP_AVAILABLE: except (ValueError, TypeError): status_code = 500 + error_msg = error_prefix or "MCP request failed" + if error_prefix is None and status_code in (401, 403): + error_msg = "Authentication processing failed" + error_response = JSONResponse( status_code=status_code, - content={"error": "MCP request failed", "details": detail}, + content={"error": error_msg, "details": detail}, + headers=headers, ) await error_response(scope, receive, send) @@ -2773,7 +2785,9 @@ if MCP_AVAILABLE: _captured_session_id_container_var.reset(_capture_token) except Exception as e: - await _render_mcp_error(e, scope, receive, send) + await _render_mcp_error( + e, scope, receive, send, error_prefix="Authentication processing failed" + ) # No need to return Response for raw ASGI app. async def handle_sse_post_messages( @@ -2835,7 +2849,9 @@ if MCP_AVAILABLE: ) await sse.handle_post_message(scope, receive, send) except Exception as e: - await _render_mcp_error(e, scope, receive, send) + await _render_mcp_error( + e, scope, receive, send, error_prefix="Authentication processing failed" + ) def get_active_mcp_session() -> Optional[_McpServerSession]: """Get the active downstream MCP session from the current context.""" diff --git a/tests/mcp_tests/test_sampling_handler.py b/tests/mcp_tests/test_sampling_handler.py index 846a1861c7f..1f7c58eb236 100644 --- a/tests/mcp_tests/test_sampling_handler.py +++ b/tests/mcp_tests/test_sampling_handler.py @@ -13,42 +13,89 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import ( # Mock MCP types if not available try: from mcp.types import ( - ModelPreferences, ModelHint, SamplingMessage, TextContent, - ImageContent, Tool, ToolChoice, CreateMessageRequestParams, - ToolUseContent, ToolResultContent + ModelPreferences, + ModelHint, + SamplingMessage, + TextContent, + ImageContent, + Tool, + ToolChoice, + CreateMessageRequestParams, + ToolUseContent, + ToolResultContent, ) except ImportError: # Minimal mocks for testing when mcp package is not installed class ModelHint: - def __init__(self, name=None): self.name = name + def __init__(self, name=None): + self.name = name + class ModelPreferences: - def __init__(self, hints=None): self.hints = hints + def __init__(self, hints=None): + self.hints = hints + 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 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 Tool: def __init__(self, name, description=None, inputSchema=None): - self.name = name; self.description = description; self.inputSchema = inputSchema + self.name = name + self.description = description + self.inputSchema = inputSchema + class ToolChoice: - def __init__(self, mode="auto"): self.mode = mode + def __init__(self, mode="auto"): + self.mode = mode + class CreateMessageRequestParams: - def __init__(self, messages, modelPreferences=None, systemPrompt=None, - maxTokens=None, temperature=None, stopSequences=None, - tools=None, toolChoice=None, metadata=None): - self.messages = messages; self.modelPreferences = modelPreferences - self.systemPrompt = systemPrompt; self.maxTokens = maxTokens - self.temperature = temperature; self.stopSequences = stopSequences - self.tools = tools; self.toolChoice = toolChoice; self.metadata = metadata + def __init__( + self, + messages, + modelPreferences=None, + systemPrompt=None, + maxTokens=None, + temperature=None, + stopSequences=None, + tools=None, + toolChoice=None, + metadata=None, + ): + self.messages = messages + self.modelPreferences = modelPreferences + self.systemPrompt = systemPrompt + self.maxTokens = maxTokens + self.temperature = temperature + self.stopSequences = stopSequences + self.tools = tools + self.toolChoice = toolChoice + self.metadata = metadata + 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 + def test_resolve_model_from_preferences(): # Test 1: Direct match @@ -66,6 +113,7 @@ def test_resolve_model_from_preferences(): # Test 3: Default fallback assert _resolve_model_from_preferences(None, default_model="fallback") == "fallback" + def test_convert_mcp_content_to_openai(): # Text content text = TextContent(type="text", text="hello") @@ -75,7 +123,7 @@ def test_convert_mcp_content_to_openai(): img = ImageContent(type="image", data="base64data", mimeType="image/jpeg") assert _convert_mcp_content_to_openai(img) == { "type": "image_url", - "image_url": {"url": "data:image/jpeg;base64,base64data"} + "image_url": {"url": "data:image/jpeg;base64,base64data"}, } # List of content @@ -85,10 +133,13 @@ def test_convert_mcp_content_to_openai(): assert result[0]["type"] == "text" assert result[1]["type"] == "image_url" + def test_convert_mcp_messages_to_openai(): msg1 = SamplingMessage(role="user", content=TextContent(type="text", text="hi")) - msg2 = SamplingMessage(role="assistant", content=TextContent(type="text", text="hello")) - + msg2 = SamplingMessage( + role="assistant", content=TextContent(type="text", text="hello") + ) + # Standard messages openai_msgs = _convert_mcp_messages_to_openai([msg1, msg2], system_prompt="system") assert len(openai_msgs) == 3 @@ -97,9 +148,14 @@ def test_convert_mcp_messages_to_openai(): assert openai_msgs[2]["role"] == "assistant" # Tool use/result conversion - tool_use = ToolUseContent(type="tool_use", id="call_1", name="search", input={"q": "test"}) - msg_tool_use = SamplingMessage(role="assistant", content=[TextContent(type="text", text="searching..."), tool_use]) - + tool_use = ToolUseContent( + type="tool_use", id="call_1", name="search", input={"q": "test"} + ) + msg_tool_use = SamplingMessage( + role="assistant", + content=[TextContent(type="text", text="searching..."), tool_use], + ) + openai_msgs = _convert_mcp_messages_to_openai([msg_tool_use]) assert len(openai_msgs) == 1 assert openai_msgs[0]["role"] == "assistant" @@ -107,7 +163,11 @@ def test_convert_mcp_messages_to_openai(): assert openai_msgs[0]["tool_calls"][0]["function"]["name"] == "search" assert openai_msgs[0]["content"] == "searching..." - tool_result = ToolResultContent(type="tool_result", toolUseId="call_1", content=[TextContent(type="text", text="found it")]) + tool_result = ToolResultContent( + type="tool_result", + toolUseId="call_1", + content=[TextContent(type="text", text="found it")], + ) msg_tool_result = SamplingMessage(role="user", content=[tool_result]) openai_msgs = _convert_mcp_messages_to_openai([msg_tool_result]) assert len(openai_msgs) == 1 @@ -115,6 +175,7 @@ def test_convert_mcp_messages_to_openai(): assert openai_msgs[0]["tool_call_id"] == "call_1" assert openai_msgs[0]["content"] == "found it" + def test_convert_mcp_tools_to_openai(): mcp_tool = Tool(name="my_tool", description="desc", inputSchema={"type": "object"}) openai_tools = _convert_mcp_tools_to_openai([mcp_tool]) @@ -122,47 +183,66 @@ def test_convert_mcp_tools_to_openai(): assert openai_tools[0]["type"] == "function" assert openai_tools[0]["function"]["name"] == "my_tool" + def test_convert_mcp_tool_choice_to_openai(): assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="auto")) == "auto" assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="required")) == "required" assert _convert_mcp_tool_choice_to_openai(ToolChoice(mode="none")) == "none" + @pytest.mark.asyncio async def test_handle_sampling_create_message_success(): 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, ) - + mock_response = MagicMock() - mock_response.choices = [MagicMock(message=MagicMock(content="hello response", tool_calls=None), finish_reason="stop")] + mock_response.choices = [ + MagicMock( + message=MagicMock(content="hello response", tool_calls=None), + finish_reason="stop", + ) + ] mock_response.model = "gpt-4o-mini" with patch("litellm.acompletion", new_callable=AsyncMock) as mock_completion: mock_completion.return_value = mock_response result = await handle_sampling_create_message(context=None, params=params) - + assert result.role == "assistant" assert result.content.text == "hello response" assert result.model == "gpt-4o-mini" + @pytest.mark.asyncio async def test_handle_sampling_with_auth_cost_tracking(): from litellm.proxy._types import UserAPIKeyAuth + 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, ) user_auth = UserAPIKeyAuth(api_key="sk-123", user_id="user-456", team_id="team-789") - + mock_response = MagicMock() - mock_response.choices = [MagicMock(message=MagicMock(content="ok", tool_calls=None), finish_reason="stop")] + mock_response.choices = [ + MagicMock( + message=MagicMock(content="ok", tool_calls=None), finish_reason="stop" + ) + ] mock_response.model = "gpt-4o-mini" with patch("litellm.acompletion", new_callable=AsyncMock) as mock_completion: mock_completion.return_value = mock_response - await handle_sampling_create_message(context=None, params=params, user_api_key_auth=user_auth) - + await handle_sampling_create_message( + context=None, params=params, user_api_key_auth=user_auth + ) + # Verify auth was injected into metadata kwargs = mock_completion.call_args.kwargs assert kwargs["user"] == "user-456" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py index 44d9bcf07f7..f541af3ec9f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py @@ -3,28 +3,37 @@ Unit tests for the MCP Elicitation Handler. Tests the elicitation/create handler that relays elicitation requests from upstream MCP servers to downstream clients or declines them. """ + from unittest.mock import AsyncMock, MagicMock, patch import pytest + + # ───────────────────────────────────────────────────────────── # Helper factories # ───────────────────────────────────────────────────────────── def _make_form_params(message="Please provide info", schema=None): """Create a mock ElicitRequestFormParams.""" from mcp.types import ElicitRequestFormParams + params = MagicMock(spec=ElicitRequestFormParams) params.mode = "form" params.message = message params.requestedSchema = schema return params + + def _make_url_params(message="Click the link", url="https://auth.example.com"): """Create a mock ElicitRequestURLParams.""" from mcp.types import ElicitRequestURLParams + params = MagicMock(spec=ElicitRequestURLParams) params.mode = "url" params.message = message params.url = url params.elicitationId = "elicit-123" return params + + def _make_capabilities(form=True, url=True): """Create mock client capabilities with elicitation support.""" caps = MagicMock() @@ -33,21 +42,27 @@ def _make_capabilities(form=True, url=True): elicit.url = MagicMock() if url else None caps.elicitation = elicit return caps + + def _make_capabilities_no_elicitation(): """Create mock client capabilities without elicitation.""" caps = MagicMock() caps.elicitation = None return caps + + # ───────────────────────────────────────────────────────────── # Tests: No downstream session (Tool Bridge mode) # ───────────────────────────────────────────────────────────── class TestElicitationNoDownstream: """Tests when no downstream client is available.""" + @pytest.mark.asyncio async def test_should_decline_when_no_downstream_session(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params() result = await handle_elicitation_request( context=MagicMock(), @@ -55,16 +70,20 @@ class TestElicitationNoDownstream: downstream_session=None, ) assert result.action == "decline" + + # ───────────────────────────────────────────────────────────── # Tests: Downstream session relay # ───────────────────────────────────────────────────────────── class TestElicitationRelay: """Tests for relaying elicitation to downstream clients.""" + @pytest.mark.asyncio async def test_should_relay_form_mode(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params(message="Enter your name") mock_session = AsyncMock() mock_result = MagicMock() @@ -80,11 +99,13 @@ class TestElicitationRelay: ) mock_session.elicit_form.assert_called_once() assert result.action == "submit" + @pytest.mark.asyncio async def test_should_relay_url_mode(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_url_params( message="Authenticate", url="https://oauth.example.com" ) @@ -101,11 +122,13 @@ class TestElicitationRelay: ) mock_session.elicit_url.assert_called_once() assert result.action == "submit" + @pytest.mark.asyncio async def test_should_decline_when_client_lacks_elicitation(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params() mock_session = AsyncMock() caps = _make_capabilities_no_elicitation() @@ -116,11 +139,13 @@ class TestElicitationRelay: downstream_capabilities=caps, ) assert result.action == "decline" + @pytest.mark.asyncio async def test_should_decline_when_client_lacks_url_mode(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_url_params() mock_session = AsyncMock() caps = _make_capabilities(form=True, url=False) @@ -131,11 +156,13 @@ class TestElicitationRelay: downstream_capabilities=caps, ) assert result.action == "decline" + @pytest.mark.asyncio async def test_should_decline_when_client_lacks_form_mode(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params() mock_session = AsyncMock() caps = _make_capabilities(form=False, url=True) @@ -146,11 +173,13 @@ class TestElicitationRelay: downstream_capabilities=caps, ) assert result.action == "decline" + @pytest.mark.asyncio async def test_should_decline_gracefully_on_relay_failure(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params() mock_session = AsyncMock() mock_session.elicit_form.side_effect = Exception("Connection lost") @@ -162,16 +191,20 @@ class TestElicitationRelay: downstream_capabilities=caps, ) assert result.action == "decline" + + # ───────────────────────────────────────────────────────────── # Tests: Error handling # ───────────────────────────────────────────────────────────── class TestElicitationErrorHandling: """Tests for error handling in the elicitation handler.""" + @pytest.mark.asyncio async def test_should_relay_without_capability_check_when_caps_none(self): from litellm.proxy._experimental.mcp_server.elicitation_handler import ( handle_elicitation_request, ) + params = _make_form_params() mock_session = AsyncMock() mock_result = MagicMock() @@ -185,4 +218,4 @@ class TestElicitationErrorHandling: downstream_capabilities=None, ) mock_session.elicit_form.assert_called_once() - assert result.action == "submit" \ No newline at end of file + assert result.action == "submit" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py index 577ee0f22f7..da12453c0ad 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sampling_handler.py @@ -201,7 +201,9 @@ class TestResolveModel: with patch("litellm.model_list", []): # Simulate configured default model - with patch.object(litellm, "default_mcp_sampling_model", "claude-3-haiku", create=True): + with patch.object( + litellm, "default_mcp_sampling_model", "claude-3-haiku", create=True + ): result = _resolve_model_from_preferences(None) assert result == "claude-3-haiku" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index e3c4c957ab1..4ce3ca90b36 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -317,6 +317,7 @@ class TestMCPServerManager: mcp_auth_header=None, mcp_protocol_version=None, raw_headers=None, + **kwargs, ): if server.name == "github": tool1 = MagicMock() @@ -371,6 +372,7 @@ class TestMCPServerManager: mcp_auth_header=None, mcp_protocol_version=None, raw_headers=None, + **kwargs, ): assert mcp_auth_header == "legacy-token" # Should use legacy header tool = MagicMock() @@ -409,6 +411,7 @@ class TestMCPServerManager: mcp_auth_header=None, mcp_protocol_version=None, raw_headers=None, + **kwargs, ): assert ( mcp_auth_header == "server-specific-token" @@ -990,6 +993,7 @@ class TestMCPServerManager: mcp_auth_header=None, mcp_protocol_version=None, raw_headers=None, + **kwargs, ): assert ( mcp_auth_header == "server-specific-token" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index a0bfbff4222..27f57edbccb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -497,37 +497,52 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value=None, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: - with pytest.raises(HTTPException) as exc_info: - await handle_streamable_http_mcp(scope, receive, send) + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value=None, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): + await handle_streamable_http_mcp(scope, receive, send) - exc = exc_info.value - assert exc.status_code == 401 - assert "www-authenticate" in exc.headers + # Verify 401 response was sent + assert send.called + # Extract status code from mock_send + response_start = next( + call.args[0] + for call in send.mock_calls + if call.args[0].get("type") == "http.response.start" + ) + assert response_start["status"] == 401 + headers = dict(response_start["headers"]) + assert b"www-authenticate" in headers assert mock_get_stored_token.await_count == 1 assert mock_handle_request.await_count == 0 @@ -562,31 +577,39 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value={"Authorization": "Bearer cached-token"}, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value={"Authorization": "Bearer cached-token"}, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): await handle_streamable_http_mcp(scope, receive, send) assert mock_get_stored_token.await_count == 1 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 2558df8533b..2888077a23a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -489,9 +489,7 @@ class TestGetToolsByNames: {"name": "send_email", "description": "send mail"}, ] - matched = filter_instance._get_tools_by_names( - ["send_email"], available_tools - ) + matched = filter_instance._get_tools_by_names(["send_email"], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "send_email" @@ -503,9 +501,7 @@ class TestGetToolsByNames: client_name = "litellm_" + canonical available_tools = [{"name": client_name, "description": "scrape"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 # Must return the incoming tool unchanged so the client-facing @@ -516,13 +512,9 @@ class TestGetToolsByNames: """Some clients use dash as alias separator; accept that too.""" filter_instance = self._make_filter() canonical = "weather_svc-get_weather" - available_tools = [ - {"name": "mcp-" + canonical, "description": "weather"} - ] + available_tools = [{"name": "mcp-" + canonical, "description": "weather"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "mcp-" + canonical @@ -552,9 +544,7 @@ class TestGetToolsByNames: {"name": "litellm_" + canonical, "description": "wrapped"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == canonical @@ -567,9 +557,7 @@ class TestGetToolsByNames: separator-anchored suffixes of ``litellm_api-fs-read_file``. """ filter_instance = self._make_filter() - available_tools = [ - {"name": "litellm_api-fs-read_file", "description": "read"} - ] + available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}] matched = filter_instance._get_tools_by_names( ["fs-read_file", "api-fs-read_file"], available_tools @@ -590,9 +578,7 @@ class TestGetToolsByNames: {"name": "my_" + canonical, "description": "plain search"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "my_" + canonical