diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index af18c666679..82475ee8d8a 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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): diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a05af66118c..8af6ee98a5a 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index bcbaea54457..e27589ff567 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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() diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1d45fbefa8f..c8836349cea 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7327f72bf40..3515227799e 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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 ): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 927e7caf7da..06926ecf9eb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 = {} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index eea5b4aeaeb..4a9dd9e247d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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__]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index f4feac68fcc..3f171390497 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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(