chore(static-assets): /simplify pass — top-level Response import + cleaner test fixture

Two cleanups from the /simplify review pass:

* ``Response`` was imported inside the ``except OSError`` branch in
  ``/get_image`` and at the top of ``/get_favicon``. Per the project's
  no-inline-imports rule (CLAUDE.md), hoisted to the existing
  ``from fastapi.responses import (...)`` block at the top of
  ``proxy_server.py``.

* The test class's ``_patches()`` helper returned a 2-element list of
  patch context managers and tests indexed into them via
  ``self._patches(...)[0], self._patches()[1]`` — two distinct calls
  with confusing aliasing semantics. Restructured to:
    - module-level ``_patch_async_safe_get(...)`` that returns a single
      patch context manager
    - autouse fixture that patches ``get_async_httpx_client`` for every
      test in the file (it's the same patch in every case)
    - small ``_image_response(...)`` factory to deduplicate Mock setup

  Tests now read as ``with _patch_async_safe_get(return_value=...):``
  with no list-indexing or duplicate Mock construction.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
user 2026-04-29 22:01:57 +00:00
parent 75d1a0116e
commit c112bdf2c1
No known key found for this signature in database
2 changed files with 58 additions and 68 deletions

View file

@ -605,6 +605,7 @@ from fastapi.responses import (
JSONResponse,
ORJSONResponse,
RedirectResponse,
Response,
StreamingResponse,
)
from fastapi.routing import APIRouter
@ -12321,8 +12322,6 @@ async def get_image():
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
)
@ -12335,8 +12334,6 @@ async def get_image():
@app.get("/get_favicon", include_in_schema=False)
async def get_favicon():
"""Get custom favicon for the admin UI."""
from fastapi.responses import Response
from litellm.proxy.common_utils.static_asset_utils import (
fetch_validated_image_bytes,
resolve_local_asset_path,

View file

@ -100,89 +100,86 @@ class TestResolveLocalAssetPath:
assert result == str(logo.resolve())
def _image_response(*, status_code=200, content_type="image/png", body=b"image-bytes"):
response = MagicMock()
response.status_code = status_code
response.headers = {"content-type": content_type}
response.content = body
return response
def _patch_async_safe_get(*, return_value=None, side_effect=None):
return patch(
"litellm.proxy.common_utils.static_asset_utils.async_safe_get",
new_callable=AsyncMock,
return_value=return_value,
side_effect=side_effect,
)
@pytest.fixture(autouse=True)
def _patch_httpx_client():
# The helper builds the client first, then hands it to async_safe_get
# — patch it once for every test so we never accidentally instantiate
# a real client.
with patch(
"litellm.proxy.common_utils.static_asset_utils.get_async_httpx_client",
return_value=MagicMock(),
):
yield
class TestFetchValidatedImageBytes:
"""
The helper now delegates to ``async_safe_get`` for the SSRF guard +
The helper 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.
"""
@staticmethod
def _patches(*, async_safe_get_return=None, async_safe_get_side_effect=None):
return [
patch(
"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=MagicMock(),
),
]
@pytest.mark.asyncio
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],
):
# and on redirect hops to those targets — closes the SSRF half of
# GHSA-pjc9-2hw6-78rr including the redirect-bypass variant.
with _patch_async_safe_get(side_effect=SSRFError("blocked: 169.254.169.254")):
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 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": "..."}'
with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]:
# Without this, an upstream that returns ``application/json`` AWS
# creds would be tunneled through the ``image/jpeg`` response
# wrapper.
with _patch_async_safe_get(
return_value=_image_response(
content_type="application/json", body=b'{"AccessKeyId": "..."}'
),
):
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):
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
with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]:
with _patch_async_safe_get(
return_value=_image_response(
content_type="image/png; charset=binary", body=png_bytes
),
):
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_response = MagicMock()
mock_response.status_code = 404
mock_response.headers = {"content-type": "image/png"}
with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]:
with _patch_async_safe_get(return_value=_image_response(status_code=404)):
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):
with (
self._patches(async_safe_get_side_effect=Exception("connection reset"))[0],
self._patches()[1],
):
with _patch_async_safe_get(side_effect=Exception("connection reset")):
result = await fetch_validated_image_bytes("https://cdn.example/logo")
assert result is None
@ -196,12 +193,12 @@ class TestFetchValidatedImageBytes:
# ``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"<svg><script>alert(1)</script></svg>"
with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]:
with _patch_async_safe_get(
return_value=_image_response(
content_type="image/svg+xml",
body=b"<svg><script>alert(1)</script></svg>",
),
):
result = await fetch_validated_image_bytes("https://cdn.example/x.svg")
assert result is None
@ -211,12 +208,8 @@ class TestFetchValidatedImageBytes:
)
@pytest.mark.asyncio
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"
with self._patches(async_safe_get_return=mock_response)[0], self._patches()[1]:
with _patch_async_safe_get(
return_value=_image_response(content_type=content_type),
):
result = await fetch_validated_image_bytes("https://cdn.example/logo")
assert result == b"image-bytes"