mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
feat(proxy): gzip buffered responses for clients that accept it (#44052)
* feat(proxy): gzip buffered responses for clients that accept it Large JSON reads like /user/daily/activity/aggregated shipped tens of MB uncompressed. Compress single-message bodies of 500B or more when the client's Accept-Encoding allows gzip (q-values and the wildcard honored). Streamed and etagged responses pass through untouched, every negotiable response carries Vary: Accept-Encoding, and bodies of 1MB or more are compressed in a worker thread * fix(proxy): skip partial and no-transform responses in gzip and always release the held start The gzip gate now also skips 206 Partial Content and Cache-Control: no-transform, since compressing either breaks byte ranges or ignores an explicit ban on transforms. A response start without a headers key no longer raises, and a start the app never follows with a body message is forwarded when the app returns instead of being dropped. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
parent
008fcb4fe3
commit
eb103334ee
3 changed files with 311 additions and 0 deletions
96
litellm/proxy/middleware/gzip_middleware.py
Normal file
96
litellm/proxy/middleware/gzip_middleware.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
import gzip
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import anyio.to_thread
|
||||
from starlette.datastructures import Headers, MutableHeaders
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
MINIMUM_SIZE_BYTES: Final = 500
|
||||
OFF_LOOP_SIZE_BYTES: Final = 1024 * 1024
|
||||
COMPRESS_LEVEL: Final = 6
|
||||
|
||||
|
||||
def _coding_weight(part: str) -> tuple[str, float]:
|
||||
coding, _, params = part.partition(";")
|
||||
qvalue: Final = next((p.strip()[2:] for p in params.split(";") if p.strip().lower().startswith("q=")), "1")
|
||||
try:
|
||||
return coding.strip().lower(), float(qvalue)
|
||||
except ValueError:
|
||||
return coding.strip().lower(), 0.0
|
||||
|
||||
|
||||
def accepts_gzip(accept_encoding: str) -> bool:
|
||||
weights: Final = MappingProxyType(dict(_coding_weight(part) for part in accept_encoding.split(",") if part.strip()))
|
||||
return weights.get("gzip", weights.get("x-gzip", weights.get("*", 0.0))) > 0
|
||||
|
||||
|
||||
async def _compress(body: bytes) -> bytes:
|
||||
if len(body) < OFF_LOOP_SIZE_BYTES:
|
||||
return gzip.compress(body, compresslevel=COMPRESS_LEVEL)
|
||||
return await anyio.to_thread.run_sync(gzip.compress, body, COMPRESS_LEVEL)
|
||||
|
||||
|
||||
class _BufferedBodyGzipResponder:
|
||||
"""Holds the response start until the first body message shows the body is complete, so streams are never delayed."""
|
||||
|
||||
def __init__(self, send: Send, gzip_accepted: bool) -> None:
|
||||
self.send = send
|
||||
self.gzip_accepted = gzip_accepted
|
||||
self.held_start: Message | None = None
|
||||
self.decided = False
|
||||
|
||||
async def __call__(self, message: Message) -> None:
|
||||
if self.decided:
|
||||
await self.send(message)
|
||||
return
|
||||
if message["type"] == "http.response.start":
|
||||
self.held_start = message
|
||||
return
|
||||
self.decided = True
|
||||
start: Final = self.held_start
|
||||
if start is None:
|
||||
await self.send(message)
|
||||
return
|
||||
body: Final[bytes] = message.get("body", b"")
|
||||
start.setdefault("headers", ())
|
||||
headers: Final = MutableHeaders(scope=start)
|
||||
negotiable: Final = (
|
||||
message["type"] == "http.response.body"
|
||||
and not message.get("more_body", False)
|
||||
and len(body) >= MINIMUM_SIZE_BYTES
|
||||
and "content-encoding" not in headers
|
||||
and "etag" not in headers
|
||||
and start["status"] != 206
|
||||
and "no-transform" not in headers.get("cache-control", "").lower()
|
||||
)
|
||||
if negotiable:
|
||||
headers.add_vary_header("Accept-Encoding")
|
||||
if not (negotiable and self.gzip_accepted):
|
||||
await self.send(start)
|
||||
await self.send(message)
|
||||
return
|
||||
compressed: Final = await _compress(body)
|
||||
headers["content-encoding"] = "gzip"
|
||||
headers["content-length"] = str(len(compressed))
|
||||
await self.send(start)
|
||||
await self.send({**message, "body": compressed})
|
||||
|
||||
async def release_held_start(self) -> None:
|
||||
if not self.decided and self.held_start is not None:
|
||||
self.decided = True
|
||||
await self.send(self.held_start)
|
||||
|
||||
|
||||
class GZipBufferedResponseMiddleware:
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
gzip_accepted: Final = accepts_gzip(Headers(scope=scope).get("accept-encoding", ""))
|
||||
responder: Final = _BufferedBodyGzipResponder(send, gzip_accepted)
|
||||
await self.app(scope, receive, responder)
|
||||
await responder.release_held_start()
|
||||
|
|
@ -725,6 +725,7 @@ from litellm.proxy.middleware.admission_control_middleware import (
|
|||
admission_control_state,
|
||||
get_admission_control_settings,
|
||||
)
|
||||
from litellm.proxy.middleware.gzip_middleware import GZipBufferedResponseMiddleware
|
||||
from litellm.proxy.middleware.in_flight_requests_middleware import (
|
||||
InFlightRequestsMiddleware,
|
||||
)
|
||||
|
|
@ -2444,6 +2445,7 @@ app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_b
|
|||
app.add_middleware(RedisRequestBatchMiddleware)
|
||||
app.add_middleware(InFlightRequestsMiddleware)
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
app.add_middleware(GZipBufferedResponseMiddleware)
|
||||
|
||||
|
||||
def mount_swagger_ui():
|
||||
|
|
|
|||
213
tests/unit/proxy/middleware/test_gzip_middleware.py
Normal file
213
tests/unit/proxy/middleware/test_gzip_middleware.py
Normal file
|
|
@ -0,0 +1,213 @@
|
|||
import asyncio
|
||||
import gzip
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, Response, StreamingResponse
|
||||
from starlette.routing import Route
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy.middleware.gzip_middleware import (
|
||||
MINIMUM_SIZE_BYTES,
|
||||
OFF_LOOP_SIZE_BYTES,
|
||||
GZipBufferedResponseMiddleware,
|
||||
)
|
||||
|
||||
LARGE_PAYLOAD = {"rows": [{"date": f"2026-09-{day:02d}", "spend": day * 1.5} for day in range(1, 31)] * 20}
|
||||
STREAM_CHUNKS = tuple(json.dumps({"part": part, "pad": "x" * MINIMUM_SIZE_BYTES}).encode() for part in range(3))
|
||||
|
||||
|
||||
async def _large_json(request: Request) -> Response:
|
||||
return JSONResponse(LARGE_PAYLOAD)
|
||||
|
||||
|
||||
async def _small_json(request: Request) -> Response:
|
||||
return JSONResponse({"ok": True})
|
||||
|
||||
|
||||
async def _already_encoded(request: Request) -> Response:
|
||||
return Response(b"x" * (MINIMUM_SIZE_BYTES * 4), headers={"content-encoding": "br"})
|
||||
|
||||
|
||||
async def _with_etag(request: Request) -> Response:
|
||||
return Response(b"y" * (MINIMUM_SIZE_BYTES * 4), headers={"etag": '"v1"'})
|
||||
|
||||
|
||||
async def _partial(request: Request) -> Response:
|
||||
return Response(b"p" * (MINIMUM_SIZE_BYTES * 4), status_code=206, headers={"content-range": "bytes 0-1999/9000"})
|
||||
|
||||
|
||||
async def _no_transform(request: Request) -> Response:
|
||||
return Response(b"n" * (MINIMUM_SIZE_BYTES * 4), headers={"cache-control": "public, no-transform"})
|
||||
|
||||
|
||||
async def _huge(request: Request) -> Response:
|
||||
return Response(b"z" * (OFF_LOOP_SIZE_BYTES * 2), media_type="application/json")
|
||||
|
||||
|
||||
async def _json_stream(request: Request) -> Response:
|
||||
async def chunks():
|
||||
for chunk in STREAM_CHUNKS:
|
||||
yield chunk
|
||||
|
||||
return StreamingResponse(chunks(), media_type="application/json")
|
||||
|
||||
|
||||
APP = Starlette(
|
||||
routes=[
|
||||
Route("/large", _large_json),
|
||||
Route("/small", _small_json),
|
||||
Route("/encoded", _already_encoded),
|
||||
Route("/stream", _json_stream),
|
||||
Route("/etag", _with_etag),
|
||||
Route("/huge", _huge),
|
||||
Route("/partial", _partial),
|
||||
Route("/no-transform", _no_transform),
|
||||
]
|
||||
)
|
||||
APP.add_middleware(GZipBufferedResponseMiddleware)
|
||||
|
||||
|
||||
async def _send_messages(path: str, accept_encoding: str | None, app: ASGIApp = APP) -> tuple[Message, ...]:
|
||||
headers = [(b"accept-encoding", accept_encoding.encode())] if accept_encoding is not None else []
|
||||
scope = {"type": "http", "method": "GET", "path": path, "query_string": b"", "headers": headers}
|
||||
sent: list[Message] = [] # mutable-ok: ASGI send callback collects messages in order
|
||||
requests: Final = iter(({"type": "http.request", "body": b"", "more_body": False},))
|
||||
never_disconnects: Final = asyncio.Event()
|
||||
|
||||
async def receive() -> Message:
|
||||
request: Final = next(requests, None)
|
||||
if request is not None:
|
||||
return request
|
||||
await never_disconnects.wait()
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
sent.append(message)
|
||||
|
||||
await app(scope, receive, send)
|
||||
return tuple(sent)
|
||||
|
||||
|
||||
def _headers(messages: tuple[Message, ...]) -> dict[str, str]:
|
||||
return {k.decode(): v.decode() for k, v in messages[0]["headers"]}
|
||||
|
||||
|
||||
def _body(messages: tuple[Message, ...]) -> bytes:
|
||||
return b"".join(m.get("body", b"") for m in messages[1:])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("accept_encoding", ["gzip, deflate, br", "GZIP", "br;q=1, gzip;q=0.5", "x-gzip", "*"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_large_buffered_json_is_gzipped_and_round_trips(accept_encoding):
|
||||
messages = await _send_messages("/large", accept_encoding)
|
||||
headers = _headers(messages)
|
||||
body = _body(messages)
|
||||
|
||||
assert headers["content-encoding"] == "gzip"
|
||||
assert headers["vary"] == "Accept-Encoding"
|
||||
assert int(headers["content-length"]) == len(body)
|
||||
assert json.loads(gzip.decompress(body)) == LARGE_PAYLOAD
|
||||
assert len(body) < len(json.dumps(LARGE_PAYLOAD))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_body_above_off_loop_threshold_round_trips():
|
||||
messages = await _send_messages("/huge", "gzip")
|
||||
|
||||
assert _headers(messages)["content-encoding"] == "gzip"
|
||||
assert gzip.decompress(_body(messages)) == b"z" * (OFF_LOOP_SIZE_BYTES * 2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "accept_encoding", "expected_vary"),
|
||||
[
|
||||
("/large", None, "Accept-Encoding"),
|
||||
("/large", "gzip;q=0", "Accept-Encoding"),
|
||||
("/small", "gzip", None),
|
||||
("/etag", "gzip", None),
|
||||
("/stream", "gzip", None),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_vary_marks_every_negotiable_variant(path, accept_encoding, expected_vary):
|
||||
messages = await _send_messages(path, accept_encoding)
|
||||
|
||||
assert _headers(messages).get("vary") == expected_vary
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "accept_encoding", "expected_encoding"),
|
||||
[
|
||||
("/large", None, None),
|
||||
("/large", "identity", None),
|
||||
("/large", "gzip;q=0", None),
|
||||
("/large", "br, gzip; q=0.0", None),
|
||||
("/large", "*;q=0", None),
|
||||
("/large", "*, gzip;q=0", None),
|
||||
("/large", "gzip;q=invalid", None),
|
||||
("/small", "gzip", None),
|
||||
("/encoded", "gzip", "br"),
|
||||
("/etag", "gzip", None),
|
||||
("/stream", "gzip", None),
|
||||
("/partial", "gzip", None),
|
||||
("/no-transform", "gzip", None),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_passes_through_unmodified(path, accept_encoding, expected_encoding):
|
||||
with_header = await _send_messages(path, accept_encoding)
|
||||
without_header = await _send_messages(path, None)
|
||||
|
||||
assert _headers(with_header).get("content-encoding") == expected_encoding
|
||||
assert _body(with_header) == _body(without_header)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streamed_chunks_are_forwarded_one_by_one():
|
||||
messages = await _send_messages("/stream", "gzip")
|
||||
chunks = tuple(m["body"] for m in messages[1:] if m.get("body"))
|
||||
|
||||
assert [m["type"] for m in messages].count("http.response.start") == 1
|
||||
assert chunks == STREAM_CHUNKS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_message_without_headers_key_is_still_gzipped():
|
||||
body: Final = b"h" * (MINIMUM_SIZE_BYTES * 4)
|
||||
|
||||
async def headerless_app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
await send({"type": "http.response.start", "status": 200})
|
||||
await send({"type": "http.response.body", "body": body})
|
||||
|
||||
messages = await _send_messages("/", "gzip", GZipBufferedResponseMiddleware(headerless_app))
|
||||
|
||||
assert _headers(messages)["content-encoding"] == "gzip"
|
||||
assert gzip.decompress(_body(messages)) == body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_without_a_body_message_is_still_forwarded():
|
||||
async def start_only_app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
await send({"type": "http.response.start", "status": 204, "headers": [(b"x-done", b"1")]})
|
||||
|
||||
messages = await _send_messages("/", "gzip", GZipBufferedResponseMiddleware(start_only_app))
|
||||
|
||||
assert messages == ({"type": "http.response.start", "status": 204, "headers": [(b"x-done", b"1")]},)
|
||||
|
||||
|
||||
def test_proxy_app_gzips_large_responses_for_clients_that_accept_it():
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
client = TestClient(app)
|
||||
compressed = client.get("/openapi.json", headers={"accept-encoding": "gzip"})
|
||||
identity = client.get("/openapi.json", headers={"accept-encoding": "identity"})
|
||||
|
||||
assert compressed.headers["content-encoding"] == "gzip"
|
||||
assert int(compressed.headers["content-length"]) < int(identity.headers["content-length"])
|
||||
assert compressed.json() == identity.json()
|
||||
Loading…
Add table
Reference in a new issue