fix(security): add SSRF protection to custom code guardrail HTTP primitives

The `http_request`, `http_get`, and `http_post` primitives in the custom
code guardrail sandbox only validated URL syntax via `is_valid_url()`,
which merely checks for a scheme and netloc. This allowed guardrail code
to reach internal services and cloud metadata endpoints
(e.g. 169.254.169.254), leading to Server-Side Request Forgery (SSRF).

This commit adds `_validate_url_for_ssrf()` which:
- Blocks requests to all RFC 1918 private ranges (10/8, 172.16/12, 192.168/16)
- Blocks loopback (127/8, ::1), link-local (169.254/16, fe80::/10)
- Blocks cloud metadata endpoint 169.254.169.254
- Blocks carrier-grade NAT (100.64/10), multicast, broadcast, etc.
- Resolves hostnames via DNS and validates all resulting IPs to prevent
  DNS rebinding attacks
- Blocks IPv4-mapped IPv6 addresses (::ffff:0:0/96)

Includes unit tests covering private IPs, metadata endpoints, IPv6
loopback, DNS rebinding scenarios, and public IP allowlisting.

Ref: #21259
This commit is contained in:
1Ckpwee 2026-03-26 22:19:37 +08:00
parent bdf4acc472
commit e7ae78dbbe
2 changed files with 240 additions and 1 deletions

View file

