fix(proxy): add SSRF protection via resolve-and-rewrite for user-supplied URLs

Add validate_url() utility that resolves DNS once, validates all IPs
against private network ranges, and rewrites the URL to connect to the
validated IP directly. Prevents DNS rebinding by pinning to the resolved
IP. Disable follow_redirects to prevent redirect-based SSRF bypasses.

Applied to all user-supplied URL entry points:
- Image URL fetching in chat completions
- Token counter image dimension fetching
- RAG file ingestion
- MCP OpenAPI spec loading
This commit is contained in:
user 2026-04-16 04:13:54 +00:00
parent 7279dca929
commit 9363f36481
No known key found for this signature in database
5 changed files with 162 additions and 8 deletions

View file

@ -10,6 +10,7 @@ import litellm
from litellm import verbose_logger
from litellm.caching.caching import InMemoryCache
from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB
from litellm.proxy.common_utils.url_utils import SSRFError, validate_url
MAX_IMGS_IN_MEMORY = 10
@ -81,10 +82,17 @@ async def async_convert_url_to_base64(url: str) -> str:
if cached_result:
return cached_result
# Resolve DNS once, validate IPs, rewrite URL to validated IP
validated_url, original_host = validate_url(url)
client = litellm.module_level_aclient
for _ in range(3):
try:
response = await client.get(url, follow_redirects=True)
response = await client.get(
validated_url,
headers={"Host": original_host},
follow_redirects=False,
)
return _process_image_response(response, url)
except litellm.ImageFetchError:
raise
@ -106,10 +114,17 @@ def convert_url_to_base64(url: str) -> str:
if cached_result:
return cached_result
# Resolve DNS once, validate IPs, rewrite URL to validated IP
validated_url, original_host = validate_url(url)
client = litellm.module_level_client
for _ in range(3):
try:
response = client.get(url, follow_redirects=True)
response = client.get(
validated_url,
headers={"Host": original_host},
follow_redirects=False,
)
return _process_image_response(response, url)
except litellm.ImageFetchError:
raise

View file

