mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-23 00:41:40 +00:00
chore(mcp): close fail-open variants in get_filtered_registry and filter_server_ids_by_ip
Veria-AI flagged that two more wrappers — get_filtered_registry and filter_server_ids_by_ip(_with_info) — still treated client_ip=None as "no filter, return everything." That left request paths that use them (_resolve_oauth2_server_for_root_endpoints when get_mcp_client_ip returned None, /public/mcp_hub registry list, server.py's MCP request preflight) able to auto-select or list internal-only servers when IP extraction failed. Both wrappers now follow the same contract as _is_server_accessible_from_ip and get_mcp_server_by_name: None fails closed (empty result), INTERNAL_REQUEST bypasses gating, real IPs apply the existing filter. All 5 callers already extract client_ip from the request and pass it, so behaviour for valid external requests is unchanged. The change only affects the previously-fail-open path where extraction returned None.
This commit is contained in:
parent
6df38ecc1d
commit
cf4156f866
2 changed files with 75 additions and 11 deletions
|
|
@ -1060,28 +1060,42 @@ class MCPServerManager:
|
|||
return toolset
|
||||
|
||||
def filter_server_ids_by_ip(
|
||||
self, server_ids: List[str], client_ip: Optional[str]
|
||||
self,
|
||||
server_ids: List[str],
|
||||
client_ip: Union[str, "_InternalRequest", None],
|
||||
) -> List[str]:
|
||||
"""
|
||||
Filter server IDs by client IP — external callers only see public servers.
|
||||
|
||||
Returns server_ids unchanged when client_ip is None (no filtering).
|
||||
See ``filter_server_ids_by_ip_with_info`` for the contract on
|
||||
``client_ip``: ``None`` fails closed, ``INTERNAL_REQUEST`` bypasses
|
||||
gating, real IPs apply the existing rules.
|
||||
"""
|
||||
filtered, _ = self.filter_server_ids_by_ip_with_info(server_ids, client_ip)
|
||||
return filtered
|
||||
|
||||
def filter_server_ids_by_ip_with_info(
|
||||
self, server_ids: List[str], client_ip: Optional[str]
|
||||
self,
|
||||
server_ids: List[str],
|
||||
client_ip: Union[str, "_InternalRequest", None],
|
||||
) -> Tuple[List[str], int]:
|
||||
"""
|
||||
Filter server IDs by client IP — external callers only see public servers.
|
||||
|
||||
Returns (filtered_ids, ip_blocked_count) where ip_blocked_count is the number
|
||||
of servers that were blocked because the client IP is not allowed to access them.
|
||||
Returns server_ids unchanged (with 0 blocked) when client_ip is None.
|
||||
Returns (filtered_ids, ip_blocked_count) where ip_blocked_count is the
|
||||
number of servers that were blocked because the client IP is not
|
||||
allowed to access them. ``None`` fails closed: external request
|
||||
handlers must extract a real IP via
|
||||
``IPAddressUtils.get_mcp_client_ip(request)``. Internal callers
|
||||
(admin debug, registry maintenance) should pass
|
||||
``INTERNAL_REQUEST`` to bypass IP gating.
|
||||
"""
|
||||
if client_ip is None:
|
||||
if client_ip is INTERNAL_REQUEST:
|
||||
return server_ids, 0
|
||||
if client_ip is None:
|
||||
# Fail closed: don't expose internal-only servers when the IP
|
||||
# couldn't be attributed.
|
||||
return [], len(server_ids)
|
||||
allowed = []
|
||||
blocked = 0
|
||||
for sid in server_ids:
|
||||
|
|
@ -3253,18 +3267,22 @@ class MCPServerManager:
|
|||
return None
|
||||
|
||||
def get_filtered_registry(
|
||||
self, client_ip: Optional[str] = None
|
||||
self, client_ip: Union[str, "_InternalRequest", None] = None
|
||||
) -> Dict[str, MCPServer]:
|
||||
"""
|
||||
Get registry filtered by client IP access control.
|
||||
|
||||
Args:
|
||||
client_ip: Optional client IP. When provided, non-public servers
|
||||
are hidden from external IPs. When None, returns all servers.
|
||||
client_ip: Real client IP (filter applied), ``INTERNAL_REQUEST``
|
||||
(return full registry, internal callers only), or
|
||||
``None`` (fails closed: returns empty registry).
|
||||
External request handlers must pass a real IP.
|
||||
"""
|
||||
registry = self.get_registry()
|
||||
if client_ip is None:
|
||||
if client_ip is INTERNAL_REQUEST:
|
||||
return registry
|
||||
if client_ip is None:
|
||||
return {}
|
||||
return {
|
||||
k: v
|
||||
for k, v in registry.items()
|
||||
|
|
|
|||
|
|
@ -3377,6 +3377,52 @@ class TestIPGatingFailClosed:
|
|||
}[client_ip_kind]
|
||||
assert manager._is_server_accessible_from_ip(server, client_ip) is expected
|
||||
|
||||
def test_get_filtered_registry_fails_closed_on_none(self):
|
||||
# The wrapper used to return the full registry on None client_ip,
|
||||
# which let unauth callers (e.g. _resolve_oauth2_server_for_root_endpoints
|
||||
# when get_mcp_client_ip returned None) auto-select internal-only
|
||||
# servers. Now fails closed; INTERNAL_REQUEST is the explicit bypass.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
INTERNAL_REQUEST,
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = self._internal_server()
|
||||
manager.registry[server.server_id] = server
|
||||
|
||||
assert manager.get_filtered_registry() == {}
|
||||
assert manager.get_filtered_registry(client_ip=None) == {}
|
||||
assert manager.get_filtered_registry(client_ip="8.8.8.8") == {}
|
||||
assert manager.get_filtered_registry(client_ip=INTERNAL_REQUEST) == {
|
||||
server.server_id: server
|
||||
}
|
||||
|
||||
def test_filter_server_ids_by_ip_fails_closed_on_none(self):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
INTERNAL_REQUEST,
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = self._internal_server()
|
||||
manager.registry[server.server_id] = server
|
||||
sids = [server.server_id]
|
||||
|
||||
# None fails closed: zero allowed, all reported as blocked.
|
||||
allowed, blocked = manager.filter_server_ids_by_ip_with_info(sids, None)
|
||||
assert (allowed, blocked) == ([], 1)
|
||||
|
||||
# External IP applies filter — internal server blocked.
|
||||
allowed, blocked = manager.filter_server_ids_by_ip_with_info(sids, "8.8.8.8")
|
||||
assert (allowed, blocked) == ([], 1)
|
||||
|
||||
# INTERNAL_REQUEST bypasses gating.
|
||||
allowed, blocked = manager.filter_server_ids_by_ip_with_info(
|
||||
sids, INTERNAL_REQUEST
|
||||
)
|
||||
assert (allowed, blocked) == (sids, 0)
|
||||
|
||||
def test_get_mcp_server_by_name_fails_closed_on_none(self):
|
||||
# External request handlers must extract a real client IP. Passing
|
||||
# None silently bypassed gating before; now it fails closed.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue