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:
mateo 2026-09-26 23:10:23 +00:00
parent d6d2473668
commit b234de051b
2 changed files with 120 additions and 7 deletions

View file

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

View file

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