[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
+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(())
}