mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): reuse safe URL fetch for OAuth discovery
This commit is contained in:
parent
1bb04cdc72
commit
51a3e90451
2 changed files with 218 additions and 131 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue