chore(mcp): close redirect-bypass + name-lookup variants of the SSRF/IP gating

Two follow-ups to address Veria-AI findings on the PR.

(1) Disable redirect-following on the OAuth outbound POSTs.
validate_url only inspects the initial URL; httpx clients default to
follow_redirects=True, so a malicious 30x from the validated host could
bounce the proxy to an internal target. Add follow_redirects to the
AsyncHTTPHandler.post wrapper (mirroring the existing get) and pass
follow_redirects=False from both /token and /register flows.

(2) Make get_mcp_server_by_name fail closed on None client_ip. The
wrapper previously translated None to INTERNAL_REQUEST internally to
preserve the "None means internal" convention, but that meant any
external request handler that forgot to pass an IP silently bypassed
gating. Update the wrapper to require an explicit sentinel for
internal callers; update the four external callers in
auth/user_api_key_auth_mcp.py and rest_endpoints.py to pass
INTERNAL_REQUEST (where the lookup is metadata-only) or the real
extracted client_ip.

Adjust the test fixture in test_discoverable_endpoints.py to return
INTERNAL_REQUEST instead of None so OAuth flow tests bypass IP gating
explicitly. Update two stub lambdas in test_rest_endpoints.py to accept
the new client_ip kwarg (CLAUDE.md: keep monkeypatch stubs in sync with
real signatures).
This commit is contained in:
user 2026-05-10 02:51:54 +00:00
parent fc018a9fca
commit 6df38ecc1d
No known key found for this signature in database
8 changed files with 103 additions and 33 deletions

View file

@ -611,6 +611,7 @@ class AsyncHTTPHandler:
logging_obj: Optional[LiteLLMLoggingObject] = None,
files: Optional[RequestFiles] = None,
content: Any = None,
follow_redirects: Optional[bool] = None,
):
start_time = time.time()
try:
@ -633,7 +634,10 @@ class AsyncHTTPHandler:
files=files,
content=request_content,
)
response = await self.client.send(req, stream=stream)
send_kwargs: Dict[str, Any] = {"stream": stream}
if follow_redirects is not None:
send_kwargs["follow_redirects"] = follow_redirects
response = await self.client.send(req, **send_kwargs)
response.raise_for_status()
return response
except (httpx.RemoteProtocolError, httpx.ConnectError):

View file

@ -213,6 +213,7 @@ class MCPRequestHandler:
# Inline imports avoid a circular dependency: mcp_server_manager imports
# from this module.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
INTERNAL_REQUEST,
global_mcp_server_manager,
)
from litellm.types.mcp import MCPAuth
@ -229,7 +230,12 @@ class MCPRequestHandler:
return False
for name in target_names:
server = global_mcp_server_manager.get_mcp_server_by_name(name)
# Metadata lookup ("does this server use OAuth2?") used to decide
# whether anonymous OAuth2 fallback is allowed for this path.
# Not an access check — the IP gate doesn't apply.
server = global_mcp_server_manager.get_mcp_server_by_name(
name, client_ip=INTERNAL_REQUEST
)
if server is None or server.auth_type != MCPAuth.oauth2:
return False
return True

View file

