[passthrough] engine: zstd request-body decompression + header overrides (#29684)

This commit is contained in:
Jialin Ouyang
2026-06-30 23:48:56 -07:00
committed by GitHub
parent 47ae1241d3
commit 40594bd381
7 changed files with 341 additions and 0 deletions
+1
View File
@@ -85,6 +85,7 @@ dependencies = [
"uvloop",
"watchfiles",
"xgrammar==0.2.1",
"zstandard",
]
[[tool.uv.index]]
@@ -0,0 +1,95 @@
"""Pure-ASGI middleware that decompresses compressed request bodies.
Gated on `SGLANG_ENABLE_REQUEST_DECOMPRESSION` and request header
`x-body-compressed`, whose value names the method. For example, a caller that
compressed the body with zstd sets the `x-body-compressed: zstd` header.
"""
import asyncio
import io
import logging
import zstandard
from fastapi.responses import Response
from starlette.datastructures import Headers
logger = logging.getLogger(__name__)
def _zstd_decompress(raw: bytes) -> bytes:
return zstandard.ZstdDecompressor().stream_reader(io.BytesIO(raw)).read()
_DECOMPRESSORS = {"zstd": _zstd_decompress}
def _rewrite_headers(headers, new_len):
"""Update headers to reflect body status after decompression."""
out = [
(k, v)
for (k, v) in headers
if k not in (b"content-length", b"x-body-compressed")
]
out.append((b"content-length", str(new_len).encode()))
return out
class RequestDecompressionMiddleware:
"""Decompress request body per request header `x-body-compressed`."""
def __init__(self, app):
self.app = app
async def __call__(self, scope, receive, send):
# No-op passthrough for any request without the compression header.
if scope["type"] != "http":
return await self.app(scope, receive, send)
method = Headers(scope=scope).get("x-body-compressed")
if method is None:
return await self.app(scope, receive, send)
# Fail loud on an unsupported compression method.
decompress = _DECOMPRESSORS.get(method)
if decompress is None:
return await Response(
f"unsupported x-body-compressed {method!r}; "
f"supported: {sorted(_DECOMPRESSORS)}",
status_code=400,
)(scope, receive, send)
# Collect request body.
body = b""
more_body = True
while more_body:
message = await receive()
# Incomplete body (e.g. client disconnect); hand off to later stages.
if message["type"] != "http.request":
return await self.app(scope, receive, send)
body += message.get("body", b"")
more_body = message.get("more_body", False)
# Decompress off the event loop by releasing the GIL around the C decompress.
try:
loop = asyncio.get_running_loop()
body = await loop.run_in_executor(None, decompress, body)
except Exception as e:
logger.warning("request body decompress failed: %s", e)
return await Response("decompress failed", status_code=400)(
scope, receive, send
)
# Update the headers after decompression
scope = dict(scope)
scope["headers"] = _rewrite_headers(scope["headers"], len(body))
# Fake receiver to let later stages see the decompressed body.
body_sent = False
async def wrapped_receive():
nonlocal body_sent
if not body_sent:
body_sent = True
return {"type": "http.request", "body": body, "more_body": False}
return await receive()
await self.app(scope, wrapped_receive, send)
@@ -104,6 +104,7 @@ from sglang.srt.entrypoints.openai.serving_tokenize import (
from sglang.srt.entrypoints.openai.serving_transcription import (
OpenAIServingTranscription,
)
from sglang.srt.entrypoints.request_headers import apply_header_overrides
from sglang.srt.entrypoints.warmup import execute_warmups
from sglang.srt.environ import envs
from sglang.srt.function_call.function_call_parser import FunctionCallParser
@@ -403,6 +404,13 @@ app.add_middleware(
allow_headers=["*"],
)
if envs.SGLANG_ENABLE_REQUEST_DECOMPRESSION.get():
from sglang.srt.entrypoints.http_request_decompression import (
RequestDecompressionMiddleware,
)
app.add_middleware(RequestDecompressionMiddleware)
# Include routers
from sglang.srt.entrypoints.v1_loads import router as v1_loads_router
@@ -781,6 +789,8 @@ if os.environ.get("DUMPER_SERVER_PORT") == "reuse":
)
async def generate_request(obj: GenerateReqInput, request: Request):
"""Handle a generate request."""
if envs.SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES.get():
apply_header_overrides(obj, request.headers)
if obj.stream:
async def stream_results() -> AsyncIterator[bytes]:
@@ -0,0 +1,32 @@
"""Override object fields based on _HEADER_OVERRIDES from header values.
This mechanism allows upstream callers to leave the body opaque
(no parse/merge/re-serialize).
"""
from fastapi import HTTPException
# request header -> (target attribute, value type)
_HEADER_OVERRIDES = {
"x-override-rid": ("rid", str),
"x-override-bootstrap-host": ("bootstrap_host", str),
"x-override-bootstrap-port": ("bootstrap_port", int),
"x-override-bootstrap-room": ("bootstrap_room", int),
"x-override-conversation-id": ("conversation_id", str),
"x-override-routed-dp-rank": ("routed_dp_rank", int),
"x-override-disagg-prefill-dp-rank": ("disagg_prefill_dp_rank", int),
}
def apply_header_overrides(obj, headers) -> None:
"""Override request based on header values. Fail the request when any override has issues."""
for header, (attr, cast) in _HEADER_OVERRIDES.items():
value = headers.get(header)
if value is None:
continue
try:
setattr(obj, attr, cast(value))
except ValueError as e:
raise HTTPException(
status_code=400, detail=f"invalid {header} header {value!r}: {e}"
) from e
+6
View File
@@ -232,6 +232,12 @@ class Envs:
SGLANG_PREFETCH_BLOCK_SIZE_MB = EnvInt(16)
SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION = EnvBool(False)
# HTTP server
# Decompress request bodies tagged with `x-body-compressed`.
SGLANG_ENABLE_REQUEST_DECOMPRESSION = EnvBool(False)
# Override parsed request fields from headers.
SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES = EnvBool(False)
# Logging Options
SGLANG_LOG_GC = EnvBool(False)
SGLANG_LOG_FORWARD_ITERS = EnvBool(False)