[passthrough] engine: zstd request-body decompression + header overrides (#29684)
This commit is contained in:
@@ -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
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user