@ -423,10 +423,14 @@ async def exchange_token_with_server(
target_url, host_header = _validate_mcp_oauth_outbound_url(
mcp_server.token_url, role="token"
)
# Disable redirect-following so a malicious 30x from the validated host
# can't bounce the proxy to an internal target — validate_url only
# checked the initial URL.
response = await async_client.post(
target_url,
headers={"Accept": "application/json", "Host": host_header},
data=token_data,
follow_redirects=False,
)
response.raise_for_status()
@ -531,10 +535,14 @@ async def register_client_with_server(
target_url, host_header = _validate_mcp_oauth_outbound_url(
mcp_server.registration_url, role="registration"
)
# Disable redirect-following so a malicious 30x from the validated host
# can't bounce the proxy to an internal target — validate_url only
# checked the initial URL.
response = await async_client.post(
target_url,
headers={**headers, "Host": host_header},
json=register_data,
follow_redirects=False,
)
response.raise_for_status()

View file

@ -3207,7 +3207,9 @@ class MCPServerManager:
return result
def get_mcp_server_by_name(
self, server_name: str, client_ip: Optional[str] = None
self,
server_name: str,
client_ip: Union[str, _InternalRequest, None] = None,
) -> Optional[MCPServer]:
"""
Get the MCP Server from the server name.
@ -3219,33 +3221,33 @@ class MCPServerManager:
Args:
server_name: The server name to look up.
client_ip: Optional client IP for access control. When provided,
non-public servers are hidden from external IPs.
``None`` is treated as "internal context, no IP gating"
to preserve the existing contract for internal callers
(auth, debug, registry maintenance). External request
handlers should pass a real IP.
client_ip: External request IP for IP-based access control, or
``INTERNAL_REQUEST`` for internal callers (admin debug,
registry maintenance) that intentionally bypass IP
gating. ``None`` fails closed: external request handlers
must extract a real IP via
``IPAddressUtils.get_mcp_client_ip(request)``. Earlier
behaviour silently bypassed gating on ``None``, which
let request handlers that forgot to pass an IP reach
internal-only servers.
"""
# The gate fails closed on None; preserve the wrapper's existing
# "None means internal" convention by routing through the sentinel.
gate_arg = INTERNAL_REQUEST if client_ip is None else client_ip
registry = self.get_registry()
# Pass 1: Match by alias (highest priority)
for server in registry.values():
if server.alias == server_name:
if not self._is_server_accessible_from_ip(server, gate_arg):
if not self._is_server_accessible_from_ip(server, client_ip):
return None
return server
# Pass 2: Match by server_name
for server in registry.values():
if server.server_name == server_name:
if not self._is_server_accessible_from_ip(server, gate_arg):
if not self._is_server_accessible_from_ip(server, client_ip):
return None
return server
# Pass 3: Match by name (lowest priority)
for server in registry.values():
if server.name == server_name:
if not self._is_server_accessible_from_ip(server, gate_arg):
if not self._is_server_accessible_from_ip(server, client_ip):
return None
return server
return None

View file

@ -380,7 +380,9 @@ if MCP_AVAILABLE:
# Resolve a server name to its UUID if needed
_name_resolved = None
if server_id not in allowed_server_ids:
_name_resolved = global_mcp_server_manager.get_mcp_server_by_name(server_id)
_name_resolved = global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=rest_client_ip
)
if _name_resolved is not None and _name_resolved.server_id in set(
allowed_server_ids
):
@ -482,7 +484,9 @@ if MCP_AVAILABLE:
# Resolve a server name to its UUID if needed
_name_resolved = None
if server_id not in allowed_server_ids:
_name_resolved = global_mcp_server_manager.get_mcp_server_by_name(server_id)
_name_resolved = global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=rest_client_ip
)
if _name_resolved is not None and _name_resolved.server_id in set(
allowed_server_ids
):

View file

