mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(http): decode a compressed response through output-bounded decoders so the size cap holds before inflation
httpx.Response.aiter_bytes inflates each raw chunk in full before the bounded reader can measure it, so a small gzip, deflate, or brotli body could allocate gigabytes under a 1 MiB cap. The capped path now feeds the wire bytes through zlib.decompressobj(max_length) or brotli.Decompressor.process(output_buffer_limit) and refuses encodings it cannot bound (zstd, or br without the brotli package) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
d6d2473668
commit
b234de051b
2 changed files with 120 additions and 7 deletions
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue