mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
bdf4acc472
commit
e7ae78dbbe
2 changed files with 240 additions and 1 deletions
|
|
@ -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"}
|
||||
|
|
|
|||
148
tests/litellm/proxy/guardrails/test_custom_code_ssrf.py
Normal file
148
tests/litellm/proxy/guardrails/test_custom_code_ssrf.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue