[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:
co-authored by
ishandhanani
parent
a407915c17
commit
0214954f26
@@ -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
@@ -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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user