From 5c1ab42498ca6c97511edc1b85ef90d5a1f0c3dd Mon Sep 17 00:00:00 2001 From: 1Ckpwee <972193026zy@gmail.com> Date: Thu, 26 Mar 2026 22:41:18 +0800 Subject: [PATCH] fix: address code review feedback on SSRF protection 1. Fix TOCTOU DNS rebinding bypass: _validate_url_for_ssrf now returns the validated IP address. http_request rewrites the URL to connect directly to the pinned IP (via _build_pinned_url) and sets the original Host header, so httpx never re-resolves DNS independently. 2. Add missing IPv6 multicast range (ff00::/8) to _BLOCKED_NETWORKS for parity with IPv4 multicast (224.0.0.0/4). 3. Fix test_blocks_unresolvable_host: mock now raises socket.gaierror (the subclass caught by production code) instead of bare OSError. 4. Add tests for _build_pinned_url and IPv6 multicast blocking. --- .../guardrail_hooks/custom_code/primitives.py | 71 +++++++++++++++---- .../proxy/guardrails/test_custom_code_ssrf.py | 64 +++++++++++++---- 2 files changed, 110 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py index 092086c79db..3e17e86d921 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -382,6 +382,7 @@ _BLOCKED_NETWORKS = [ ipaddress.ip_network("fc00::/7"), # Unique local ipaddress.ip_network("fe80::/10"), # Link-local ipaddress.ip_network("::ffff:0:0/96"), # IPv4-mapped IPv6 + ipaddress.ip_network("ff00::/8"), # IPv6 multicast ] @@ -390,24 +391,37 @@ def _is_private_ip(addr: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool: return any(addr in network for network in _BLOCKED_NETWORKS) -def _validate_url_for_ssrf(url: str) -> Optional[str]: +def _validate_url_for_ssrf( + url: str, +) -> Tuple[Optional[str], Optional[str]]: """ Validate a URL against SSRF attacks by resolving the hostname and checking that none of the resolved IP addresses are private/reserved. - Returns None if the URL is safe, or an error message string if blocked. + To prevent TOCTOU / DNS-rebinding attacks, this function returns the + first validated public IP so the caller can connect directly to it + instead of letting httpx re-resolve the hostname. + + Returns: + (error, validated_ip) — *error* is a human-readable message when the + URL is blocked (validated_ip will be None), or None when the URL is + safe (validated_ip will contain the resolved address to use). """ parsed = urlparse(url) hostname = parsed.hostname if not hostname: - return "URL has no hostname" + return "URL has no hostname", None # Block raw IP addresses in private ranges (skip DNS) try: addr = ipaddress.ip_address(hostname) if _is_private_ip(addr): - return f"Requests to private/reserved IP address {hostname} are not allowed" - return None + return ( + f"Requests to private/reserved IP address {hostname} are not allowed", + None, + ) + # Raw public IP — no DNS needed, pin to itself + return None, hostname except ValueError: pass # Not a raw IP — resolve via DNS below @@ -416,10 +430,10 @@ def _validate_url_for_ssrf(url: str) -> Optional[str]: try: addrinfos = socket.getaddrinfo(hostname, port, proto=socket.IPPROTO_TCP) except socket.gaierror: - return f"Could not resolve hostname: {hostname}" + return f"Could not resolve hostname: {hostname}", None if not addrinfos: - return f"No addresses found for hostname: {hostname}" + return f"No addresses found for hostname: {hostname}", None for family, _type, _proto, _canonname, sockaddr in addrinfos: ip_str = sockaddr[0] @@ -429,11 +443,32 @@ def _validate_url_for_ssrf(url: str) -> Optional[str]: return ( f"Hostname {hostname} resolves to private/reserved address " f"{ip_str}, request blocked" - ) + ), None except ValueError: continue - return None + # All addresses are public — return the first one to pin the connection + first_ip = addrinfos[0][4][0] + return None, first_ip + + +def _build_pinned_url(original_url: str, validated_ip: str) -> Tuple[str, str]: + """ + Rewrite *original_url* so the hostname is replaced by *validated_ip*, + and return (rewritten_url, original_hostname) so the caller can set + a ``Host`` header preserving the original hostname for TLS / vhosts. + """ + parsed = urlparse(original_url) + original_host = parsed.hostname or "" + # Bracket IPv6 addresses for URL syntax + ip_host = f"[{validated_ip}]" if ":" in validated_ip else validated_ip + # Preserve explicit port if present + if parsed.port: + netloc = f"{ip_host}:{parsed.port}" + else: + netloc = ip_host + pinned = parsed._replace(netloc=netloc).geturl() + return pinned, original_host # ============================================================================= @@ -539,14 +574,26 @@ async def http_request( if not is_valid_url(url): return _http_error_response(f"Invalid URL: {url}") - # SSRF protection: block requests to private/reserved IP ranges - ssrf_error = _validate_url_for_ssrf(url) + # SSRF protection: block requests to private/reserved IP ranges. + # _validate_url_for_ssrf resolves DNS once and returns a validated IP + # so that we can pin the connection to it, preventing TOCTOU / + # DNS-rebinding attacks where the hostname re-resolves to a different + # (private) address between check and use. + ssrf_error, validated_ip = _validate_url_for_ssrf(url) if ssrf_error: verbose_proxy_logger.warning( "Custom code guardrail SSRF blocked: %s (url=%s)", ssrf_error, url ) return _http_error_response(ssrf_error) + # Rewrite the URL to connect directly to the validated IP, and + # preserve the original Host header for TLS SNI / virtual hosts. + pinned_url, original_host = _build_pinned_url(url, validated_ip or "") + if headers is None: + headers = {} + if original_host and "Host" not in headers: + headers["Host"] = original_host + # Validate and normalize method method = method.upper() allowed_methods = {"GET", "POST", "PUT", "DELETE", "PATCH"} @@ -569,7 +616,7 @@ async def http_request( try: response = await _execute_http_request( - client, method, url, headers, body, timeout + client, method, pinned_url, headers, body, timeout ) return _http_success_response(response) diff --git a/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py b/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py index 13d69f7b0e7..9708ee608ff 100644 --- a/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py +++ b/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py @@ -3,11 +3,13 @@ Tests for SSRF protection in custom code guardrail HTTP primitives. """ import ipaddress +import socket from unittest.mock import patch import pytest from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import ( + _build_pinned_url, _is_private_ip, _validate_url_for_ssrf, http_request, @@ -35,6 +37,7 @@ class TestIsPrivateIp: "::1", # IPv6 loopback "fc00::1", # IPv6 unique-local "fe80::1", # IPv6 link-local + "ff02::1", # IPv6 multicast ], ) def test_private_ips_blocked(self, ip): @@ -56,7 +59,7 @@ class TestIsPrivateIp: # --------------------------------------------------------------------------- -# _validate_url_for_ssrf +# _validate_url_for_ssrf (now returns (error, validated_ip) tuple) # --------------------------------------------------------------------------- @@ -64,26 +67,30 @@ class TestValidateUrlForSsrf: """URL-level SSRF validation.""" def test_blocks_raw_private_ipv4(self): - err = _validate_url_for_ssrf("http://127.0.0.1/admin") + err, ip = _validate_url_for_ssrf("http://127.0.0.1/admin") assert err is not None + assert ip is None assert "private" in err.lower() or "reserved" in err.lower() def test_blocks_metadata_endpoint(self): - err = _validate_url_for_ssrf( + err, ip = _validate_url_for_ssrf( "http://169.254.169.254/latest/meta-data/iam/security-credentials/" ) assert err is not None + assert ip is None def test_blocks_raw_private_ipv6(self): - err = _validate_url_for_ssrf("http://[::1]/secret") + err, ip = _validate_url_for_ssrf("http://[::1]/secret") assert err is not None + assert ip is None - def test_allows_public_ip(self): - err = _validate_url_for_ssrf("https://8.8.8.8/dns-query") + def test_allows_public_ip_and_returns_it(self): + err, ip = _validate_url_for_ssrf("https://8.8.8.8/dns-query") assert err is None + assert ip == "8.8.8.8" def test_blocks_no_hostname(self): - err = _validate_url_for_ssrf("file:///etc/passwd") + err, ip = _validate_url_for_ssrf("file:///etc/passwd") assert err is not None @patch("socket.getaddrinfo") @@ -92,26 +99,57 @@ class TestValidateUrlForSsrf: mock_getaddrinfo.return_value = [ (2, 1, 6, "", ("127.0.0.1", 80)), ] - err = _validate_url_for_ssrf("http://evil.example.com/steal") + err, ip = _validate_url_for_ssrf("http://evil.example.com/steal") assert err is not None + assert ip is None assert "private" in err.lower() or "reserved" in err.lower() @patch("socket.getaddrinfo") - def test_allows_dns_to_public(self, mock_getaddrinfo): - """Hostname resolves to a public IP — should be allowed.""" + def test_allows_dns_to_public_and_returns_pinned_ip(self, mock_getaddrinfo): + """Hostname resolves to a public IP — return it for pinning.""" mock_getaddrinfo.return_value = [ (2, 1, 6, "", ("151.101.1.140", 443)), ] - err = _validate_url_for_ssrf("https://api.example.com/check") + err, ip = _validate_url_for_ssrf("https://api.example.com/check") assert err is None + assert ip == "151.101.1.140" - @patch("socket.getaddrinfo", side_effect=OSError("DNS failure")) + @patch("socket.getaddrinfo", side_effect=socket.gaierror("DNS failure")) def test_blocks_unresolvable_host(self, mock_getaddrinfo): - err = _validate_url_for_ssrf("http://doesnotexist.invalid/path") + err, ip = _validate_url_for_ssrf("http://doesnotexist.invalid/path") assert err is not None + assert ip is None assert "resolve" in err.lower() +# --------------------------------------------------------------------------- +# _build_pinned_url +# --------------------------------------------------------------------------- + + +class TestBuildPinnedUrl: + """Verify URL rewriting for IP pinning.""" + + def test_replaces_hostname_with_ip(self): + pinned, host = _build_pinned_url("https://example.com/path", "93.184.216.34") + assert "93.184.216.34" in pinned + assert host == "example.com" + + def test_preserves_explicit_port(self): + pinned, host = _build_pinned_url( + "http://example.com:8080/api", "93.184.216.34" + ) + assert "93.184.216.34:8080" in pinned + assert host == "example.com" + + def test_brackets_ipv6(self): + pinned, host = _build_pinned_url( + "https://example.com/path", "2607:f8b0:4004:800::200e" + ) + assert "[2607:f8b0:4004:800::200e]" in pinned + assert host == "example.com" + + # --------------------------------------------------------------------------- # http_request integration # ---------------------------------------------------------------------------