mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
[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:
parent
a7d0b122b9
commit
83b0c4cba7
2 changed files with 78 additions and 14 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue