[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 Tokenize(TokenizeRequest) returns (TokenizeResponse);
|
||||||
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
|
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
|
||||||
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
||||||
rpc GetIsReady(GetIsReadyRequest) returns (GetIsReadyResponse);
|
rpc WatchEngineState(WatchEngineStateRequest) returns (stream EngineStateSnapshot);
|
||||||
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
||||||
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
||||||
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
|
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
|
||||||
@@ -175,15 +175,21 @@ message HealthCheckResponse {
|
|||||||
bool healthy = 1;
|
bool healthy = 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---- Readiness ----
|
// ---- Engine state ----
|
||||||
|
|
||||||
message GetIsReadyRequest {}
|
message WatchEngineStateRequest {}
|
||||||
|
|
||||||
message GetIsReadyResponse {
|
// A complete discovery and lifecycle snapshot. The JSON discovery payloads
|
||||||
// True when the server is ready to receive new requests.
|
// intentionally retain the same wire shape as GetModelInfo/GetServerInfo.
|
||||||
bool is_ready = 1;
|
message EngineStateSnapshot {
|
||||||
// Values are JSON-encoded, matching the existing meta_info convention.
|
// Unix time in nanoseconds, captured once for this engine process.
|
||||||
map<string, string> metadata = 2;
|
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 ----
|
// ---- Model info ----
|
||||||
|
|||||||
@@ -84,6 +84,9 @@ class RuntimeHandle:
|
|||||||
self.tokenizer_manager.auto_create_handle_loop()
|
self.tokenizer_manager.auto_create_handle_loop()
|
||||||
self._event_loop = self.tokenizer_manager.event_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
|
@property
|
||||||
def _tm_loop(self):
|
def _tm_loop(self):
|
||||||
"""Return the TokenizerManager loop used by communicator RPCs."""
|
"""Return the TokenizerManager loop used by communicator RPCs."""
|
||||||
@@ -449,8 +452,9 @@ class RuntimeHandle:
|
|||||||
ServerStatus.UnHealthy,
|
ServerStatus.UnHealthy,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_is_ready(self) -> bool:
|
def is_pause(self) -> bool:
|
||||||
return self.tokenizer_manager.is_ready()
|
"""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:
|
def tokenize(self, text: str, add_special_tokens: bool = True) -> str:
|
||||||
tokenizer = self.tokenizer_manager.tokenizer
|
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,
|
# 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.
|
# which own no server. Class-level to leave the frozen __init__ alone.
|
||||||
_server_stop_hook: Optional[Callable[[], None]] = None
|
_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:
|
def set_server_stop_hook(self, hook: Callable[[], None]) -> None:
|
||||||
self._server_stop_hook = hook
|
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
|
@property
|
||||||
def serving_chat_class(self):
|
def serving_chat_class(self):
|
||||||
"""Return the serving chat class for OpenAI API.
|
"""Return the serving chat class for OpenAI API.
|
||||||
|
|||||||
@@ -77,6 +77,20 @@ pub enum ChunkSendStatus {
|
|||||||
Closed,
|
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> {
|
fn lock_or_recover<'a, T>(mutex: &'a Mutex<T>, name: &'static str) -> MutexGuard<'a, T> {
|
||||||
mutex.lock().unwrap_or_else(|poisoned| {
|
mutex.lock().unwrap_or_else(|poisoned| {
|
||||||
tracing::warn!(mutex = name, "Recovering from poisoned gRPC bridge mutex");
|
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| {
|
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)
|
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).
|
/// Tokenize via Python (fallback when Rust tokenizer unavailable).
|
||||||
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
|
pub fn tokenize_py(&self, text: &str, add_special_tokens: bool) -> PyResult<String> {
|
||||||
Python::attach(|py| {
|
Python::attach(|py| {
|
||||||
|
|||||||
+135
-14
@@ -1,11 +1,12 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
use pyo3::PyErr;
|
use pyo3::PyErr;
|
||||||
use pyo3::Python;
|
use pyo3::Python;
|
||||||
use pyo3::exceptions::{PyTypeError, PyValueError};
|
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::time::{Duration, timeout};
|
||||||
use tokio_stream::Stream;
|
use tokio_stream::Stream;
|
||||||
use tokio_stream::wrappers::TcpListenerStream;
|
use tokio_stream::wrappers::TcpListenerStream;
|
||||||
@@ -21,11 +22,92 @@ use crate::utils::{
|
|||||||
pub struct SglangServiceImpl {
|
pub struct SglangServiceImpl {
|
||||||
pub bridge: Arc<PyBridge>,
|
pub bridge: Arc<PyBridge>,
|
||||||
pub response_timeout: Duration,
|
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>>;
|
type StreamResult<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + Send + 'static>>;
|
||||||
pub const DEFAULT_RESPONSE_TIMEOUT_SECS: u64 = 300;
|
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,
|
/// 64 MiB — leaves headroom for multimodal inputs and OpenAI JSON pass-through bodies,
|
||||||
/// well above tonic's 4 MiB decode default.
|
/// well above tonic's 4 MiB decode default.
|
||||||
pub const DEFAULT_GRPC_MAX_MESSAGE_SIZE: usize = 64 * 1024 * 1024;
|
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 }))
|
Ok(Response::new(proto::HealthCheckResponse { healthy }))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn get_is_ready(
|
type WatchEngineStateStream = StreamResult<proto::EngineStateSnapshot>;
|
||||||
&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?;
|
|
||||||
|
|
||||||
Ok(Response::new(proto::GetIsReadyResponse {
|
async fn watch_engine_state(
|
||||||
is_ready,
|
&self,
|
||||||
metadata: HashMap::new(),
|
_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(
|
async fn get_model_info(
|
||||||
@@ -989,11 +1092,26 @@ pub async fn run_grpc_server(
|
|||||||
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||||
let addr = listener.local_addr()?;
|
let addr = listener.local_addr()?;
|
||||||
let listener = tokio::net::TcpListener::from_std(listener)?;
|
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 {
|
let service = SglangServiceImpl {
|
||||||
bridge,
|
bridge,
|
||||||
response_timeout,
|
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 max_message_size = resolve_max_message_size();
|
||||||
let svc = proto::sglang_service_server::SglangServiceServer::new(service)
|
let svc = proto::sglang_service_server::SglangServiceServer::new(service)
|
||||||
.max_decoding_message_size(max_message_size)
|
.max_decoding_message_size(max_message_size)
|
||||||
@@ -1001,13 +1119,16 @@ pub async fn run_grpc_server(
|
|||||||
|
|
||||||
tracing::info!("gRPC server listening on {}", addr);
|
tracing::info!("gRPC server listening on {}", addr);
|
||||||
|
|
||||||
tonic::transport::Server::builder()
|
let result = tonic::transport::Server::builder()
|
||||||
.add_service(svc)
|
.add_service(svc)
|
||||||
.serve_with_incoming_shutdown(TcpListenerStream::new(listener), async move {
|
.serve_with_incoming_shutdown(TcpListenerStream::new(listener), async move {
|
||||||
shutdown.notified().await;
|
shutdown.notified().await;
|
||||||
|
stream_shutdown_tx.send_replace(true);
|
||||||
tracing::info!("gRPC server shutting down");
|
tracing::info!("gRPC server shutting down");
|
||||||
})
|
})
|
||||||
.await?;
|
.await;
|
||||||
|
monitor.abort();
|
||||||
|
result?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import unittest
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -43,18 +44,6 @@ def _make_runtime_handle(responses):
|
|||||||
return handle
|
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):
|
class TestNativeGrpcParallelResponses(CustomTestCase):
|
||||||
def test_non_streaming_returns_every_choice_before_finishing(self):
|
def test_non_streaming_returns_every_choice_before_finishing(self):
|
||||||
callback = _RecordingCallback()
|
callback = _RecordingCallback()
|
||||||
@@ -123,5 +112,51 @@ class TestNativeGrpcParallelResponses(CustomTestCase):
|
|||||||
self.assertEqual([call[1] for call in callback.calls], [False, False, True])
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main(verbosity=2)
|
unittest.main(verbosity=2)
|
||||||
|
|||||||
Reference in New Issue
Block a user