[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
@@ -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 ----
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user