fix(guardrails): restore SSRF protection in custom code HTTP primitives

The custom code guardrail sandbox exposes http_get/http_post/http_request
so guardrails can call external moderation APIs. The SSRF controls from
#25004 (_validate_url_for_ssrf + follow_redirects=False) were dropped in
the Starlark -> RestrictedPython sandbox migration, so guardrail code
could reach loopback, RFC1918, and cloud metadata endpoints again, and
all five HTTP methods followed redirects.

Route the primitives through the canonical validator instead of
restoring the old hand-rolled check:

- validate_url() blocks non-globally-routable targets before connecting
  and eliminates DNS rebinding for plain http via resolve-and-rewrite
- operators can allow specific internal hosts via user_url_allowed_hosts
  in general_settings, or disable validation with
  litellm.user_url_validation = False
- follow_redirects=False on GET/POST/PUT/DELETE/PATCH so a 3xx cannot
  bypass validation

Adds regression tests to test_custom_code_security.py covering blocked
targets, the allowlist and master-switch escape hatches, Host header
rewriting, and redirect refusal on every method.
This commit is contained in:
maycuatroi1 2026-08-15 06:23:20 +07:00
parent bd0d13566e
commit d8ce786334
2 changed files with 212 additions and 33 deletions

View file

@ -12,7 +12,9 @@ from urllib.parse import urlparse
import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -417,6 +419,14 @@ async def http_request(
Uses LiteLLM's global cached AsyncHTTPHandler for connection pooling
and better performance.
SSRF protection: the URL is resolved and validated (via
``litellm_core_utils.url_utils.validate_url``) before connecting, and
redirects are never followed. Requests targeting loopback, private,
link-local, or cloud-metadata addresses are blocked. Operators can
allow specific internal hosts via ``user_url_allowed_hosts`` in
``general_settings``, or disable validation with
``litellm.user_url_validation = False``.
Args:
url: The URL to request
method: HTTP method (GET, POST, PUT, DELETE, PATCH). Defaults to GET.
@ -450,6 +460,19 @@ async def http_request(
if not is_valid_url(url):
return _http_error_response(f"Invalid URL: {url}")
# SSRF protection: block requests that resolve to internal/private/metadata
# targets before any connection is made. Legitimate internal services can
# be reached by adding them to `user_url_allowed_hosts` in general_settings.
# `litellm.user_url_validation = False` disables this check entirely.
validated_url = url
if getattr(litellm, "user_url_validation", True):
try:
validated_url, host_header = validate_url(url)
headers = {**(headers or {}), "Host": host_header}
except SSRFError as e:
verbose_proxy_logger.warning("Custom code http_request SSRF blocked: %s", e)
return _http_error_response(f"Blocked: {e}")
# Validate and normalize method
method = method.upper()
allowed_methods: Final = {"GET", "POST", "PUT", "DELETE", "PATCH"}
@ -469,7 +492,7 @@ async def http_request(
)
try:
response: Final = await _execute_http_request(client, method, url, headers, body, timeout)
response: Final = await _execute_http_request(client, method, validated_url, headers, body, timeout)
return _http_success_response(response)
except httpx.TimeoutException as e:
@ -494,19 +517,51 @@ async def _execute_http_request(
body: Any | None,
timeout: float,
) -> httpx.Response:
"""Execute the HTTP request using the appropriate client method."""
"""Execute the HTTP request using the appropriate client method.
Redirects are disabled on every method so a 3xx response cannot bypass
the SSRF validation performed by the caller.
"""
json_body, data_body = _prepare_http_body(body)
if method == "GET":
return await client.get(url=url, headers=headers)
return await client.get(url=url, headers=headers, follow_redirects=False)
elif method == "POST":
return await client.post(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
return await client.post(
url=url,
headers=headers,
json=json_body,
data=data_body,
timeout=timeout,
follow_redirects=False,
)
elif method == "PUT":
return await client.put(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
return await client.put(
url=url,
headers=headers,
json=json_body,
data=data_body,
timeout=timeout,
follow_redirects=False,
)
elif method == "DELETE":
return await client.delete(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
return await client.delete(
url=url,
headers=headers,
json=json_body,
data=data_body,
timeout=timeout,
follow_redirects=False,
)
elif method == "PATCH":
return await client.patch(url=url, headers=headers, json=json_body, data=data_body, timeout=timeout)
return await client.patch(
url=url,
headers=headers,
json=json_body,
data=data_body,
timeout=timeout,
follow_redirects=False,
)
else:
raise ValueError(f"Unsupported HTTP method: {method}")

View file

@ -1,6 +1,10 @@
import socket
import httpx
import pytest
from fastapi import HTTPException
import litellm
from litellm.exceptions import ModifyResponseException
from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import (
CustomCodeCompilationError,
@ -77,18 +81,14 @@ def test_nfkc_homoglyph_rejected_at_compile():
[
# Literal dunder attribute access.
"def apply_guardrail(i, r, t):\n return str.__class__\n",
"def apply_guardrail(i, r, t):\n"
" return ().__class__.__bases__[0].__subclasses__()\n",
"def apply_guardrail(i, r, t):\n return ().__class__.__bases__[0].__subclasses__()\n",
# gi_code — on the transformer's restricted-names list.
"def apply_guardrail(i, r, t):\n"
" def g():\n yield 1\n"
" return g().gi_code\n",
"def apply_guardrail(i, r, t):\n def g():\n yield 1\n return g().gi_code\n",
# Import forms.
"import os\ndef apply_guardrail(i, r, t):\n return allow()\n",
"from subprocess import call\n"
"def apply_guardrail(i, r, t):\n return allow()\n",
"from subprocess import call\ndef apply_guardrail(i, r, t):\n return allow()\n",
# __import__ is rejected as an underscore-prefixed name.
"def apply_guardrail(i, r, t):\n" ' return __import__("os")\n',
'def apply_guardrail(i, r, t):\n return __import__("os")\n',
],
)
def test_compile_time_rejections(snippet: str):
@ -100,8 +100,7 @@ def test_compile_time_rejections(snippet: str):
"snippet",
[
# getattr is not in the sandbox builtins — NameError at call time.
"def apply_guardrail(i, r, t):\n"
' return getattr(str, "_"+"_class_"+"_")\n',
'def apply_guardrail(i, r, t):\n return getattr(str, "_"+"_class_"+"_")\n',
# setattr is guarded_setattr + full_write_guard — setting any attribute
# on a user-defined object raises TypeError, whether the name is a
# dunder or not.
@ -139,10 +138,7 @@ def test_documented_ssn_example_compiles_and_runs():
@pytest.mark.asyncio
async def test_async_guardrail_compiles_and_runs():
code = (
"async def apply_guardrail(inputs, request_data, input_type):\n"
" return allow()\n"
)
code = "async def apply_guardrail(inputs, request_data, input_type):\n return allow()\n"
guardrail = _compile(code)
from litellm.types.utils import GenericGuardrailAPIInputs
@ -156,10 +152,7 @@ async def test_async_guardrail_compiles_and_runs():
@pytest.mark.asyncio
async def test_custom_code_pre_call_block_uses_passthrough():
code = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' return block("blocked by test")\n'
)
code = 'def apply_guardrail(inputs, request_data, input_type):\n return block("blocked by test")\n'
guardrail = _compile(code)
with pytest.raises(ModifyResponseException) as exc_info:
@ -176,10 +169,7 @@ async def test_custom_code_pre_call_block_uses_passthrough():
@pytest.mark.asyncio
async def test_custom_code_post_call_block_raises_http_400():
code = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' return block("blocked by test")\n'
)
code = 'def apply_guardrail(inputs, request_data, input_type):\n return block("blocked by test")\n'
guardrail = _compile(code)
with pytest.raises(HTTPException) as exc_info:
@ -198,10 +188,7 @@ async def test_custom_code_post_call_block_raises_http_400():
def test_typical_sync_guardrail_still_works():
code = (
"def apply_guardrail(inputs, request_data, input_type):\n"
" return allow()\n"
)
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()\n"
guardrail = _compile(code)
assert guardrail._compiled_function is not None
@ -228,3 +215,140 @@ def test_augmented_assignment_works():
def test_missing_apply_guardrail_raises():
with pytest.raises(CustomCodeCompilationError, match="apply_guardrail"):
_compile("x = 1\n")
# --- SSRF protection on the HTTP primitives ---------------------------------
#
# The primitives run inside the sandbox but talk to the network with the
# proxy process's privileges. http_request/http_get/http_post must therefore
# refuse loopback, private, and cloud-metadata targets and must not follow
# redirects (a 302 to an internal address would otherwise bypass the check).
class _FakeAsyncClient:
def __init__(self, response=None):
self.response = response
self.calls = []
async def _record(self, **kwargs):
self.calls.append(kwargs)
if isinstance(self.response, Exception):
raise self.response
return self.response
async def get(self, **kwargs):
return await self._record(**kwargs)
async def post(self, **kwargs):
return await self._record(**kwargs)
async def put(self, **kwargs):
return await self._record(**kwargs)
async def delete(self, **kwargs):
return await self._record(**kwargs)
async def patch(self, **kwargs):
return await self._record(**kwargs)
def _ok_response(status_code=200):
return httpx.Response(status_code, request=httpx.Request("GET", "http://ok"))
@pytest.fixture
def _public_dns(monkeypatch):
"""Point every hostname at a globally routable IP so tests stay hermetic."""
def fake_getaddrinfo(host, port, *args, **kwargs):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port))]
monkeypatch.setattr("litellm.litellm_core_utils.url_utils.socket.getaddrinfo", fake_getaddrinfo)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"url",
[
"http://127.0.0.1:8971/",
"http://localhost/admin",
"http://10.1.2.3/",
"http://192.168.1.1/",
"http://169.254.169.254/latest/meta-data/iam/security-credentials/",
"http://[::1]/",
],
)
async def test_http_primitives_block_internal_targets(url, monkeypatch):
from litellm.proxy.guardrails.guardrail_hooks.custom_code import primitives
client = _FakeAsyncClient()
monkeypatch.setattr(primitives, "get_async_httpx_client", lambda **kwargs: client)
result = await primitives.http_request(url)
assert result["success"] is False
assert result["status_code"] == 0
assert "Blocked" in result["error"]
assert client.calls == []
@pytest.mark.asyncio
async def test_http_primitives_allow_public_url(monkeypatch, _public_dns):
from litellm.proxy.guardrails.guardrail_hooks.custom_code import primitives
client = _FakeAsyncClient(_ok_response())
monkeypatch.setattr(primitives, "get_async_httpx_client", lambda **kwargs: client)
result = await primitives.http_get("http://moderation.example.com/v1/check", headers={"Authorization": "Bearer t"})
assert result["success"] is True
assert len(client.calls) == 1
call = client.calls[0]
assert call["url"] == "http://93.184.216.34/v1/check"
assert call["headers"]["Host"] == "moderation.example.com"
assert call["headers"]["Authorization"] == "Bearer t"
assert call["follow_redirects"] is False
@pytest.mark.asyncio
async def test_http_primitives_honor_allowlisted_internal_host(monkeypatch, _public_dns):
from litellm.proxy.guardrails.guardrail_hooks.custom_code import primitives
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal-moderation.corp"], raising=False)
client = _FakeAsyncClient(_ok_response())
monkeypatch.setattr(primitives, "get_async_httpx_client", lambda **kwargs: client)
result = await primitives.http_get("http://internal-moderation.corp/check")
assert result["success"] is True
assert client.calls[0]["follow_redirects"] is False
@pytest.mark.asyncio
async def test_http_primitives_validation_can_be_disabled(monkeypatch, _public_dns):
from litellm.proxy.guardrails.guardrail_hooks.custom_code import primitives
monkeypatch.setattr(litellm, "user_url_validation", False, raising=False)
client = _FakeAsyncClient(_ok_response())
monkeypatch.setattr(primitives, "get_async_httpx_client", lambda **kwargs: client)
result = await primitives.http_get("http://10.1.2.3/internal")
assert result["success"] is True
assert client.calls[0]["url"] == "http://10.1.2.3/internal"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["POST", "PUT", "DELETE", "PATCH"])
async def test_http_primitives_do_not_follow_redirects(method, monkeypatch, _public_dns):
from litellm.proxy.guardrails.guardrail_hooks.custom_code import primitives
client = _FakeAsyncClient(_ok_response(status_code=302))
monkeypatch.setattr(primitives, "get_async_httpx_client", lambda **kwargs: client)
result = await primitives.http_request("http://moderation.example.com/v1/check", method=method, body={"text": "hi"})
assert len(client.calls) == 1
assert client.calls[0]["follow_redirects"] is False
assert result["status_code"] == 302
assert result["success"] is False