diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 52cfdeac000..5ba96f87c53 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -9,7 +9,7 @@ import threading import time import weakref import zlib -from collections.abc import AsyncGenerator, AsyncIterable, Callable, Iterable, Mapping +from collections.abc import AsyncGenerator, AsyncIterable, Callable, Iterable, Iterator, Mapping from contextlib import aclosing from http.cookiejar import CookieJar, DefaultCookiePolicy from io import BytesIO @@ -65,11 +65,11 @@ except Exception: try: from brotli import Decompressor as _BrotliInflater from brotli import error as _BrotliError - - _BROTLI_ERRORS: Final[tuple[type[Exception], ...]] = (_BrotliError,) except ImportError: _BrotliInflater = None - _BROTLI_ERRORS = () + + class _BrotliError(Exception): + """Stand-in so the except clause below type-checks; unreachable without brotli.""" # aiohttp 3.10+ exposes a `socket_factory` kwarg on TCPConnector. Older @@ -654,11 +654,17 @@ def _already_read_within(response: httpx.Response, max_bytes: int) -> httpx.Resp class _BoundedDecoder(Protocol): def decode(self, data: bytes, max_output: int) -> bytes: ... + def finish(self) -> None: + """Raise httpx.DecodingError when the wire ended before the compressed stream did.""" + class _IdentityDecoder: def decode(self, data: bytes, max_output: int) -> bytes: return data + def finish(self) -> None: + return None + class _ZlibDecoder: def __init__(self, wbits: int) -> None: @@ -679,35 +685,51 @@ class _ZlibDecoder: return head return head + self.decode(self._inflate.unconsumed_tail, max_output - len(head)) + def finish(self) -> None: + if not self._inflate.eof: + raise httpx.DecodingError("Compressed response ended before the end of the stream") + class _BrotliDecoder: def __init__(self, inflate: "_BrotliInflater") -> None: self._inflate: Final = inflate def decode(self, data: bytes, max_output: int) -> bytes: - try: - head: bytes = self._inflate.process(data, output_buffer_limit=max_output) - except _BROTLI_ERRORS as exc: - raise httpx.DecodingError(str(exc)) from exc - if len(head) >= max_output or self._inflate.can_accept_more_data(): - return head - return head + self.decode(b"", max_output - len(head)) + def drained() -> Iterator[bytes]: + produced = 0 # rebind-ok: running total of the bytes yielded so far + pending = data # rebind-ok: first call feeds the wire bytes, later calls drain buffered output + while produced < max_output: + try: + piece: bytes = self._inflate.process(pending, output_buffer_limit=max_output - produced) + except _BrotliError as exc: + raise httpx.DecodingError(str(exc)) from exc + produced += len(piece) # rebind-ok: see above + pending = b"" # rebind-ok: see above + yield piece + if self._inflate.can_accept_more_data(): + return + + return b"".join(drained()) + + def finish(self) -> None: + if not self._inflate.is_finished(): + raise httpx.DecodingError("Compressed response ended before the end of the stream") def _bounded_decoder(headers: httpx.Headers) -> _BoundedDecoder: - match headers.get("content-encoding", "identity").strip().lower(): - case "identity": - return _IdentityDecoder() - case "gzip" | "x-gzip": - return _ZlibDecoder(zlib.MAX_WBITS | 16) - case "deflate": - return _ZlibDecoder(zlib.MAX_WBITS) - case "br" if _BrotliInflater is not None: - return _BrotliDecoder(_BrotliInflater()) - case _: - raise HTTPResponseLimitError( - "Response size limits require an identity, gzip, deflate, or br encoded response" - ) + factories: Final[Mapping[str, Callable[[], _BoundedDecoder]]] = MappingProxyType( + { + "identity": _IdentityDecoder, + "gzip": lambda: _ZlibDecoder(zlib.MAX_WBITS | 16), + "x-gzip": lambda: _ZlibDecoder(zlib.MAX_WBITS | 16), + "deflate": lambda: _ZlibDecoder(zlib.MAX_WBITS), + **({} if _BrotliInflater is None else {"br": lambda: _BrotliDecoder(_BrotliInflater())}), + } + ) + factory: Final = factories.get(headers.get("content-encoding", "identity").strip().lower()) + if factory is None: + raise HTTPResponseLimitError("Response size limits require an identity, gzip, deflate, or br encoded response") + return factory() async def _decoded_within(response: httpx.Response, wire: AsyncGenerator[bytes, None], max_bytes: int) -> bytes: @@ -717,6 +739,7 @@ async def _decoded_within(response: httpx.Response, wire: AsyncGenerator[bytes, body.write(decoder.decode(chunk, max_bytes + 1 - body.tell())) if body.tell() > max_bytes: raise HTTPResponseLimitError("Response exceeds the configured size limit") + decoder.finish() return body.getvalue() diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index 918d3b6b947..d7c105bec74 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -13,7 +13,6 @@ from collections.abc import Callable, Mapping from typing import Final from unittest.mock import MagicMock, patch -import brotli import certifi import httpx import pytest @@ -1884,15 +1883,15 @@ async def test_bounded_get_bounds_an_already_read_body_on_its_decoded_length(mon @pytest.mark.asyncio @pytest.mark.parametrize( ("encoding", "compress"), - [("gzip", gzip.compress), ("deflate", zlib.compress), ("br", brotli.compress)], + [("gzip", gzip.compress), ("deflate", zlib.compress), ("br", None)], ) async def test_bounded_get_never_inflates_a_compressed_body_past_the_cap( - respx_mock, monkeypatch, encoding: str, compress: Callable[[bytes], bytes] + respx_mock, monkeypatch, encoding: str, compress: Callable[[bytes], bytes] | None ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") inflated: Final = 64 * 1024 * 1024 cap: Final = 1024 * 1024 - bomb: Final = compress(b"\0" * inflated) + bomb: Final = (compress or pytest.importorskip("brotli").compress)(b"\0" * inflated) assert len(bomb) < cap class WireStream(httpx.AsyncByteStream): @@ -1912,6 +1911,29 @@ async def test_bounded_get_never_inflates_a_compressed_body_past_the_cap( assert peak < 4 * cap, f"decoder materialized {peak} bytes for a {cap} byte cap" +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("encoding", "compress"), + [("gzip", gzip.compress), ("deflate", zlib.compress), ("br", None)], +) +async def test_bounded_get_decodes_a_compressed_body_under_the_cap_and_rejects_a_truncated_one( + respx_mock, monkeypatch, encoding: str, compress: Callable[[bytes], bytes] | None +): + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + document: Final = b"line %d of the notes\n" * 2000 % tuple(range(2000)) + wire: Final = (compress or pytest.importorskip("brotli").compress)(document) + for name, body in (("whole", wire), ("cut", wire[: len(wire) // 2])): + respx_mock.get(f"https://cdn.example/{name}").respond(200, content=body, headers={"content-encoding": encoding}) + handler = AsyncHTTPHandler() + try: + served = await handler.get("https://cdn.example/whole", max_response_bytes=len(document)) + assert served.content == document + with pytest.raises(httpx.DecodingError, match="ended before the end of the stream"): + await handler.get("https://cdn.example/cut", max_response_bytes=len(document)) + finally: + await handler.close() + + @pytest.mark.asyncio async def test_bounded_get_refuses_an_encoding_it_cannot_decode_under_the_cap(respx_mock, monkeypatch): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")