fix: address greptile feedback - client_id guards, dict spread, helper refactor, tests

- mcp_management_endpoints: raise 400 when resolved_client_id is empty in
  mcp_authorize and mcp_token instead of forwarding "" to upstream
- litellm_proxy_mcp_handler: use {**tool, "server_url": ...} spread instead
  of dict(tool) + mutation for shallow copy safety
- rest_endpoints: extract _oauth2_server_ids set comprehension to a named
  _get_oauth2_server_ids() helper for clarity; add Set to typing imports
- test_rest_endpoints: add tests for name→UUID resolution path,
  access-denied when resolved UUID not in allowed list, and OAuth2 user
  token injection for single-server requests; fix fake_get_tools signature
  to accept extra_headers kwarg
This commit is contained in:
Ishaan Jaffer 2026-03-11 21:57:13 -07:00
parent acc74afc51
commit 21cc86e93e
4 changed files with 213 additions and 13 deletions

View file

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

View file

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

View file

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

View file

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