[gRPC] Stream engine state changes (#39915)

Signed-off-by: William Arnold <7565007+Aphoh@users.noreply.github.com>
Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com>
This commit is contained in:
William Arnold
2026-09-17 19:17:25 -07:00
committed by GitHub
co-authored by ishandhanani
parent a407915c17
commit 0214954f26
6 changed files with 273 additions and 38 deletions
+14 -8
View File
@@ -11,7 +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 WatchEngineState(WatchEngineStateRequest) returns (stream EngineStateSnapshot);
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
@@ -175,15 +175,21 @@ message HealthCheckResponse {
bool healthy = 1;
}
// ---- Readiness ----
// ---- Engine state ----
message GetIsReadyRequest {}
message WatchEngineStateRequest {}
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;
// A complete discovery and lifecycle snapshot. The JSON discovery payloads
// intentionally retain the same wire shape as GetModelInfo/GetServerInfo.
message EngineStateSnapshot {
// Unix time in nanoseconds, captured once for this engine process.
uint64 instance_id = 1;
// Starts at one and increases for each snapshot from this instance.
uint64 revision = 2;
bool healthy = 3;
bool is_pause = 4;
GetModelInfoResponse model_info = 5;
GetServerInfoResponse server_info = 6;
}
// ---- Model info ----
+6 -2
View File
@@ -84,6 +84,9 @@ class RuntimeHandle:
self.tokenizer_manager.auto_create_handle_loop()
self._event_loop = self.tokenizer_manager.event_loop
def set_engine_state_changed_callback(self, callback) -> None:
self.tokenizer_manager.set_engine_state_changed_callback(callback)
@property
def _tm_loop(self):
"""Return the TokenizerManager loop used by communicator RPCs."""
@@ -449,8 +452,9 @@ class RuntimeHandle:
ServerStatus.UnHealthy,
)
def get_is_ready(self) -> bool:
return self.tokenizer_manager.is_ready()
def is_pause(self) -> bool:
"""Return the tokenizer manager's authoritative generation pause state."""
return self.tokenizer_manager.is_pause
def tokenize(self, text: str, add_special_tokens: bool = True) -> str:
tokenizer = self.tokenizer_manager.tokenizer
@@ -416,10 +416,53 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Set by whoever owns the event loop, and left None for Engine and grpc,
# which own no server. Class-level to leave the frozen __init__ alone.
_server_stop_hook: Optional[Callable[[], None]] = None
_engine_state_changed_callback: Optional[Callable[[], None]] = None
def set_server_stop_hook(self, hook: Callable[[], None]) -> None:
self._server_stop_hook = hook
def _notify_engine_state_changed(self) -> None:
callback = self._engine_state_changed_callback
if callback is None:
return
try:
callback()
except Exception:
logger.exception("Engine-state change callback failed")
def _set_engine_state_field(self, name: str, value: Any) -> None:
if value == getattr(self, name, None):
return
setattr(self, name, value)
self._notify_engine_state_changed()
def set_engine_state_changed_callback(self, callback: Callable[[], None]) -> None:
self._engine_state_changed_callback = callback
@property
def server_status(self):
return self._server_status
@server_status.setter
def server_status(self, value) -> None:
self._set_engine_state_field("_server_status", value)
@property
def gracefully_exit(self) -> bool:
return self._gracefully_exit
@gracefully_exit.setter
def gracefully_exit(self, value: bool) -> None:
self._set_engine_state_field("_gracefully_exit", value)
@property
def is_pause(self) -> bool:
return self._is_pause
@is_pause.setter
def is_pause(self, value: bool) -> None:
self._set_engine_state_field("_is_pause", value)
@property
def serving_chat_class(self):
"""Return the serving chat class for OpenAI API.
+28 -2
View File
@@ -77,6 +77,20 @@ pub enum ChunkSendStatus {
Closed,
}
#[pyclass]
struct EngineStateChangedCallback {
sender: Sender<()>,
}
#[pymethods]
impl EngineStateChangedCallback {
fn __call__(&self) {
// A full snapshot is built after the notification is received, so one
// pending signal is enough to represent any number of quick changes.
let _ = self.sender.try_send(());
}
}
fn lock_or_recover<'a, T>(mutex: &'a Mutex<T>, name: &'static str) -> MutexGuard<'a, T> {
mutex.lock().unwrap_or_else(|poisoned| {
tracing::warn!(mutex = name, "Recovering from poisoned gRPC bridge mutex");
@@ -288,13 +302,25 @@ impl PyBridge {
})
}
pub fn get_is_ready(&self) -> PyResult<bool> {
pub fn is_pause(&self) -> PyResult<bool> {
Python::attach(|py| {
let result = self.runtime_handle.call_method0(py, "get_is_ready")?;
let result = self.runtime_handle.call_method0(py, "is_pause")?;
result.extract::<bool>(py)
})
}
pub fn set_engine_state_changed_callback(&self, sender: Sender<()>) -> PyResult<()> {
Python::attach(|py| {
let callback = Py::new(py, EngineStateChangedCallback { sender })?;
self.runtime_handle.call_method1(
py,
"set_engine_state_changed_callback",
(callback,),
)?;
Ok(())
})
}
/// Tokenize via Python (fallback when Rust tokenizer unavailable).
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
Python::attach(|py| {
+135 -14
View File
@@ -1,11 +1,12 @@
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use pyo3::PyErr;
use pyo3::Python;
use pyo3::exceptions::{PyTypeError, PyValueError};
use tokio::sync::{Notify, mpsc::Receiver};
use tokio::sync::{Notify, mpsc::Receiver, watch};
use tokio::time::{Duration, timeout};
use tokio_stream::Stream;
use tokio_stream::wrappers::TcpListenerStream;
@@ -21,11 +22,92 @@ use crate::utils::{
pub struct SglangServiceImpl {
pub bridge: Arc<PyBridge>,
pub response_timeout: Duration,
engine_state: EngineStatePublisher,
stream_shutdown: watch::Receiver<bool>,
}
type StreamResult<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + Send + 'static>>;
pub const DEFAULT_RESPONSE_TIMEOUT_SECS: u64 = 300;
#[derive(Clone)]
struct EngineStatePublisher {
bridge: Arc<PyBridge>,
instance_id: u64,
sender: watch::Sender<proto::EngineStateSnapshot>,
}
impl EngineStatePublisher {
async fn new(bridge: Arc<PyBridge>) -> Result<Self, Status> {
let instance_id = u64::try_from(
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|error| {
Status::internal(format!("system clock before Unix epoch: {error}"))
})?
.as_nanos(),
)
.map_err(|_| Status::internal("engine instance timestamp does not fit in uint64"))?;
let snapshot = build_engine_state_snapshot(bridge.clone(), instance_id, 1).await?;
let (sender, _) = watch::channel(snapshot);
Ok(Self {
bridge,
instance_id,
sender,
})
}
fn subscribe(&self) -> watch::Receiver<proto::EngineStateSnapshot> {
self.sender.subscribe()
}
async fn publish_current(&self) -> Result<(), Status> {
let revision = self.sender.borrow().revision + 1;
let snapshot =
build_engine_state_snapshot(self.bridge.clone(), self.instance_id, revision).await?;
tracing::info!(
instance_id = snapshot.instance_id,
revision = snapshot.revision,
healthy = snapshot.healthy,
is_pause = snapshot.is_pause,
"publishing SGLang engine state"
);
self.sender.send_replace(snapshot);
Ok(())
}
}
async fn build_engine_state_snapshot(
bridge: Arc<PyBridge>,
instance_id: u64,
revision: u64,
) -> Result<proto::EngineStateSnapshot, Status> {
let values = tokio::task::spawn_blocking(move || {
Ok::<_, PyErr>((
bridge.health_check()?,
bridge.is_pause()?,
bridge.get_model_info()?,
bridge.get_server_info()?,
))
})
.await
.map_err(|error| Status::internal(format!("engine snapshot task failed: {error}")))?
.map_err(|error| pyerr_to_status(error, "Failed to build engine state snapshot"))?;
let (healthy, is_pause, model_json, server_json) = values;
Ok(proto::EngineStateSnapshot {
instance_id,
revision,
healthy,
is_pause,
model_info: Some(proto::GetModelInfoResponse {
model_path: extract_model_path(&model_json),
json_info: model_json,
}),
server_info: Some(proto::GetServerInfoResponse {
json_info: server_json,
}),
})
}
/// 64 MiB — leaves headroom for multimodal inputs and OpenAI JSON pass-through bodies,
/// well above tonic's 4 MiB decode default.
pub const DEFAULT_GRPC_MAX_MESSAGE_SIZE: usize = 64 * 1024 * 1024;
@@ -563,18 +645,39 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
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?;
type WatchEngineStateStream = StreamResult<proto::EngineStateSnapshot>;
Ok(Response::new(proto::GetIsReadyResponse {
is_ready,
metadata: HashMap::new(),
}))
async fn watch_engine_state(
&self,
_request: Request<proto::WatchEngineStateRequest>,
) -> Result<Response<Self::WatchEngineStateStream>, Status> {
let mut receiver = self.engine_state.subscribe();
let mut shutdown = self.stream_shutdown.clone();
let stream = async_stream::stream! {
if *shutdown.borrow_and_update() {
return;
}
let initial = receiver.borrow_and_update().clone();
yield Ok(initial);
loop {
tokio::select! {
biased;
result = shutdown.changed() => {
if result.is_err() || *shutdown.borrow_and_update() {
break;
}
}
result = receiver.changed() => {
if result.is_err() {
break;
}
let update = receiver.borrow_and_update().clone();
yield Ok(update);
}
}
}
};
Ok(Response::new(Box::pin(stream)))
}
async fn get_model_info(
@@ -989,11 +1092,26 @@ pub async fn run_grpc_server(
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let addr = listener.local_addr()?;
let listener = tokio::net::TcpListener::from_std(listener)?;
let (state_changed_tx, mut state_changed_rx) = tokio::sync::mpsc::channel(1);
bridge.set_engine_state_changed_callback(state_changed_tx)?;
let engine_state = EngineStatePublisher::new(bridge.clone()).await?;
let (stream_shutdown_tx, stream_shutdown_rx) = watch::channel(false);
let service = SglangServiceImpl {
bridge,
response_timeout,
engine_state: engine_state.clone(),
stream_shutdown: stream_shutdown_rx,
};
let monitor = tokio::spawn(async move {
while state_changed_rx.recv().await.is_some() {
while state_changed_rx.try_recv().is_ok() {}
if let Err(error) = engine_state.publish_current().await {
tracing::warn!(%error, "failed to publish SGLang engine state");
}
}
});
let max_message_size = resolve_max_message_size();
let svc = proto::sglang_service_server::SglangServiceServer::new(service)
.max_decoding_message_size(max_message_size)
@@ -1001,13 +1119,16 @@ pub async fn run_grpc_server(
tracing::info!("gRPC server listening on {}", addr);
tonic::transport::Server::builder()
let result = tonic::transport::Server::builder()
.add_service(svc)
.serve_with_incoming_shutdown(TcpListenerStream::new(listener), async move {
shutdown.notified().await;
stream_shutdown_tx.send_replace(true);
tracing::info!("gRPC server shutting down");
})
.await?;
.await;
monitor.abort();
result?;
Ok(())
}
@@ -4,6 +4,7 @@ import unittest
from types import SimpleNamespace
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
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
@@ -43,18 +44,6 @@ 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()
@@ -123,5 +112,51 @@ class TestNativeGrpcParallelResponses(CustomTestCase):
self.assertEqual([call[1] for call in callback.calls], [False, False, True])
class TestEngineStateNotifications(CustomTestCase):
def setUp(self):
self.manager = TokenizerManager.__new__(TokenizerManager)
self.manager._engine_state_changed_callback = None
self.manager._server_status = ServerStatus.Starting
self.manager._gracefully_exit = False
self.manager._is_pause = False
self.notifications = 0
self.manager.set_engine_state_changed_callback(self._notify)
def _notify(self):
self.notifications += 1
def test_observable_state_notifies_only_on_changes(self):
self.manager.is_pause = False
self.manager.server_status = ServerStatus.Starting
self.manager.gracefully_exit = False
self.assertEqual(self.notifications, 0)
self.manager.is_pause = True
self.manager.server_status = ServerStatus.Up
self.manager.gracefully_exit = True
self.assertEqual(self.notifications, 3)
def test_runtime_handle_registers_callback_with_manager(self):
handle = RuntimeHandle.__new__(RuntimeHandle)
handle.tokenizer_manager = self.manager
callback = object()
handle.set_engine_state_changed_callback(callback)
self.assertIs(self.manager._engine_state_changed_callback, callback)
def test_graceful_exit_notifies_and_changes_computed_health(self):
handle = RuntimeHandle.__new__(RuntimeHandle)
handle.tokenizer_manager = self.manager
self.manager.server_status = ServerStatus.Up
self.notifications = 0
self.assertTrue(handle.health_check())
self.manager.gracefully_exit = True
self.assertEqual(self.notifications, 1)
self.assertFalse(handle.health_check())
if __name__ == "__main__":
unittest.main(verbosity=2)