From e3c2fc3b3ca3f41466126cd2aa0b228eb8465844 Mon Sep 17 00:00:00 2001 From: Connor Carpenter Date: Mon, 27 Jul 2026 09:53:27 -0700 Subject: [PATCH] [Rust Frontend][gRPC] Add server and model discovery (#49491) Signed-off-by: Connor Carpenter Co-authored-by: Nick Hill --- rust/proto/control.proto | 53 +++++ .../{vllm_grpc.proto => inference.proto} | 16 +- rust/src/chat/src/lib.rs | 27 +++ rust/src/engine-core-client/src/client.rs | 11 + .../src/engine-core-client/src/mock_engine.rs | 7 + .../src/protocol/handshake.rs | 15 ++ .../src/tests/python_compat.py | 14 ++ rust/src/server/build.rs | 8 +- rust/src/server/src/grpc/control.rs | 106 ++++++++++ rust/src/server/src/grpc/health.rs | 8 +- rust/src/server/src/grpc/inference.rs | 156 ++++++++++++++ rust/src/server/src/grpc/mod.rs | 194 +---------------- rust/src/server/src/grpc/tests.rs | 195 ++++++++++++++---- rust/src/server/src/lib.rs | 10 +- rust/src/server/src/middleware/offload.rs | 8 +- tests/v1/engine/test_engine_core_client.py | 7 + vllm/v1/engine/__init__.py | 7 + vllm/v1/engine/core.py | 42 ++-- 18 files changed, 615 insertions(+), 269 deletions(-) create mode 100644 rust/proto/control.proto rename rust/proto/{vllm_grpc.proto => inference.proto} (93%) create mode 100644 rust/src/server/src/grpc/control.rs create mode 100644 rust/src/server/src/grpc/inference.rs diff --git a/rust/proto/control.proto b/rust/proto/control.proto new file mode 100644 index 00000000000..7b858cb3e20 --- /dev/null +++ b/rust/proto/control.proto @@ -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 {} diff --git a/rust/proto/vllm_grpc.proto b/rust/proto/inference.proto similarity index 93% rename from rust/proto/vllm_grpc.proto rename to rust/proto/inference.proto index c10da4818e3..b1c08ae76e6 100644 --- a/rust/proto/vllm_grpc.proto +++ b/rust/proto/inference.proto @@ -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 {} diff --git a/rust/src/chat/src/lib.rs b/rust/src/chat/src/lib.rs index e584257eaa4..9c38c40be4e 100644 --- a/rust/src/chat/src/lib.rs +++ b/rust/src/chat/src/lib.rs @@ -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 { let (text_request, output_processor) = self diff --git a/rust/src/engine-core-client/src/client.rs b/rust/src/engine-core-client/src/client.rs index 9ea7b703671..705cb003e95 100644 --- a/rust/src/engine-core-client/src/client.rs +++ b/rust/src/engine-core-client/src/client.rs @@ -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 diff --git a/rust/src/engine-core-client/src/mock_engine.rs b/rust/src/engine-core-client/src/mock_engine.rs index e8f277b128d..5540ae8bc11 100644 --- a/rust/src/engine-core-client/src/mock_engine.rs +++ b/rust/src/engine-core-client/src/mock_engine.rs @@ -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, } diff --git a/rust/src/engine-core-client/src/protocol/handshake.rs b/rust/src/engine-core-client/src/protocol/handshake.rs index c8545017a96..125a82b39d2 100644 --- a/rust/src/engine-core-client/src/protocol/handshake.rs +++ b/rust/src/engine-core-client/src/protocol/handshake.rs @@ -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, /// Maximum achievable request concurrency given the KV cache, if reported. diff --git a/rust/src/engine-core-client/src/tests/python_compat.py b/rust/src/engine-core-client/src/tests/python_compat.py index 0005ff20f92..ed16f2bd8b3 100755 --- a/rust/src/engine-core-client/src/tests/python_compat.py +++ b/rust/src/engine-core-client/src/tests/python_compat.py @@ -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()) diff --git a/rust/src/server/build.rs b/rust/src/server/build.rs index c20ff1c86b8..585c3c70b99 100644 --- a/rust/src/server/build.rs +++ b/rust/src/server/build.rs @@ -9,7 +9,13 @@ fn main() -> Result<(), Box> { .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(()) } diff --git a/rust/src/server/src/grpc/control.rs b/rust/src/server/src/grpc/control.rs new file mode 100644 index 00000000000..e989e6abaed --- /dev/null +++ b/rust/src/server/src/grpc/control.rs @@ -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; + +/// gRPC control service backed by the shared application state. +pub struct ControlServiceImpl { + state: Arc, +} + +impl ControlServiceImpl { + pub fn new(state: Arc) -> 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, + ) -> Result, 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, + ) -> Result, 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, + ) -> 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 {})) + } +} diff --git a/rust/src/server/src/grpc/health.rs b/rust/src/server/src/grpc/health.rs index 554dfae675f..cd35aec972e 100644 --- a/rust/src/server/src/grpc/health.rs +++ b/rust/src/server/src/grpc/health.rs @@ -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, 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::().await; + health_reporter.set_not_serving::().await; health_reporter.set_not_serving::().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; } diff --git a/rust/src/server/src/grpc/inference.rs b/rust/src/server/src/grpc/inference.rs new file mode 100644 index 00000000000..56fa40924cf --- /dev/null +++ b/rust/src/server/src/grpc/inference.rs @@ -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; + +/// gRPC inference service backed by the shared application state. +pub struct InferenceServiceImpl { + state: Arc, +} + +impl InferenceServiceImpl { + pub fn new(state: Arc) -> Self { + Self { state } + } +} + +#[tonic::async_trait] +impl pb::inference_server::Inference for InferenceServiceImpl { + type GenerateStreamStream = + Pin> + Send>>; + + /// Unary generate: collect all output and return a single response. + async fn generate( + &self, + request: Request, + ) -> Result, 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, + ) -> Result, 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) + } +} diff --git a/rust/src/server/src/grpc/mod.rs b/rust/src/server/src/grpc/mod.rs index f5d2f347bb3..06c7c9259eb 100644 --- a/rust/src/server/src/grpc/mod.rs +++ b/rust/src/server/src/grpc/mod.rs @@ -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; -pub(crate) type GenerateGrpcService = GenerateServer; +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, -} - -impl GenerateServiceImpl { - pub fn new(state: Arc) -> Self { - Self { state } - } -} - -/// 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 = - Pin> + Send>>; - - /// Unary generate: collect all output and return a single response. - async fn generate( - &self, - request: Request, - ) -> Result, 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, - ) -> Result, 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) - } -} diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 4cd67f4a4b3..75da670ec3b 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -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, output_specs: Vec<(Vec, Option)>, ) -> ( - GenerateServer, + InferenceServer, ControlServer, tokio::sync::watch::Receiver, 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, output_specs: Vec<(Vec, Option)>, ) -> ( - GenerateClient, + InferenceClient, 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, + inference_service: InferenceServer, 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; health_reporter.set_serving::>().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, tonic::transport::Error> { +) -> Result, 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, keepalive: Option, ) -> (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, + ); + 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) diff --git a/rust/src/server/src/lib.rs b/rust/src/server/src/lib.rs index 1f3468fe46d..470c424cdb8 100644 --- a/rust/src/server/src/lib.rs +++ b/rust/src/server/src/lib.rs @@ -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::().await; + 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 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 { diff --git a/rust/src/server/src/middleware/offload.rs b/rust/src/server/src/middleware/offload.rs index edc6560bf20..c310ab442fc 100644 --- a/rust/src/server/src/middleware/offload.rs +++ b/rust/src/server/src/middleware/offload.rs @@ -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] diff --git a/tests/v1/engine/test_engine_core_client.py b/tests/v1/engine/test_engine_core_client.py index 0bdca7ada99..64adf7a8b3c 100644 --- a/tests/v1/engine/test_engine_core_client.py +++ b/tests/v1/engine/test_engine_core_client.py @@ -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) diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 4ac27be5068..e80be0e45d7 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -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 diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 39393f64ae4..9817c474343 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -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