@ -6,19 +6,19 @@ import pytest
from fastapi import HTTPException
# Fixture to mock IP address check for all MCP tests
# This prevents tests from failing due to IP-based access control
# Fixture to bypass MCP IP-based access control for all OAuth flow tests.
# Mock requests don't carry a real client IP context; the bypass uses the
# explicit INTERNAL_REQUEST sentinel because passing None now fails closed
# in the gate function (see MCPServerManager._is_server_accessible_from_ip).
@pytest.fixture(autouse=True)
def mock_mcp_client_ip():
"""Mock IPAddressUtils.get_mcp_client_ip to return None for all tests.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
INTERNAL_REQUEST,
)
This bypasses IP-based access control in tests, since the MCP server's
available_on_public_internet defaults to False and mock requests don't
have proper client IP context.
"""
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value=None,
return_value=INTERNAL_REQUEST,
):
yield
@ -91,6 +91,7 @@ async def test_authorize_endpoint_includes_response_type():
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -156,6 +157,7 @@ async def test_authorize_endpoint_preserves_existing_query_params():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -223,6 +225,7 @@ async def test_authorize_endpoint_forwards_pkce_parameters():
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm-proxy.example.com/"
mock_request.headers = {}
@ -294,6 +297,7 @@ async def test_token_endpoint_forwards_code_verifier():
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm-proxy.example.com/"
mock_request.headers = {}
@ -371,6 +375,7 @@ async def test_register_client_without_mcp_server_name_returns_dummy():
global_mcp_server_manager.registry.clear()
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
with patch(
@ -419,6 +424,7 @@ async def test_register_client_returns_existing_server_credentials():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
@ -474,6 +480,7 @@ async def test_register_client_remote_registration_success():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
@ -581,6 +588,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto():
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://litellm.example.com/" # HTTP
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
@ -649,6 +657,7 @@ async def test_token_endpoint_respects_x_forwarded_proto():
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://litellm-proxy.example.com/" # HTTP
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
@ -727,6 +736,7 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto():
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://litellm.example.com/" # HTTP
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
@ -782,6 +792,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto():
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://litellm.example.com/" # HTTP
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
@ -820,6 +831,7 @@ async def test_register_client_respects_x_forwarded_proto():
# Mock request with http base_url but X-Forwarded-Proto: https
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://proxy.litellm.example/" # HTTP
mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy
@ -879,6 +891,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host():
# Internal: http://localhost:8888/github/mcp
# External: https://proxy.example.com/github/mcp
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://localhost:8888/github/mcp"
mock_request.headers = {
"X-Forwarded-Proto": "https",
@ -951,6 +964,7 @@ async def test_token_endpoint_respects_x_forwarded_host():
# Mock request simulating nginx proxy without port in host
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://localhost:8888/github/mcp"
mock_request.headers = {
"X-Forwarded-Proto": "https",
@ -1131,6 +1145,7 @@ def test_get_request_base_url_comprehensive(
pytest.skip("MCP discoverable endpoints not available")
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = base_url
headers = {}
@ -1217,6 +1232,7 @@ def test_get_request_base_url_xff_trust_gate(
pytest.skip("MCP discoverable endpoints not available")
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://localhost:4000/"
mock_request.client = MagicMock()
mock_request.client.host = direct_ip
@ -1261,6 +1277,7 @@ def test_xff_misconfig_warning_emitted_once(caplog):
ip_address_utils._warned_xff_without_trusted_ranges = False
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "http://localhost:4000/"
mock_request.client = MagicMock()
mock_request.client.host = "203.0.113.5"
@ -1331,6 +1348,7 @@ async def test_oauth_protected_resource_returns_empty_scopes_when_none():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -1385,6 +1403,7 @@ async def test_oauth_authorization_server_returns_empty_scopes_when_none():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -1453,6 +1472,7 @@ async def test_authorize_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1506,6 +1526,7 @@ async def test_authorize_root_fails_with_multiple_oauth2_servers():
global_mcp_server_manager.registry[server2.server_id] = server2
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1544,6 +1565,7 @@ async def test_authorize_root_does_not_resolve_private_server_for_external_clien
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1585,6 +1607,7 @@ async def test_token_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1650,6 +1673,7 @@ async def test_token_root_does_not_resolve_private_server_for_external_client():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1694,6 +1718,7 @@ async def test_register_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1731,6 +1756,7 @@ async def test_register_root_does_not_resolve_private_server_for_external_client
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1773,6 +1799,7 @@ async def test_discovery_root_includes_server_name_prefix():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -1813,6 +1840,7 @@ async def test_discovery_root_does_not_expose_private_server_for_external_client
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
@ -2014,6 +2042,7 @@ async def test_oauth_authorize_includes_scopes_from_server_config():
)
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -2072,6 +2101,7 @@ async def test_oauth_authorize_prefers_request_scope_over_server_config():
)
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -2141,6 +2171,7 @@ async def test_token_endpoint_refresh_token_grant():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
@ -2275,6 +2306,7 @@ async def test_authorize_endpoint_rejects_non_loopback_redirect_uri():
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -2322,6 +2354,7 @@ async def test_authorize_endpoint_accepts_ipv4_loopback_range_and_ipv6_full_form
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -2451,6 +2484,7 @@ async def test_token_endpoint_sets_no_store_cache_control():
token_url="https://provider.com/oauth/token",
)
mock_request = MagicMock(spec=Request)
mock_request.client = MagicMock(host="127.0.0.1")
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}

View file

@ -3377,11 +3377,13 @@ class TestIPGatingFailClosed:
}[client_ip_kind]
assert manager._is_server_accessible_from_ip(server, client_ip) is expected
def test_get_mcp_server_by_name_preserves_internal_contract(self):
# Internal callers historically passed client_ip=None to mean "no IP
# gating." get_mcp_server_by_name translates None → INTERNAL_REQUEST
# so those callers keep working after the gate change.
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.
# Internal callers (admin debug, registry maintenance) must use
# INTERNAL_REQUEST explicitly.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
INTERNAL_REQUEST,
MCPServerManager,
)
@ -3389,12 +3391,18 @@ class TestIPGatingFailClosed:
server = self._internal_server()
manager.registry[server.server_id] = server
result = manager.get_mcp_server_by_name("internal", client_ip=None)
assert result is server
# Default client_ip is None — fails closed for internal-only servers.
result = manager.get_mcp_server_by_name("internal")
assert result is None
# External IP can't reach a non-public server.
result = manager.get_mcp_server_by_name("internal", client_ip="8.8.8.8")
assert result is None
# INTERNAL_REQUEST sentinel bypasses gating for internal callers.
result = manager.get_mcp_server_by_name("internal", client_ip=INTERNAL_REQUEST)
assert result is server
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -595,7 +595,9 @@ class TestListToolsRestAPI:
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_name",
lambda name: stub_server if name == "my-server" else None,
lambda name, client_ip=None, **kwargs: (
stub_server if name == "my-server" else None
),
raising=False,
)
monkeypatch.setattr(
@ -658,7 +660,9 @@ class TestListToolsRestAPI:
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_name",
lambda name: stub_server if name == "restricted-server" else None,
lambda name, client_ip=None, **kwargs: (
stub_server if name == "restricted-server" else None
),
raising=False,
)
monkeypatch.setattr(