mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
94797e0566
commit
03c95899ef
4 changed files with 78 additions and 2 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue