diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 3d347f4a470..f10263ba571 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1,6 +1,6 @@ import importlib from datetime import datetime, timezone -from typing import Any, Awaitable, Callable, Dict, List, Optional, Union +from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, Union from fastapi import APIRouter, Depends, HTTPException, Query, Request @@ -69,6 +69,21 @@ if MCP_AVAILABLE: return server_auth return mcp_auth_header + def _get_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]: + """Return the subset of *allowed_server_ids* whose servers use OAuth2 auth. + + Used as a cheap pre-flight check to skip bulk credential fetching when no + OAuth2 servers are involved in the current request. + """ + return { + sid + for sid in allowed_server_ids + if getattr( + global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None + ) + == MCPAuth.oauth2 + } + async def _get_user_oauth_extra_headers( server, user_api_key_dict: UserAPIKeyAuth, @@ -506,15 +521,9 @@ if MCP_AVAILABLE: # Pre-fetch OAuth credentials only when at least one allowed server uses OAuth2, # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. - _oauth2_server_ids = { - sid for sid in allowed_server_ids - if getattr( - global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None - ) == MCPAuth.oauth2 - } prefetched_oauth_creds = ( await _prefetch_user_oauth_creds(user_api_key_dict) - if _oauth2_server_ids + if _get_oauth2_server_ids(allowed_server_ids) else {} ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 106e3618f0f..a9ff61dab5b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1342,6 +1342,17 @@ if MCP_AVAILABLE: mcp_server = _get_cached_temporary_mcp_server_or_404(server_id) # Use the server's stored client_id when the caller doesn't supply one resolved_client_id = mcp_server.client_id or client_id or "" + if not resolved_client_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "missing_client_id", + "message": ( + "No client_id available for this MCP server. " + "Either configure the server with a client_id or supply one in the request." + ), + }, + ) return await authorize_with_server( request=request, mcp_server=mcp_server, @@ -1370,6 +1381,17 @@ if MCP_AVAILABLE: ): mcp_server = _get_cached_temporary_mcp_server_or_404(server_id) resolved_client_id = mcp_server.client_id or client_id or "" + if not resolved_client_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "missing_client_id", + "message": ( + "No client_id available for this MCP server. " + "Either configure the server with a client_id or supply one in the request." + ), + }, + ) return await exchange_token_with_server( request=request, mcp_server=mcp_server, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 254faa57b6e..7a3934ffdaa 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -105,10 +105,10 @@ class LiteLLM_Proxy_MCP_Handler: # by rewriting them to the internal litellm_proxy format. m = _PROXY_MCP_PATH_RE.match(server_url) if m: - rewritten = {**tool, "server_url": ( - f"{LITELLM_PROXY_MCP_SERVER_URL_PREFIX}{m.group(1)}" - )} - mcp_tools_with_litellm_proxy.append(rewritten) + rewritten = { + **tool, + "server_url": f"{LITELLM_PROXY_MCP_SERVER_URL_PREFIX}{m.group(1)}", + } mcp_tools_with_litellm_proxy.append(rewritten) else: other_tools.append(tool) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 4f93270c162..1d296f0440c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -484,7 +484,7 @@ class TestListToolsRestAPI: captured = {"called": False} async def fake_get_tools( - server, server_auth_header, raw_headers=None, user_api_key_auth=None + server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None ): captured["called"] = True captured["server"] = server @@ -529,6 +529,175 @@ class TestListToolsRestAPI: assert result["error"] is None assert result["message"] == "Successfully retrieved tools" + async def test_name_resolution_finds_server_by_uuid(self, monkeypatch): + """When server_id is a name string, it should be resolved to its UUID + and used for the tools lookup when the UUID is in allowed_server_ids.""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + stub_server = MCPServer( + server_id="uuid-abc-123", + name="my-server", + transport=MCPTransport.sse, + ) + stub_server.alias = "my-server" + stub_server.server_name = "my-server" + stub_server.available_on_public_internet = True + stub_server.allowed_tools = None + stub_server.mcp_info = {"server_name": "my-server"} + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + # Allowed list contains the UUID, not the name + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["uuid-abc-123"] + + captured = {"called": False, "server_arg": None} + + async def fake_get_tools(server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None): + captured["called"] = True + captured["server_arg"] = server + return ["tool-x"] + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_name", + lambda name: stub_server if name == "my-server" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id", + lambda sid: stub_server if sid == "uuid-abc-123" else None, + raising=False, + ) + monkeypatch.setattr(rest_endpoints, "_get_tools_for_single_server", fake_get_tools, raising=False) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id="my-server", # pass name, not UUID + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert captured["called"] is True + assert captured["server_arg"] is stub_server + assert result["tools"] == ["tool-x"] + assert result["error"] is None + + async def test_name_not_in_allowed_returns_access_denied(self, monkeypatch): + """When name resolves to a server whose UUID is NOT in allowed_server_ids, + the result should be an access_denied error (not a crash or silent pass).""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + stub_server = MCPServer( + server_id="uuid-xyz-999", + name="restricted-server", + transport=MCPTransport.sse, + ) + stub_server.available_on_public_internet = True + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + # No allowed servers for this key + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return [] + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_name", + lambda name: stub_server if name == "restricted-server" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id", + lambda sid: stub_server if sid == "uuid-xyz-999" else None, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id="restricted-server", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result["tools"] == [] + assert result["error"] == "unexpected_error" + assert "access_denied" in result["message"] + + async def test_oauth2_user_token_injected_for_single_server(self, monkeypatch): + """For a single-server OAuth2 request, _get_user_oauth_extra_headers is called + and the returned headers are forwarded to _get_tools_for_single_server.""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + stub_server = MCPServer( + server_id="oauth-server-id", + name="oauth-server", + transport=MCPTransport.sse, + ) + stub_server.alias = "oauth-server" + stub_server.server_name = "oauth-server" + stub_server.available_on_public_internet = True + stub_server.allowed_tools = None + stub_server.mcp_info = {"server_name": "oauth-server"} + stub_server.auth_type = MCPAuth.oauth2 + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["oauth-server-id"] + + oauth_headers = {"Authorization": "Bearer user-oauth-token"} + + async def fake_get_user_oauth_extra_headers(server, user_api_key_dict, prefetched_creds=None): + return oauth_headers + + captured = {} + + async def fake_get_tools(server, server_auth_header, raw_headers=None, user_api_key_auth=None, extra_headers=None): + captured["server"] = server + captured["auth_header"] = server_auth_header + return ["oauth-tool"] + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, "get_mcp_server_by_id", + lambda sid: stub_server if sid == "oauth-server-id" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, "_get_user_oauth_extra_headers", + fake_get_user_oauth_extra_headers, raising=False, + ) + monkeypatch.setattr(rest_endpoints, "_get_tools_for_single_server", fake_get_tools, raising=False) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id="oauth-server-id", + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + + assert result["tools"] == ["oauth-tool"] + assert result["error"] is None + class TestCallToolRestAPI: pytestmark = pytest.mark.asyncio