forked from Karylab-cklius/vllm
[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:
co-authored by
Nick Hill
parent
2b465b2c42
commit
e3c2fc3b3c
@@ -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 {}
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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 {}))
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user