mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): stop one unauthenticated server from emptying the aggregate tools/list (#31684)
* fix(mcp): stop one unauthenticated server from emptying the aggregate tools/list On the aggregate MCP route (/mcp), the gateway fans out to every server the caller can access and flattens their tools. _fetch_and_filter_server_tools re-raises MCPUpstreamAuthError unconditionally (added with the OAuth passthrough feature in #28356) so it surfaces a 401 on single-server routes, but on the aggregate route that exception propagates through the asyncio.gather fan-out and the outer handler turns it into an empty list. The result: a single delegate/passthrough OAuth server the user has not authenticated (e.g. a delegate-auth server) zeroes the tools of every other server, including the ones that resolve fine, so the client connects and sees no tools. Surface the upstream auth error only when a single server was explicitly targeted (so that route still drives the upstream OAuth flow); across the aggregate, absorb it to [] for that one server so the rest still list their tools. This restores the graceful per-server degradation that predated #28356. Adds regression tests: the aggregate keeps a healthy server's tools when a sibling raises MCPUpstreamAuthError, and a single-server listing still surfaces it. * fix(mcp): decide aggregate vs single-server listing by route scope, not server count Addresses review: keying the surface-vs-absorb decision off the server count (len(allowed_mcp_servers), and even len(mcp_servers)) misclassifies an aggregate /mcp request from a key that can access exactly one server as a targeted single-server listing, so that one server's MCPUpstreamAuthError re-raises and empties the aggregate again for one-server permission sets. Use the path-derived single-server scope instead: _mcp_gateway_server_name, set by _gateway_initialize_instructions_request_scope only when the request path names exactly one upstream server (/<server>/mcp) and never from client headers, is None on the aggregate route (/mcp) regardless of how many servers the key can access. Single-server routes still surface the upstream-auth challenge; the aggregate absorbs it per server. Adds a regression test that an aggregate request with a single accessible server still absorbs, plus renames the single-server test to drive the route scope explicitly. The new test fails on the count-based logic. * fixing aggregation error * style(mcp): collapse single-line debug log to satisfy ruff format
This commit is contained in:
parent
fecaf5c9e5
commit
87de0e80a8
2 changed files with 135 additions and 6 deletions
|
|
@ -1699,12 +1699,13 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return filtered_tools
|
||||
except MCPUpstreamAuthError:
|
||||
# Surface upstream 401/403 to the outer handler so the
|
||||
# client receives a proper WWW-Authenticate challenge
|
||||
# instead of a silently empty tool list. Without this
|
||||
# re-raise the broad ``except Exception`` below would
|
||||
# swallow the auth error.
|
||||
raise
|
||||
# Absorb so one unauthenticated server does not empty every other server's
|
||||
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
|
||||
# intentionally not done here: raising from this list handler cannot produce a
|
||||
# 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC
|
||||
# error), so that belongs in a request-scope preemptive check, tracked separately.
|
||||
verbose_logger.debug(f"MCP list_tools: omitting {server.name}; it needs upstream auth")
|
||||
return []
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}")
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -265,3 +265,131 @@ async def test_fetch_tools_from_gateway_managed_swallows_errors():
|
|||
)
|
||||
assert tools == []
|
||||
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
|
||||
|
||||
|
||||
def _http_server(server_id: str, name: str, **kwargs) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
url=f"https://{name}/mcp",
|
||||
transport=MCPTransport.http,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_list_tools_absorbs_one_unauthenticated_server():
|
||||
"""Regression: across the aggregate (/mcp), a delegate/passthrough server that raises
|
||||
MCPUpstreamAuthError must not empty every other server's tools. Re-raising it on the
|
||||
aggregate path (introduced with the passthrough feature) zeroed the whole list because the
|
||||
fan-out gather propagated it."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
delegate = _http_server(
|
||||
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
|
||||
)
|
||||
working = _http_server("s2", "working_docs", auth_type=MCPAuth.none)
|
||||
good_tool = MCPTool(name="working_docs-read", description="d", inputSchema={"type": "object"})
|
||||
|
||||
async def fake_get_tools(server, **kwargs):
|
||||
if server.server_id == delegate.server_id:
|
||||
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
|
||||
return [good_tool]
|
||||
|
||||
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate, working])), patch.object(
|
||||
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
|
||||
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
|
||||
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
|
||||
), patch.object(
|
||||
mcp_server, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools)
|
||||
), patch.object(
|
||||
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
|
||||
):
|
||||
tools = await mcp_server._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
)
|
||||
|
||||
assert [t.name for t in tools] == ["working_docs-read"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_server_route_also_absorbs_upstream_auth_error():
|
||||
"""A single-server route (/<server>/mcp) absorbs an upstream-auth error just like the aggregate:
|
||||
the failing server is omitted (empty list) rather than re-raised. Surfacing it to the client as a
|
||||
401 + WWW-Authenticate challenge cannot be done from this list handler — the MCP session manager
|
||||
serializes a raise into a JSON-RPC error, not an HTTP 401 — so re-auth surfacing is handled by a
|
||||
request-scope preemptive check, tracked separately."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_gateway_server_name
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
delegate = _http_server(
|
||||
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
|
||||
)
|
||||
|
||||
async def fake_get_tools(server, **kwargs):
|
||||
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
|
||||
|
||||
# /<server>/mcp sets the path-derived single-server scope; absorption must hold even then.
|
||||
token = _mcp_gateway_server_name.set("delegate_docs")
|
||||
try:
|
||||
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object(
|
||||
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
|
||||
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
|
||||
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
|
||||
), patch.object(
|
||||
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
|
||||
):
|
||||
tools = await mcp_server._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=["delegate_docs"],
|
||||
)
|
||||
assert tools == []
|
||||
finally:
|
||||
_mcp_gateway_server_name.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_with_single_accessible_server_still_absorbs():
|
||||
"""Regression for the route-misclassification: an aggregate request (/mcp, mcp_servers=None)
|
||||
from a key that can access exactly one server must still absorb that server's
|
||||
MCPUpstreamAuthError, not surface it. Keying the surface decision off the allowed count rather
|
||||
than the request filter would re-raise here and leave the aggregate broken for one-server
|
||||
permission sets."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
delegate = _http_server(
|
||||
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
|
||||
)
|
||||
|
||||
async def fake_get_tools(server, **kwargs):
|
||||
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
|
||||
|
||||
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object(
|
||||
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
|
||||
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
|
||||
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
|
||||
), patch.object(
|
||||
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
|
||||
):
|
||||
# Aggregate route: no explicit server filter, even though only one server is accessible.
|
||||
tools = await mcp_server._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
)
|
||||
|
||||
assert tools == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue