From 0214954f26195b890bb20f08a24d5441247edc1b Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:17:25 +0900 Subject: [PATCH] [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> --- proto/sglang/runtime/v1/sglang.proto | 22 ++- python/sglang/srt/entrypoints/grpc_bridge.py | 8 +- .../sglang/srt/managers/tokenizer_manager.py | 43 +++++ rust/sglang-grpc/src/bridge.rs | 30 +++- rust/sglang-grpc/src/server.rs | 149 ++++++++++++++++-- .../unit/entrypoints/test_grpc_bridge.py | 59 +++++-- 6 files changed, 273 insertions(+), 38 deletions(-) diff --git a/proto/sglang/runtime/v1/sglang.proto b/proto/sglang/runtime/v1/sglang.proto index 15341c34f..15251ef80 100644 --- a/proto/sglang/runtime/v1/sglang.proto +++ b/proto/sglang/runtime/v1/sglang.proto @@ -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 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 ---- diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index f0247abf2..e8e9c3073 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -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 diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 628d94be0..24da7ed13 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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. diff --git a/rust/sglang-grpc/src/bridge.rs b/rust/sglang-grpc/src/bridge.rs index 69bb86e84..ab40bb76c 100644 --- a/rust/sglang-grpc/src/bridge.rs +++ b/rust/sglang-grpc/src/bridge.rs @@ -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, 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 { + pub fn is_pause(&self) -> PyResult { 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::(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 { Python::attach(|py| { diff --git a/rust/sglang-grpc/src/server.rs b/rust/sglang-grpc/src/server.rs index 7e5db04b1..c1f939e7e 100644 --- a/rust/sglang-grpc/src/server.rs +++ b/rust/sglang-grpc/src/server.rs @@ -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, pub response_timeout: Duration, + engine_state: EngineStatePublisher, + stream_shutdown: watch::Receiver, } type StreamResult = Pin> + Send + 'static>>; pub const DEFAULT_RESPONSE_TIMEOUT_SECS: u64 = 300; +#[derive(Clone)] +struct EngineStatePublisher { + bridge: Arc, + instance_id: u64, + sender: watch::Sender, +} + +impl EngineStatePublisher { + async fn new(bridge: Arc) -> Result { + 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 { + 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, + instance_id: u64, + revision: u64, +) -> Result { + 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, - ) -> Result, Status> { - let is_ready = self - .blocking_bridge_call("Failed to get readiness", PyBridge::get_is_ready) - .await?; + type WatchEngineStateStream = StreamResult; - Ok(Response::new(proto::GetIsReadyResponse { - is_ready, - metadata: HashMap::new(), - })) + async fn watch_engine_state( + &self, + _request: Request, + ) -> Result, 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> { 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(()) } diff --git a/test/registered/unit/entrypoints/test_grpc_bridge.py b/test/registered/unit/entrypoints/test_grpc_bridge.py index 7f2d437e9..b48abe8b9 100644 --- a/test/registered/unit/entrypoints/test_grpc_bridge.py +++ b/test/registered/unit/entrypoints/test_grpc_bridge.py @@ -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)