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:
mateo 2026-09-26 23:31:07 +00:00
parent b234de051b
commit 3c63bd87a8
2 changed files with 73 additions and 28 deletions

View file

@ -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()

View file

@ -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")