From 51a3e90451f0b23cd2ee650f108a39c31b706b80 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 30 Apr 2026 12:42:52 -0700 Subject: [PATCH] fix(mcp): reuse safe URL fetch for OAuth discovery --- .../mcp_server/mcp_server_manager.py | 122 ++++------ .../mcp_server/test_mcp_server_manager.py | 227 ++++++++++++++---- 2 files changed, 218 insertions(+), 131 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 28928210814..f96350500db 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -12,7 +12,6 @@ import hashlib import json import os import re -import socket from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast from urllib.parse import urlparse @@ -42,7 +41,7 @@ from litellm.constants import ( MCP_TOOL_LISTING_TIMEOUT, ) from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException -from litellm.litellm_core_utils.url_utils import _is_blocked_ip +from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( @@ -1502,27 +1501,15 @@ class MCPServerManager: return await client.get_prompt(get_prompt_request_params) @staticmethod - def _is_safe_metadata_url(url: str, server_url: str) -> bool: + def _is_same_authority_metadata_url(url: str, server_url: str) -> bool: """ - Whether ``url`` is safe to fetch during OAuth discovery for ``server_url``. + Whether ``url`` shares scheme, host, and port with ``server_url``. - OAuth metadata discovery follows attacker-influenceable URLs from a - WWW-Authenticate header and from the protected-resource-metadata JSON. - Without a guard those become an SSRF primitive: a malicious MCP server - can point the proxy at cloud metadata services, internal admin panels, - or loopback debug endpoints. - - A URL is allowed when: - - it shares (scheme, host, port) with ``server_url`` — well-known - endpoints constructed from the admin's URL, and PRM published at - the resource server itself per RFC 9728 §3.3, OR - - it resolves to one or more public IPs only — covers federated - authorization servers (Azure Entra, Google, Okta, GitHub) hosted - cross-origin from the resource. - - URLs that resolve to private / loopback / link-local / cloud-metadata - addresses, or that don't resolve at all, are rejected. ``http`` and - ``https`` are the only schemes accepted. + Same-authority metadata URLs are produced by our well-known discovery + construction and by resource servers that publish protected-resource + metadata on the resource origin. These must keep working for + administrator-configured internal MCP servers, so they are fetched + directly. Cross-origin URLs are fetched through ``async_safe_get``. """ try: target = urlparse(url) @@ -1535,37 +1522,24 @@ class MCPServerManager: target_port = target.port or (443 if target.scheme == "https" else 80) base_port = base.port or (443 if base.scheme == "https" else 80) - same_authority = ( + return ( base.scheme == target.scheme and (base.hostname or "").lower() == target.hostname.lower() and base_port == target_port ) - if same_authority: - return True - try: - infos = socket.getaddrinfo( - target.hostname, target_port, type=socket.SOCK_STREAM - ) - except socket.gaierror: - return False - - if not infos: - return False - - # Reuse the proxy-wide outbound block list (private / loopback / - # link-local / multicast / reserved / cloud-fabric IPs). Defence - # in depth only — the resolution here and the one httpx performs at - # request time leave a small DNS-rebinding window; an attacker who - # also controls the resource server is already in scope, so the - # primary mitigation is the same-authority pin above. - for info in infos: - sockaddr_host = info[4][0] - if not isinstance(sockaddr_host, str): - return False - if _is_blocked_ip(sockaddr_host): - return False - return True + async def _fetch_oauth_discovery_url(self, url: str, server_url: str) -> Any: + client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.MCP, + params={"timeout": MCP_METADATA_TIMEOUT}, + ) + if self._is_same_authority_metadata_url(url, server_url): + # Same-authority URLs may point at administrator-configured + # internal MCP servers. Do not run them through user URL + # validation, but also do not follow redirects because the + # redirect target would not inherit the same-authority guarantee. + return await client.get(url, follow_redirects=False) + return await async_safe_get(client, url) async def _descovery_metadata( self, @@ -1689,26 +1663,21 @@ class MCPServerManager: if not resource_metadata_url: return [], None - if not self._is_safe_metadata_url(resource_metadata_url, server_url): - verbose_logger.warning( - "MCP OAuth discovery: refusing to fetch resource metadata from %s " - "(rejected by SSRF guard for server %s)", - resource_metadata_url, - server_url, - ) - return [], None - try: - client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.MCP, - params={"timeout": MCP_METADATA_TIMEOUT}, + response = await self._fetch_oauth_discovery_url( + resource_metadata_url, server_url ) - # Redirects bypass the SSRF guard (the new ``Location`` is not - # re-checked against ``server_url``), so refuse to follow them. - # Spec-compliant OAuth metadata endpoints serve the JSON directly. - response = await client.get(resource_metadata_url, follow_redirects=False) response.raise_for_status() data = response.json() + except SSRFError as exc: + verbose_logger.warning( + "MCP OAuth discovery: refusing to fetch resource metadata from %s " + "(rejected by SSRF guard for server %s): %s", + resource_metadata_url, + server_url, + exc, + ) + return [], None except Exception as exc: # pragma: no cover - network issues verbose_logger.debug( "Failed to fetch MCP OAuth metadata from %s: %s", @@ -1802,26 +1771,19 @@ class MCPServerManager: candidate_urls.append(issuer_url.rstrip("/")) for url in candidate_urls: - if not self._is_safe_metadata_url(url, server_url): - verbose_logger.warning( - "MCP OAuth discovery: refusing to fetch authorization-server " - "metadata from %s (rejected by SSRF guard for server %s)", - url, - server_url, - ) - continue try: - client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.MCP, - params={"timeout": MCP_METADATA_TIMEOUT}, - ) - # Disable redirects: a redirect to a private IP would bypass - # the SSRF guard (the ``Location`` target is not re-checked - # against ``server_url``). Spec-compliant OAuth/OIDC - # metadata endpoints serve the JSON directly. - response = await client.get(url, follow_redirects=False) + response = await self._fetch_oauth_discovery_url(url, server_url) response.raise_for_status() data = response.json() + except SSRFError as exc: + verbose_logger.warning( + "MCP OAuth discovery: refusing to fetch authorization-server " + "metadata from %s (rejected by SSRF guard for server %s): %s", + url, + server_url, + exc, + ) + continue except Exception as exc: # pragma: no cover - network issues verbose_logger.debug( "Failed to fetch authorization metadata from %s: %s", 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 28eab1e7a00..6cbdcd28208 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 @@ -4,6 +4,7 @@ import logging import os import sys from datetime import datetime +from typing import Any, Dict from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -2971,7 +2972,7 @@ class TestOAuthDiscoverySSRFGuard: """Patch ``socket.getaddrinfo`` for a deterministic SSRF-guard test. ``mapping`` is ``{hostname: [ip-string, ...]}``; unknown hosts raise - ``gaierror`` (treated as "unresolvable" -> blocked). + ``gaierror`` (treated as "unresolvable" -> blocked by async_safe_get). """ import socket as _socket @@ -2984,26 +2985,21 @@ class TestOAuthDiscoverySSRFGuard: ] monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.socket.getaddrinfo", + "litellm.litellm_core_utils.url_utils.socket.getaddrinfo", fake_getaddrinfo, ) - def test_same_authority_url_is_safe(self): + def test_same_authority_url_is_direct_fetch_eligible(self): # Same scheme + host + port skips DNS entirely — the well-known # endpoint construction in _attempt_well_known_discovery always # produces same-authority URLs against the admin's server_url. - assert MCPServerManager._is_safe_metadata_url( + assert MCPServerManager._is_same_authority_metadata_url( "https://example.com/.well-known/oauth-protected-resource", "https://example.com/mcp", ) - def test_same_host_different_port_blocked_when_resolves_to_private_ip( - self, monkeypatch - ): - # Cross-authority (different port). Falls through to DNS check. - # If the resolved IP is private, the guard rejects. - self._patch_resolves(monkeypatch, {"example.com": ["10.1.2.3"]}) - assert not MCPServerManager._is_safe_metadata_url( + def test_same_host_different_port_uses_safe_fetch_path(self): + assert not MCPServerManager._is_same_authority_metadata_url( "https://example.com:9999/.well-known/oauth-protected-resource", "https://example.com/mcp", ) @@ -3023,48 +3019,133 @@ class TestOAuthDiscoverySSRFGuard: "fc00::1", # IPv6 ULA ], ) - def test_cross_origin_blocked_when_resolves_to_unsafe_ip(self, monkeypatch, ip): + @pytest.mark.asyncio + async def test_cross_origin_blocked_when_resolves_to_unsafe_ip( + self, monkeypatch, ip + ): self._patch_resolves(monkeypatch, {"attacker.example.com": [ip]}) - assert not MCPServerManager._is_safe_metadata_url( - f"https://attacker.example.com/.well-known/oauth-authorization-server", - "https://legit-mcp.example.com/mcp", - ) + manager = MCPServerManager() - def test_cross_origin_allowed_when_resolves_to_public_ip(self, monkeypatch): + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + result = await manager._fetch_single_authorization_server_metadata( + "https://attacker.example.com", + "https://legit-mcp.example.com/mcp", + ) + + assert result is None + mock_client.get.assert_not_called() + + @pytest.mark.asyncio + async def test_cross_origin_allowed_when_resolves_to_public_ip(self, monkeypatch): self._patch_resolves( monkeypatch, {"login.microsoftonline.com": ["20.190.151.7"]} ) - assert MCPServerManager._is_safe_metadata_url( - "https://login.microsoftonline.com/tenant/v2.0/.well-known/openid-configuration", - "https://atlassian-mcp.example.com/mcp", + manager = MCPServerManager() + + mock_response = MagicMock() + mock_response.is_redirect = False + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "authorization_servers": ["https://login.microsoftonline.com/tenant/v2.0"], + "scopes_supported": ["mcp.read"], + } + + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + servers, scopes = await manager._fetch_oauth_metadata_from_resource( + "https://login.microsoftonline.com/tenant/v2.0/.well-known/openid-configuration", + "https://atlassian-mcp.example.com/mcp", + ) + + assert servers == ["https://login.microsoftonline.com/tenant/v2.0"] + assert scopes == ["mcp.read"] + mock_client.get.assert_awaited_once() + assert mock_client.get.await_args.kwargs["follow_redirects"] is False + assert ( + mock_client.get.await_args.kwargs["headers"]["Host"] + == "login.microsoftonline.com" ) - def test_cross_origin_blocked_when_unresolvable(self, monkeypatch): + @pytest.mark.asyncio + async def test_cross_origin_blocked_when_unresolvable(self, monkeypatch): self._patch_resolves(monkeypatch, {}) - assert not MCPServerManager._is_safe_metadata_url( - "https://nope.example.invalid/.well-known/oauth-authorization-server", - "https://legit-mcp.example.com/mcp", - ) + manager = MCPServerManager() - def test_non_http_scheme_is_not_safe(self): - assert not MCPServerManager._is_safe_metadata_url( - "file:///etc/passwd", "https://example.com/mcp" - ) - assert not MCPServerManager._is_safe_metadata_url( - "gopher://example.com/", "https://example.com/mcp" - ) + mock_client = MagicMock() + mock_client.get = AsyncMock() - def test_dual_resolution_blocked_if_any_ip_unsafe(self, monkeypatch): + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + servers, scopes = await manager._fetch_oauth_metadata_from_resource( + "https://nope.example.invalid/.well-known/oauth-authorization-server", + "https://legit-mcp.example.com/mcp", + ) + + assert servers == [] + assert scopes is None + mock_client.get.assert_not_called() + + @pytest.mark.asyncio + async def test_non_http_scheme_is_not_safe(self): + manager = MCPServerManager() + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + servers, scopes = await manager._fetch_oauth_metadata_from_resource( + "file:///etc/passwd", + "https://example.com/mcp", + ) + result = await manager._fetch_single_authorization_server_metadata( + "gopher://example.com/", + "https://example.com/mcp", + ) + + assert servers == [] + assert scopes is None + assert result is None + mock_client.get.assert_not_called() + + @pytest.mark.asyncio + async def test_dual_resolution_blocked_if_any_ip_unsafe(self, monkeypatch): # If the attacker controls a DNS record returning multiple A records, - # one of which is private, the guard must reject — pinning to the - # safe IP would be a TOCTOU window. + # one of which is private, async_safe_get rejects before any network call. self._patch_resolves( monkeypatch, {"dual-stack.example.com": ["8.8.8.8", "127.0.0.1"]} ) - assert not MCPServerManager._is_safe_metadata_url( - "https://dual-stack.example.com/.well-known/oauth-authorization-server", - "https://legit-mcp.example.com/mcp", - ) + manager = MCPServerManager() + + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + servers, scopes = await manager._fetch_oauth_metadata_from_resource( + "https://dual-stack.example.com/.well-known/oauth-authorization-server", + "https://legit-mcp.example.com/mcp", + ) + + assert servers == [] + assert scopes is None + mock_client.get.assert_not_called() @pytest.mark.asyncio async def test_fetch_oauth_metadata_refuses_unsafe_url(self, monkeypatch): @@ -3089,23 +3170,67 @@ class TestOAuthDiscoverySSRFGuard: assert scopes is None mock_client.get.assert_not_called() - def test_empty_getaddrinfo_result_blocks_url(self, monkeypatch): + @pytest.mark.asyncio + async def test_empty_getaddrinfo_result_blocks_url(self, monkeypatch): # POSIX doesn't strictly forbid an empty success-list from getaddrinfo. - # The guard must fail closed rather than fall through to ``return True``. + # async_safe_get must fail closed rather than making a network call. monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.socket.getaddrinfo", + "litellm.litellm_core_utils.url_utils.socket.getaddrinfo", lambda *a, **k: [], ) - assert not MCPServerManager._is_safe_metadata_url( - "https://no-records.example.com/.well-known/oauth-authorization-server", - "https://legit-mcp.example.com/mcp", - ) + manager = MCPServerManager() + + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + servers, scopes = await manager._fetch_oauth_metadata_from_resource( + "https://no-records.example.com/.well-known/oauth-authorization-server", + "https://legit-mcp.example.com/mcp", + ) + + assert servers == [] + assert scopes is None + mock_client.get.assert_not_called() @pytest.mark.asyncio - async def test_fetch_oauth_metadata_does_not_follow_redirects(self): - # If the validated origin redirects to a loopback or other unsafe - # address, httpx must NOT follow — the new ``Location`` would not - # be re-checked against ``server_url``. + async def test_cross_origin_redirect_is_revalidated(self, monkeypatch): + self._patch_resolves( + monkeypatch, + { + "provider.example.com": ["8.8.8.8"], + "127.0.0.1": ["127.0.0.1"], + }, + ) + manager = MCPServerManager() + + redirect_response = MagicMock() + redirect_response.is_redirect = True + redirect_response.headers = {"location": "http://127.0.0.1/admin"} + + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=redirect_response) + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + result = await manager._fetch_single_authorization_server_metadata( + "https://provider.example.com", + "https://legit-mcp.example.com/mcp", + ) + + assert result is None + assert mock_client.get.await_count == 3 + + @pytest.mark.asyncio + async def test_same_authority_fetch_does_not_follow_redirects(self): + # Same-authority URLs may be internal admin-configured MCP servers, so + # they are fetched directly. Redirects are still disabled because a + # Location target would not inherit the same-authority guarantee. manager = MCPServerManager() mock_response = MagicMock() @@ -3134,7 +3259,7 @@ class TestOAuthDiscoverySSRFGuard: assert captured_kwargs.get("follow_redirects") is False @pytest.mark.asyncio - async def test_fetch_single_auth_server_does_not_follow_redirects(self): + async def test_same_authority_auth_server_fetch_does_not_follow_redirects(self): # Same redirect-bypass concern for the authorization-server fetch path. manager = MCPServerManager()