diff --git a/rust/proto/vllm_grpc.proto b/rust/proto/vllm_grpc.proto index 2509d5071b6..c10da4818e3 100644 --- a/rust/proto/vllm_grpc.proto +++ b/rust/proto/vllm_grpc.proto @@ -14,6 +14,10 @@ service Generate { rpc GenerateStream (GenerateRequest) returns (stream GenerateResponse) {} } +service Control { + rpc Abort (AbortRequest) returns (AbortResponse) {} +} + // ====================================================================================== // Generate Request // ====================================================================================== @@ -201,3 +205,12 @@ message TokenIds { repeated uint32 ids = 1; } +// ====================================================================================== +// Control +// ====================================================================================== + +message AbortRequest { + repeated string request_ids = 1; +} + +message AbortResponse {} diff --git a/rust/src/server/src/grpc/health.rs b/rust/src/server/src/grpc/health.rs index 70bb81da9f9..554dfae675f 100644 --- a/rust/src/server/src/grpc/health.rs +++ b/rust/src/server/src/grpc/health.rs @@ -8,7 +8,7 @@ use tonic_health::ServingStatus; use tonic_health::server::HealthReporter; use tracing::{info, warn}; -use super::GenerateGrpcService; +use super::{ControlGrpcService, GenerateGrpcService}; pub(crate) async fn monitor_health( mut health_reporter: HealthReporter, @@ -16,21 +16,18 @@ pub(crate) async fn monitor_health( shutdown: CancellationToken, ) { let generate_service = GenerateGrpcService::NAME; + let control_service = ControlGrpcService::NAME; let status = ServingStatus::NotServing; let health_event_first = tokio::select! { result = engine_health.wait_for(|healthy| !*healthy) => { match result { Ok(_) => warn!( - generate_service, - overall_service = true, status = ?status, reason = "engine_unhealthy", "marking gRPC health services as not serving" ), Err(error) => warn!( %error, - generate_service, - overall_service = true, status = ?status, reason = "health_channel_closed", "engine health channel closed; marking gRPC health services as not serving" @@ -40,8 +37,6 @@ pub(crate) async fn monitor_health( } _ = shutdown.cancelled() => { info!( - generate_service, - overall_service = true, status = ?status, reason = "server_shutdown", "server shutting down; marking gRPC health services as not serving" @@ -51,20 +46,20 @@ pub(crate) async fn monitor_health( }; health_reporter.set_not_serving::().await; - // Generate is currently the only engine-backed gRPC service, so overall - // server health intentionally mirrors it. + health_reporter.set_not_serving::().await; + // Both gRPC services use the same engine client, so overall server health + // mirrors their shared engine health. health_reporter.set_service_status("", status).await; if health_event_first { shutdown.cancelled().await; info!( - generate_service, - overall_service = true, reason = "server_shutdown", "server shutting down; closing gRPC health watches" ); } health_reporter.clear_service_status(generate_service).await; + health_reporter.clear_service_status(control_service).await; health_reporter.clear_service_status("").await; } diff --git a/rust/src/server/src/grpc/mod.rs b/rust/src/server/src/grpc/mod.rs index 4170e5bb895..f5d2f347bb3 100644 --- a/rust/src/server/src/grpc/mod.rs +++ b/rust/src/server/src/grpc/mod.rs @@ -26,8 +26,10 @@ pub mod pb { } pub(crate) use health::monitor_health; +pub use pb::control_server::ControlServer; pub use pb::generate_server::GenerateServer; +pub(crate) type ControlGrpcService = ControlServer; pub(crate) type GenerateGrpcService = GenerateServer; #[cfg(test)] @@ -44,6 +46,36 @@ impl GenerateServiceImpl { } } +/// gRPC control service backed by the shared application state. +pub struct ControlServiceImpl { + state: Arc, +} + +impl ControlServiceImpl { + pub fn new(state: Arc) -> Self { + Self { state } + } +} + +#[tonic::async_trait] +impl pb::control_server::Control for ControlServiceImpl { + async fn abort( + &self, + request: Request, + ) -> Result, Status> { + let request_ids = request.into_inner().request_ids; + if request_ids.is_empty() { + return Ok(Response::new(pb::AbortResponse {})); + } + self.state + .chat + .abort(&request_ids) + .await + .map_err(|error| Status::internal(error.to_report_string()))?; + Ok(Response::new(pb::AbortResponse {})) + } +} + #[tonic::async_trait] impl pb::generate_server::Generate for GenerateServiceImpl { type GenerateStreamStream = diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 32412dd9ab7..4cd67f4a4b3 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -38,8 +38,9 @@ use vllm_tokenizer::test_utils::TestTokenizer; use zeromq::prelude::{SocketRecv, SocketSend}; use zeromq::{DealerSocket, PushSocket, ZmqMessage}; +use super::pb::control_client::ControlClient; use super::pb::generate_client::GenerateClient; -use super::{GenerateServer, GenerateServiceImpl, pb}; +use super::{ControlServer, ControlServiceImpl, GenerateServer, GenerateServiceImpl, pb}; use crate::listener::{Listener, MaybeTlsListener}; use crate::state::AppState; use crate::tls; @@ -153,10 +154,6 @@ async fn recv_engine_message(dealer: &mut DealerSocket) -> Vec { dealer.recv().await.expect("recv engine message").into_vec() } -fn test_llm(client: EngineCoreClient) -> Llm { - Llm::new(client).with_request_id_randomization(false) -} - #[derive(Clone, Debug)] struct FakeTextBackend; @@ -206,6 +203,7 @@ async fn setup_grpc_service( output_specs: Vec<(Vec, Option)>, ) -> ( GenerateServer, + ControlServer, tokio::sync::watch::Receiver, MockEngineTask, ) { @@ -243,12 +241,13 @@ async fn setup_grpc_service( let engine_health = client.subscribe_health(); let chat = ChatLlm::from_shared_backend( - test_llm(client), + Llm::new(client), Arc::new(FakeTextBackend) as Arc, ); let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat)); ( - GenerateServer::new(GenerateServiceImpl::new(state)), + GenerateServer::new(GenerateServiceImpl::new(state.clone())), + ControlServer::new(ControlServiceImpl::new(state)), engine_health, engine_task, ) @@ -264,9 +263,11 @@ async fn grpc_test_server( tokio::task::JoinHandle<()>, MockEngineTask, ) { - let (svc, engine_health, engine_task) = setup_grpc_service(engine_id, output_specs).await; + let (generate_service, control_service, engine_health, engine_task) = + setup_grpc_service(engine_id, output_specs).await; let (channel, server_task) = start_grpc_test_server( - svc, + generate_service, + control_service, engine_health, tokio_util::sync::CancellationToken::new(), ) @@ -276,11 +277,13 @@ async fn grpc_test_server( async fn start_grpc_test_server( generate_service: GenerateServer, + control_service: ControlServer, engine_health: tokio::sync::watch::Receiver, shutdown: tokio_util::sync::CancellationToken, ) -> (Channel, tokio::task::JoinHandle<()>) { let (health_reporter, health_service) = health_reporter(); health_reporter.set_serving::>().await; + health_reporter.set_serving::>().await; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener"); let addr = listener.local_addr().expect("local addr"); @@ -289,6 +292,7 @@ async fn start_grpc_test_server( let incoming = MaybeTlsListener::plain(Listener::Tcp(listener)); let server = TonicServer::builder() .add_service(health_service) + .add_service(control_service) .add_service(generate_service) .serve_with_incoming_shutdown(incoming, shutdown.clone().cancelled_owned()); let health_monitor = @@ -319,7 +323,8 @@ async fn grpc_tls_test_server( certs: &TestCerts, cert_reqs: i32, ) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) { - let (svc, _engine_health, engine_task) = setup_grpc_service(engine_id, output_specs).await; + let (generate_service, control_service, _engine_health, engine_task) = + setup_grpc_service(engine_id, output_specs).await; let context = tls::build_grpc_server_config(&server_tls(certs, cert_reqs)) .expect("build grpc tls config"); @@ -329,7 +334,8 @@ async fn grpc_tls_test_server( let server_task = tokio::spawn(async move { let incoming = MaybeTlsListener::tls(Listener::Tcp(listener), context); TonicServer::builder() - .add_service(svc) + .add_service(control_service) + .add_service(generate_service) .serve_with_incoming(incoming) .await .expect("grpc tls server"); @@ -409,7 +415,7 @@ async fn grpc_server_with_keepalive( engine_id: impl Into, keepalive: Option, ) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) { - let (svc, _engine_health, engine_task) = + let (generate_service, control_service, _engine_health, engine_task) = setup_grpc_service(engine_id, default_stream_output_specs()).await; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener"); @@ -425,7 +431,8 @@ async fn grpc_server_with_keepalive( let server_task = tokio::spawn(async move { let incoming = MaybeTlsListener::plain(Listener::Tcp(listener)); builder - .add_service(svc) + .add_service(control_service) + .add_service(generate_service) .serve_with_incoming(incoming) .await .expect("grpc server"); @@ -1073,14 +1080,106 @@ async fn grpc_without_keepalive_keeps_unresponsive_connection_open() { server_task.abort(); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn control_abort_resolves_external_id_and_empty_is_noop() { + let (generate_service, control_service, engine_health, engine_task) = + setup_grpc_service(b"engine-grpc-abort-active", vec![(vec![b'h' as u32], None)]).await; + let (channel, server_task) = start_grpc_test_server( + generate_service, + control_service, + engine_health, + tokio_util::sync::CancellationToken::new(), + ) + .await; + let mut generate_client = GenerateClient::new(channel.clone()); + let mut control_client = ControlClient::new(channel); + let request_id = "test-abort-active"; + + let mut stream = generate_client + .generate_stream(pb::GenerateRequest { + request_id: request_id.to_string(), + model: "test-model".to_string(), + prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())), + stopping: Some(pb::StoppingCriteria { + max_new_tokens: 10, + ..Default::default() + }), + ..Default::default() + }) + .await + .expect("start generation") + .into_inner(); + + loop { + let response = tokio::time::timeout(Duration::from_secs(2), stream.message()) + .await + .expect("timed out waiting for active generation output") + .expect("read active generation output") + .expect("generation ended before producing output"); + if let Some(output) = response.outputs { + assert!( + output.finish_info.is_none(), + "generation finished before abort behavior was exercised" + ); + break; + } + } + + control_client + .abort(pb::AbortRequest::default()) + .await + .expect("empty abort should be a no-op"); + assert!( + tokio::time::timeout(Duration::from_millis(100), stream.message()) + .await + .is_err(), + "empty abort unexpectedly ended the active generation" + ); + + control_client + .abort(pb::AbortRequest { + request_ids: vec![ + request_id.to_string(), + request_id.to_string(), + "unknown".to_string(), + ], + }) + .await + .expect("abort active generation"); + + let finish_reason = loop { + let response = tokio::time::timeout(Duration::from_secs(2), stream.message()) + .await + .expect("timed out waiting for aborted generation") + .expect("read aborted generation") + .expect("generation ended without an aborted response"); + if let Some(finish_info) = response.outputs.and_then(|output| output.finish_info) { + break finish_info.finish_reason; + } + }; + assert_eq!(finish_reason, pb::finish_info::FinishReason::Aborted as i32); + + control_client + .abort(pb::AbortRequest { + request_ids: vec![request_id.to_string()], + }) + .await + .expect("repeated abort should be idempotent"); + + engine_task.await.expect("mock engine task"); + server_task.abort(); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() { - let (generate_service, _connected_engine_health, _engine_task) = + let (generate_service, control_service, _connected_engine_health, _engine_task) = setup_grpc_service(b"engine-grpc-health-failure", default_stream_output_specs()).await; let (engine_health_tx, engine_health) = tokio::sync::watch::channel(true); let (channel, server_task) = start_grpc_test_server( generate_service, + control_service, engine_health, tokio_util::sync::CancellationToken::new(), ) @@ -1088,7 +1187,7 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() let mut health_client = HealthClient::new(channel); let mut health_streams = Vec::new(); - for service in ["vllm.Generate", ""] { + for service in ["vllm.Generate", "vllm.Control", ""] { let service_label = if service.is_empty() { "overall" } else { @@ -1143,14 +1242,19 @@ async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn grpc_health_watch_closes_on_graceful_shutdown() { - let (generate_service, engine_health, _engine_task) = setup_grpc_service( + let (generate_service, control_service, engine_health, _engine_task) = setup_grpc_service( b"engine-grpc-health-shutdown", default_stream_output_specs(), ) .await; let shutdown = tokio_util::sync::CancellationToken::new(); - let (channel, server_task) = - start_grpc_test_server(generate_service, engine_health, shutdown.clone()).await; + let (channel, server_task) = start_grpc_test_server( + generate_service, + control_service, + engine_health, + shutdown.clone(), + ) + .await; let mut health_client = HealthClient::new(channel); let mut stream = health_client .watch(HealthCheckRequest { diff --git a/rust/src/server/src/lib.rs b/rust/src/server/src/lib.rs index 220da778af8..1f3468fe46d 100644 --- a/rust/src/server/src/lib.rs +++ b/rust/src/server/src/lib.rs @@ -207,6 +207,9 @@ where let (health_reporter, health_service) = health_reporter(); let engine_health = state.engine_core_client().subscribe_health(); health_reporter.set_serving::().await; + health_reporter.set_serving::().await; + let control_service = + grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone())); let generate_service = grpc::GenerateGrpcService::new(grpc::GenerateServiceImpl::new(state.clone())); let svc = TonicServer::builder() @@ -214,6 +217,7 @@ where .http2_keepalive_timeout(Some(GRPC_KEEPALIVE_TIMEOUT)) .layer(middleware::request_runtime_layer(state.clone())) .add_service(health_service) + .add_service(control_service) .add_service(generate_service); info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server"); Some((grpc_listener, svc, grpc_tls, health_reporter, engine_health))