@ -5,8 +5,10 @@ These functions are injected into the custom code execution environment
and provide safe, sandboxed functionality for common guardrail operations.
"""
import ipaddress
import json
import re
import socket
from typing import Any, Dict, List, Optional, Tuple, Type, Union
from urllib.parse import urlparse
@ -353,6 +355,87 @@ def get_url_domain(url: str) -> Optional[str]:
return None
# =============================================================================
# SSRF Protection
# =============================================================================
# Private/reserved IP networks that must not be reachable from guardrail code.
_BLOCKED_NETWORKS = [
ipaddress.ip_network("0.0.0.0/8"), # "This" network
ipaddress.ip_network("10.0.0.0/8"), # RFC 1918
ipaddress.ip_network("100.64.0.0/10"), # Carrier-grade NAT
ipaddress.ip_network("127.0.0.0/8"), # Loopback
ipaddress.ip_network("169.254.0.0/16"), # Link-local / cloud metadata
ipaddress.ip_network("172.16.0.0/12"), # RFC 1918
ipaddress.ip_network("192.0.0.0/24"), # IETF protocol assignments
ipaddress.ip_network("192.0.2.0/24"), # TEST-NET-1
ipaddress.ip_network("192.88.99.0/24"), # 6to4 relay anycast
ipaddress.ip_network("192.168.0.0/16"), # RFC 1918
ipaddress.ip_network("198.18.0.0/15"), # Benchmarking
ipaddress.ip_network("198.51.100.0/24"), # TEST-NET-2
ipaddress.ip_network("203.0.113.0/24"), # TEST-NET-3
ipaddress.ip_network("224.0.0.0/4"), # Multicast
ipaddress.ip_network("240.0.0.0/4"), # Reserved for future use
ipaddress.ip_network("255.255.255.255/32"), # Broadcast
# IPv6
ipaddress.ip_network("::1/128"), # Loopback
ipaddress.ip_network("fc00::/7"), # Unique local
ipaddress.ip_network("fe80::/10"), # Link-local
ipaddress.ip_network("::ffff:0:0/96"), # IPv4-mapped IPv6
]
def _is_private_ip(addr: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
"""Return True if *addr* belongs to any blocked network."""
return any(addr in network for network in _BLOCKED_NETWORKS)
def _validate_url_for_ssrf(url: 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.
"""
parsed = urlparse(url)
hostname = parsed.hostname
if not hostname:
return "URL has no hostname"
# 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
except ValueError:
pass # Not a raw IP — resolve via DNS below
# Resolve hostname and check every resulting address
port = parsed.port or (443 if parsed.scheme == "https" else 80)
try:
addrinfos = socket.getaddrinfo(hostname, port, proto=socket.IPPROTO_TCP)
except socket.gaierror:
return f"Could not resolve hostname: {hostname}"
if not addrinfos:
return f"No addresses found for hostname: {hostname}"
for family, _type, _proto, _canonname, sockaddr in addrinfos:
ip_str = sockaddr[0]
try:
addr = ipaddress.ip_address(ip_str)
if _is_private_ip(addr):
return (
f"Hostname {hostname} resolves to private/reserved address "
f"{ip_str}, request blocked"
)
except ValueError:
continue
return None
# =============================================================================
# HTTP Request Primitives (Async)
# =============================================================================
@ -452,10 +535,18 @@ async def http_request(
body={"text": "content to check"}
)
"""
# Validate URL
# Validate URL syntax
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)
if ssrf_error:
verbose_proxy_logger.warning(
"Custom code guardrail SSRF blocked: %s (url=%s)", ssrf_error, url
)
return _http_error_response(ssrf_error)
# Validate and normalize method
method = method.upper()
allowed_methods = {"GET", "POST", "PUT", "DELETE", "PATCH"}

View file

@ -0,0 +1,148 @@
"""
Tests for SSRF protection in custom code guardrail HTTP primitives.
"""
import ipaddress
from unittest.mock import patch
import pytest
from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import (
_is_private_ip,
_validate_url_for_ssrf,
http_request,
)
# ---------------------------------------------------------------------------
# _is_private_ip
# ---------------------------------------------------------------------------
class TestIsPrivateIp:
"""Verify that private/reserved addresses are correctly identified."""
@pytest.mark.parametrize(
"ip",
[
"127.0.0.1",
"10.0.0.1",
"172.16.0.1",
"192.168.1.1",
"169.254.169.254", # AWS/GCP metadata
"0.0.0.0",
"100.64.0.1", # Carrier-grade NAT
"::1", # IPv6 loopback
"fc00::1", # IPv6 unique-local
"fe80::1", # IPv6 link-local
],
)
def test_private_ips_blocked(self, ip):
addr = ipaddress.ip_address(ip)
assert _is_private_ip(addr) is True
@pytest.mark.parametrize(
"ip",
[
"8.8.8.8",
"1.1.1.1",
"151.101.1.140",
"2607:f8b0:4004:800::200e", # Google public IPv6
],
)
def test_public_ips_allowed(self, ip):
addr = ipaddress.ip_address(ip)
assert _is_private_ip(addr) is False
# ---------------------------------------------------------------------------
# _validate_url_for_ssrf
# ---------------------------------------------------------------------------
class TestValidateUrlForSsrf:
"""URL-level SSRF validation."""
def test_blocks_raw_private_ipv4(self):
err = _validate_url_for_ssrf("http://127.0.0.1/admin")
assert err is not None
assert "private" in err.lower() or "reserved" in err.lower()
def test_blocks_metadata_endpoint(self):
err = _validate_url_for_ssrf(
"http://169.254.169.254/latest/meta-data/iam/security-credentials/"
)
assert err is not None
def test_blocks_raw_private_ipv6(self):
err = _validate_url_for_ssrf("http://[::1]/secret")
assert err is not None
def test_allows_public_ip(self):
err = _validate_url_for_ssrf("https://8.8.8.8/dns-query")
assert err is None
def test_blocks_no_hostname(self):
err = _validate_url_for_ssrf("file:///etc/passwd")
assert err is not None
@patch("socket.getaddrinfo")
def test_blocks_dns_rebinding_to_private(self, mock_getaddrinfo):
"""Hostname resolves to a private IP — must be blocked."""
mock_getaddrinfo.return_value = [
(2, 1, 6, "", ("127.0.0.1", 80)),
]
err = _validate_url_for_ssrf("http://evil.example.com/steal")
assert err is not 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."""
mock_getaddrinfo.return_value = [
(2, 1, 6, "", ("151.101.1.140", 443)),
]
err = _validate_url_for_ssrf("https://api.example.com/check")
assert err is None
@patch("socket.getaddrinfo", side_effect=OSError("DNS failure"))
def test_blocks_unresolvable_host(self, mock_getaddrinfo):
err = _validate_url_for_ssrf("http://doesnotexist.invalid/path")
assert err is not None
assert "resolve" in err.lower()
# ---------------------------------------------------------------------------
# http_request integration
# ---------------------------------------------------------------------------
class TestHttpRequestSsrf:
"""End-to-end: http_request must reject SSRF attempts."""
@pytest.mark.asyncio
async def test_http_request_blocks_localhost(self):
result = await http_request("http://127.0.0.1:8080/admin")
assert result["success"] is False
assert result["error"] is not None
assert "private" in result["error"].lower() or "reserved" in result["error"].lower()
@pytest.mark.asyncio
async def test_http_request_blocks_metadata(self):
result = await http_request(
"http://169.254.169.254/latest/meta-data/"
)
assert result["success"] is False
assert result["error"] is not None
@pytest.mark.asyncio
async def test_http_request_blocks_internal_network(self):
result = await http_request("http://10.0.0.1/internal-api")
assert result["success"] is False
assert result["error"] is not None
@pytest.mark.asyncio
async def test_http_request_blocks_ipv6_loopback(self):
result = await http_request("http://[::1]/secret")
assert result["success"] is False
assert result["error"] is not None