[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", "uvloop",
"watchfiles", "watchfiles",
"xgrammar==0.2.1", "xgrammar==0.2.1",
"zstandard",
] ]
[[tool.uv.index]] [[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 ( from sglang.srt.entrypoints.openai.serving_transcription import (
OpenAIServingTranscription, OpenAIServingTranscription,
) )
from sglang.srt.entrypoints.request_headers import apply_header_overrides
from sglang.srt.entrypoints.warmup import execute_warmups from sglang.srt.entrypoints.warmup import execute_warmups
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.function_call.function_call_parser import FunctionCallParser from sglang.srt.function_call.function_call_parser import FunctionCallParser
@@ -403,6 +404,13 @@ app.add_middleware(
allow_headers=["*"], allow_headers=["*"],
) )
if envs.SGLANG_ENABLE_REQUEST_DECOMPRESSION.get():
from sglang.srt.entrypoints.http_request_decompression import (
RequestDecompressionMiddleware,
)
app.add_middleware(RequestDecompressionMiddleware)
# Include routers # Include routers
from sglang.srt.entrypoints.v1_loads import router as v1_loads_router 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): async def generate_request(obj: GenerateReqInput, request: Request):
"""Handle a generate request.""" """Handle a generate request."""
if envs.SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES.get():
apply_header_overrides(obj, request.headers)
if obj.stream: if obj.stream:
async def stream_results() -> AsyncIterator[bytes]: 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_PREFETCH_BLOCK_SIZE_MB = EnvInt(16)
SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION = EnvBool(False) 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 # Logging Options
SGLANG_LOG_GC = EnvBool(False) SGLANG_LOG_GC = EnvBool(False)
SGLANG_LOG_FORWARD_ITERS = EnvBool(False) SGLANG_LOG_FORWARD_ITERS = EnvBool(False)
@@ -0,0 +1,116 @@
import asyncio
import unittest
import zstandard
from sglang.srt.entrypoints.http_request_decompression import (
RequestDecompressionMiddleware,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
PAYLOAD = b'{"text":"hello world","n":7}'
COMPRESSED = zstandard.ZstdCompressor().compress(PAYLOAD)
def _drive(scope, body_chunks):
"""Drive the middleware once. Returns (seen, sent): `seen` is what the inner
app received ({scope, body}) or None if the app was never called; `sent` is
the list of ASGI messages the middleware emitted directly."""
seen = {}
sent = []
chunks = list(body_chunks)
async def receive():
if chunks:
chunk, more = chunks.pop(0)
return {"type": "http.request", "body": chunk, "more_body": more}
return {"type": "http.disconnect"}
async def send(message):
sent.append(message)
async def app(inner_scope, inner_receive, inner_send):
body = b""
more = True
while more:
message = await inner_receive()
if message["type"] != "http.request":
break
body += message.get("body", b"")
more = message.get("more_body", False)
seen["scope"] = inner_scope
seen["body"] = body
asyncio.run(RequestDecompressionMiddleware(app)(scope, receive, send))
return (seen or None), sent
class TestRequestDecompressionMiddleware(unittest.TestCase):
def test_passthrough_when_header_absent(self):
scope = {"type": "http", "headers": [(b"content-length", b"4")]}
seen, sent = _drive(scope, [(b"abcd", False)])
self.assertEqual(seen["body"], b"abcd")
self.assertEqual(seen["scope"]["headers"], [(b"content-length", b"4")])
self.assertEqual(sent, [])
def test_decompresses_zstd_body(self):
scope = {
"type": "http",
"headers": [
(b"x-body-compressed", b"zstd"),
(b"content-length", str(len(COMPRESSED)).encode()),
],
}
seen, sent = _drive(scope, [(COMPRESSED, False)])
self.assertEqual(seen["body"], PAYLOAD)
self.assertEqual(sent, [])
def test_strips_header_and_fixes_content_length(self):
scope = {
"type": "http",
"headers": [
(b"x-body-compressed", b"zstd"),
(b"content-length", str(len(COMPRESSED)).encode()),
(b"content-type", b"application/json"),
],
}
seen, _ = _drive(scope, [(COMPRESSED, False)])
self.assertEqual(
seen["scope"]["headers"],
[
(b"content-type", b"application/json"),
(b"content-length", str(len(PAYLOAD)).encode()),
],
)
def test_unsupported_method_returns_400(self):
scope = {"type": "http", "headers": [(b"x-body-compressed", b"gzip")]}
seen, sent = _drive(scope, [(b"abcd", False)])
self.assertIsNone(seen)
self.assertEqual(sent[0]["type"], "http.response.start")
self.assertEqual(sent[0]["status"], 400)
def test_chunked_body_reassembled(self):
half = len(COMPRESSED) // 2
scope = {"type": "http", "headers": [(b"x-body-compressed", b"zstd")]}
seen, _ = _drive(scope, [(COMPRESSED[:half], True), (COMPRESSED[half:], False)])
self.assertEqual(seen["body"], PAYLOAD)
def test_bad_body_returns_400(self):
scope = {"type": "http", "headers": [(b"x-body-compressed", b"zstd")]}
seen, sent = _drive(scope, [(b"not-zstd-data", False)])
self.assertIsNone(seen)
self.assertEqual(sent[0]["type"], "http.response.start")
self.assertEqual(sent[0]["status"], 400)
def test_non_http_scope_passthrough(self):
scope = {"type": "lifespan", "headers": []}
seen, sent = _drive(scope, [])
self.assertEqual(seen["scope"]["type"], "lifespan")
self.assertEqual(sent, [])
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,81 @@
import unittest
from types import SimpleNamespace
from fastapi import HTTPException
from starlette.datastructures import Headers
from sglang.srt.entrypoints.request_headers import apply_header_overrides
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
def _obj():
return SimpleNamespace(
rid=None,
bootstrap_host=None,
bootstrap_port=None,
bootstrap_room=None,
conversation_id=None,
routed_dp_rank=None,
disagg_prefill_dp_rank=None,
)
class TestApplyRoutingHeaders(unittest.TestCase):
def test_sets_all_fields_with_types(self):
obj = _obj()
apply_header_overrides(
obj,
Headers(
{
"x-override-rid": "r1",
"x-override-bootstrap-host": "prefill1",
"x-override-bootstrap-port": "8998",
"x-override-bootstrap-room": "18446744073709551615",
"x-override-conversation-id": "c1",
"x-override-routed-dp-rank": "3",
"x-override-disagg-prefill-dp-rank": "5",
}
),
)
self.assertEqual(obj.rid, "r1")
self.assertEqual(obj.bootstrap_host, "prefill1")
self.assertEqual(obj.bootstrap_port, 8998)
self.assertEqual(obj.bootstrap_room, 18446744073709551615)
self.assertEqual(obj.conversation_id, "c1")
self.assertEqual(obj.routed_dp_rank, 3)
self.assertEqual(obj.disagg_prefill_dp_rank, 5)
def test_absent_headers_leave_obj_unchanged(self):
obj = _obj()
apply_header_overrides(obj, Headers({}))
self.assertIsNone(obj.rid)
self.assertIsNone(obj.bootstrap_host)
self.assertIsNone(obj.routed_dp_rank)
def test_header_overrides_existing_value(self):
obj = _obj()
obj.rid = "from-body"
apply_header_overrides(obj, Headers({"x-override-rid": "from-header"}))
self.assertEqual(obj.rid, "from-header")
def test_partial_headers_set_only_present(self):
obj = _obj()
apply_header_overrides(
obj, Headers({"x-override-rid": "r1", "x-override-routed-dp-rank": "2"})
)
self.assertEqual(obj.rid, "r1")
self.assertEqual(obj.routed_dp_rank, 2)
self.assertIsNone(obj.bootstrap_host)
def test_invalid_int_fails_loud(self):
obj = _obj()
with self.assertRaises(HTTPException):
apply_header_overrides(
obj, Headers({"x-override-bootstrap-port": "not-an-int"})
)
if __name__ == "__main__":
unittest.main()