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"