diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index b75dcc050d1..52cfdeac000 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -8,12 +8,13 @@ import sys import threading import time import weakref +import zlib from collections.abc import AsyncGenerator, AsyncIterable, Callable, Iterable, Mapping from contextlib import aclosing from http.cookiejar import CookieJar, DefaultCookiePolicy from io import BytesIO from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict, TypeVar +from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, Protocol, TypeAlias, TypedDict, TypeVar import certifi import httpx @@ -61,6 +62,15 @@ try: except Exception: version = "0.0.0" +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 = () + # aiohttp 3.10+ exposes a `socket_factory` kwarg on TCPConnector. Older # versions don't — detect once and skip the keep-alive wiring there. @@ -641,15 +651,72 @@ def _already_read_within(response: httpx.Response, max_bytes: int) -> httpx.Resp ) +class _BoundedDecoder(Protocol): + def decode(self, data: bytes, max_output: int) -> bytes: ... + + +class _IdentityDecoder: + def decode(self, data: bytes, max_output: int) -> bytes: + return data + + +class _ZlibDecoder: + def __init__(self, wbits: int) -> None: + self._inflate = zlib.decompressobj(wbits) + self._raw_fallback = wbits == zlib.MAX_WBITS + + def decode(self, data: bytes, max_output: int) -> bytes: + try: + head: Final = self._inflate.decompress(data, max_output) + except zlib.error as exc: + if not self._raw_fallback: + raise httpx.DecodingError(str(exc)) from exc + self._inflate = zlib.decompressobj(-zlib.MAX_WBITS) + self._raw_fallback = False + return self.decode(data, max_output) + self._raw_fallback = False + if len(head) >= max_output or not self._inflate.unconsumed_tail: + return head + return head + self.decode(self._inflate.unconsumed_tail, max_output - len(head)) + + +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 _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" + ) + + async def _decoded_within(response: httpx.Response, wire: AsyncGenerator[bytes, None], max_bytes: int) -> bytes: - decoding: Final = httpx.Response( - response.status_code, headers=response.headers, content=wire, request=response.request - ) + decoder: Final = _bounded_decoder(response.headers) with BytesIO() as body: - async for chunk in decoding.aiter_bytes(chunk_size=65536): - if body.tell() + len(chunk) > max_bytes: + async for chunk in wire: + body.write(decoder.decode(chunk, max_bytes + 1 - body.tell())) + if body.tell() > max_bytes: raise HTTPResponseLimitError("Response exceeds the configured size limit") - body.write(chunk) 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 35e5728db42..918d3b6b947 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -6,11 +6,14 @@ import os import pathlib import ssl import threading +import tracemalloc import weakref +import zlib from collections.abc import Callable, Mapping from typing import Final from unittest.mock import MagicMock, patch +import brotli import certifi import httpx import pytest @@ -1878,6 +1881,49 @@ async def test_bounded_get_bounds_an_already_read_body_on_its_decoded_length(mon await handler.close() +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("encoding", "compress"), + [("gzip", gzip.compress), ("deflate", zlib.compress), ("br", brotli.compress)], +) +async def test_bounded_get_never_inflates_a_compressed_body_past_the_cap( + respx_mock, monkeypatch, encoding: str, compress: Callable[[bytes], bytes] +): + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + inflated: Final = 64 * 1024 * 1024 + cap: Final = 1024 * 1024 + bomb: Final = compress(b"\0" * inflated) + assert len(bomb) < cap + + class WireStream(httpx.AsyncByteStream): + async def __aiter__(self): + yield bomb + + respx_mock.get("https://cdn.example/bomb").respond(200, stream=WireStream(), headers={"content-encoding": encoding}) + handler = AsyncHTTPHandler() + tracemalloc.start() + try: + with pytest.raises(HTTPResponseLimitError, match="size limit"): + await handler.get("https://cdn.example/bomb", max_response_bytes=cap) + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + await handler.close() + assert peak < 4 * cap, f"decoder materialized {peak} bytes for a {cap} byte cap" + + +@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") + respx_mock.get("https://cdn.example/zst").respond(200, content=b"\x28\xb5\x2f\xfd", headers={"content-encoding": "zstd"}) + handler = AsyncHTTPHandler() + try: + with pytest.raises(HTTPResponseLimitError, match="identity, gzip, deflate, or br"): + await handler.get("https://cdn.example/zst", max_response_bytes=1024) + finally: + await handler.close() + + @pytest.mark.asyncio async def test_bounded_get_rejects_a_declared_compressed_length_over_the_wire_limit(respx_mock, monkeypatch): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")