diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 82475ee8d8a..6d716bd9f52 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0c386666745..c4f0d10ca2a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 26f50e8e492..47170163392 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -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() diff --git a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py index 6472a52fa22..604af05dd01 100644 --- a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py +++ b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py @@ -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."""