[MCP Gateway] Allow MCP sse and http to have namespaced url for better segregation LIT-304 (#12658)

* fix tools fetch for keys

* Add namespacing in url

* add test for namespacing url

* helper method

* fix test
This commit is contained in:
Jugal D. Bhatt 2025-07-17 03:17:45 +05:30 • committed by GitHub
parent a7d0b122b9
commit 83b0c4cba7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 78 additions and 14 deletions

View file

@ -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)

View file

@ -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"]