diff --git a/litellm/proxy/common_utils/static_asset_utils.py b/litellm/proxy/common_utils/static_asset_utils.py index ffead95dcfe..74fe9939aab 100644 --- a/litellm/proxy/common_utils/static_asset_utils.py +++ b/litellm/proxy/common_utils/static_asset_utils.py @@ -18,7 +18,7 @@ import os from typing import List, Optional from litellm._logging import verbose_proxy_logger -from litellm.litellm_core_utils.url_utils import SSRFError, validate_url +from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider @@ -26,13 +26,19 @@ from litellm.types.llms.custom_http import httpxSpecialProvider # without this, an admin-configured URL whose upstream returns # ``application/json`` (e.g. cloud metadata, internal API) would still be # served back to the caller verbatim. +# +# ``image/svg+xml`` is intentionally NOT in this list: SVG is the only +# common image format that can embed JavaScript, and the endpoint is +# unauthenticated. An admin-configured CDN serving a crafted SVG would +# otherwise reach unauthenticated callers; removing SVG closes the +# residual XSS surface even though the response is served with a +# hardcoded ``image/jpeg`` / ``image/x-icon`` media type. ALLOWED_IMAGE_CONTENT_TYPES = frozenset( { "image/jpeg", "image/jpg", "image/png", "image/gif", - "image/svg+xml", "image/webp", "image/x-icon", "image/vnd.microsoft.icon", @@ -76,19 +82,27 @@ async def fetch_validated_image_bytes( url: str, *, timeout_s: float = 5.0 ) -> Optional[bytes]: """ - Fetch ``url`` with SSRF protection (always-on) and Content-Type - validation. Returns the raw bytes on success, ``None`` on any - failure (blocked target, non-200, or non-image response). + Fetch ``url`` with SSRF protection and Content-Type validation. + Returns the raw bytes on success, ``None`` on any failure (blocked + target, redirect to a blocked target, non-200, or non-image + response). - The SSRF guard is enforced unconditionally — these endpoints are - unauthenticated, so the admin-facing ``litellm.user_url_validation`` - toggle does not apply. An admin who opted out of URL validation for - LLM provider paths should not also expose ``/get_image`` to SSRF. + Delegates to ``async_safe_get`` so each redirect hop is re-validated + against ``BLOCKED_NETWORKS`` (a 3xx to ``169.254.169.254`` is + rejected, not followed). Honours ``litellm.user_url_validation`` + like every other SSRF-aware fetch in the codebase; the toggle + defaults to True, and an admin who has explicitly disabled URL + validation has opted out of SSRF protection globally. """ if not url: return None + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.UI, + params={"timeout": timeout_s}, + ) try: - rewritten_url, host_header = validate_url(url) + response = await async_safe_get(async_client, url) except SSRFError as exc: verbose_proxy_logger.warning( "Blocked unauthenticated asset fetch — SSRF guard rejected %r: %s", @@ -96,22 +110,6 @@ async def fetch_validated_image_bytes( exc, ) return None - - # ``validate_url`` rewrites HTTP URLs to point at a validated IP and - # returns the original hostname for the Host header. For HTTPS with - # ssl_verify enabled, it returns the URL unchanged (TLS hostname - # validation handles DNS rebinding). - async_client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.UI, - params={"timeout": timeout_s}, - ) - try: - if rewritten_url != url: - response = await async_client.get( - rewritten_url, headers={"host": host_header} - ) - else: - response = await async_client.get(rewritten_url) except Exception as exc: verbose_proxy_logger.debug("Asset fetch failed for %r: %s", url, exc) return None diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ef17aab73e6..736f259f23a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -12312,16 +12312,21 @@ async def get_image(): # SSRF + content-type validation — the helper rejects # private/internal/cloud-metadata targets and non-image responses. image_bytes = await fetch_validated_image_bytes(logo_path) - if image_bytes is not None: - try: - with open(cache_path, "wb") as f: - f.write(image_bytes) - return FileResponse(cache_path, media_type="image/jpeg") - except OSError as e: - verbose_proxy_logger.debug( - "Could not write logo cache to %s: %s", cache_path, e - ) - return FileResponse(default_logo, media_type="image/jpeg") + if image_bytes is None: + return FileResponse(default_logo, media_type="image/jpeg") + try: + with open(cache_path, "wb") as f: + f.write(image_bytes) + return FileResponse(cache_path, media_type="image/jpeg") + except OSError as e: + # Read-only assets dir: serve the validated bytes inline + # rather than dropping them and returning the default logo. + from fastapi.responses import Response + + verbose_proxy_logger.debug( + "Could not write logo cache to %s: %s — serving inline", cache_path, e + ) + return Response(content=image_bytes, media_type="image/jpeg") else: # Default logo (resolved from the bundled asset, not user-controlled). return FileResponse(logo_path, media_type="image/jpeg") diff --git a/tests/proxy_unit_tests/test_get_image.py b/tests/proxy_unit_tests/test_get_image.py index f14c6da5539..bdc7743faac 100644 --- a/tests/proxy_unit_tests/test_get_image.py +++ b/tests/proxy_unit_tests/test_get_image.py @@ -65,12 +65,14 @@ async def test_get_image_cache_logic(): os.remove(cache_path) # Mock response — set headers explicitly so the Content-Type - # validation added for GHSA-pjc9-2hw6-78rr accepts the response - # as a legitimate image. + # validation accepts the response as a legitimate image, and set + # ``is_redirect=False`` so ``async_safe_get`` doesn't try to walk + # a redirect chain. mock_response = mock.Mock() mock_response.status_code = 200 mock_response.content = b"fake image data" mock_response.headers = {"content-type": "image/jpeg"} + mock_response.is_redirect = False with mock.patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get" diff --git a/tests/test_litellm/proxy/common_utils/test_static_asset_utils.py b/tests/test_litellm/proxy/common_utils/test_static_asset_utils.py index 6fe3b04bf02..0c6b7f18973 100644 --- a/tests/test_litellm/proxy/common_utils/test_static_asset_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_static_asset_utils.py @@ -101,121 +101,89 @@ class TestResolveLocalAssetPath: class TestFetchValidatedImageBytes: - @pytest.fixture - def mock_async_client(self): - client = MagicMock() - client.get = AsyncMock() - return client + """ + The helper now delegates to ``async_safe_get`` for the SSRF guard + + redirect handling. Tests mock ``async_safe_get`` directly so they + exercise the helper's contract (Content-Type validation, status code + handling, exception fallthrough) without depending on the SSRF + primitive's internals. + """ - @pytest.mark.asyncio - async def test_blocks_private_ip_via_validate_url(self, mock_async_client): - # The SSRF half of GHSA-pjc9-2hw6-78rr — admin sets logo URL to - # http://169.254.169.254/iam, attacker hits /get_image, exfils creds. - with ( + @staticmethod + def _patches(*, async_safe_get_return=None, async_safe_get_side_effect=None): + return [ patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - side_effect=SSRFError("blocked: 169.254.169.254"), + "litellm.proxy.common_utils.static_asset_utils.async_safe_get", + new_callable=AsyncMock, + return_value=async_safe_get_return, + side_effect=async_safe_get_side_effect, ), patch( "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, + return_value=MagicMock(), ), - ): - result = await fetch_validated_image_bytes("http://169.254.169.254/iam") - - assert result is None - # The fetch must not be attempted when the URL is rejected. - mock_async_client.get.assert_not_called() + ] @pytest.mark.asyncio - async def test_rejects_non_image_content_type(self, mock_async_client): + async def test_blocks_ssrf_target(self): + # The SSRF half of GHSA-pjc9-2hw6-78rr — admin sets logo URL to + # http://169.254.169.254/iam, attacker hits /get_image, exfils creds. + # ``async_safe_get`` raises SSRFError on private/metadata targets + # and on redirect hops to those targets (covers the redirect + # bypass Veria flagged on the previous iteration). + with ( + self._patches( + async_safe_get_side_effect=SSRFError("blocked: 169.254.169.254") + )[0], + self._patches()[1], + ): + result = await fetch_validated_image_bytes("http://169.254.169.254/iam") + assert result is None + + @pytest.mark.asyncio + async def test_rejects_non_image_content_type(self): # Even when the URL passes SSRF, the upstream response must be an - # image. Otherwise an attacker could redirect to an upstream that + # image. Otherwise an attacker could point at an upstream that # returns ``application/json`` AWS creds and have them tunneled # through the ``image/jpeg`` response wrapper. mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {"content-type": "application/json"} mock_response.content = b'{"AccessKeyId": "..."}' - mock_async_client.get.return_value = mock_response - with ( - patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - return_value=("http://cdn.example/logo", "cdn.example"), - ), - patch( - "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, - ), - ): + with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]: result = await fetch_validated_image_bytes("http://cdn.example/logo") - assert result is None @pytest.mark.asyncio - async def test_returns_bytes_for_valid_image_response(self, mock_async_client): + async def test_returns_bytes_for_valid_image_response(self): png_bytes = b"\x89PNG\r\n\x1a\nfake png body" mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {"content-type": "image/png; charset=binary"} mock_response.content = png_bytes - mock_async_client.get.return_value = mock_response - with ( - patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - return_value=( - "https://cdn.example/logo.png", - "cdn.example", - ), - ), - patch( - "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, - ), - ): + with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]: result = await fetch_validated_image_bytes("https://cdn.example/logo.png") - assert result == png_bytes @pytest.mark.asyncio - async def test_returns_none_on_non_200_response(self, mock_async_client): + async def test_returns_none_on_non_200_response(self): mock_response = MagicMock() mock_response.status_code = 404 mock_response.headers = {"content-type": "image/png"} - mock_async_client.get.return_value = mock_response - with ( - patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - return_value=("https://cdn.example/logo", "cdn.example"), - ), - patch( - "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, - ), - ): + with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]: result = await fetch_validated_image_bytes("https://cdn.example/logo") - assert result is None @pytest.mark.asyncio - async def test_returns_none_on_fetch_exception(self, mock_async_client): - mock_async_client.get.side_effect = Exception("connection reset") - + async def test_returns_none_on_fetch_exception(self): with ( - patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - return_value=("https://cdn.example/logo", "cdn.example"), - ), - patch( - "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, - ), + self._patches(async_safe_get_side_effect=Exception("connection reset"))[0], + self._patches()[1], ): result = await fetch_validated_image_bytes("https://cdn.example/logo") - assert result is None @pytest.mark.asyncio @@ -223,30 +191,32 @@ class TestFetchValidatedImageBytes: result = await fetch_validated_image_bytes("") assert result is None + @pytest.mark.asyncio + async def test_rejects_svg_content_type(self): + # ``image/svg+xml`` is intentionally NOT in the allowlist for + # unauthenticated endpoints — SVG is the only common image + # format that can embed JavaScript. + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "image/svg+xml"} + mock_response.content = b"" + + with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]: + result = await fetch_validated_image_bytes("https://cdn.example/x.svg") + assert result is None + @pytest.mark.parametrize( "content_type", sorted(ALLOWED_IMAGE_CONTENT_TYPES), ) @pytest.mark.asyncio - async def test_accepts_each_allowed_image_content_type( - self, mock_async_client, content_type - ): + async def test_accepts_each_allowed_image_content_type(self, content_type): mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {"content-type": content_type} mock_response.content = b"image-bytes" - mock_async_client.get.return_value = mock_response - with ( - patch( - "litellm.proxy.common_utils.static_asset_utils.validate_url", - return_value=("https://cdn.example/logo", "cdn.example"), - ), - patch( - "litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client", - return_value=mock_async_client, - ), - ): + with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]: result = await fetch_validated_image_bytes("https://cdn.example/logo") assert result == b"image-bytes"