diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 3a7e06d085a..c24d2e806a3 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1,11 +1,15 @@ import base64 import json import os +import socket from unittest.mock import MagicMock, patch import pytest +import httpx + import litellm +from litellm.litellm_core_utils.url_utils import PayloadTooLargeError from litellm.litellm_core_utils.prompt_templates.factory import ( BAD_MESSAGE_ERROR_STR, BEDROCK_DOCUMENT_PLACEHOLDER_TEXT, @@ -3516,3 +3520,71 @@ async def test_bedrock_converse_pdf_only_user_message_gets_text_block_async(): assert len(result) == 1 assert any("document" in block for block in result[0]["content"]) assert _text_blocks(result[0]) == [BEDROCK_DOCUMENT_PLACEHOLDER_TEXT] + + +def _resolve_to_public(host, port, *args, **kwargs): + """Keep validate_url's DNS lookup off the network without faking the fetch.""" + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port or 443))] + + +class TestBedrockImageProcessorMaxBytes: + """`max_bytes` is threaded from the caller down to the fetch. + + The Bedrock guardrail is the only caller that sets it. Everything else, the + model-call image paths included, must keep the previous unbounded fetch, and the + keyword has to be absent from the call rather than merely defaulted -- an + override or stub written against the old signature would otherwise break. + """ + + _REMOTE_URL = "https://93.184.216.34/a.png" + + @staticmethod + def _fake_stream(chunks): + import contextlib + + @contextlib.asynccontextmanager + async def _stream(self, method, url, **kwargs): + async def aiter_bytes(): + for chunk in chunks: + yield chunk + + response = MagicMock() + response.status_code = 200 + response.headers = httpx.Headers({"content-type": "image/png"}) + response.request = httpx.Request("GET", str(url)) + response.aiter_bytes = aiter_bytes + yield response + + return _stream + + @pytest.mark.asyncio + async def test_a_remote_fetch_is_capped_when_max_bytes_is_given(self, monkeypatch): + monkeypatch.setattr(socket, "getaddrinfo", _resolve_to_public, raising=False) + + with patch.object(httpx.AsyncClient, "stream", new=self._fake_stream([b"\0" * 8192])): + with pytest.raises(PayloadTooLargeError): + await BedrockImageProcessor.get_image_details_async(self._REMOTE_URL, max_bytes=1024) + + @pytest.mark.asyncio + async def test_omitting_max_bytes_leaves_the_call_as_it_was(self, monkeypatch): + """A stub written against the previous one-parameter signature still works. + + This is what test_url_with_format_param asserts through the model path; here + it is pinned on the helper itself so the plumbing cannot start passing the + keyword unconditionally again. + """ + monkeypatch.setattr(socket, "getaddrinfo", _resolve_to_public, raising=False) + seen: list = [] + + async def one_parameter_stub(image_url): + seen.append(image_url) + return "ZmFrZQ==", "image/png" + + monkeypatch.setattr( + BedrockImageProcessor, "get_image_details_async", staticmethod(one_parameter_stub) + ) + + block = await BedrockImageProcessor.process_image_async(image_url=self._REMOTE_URL, format=None) + + assert seen == [self._REMOTE_URL] + assert block["image"]["source"]["bytes"] == "ZmFrZQ==" diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py index aaaa43a0dc4..e285f749ebc 100644 --- a/tests/test_litellm/litellm_core_utils/test_url_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -1,11 +1,17 @@ +import contextlib import socket +from types import SimpleNamespace +from unittest.mock import MagicMock, patch +import httpx import pytest import litellm from litellm.litellm_core_utils import url_utils from litellm.litellm_core_utils.url_utils import ( + PayloadTooLargeError, SSRFError, + _underlying_httpx_client, _is_blocked_ip, assert_same_origin, encode_url_path_segment, @@ -535,3 +541,91 @@ def test_assert_same_origin_error_message_does_not_leak_hostnames(): detail = str(exc.value) assert "attacker.example.com" not in detail assert "api.internal-corp.example" not in detail + + +class TestCappedFetch: + """`async_safe_get(max_bytes=...)` streams and aborts past the cap. + + `client.get` buffers the whole body first, so a caller-supplied url serving an + arbitrarily large or indefinitely chunked response is an unbounded allocation. + """ + + def test_a_client_with_no_httpx_client_is_rejected(self): + """Streaming needs the wrapped httpx client. + + AsyncHTTPHandler forwards get/post but not stream, so the wrapped `.client` + is what gets used. Something with neither is a programming error and says so, + rather than failing later inside httpx with nothing pointing back here. + """ + with pytest.raises(TypeError) as exc: + _underlying_httpx_client(object()) + + assert "no httpx client" in str(exc.value) + + def test_a_raw_httpx_client_is_used_as_is(self): + client = httpx.AsyncClient() + + assert _underlying_httpx_client(client) is client + + def test_a_wrapped_client_resolves_to_the_one_it_wraps(self): + inner = httpx.AsyncClient() + wrapper = SimpleNamespace(client=inner) + + assert _underlying_httpx_client(wrapper) is inner + + @pytest.mark.asyncio + async def test_the_body_is_cut_off_once_it_passes_the_cap(self, mock_dns_public): + """Asserting on how much was pulled is what separates a streamed abort from + buffering everything and rejecting afterwards.""" + served: list[int] = [] + chunks = [b"\0" * 1024 for _ in range(100)] + + @contextlib.asynccontextmanager + async def fake_stream(self, method, url, **kwargs): + async def aiter_bytes(): + for chunk in chunks: + served.append(len(chunk)) + yield chunk + + response = MagicMock() + response.status_code = 200 + response.headers = httpx.Headers({"content-type": "image/png"}) + response.request = httpx.Request("GET", str(url)) + response.aiter_bytes = aiter_bytes + yield response + + client = httpx.AsyncClient() + with patch.object(httpx.AsyncClient, "stream", new=fake_stream): + with pytest.raises(PayloadTooLargeError): + await url_utils.async_safe_get(client, "https://93.184.216.34/a.png", max_bytes=4096) + + assert sum(served) <= 5 * 1024, f"pulled {sum(served)} bytes past a 4096 byte cap" + assert len(served) < len(chunks), "the whole body was read before rejecting it" + + @pytest.mark.asyncio + async def test_a_body_inside_the_cap_comes_back_whole(self, mock_dns_public): + """The rebuilt response drops content-encoding and content-length: aiter_bytes + yields decoded bytes, so carrying those over would describe the body wrongly.""" + + @contextlib.asynccontextmanager + async def fake_stream(self, method, url, **kwargs): + async def aiter_bytes(): + yield b"tiny-image" + + response = MagicMock() + response.status_code = 200 + response.headers = httpx.Headers( + {"content-type": "image/png", "content-length": "999", "content-encoding": "gzip"} + ) + response.request = httpx.Request("GET", str(url)) + response.aiter_bytes = aiter_bytes + yield response + + client = httpx.AsyncClient() + with patch.object(httpx.AsyncClient, "stream", new=fake_stream): + result = await url_utils.async_safe_get(client, "https://93.184.216.34/a.png", max_bytes=4096) + + assert result.content == b"tiny-image" + assert result.headers.get("content-type") == "image/png" + assert "content-encoding" not in result.headers + assert result.headers.get("content-length") == str(len(b"tiny-image"))