@ -30,6 +30,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.default_encoding import encoding as default_encoding
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
from litellm.proxy.common_utils.url_utils import validate_url
from litellm.types.llms.anthropic import (
AnthropicMessagesToolResultParam,
AnthropicMessagesToolUseParam,
@ -211,9 +212,14 @@ def get_image_dimensions(
"""
img_data = None
try:
# Try to open as URL
# Try to open as URL — validate and pin to resolved IP
validated_url, original_host = validate_url(data)
client = _get_httpx_client()
response = client.get(data)
response = client.get(
validated_url,
headers={"Host": original_host},
follow_redirects=False,
)
img_data = response.read()
except Exception:
# If not URL, assume it's base64

View file

@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.common_utils.url_utils import validate_url
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
@ -74,10 +75,13 @@ def load_openapi_spec(filepath: str) -> Dict[str, Any]:
async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]:
if filepath.startswith("http://") or filepath.startswith("https://"):
validated_url, original_host = validate_url(filepath)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
# NOTE: do not close shared client if get_async_httpx_client returns a shared singleton.
# If it returns a new client each time, consider wrapping it in an async context manager.
r = await client.get(filepath)
r = await client.get(
validated_url,
headers={"Host": original_host},
follow_redirects=False,
)
r.raise_for_status()
return r.json()

View file

@ -0,0 +1,123 @@
"""
URL validation for user-controlled URLs.
Use validate_url() before fetching any URL that originates from user
input (image_url, file_url, spec_path, etc.) to prevent SSRF attacks.
The function resolves DNS once, validates all IPs, and rewrites the URL
to connect to the validated IP directly — no TOCTOU gap, no DNS rebinding.
Callers should also set follow_redirects=False to prevent redirect-based
SSRF bypasses.
"""
import ipaddress
import socket
from ipaddress import ip_address, ip_network
from typing import Optional, Tuple
from urllib.parse import urlparse, urlunparse
_BLOCKED_NETWORKS = [
ip_network("0.0.0.0/8"),
ip_network("10.0.0.0/8"),
ip_network("100.64.0.0/10"),
ip_network("127.0.0.0/8"),
ip_network("169.254.0.0/16"),
ip_network("172.16.0.0/12"),
ip_network("192.0.0.0/24"),
ip_network("192.168.0.0/16"),
ip_network("198.18.0.0/15"),
ip_network("::1/128"),
ip_network("fc00::/7"),
ip_network("fe80::/10"),
]
_ALLOWED_SCHEMES = ("http", "https")
class SSRFError(ValueError):
"""Raised when a URL targets a blocked network."""
pass
def _is_blocked_ip(addr: str) -> bool:
try:
ip = ip_address(addr)
except ValueError:
return False
if ip.version == 6 and hasattr(ip, "ipv4_mapped") and ip.ipv4_mapped:
ip = ip.ipv4_mapped
return any(ip in net for net in _BLOCKED_NETWORKS)
def validate_url(url: str) -> Tuple[str, str]:
"""
Validate a user-supplied URL and rewrite it to connect to a validated IP.
Resolves the hostname, checks all resolved IPs against blocked networks,
then returns a rewritten URL that points to the validated IP along with
the original hostname (for use in the Host header).
This eliminates DNS rebinding because the caller connects to the IP we
validated, not the hostname that could rebind. Callers should also disable
follow_redirects to prevent redirect-based SSRF bypasses.
Args:
url: The user-supplied URL to validate.
Returns:
Tuple of (rewritten_url, original_hostname).
The rewritten URL has the hostname replaced with the validated IP.
The original hostname should be set as the Host header.
Raises:
SSRFError: If the URL scheme is invalid or the hostname resolves
to a private/internal IP address.
"""
parsed = urlparse(url)
if parsed.scheme not in _ALLOWED_SCHEMES:
raise SSRFError(f"URL scheme '{parsed.scheme}' is not allowed")
hostname = parsed.hostname
if not hostname:
raise SSRFError("URL has no hostname")
port = parsed.port
default_port = 443 if parsed.scheme == "https" else 80
# Resolve hostname and validate ALL addresses
try:
addrinfo = socket.getaddrinfo(
hostname, port or default_port, proto=socket.IPPROTO_TCP
)
except socket.gaierror as e:
raise SSRFError(f"DNS resolution failed for '{hostname}': {e}")
if not addrinfo:
raise SSRFError(f"No addresses found for '{hostname}'")
for family, type_, proto, canonname, sockaddr in addrinfo:
if _is_blocked_ip(sockaddr[0]):
raise SSRFError(
f"URL targets a blocked address ({sockaddr[0]}). "
"If this is a legitimate internal service, use a direct "
"provider configuration instead of a user-supplied URL."
)
# Rewrite URL to connect to the first validated IP
validated_ip = addrinfo[0][4][0]
is_ipv6 = addrinfo[0][0] == socket.AF_INET6
ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip
# Reconstruct netloc with IP instead of hostname
if port:
new_netloc = f"{ip_host}:{port}"
else:
new_netloc = ip_host
rewritten = urlunparse(
(parsed.scheme, new_netloc, parsed.path, parsed.params, parsed.query, "")
)
return rewritten, hostname

View file

@ -24,6 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.common_utils.url_utils import validate_url
from litellm.rag.ingestion.file_parsers import extract_text_from_pdf
from litellm.rag.text_splitters import RecursiveCharacterTextSplitter
from litellm.types.rag import RAGIngestOptions, RAGIngestResponse
@ -111,8 +112,13 @@ class BaseRAGIngestion(ABC):
return filename, file_content, content_type, None
if file_url:
validated_url, original_host = validate_url(file_url)
http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG)
response = await http_client.get(file_url)
response = await http_client.get(
validated_url,
headers={"Host": original_host},
follow_redirects=False,
)
response.raise_for_status()
file_content = response.content
filename = file_url.split("/")[-1] or "document"