diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py index de7690635d8..092086c79db 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -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"} diff --git a/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py b/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py new file mode 100644 index 00000000000..13d69f7b0e7 --- /dev/null +++ b/tests/litellm/proxy/guardrails/test_custom_code_ssrf.py @@ -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