[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 Tokenize(TokenizeRequest) returns (TokenizeResponse);
|
||||||
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
|
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
|
||||||
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
||||||
|
rpc GetIsReady(GetIsReadyRequest) returns (GetIsReadyResponse);
|
||||||
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
||||||
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
||||||
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
|
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
|
||||||
@@ -174,6 +175,17 @@ message HealthCheckResponse {
|
|||||||
bool healthy = 1;
|
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 ----
|
// ---- Model info ----
|
||||||
|
|
||||||
message GetModelInfoRequest {}
|
message GetModelInfoRequest {}
|
||||||
|
|||||||
@@ -446,6 +446,9 @@ class RuntimeHandle:
|
|||||||
ServerStatus.UnHealthy,
|
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:
|
def tokenize(self, text: str, add_special_tokens: bool = True) -> str:
|
||||||
tokenizer = self.tokenizer_manager.tokenizer
|
tokenizer = self.tokenizer_manager.tokenizer
|
||||||
tokens = tokenizer.encode(text, add_special_tokens=add_special_tokens)
|
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 #####
|
##### 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")
|
||||||
@app.get("/health_generate")
|
@app.get("/health_generate")
|
||||||
async def health_generate(request: Request) -> Response:
|
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
|
# Subprocess liveness watchdog — set by Engine or http_server after construction
|
||||||
self._subprocess_watchdog = None
|
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):
|
def init_request_logging_and_dumping(self):
|
||||||
# TODO: Refactor and organize the log export code.
|
# TODO: Refactor and organize the log export code.
|
||||||
# Request logging
|
# Request logging
|
||||||
|
|||||||
@@ -90,14 +90,19 @@ def decide_request_auth(
|
|||||||
it must be rejected (403) even if api_key is provided.
|
it must be rejected (403) even if api_key is provided.
|
||||||
|
|
||||||
NOTE :
|
NOTE :
|
||||||
- Health/metrics endpoints are always allowed (even when api_key/admin_api_key is set),
|
- Health/readiness/metrics endpoints are always allowed (even when
|
||||||
to support k8s/liveness/readiness and Prometheus scraping without embedding secrets.
|
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.
|
- We match them by prefix to cover common variants like /health_generate.
|
||||||
"""
|
"""
|
||||||
if method == "OPTIONS":
|
if method == "OPTIONS":
|
||||||
return AuthDecision(allowed=True)
|
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)
|
return AuthDecision(allowed=True)
|
||||||
|
|
||||||
def _check_bearer_token(
|
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).
|
/// Tokenize via Python (fallback when Rust tokenizer unavailable).
|
||||||
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
|
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
|
||||||
Python::attach(|py| {
|
Python::attach(|py| {
|
||||||
|
|||||||
@@ -494,14 +494,12 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Fallback to Python
|
// Fallback to Python
|
||||||
let json_str = tokio::task::spawn_blocking({
|
let text = req.text.clone();
|
||||||
let bridge = self.bridge.clone();
|
let json_str = self
|
||||||
let text = req.text.clone();
|
.blocking_bridge_call("Tokenize failed", move |bridge| {
|
||||||
move || bridge.tokenize_py(&text, add_special)
|
bridge.tokenize_py(&text, add_special)
|
||||||
})
|
})
|
||||||
.await
|
.await?;
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Tokenize failed"))?;
|
|
||||||
|
|
||||||
let v: serde_json::Value = serde_json::from_str(&json_str)
|
let v: serde_json::Value = serde_json::from_str(&json_str)
|
||||||
.map_err(|e| Status::internal(format!("Failed to parse JSON response: {}", e)))?;
|
.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
|
// Fallback to Python
|
||||||
let json_str = tokio::task::spawn_blocking({
|
let json_str = self
|
||||||
let bridge = self.bridge.clone();
|
.blocking_bridge_call("Detokenize failed", move |bridge| {
|
||||||
let tokens = req.tokens;
|
bridge.detokenize_py(req.tokens)
|
||||||
move || bridge.detokenize_py(tokens)
|
})
|
||||||
})
|
.await?;
|
||||||
.await
|
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Detokenize failed"))?;
|
|
||||||
|
|
||||||
let v: serde_json::Value = serde_json::from_str(&json_str)
|
let v: serde_json::Value = serde_json::from_str(&json_str)
|
||||||
.map_err(|e| Status::internal(format!("Failed to parse JSON response: {}", e)))?;
|
.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,
|
&self,
|
||||||
_request: Request<proto::HealthCheckRequest>,
|
_request: Request<proto::HealthCheckRequest>,
|
||||||
) -> Result<Response<proto::HealthCheckResponse>, Status> {
|
) -> Result<Response<proto::HealthCheckResponse>, Status> {
|
||||||
let healthy = tokio::task::spawn_blocking({
|
let healthy = self
|
||||||
let bridge = self.bridge.clone();
|
.blocking_bridge_call("Health check failed", PyBridge::health_check)
|
||||||
move || bridge.health_check()
|
.await?;
|
||||||
})
|
|
||||||
.await
|
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Health check failed"))?;
|
|
||||||
|
|
||||||
Ok(Response::new(proto::HealthCheckResponse { healthy }))
|
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(
|
async fn get_model_info(
|
||||||
&self,
|
&self,
|
||||||
_request: Request<proto::GetModelInfoRequest>,
|
_request: Request<proto::GetModelInfoRequest>,
|
||||||
) -> Result<Response<proto::GetModelInfoResponse>, Status> {
|
) -> Result<Response<proto::GetModelInfoResponse>, Status> {
|
||||||
let json_info = tokio::task::spawn_blocking({
|
let json_info = self
|
||||||
let bridge = self.bridge.clone();
|
.blocking_bridge_call("Failed to get model info", PyBridge::get_model_info)
|
||||||
move || bridge.get_model_info()
|
.await?;
|
||||||
})
|
|
||||||
.await
|
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Failed to get model info"))?;
|
|
||||||
|
|
||||||
Ok(Response::new(proto::GetModelInfoResponse {
|
Ok(Response::new(proto::GetModelInfoResponse {
|
||||||
model_path: extract_model_path(&json_info),
|
model_path: extract_model_path(&json_info),
|
||||||
@@ -594,13 +595,9 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
|||||||
&self,
|
&self,
|
||||||
_request: Request<proto::GetServerInfoRequest>,
|
_request: Request<proto::GetServerInfoRequest>,
|
||||||
) -> Result<Response<proto::GetServerInfoResponse>, Status> {
|
) -> Result<Response<proto::GetServerInfoResponse>, Status> {
|
||||||
let json_info = tokio::task::spawn_blocking({
|
let json_info = self
|
||||||
let bridge = self.bridge.clone();
|
.blocking_bridge_call("Failed to get server info", PyBridge::get_server_info)
|
||||||
move || bridge.get_server_info()
|
.await?;
|
||||||
})
|
|
||||||
.await
|
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Failed to get server info"))?;
|
|
||||||
|
|
||||||
Ok(Response::new(proto::GetServerInfoResponse { json_info }))
|
Ok(Response::new(proto::GetServerInfoResponse { json_info }))
|
||||||
}
|
}
|
||||||
@@ -609,13 +606,9 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
|||||||
&self,
|
&self,
|
||||||
_request: Request<proto::ListModelsRequest>,
|
_request: Request<proto::ListModelsRequest>,
|
||||||
) -> Result<Response<proto::ListModelsResponse>, Status> {
|
) -> Result<Response<proto::ListModelsResponse>, Status> {
|
||||||
let json_str = tokio::task::spawn_blocking({
|
let json_str = self
|
||||||
let bridge = self.bridge.clone();
|
.blocking_bridge_call("Failed to list models", PyBridge::list_models)
|
||||||
move || bridge.list_models()
|
.await?;
|
||||||
})
|
|
||||||
.await
|
|
||||||
.map_err(|e| Status::internal(format!("Task join error: {}", e)))?
|
|
||||||
.map_err(|e| pyerr_to_status(e, "Failed to list models"))?;
|
|
||||||
|
|
||||||
let models_arr: Vec<serde_json::Value> = serde_json::from_str(&json_str)
|
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)))?;
|
.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 {
|
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(
|
async fn openai_streaming_rpc(
|
||||||
&self,
|
&self,
|
||||||
request: Request<proto::OpenAiRequest>,
|
request: Request<proto::OpenAiRequest>,
|
||||||
|
|||||||
@@ -43,6 +43,18 @@ def _make_runtime_handle(responses):
|
|||||||
return handle
|
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):
|
class TestNativeGrpcParallelResponses(CustomTestCase):
|
||||||
def test_non_streaming_returns_every_choice_before_finishing(self):
|
def test_non_streaming_returns_every_choice_before_finishing(self):
|
||||||
callback = _RecordingCallback()
|
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)
|
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):
|
def test_metrics_path_always_allowed(self):
|
||||||
decision = decide_request_auth(
|
decision = decide_request_auth(
|
||||||
method="GET",
|
method="GET",
|
||||||
|
|||||||
@@ -274,6 +274,7 @@ class TestHttpServerAdminAuth(unittest.TestCase):
|
|||||||
paths_allowed = [
|
paths_allowed = [
|
||||||
"/health",
|
"/health",
|
||||||
"/health_generate",
|
"/health_generate",
|
||||||
|
"/ready",
|
||||||
"/metrics",
|
"/metrics",
|
||||||
"/metrics/",
|
"/metrics/",
|
||||||
"/metrics/prometheus",
|
"/metrics/prometheus",
|
||||||
|
|||||||
Reference in New Issue
Block a user