[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 --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 ### Cloud Storage
SGLang Diffusion can upload generated images and videos to S3-compatible object storage after generation. 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 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. - 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 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. `--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 httpx
import torch import torch
from fastapi import APIRouter, FastAPI, Request from fastapi import APIRouter, FastAPI, Request, Response
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams 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") VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate")
SERVER_WARMUP_BYPASS_PATHS = ( SERVER_WARMUP_BYPASS_PATHS = (
"/liveness",
"/health", "/health",
"/health_generate", "/health_generate",
"/model_info", "/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""" """for server warmup"""
health_url = f"{server_args.url()}/health" liveness_url = f"{server_args.url()}/liveness"
# Probe the local server directly: a loopback readiness check must never be # Probe the local server directly: a loopback liveness check must never be
# routed through an HTTP proxy. trust_env=False also avoids crashing startup # 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 # 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. # client is constructed (raising httpx.InvalidURL before any request). See #28493.
async with httpx.AsyncClient(trust_env=False) as client: async with httpx.AsyncClient(trust_env=False) as client:
for _ in range(120): for _ in range(120):
try: try:
response = await client.get(health_url, timeout=5.0) response = await client.get(liveness_url, timeout=5.0)
if response.status_code == 200: if response.status_code == 200:
return return
except httpx.HTTPError: except httpx.HTTPError:
pass pass
await asyncio.sleep(1.0) 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 server_args: ServerArgs, warmup_done: asyncio.Event
) -> None: ) -> None:
try: try:
@@ -86,7 +87,7 @@ async def _run_server_warmup_after_http_ready(
warmup_done.set() warmup_done.set()
return return
await _wait_until_http_ready(server_args) await _wait_until_http_live(server_args)
await run_async_client_warmup( await run_async_client_warmup(
server_args, server_args,
@@ -120,7 +121,7 @@ async def lifespan(app: FastAPI):
warmup_task = None warmup_task = None
if server_args.warmup_mode == "server": if server_args.warmup_mode == "server":
warmup_task = asyncio.create_task( 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: else:
warmup_done.set() warmup_done.set()
@@ -143,8 +144,17 @@ async def lifespan(app: FastAPI):
health_router = APIRouter() 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") @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"} return {"status": "ok"}
@@ -236,9 +246,9 @@ async def model_info_endpoint(request: Request):
@health_router.get("/health_generate") @health_router.get("/health_generate")
async def health_generate(): async def health_generate(request: Request):
# TODO : health generate endpoint """Compatibility readiness endpoint; no generation is issued."""
return {"status": "ok"} return await health(request)
@health_router.get("/stats") @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()