diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py index 864ec052543..1553487ee67 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -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}") diff --git a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py index f93ecfc3010..b6e7f3b7025 100644 --- a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py +++ b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py @@ -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