From 4b1eb39b4da46ca9301f2704b5e38e3611a727dc Mon Sep 17 00:00:00 2001 From: Yug Date: Wed, 29 Apr 2026 13:37:17 +0530 Subject: [PATCH] erros resolve --- .../proxy/_experimental/mcp_server/server.py | 2 + tests/mcp_tests/test_mcp_server.py | 108 +++++++++++++++--- .../mcp_server/test_mcp_hook_extra_headers.py | 1 - 3 files changed, 96 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a7f68f00b1c..f49ee7fda4e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2710,6 +2710,8 @@ if MCP_AVAILABLE: await server.run(streams[0], streams[1], options) except Exception as session_e: verbose_logger.exception(f"Error in SSE session: {session_e}") + except HTTPException: + raise except Exception as e: verbose_logger.exception(f"Error handling MCP request: {e}") # Instead of re-raising, try to send a graceful error response diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index db1830a3341..19ea5715e57 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -3,7 +3,6 @@ import os import sys import pytest from unittest.mock import AsyncMock, MagicMock, patch -from contextlib import asynccontextmanager sys.path.insert( 0, os.path.abspath("../../..") @@ -15,7 +14,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPTransport, ) from litellm.proxy._types import LiteLLM_ObjectPermissionTable -from mcp.types import Tool as MCPTool, CallToolResult, ListToolsResult +from mcp.types import Tool as MCPTool, CallToolResult from mcp.types import TextContent @@ -32,13 +31,11 @@ async def test_mcp_server_manager(): } } ) - tools = await mcp_server_manager.list_tools() - print("TOOLS FROM MCP SERVER MANAGER== ", tools) + await mcp_server_manager.list_tools() - result = await mcp_server_manager.call_tool( + await mcp_server_manager.call_tool( name="gmail_send_email", arguments={"body": "Test"}, proxy_logging_obj=None ) - print("RESULT FROM CALLING TOOL FROM MCP SERVER MANAGER== ", result) @pytest.mark.asyncio @@ -96,7 +93,6 @@ async def test_mcp_server_manager_https_server(): new=AsyncMock(return_value=allowed_server_ids), ): tools = await mcp_server_manager.list_tools() - print("TOOLS FROM MCP SERVER MANAGER== ", tools) # Verify tools were returned and properly prefixed assert len(tools) == 1 @@ -122,7 +118,6 @@ async def test_mcp_server_manager_https_server(): }, proxy_logging_obj=None, ) - print("RESULT FROM CALLING TOOL FROM MCP SERVER MANAGER== ", result) # Verify result assert result.isError is False @@ -395,7 +390,6 @@ async def test_mcp_http_transport_tool_not_found(): @pytest.mark.asyncio async def test_streamable_http_mcp_handler_mock(): """Test the streamable HTTP MCP handler functionality""" - from litellm.proxy._types import UserAPIKeyAuth # Mock the session manager and its methods mock_session_manager = AsyncMock() @@ -445,6 +439,96 @@ async def test_streamable_http_mcp_handler_mock(): mock_session_manager.handle_request.assert_called_once() +@pytest.mark.asyncio +async def test_sse_mcp_handler_mock(): + """Test the SSE MCP handler functionality""" + mock_sse_server = AsyncMock() + mock_sse_server.handle_sse = AsyncMock() + + mock_scope = { + "type": "http", + "method": "GET", + "path": "/sse", + "headers": [(b"accept", b"text/event-stream")], + "query_string": b"", + "server": ("localhost", 8000), + "scheme": "http", + } + mock_receive = AsyncMock() + mock_send = AsyncMock() + + mock_auth_context = (None, None, None, {}, {}, {}) + + mock_sse = MagicMock() + mock_sse.connect_sse = MagicMock() + + # Mock connect_sse to return an async context manager yielding dummy streams + mock_context_manager = AsyncMock() + mock_context_manager.__aenter__.return_value = (AsyncMock(), AsyncMock()) + mock_sse.connect_sse.return_value = mock_context_manager + + mock_server = MagicMock() + mock_server.run = AsyncMock() + mock_server.create_initialization_options = MagicMock() + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.sse", + mock_sse, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.server", + mock_server, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + AsyncMock(return_value=mock_auth_context), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + ): + from litellm.proxy._experimental.mcp_server.server import ( + handle_sse_mcp_endpoint, + ) + + await handle_sse_mcp_endpoint(mock_scope, mock_receive, mock_send) + mock_sse.connect_sse.assert_called_once_with(mock_scope, mock_receive, mock_send) + mock_server.run.assert_called_once() + + +@pytest.mark.asyncio +async def test_sse_post_messages_auth_failure(): + """Test that handle_sse_post_messages correctly propagates HTTPException on auth failure""" + mock_scope = { + "type": "http", + "method": "POST", + "path": "/messages", + "headers": [(b"content-type", b"application/json")], + "query_string": b"sessionId=test-id", + "server": ("localhost", 8000), + "scheme": "http", + } + mock_receive = AsyncMock() + mock_send = AsyncMock() + + from fastapi import HTTPException + + # Make extract_mcp_auth_context raise an HTTPException + with patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + AsyncMock(side_effect=HTTPException(status_code=401, detail="Unauthorized")), + ): + from litellm.proxy._experimental.mcp_server.server import ( + handle_sse_post_messages, + ) + + with pytest.raises(HTTPException) as exc_info: + await handle_sse_post_messages(mock_scope, mock_receive, mock_send) + + assert exc_info.value.status_code == 401 + assert exc_info.value.detail == "Unauthorized" + def test_generate_stable_server_id(): """ @@ -603,9 +687,7 @@ async def test_list_tools_rest_api_server_not_found(): """Test the list_tools REST API when server is not found""" from litellm.proxy._experimental.mcp_server.rest_endpoints import ( list_tool_rest_api, - global_mcp_server_manager, ) - from fastapi import Query from litellm.proxy._types import UserAPIKeyAuth # Mock UserAPIKeyAuth with explicit permission to access the requested server id @@ -661,7 +743,6 @@ async def test_list_tools_rest_api_success(): from litellm.proxy._experimental.mcp_server.server import ( ListMCPToolsRestAPIResponseObject, ) - from fastapi import Query from litellm.proxy._types import UserAPIKeyAuth # Mock successful tools @@ -729,7 +810,7 @@ async def test_list_tools_rest_api_success(): @pytest.mark.asyncio -async def test_get_tools_from_mcp_servers(): +async def test_get_tools_from_mcp_servers(): # noqa: PLR0915 """Test _get_tools_from_mcp_servers function with both specific and no server filters""" from litellm.proxy._experimental.mcp_server.server import ( _get_tools_from_mcp_servers, @@ -1009,7 +1090,6 @@ async def test_mcp_server_manager_access_groups_from_config(): mcp_server_manager_mod.global_mcp_server_manager = test_manager try: # Should find config_server for group-a, both for group-b, other_server for group-c - import asyncio server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups( ["group-a"] 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 84c556b8ddc..084bc75c07d 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 @@ -16,7 +16,6 @@ from typing import Any, Dict, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest -from fastapi import HTTPException from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.proxy._types import UserAPIKeyAuth