mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
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:
parent
acc74afc51
commit
21cc86e93e
4 changed files with 213 additions and 13 deletions
|
|
@ -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 {}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue