mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
7279dca929
commit
9363f36481
5 changed files with 162 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
123
litellm/proxy/common_utils/url_utils.py
Normal file
123
litellm/proxy/common_utils/url_utils.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue