fix(mcp): route a scoped name to the caller's granted server before the registry's pick

A key granted only the server named `docs` was refused with 403 on
/mcp/docs when an ungranted server held `docs` as its alias, and a key
granted only `github` was refused on /mcp/GITHUB when an ungranted
`GitHub` existed: the scoped router took the registry-wide pick and
answered "denied" whenever that pick was not among the caller's servers.

The router now runs the registry's own pass order (exact alias,
server_name, name, server_id; the same case-insensitively; prefix forms)
over the caller's granted servers first, through
get_mcp_server_answering_to(among=...), with IP hiding applied at the
pass that found the server. A name hidden from the client IP stays denied
before any grant lookup, and a name the registry places only on an
ungranted server stays denied rather than being retried as an access
group.

The connect preflight reuses the router's selection for a single scoped
name as the server it challenges, signs in and exchanges for, so a 401
names the granted server; an ungranted caller keeps the registry pick and
the downstream 403, and the no-key path is unchanged. Unauthenticated
RFC 9728 discovery has no grant list and stays on the registry pick.

Tests: the `only_b` assertion in
test_scoped_router_selects_the_server_the_connect_preflight_resolves and
the no-IP `internal` assertion in
test_scoped_router_hides_a_private_server_from_an_external_ip_like_the_connect_preflight
now expect the granted server, which is what the base branch returned for
both shapes; the unentitled-key exchange test's fixture becomes the empty
selection the router returns for such a key; listing tests that stub the
manager now point its lookup at a real empty manager so the `among` pass
runs. New tests cover the grant-first router, the manager `among` pass
order and IP hiding, and the connect preflight under an alias collision.
This commit is contained in:
Yucheng He 2026-10-01 18:00:25 -07:00
parent 608fff7b08
commit 86c3b910a6
5 changed files with 270 additions and 36 deletions

View file

