[Rust Frontend][gRPC] Add server and model discovery (#49491)

Signed-off-by: Connor Carpenter <connorc@nvidia.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
This commit is contained in:
Connor Carpenter
2026-07-27 09:53:27 -07:00
committed by GitHub
co-authored by Nick Hill
parent 2b465b2c42
commit e3c2fc3b3c
18 changed files with 615 additions and 269 deletions
+53
View File
@@ -0,0 +1,53 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
syntax = "proto3";
package vllm;
service Control {
rpc GetServerInfo (GetServerInfoRequest) returns (ServerInfo) {}
rpc GetModelInfo (GetModelInfoRequest) returns (ModelInfo) {}
rpc Abort (AbortRequest) returns (AbortResponse) {}
}
message GetServerInfoRequest {}
message ServerInfo {
string engine_version = 1;
string api_version = 2;
string instance_id = 3;
ParallelismInfo parallelism = 4;
uint32 max_model_len = 5;
uint32 kv_block_size = 6;
uint64 total_kv_blocks = 7;
uint64 max_running_requests = 8;
uint64 max_batched_tokens = 9;
}
message ParallelismInfo {
uint32 tensor_parallel_size = 1;
uint32 pipeline_parallel_size = 2;
uint32 data_parallel_size = 3;
uint32 data_parallel_rank = 4;
uint32 decode_context_parallel_size = 5;
}
message GetModelInfoRequest {}
message ModelInfo {
string model_id = 1;
string served_model_name = 2;
repeated string served_model_aliases = 3;
bool supports_text_input = 20;
bool supports_token_ids_input = 21;
bool supports_multimodal = 23;
string reasoning_parser = 24;
string tool_call_parser = 25;
}
message AbortRequest {
repeated string request_ids = 1;
}
message AbortResponse {}
@@ -7,17 +7,13 @@ package vllm;
import "google/protobuf/struct.proto";
service Generate {
service Inference {
// Generates text given a prompt
rpc Generate (GenerateRequest) returns (GenerateResponse) {}
// Generates text given a prompt, streaming the outputs
rpc GenerateStream (GenerateRequest) returns (stream GenerateResponse) {}
}
service Control {
rpc Abort (AbortRequest) returns (AbortResponse) {}
}
// ======================================================================================
// Generate Request
// ======================================================================================
@@ -204,13 +200,3 @@ message CandidateTokenInfo {
message TokenIds {
repeated uint32 ids = 1;
}
// ======================================================================================
// Control
// ======================================================================================
message AbortRequest {
repeated string request_ids = 1;
}
message AbortResponse {}
+27
View File
@@ -255,6 +255,33 @@ impl ChatLlm {
self.text.engine_core_client()
}
/// Whether the loaded backend has a registered multimodal processor.
pub fn supports_multimodal(&self) -> bool {
self.processor.backend.multimodal_model_info().is_some()
}
/// Effective tool-call parser name for this model, if parsing is enabled.
pub fn tool_call_parser_name(&self) -> Option<&str> {
match &self.tool_call_parser {
ParserSelection::Auto => {
ToolParserFactory::global().resolve_name_for_model(self.model_id())
}
ParserSelection::None => None,
ParserSelection::Explicit(name) => Some(name),
}
}
/// Effective reasoning parser name for this model, if parsing is enabled.
pub fn reasoning_parser_name(&self) -> Option<&str> {
match &self.reasoning_parser {
ParserSelection::Auto => {
ReasoningParserFactory::global().resolve_name_for_model(self.model_id())
}
ParserSelection::None => None,
ParserSelection::Explicit(name) => Some(name),
}
}
/// Render, tokenize, and submit one chat request.
pub async fn chat(&self, request: ChatRequest) -> Result<ChatEventStream> {
let (text_request, output_processor) = self
+11
View File
@@ -394,6 +394,17 @@ impl EngineCoreClient {
self.engines.iter().map(|engine| &engine.ready_response).collect()
}
/// Return the first engine's ready response.
///
/// Per-engine fields such as `data_parallel_rank` should be read through
/// [`ready_responses`](Self::ready_responses).
pub fn ready_response(&self) -> &EngineCoreReadyResponse {
&self
.engines
.first()
.expect("engine core client requires at least one engine")
.ready_response
}
/// Return the engine-reported effective model dtype.
pub fn model_dtype(&self) -> ModelDtype {
self.engines
@@ -57,6 +57,13 @@ pub fn default_ready_response() -> EngineCoreReadyResponse {
vllm_version: "test-vllm-version".to_string(),
world_size: 1,
data_parallel_size: 1,
tensor_parallel_size: 1,
pipeline_parallel_size: 1,
decode_context_parallel_size: 1,
data_parallel_rank: 0,
max_num_seqs: 256,
max_num_batched_tokens: 8192,
instance_id: "test-instance".to_string(),
kv_cache_size_tokens: None,
kv_cache_max_concurrency: None,
}
@@ -52,6 +52,21 @@ pub struct EngineCoreReadyResponse {
pub world_size: u64,
/// Data parallelism size from the parallel config.
pub data_parallel_size: u64,
// Required discovery metadata; EngineCore and client versions must match.
/// Tensor-parallel size of this engine.
pub tensor_parallel_size: u32,
/// Pipeline-parallel size of this engine.
pub pipeline_parallel_size: u32,
/// Decode-context-parallel size of this engine.
pub decode_context_parallel_size: u32,
/// This engine's data-parallel rank.
pub data_parallel_rank: u32,
/// Scheduler cap on concurrently running sequences.
pub max_num_seqs: u64,
/// Scheduler cap on batched tokens per step.
pub max_num_batched_tokens: u64,
/// Unique identifier for this server instance.
pub instance_id: String,
/// Total KV cache capacity in tokens, if reported.
pub kv_cache_size_tokens: Option<u64>,
/// Maximum achievable request concurrency given the KV cache, if reported.
@@ -363,6 +363,13 @@ class EngineCoreReadyResponse:
vllm_version: str
world_size: int
data_parallel_size: int
tensor_parallel_size: int
pipeline_parallel_size: int
decode_context_parallel_size: int
data_parallel_rank: int
max_num_seqs: int
max_num_batched_tokens: int
instance_id: str
kv_cache_size_tokens: int | None = None
kv_cache_max_concurrency: float | None = None
@@ -376,6 +383,13 @@ ready_response = EngineCoreReadyResponse(
vllm_version="0.0.0",
data_parallel_size=1,
world_size=1,
tensor_parallel_size=1,
pipeline_parallel_size=1,
decode_context_parallel_size=1,
data_parallel_rank=0,
max_num_seqs=256,
max_num_batched_tokens=8192,
instance_id="test-instance",
)
print(msgspec.msgpack.encode(request).hex())
+7 -1
View File
@@ -9,7 +9,13 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
.build_server(true)
.build_client(true)
.protoc_arg("--experimental_allow_proto3_optional") // be compatible with old compilers
.compile_protos(&[format!("{proto_dir}/vllm_grpc.proto")], &[proto_dir])?;
.compile_protos(
&[
format!("{proto_dir}/control.proto"),
format!("{proto_dir}/inference.proto"),
],
&[proto_dir],
)?;
Ok(())
}
+106
View File
@@ -0,0 +1,106 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
use std::sync::Arc;
use thiserror_ext::AsReport as _;
use tonic::{Request, Response, Status};
use vllm_engine_core_client::protocol::handshake::EngineCoreReadyResponse;
use super::{ControlServer, pb};
use crate::state::AppState;
pub(crate) type ControlGrpcService = ControlServer<ControlServiceImpl>;
/// gRPC control service backed by the shared application state.
pub struct ControlServiceImpl {
state: Arc<AppState>,
}
impl ControlServiceImpl {
pub fn new(state: Arc<AppState>) -> Self {
Self { state }
}
fn ready(&self) -> &EngineCoreReadyResponse {
self.state.engine_core_client().ready_response()
}
fn parallelism_info(&self) -> pb::ParallelismInfo {
let ready = self.ready();
pb::ParallelismInfo {
tensor_parallel_size: ready.tensor_parallel_size,
pipeline_parallel_size: ready.pipeline_parallel_size,
data_parallel_size: ready.data_parallel_size.min(u64::from(u32::MAX)) as u32,
data_parallel_rank: ready.data_parallel_rank,
decode_context_parallel_size: ready.decode_context_parallel_size,
}
}
}
const GRPC_API_VERSION: &str = "vllm";
#[tonic::async_trait]
impl pb::control_server::Control for ControlServiceImpl {
async fn get_server_info(
&self,
_request: Request<pb::GetServerInfoRequest>,
) -> Result<Response<pb::ServerInfo>, Status> {
let ready = self.ready();
Ok(Response::new(pb::ServerInfo {
engine_version: ready.vllm_version.clone(),
api_version: GRPC_API_VERSION.to_string(),
instance_id: ready.instance_id.clone(),
parallelism: Some(self.parallelism_info()),
max_model_len: self.state.engine_core_client().max_model_len(),
kv_block_size: ready.block_size.min(u64::from(u32::MAX)) as u32,
total_kv_blocks: self.state.engine_core_client().total_num_gpu_blocks(),
max_running_requests: ready.max_num_seqs,
max_batched_tokens: ready.max_num_batched_tokens,
}))
}
async fn get_model_info(
&self,
_request: Request<pb::GetModelInfoRequest>,
) -> Result<Response<pb::ModelInfo>, Status> {
let served = self.state.served_model_names();
Ok(Response::new(pb::ModelInfo {
model_id: self.state.chat.text().model_id().to_string(),
served_model_name: self.state.primary_model_name().to_string(),
served_model_aliases: served.iter().skip(1).cloned().collect(),
// GenerateRequest accepts both prompt representations.
supports_text_input: true,
supports_token_ids_input: true,
supports_multimodal: self.state.chat.supports_multimodal(),
reasoning_parser: self
.state
.chat
.reasoning_parser_name()
.unwrap_or_default()
.to_string(),
tool_call_parser: self
.state
.chat
.tool_call_parser_name()
.unwrap_or_default()
.to_string(),
}))
}
async fn abort(
&self,
request: Request<pb::AbortRequest>,
) -> Result<Response<pb::AbortResponse>, 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 {}))
}
}
+4 -4
View File
@@ -8,14 +8,14 @@ use tonic_health::ServingStatus;
use tonic_health::server::HealthReporter;
use tracing::{info, warn};
use super::{ControlGrpcService, GenerateGrpcService};
use super::{ControlGrpcService, InferenceGrpcService};
pub(crate) async fn monitor_health(
mut health_reporter: HealthReporter,
mut engine_health: watch::Receiver<bool>,
shutdown: CancellationToken,
) {
let generate_service = GenerateGrpcService::NAME;
let inference_service = InferenceGrpcService::NAME;
let control_service = ControlGrpcService::NAME;
let status = ServingStatus::NotServing;
let health_event_first = tokio::select! {
@@ -45,7 +45,7 @@ pub(crate) async fn monitor_health(
}
};
health_reporter.set_not_serving::<GenerateGrpcService>().await;
health_reporter.set_not_serving::<InferenceGrpcService>().await;
health_reporter.set_not_serving::<ControlGrpcService>().await;
// Both gRPC services use the same engine client, so overall server health
// mirrors their shared engine health.
@@ -59,7 +59,7 @@ pub(crate) async fn monitor_health(
);
}
health_reporter.clear_service_status(generate_service).await;
health_reporter.clear_service_status(inference_service).await;
health_reporter.clear_service_status(control_service).await;
health_reporter.clear_service_status("").await;
}
+156
View File
@@ -0,0 +1,156 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
use std::pin::Pin;
use std::sync::Arc;
use futures::{Stream, StreamExt as _};
use thiserror_ext::AsReport as _;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status};
use tracing::info;
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
use super::convert::{self, ResponseOpts};
use super::{InferenceServer, pb};
use crate::state::AppState;
pub(crate) type InferenceGrpcService = InferenceServer<InferenceServiceImpl>;
/// gRPC inference service backed by the shared application state.
pub struct InferenceServiceImpl {
state: Arc<AppState>,
}
impl InferenceServiceImpl {
pub fn new(state: Arc<AppState>) -> Self {
Self { state }
}
}
#[tonic::async_trait]
impl pb::inference_server::Inference for InferenceServiceImpl {
type GenerateStreamStream =
Pin<Box<dyn Stream<Item = Result<pb::GenerateResponse, Status>> + Send>>;
/// Unary generate: collect all output and return a single response.
async fn generate(
&self,
request: Request<pb::GenerateRequest>,
) -> Result<Response<pb::GenerateResponse>, Status> {
let proto_req = request.into_inner();
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
let text_request =
convert::to_text_request(proto_req, false, self.state.served_model_names())?;
let request_id = text_request.request_id.clone();
info!(%request_id, "grpc generate (unary)");
let stream = self.state.chat.text().generate(text_request).await;
let stream = stream.map_err(text_error_to_status)?;
let collected = stream.collect_output().await.map_err(text_error_to_status)?;
// Build the single aggregated response.
let prompt_info = convert::to_prompt_info(
&collected.prompt_token_ids,
collected.prompt_logprobs.as_ref(),
&response_opts,
);
let finish_info = vllm_text::Finished {
usage: collected.usage,
finish_reason: collected.finish_reason,
kv_transfer_params: collected.kv_transfer_params,
ec_transfer_params: collected.ec_transfer_params,
};
let outputs = convert::to_sequence_output(
&collected.text,
&collected.token_ids,
collected.logprobs.as_ref(),
Some(&finish_info),
&response_opts,
);
Ok(Response::new(pb::GenerateResponse {
prompt_info: Some(prompt_info),
outputs: Some(outputs),
}))
}
/// Streaming generate: yield incremental responses as tokens are produced.
async fn generate_stream(
&self,
request: Request<pb::GenerateRequest>,
) -> Result<Response<Self::GenerateStreamStream>, Status> {
let proto_req = request.into_inner();
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
let text_request =
convert::to_text_request(proto_req, true, self.state.served_model_names())?;
let request_id = text_request.request_id.clone();
info!(%request_id, "grpc generate (stream)");
let stream = self.state.chat.text().generate(text_request).await;
let stream = stream.map_err(text_error_to_status)?;
let (tx, rx) = mpsc::channel(32);
tokio::spawn(async move {
futures::pin_mut!(stream);
while let Some(event) = stream.next().await {
let response = match event {
Err(e) => Err(text_error_to_status(e)),
Ok(DecodedTextEvent::Start {
prompt_token_ids,
prompt_logprobs,
}) => {
let prompt_info = convert::to_prompt_info(
&prompt_token_ids,
prompt_logprobs.as_ref(),
&response_opts,
);
Ok(pb::GenerateResponse {
prompt_info: Some(prompt_info),
outputs: None,
})
}
Ok(DecodedTextEvent::TextDelta {
delta,
token_ids,
logprobs,
finished,
}) => Ok(pb::GenerateResponse {
prompt_info: None,
outputs: Some(convert::to_sequence_output(
&delta,
&token_ids,
logprobs.as_ref(),
finished.as_ref(),
&response_opts,
)),
}),
};
if tx.send(response).await.is_err() {
// Client disconnected.
break;
}
}
});
let response_stream = ReceiverStream::new(rx);
Ok(Response::new(Box::pin(response_stream)))
}
}
fn text_error_to_status(error: vllm_text::Error) -> Status {
let message = error.to_report_string();
if error.is_request_validation_error() {
Status::invalid_argument(message)
} else {
Status::internal(message)
}
}
+8 -186
View File
@@ -1,203 +1,25 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
//! gRPC Generate service backed by the shared [`vllm_text::TextLlm`] facade.
//! gRPC services backed by the shared application state.
mod control;
mod convert;
mod health;
use std::pin::Pin;
use std::sync::Arc;
use futures::{Stream, StreamExt as _};
use thiserror_ext::AsReport as _;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tonic::{Request, Response, Status};
use tracing::info;
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
use self::convert::ResponseOpts;
use crate::state::AppState;
mod inference;
/// Generated protobuf/gRPC types for the `vllm` package.
pub mod pb {
tonic::include_proto!("vllm");
}
pub(crate) use control::ControlGrpcService;
pub use control::ControlServiceImpl;
pub(crate) use health::monitor_health;
pub(crate) use inference::InferenceGrpcService;
pub use inference::InferenceServiceImpl;
pub use pb::control_server::ControlServer;
pub use pb::generate_server::GenerateServer;
pub(crate) type ControlGrpcService = ControlServer<ControlServiceImpl>;
pub(crate) type GenerateGrpcService = GenerateServer<GenerateServiceImpl>;
pub use pb::inference_server::InferenceServer;
#[cfg(test)]
mod tests;
/// gRPC Generate service implementation backed by the shared application state.
pub struct GenerateServiceImpl {
state: Arc<AppState>,
}
impl GenerateServiceImpl {
pub fn new(state: Arc<AppState>) -> Self {
Self { state }
}
}
/// gRPC control service backed by the shared application state.
pub struct ControlServiceImpl {
state: Arc<AppState>,
}
impl ControlServiceImpl {
pub fn new(state: Arc<AppState>) -> Self {
Self { state }
}
}
#[tonic::async_trait]
impl pb::control_server::Control for ControlServiceImpl {
async fn abort(
&self,
request: Request<pb::AbortRequest>,
) -> Result<Response<pb::AbortResponse>, 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 =
Pin<Box<dyn Stream<Item = Result<pb::GenerateResponse, Status>> + Send>>;
/// Unary generate: collect all output and return a single response.
async fn generate(
&self,
request: Request<pb::GenerateRequest>,
) -> Result<Response<pb::GenerateResponse>, Status> {
let proto_req = request.into_inner();
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
let text_request =
convert::to_text_request(proto_req, false, self.state.served_model_names())?;
let request_id = text_request.request_id.clone();
info!(%request_id, "grpc generate (unary)");
let stream = self.state.chat.text().generate(text_request).await;
let stream = stream.map_err(text_error_to_status)?;
let collected = stream.collect_output().await.map_err(text_error_to_status)?;
// Build the single aggregated response.
let prompt_info = convert::to_prompt_info(
&collected.prompt_token_ids,
collected.prompt_logprobs.as_ref(),
&response_opts,
);
let finish_info = vllm_text::Finished {
usage: collected.usage,
finish_reason: collected.finish_reason,
kv_transfer_params: collected.kv_transfer_params,
ec_transfer_params: collected.ec_transfer_params,
};
let outputs = convert::to_sequence_output(
&collected.text,
&collected.token_ids,
collected.logprobs.as_ref(),
Some(&finish_info),
&response_opts,
);
Ok(Response::new(pb::GenerateResponse {
prompt_info: Some(prompt_info),
outputs: Some(outputs),
}))
}
/// Streaming generate: yield incremental responses as tokens are produced.
async fn generate_stream(
&self,
request: Request<pb::GenerateRequest>,
) -> Result<Response<Self::GenerateStreamStream>, Status> {
let proto_req = request.into_inner();
let response_opts = ResponseOpts::from_proto(proto_req.response.as_ref());
let text_request =
convert::to_text_request(proto_req, true, self.state.served_model_names())?;
let request_id = text_request.request_id.clone();
info!(%request_id, "grpc generate (stream)");
let stream = self.state.chat.text().generate(text_request).await;
let stream = stream.map_err(text_error_to_status)?;
let (tx, rx) = mpsc::channel(32);
tokio::spawn(async move {
futures::pin_mut!(stream);
while let Some(event) = stream.next().await {
let response = match event {
Err(e) => Err(text_error_to_status(e)),
Ok(DecodedTextEvent::Start {
prompt_token_ids,
prompt_logprobs,
}) => {
let prompt_info = convert::to_prompt_info(
&prompt_token_ids,
prompt_logprobs.as_ref(),
&response_opts,
);
Ok(pb::GenerateResponse {
prompt_info: Some(prompt_info),
outputs: None,
})
}
Ok(DecodedTextEvent::TextDelta {
delta,
token_ids,
logprobs,
finished,
}) => Ok(pb::GenerateResponse {
prompt_info: None,
outputs: Some(convert::to_sequence_output(
&delta,
&token_ids,
logprobs.as_ref(),
finished.as_ref(),
&response_opts,
)),
}),
};
if tx.send(response).await.is_err() {
// Client disconnected.
break;
}
}
});
let response_stream = ReceiverStream::new(rx);
Ok(Response::new(Box::pin(response_stream)))
}
}
fn text_error_to_status(error: vllm_text::Error) -> Status {
let message = error.to_report_string();
if error.is_request_validation_error() {
Status::invalid_argument(message)
} else {
Status::internal(message)
}
}
+157 -38
View File
@@ -25,12 +25,18 @@ use vllm_chat::{
ChatBackend, ChatLlm, ChatRenderer, ChatRequest, ChatTextBackend, DefaultChatOutputProcessor,
DynChatOutputProcessor, DynChatRenderer, NewChatOutputProcessorOptions, RenderedPrompt,
};
use vllm_engine_core_client::mock_engine::{
DEFAULT_MOCK_BLOCK_SIZE, DEFAULT_MOCK_MAX_MODEL_LEN, DEFAULT_MOCK_NUM_GPU_BLOCKS,
default_ready_response,
};
use vllm_engine_core_client::protocol::output::{
EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, RequestBatchOutputs,
};
use vllm_engine_core_client::protocol::request::EngineCoreRequest;
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId};
use vllm_engine_core_client::test_utils::{
IpcNamespace, spawn_mock_engine_task, spawn_mock_engine_task_with_ready,
};
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId, TransportMode};
use vllm_llm::Llm;
use vllm_text::tokenizer::DynTokenizer;
use vllm_text::{Prompt, TextBackend};
@@ -39,8 +45,8 @@ use zeromq::prelude::{SocketRecv, SocketSend};
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
use super::pb::control_client::ControlClient;
use super::pb::generate_client::GenerateClient;
use super::{ControlServer, ControlServiceImpl, GenerateServer, GenerateServiceImpl, pb};
use super::pb::inference_client::InferenceClient;
use super::{ControlServer, ControlServiceImpl, InferenceServer, InferenceServiceImpl, pb};
use crate::listener::{Listener, MaybeTlsListener};
use crate::state::AppState;
use crate::tls;
@@ -202,7 +208,7 @@ async fn setup_grpc_service(
engine_id: impl Into<EngineId>,
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
) -> (
GenerateServer<GenerateServiceImpl>,
InferenceServer<InferenceServiceImpl>,
ControlServer<ControlServiceImpl>,
tokio::sync::watch::Receiver<bool>,
MockEngineTask,
@@ -246,7 +252,7 @@ async fn setup_grpc_service(
);
let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat));
(
GenerateServer::new(GenerateServiceImpl::new(state.clone())),
InferenceServer::new(InferenceServiceImpl::new(state.clone())),
ControlServer::new(ControlServiceImpl::new(state)),
engine_health,
engine_task,
@@ -259,30 +265,30 @@ async fn grpc_test_server(
engine_id: impl Into<EngineId>,
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
) -> (
GenerateClient<tonic::transport::Channel>,
InferenceClient<tonic::transport::Channel>,
tokio::task::JoinHandle<()>,
MockEngineTask,
) {
let (generate_service, control_service, engine_health, engine_task) =
let (inference_service, control_service, engine_health, engine_task) =
setup_grpc_service(engine_id, output_specs).await;
let (channel, server_task) = start_grpc_test_server(
generate_service,
inference_service,
control_service,
engine_health,
tokio_util::sync::CancellationToken::new(),
)
.await;
(GenerateClient::new(channel), server_task, engine_task)
(InferenceClient::new(channel), server_task, engine_task)
}
async fn start_grpc_test_server(
generate_service: GenerateServer<GenerateServiceImpl>,
inference_service: InferenceServer<InferenceServiceImpl>,
control_service: ControlServer<ControlServiceImpl>,
engine_health: tokio::sync::watch::Receiver<bool>,
shutdown: tokio_util::sync::CancellationToken,
) -> (Channel, tokio::task::JoinHandle<()>) {
let (health_reporter, health_service) = health_reporter();
health_reporter.set_serving::<GenerateServer<GenerateServiceImpl>>().await;
health_reporter.set_serving::<InferenceServer<InferenceServiceImpl>>().await;
health_reporter.set_serving::<ControlServer<ControlServiceImpl>>().await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
@@ -293,7 +299,7 @@ async fn start_grpc_test_server(
let server = TonicServer::builder()
.add_service(health_service)
.add_service(control_service)
.add_service(generate_service)
.add_service(inference_service)
.serve_with_incoming_shutdown(incoming, shutdown.clone().cancelled_owned());
let health_monitor =
super::monitor_health(health_reporter, engine_health, shutdown.clone());
@@ -323,7 +329,7 @@ async fn grpc_tls_test_server(
certs: &TestCerts,
cert_reqs: i32,
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
let (generate_service, control_service, _engine_health, engine_task) =
let (inference_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");
@@ -335,7 +341,7 @@ async fn grpc_tls_test_server(
let incoming = MaybeTlsListener::tls(Listener::Tcp(listener), context);
TonicServer::builder()
.add_service(control_service)
.add_service(generate_service)
.add_service(inference_service)
.serve_with_incoming(incoming)
.await
.expect("grpc tls server");
@@ -351,7 +357,7 @@ async fn grpc_tls_client(
certs: &TestCerts,
addr: &str,
identity: Option<&str>,
) -> Result<GenerateClient<Channel>, tonic::transport::Error> {
) -> Result<InferenceClient<Channel>, tonic::transport::Error> {
let ca = certs.path("ca.pem");
let identity = identity.map(|name| {
(
@@ -388,7 +394,7 @@ async fn grpc_tls_client(
.expect("grpc endpoint")
.connect_with_connector(connector)
.await?;
Ok(GenerateClient::new(channel))
Ok(InferenceClient::new(channel))
}
/// Complete a raw TLS handshake against the gRPC port (offering ALPN `h2`) for
@@ -415,7 +421,7 @@ async fn grpc_server_with_keepalive(
engine_id: impl Into<EngineId>,
keepalive: Option<Duration>,
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
let (generate_service, control_service, _engine_health, engine_task) =
let (inference_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");
@@ -432,7 +438,7 @@ async fn grpc_server_with_keepalive(
let incoming = MaybeTlsListener::plain(Listener::Tcp(listener));
builder
.add_service(control_service)
.add_service(generate_service)
.add_service(inference_service)
.serve_with_incoming(incoming)
.await
.expect("grpc server");
@@ -1083,20 +1089,20 @@ async fn grpc_without_keepalive_keeps_unresponsive_connection_open() {
#[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) =
let (inference_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,
inference_service,
control_service,
engine_health,
tokio_util::sync::CancellationToken::new(),
)
.await;
let mut generate_client = GenerateClient::new(channel.clone());
let mut inference_client = InferenceClient::new(channel.clone());
let mut control_client = ControlClient::new(channel);
let request_id = "test-abort-active";
let mut stream = generate_client
let mut stream = inference_client
.generate_stream(pb::GenerateRequest {
request_id: request_id.to_string(),
model: "test-model".to_string(),
@@ -1171,14 +1177,127 @@ async fn control_abort_resolves_external_id_and_empty_is_noop() {
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn control_reports_server_and_model_info() {
let (generate_service, control_service, engine_health, _engine_task) =
setup_grpc_service(b"engine-grpc-info", default_stream_output_specs()).await;
let (channel, server_task) = start_grpc_test_server(
generate_service,
control_service,
engine_health,
tokio_util::sync::CancellationToken::new(),
)
.await;
let mut client = ControlClient::new(channel);
let server = client
.get_server_info(pb::GetServerInfoRequest {})
.await
.expect("get server info")
.into_inner();
assert_eq!(server.engine_version, "test-vllm-version");
assert_eq!(server.api_version, "vllm");
assert_eq!(server.instance_id, "test-instance");
assert_eq!(server.max_model_len, DEFAULT_MOCK_MAX_MODEL_LEN as u32);
assert_eq!(server.kv_block_size, DEFAULT_MOCK_BLOCK_SIZE as u32);
assert_eq!(server.total_kv_blocks, DEFAULT_MOCK_NUM_GPU_BLOCKS);
assert_eq!(server.max_running_requests, 256);
assert_eq!(server.max_batched_tokens, 8_192);
let parallelism = server.parallelism.expect("parallelism metadata");
assert_eq!(parallelism.tensor_parallel_size, 1);
assert_eq!(parallelism.pipeline_parallel_size, 1);
assert_eq!(parallelism.data_parallel_size, 1);
assert_eq!(parallelism.data_parallel_rank, 0);
assert_eq!(parallelism.decode_context_parallel_size, 1);
let model = client
.get_model_info(pb::GetModelInfoRequest {})
.await
.expect("get model info")
.into_inner();
assert_eq!(model.model_id, "test-model");
assert_eq!(model.served_model_name, "test-model");
assert!(model.served_model_aliases.is_empty());
assert!(model.supports_text_input);
assert!(model.supports_token_ids_input);
assert!(!model.supports_multimodal);
assert!(model.reasoning_parser.is_empty());
assert!(model.tool_call_parser.is_empty());
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn control_aggregates_multi_engine_capacity() {
let ipc = IpcNamespace::new().expect("create ipc namespace");
let handshake_address = ipc.handshake_endpoint();
let mut ready_0 = default_ready_response();
ready_0.max_model_len = 8_192;
ready_0.num_gpu_blocks = 10;
ready_0.data_parallel_size = 2;
let mut ready_1 = default_ready_response();
ready_1.max_model_len = 4_096;
ready_1.num_gpu_blocks = 20;
ready_1.data_parallel_size = 2;
ready_1.data_parallel_rank = 1;
let engine_tasks = [ready_0, ready_1].map(|ready| {
let engine_id = EngineId::from_engine_index(ready.data_parallel_rank);
MockEngineTask::new(spawn_mock_engine_task_with_ready(
handshake_address.clone(),
engine_id,
ready,
|_, _| boxed_test_future(async {}),
))
});
let client = EngineCoreClient::connect(EngineCoreClientConfig {
transport_mode: TransportMode::HandshakeOwner {
handshake_address,
advertised_host: "127.0.0.1".to_string(),
engine_count: 2,
ready_timeout: Duration::from_secs(2),
local_input_address: Some(ipc.input_endpoint()),
local_output_address: Some(ipc.output_endpoint()),
},
coordinator_mode: None,
model_name: "test-model".to_string(),
client_index: 0,
})
.await
.expect("connect multi-engine client");
let chat = ChatLlm::from_shared_backend(
Llm::new(client),
Arc::new(FakeTextBackend) as Arc<dyn ChatTextBackend>,
);
let service = ControlServiceImpl::new(Arc::new(AppState::new(
vec!["test-model".to_string()],
chat,
)));
let server = pb::control_server::Control::get_server_info(
&service,
tonic::Request::new(pb::GetServerInfoRequest {}),
)
.await
.expect("get server info")
.into_inner();
assert_eq!(server.max_model_len, 4_096);
assert_eq!(server.total_kv_blocks, 30);
drop(engine_tasks);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_health_transitions_to_not_serving_when_engine_becomes_unhealthy() {
let (generate_service, control_service, _connected_engine_health, _engine_task) =
let (inference_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,
inference_service,
control_service,
engine_health,
tokio_util::sync::CancellationToken::new(),
@@ -1187,7 +1306,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", "vllm.Control", ""] {
for service in ["vllm.Inference", "vllm.Control", ""] {
let service_label = if service.is_empty() {
"overall"
} else {
@@ -1242,14 +1361,14 @@ 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, control_service, engine_health, _engine_task) = setup_grpc_service(
let (inference_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,
inference_service,
control_service,
engine_health,
shutdown.clone(),
@@ -1258,43 +1377,43 @@ async fn grpc_health_watch_closes_on_graceful_shutdown() {
let mut health_client = HealthClient::new(channel);
let mut stream = health_client
.watch(HealthCheckRequest {
service: "vllm.Generate".to_string(),
service: "vllm.Inference".to_string(),
})
.await
.expect("start health watch for vllm.Generate")
.expect("start health watch for vllm.Inference")
.into_inner();
let initial = stream
.message()
.await
.expect("read initial health status for vllm.Generate")
.expect("read initial health status for vllm.Inference")
.expect("health watch ended before its initial status");
assert_eq!(
initial.status,
HealthServingStatus::Serving as i32,
"unexpected initial health status for vllm.Generate"
"unexpected initial health status for vllm.Inference"
);
shutdown.cancel();
let update = tokio::time::timeout(Duration::from_secs(2), stream.message())
.await
.expect("timed out waiting for shutdown health update for vllm.Generate")
.expect("failed to read shutdown health update for vllm.Generate")
.expect("timed out waiting for shutdown health update for vllm.Inference")
.expect("failed to read shutdown health update for vllm.Inference")
.expect("health watch ended before its shutdown update");
assert_eq!(
update.status,
HealthServingStatus::NotServing as i32,
"unexpected shutdown health status for vllm.Generate"
"unexpected shutdown health status for vllm.Inference"
);
let stream_end = tokio::time::timeout(Duration::from_secs(2), stream.message())
.await
.expect("timed out waiting for vllm.Generate health watch to close")
.expect("failed while closing vllm.Generate health watch");
.expect("timed out waiting for vllm.Inference health watch to close")
.expect("failed while closing vllm.Inference health watch");
assert!(
stream_end.is_none(),
"vllm.Generate health watch remained open"
"vllm.Inference health watch remained open"
);
tokio::time::timeout(Duration::from_secs(2), server_task)
+5 -5
View File
@@ -187,7 +187,7 @@ where
let model = state.primary_model_name().to_owned();
let app = extend_router(build_router(state.clone()));
// Optionally bind the gRPC Generate server on a separate port. Bind
// Optionally bind the gRPC Inference server on a separate port. Bind
// synchronously here so bind errors (port in use, permission denied, ...)
// surface before serving rather than being deferred until shutdown.
let grpc_setup = if let Some(grpc_port) = config.grpc_port {
@@ -206,19 +206,19 @@ where
.context("invalid gRPC TLS configuration")?;
let (health_reporter, health_service) = health_reporter();
let engine_health = state.engine_core_client().subscribe_health();
health_reporter.set_serving::<grpc::GenerateGrpcService>().await;
health_reporter.set_serving::<grpc::InferenceGrpcService>().await;
health_reporter.set_serving::<grpc::ControlGrpcService>().await;
let control_service =
grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone()));
let generate_service =
grpc::GenerateGrpcService::new(grpc::GenerateServiceImpl::new(state.clone()));
let inference_service =
grpc::InferenceGrpcService::new(grpc::InferenceServiceImpl::new(state.clone()));
let svc = TonicServer::builder()
.http2_keepalive_interval(Some(GRPC_KEEPALIVE_INTERVAL))
.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);
.add_service(inference_service);
info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server");
Some((grpc_listener, svc, grpc_tls, health_reporter, engine_health))
} else {
+4 -4
View File
@@ -30,8 +30,8 @@ const OFFLOADED_PATHS: &[&str] = &[
"/detokenize",
"/inference/v1/generate",
// gRPC routes:
"/vllm.Generate/Generate",
"/vllm.Generate/GenerateStream",
"/vllm.Inference/Generate",
"/vllm.Inference/GenerateStream",
];
/// Return a Tower layer that runs selected data-plane requests on the request runtime,
@@ -124,8 +124,8 @@ mod tests {
assert!(should_offload("/tokenize"));
assert!(should_offload("/detokenize"));
assert!(should_offload("/inference/v1/generate"));
assert!(should_offload("/vllm.Generate/Generate"));
assert!(should_offload("/vllm.Generate/GenerateStream"));
assert!(should_offload("/vllm.Inference/Generate"));
assert!(should_offload("/vllm.Inference/GenerateStream"));
}
#[test]
@@ -328,6 +328,13 @@ def test_apply_ready_response_syncs_block_size():
vllm_version="test",
world_size=1,
data_parallel_size=1,
tensor_parallel_size=1,
pipeline_parallel_size=1,
decode_context_parallel_size=1,
data_parallel_rank=0,
max_num_seqs=256,
max_num_batched_tokens=8192,
instance_id="test-instance",
)
)
client._apply_ready_response(payload)
+7
View File
@@ -82,6 +82,13 @@ class EngineCoreReadyResponse:
vllm_version: str
world_size: int
data_parallel_size: int
tensor_parallel_size: int
pipeline_parallel_size: int
decode_context_parallel_size: int
data_parallel_rank: int
max_num_seqs: int
max_num_batched_tokens: int
instance_id: str
# KV cache capacity (None for encoder-only/attention-free models).
kv_cache_size_tokens: int | None = None
kv_cache_max_concurrency: float | None = None
+26 -16
View File
@@ -1613,6 +1613,31 @@ class EngineCoreProc(EngineCore):
"to send. Please report this issue."
)
def _make_ready_response(self) -> EngineCoreReadyResponse:
parallel_config = self.vllm_config.parallel_config
scheduler_config = self.vllm_config.scheduler_config
return EngineCoreReadyResponse(
max_model_len=self.vllm_config.model_config.max_model_len,
num_gpu_blocks=self.vllm_config.cache_config.num_gpu_blocks or 0,
block_size=self.vllm_config.cache_config.block_size,
dp_stats_address=self.frontend_stats_publish_address,
dtype=str(self.vllm_config.model_config.dtype).removeprefix("torch."),
vllm_version=VLLM_VERSION,
world_size=self.vllm_config.parallel_config.world_size,
data_parallel_size=parallel_config.data_parallel_size,
kv_cache_size_tokens=self.vllm_config.cache_config.kv_cache_size_tokens,
kv_cache_max_concurrency=(
self.vllm_config.cache_config.kv_cache_max_concurrency
),
tensor_parallel_size=parallel_config.tensor_parallel_size,
pipeline_parallel_size=parallel_config.pipeline_parallel_size,
decode_context_parallel_size=parallel_config.decode_context_parallel_size,
data_parallel_rank=self.engine_index,
max_num_seqs=scheduler_config.max_num_seqs,
max_num_batched_tokens=scheduler_config.max_num_batched_tokens,
instance_id=self.vllm_config.instance_id,
)
def process_input_sockets(
self,
input_addresses: list[str],
@@ -1654,22 +1679,7 @@ class EngineCoreProc(EngineCore):
# Register sockets with poller.
poller = zmq.Poller()
ready_response = EngineCoreReadyResponse(
max_model_len=self.vllm_config.model_config.max_model_len,
num_gpu_blocks=self.vllm_config.cache_config.num_gpu_blocks or 0,
block_size=self.vllm_config.cache_config.block_size,
dp_stats_address=self.frontend_stats_publish_address,
dtype=str(self.vllm_config.model_config.dtype).removeprefix("torch."),
vllm_version=VLLM_VERSION,
world_size=self.vllm_config.parallel_config.world_size,
data_parallel_size=self.vllm_config.parallel_config.data_parallel_size,
kv_cache_size_tokens=(
self.vllm_config.cache_config.kv_cache_size_tokens
),
kv_cache_max_concurrency=(
self.vllm_config.cache_config.kv_cache_max_concurrency
),
)
ready_response = self._make_ready_response()
ready_payload = msgspec.msgpack.encode(ready_response)
for input_socket in input_sockets:
# Send initial message to each input socket - this is required