fix(mcp): reuse safe URL fetch for OAuth discovery

This commit is contained in:
user 2026-04-30 12:42:52 -07:00
parent 1bb04cdc72
commit 51a3e90451
2 changed files with 218 additions and 131 deletions

View file

@ -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",

View file

@ -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()