@ -7199,12 +7199,18 @@ class MCPServerManager:
return server
return None
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
def get_mcp_server_answering_to(
self, name: str, client_ip: str | None = None, *, among: Sequence[MCPServer] | None = None
) -> MCPServer | None:
"""The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped
router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name``
priority first, then the exact ``server_id``, then the name priority case-insensitively, then any prefix
form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the
pass that found it: it never falls through to a looser pass that could name another server."""
pass that found it: it never falls through to a looser pass that could name another server. ``among``
runs the same passes over those servers alone instead of the registry, which is how the scoped router
picks the caller's granted server answering to ``name``."""
if among is not None:
return self._server_among_answering_to(name, tuple(among), client_ip)
exact: Final = self.get_mcp_server_by_name(name)
if exact is not None:
return exact if self._is_server_accessible_from_ip(exact, client_ip) else None
@ -7230,6 +7236,33 @@ class MCPServerManager:
None,
)
def _server_among_answering_to(
self, name: str, servers: Sequence[MCPServer], client_ip: str | None
) -> MCPServer | None:
"""``get_mcp_server_answering_to`` over ``servers`` instead of the registry: the same passes in the
same order, with a server hidden from ``client_ip`` resolving to ``None`` at the pass that found it."""
requested: Final = name.lower()
passes: Final[tuple[Callable[[MCPServer], bool], ...]] = (
lambda server: server.alias == name,
lambda server: server.server_name == name,
lambda server: server.name == name,
lambda server: server.server_id == name,
lambda server: (server.alias or "").lower() == requested,
lambda server: (server.server_name or "").lower() == requested,
lambda server: (server.name or "").lower() == requested,
)
for matches in passes:
if (found := next((server for server in servers if matches(server)), None)) is not None:
return found if self._is_server_accessible_from_ip(found, client_ip) else None
return next(
(
server
for server in servers
if self._is_server_accessible_from_ip(server, client_ip) and server_answers_to_name(server, name)
),
None,
)
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
"""
Get registry filtered by client IP access control.

View file

@ -501,17 +501,22 @@ def _server_answers_to(server: MCPServer, name: str) -> bool:
def _scoped_server(
name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None
) -> MCPServer | Literal["denied"] | None:
"""The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made
with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"``
when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name
is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry
cannot place the name for any caller, after trying the granted servers answering to it."""
"""The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's
``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery
use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: the registry's
own pass order run over ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or
case variant the registry would pick. With no granted server answering: ``"denied"`` when the registry
places the name on a server the caller does not hold, so it is not retried as an access group; ``None``
when the registry cannot place the name for any caller."""
registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip)
if registry_pick is not None:
return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), "denied")
if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None:
if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None:
return "denied"
return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None)
granted: Final = global_mcp_server_manager.get_mcp_server_answering_to(
name, client_ip=client_ip, among=allowed_mcp_servers
)
if granted is not None:
return granted
return "denied" if registry_pick is not None else None
async def raise_denied_scoped_mcp_access(

View file

@ -1603,7 +1603,24 @@ if MCP_AVAILABLE:
a server it will be 403'd on immediately after authentication.
"""
for server_name in mcp_servers or []:
server = operations.global_mcp_server_manager.get_mcp_server_answering_to(server_name, client_ip=client_ip)
registry_pick = operations.global_mcp_server_manager.get_mcp_server_answering_to(
server_name, client_ip=client_ip
)
obo_without_subject = (
registry_pick is not None
and registry_pick.auth_type == MCPAuth.oauth2_token_exchange
and not oauth2_headers
)
allowed_single = (
await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
if registry_pick and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1
else ()
)
granted = next(iter(allowed_single), None)
server = granted if granted is not None else registry_pick
granted_single = granted is not None
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
# Caller's narrowed scope excludes this server — skip the
# preemptive challenge and let downstream authorization
@ -1703,7 +1720,7 @@ if MCP_AVAILABLE:
# JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate;
# guardrail-only gates fire only on a single-server connect the key's grant admits, so a
# key without access gets the grant's 403 instead of a sign-in it could not use. The one
# admission lookup below serves the challenge, the sign-in preflight and the exchange.
# admission lookup above serves the challenge, the sign-in preflight and the exchange.
sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None
subject_token = (
operations.global_mcp_server_manager._extract_subject_token( # pyright: ignore[reportPrivateUsage] # the manager owns the subject/admission filter shared with the preflight
@ -1712,19 +1729,6 @@ if MCP_AVAILABLE:
if server is not None
else None
)
obo_without_subject = (
server is not None and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers
)
allowed_single = (
await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
if server and not obo_without_subject and mcp_servers is not None and len(mcp_servers) == 1
else ()
)
granted_single = server is not None and any(
allowed.server_id == server.server_id for allowed in allowed_single
)
if server and sign_in is not None and subject_token is None and (obo_without_subject or granted_single):
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,

View file

@ -924,7 +924,7 @@ async def test_get_tools_from_mcp_servers():
mock_manager.get_mcp_server_by_id = lambda server_id: (
mock_server_1 if server_id == "server1_id" else mock_server_2
)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(

View file

@ -32,7 +32,7 @@ import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPTransport,
@ -1279,7 +1279,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
mock_manager.get_mcp_server_by_id = lambda server_id: (
working_server if server_id == "working_server" else failing_server
)
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
@ -6765,7 +6765,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
):
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0))
mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers)
@ -8620,7 +8620,9 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a
finally:
global_mcp_server_manager.registry.clear()
assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another"
assert [s.server_id for s in only_b] == ["b-id"], (
"the granted server answering to the name wins over the registry's ungranted alias holder"
)
@pytest.mark.asyncio
@ -8691,10 +8693,202 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the
global_mcp_server_manager.registry.clear()
assert external == [], "a name the preflight hides from this IP must not reroute to a case variant"
assert internal == []
assert [s.server_id for s in internal] == ["u-id"], "with no IP hiding in play the granted case variant wins"
assert [s.server_id for s in by_own_name] == ["u-id"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("registry", "scope", "granted", "expected"),
[
pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"),
pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"),
pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"),
pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"),
],
)
async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names
servers: Final = {
"a-id": MCPServer(
server_id="a-id", name="a_docs", server_name="a_docs", alias="docs", transport=MCPTransport.http
),
"d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http),
"g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http),
"g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http),
"p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http),
"m-id": MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http),
}
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update({server_id: servers[server_id] for server_id in registry})
try:
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["m-id"],
) as groups:
selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted]
)
finally:
global_mcp_server_manager.registry.clear()
assert [s.server_id for s in selected] == expected
groups.assert_not_awaited()
def test_get_mcp_server_answering_to_among_applies_the_registry_pass_order_and_ip_hiding():
manager: Final = MCPServerManager()
by_alias: Final = MCPServer(server_id="a-id", name="a", server_name="a", alias="svc", transport=MCPTransport.http)
by_server_name: Final = MCPServer(server_id="b-id", name="b", server_name="svc", transport=MCPTransport.http)
by_name: Final = MCPServer(server_id="c-id", name="svc", server_name="c", transport=MCPTransport.http)
by_id: Final = MCPServer(server_id="svc", name="d", server_name="d", transport=MCPTransport.http)
folded: Final = MCPServer(server_id="e-id", name="e", server_name="SVC", transport=MCPTransport.http)
hidden: Final = MCPServer(
server_id="h-id",
name="h",
server_name="h",
alias="svc",
transport=MCPTransport.http,
available_on_public_internet=False,
)
assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name, by_alias]) is by_alias
assert manager.get_mcp_server_answering_to("svc", among=[by_name, by_server_name]) is by_server_name
assert manager.get_mcp_server_answering_to("svc", among=[by_id, by_name]) is by_name
assert manager.get_mcp_server_answering_to("svc", among=[folded, by_id]) is by_id
assert manager.get_mcp_server_answering_to("svc", among=[folded]) is folded
assert manager.get_mcp_server_answering_to("E-ID", among=[folded]) is folded
assert manager.get_mcp_server_answering_to("svc", among=()) is None
assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[hidden, folded]) is None
assert manager.get_mcp_server_answering_to("svc", client_ip="10.0.0.7", among=[hidden, folded]) is hidden
assert manager.get_mcp_server_answering_to("svc", client_ip="203.0.113.7", among=[folded]) is folded
assert manager.get_mcp_server_answering_to("svc") is None
manager.registry = {"e-id": folded, "b-id": by_server_name, "a-id": by_alias}
assert manager.get_mcp_server_answering_to("svc") is by_alias
assert manager.get_mcp_server_answering_to("svc") is manager.get_mcp_server_answering_to(
"svc", among=tuple(manager.registry.values())
)
class _GrantedServerSignInGuardrail(CustomGuardrail):
"""Requires caller sign-in on one server only and records every server it is asked about."""
def __init__(self, *args, gated_server_id: str, **kwargs):
super().__init__(*args, **kwargs)
self._gated_server_id = gated_server_id
self.asked_about = [] # mutable-ok: call recorder
def caller_sign_in(self, server, user_api_key_auth):
from litellm.proxy._experimental.mcp_server.caller_sign_in import CallerSignIn
self.asked_about.append(server.server_id)
if server.server_id != self._gated_server_id:
return None
return CallerSignIn(issuers=("https://idp.test",), scopes=("scope-a",))
async def preflight_caller_sign_in(self, server, user_api_key_auth, subject_token):
from litellm.proxy._experimental.mcp_server.caller_sign_in import SignedIn
return SignedIn()
class TestConnectPreflightRoutesLikeTheScopedRouter:
"""A granted key connecting to ``/mcp/{name}`` is pre-flighted for the server the scoped router routes
it to, so an ungranted server holding the name as an alias neither hides the granted server's sign-in
challenge nor skips its connect-time exchange."""
@staticmethod
def _register_alias_collision() -> None:
alias_holder: Final = MCPServer(
server_id="a-id",
name="a_docs",
server_name="a_docs",
alias="docs",
url="https://a.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
granted: Final = MCPServer(
server_id="d-id",
name="docs",
server_name="docs",
url="https://d.test/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
mcp_operations.global_mcp_server_manager.registry.update({"a-id": alias_holder, "d-id": granted})
@staticmethod
async def _connect_to_docs() -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp/docs", "headers": []},
mcp_servers=["docs"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"),
client_ip=None,
raw_headers={"x-litellm-api-key": "sk-granted-docs"},
)
@pytest.mark.asyncio
async def test_exchange_runs_for_the_granted_server_not_the_alias_holder(self):
self._register_alias_collision()
async def refuse_exchange(server, **kwargs):
raise HTTPException(
status_code=401, detail=f"exchange refused for {server.server_id} as {kwargs['connected_as']}"
)
with (
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
mcp_operations.global_mcp_server_manager, "get_allowed_mcp_servers", AsyncMock(return_value=["d-id"])
),
patch.object( # test-quality-ok: a real exchanger would call an IdP; this one reports which server it ran for
mcp_operations.global_mcp_server_manager, "preflight_token_exchange", refuse_exchange
),
pytest.raises(HTTPException) as exc,
):
await self._connect_to_docs()
assert exc.value.status_code == 401
assert exc.value.detail == "exchange refused for d-id as docs"
@pytest.mark.asyncio
async def test_sign_in_challenge_names_the_granted_server_not_the_alias_holder(self, monkeypatch):
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
self._register_alias_collision()
guardrail: Final = _GrantedServerSignInGuardrail(guardrail_name="sign-in-stub", gated_server_id="d-id")
litellm.logging_callback_manager.add_litellm_callback(guardrail)
try:
with (
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
mcp_operations.global_mcp_server_manager,
"get_allowed_mcp_servers",
AsyncMock(return_value=["d-id"]),
),
pytest.raises(HTTPException) as exc,
):
await self._connect_to_docs()
finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
litellm.callbacks, guardrail, require_self=False
)
assert exc.value.status_code == 401
assert ((exc.value.headers or {}).get("WWW-Authenticate") or "").startswith(
'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/docs"'
)
assert guardrail.asked_about == ["d-id"]
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():
"""
@ -9836,7 +10030,7 @@ async def test_aggregate_listing_reports_per_server_outcomes():
mock_manager.get_mcp_server_by_id = lambda server_id: (
working_server if server_id == "working_server" else broken_server
)
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
mock_manager.get_mcp_server_answering_to = MCPServerManager().get_mcp_server_answering_to
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(server, **kwargs):
@ -10351,9 +10545,7 @@ class TestOboPreflightScopedToAllowedServers:
requested = _make_obo_server("obo_tools")
key = UserAPIKeyAuth(api_key="sk-plain-only")
allowed_lookup, preflight = await self._run(
requested, allowed=[_make_obo_server("plain_tools")], user_api_key_auth=key
)
allowed_lookup, preflight = await self._run(requested, allowed=[], user_api_key_auth=key)
preflight.assert_not_awaited()
allowed_lookup.assert_awaited_once_with(