mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
e7ae78dbbe
commit
5c1ab42498
2 changed files with 110 additions and 25 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue