[diffusion] feat: gate /health and /health_generate on warmup completion and add liveness endpoint (#33787)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Lennox Fu
2026-08-07 17:59:23 +08:00
committed by GitHub
co-authored by Mick
parent a42683eb62
commit 7af3d000f2
4 changed files with 162 additions and 13 deletions
+18
View File
@@ -218,6 +218,24 @@ sglang serve \
--port 30010
```
### Health endpoints
SGLang Diffusion separates process liveness from inference readiness:
| Endpoint | Success condition | Recommended use |
| --- | --- | --- |
| `GET /liveness` | The HTTP server is accepting requests. It remains `200` during server warmup. | Kubernetes liveness probe |
| `GET /health` | The server is ready for normal inference traffic. It returns `503` while server-based synthetic warmup is running and `200` after it completes. | Startup and readiness probes |
| `GET /health_generate` | Compatibility alias for `/health`. It does not currently issue a generation request in SGLang Diffusion. | Existing integrations only |
`/health` gates only server-based warmup. With `--warmup-mode off` or
`--warmup-mode request`, it returns `200` once the HTTP server starts; those modes
do not promise that compilation or other first-request work has completed. If
server-based warmup fails, the server terminates instead of reporting ready.
Do not use `/health` as a liveness probe: a long server warmup can legitimately
keep it at `503` for several minutes.
### Cloud Storage
SGLang Diffusion can upload generated images and videos to S3-compatible object storage after generation.
@@ -50,6 +50,33 @@ Base the decision on available memory on the selected GPU(s).
- For multi-GPU deployment: the least-free selected GPU is the bottleneck. A busy 80GiB GPU can behave like a much smaller GPU.
- For single-GPU deployment: FSDP shards DiT weights across multiple GPUs. It is not useful for keeping a single-GPU deployment on one GPU; for that case use CPU offload.
## Health Probes
Use `/liveness` to check that the HTTP process is alive and `/health` to check
that the server is ready for inference. During server-based warmup, `/liveness`
returns `200` while `/health` returns `503`. Configure the startup probe with a
failure budget large enough for model loading and compilation:
```yaml
startupProbe:
httpGet:
path: /health
port: 30010
periodSeconds: 10
failureThreshold: 180
readinessProbe:
httpGet:
path: /health
port: 30010
livenessProbe:
httpGet:
path: /liveness
port: 30010
```
See [Health endpoints](/docs/sglang-diffusion/api/cli#health-endpoints) for the
status-code contract and warmup-mode behavior.
## Performance Modes
`--performance-mode` applies safe presets without overriding explicit offload, FSDP, or parallelism flags. `auto` is the default. Use `manual` when you need to keep performance-related server args under explicit user control. `--mode` is a short alias.
@@ -10,7 +10,7 @@ from typing import TYPE_CHECKING
import httpx
import torch
from fastapi import APIRouter, FastAPI, Request
from fastapi import APIRouter, FastAPI, Request, Response
from fastapi.middleware.cors import CORSMiddleware
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@@ -52,6 +52,7 @@ logger = init_logger(__name__)
VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate")
SERVER_WARMUP_BYPASS_PATHS = (
"/liveness",
"/health",
"/health_generate",
"/model_info",
@@ -59,26 +60,26 @@ SERVER_WARMUP_BYPASS_PATHS = (
)
async def _wait_until_http_ready(server_args: ServerArgs) -> None:
async def _wait_until_http_live(server_args: ServerArgs) -> None:
"""for server warmup"""
health_url = f"{server_args.url()}/health"
# Probe the local server directly: a loopback readiness check must never be
liveness_url = f"{server_args.url()}/liveness"
# Probe the local server directly: a loopback liveness check must never be
# routed through an HTTP proxy. trust_env=False also avoids crashing startup
# on a malformed proxy env var, since httpx parses *_PROXY/NO_PROXY when the
# client is constructed (raising httpx.InvalidURL before any request). See #28493.
async with httpx.AsyncClient(trust_env=False) as client:
for _ in range(120):
try:
response = await client.get(health_url, timeout=5.0)
response = await client.get(liveness_url, timeout=5.0)
if response.status_code == 200:
return
except httpx.HTTPError:
pass
await asyncio.sleep(1.0)
raise RuntimeError(f"HTTP server did not become ready at {health_url}")
raise RuntimeError(f"HTTP server did not become live at {liveness_url}")
async def _run_server_warmup_after_http_ready(
async def _run_server_warmup_after_http_live(
server_args: ServerArgs, warmup_done: asyncio.Event
) -> None:
try:
@@ -86,7 +87,7 @@ async def _run_server_warmup_after_http_ready(
warmup_done.set()
return
await _wait_until_http_ready(server_args)
await _wait_until_http_live(server_args)
await run_async_client_warmup(
server_args,
@@ -120,7 +121,7 @@ async def lifespan(app: FastAPI):
warmup_task = None
if server_args.warmup_mode == "server":
warmup_task = asyncio.create_task(
_run_server_warmup_after_http_ready(server_args, warmup_done)
_run_server_warmup_after_http_live(server_args, warmup_done)
)
else:
warmup_done.set()
@@ -143,8 +144,17 @@ async def lifespan(app: FastAPI):
health_router = APIRouter()
@health_router.get("/liveness")
async def liveness():
"""Report that the HTTP server is accepting requests."""
return {"status": "ok"}
@health_router.get("/health")
async def health():
async def health(request: Request):
"""Report readiness for normal inference traffic."""
if not request.app.state.server_warmup_done.is_set():
return Response(status_code=503)
return {"status": "ok"}
@@ -236,9 +246,9 @@ async def model_info_endpoint(request: Request):
@health_router.get("/health_generate")
async def health_generate():
# TODO : health generate endpoint
return {"status": "ok"}
async def health_generate(request: Request):
"""Compatibility readiness endpoint; no generation is issued."""
return await health(request)
@health_router.get("/stats")
@@ -0,0 +1,94 @@
"""Unit tests for diffusion server liveness and readiness endpoints.
`/liveness` reports HTTP availability independently of model warmup.
`/health` and `/health_generate` report readiness for inference traffic.
"""
import asyncio
import unittest
from types import SimpleNamespace
from unittest import mock
from sglang.multimodal_gen.runtime.entrypoints import http_server
from sglang.multimodal_gen.runtime.entrypoints.http_server import (
health,
health_generate,
liveness,
)
def _make_request(warmup_done) -> SimpleNamespace:
state = SimpleNamespace(server_warmup_done=warmup_done)
return SimpleNamespace(app=SimpleNamespace(state=state))
class TestHealthWarmupGate(unittest.IsolatedAsyncioTestCase):
async def test_liveness_returns_200_before_warmup(self):
self.assertEqual(await liveness(), {"status": "ok"})
async def test_health_returns_503_before_warmup(self):
warmup_done = asyncio.Event()
resp = await health(_make_request(warmup_done))
self.assertEqual(resp.status_code, 503)
async def test_health_returns_200_after_warmup(self):
warmup_done = asyncio.Event()
warmup_done.set()
resp = await health(_make_request(warmup_done))
self.assertEqual(resp, {"status": "ok"})
async def test_health_generate_returns_503_before_warmup(self):
warmup_done = asyncio.Event()
resp = await health_generate(_make_request(warmup_done))
self.assertEqual(resp.status_code, 503)
async def test_health_generate_returns_200_after_warmup(self):
warmup_done = asyncio.Event()
warmup_done.set()
resp = await health_generate(_make_request(warmup_done))
self.assertEqual(resp, {"status": "ok"})
class _FakeResponse:
def __init__(self, status_code: int):
self.status_code = status_code
class _FakeAsyncClient:
def __init__(self, status_codes: list[int]):
self._status_codes = iter(status_codes)
self.get_calls = 0
self.urls = []
def __call__(self, *args, **kwargs):
return self
async def __aenter__(self):
return self
async def __aexit__(self, *exc_info):
return False
async def get(self, url, timeout=None):
self.get_calls += 1
self.urls.append(url)
return _FakeResponse(next(self._status_codes))
class TestWaitUntilHttpLive(unittest.IsolatedAsyncioTestCase):
async def test_waits_for_liveness_200(self):
fake_client = _FakeAsyncClient([503, 200])
server_args = SimpleNamespace(url=lambda: "http://127.0.0.1:11000")
with (
mock.patch.object(http_server.httpx, "AsyncClient", fake_client),
mock.patch.object(http_server.asyncio, "sleep", mock.AsyncMock()),
):
await asyncio.wait_for(
http_server._wait_until_http_live(server_args), timeout=5.0
)
self.assertEqual(fake_client.get_calls, 2)
self.assertEqual(fake_client.urls, ["http://127.0.0.1:11000/liveness"] * 2)
if __name__ == "__main__":
unittest.main()