mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(http): bound the wire bytes of a compressed response under the size limit
This commit is contained in:
parent
7d423b51a2
commit
1e8672ded1
3 changed files with 71 additions and 13 deletions
|
|
@ -8,7 +8,8 @@ import sys
|
|||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
|
||||
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
|
||||
|
|
@ -554,6 +555,8 @@ class HTTPResponseLimitError(ValueError):
|
|||
|
||||
|
||||
_WIRE_BODY_HEADERS: Final = frozenset({"content-encoding", "content-length"})
|
||||
_ENCODED_OVERHEAD_DIVISOR: Final = 4096
|
||||
_ENCODED_OVERHEAD_BYTES: Final = 64
|
||||
|
||||
|
||||
class MaskedHTTPStatusError(httpx.HTTPStatusError):
|
||||
|
|
@ -614,10 +617,29 @@ def _headers_of_the_decoded_body(headers: httpx.Headers) -> httpx.Headers:
|
|||
)
|
||||
|
||||
|
||||
def _declared_decoded_length(headers: httpx.Headers) -> int:
|
||||
if headers.get("content-encoding", "identity") != "identity":
|
||||
return 0
|
||||
return int(headers.get("content-length", "0"))
|
||||
def _wire_byte_limit(headers: httpx.Headers, max_bytes: int) -> int:
|
||||
if headers.get("content-encoding", "identity").strip().lower() == "identity":
|
||||
return max_bytes
|
||||
return max_bytes + max_bytes // _ENCODED_OVERHEAD_DIVISOR + _ENCODED_OVERHEAD_BYTES
|
||||
|
||||
|
||||
async def _wire_bounded(response: httpx.Response, limit: int) -> AsyncGenerator[bytes, None]:
|
||||
async for chunk in response.aiter_raw():
|
||||
if response.num_bytes_downloaded > limit:
|
||||
raise HTTPResponseLimitError("Response exceeds the configured size limit")
|
||||
yield chunk
|
||||
|
||||
|
||||
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
|
||||
)
|
||||
with BytesIO() as body:
|
||||
async for chunk in decoding.aiter_bytes(chunk_size=65536):
|
||||
if body.tell() + len(chunk) > max_bytes:
|
||||
raise HTTPResponseLimitError("Response exceeds the configured size limit")
|
||||
body.write(chunk)
|
||||
return body.getvalue()
|
||||
|
||||
|
||||
class AsyncHTTPHandler:
|
||||
|
|
@ -795,17 +817,14 @@ class AsyncHTTPHandler:
|
|||
)
|
||||
if response.is_redirect or response.is_error:
|
||||
return httpx.Response(response.status_code, headers=response.headers, request=response.request)
|
||||
if _declared_decoded_length(response.headers) > max_bytes:
|
||||
wire_limit: Final = _wire_byte_limit(response.headers, max_bytes)
|
||||
if int(response.headers.get("content-length", "0")) > wire_limit:
|
||||
raise HTTPResponseLimitError("Response exceeds the configured size limit")
|
||||
with BytesIO() as body:
|
||||
async for chunk in response.aiter_bytes(chunk_size=65536):
|
||||
if body.tell() + len(chunk) > max_bytes:
|
||||
raise HTTPResponseLimitError("Response exceeds the configured size limit")
|
||||
body.write(chunk)
|
||||
async with aclosing(_wire_bounded(response, wire_limit)) as wire:
|
||||
return httpx.Response(
|
||||
response.status_code,
|
||||
headers=_headers_of_the_decoded_body(response.headers),
|
||||
content=body.getvalue(),
|
||||
content=await _decoded_within(response, wire, max_bytes),
|
||||
request=response.request,
|
||||
)
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -1571,7 +1571,11 @@ class TestBoundedOpenAPISpecLoading:
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"headers",
|
||||
[{"content-length": "1000000"}, {"content-length": "1000000", "content-encoding": "identity"}],
|
||||
[
|
||||
{"content-length": "1000000"},
|
||||
{"content-length": "1000000", "content-encoding": "identity"},
|
||||
{"content-length": "1000000", "content-encoding": "gzip"},
|
||||
],
|
||||
)
|
||||
async def test_unsafe_response_headers_reject_before_reading(self, respx_mock, monkeypatch, headers):
|
||||
import httpx
|
||||
|
|
|
|||
|
|
@ -1826,6 +1826,41 @@ async def test_bounded_get_caps_the_decoded_size_of_a_compressed_body(respx_mock
|
|||
await handler.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bounded_get_caps_the_wire_bytes_of_a_compressed_body_that_decodes_to_nothing(respx_mock, monkeypatch):
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
padded = gzip.compress(b"hello") + b"\x00" * 100_000
|
||||
respx_mock.get("https://cdn.example/padded.txt").respond(200, content=padded, headers={"content-encoding": "gzip"})
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
with pytest.raises(HTTPResponseLimitError, match="size limit"):
|
||||
await handler.get("https://cdn.example/padded.txt", 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")
|
||||
consumed = []
|
||||
|
||||
class RecordingStream(httpx.AsyncByteStream):
|
||||
async def __aiter__(self):
|
||||
consumed.append(True)
|
||||
yield b""
|
||||
|
||||
respx_mock.get("https://cdn.example/big.gz").respond(
|
||||
200, headers={"content-encoding": "gzip", "content-length": "1000000"}, stream=RecordingStream()
|
||||
)
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
with pytest.raises(HTTPResponseLimitError, match="size limit"):
|
||||
await handler.get("https://cdn.example/big.gz", max_response_bytes=1024)
|
||||
finally:
|
||||
await handler.close()
|
||||
assert consumed == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bounded_get_closes_stream_on_cancellation(respx_mock, monkeypatch):
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue