diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index df930f32224..7aab5eef987 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -478,16 +478,38 @@ if MCP_AVAILABLE: except Exception as e: return [TextContent(text=f"Error: {str(e)}", type="text")] + async def extract_mcp_auth_context(scope, path): + """ + Extracts mcp_servers from the path and processes the MCP request for auth context. + Returns: (user_api_key_auth, mcp_auth_header, mcp_servers) + """ + import re + mcp_servers_from_path = None + mcp_path_match = re.match(r"^/mcp/([^/]+)(/.*)?$", path) + if mcp_path_match: + mcp_servers_str = mcp_path_match.group(1) + if mcp_servers_str: + mcp_servers_from_path = [s.strip() for s in mcp_servers_str.split(",") if s.strip()] + + if mcp_servers_from_path is not None: + user_api_key_auth, mcp_auth_header, _ = ( + await MCPRequestHandler.process_mcp_request(scope) + ) + mcp_servers = mcp_servers_from_path + else: + user_api_key_auth, mcp_auth_header, mcp_servers = ( + await MCPRequestHandler.process_mcp_request(scope) + ) + return user_api_key_auth, mcp_auth_header, mcp_servers + async def handle_streamable_http_mcp( scope: Scope, receive: Receive, send: Send ) -> None: """Handle MCP requests through StreamableHTTP.""" try: - # Validate headers and log request info - user_api_key_auth, mcp_auth_header, mcp_servers = ( - await MCPRequestHandler.process_mcp_request(scope) - ) - verbose_logger.debug(f"MCP request headers - mcp_servers: {mcp_servers}") + path = scope.get("path", "") + user_api_key_auth, mcp_auth_header, mcp_servers = await extract_mcp_auth_context(scope, path) + verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}") # Set the auth context variable for easy access in MCP functions set_auth_context( user_api_key_auth=user_api_key_auth, @@ -509,22 +531,17 @@ if MCP_AVAILABLE: async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None: """Handle MCP requests through SSE.""" try: - # Validate headers and log request info - user_api_key_auth, mcp_auth_header, mcp_servers = ( - await MCPRequestHandler.process_mcp_request(scope) - ) - verbose_logger.debug(f"MCP request headers - mcp_servers: {mcp_servers}") - # Set the auth context variable for easy access in MCP functions + path = scope.get("path", "") + user_api_key_auth, mcp_auth_header, mcp_servers = await extract_mcp_auth_context(scope, path) + verbose_logger.debug(f"MCP request mcp_servers (header/path): {mcp_servers}") set_auth_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, ) - # Ensure session managers are initialized if not _SESSION_MANAGERS_INITIALIZED: await initialize_session_managers() - # Give it a moment to start up await asyncio.sleep(0.1) await sse_session_manager.handle_request(scope, receive, send) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 05e2b4142c7..2739e6f45a4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import orjson import pytest -from fastapi import Request +from fastapi import Request, FastAPI from fastapi.testclient import TestClient sys.path.insert( @@ -617,3 +617,50 @@ class TestMCPAccessGroupsE2E: # Verify the header parsing worked correctly assert auth_result.api_key == "test-api-key" assert mcp_servers == ["zapier-server", "dev-group"] # Should contain both server name and access group + + +@pytest.mark.asyncio +def test_mcp_path_based_server_segregation(monkeypatch): + # Import the MCP server FastAPI app and context getter + from litellm.proxy._experimental.mcp_server.server import app, get_auth_context + + captured_mcp_servers = {} + + # Patch the session manager to send a dummy response and capture context + async def dummy_handle_request(scope, receive, send): + from litellm.proxy._experimental.mcp_server.server import get_auth_context + _, _, mcp_servers = get_auth_context() + captured_mcp_servers[id(scope)] = mcp_servers + await send({ + "type": "http.response.start", + "status": 200, + "headers": [(b"content-type", b"application/json")], + }) + await send({ + "type": "http.response.body", + "body": b'{"ok": true}', + }) + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.session_manager", + MagicMock(handle_request=dummy_handle_request) + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.initialize_session_managers", + AsyncMock() + ) + + # Patch user_api_key_auth to always return a dummy user + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + AsyncMock(return_value=UserAPIKeyAuth(api_key="test", user_id="user")) + ) + + # Use TestClient to make a request to /mcp/zapier,group1/tools + client = TestClient(app) + response = client.get("/mcp/zapier,group1/tools", headers={"x-litellm-api-key": "test"}) + assert response.status_code == 200 + assert response.json() == {"ok": True} + + # The context should have mcp_servers set to ["zapier", "group1"] + assert list(captured_mcp_servers.values())[0] == ["zapier", "group1"]