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.
This commit is contained in:
1Ckpwee 2026-03-26 22:41:18 +08:00
parent e7ae78dbbe
commit 5c1ab42498
2 changed files with 110 additions and 25 deletions

View file

@ -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)

View file

@ -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
# ---------------------------------------------------------------------------