[gRPC] Expose native pause status (#37488)
Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com>
This commit is contained in:
co-authored by
ishandhanani
parent
35b7589e1a
commit
f45aad44bd
@@ -11,6 +11,7 @@ service SglangService {
|
||||
rpc Tokenize(TokenizeRequest) returns (TokenizeResponse);
|
||||
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
|
||||
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
||||
rpc GetIsReady(GetIsReadyRequest) returns (GetIsReadyResponse);
|
||||
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
||||
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
||||
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
|
||||
@@ -174,6 +175,17 @@ message HealthCheckResponse {
|
||||
bool healthy = 1;
|
||||
}
|
||||
|
||||
// ---- Readiness ----
|
||||
|
||||
message GetIsReadyRequest {}
|
||||
|
||||
message GetIsReadyResponse {
|
||||
// True when the server is ready to receive new requests.
|
||||
bool is_ready = 1;
|
||||
// Values are JSON-encoded, matching the existing meta_info convention.
|
||||
map<string, string> metadata = 2;
|
||||
}
|
||||
|
||||
// ---- Model info ----
|
||||
|
||||
message GetModelInfoRequest {}
|
||||
|
||||
@@ -446,6 +446,9 @@ class RuntimeHandle:
|
||||
ServerStatus.UnHealthy,
|
||||
)
|
||||
|
||||
def get_is_ready(self) -> bool:
|
||||
return self.tokenizer_manager.is_ready()
|
||||
|
||||
def tokenize(self, text: str, add_special_tokens: bool = True) -> str:
|
||||
tokenizer = self.tokenizer_manager.tokenizer
|
||||
tokens = tokenizer.encode(text, add_special_tokens=add_special_tokens)
|
||||
|
||||
@@ -659,6 +659,13 @@ async def validate_json_request(raw_request: Request):
|
||||
##### Native API endpoints #####
|
||||
|
||||
|
||||
@app.get("/ready")
|
||||
async def ready() -> Response:
|
||||
"""Report whether the server is ready to receive new requests."""
|
||||
status_code = 200 if _global_state.tokenizer_manager.is_ready() else 503
|
||||
return Response(status_code=status_code)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
@app.get("/health_generate")
|
||||
async def health_generate(request: Request) -> Response:
|
||||
|
||||
@@ -626,6 +626,14 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Subprocess liveness watchdog — set by Engine or http_server after construction
|
||||
self._subprocess_watchdog = None
|
||||
|
||||
def is_ready(self) -> bool:
|
||||
"""Return whether this server should receive new requests."""
|
||||
return (
|
||||
not self.is_pause
|
||||
and not self.gracefully_exit
|
||||
and self.server_status == ServerStatus.Up
|
||||
)
|
||||
|
||||
def init_request_logging_and_dumping(self):
|
||||
# TODO: Refactor and organize the log export code.
|
||||
# Request logging
|
||||
|
||||
@@ -90,14 +90,19 @@ def decide_request_auth(
|
||||
it must be rejected (403) even if api_key is provided.
|
||||
|
||||
NOTE :
|
||||
- Health/metrics endpoints are always allowed (even when api_key/admin_api_key is set),
|
||||
to support k8s/liveness/readiness and Prometheus scraping without embedding secrets.
|
||||
- Health/readiness/metrics endpoints are always allowed (even when
|
||||
api_key/admin_api_key is set), to support k8s probes and Prometheus
|
||||
scraping without embedding secrets.
|
||||
- We match them by prefix to cover common variants like /health_generate.
|
||||
"""
|
||||
if method == "OPTIONS":
|
||||
return AuthDecision(allowed=True)
|
||||
|
||||
if path.startswith("/health") or path.startswith("/metrics"):
|
||||
if (
|
||||
path.startswith("/health")
|
||||
or path.startswith("/ready")
|
||||
or path.startswith("/metrics")
|
||||
):
|
||||
return AuthDecision(allowed=True)
|
||||
|
||||
def _check_bearer_token(
|
||||
|
||||
@@ -288,6 +288,13 @@ impl PyBridge {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn get_is_ready(&self) -> PyResult<bool> {
|
||||
Python::attach(|py| {
|
||||
let result = self.runtime_handle.call_method0(py, "get_is_ready")?;
|
||||
result.extract::<bool>(py)
|
||||
})
|
||||
}
|
||||
|
||||
/// Tokenize via Python (fallback when Rust tokenizer unavailable).
|
||||
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
|
||||
Python::attach(|py| {
|
||||
|
||||
@@ -494,14 +494,12 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
}
|
||||
|
||||
// Fallback to Python
|
||||
let json_str = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
let text = req.text.clone();
|
||||
move || bridge.tokenize_py(&text, add_special)
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Tokenize failed"))?;
|
||||
let text = req.text.clone();
|
||||
let json_str = self
|
||||
.blocking_bridge_call("Tokenize failed", move |bridge| {
|
||||
bridge.tokenize_py(&text, add_special)
|
||||
})
|
||||
.await?;
|
||||
|
||||
let v: serde_json::Value = serde_json::from_str(&json_str)
|
||||
.map_err(|e| Status::internal(format!("Failed to parse JSON response: {}", e)))?;
|
||||
@@ -539,14 +537,11 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
}
|
||||
|
||||
// Fallback to Python
|
||||
let json_str = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
let tokens = req.tokens;
|
||||
move || bridge.detokenize_py(tokens)
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Detokenize failed"))?;
|
||||
let json_str = self
|
||||
.blocking_bridge_call("Detokenize failed", move |bridge| {
|
||||
bridge.detokenize_py(req.tokens)
|
||||
})
|
||||
.await?;
|
||||
|
||||
let v: serde_json::Value = serde_json::from_str(&json_str)
|
||||
.map_err(|e| Status::internal(format!("Failed to parse JSON response: {}", e)))?;
|
||||
@@ -561,28 +556,34 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
&self,
|
||||
_request: Request<proto::HealthCheckRequest>,
|
||||
) -> Result<Response<proto::HealthCheckResponse>, Status> {
|
||||
let healthy = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
move || bridge.health_check()
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Health check failed"))?;
|
||||
let healthy = self
|
||||
.blocking_bridge_call("Health check failed", PyBridge::health_check)
|
||||
.await?;
|
||||
|
||||
Ok(Response::new(proto::HealthCheckResponse { healthy }))
|
||||
}
|
||||
|
||||
async fn get_is_ready(
|
||||
&self,
|
||||
_request: Request<proto::GetIsReadyRequest>,
|
||||
) -> Result<Response<proto::GetIsReadyResponse>, Status> {
|
||||
let is_ready = self
|
||||
.blocking_bridge_call("Failed to get readiness", PyBridge::get_is_ready)
|
||||
.await?;
|
||||
|
||||
Ok(Response::new(proto::GetIsReadyResponse {
|
||||
is_ready,
|
||||
metadata: HashMap::new(),
|
||||
}))
|
||||
}
|
||||
|
||||
async fn get_model_info(
|
||||
&self,
|
||||
_request: Request<proto::GetModelInfoRequest>,
|
||||
) -> Result<Response<proto::GetModelInfoResponse>, Status> {
|
||||
let json_info = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
move || bridge.get_model_info()
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Failed to get model info"))?;
|
||||
let json_info = self
|
||||
.blocking_bridge_call("Failed to get model info", PyBridge::get_model_info)
|
||||
.await?;
|
||||
|
||||
Ok(Response::new(proto::GetModelInfoResponse {
|
||||
model_path: extract_model_path(&json_info),
|
||||
@@ -594,13 +595,9 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
&self,
|
||||
_request: Request<proto::GetServerInfoRequest>,
|
||||
) -> Result<Response<proto::GetServerInfoResponse>, Status> {
|
||||
let json_info = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
move || bridge.get_server_info()
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Failed to get server info"))?;
|
||||
let json_info = self
|
||||
.blocking_bridge_call("Failed to get server info", PyBridge::get_server_info)
|
||||
.await?;
|
||||
|
||||
Ok(Response::new(proto::GetServerInfoResponse { json_info }))
|
||||
}
|
||||
@@ -609,13 +606,9 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
&self,
|
||||
_request: Request<proto::ListModelsRequest>,
|
||||
) -> Result<Response<proto::ListModelsResponse>, Status> {
|
||||
let json_str = tokio::task::spawn_blocking({
|
||||
let bridge = self.bridge.clone();
|
||||
move || bridge.list_models()
|
||||
})
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, "Failed to list models"))?;
|
||||
let json_str = self
|
||||
.blocking_bridge_call("Failed to list models", PyBridge::list_models)
|
||||
.await?;
|
||||
|
||||
let models_arr: Vec<serde_json::Value> = serde_json::from_str(&json_str)
|
||||
.map_err(|e| Status::internal(format!("Failed to parse models JSON: {}", e)))?;
|
||||
@@ -847,8 +840,20 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
}
|
||||
}
|
||||
|
||||
// Helper methods for OpenAI pass-through RPCs.
|
||||
// Shared RPC helpers.
|
||||
impl SglangServiceImpl {
|
||||
async fn blocking_bridge_call<T, F>(&self, context: &str, call: F) -> Result<T, Status>
|
||||
where
|
||||
T: Send + 'static,
|
||||
F: FnOnce(&PyBridge) -> Result<T, PyErr> + Send + 'static,
|
||||
{
|
||||
let bridge = self.bridge.clone();
|
||||
tokio::task::spawn_blocking(move || call(&bridge))
|
||||
.await
|
||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
||||
.map_err(|e| pyerr_to_status(e, context))
|
||||
}
|
||||
|
||||
async fn openai_streaming_rpc(
|
||||
&self,
|
||||
request: Request<proto::OpenAiRequest>,
|
||||
|
||||
@@ -43,6 +43,18 @@ def _make_runtime_handle(responses):
|
||||
return handle
|
||||
|
||||
|
||||
class TestNativeGrpcReadiness(CustomTestCase):
|
||||
def test_readiness_comes_from_tokenizer_manager(self):
|
||||
tokenizer_manager = SimpleNamespace(is_ready=lambda: True)
|
||||
handle = RuntimeHandle.__new__(RuntimeHandle)
|
||||
handle.tokenizer_manager = tokenizer_manager
|
||||
|
||||
self.assertTrue(handle.get_is_ready())
|
||||
|
||||
tokenizer_manager.is_ready = lambda: False
|
||||
self.assertFalse(handle.get_is_ready())
|
||||
|
||||
|
||||
class TestNativeGrpcParallelResponses(CustomTestCase):
|
||||
def test_non_streaming_returns_every_choice_before_finishing(self):
|
||||
callback = _RecordingCallback()
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
import asyncio
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.entrypoints import http_server
|
||||
from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestReadyEndpoint(CustomTestCase):
|
||||
def call_ready(
|
||||
self,
|
||||
*,
|
||||
is_pause: bool = False,
|
||||
gracefully_exit: bool = False,
|
||||
server_status: ServerStatus = ServerStatus.Up,
|
||||
):
|
||||
tokenizer_manager = TokenizerManager.__new__(TokenizerManager)
|
||||
tokenizer_manager.is_pause = is_pause
|
||||
tokenizer_manager.gracefully_exit = gracefully_exit
|
||||
tokenizer_manager.server_status = server_status
|
||||
|
||||
prior_state = http_server.get_global_state()
|
||||
http_server.set_global_state(
|
||||
SimpleNamespace(tokenizer_manager=tokenizer_manager)
|
||||
)
|
||||
try:
|
||||
return asyncio.run(http_server.ready())
|
||||
finally:
|
||||
http_server._global_state = prior_state
|
||||
|
||||
def test_ready_while_accepting_requests(self):
|
||||
self.assertEqual(self.call_ready().status_code, 200)
|
||||
|
||||
def test_not_ready_while_paused(self):
|
||||
self.assertEqual(self.call_ready(is_pause=True).status_code, 503)
|
||||
|
||||
def test_not_ready_while_starting(self):
|
||||
self.assertEqual(
|
||||
self.call_ready(server_status=ServerStatus.Starting).status_code, 503
|
||||
)
|
||||
|
||||
def test_not_ready_while_unhealthy(self):
|
||||
self.assertEqual(
|
||||
self.call_ready(server_status=ServerStatus.UnHealthy).status_code, 503
|
||||
)
|
||||
|
||||
def test_not_ready_while_gracefully_exiting(self):
|
||||
self.assertEqual(self.call_ready(gracefully_exit=True).status_code, 503)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
@@ -80,6 +80,17 @@ class TestDecideRequestAuth(CustomTestCase):
|
||||
)
|
||||
self.assertTrue(decision.allowed)
|
||||
|
||||
def test_ready_path_always_allowed(self):
|
||||
decision = decide_request_auth(
|
||||
method="GET",
|
||||
path="/ready",
|
||||
authorization_header=None,
|
||||
api_key="secret",
|
||||
admin_api_key=None,
|
||||
auth_level=AuthLevel.NORMAL,
|
||||
)
|
||||
self.assertTrue(decision.allowed)
|
||||
|
||||
def test_metrics_path_always_allowed(self):
|
||||
decision = decide_request_auth(
|
||||
method="GET",
|
||||
|
||||
@@ -274,6 +274,7 @@ class TestHttpServerAdminAuth(unittest.TestCase):
|
||||
paths_allowed = [
|
||||
"/health",
|
||||
"/health_generate",
|
||||
"/ready",
|
||||
"/metrics",
|
||||
"/metrics/",
|
||||
"/metrics/prometheus",
|
||||
|
||||
Reference in New Issue
Block a user