From 9363f36481b8602e7866398bd054a11dd9842e8e Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 04:13:54 +0000 Subject: [PATCH] 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 --- .../prompt_templates/image_handling.py | 19 ++- litellm/litellm_core_utils/token_counter.py | 10 +- .../mcp_server/openapi_to_mcp_generator.py | 10 +- litellm/proxy/common_utils/url_utils.py | 123 ++++++++++++++++++ litellm/rag/ingestion/base_ingestion.py | 8 +- 5 files changed, 162 insertions(+), 8 deletions(-) create mode 100644 litellm/proxy/common_utils/url_utils.py diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index eaf78b7bcf5..c0727699167 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -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 diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 09c62f2eb55..e2d2a56c698 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 4b4818892bb..68f52f34395 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -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() diff --git a/litellm/proxy/common_utils/url_utils.py b/litellm/proxy/common_utils/url_utils.py new file mode 100644 index 00000000000..97fee0c0965 --- /dev/null +++ b/litellm/proxy/common_utils/url_utils.py @@ -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 diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 0d12bdfffc1..1d868d35c7a 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -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"