mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(http): fail truncated compressed downloads, drain brotli iteratively, and keep brotli optional in tests
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b234de051b
commit
3c63bd87a8
2 changed files with 73 additions and 28 deletions
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue