mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): keep a route hidden from a client ip from rerouting to a case variant of its name
get_mcp_server_answering_to stops at the pass that finds an exact name or id and hides it from client_ip instead of falling through to the case-insensitive and prefix passes, and _scoped_server treats a name the registry knows for some caller but not this one as denied, so connect, discovery and scoped routing all refuse the hidden route instead of serving a public server whose alias only differs by case Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2172b8b60e
commit
f2a876f23f
6 changed files with 92 additions and 18 deletions
|
|
@ -7203,13 +7203,14 @@ class MCPServerManager:
|
|||
"""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."""
|
||||
exact: Final = self.get_mcp_server_by_name(name, client_ip=client_ip)
|
||||
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."""
|
||||
exact: Final = self.get_mcp_server_by_name(name)
|
||||
if exact is not None:
|
||||
return exact
|
||||
by_id: Final = self.get_mcp_server_by_id(name, client_ip=client_ip)
|
||||
return exact if self._is_server_accessible_from_ip(exact, client_ip) else None
|
||||
by_id: Final = self.get_mcp_server_by_id(name)
|
||||
if by_id is not None:
|
||||
return by_id
|
||||
return by_id if self._is_server_accessible_from_ip(by_id, client_ip) else None
|
||||
requested: Final = name.lower()
|
||||
servers: Final = tuple(self.get_registry().values())
|
||||
identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = (
|
||||
|
|
|
|||
|
|
@ -503,11 +503,14 @@ def _scoped_server(
|
|||
) -> 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, so the name is not retried as an access group.
|
||||
``None`` when the registry cannot place the name, after trying the granted servers answering to it."""
|
||||
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."""
|
||||
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:
|
||||
return "denied"
|
||||
return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -182,6 +182,54 @@ def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_a
|
|||
assert _advertised(candidate, stem) != _advertised(candidate, cased)
|
||||
|
||||
|
||||
def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_variant(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
stem: Final = "gh" + uuid.uuid4().hex[:6]
|
||||
cased: Final = stem.capitalize()
|
||||
external: Final = {"X-Forwarded-For": "203.0.113.7"}
|
||||
with mcp_peer() as hidden_peer, mcp_peer() as public_peer:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"] = {
|
||||
**config.get("general_settings", {}),
|
||||
"use_x_forwarded_for": True,
|
||||
"mcp_trusted_proxy_ranges": ["127.0.0.0/8"],
|
||||
}
|
||||
config["mcp_servers"] = {
|
||||
stem: {"transport": "http", "url": hidden_peer.url, "available_on_public_internet": False}
|
||||
}
|
||||
path: Final = tmp_path / "external-ip.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
public_id: Final = register_mcp(scenario, public_peer, cased)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [public_id]})
|
||||
|
||||
hidden_name: Final = _rpc(candidate, f"/mcp/{stem}", key, external)
|
||||
assert hidden_name.status_code == 403, hidden_name.text
|
||||
assert "www-authenticate" not in hidden_name.headers, hidden_name.headers
|
||||
assert (
|
||||
candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{stem}", headers=external).status_code
|
||||
== 404
|
||||
)
|
||||
assert _rpc(candidate, f"/mcp/{stem}", key, {}).status_code == 403
|
||||
|
||||
own_name: Final = _rpc(candidate, f"/mcp/{cased}", key, external)
|
||||
assert own_name.status_code == 200, own_name.text
|
||||
listed: Final = _rpc(candidate, f"/mcp/{cased}", key, external, method="tools/list")
|
||||
add: Final = next(
|
||||
tool["name"]
|
||||
for tool in json.loads(_sse_data(listed))["result"]["tools"]
|
||||
if tool["name"].endswith("add")
|
||||
)
|
||||
called: Final = _rpc(
|
||||
candidate, f"/mcp/{cased}", key, external, method="tools/call", params={"name": add, "arguments": ADD}
|
||||
)
|
||||
assert called.status_code == 200, called.text
|
||||
assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text
|
||||
assert len(tool_calls(public_peer.drain())) == 1
|
||||
assert tool_calls(hidden_peer.drain()) == ()
|
||||
|
||||
|
||||
def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, tmp_path: Path) -> None:
|
||||
signer_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
jwk: Final = json.loads(jwt_algorithms.RSAAlgorithm.to_jwk(signer_key.public_key()))
|
||||
|
|
|
|||
|
|
@ -3642,8 +3642,8 @@ async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
|
|||
|
||||
assert response.status_code == 307
|
||||
assert "https://provider.com/oauth/authorize" in response.headers["location"]
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3679,8 +3679,8 @@ async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
|
|||
)
|
||||
|
||||
assert json.loads(result.body)["access_token"] == "token"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3711,8 +3711,8 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
|
|||
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
|
||||
|
||||
assert json.loads(result.body)["client_id"] == "registered-client"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3739,8 +3739,8 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku
|
|||
|
||||
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
|
||||
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3815,8 +3815,8 @@ def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fa
|
|||
|
||||
assert result["scopes_supported"] == server.scopes
|
||||
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_name.assert_called_once_with(server.server_id)
|
||||
by_id.assert_called_once_with(server.server_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -8682,13 +8682,17 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the
|
|||
internal = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip=None
|
||||
)
|
||||
by_own_name = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=["Gh"], allowed_mcp_servers=[public], client_ip="203.0.113.7"
|
||||
)
|
||||
assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None
|
||||
assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip=None) is private
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert [s.server_id for s in external] == ["u-id"], "the router must apply the connect preflight's IP filter"
|
||||
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 by_own_name] == ["u-id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -6253,6 +6253,24 @@ class TestMCPServerManager:
|
|||
assert manager.get_mcp_server_answering_to("gh") is by_alias
|
||||
assert manager.get_mcp_server_answering_to("GH") is by_alias
|
||||
|
||||
@pytest.mark.parametrize("hidden_first", [True, False], ids=["hidden-listed-first", "public-listed-first"])
|
||||
def test_answering_to_never_reroutes_a_name_hidden_from_an_ip_to_a_case_variant(self, hidden_first):
|
||||
manager = MCPServerManager()
|
||||
hidden = MCPServer(
|
||||
server_id="p-id",
|
||||
name="gh",
|
||||
server_name="gh",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http)
|
||||
manager.registry = {"p-id": hidden, "u-id": public} if hidden_first else {"u-id": public, "p-id": hidden}
|
||||
|
||||
assert manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None
|
||||
assert manager.get_mcp_server_answering_to("p-id", client_ip="203.0.113.7") is None
|
||||
assert manager.get_mcp_server_answering_to("gh", client_ip="10.0.0.7") is hidden
|
||||
assert manager.get_mcp_server_answering_to("Gh", client_ip="203.0.113.7") is public
|
||||
|
||||
@pytest.mark.parametrize("pinned_first", [True, False], ids=["pinned-id-listed-first", "alias-listed-first"])
|
||||
def test_answering_to_and_discovery_agree_on_a_pinned_id_that_another_alias_case_folds_to(self, pinned_first):
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue