fix(mcp): add unknown-IP opt-out + propagate follow_redirects on retry

Two compat/safety items flagged by Greptile in PR review:

1. Fail-closed IP gating had no opt-out flag. Operators behind ASGI
   middleware or load-balancers where request.client is legitimately
   None will get unexpected 403s with no recourse short of code
   changes. Add general_settings.mcp_allow_unknown_client_ip (default
   false). When true, missing client_ip is promoted to INTERNAL_REQUEST
   at the gate. Centralized in _resolve_unknown_client_ip and applied
   at the three entry points (_is_server_accessible_from_ip,
   filter_server_ids_by_ip_with_info, get_filtered_registry).

2. AsyncHTTPHandler.post() now passes follow_redirects through to
   single_connection_post_request on the connection-error retry path.
   Without this, a transient RemoteProtocolError followed by a 30x on
   reconnect could bypass the SSRF redirect block on the MCP OAuth
   /token and /register flows.
This commit is contained in:
user 2026-05-10 07:08:40 +00:00
parent 94797e0566
commit 03c95899ef
No known key found for this signature in database
4 changed files with 78 additions and 2 deletions

View file

@ -654,6 +654,7 @@ class AsyncHTTPHandler:
params=params,
headers=headers,
stream=stream,
follow_redirects=follow_redirects,
)
finally:
await new_client.aclose()
@ -857,6 +858,7 @@ class AsyncHTTPHandler:
headers: Optional[dict] = None,
stream: bool = False,
content: Any = None,
follow_redirects: Optional[bool] = None,
):
"""
Making POST request for a single connection client.
@ -869,7 +871,10 @@ class AsyncHTTPHandler:
req = client.build_request(
"POST", url, data=request_data, json=json, params=params, headers=headers, content=request_content # type: ignore
)
response = await 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 client.send(req, **send_kwargs)
response.raise_for_status()
return response

View file

@ -1074,6 +1074,20 @@ class MCPServerManager:
filtered, _ = self.filter_server_ids_by_ip_with_info(server_ids, client_ip)
return filtered
def _resolve_unknown_client_ip(
self, client_ip: Union[str, "_InternalRequest", None]
) -> Union[str, "_InternalRequest", None]:
"""Promote unknown ``client_ip=None`` to ``INTERNAL_REQUEST`` when the
operator has explicitly opted in via
``general_settings.mcp_allow_unknown_client_ip: true``. Lets
deployments behind ASGI middleware where ``request.client`` is
legitimately ``None`` opt back to the previous fail-open behavior."""
if client_ip is None:
general_settings = self._get_general_settings()
if general_settings.get("mcp_allow_unknown_client_ip", False):
return INTERNAL_REQUEST
return client_ip
def filter_server_ids_by_ip_with_info(
self,
server_ids: List[str],
@ -1090,6 +1104,7 @@ class MCPServerManager:
(admin debug, registry maintenance) should pass
``INTERNAL_REQUEST`` to bypass IP gating.
"""
client_ip = self._resolve_unknown_client_ip(client_ip)
if client_ip is INTERNAL_REQUEST:
return server_ids, 0
if client_ip is None:
@ -3108,11 +3123,15 @@ class MCPServerManager:
IP via ``IPAddressUtils.get_mcp_client_ip(request)`` and reject the
request when extraction fails. Earlier behaviour treated ``None`` as
"no filter," which let external callers reach internal-only servers
when IP extraction silently failed.
when IP extraction silently failed. Operators behind ASGI middleware
or load-balancer setups where ``request.client`` is legitimately
``None`` can opt back to the previous fail-open behaviour by setting
``general_settings.mcp_allow_unknown_client_ip: true``.
- If the server has ``available_on_public_internet=True``, it's
always accessible.
- Otherwise, only internal/private IPs can access it.
"""
client_ip = self._resolve_unknown_client_ip(client_ip)
if client_ip is INTERNAL_REQUEST:
return True
if client_ip is None:
@ -3279,6 +3298,7 @@ class MCPServerManager:
External request handlers must pass a real IP.
"""
registry = self.get_registry()
client_ip = self._resolve_unknown_client_ip(client_ip)
if client_ip is INTERNAL_REQUEST:
return registry
if client_ip is None:

View file

@ -798,3 +798,40 @@ def test_get_httpx_client_applies_httpx_timeout_object_without_mocking_handler()
assert handler.client.timeout == t
finally:
handler.close()
@pytest.mark.asyncio
async def test_post_retry_propagates_follow_redirects():
"""
AsyncHTTPHandler.post() retries connection errors via
single_connection_post_request(). When a caller passes
follow_redirects=False (e.g. the MCP OAuth SSRF guard), the retry
must honor it too — otherwise a transient connection error followed
by a 30x on retry can bypass the redirect block and hit an internal
address.
"""
from unittest.mock import AsyncMock
handler = AsyncHTTPHandler()
try:
with (
patch.object(
handler.client,
"send",
side_effect=httpx.RemoteProtocolError("forced retry"),
),
patch.object(
handler,
"single_connection_post_request",
new=AsyncMock(return_value=MagicMock(status_code=200)),
) as mock_retry,
):
await handler.post("https://example.com/token", follow_redirects=False)
assert mock_retry.await_count == 1
kwargs = mock_retry.await_args.kwargs
assert kwargs.get("follow_redirects") is False, (
"follow_redirects must reach the retry path so the SSRF "
"redirect block holds across reconnect"
)
finally:
await handler.close()

View file

@ -100,6 +100,20 @@ class TestMCPServerIPFiltering:
["priv"], client_ip=INTERNAL_REQUEST
) == ["priv"]
@patch("litellm.public_mcp_servers", [])
@patch(
"litellm.proxy.proxy_server.general_settings",
{"mcp_allow_unknown_client_ip": True},
)
def test_unknown_client_ip_opt_out_restores_fail_open(self):
# Operators behind ASGI middleware where request.client is legitimately
# None can opt back to the pre-fix fail-open behavior with explicit
# consent via general_settings.mcp_allow_unknown_client_ip: true.
priv = _make_server("priv", available_on_public_internet=False)
manager = _make_manager([priv])
assert manager.filter_server_ids_by_ip(["priv"], client_ip=None) == ["priv"]
class TestFilterServerIdsByIpWithInfo:
"""Tests that filter_server_ids_by_ip_with_info returns accurate block counts."""