From 7f3ab290464ac319e867b0d011d11dd6b2ff37f4 Mon Sep 17 00:00:00 2001 From: Connor Carpenter Date: Thu, 6 Aug 2026 12:40:42 -0700 Subject: [PATCH 1/3] feat(grpc): add RL lifecycle control Signed-off-by: Connor Carpenter --- docs/training/weight_transfer/README.md | 2 + docs/usage/security.md | 7 +- rust/proto/control.proto | 81 +++++ rust/src/engine-core-client/src/client.rs | 73 +++++ .../src/engine-core-client/src/mock_engine.rs | 3 + .../src/protocol/handshake.rs | 9 + .../src/tests/python_compat.py | 6 + rust/src/server/src/grpc/control.rs | 300 +++++++++++++++++- rust/src/server/src/grpc/tests.rs | 247 ++++++++++++++ rust/src/server/src/lib.rs | 3 +- vllm/v1/engine/__init__.py | 3 + vllm/v1/engine/core.py | 7 + 12 files changed, 736 insertions(+), 5 deletions(-) diff --git a/docs/training/weight_transfer/README.md b/docs/training/weight_transfer/README.md index b8d39763181d..03629133b755 100644 --- a/docs/training/weight_transfer/README.md +++ b/docs/training/weight_transfer/README.md @@ -63,6 +63,8 @@ When running vLLM as an HTTP server, the following endpoints are available for w !!! note The HTTP weight transfer endpoints require `VLLM_SERVER_DEV_MODE=1` to be set. +The Rust frontend's optional gRPC `Control` service exposes the same pause, sleep, weight-transfer, and weight-version lifecycle for trusted sidecars. The `ServerInfo.rl_capabilities` response reports whether weight transfer and sleep mode were configured. Backend-specific `init_info` and `update_info` remain JSON metadata; model tensors continue to move over the configured NCCL, IPC, or sparse-NCCL transport. + ## Trainer-Side API Both backends provide static methods that the trainer calls to send weights. The general pattern is: diff --git a/docs/usage/security.md b/docs/usage/security.md index 7eb57c9aa24e..5174fb8c5c23 100644 --- a/docs/usage/security.md +++ b/docs/usage/security.md @@ -344,7 +344,7 @@ vLLM supports loading out-of-tree HTTP routes via the `vllm.endpoint_plugins` en ## gRPC Interface -vLLM provides an optional gRPC Generate service on a separate TCP port, enabled via the `--grpc-port` flag. When not specified, no gRPC server is started. The gRPC listener binds to the same host address as the HTTP server. +vLLM provides optional gRPC `Inference` and `Control` services on a separate TCP port, enabled via the `--grpc-port` flag. When not specified, no gRPC server is started. The gRPC listener binds to the same host address as the HTTP server. **Warning:** The gRPC interface is **insecure by default** — it does not implement authentication, authorization, or encryption. It should be considered a private, internal interface intended for use only between co-located services within a trusted network. Do not expose the gRPC port to the public internet or untrusted clients. If you enable the gRPC interface, protect it via network-level access controls such as firewall rules, network segmentation, or deployment on an isolated private network. @@ -353,8 +353,9 @@ vLLM provides an optional gRPC Generate service on a separate TCP port, enabled An attacker who can reach the gRPC port can: 1. **Run arbitrary inference** via the `Generate` and `GenerateStream` RPCs without any credentials -2. **Consume GPU and compute resources** by submitting unbounded generation requests -3. **Cause Denial of Service** by exploiting bugs in the gRPC interface that can crash vLLM. +2. **Mutate engine state** by pausing generation, sleeping the engine, or initiating configured RL weight updates through the `Control` service +3. **Consume GPU and compute resources** by submitting unbounded generation requests +4. **Cause Denial of Service** by exploiting bugs in the gRPC interface that can crash vLLM. ### Recommendations diff --git a/rust/proto/control.proto b/rust/proto/control.proto index 0e25aea26474..428ad6460499 100644 --- a/rust/proto/control.proto +++ b/rust/proto/control.proto @@ -9,6 +9,21 @@ service Control { rpc GetModelInfo (GetModelInfoRequest) returns (ModelInfo) {} rpc Abort (AbortRequest) returns (AbortResponse) {} rpc GetKvEventSources (GetKvEventSourcesRequest) returns (GetKvEventSourcesResponse) {} + + // Reinforcement-learning lifecycle and weight updates. + rpc PauseGeneration (PauseGenerationRequest) returns (PauseGenerationResponse) {} + rpc ResumeGeneration (ResumeGenerationRequest) returns (ResumeGenerationResponse) {} + rpc IsPaused (IsPausedRequest) returns (IsPausedResponse) {} + rpc Sleep (SleepRequest) returns (SleepResponse) {} + rpc WakeUp (WakeUpRequest) returns (WakeUpResponse) {} + rpc IsSleeping (IsSleepingRequest) returns (IsSleepingResponse) {} + rpc InitWeightTransferEngine (InitWeightTransferEngineRequest) returns (InitWeightTransferEngineResponse) {} + rpc StartWeightUpdate (StartWeightUpdateRequest) returns (StartWeightUpdateResponse) {} + rpc StartDraftWeightUpdate (StartDraftWeightUpdateRequest) returns (StartDraftWeightUpdateResponse) {} + rpc UpdateWeights (UpdateWeightsRequest) returns (UpdateWeightsResponse) {} + rpc FinishWeightUpdate (FinishWeightUpdateRequest) returns (FinishWeightUpdateResponse) {} + rpc UpdateWeightVersion (UpdateWeightVersionRequest) returns (UpdateWeightVersionResponse) {} + rpc GetWeightVersion (GetWeightVersionRequest) returns (GetWeightVersionResponse) {} } message GetServerInfoRequest {} @@ -23,6 +38,14 @@ message ServerInfo { uint64 total_kv_blocks = 7; uint64 max_running_requests = 8; uint64 max_batched_tokens = 9; + RlCapabilities rl_capabilities = 11; +} + +message RlCapabilities { + bool weight_transfer_enabled = 1; + string weight_transfer_backend = 2; + bool sleep_mode_enabled = 3; + bool draft_weight_updates_enabled = 4; } message ParallelismInfo { @@ -53,6 +76,64 @@ message AbortRequest { message AbortResponse {} +// ====================================================================================== +// Reinforcement-learning control +// ====================================================================================== + +enum PauseMode { + PAUSE_MODE_UNSPECIFIED = 0; + PAUSE_MODE_ABORT = 1; + PAUSE_MODE_WAIT = 2; + PAUSE_MODE_KEEP = 3; +} + +message PauseGenerationRequest { + PauseMode mode = 1; + optional bool clear_cache = 2; +} +message PauseGenerationResponse {} + +message ResumeGenerationRequest {} +message ResumeGenerationResponse {} + +message IsPausedRequest {} +message IsPausedResponse { bool paused = 1; } + +message SleepRequest { + optional uint32 level = 1; + PauseMode mode = 2; +} +message SleepResponse {} + +message WakeUpRequest { repeated string tags = 1; } +message WakeUpResponse {} + +message IsSleepingRequest {} +message IsSleepingResponse { bool sleeping = 1; } + +// The payloads are backend-specific JSON objects. Tensor data remains on the +// configured NCCL, IPC, or sparse-NCCL transport rather than crossing gRPC. +message InitWeightTransferEngineRequest { bytes init_info_json = 1; } +message InitWeightTransferEngineResponse {} + +message StartWeightUpdateRequest {} +message StartWeightUpdateResponse {} + +message StartDraftWeightUpdateRequest {} +message StartDraftWeightUpdateResponse {} + +message UpdateWeightsRequest { bytes update_info_json = 1; } +message UpdateWeightsResponse {} + +message FinishWeightUpdateRequest { optional string weight_version = 1; } +message FinishWeightUpdateResponse {} + +message UpdateWeightVersionRequest { string weight_version = 1; } +message UpdateWeightVersionResponse {} + +message GetWeightVersionRequest {} +message GetWeightVersionResponse { string weight_version = 1; } + // ====================================================================================== // KV discovery // ====================================================================================== diff --git a/rust/src/engine-core-client/src/client.rs b/rust/src/engine-core-client/src/client.rs index a5709e52d67f..c6d775c36c52 100644 --- a/rust/src/engine-core-client/src/client.rs +++ b/rust/src/engine-core-client/src/client.rs @@ -1,12 +1,14 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright contributors to the vLLM project +use std::collections::BTreeMap; use std::sync::Arc; use std::time::Duration; use futures::future::{join_all, try_join_all}; use itertools::Itertools; use serde::Serialize; +use serde_json::Value as JsonValue; use tokio::sync::mpsc; use tokio_util::task::AbortOnDropHandle; use tracing::{debug, info, trace}; @@ -683,6 +685,77 @@ impl EngineCoreClient { .collect()) } + /// Initialize the configured RL weight-transfer backend. + pub async fn init_weight_transfer_engine(&self, init_info: JsonValue) -> Result<()> { + self.collective_rpc( + "init_weight_transfer_engine", + None, + Vec::::new(), + BTreeMap::from([("init_info".to_string(), init_info)]), + ) + .await?; + Ok(()) + } + + /// Start a weight update for the base model. + pub async fn start_weight_update(&self) -> Result<()> { + self.collective_rpc( + "start_weight_update", + None, + Vec::::new(), + BTreeMap::::new(), + ) + .await?; + Ok(()) + } + + /// Start a weight update for the speculative draft model. + pub async fn start_draft_weight_update(&self) -> Result<()> { + self.collective_rpc( + "start_draft_weight_update", + None, + Vec::::new(), + BTreeMap::::new(), + ) + .await?; + Ok(()) + } + + /// Apply one backend-specific weight metadata chunk. + pub async fn update_weights(&self, update_info: JsonValue) -> Result<()> { + self.collective_rpc( + "update_weights", + None, + Vec::::new(), + BTreeMap::from([("update_info".to_string(), update_info)]), + ) + .await?; + Ok(()) + } + + /// Finish the current weight update. + pub async fn finish_weight_update(&self) -> Result<()> { + self.collective_rpc( + "finish_weight_update", + None, + Vec::::new(), + BTreeMap::::new(), + ) + .await?; + Ok(()) + } + + /// Set the committed weight version on every connected engine. + pub async fn set_weight_version(&self, weight_version: &str) -> Result<()> { + self.call_utility::<(), _>("set_weight_version", (weight_version,)).await?; + Ok(()) + } + + /// Return the committed weight version agreed on by every connected engine. + pub async fn get_weight_version(&self) -> Result { + self.call_utility_consensus("get_weight_version", ()).await + } + /// Return whether the engine is currently sleeping at any level. pub async fn is_sleeping(&self) -> Result { self.call_utility_consensus("is_sleeping", ()).await diff --git a/rust/src/engine-core-client/src/mock_engine.rs b/rust/src/engine-core-client/src/mock_engine.rs index ee066635ff0c..28fb4d84fe05 100644 --- a/rust/src/engine-core-client/src/mock_engine.rs +++ b/rust/src/engine-core-client/src/mock_engine.rs @@ -67,6 +67,9 @@ pub fn default_ready_response() -> EngineCoreReadyResponse { kv_cache_size_tokens: None, kv_cache_max_concurrency: None, kv_events_config: None, + weight_transfer_backend: None, + enable_sleep_mode: false, + supports_draft_weight_updates: false, } } diff --git a/rust/src/engine-core-client/src/protocol/handshake.rs b/rust/src/engine-core-client/src/protocol/handshake.rs index a4e1cebde67d..b4214a06f27f 100644 --- a/rust/src/engine-core-client/src/protocol/handshake.rs +++ b/rust/src/engine-core-client/src/protocol/handshake.rs @@ -87,6 +87,15 @@ pub struct EngineCoreReadyResponse { /// KV-event publisher configuration, if configured. #[serde(default)] pub kv_events_config: Option, + /// Configured RL weight-transfer backend, if weight transfer is enabled. + #[serde(default)] + pub weight_transfer_backend: Option, + /// Whether the engine was started with sleep mode enabled. + #[serde(default)] + pub enable_sleep_mode: bool, + /// Whether the engine has a speculative draft model that can be updated. + #[serde(default)] + pub supports_draft_weight_updates: bool, } /// Frontend-owned ZMQ addresses that are sent to the engine during startup 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 0d95dbba1950..ea090a939d52 100755 --- a/rust/src/engine-core-client/src/tests/python_compat.py +++ b/rust/src/engine-core-client/src/tests/python_compat.py @@ -387,6 +387,9 @@ class EngineCoreReadyResponse: kv_cache_size_tokens: int | None = None kv_cache_max_concurrency: float | None = None kv_events_config: KVEventsConfig | None = None + weight_transfer_backend: str | None = None + enable_sleep_mode: bool = False + supports_draft_weight_updates: bool = False ready_response = EngineCoreReadyResponse( @@ -405,6 +408,9 @@ class EngineCoreReadyResponse: max_num_seqs=256, max_num_batched_tokens=8192, instance_id="test-instance", + weight_transfer_backend="nccl", + enable_sleep_mode=True, + supports_draft_weight_updates=True, kv_events_config=KVEventsConfig( enable_kv_cache_events=True, publisher="zmq", diff --git a/rust/src/server/src/grpc/control.rs b/rust/src/server/src/grpc/control.rs index 897209822252..d9fb4d8aa747 100644 --- a/rust/src/server/src/grpc/control.rs +++ b/rust/src/server/src/grpc/control.rs @@ -3,9 +3,13 @@ use std::sync::Arc; +use serde_json::Value as JsonValue; use thiserror_ext::AsReport as _; +use tokio::sync::Mutex; use tonic::{Request, Response, Status}; +use vllm_engine_core_client::EngineCoreClient; use vllm_engine_core_client::protocol::handshake::EngineCoreReadyResponse; +use vllm_engine_core_client::protocol::utility::PauseMode as EnginePauseMode; use super::{ControlServer, pb}; use crate::state::AppState; @@ -15,17 +19,25 @@ pub(crate) type ControlGrpcService = ControlServer; /// gRPC control service backed by the shared application state. pub struct ControlServiceImpl { state: Arc, + rl_lock: Mutex<()>, } impl ControlServiceImpl { pub fn new(state: Arc) -> Self { - Self { state } + Self { + state, + rl_lock: Mutex::new(()), + } } fn ready(&self) -> &EngineCoreReadyResponse { self.state.engine_core_client().ready_response() } + fn client(&self) -> &EngineCoreClient { + self.state.engine_core_client() + } + fn parallelism_info(&self) -> pb::ParallelismInfo { let ready = self.ready(); pb::ParallelismInfo { @@ -36,10 +48,102 @@ impl ControlServiceImpl { decode_context_parallel_size: ready.decode_context_parallel_size, } } + + fn weight_transfer_backend(&self) -> Option<&str> { + let responses = self.client().ready_responses(); + let backend = responses.first()?.weight_transfer_backend.as_deref()?; + responses + .iter() + .all(|ready| ready.weight_transfer_backend.as_deref() == Some(backend)) + .then_some(backend) + } + + fn sleep_mode_enabled(&self) -> bool { + self.client().ready_responses().iter().all(|ready| ready.enable_sleep_mode) + } + + fn draft_weight_updates_enabled(&self) -> bool { + self.client() + .ready_responses() + .iter() + .all(|ready| ready.supports_draft_weight_updates) + } + + fn rl_capabilities(&self) -> pb::RlCapabilities { + let backend = self.weight_transfer_backend(); + pb::RlCapabilities { + weight_transfer_enabled: backend.is_some(), + weight_transfer_backend: backend.unwrap_or_default().to_string(), + sleep_mode_enabled: self.sleep_mode_enabled(), + draft_weight_updates_enabled: self.draft_weight_updates_enabled(), + } + } + + fn require_weight_transfer(&self) -> Result<(), Status> { + self.weight_transfer_backend().map(|_| ()).ok_or_else(|| { + Status::failed_precondition( + "weight transfer is not configured; start vLLM with --weight-transfer-config", + ) + }) + } + + fn require_sleep_mode(&self) -> Result<(), Status> { + self.sleep_mode_enabled().then_some(()).ok_or_else(|| { + Status::failed_precondition( + "sleep mode is not configured; start vLLM with --enable-sleep-mode", + ) + }) + } + + async fn require_paused(&self) -> Result<(), Status> { + let paused = self + .client() + .is_scheduler_paused() + .await + .map_err(|error| utility_status("is_scheduler_paused", error))?; + paused + .then_some(()) + .ok_or_else(|| Status::failed_precondition("pause generation before updating weights")) + } } const GRPC_API_VERSION: &str = "vllm"; +fn utility_status(method: &'static str, error: vllm_engine_core_client::Error) -> Status { + Status::internal(format!("{method} failed: {}", error.to_report_string())) +} + +fn pause_mode(mode: i32) -> Result { + match pb::PauseMode::try_from(mode) { + Ok(pb::PauseMode::Unspecified | pb::PauseMode::Abort) => Ok(EnginePauseMode::Abort), + Ok(pb::PauseMode::Wait) => Ok(EnginePauseMode::Wait), + Ok(pb::PauseMode::Keep) => Ok(EnginePauseMode::Keep), + Err(_) => Err(Status::invalid_argument("invalid pause mode")), + } +} + +fn json_object(bytes: &[u8], field: &'static str) -> Result { + let value = serde_json::from_slice::(bytes).map_err(|error| { + Status::invalid_argument(format!( + "{field} must contain valid JSON: {}", + error.to_report_string() + )) + })?; + if !value.is_object() { + return Err(Status::invalid_argument(format!( + "{field} must contain a JSON object" + ))); + } + Ok(value) +} + +fn weight_version(value: String) -> Result { + if value.trim().is_empty() { + return Err(Status::invalid_argument("weight_version must not be empty")); + } + Ok(value) +} + #[tonic::async_trait] impl pb::control_server::Control for ControlServiceImpl { async fn get_server_info( @@ -57,6 +161,7 @@ impl pb::control_server::Control for ControlServiceImpl { 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, + rl_capabilities: Some(self.rl_capabilities()), })) } @@ -112,6 +217,199 @@ impl pb::control_server::Control for ControlServiceImpl { let sources = client.ready_responses().into_iter().filter_map(kv_event_source).collect(); Ok(Response::new(pb::GetKvEventSourcesResponse { sources })) } + + async fn pause_generation( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + let mode = pause_mode(request.mode)?; + let clear_cache = request.clear_cache.unwrap_or(true); + let _guard = self.rl_lock.lock().await; + self.client() + .pause_scheduler(mode, clear_cache) + .await + .map_err(|error| utility_status("pause_generation", error))?; + Ok(Response::new(pb::PauseGenerationResponse {})) + } + + async fn resume_generation( + &self, + _request: Request, + ) -> Result, Status> { + let _guard = self.rl_lock.lock().await; + self.client() + .resume_scheduler() + .await + .map_err(|error| utility_status("resume_generation", error))?; + Ok(Response::new(pb::ResumeGenerationResponse {})) + } + + async fn is_paused( + &self, + _request: Request, + ) -> Result, Status> { + let paused = self + .client() + .is_scheduler_paused() + .await + .map_err(|error| utility_status("is_paused", error))?; + Ok(Response::new(pb::IsPausedResponse { paused })) + } + + async fn sleep( + &self, + request: Request, + ) -> Result, Status> { + self.require_sleep_mode()?; + let request = request.into_inner(); + let mode = pause_mode(request.mode)?; + let level = request.level.unwrap_or(1); + let _guard = self.rl_lock.lock().await; + self.client() + .sleep(level, mode) + .await + .map_err(|error| utility_status("sleep", error))?; + Ok(Response::new(pb::SleepResponse {})) + } + + async fn wake_up( + &self, + request: Request, + ) -> Result, Status> { + self.require_sleep_mode()?; + let tags = request.into_inner().tags; + let tags = (!tags.is_empty()).then_some(tags); + let _guard = self.rl_lock.lock().await; + self.client() + .wake_up(tags) + .await + .map_err(|error| utility_status("wake_up", error))?; + Ok(Response::new(pb::WakeUpResponse {})) + } + + async fn is_sleeping( + &self, + _request: Request, + ) -> Result, Status> { + self.require_sleep_mode()?; + let sleeping = self + .client() + .is_sleeping() + .await + .map_err(|error| utility_status("is_sleeping", error))?; + Ok(Response::new(pb::IsSleepingResponse { sleeping })) + } + + async fn init_weight_transfer_engine( + &self, + request: Request, + ) -> Result, Status> { + self.require_weight_transfer()?; + let init_info = json_object(&request.into_inner().init_info_json, "init_info_json")?; + let _guard = self.rl_lock.lock().await; + self.client() + .init_weight_transfer_engine(init_info) + .await + .map_err(|error| utility_status("init_weight_transfer_engine", error))?; + Ok(Response::new(pb::InitWeightTransferEngineResponse {})) + } + + async fn start_weight_update( + &self, + _request: Request, + ) -> Result, Status> { + self.require_weight_transfer()?; + let _guard = self.rl_lock.lock().await; + self.require_paused().await?; + self.client() + .start_weight_update() + .await + .map_err(|error| utility_status("start_weight_update", error))?; + Ok(Response::new(pb::StartWeightUpdateResponse {})) + } + + async fn start_draft_weight_update( + &self, + _request: Request, + ) -> Result, Status> { + self.require_weight_transfer()?; + if !self.draft_weight_updates_enabled() { + return Err(Status::failed_precondition( + "draft weight updates require a configured speculative draft model", + )); + } + let _guard = self.rl_lock.lock().await; + self.require_paused().await?; + self.client() + .start_draft_weight_update() + .await + .map_err(|error| utility_status("start_draft_weight_update", error))?; + Ok(Response::new(pb::StartDraftWeightUpdateResponse {})) + } + + async fn update_weights( + &self, + request: Request, + ) -> Result, Status> { + self.require_weight_transfer()?; + let update_info = json_object(&request.into_inner().update_info_json, "update_info_json")?; + let _guard = self.rl_lock.lock().await; + self.require_paused().await?; + self.client() + .update_weights(update_info) + .await + .map_err(|error| utility_status("update_weights", error))?; + Ok(Response::new(pb::UpdateWeightsResponse {})) + } + + async fn finish_weight_update( + &self, + request: Request, + ) -> Result, Status> { + self.require_weight_transfer()?; + let version = request.into_inner().weight_version.map(weight_version).transpose()?; + let _guard = self.rl_lock.lock().await; + self.require_paused().await?; + self.client() + .finish_weight_update() + .await + .map_err(|error| utility_status("finish_weight_update", error))?; + if let Some(version) = version { + self.client() + .set_weight_version(&version) + .await + .map_err(|error| utility_status("update_weight_version", error))?; + } + Ok(Response::new(pb::FinishWeightUpdateResponse {})) + } + + async fn update_weight_version( + &self, + request: Request, + ) -> Result, Status> { + let version = weight_version(request.into_inner().weight_version)?; + let _guard = self.rl_lock.lock().await; + self.client() + .set_weight_version(&version) + .await + .map_err(|error| utility_status("update_weight_version", error))?; + Ok(Response::new(pb::UpdateWeightVersionResponse {})) + } + + async fn get_weight_version( + &self, + _request: Request, + ) -> Result, Status> { + let weight_version = self + .client() + .get_weight_version() + .await + .map_err(|error| utility_status("get_weight_version", error))?; + Ok(Response::new(pb::GetWeightVersionResponse { + weight_version, + })) + } } pub(super) fn kv_event_source(response: &EngineCoreReadyResponse) -> Option { diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index fdc98333f178..ef06a976ca36 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright contributors to the vLLM project +use std::collections::BTreeMap; use std::fs; use std::future::Future; use std::io; @@ -12,6 +13,7 @@ use std::time::Duration; use futures::StreamExt as _; use hyper_util::rt::TokioIo; use openssl::ssl::{SslConnector, SslFiletype, SslMethod}; +use rmpv::Value; use serial_test::serial; use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; use tokio::net::TcpStream; @@ -30,11 +32,14 @@ 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::decode_value; use vllm_engine_core_client::protocol::handshake::KvEventsConfig; use vllm_engine_core_client::protocol::output::{ EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, RequestBatchOutputs, + UtilityCallOutput, }; use vllm_engine_core_client::protocol::request::EngineCoreRequest; +use vllm_engine_core_client::protocol::utility::{UtilityOutput, UtilityResultEnvelope}; use vllm_engine_core_client::test_utils::{ IpcNamespace, spawn_mock_engine_task, spawn_mock_engine_task_with_ready, }; @@ -163,6 +168,49 @@ async fn recv_engine_message(dealer: &mut DealerSocket) -> Vec { dealer.recv().await.expect("recv engine message").into_vec() } +async fn recv_utility_call(dealer: &mut DealerSocket) -> (u64, String, Value) { + let frames = recv_engine_message(dealer).await; + assert_eq!(frames[0].as_ref(), &[0x03]); + let payload = decode_value(&frames[1]).expect("decode utility payload"); + let fields = payload.as_array().expect("utility payload array"); + ( + fields[1].as_u64().expect("utility call id"), + fields[2].as_str().expect("utility method").to_owned(), + fields[3].clone(), + ) +} + +async fn send_utility_result(push: &mut PushSocket, call_id: u64, result: T) +where + T: serde::Serialize, +{ + send_outputs( + push, + UtilityCallOutput { + output: UtilityOutput { + call_id: call_id.into(), + failure_message: None, + result: Some(UtilityResultEnvelope::without_type_info( + rmpv::ext::to_value(result).expect("encode utility result"), + )), + }, + ..Default::default() + } + .into(), + ) + .await; +} + +fn collective_args(method: &str, kwargs: BTreeMap) -> Value { + rmpv::ext::to_value(( + method, + Option::::None, + Vec::::new(), + kwargs, + )) + .expect("encode collective arguments") +} + #[derive(Clone, Debug)] struct FakeTextBackend; @@ -346,6 +394,44 @@ where ) } +async fn setup_control_service_with_engine_script( + ready: vllm_engine_core_client::protocol::handshake::EngineCoreReadyResponse, + script: F, +) -> (ControlServiceImpl, MockEngineTask) +where + F: for<'a> FnOnce(&'a mut DealerSocket, &'a mut PushSocket) -> TestFuture<'a> + Send + 'static, +{ + let ipc = IpcNamespace::new().expect("create ipc namespace"); + let handshake_address = ipc.handshake_endpoint(); + let engine_task = MockEngineTask::new(spawn_mock_engine_task_with_ready( + handshake_address.clone(), + b"engine-grpc-rl".to_vec(), + ready, + script, + )); + let client = EngineCoreClient::connect( + EngineCoreClientConfig::new_single(handshake_address) + .with_model_name("test-model") + .with_local_input_output_addresses( + Some(ipc.input_endpoint()), + Some(ipc.output_endpoint()), + ), + ) + .await + .expect("connect client"); + let chat = ChatLlm::from_shared_backend( + Llm::new(client), + Arc::new(FakeTextBackend) as Arc, + ); + ( + ControlServiceImpl::new(Arc::new(AppState::new( + vec!["test-model".to_string()], + chat, + ))), + engine_task, + ) +} + /// Spin up a plaintext gRPC server backed by a mock engine. Returns the client, /// the gRPC server task, and the mock engine task. async fn grpc_test_server( @@ -1488,6 +1574,11 @@ async fn control_reports_server_and_model_info() { assert_eq!(parallelism.data_parallel_size, 1); assert_eq!(parallelism.data_parallel_rank, 0); assert_eq!(parallelism.decode_context_parallel_size, 1); + let rl = server.rl_capabilities.expect("RL capabilities"); + assert!(!rl.weight_transfer_enabled); + assert!(rl.weight_transfer_backend.is_empty()); + assert!(!rl.sleep_mode_enabled); + assert!(!rl.draft_weight_updates_enabled); let model = client .get_model_info(pb::GetModelInfoRequest {}) @@ -1506,6 +1597,162 @@ async fn control_reports_server_and_model_info() { server_task.abort(); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn control_executes_safe_weight_update_lifecycle() { + let mut ready = default_ready_response(); + ready.weight_transfer_backend = Some("nccl".to_string()); + let (service, engine_task) = setup_control_service_with_engine_script(ready, |dealer, push| { + boxed_test_future(async move { + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "collective_rpc"); + assert_eq!( + args, + collective_args( + "init_weight_transfer_engine", + BTreeMap::from([( + "init_info".to_string(), + serde_json::json!({"master_address": "127.0.0.1"}), + )]), + ) + ); + send_utility_result(push, call_id, vec![()]).await; + + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "pause_scheduler"); + assert_eq!( + args, + Value::Array(vec![Value::from("keep"), Value::from(true)]) + ); + send_utility_result(push, call_id, ()).await; + + for expected_method in [ + "start_weight_update", + "update_weights", + "finish_weight_update", + ] { + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "is_scheduler_paused"); + assert_eq!(args, Value::Array(Vec::new())); + send_utility_result(push, call_id, true).await; + + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "collective_rpc"); + let kwargs = if expected_method == "update_weights" { + BTreeMap::from([( + "update_info".to_string(), + serde_json::json!({"names": ["model.weight"]}), + )]) + } else { + BTreeMap::new() + }; + assert_eq!(args, collective_args(expected_method, kwargs)); + send_utility_result(push, call_id, vec![()]).await; + } + + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "set_weight_version"); + assert_eq!(args, Value::Array(vec![Value::from("step-7")])); + send_utility_result(push, call_id, ()).await; + + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "get_weight_version"); + assert_eq!(args, Value::Array(Vec::new())); + send_utility_result(push, call_id, "step-7").await; + + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "resume_scheduler"); + assert_eq!(args, Value::Array(Vec::new())); + send_utility_result(push, call_id, ()).await; + }) + }) + .await; + + pb::control_server::Control::init_weight_transfer_engine( + &service, + tonic::Request::new(pb::InitWeightTransferEngineRequest { + init_info_json: br#"{"master_address":"127.0.0.1"}"#.to_vec(), + }), + ) + .await + .expect("initialize weight transfer"); + pb::control_server::Control::pause_generation( + &service, + tonic::Request::new(pb::PauseGenerationRequest { + mode: pb::PauseMode::Keep as i32, + clear_cache: None, + }), + ) + .await + .expect("pause generation"); + pb::control_server::Control::start_weight_update( + &service, + tonic::Request::new(pb::StartWeightUpdateRequest {}), + ) + .await + .expect("start weight update"); + pb::control_server::Control::update_weights( + &service, + tonic::Request::new(pb::UpdateWeightsRequest { + update_info_json: br#"{"names":["model.weight"]}"#.to_vec(), + }), + ) + .await + .expect("update weights"); + pb::control_server::Control::finish_weight_update( + &service, + tonic::Request::new(pb::FinishWeightUpdateRequest { + weight_version: Some("step-7".to_string()), + }), + ) + .await + .expect("finish weight update"); + let version = pb::control_server::Control::get_weight_version( + &service, + tonic::Request::new(pb::GetWeightVersionRequest {}), + ) + .await + .expect("get weight version") + .into_inner(); + assert_eq!(version.weight_version, "step-7"); + pb::control_server::Control::resume_generation( + &service, + tonic::Request::new(pb::ResumeGenerationRequest {}), + ) + .await + .expect("resume generation"); + + engine_task.await.expect("mock engine task"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn control_rejects_weight_update_while_serving() { + let mut ready = default_ready_response(); + ready.weight_transfer_backend = Some("nccl".to_string()); + let (service, engine_task) = setup_control_service_with_engine_script(ready, |dealer, push| { + boxed_test_future(async move { + let (call_id, method, args) = recv_utility_call(dealer).await; + assert_eq!(method, "is_scheduler_paused"); + assert_eq!(args, Value::Array(Vec::new())); + send_utility_result(push, call_id, false).await; + }) + }) + .await; + + let error = pb::control_server::Control::update_weights( + &service, + tonic::Request::new(pb::UpdateWeightsRequest { + update_info_json: br#"{"names":["model.weight"]}"#.to_vec(), + }), + ) + .await + .expect_err("unpaused weight update should fail"); + assert_eq!(error.code(), tonic::Code::FailedPrecondition); + + engine_task.await.expect("mock engine task"); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn control_aggregates_multi_engine_capacity() { let ipc = IpcNamespace::new().expect("create ipc namespace"); diff --git a/rust/src/server/src/lib.rs b/rust/src/server/src/lib.rs index c36aaac2a19b..f9f5c423506f 100644 --- a/rust/src/server/src/lib.rs +++ b/rust/src/server/src/lib.rs @@ -216,7 +216,8 @@ where health_reporter.set_serving::().await; health_reporter.set_serving::().await; let control_service = - grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone())); + grpc::ControlGrpcService::new(grpc::ControlServiceImpl::new(state.clone())) + .max_decoding_message_size(DEFAULT_REQUEST_BODY_LIMIT_BYTES); let inference_service = grpc::InferenceGrpcService::new(grpc::InferenceServiceImpl::new(state.clone())) .max_decoding_message_size(DEFAULT_REQUEST_BODY_LIMIT_BYTES); diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 0197438e2382..386b9a0bf224 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -92,6 +92,9 @@ class EngineCoreReadyResponse: kv_cache_size_tokens: int | None = None kv_cache_max_concurrency: float | None = None kv_events_config: KVEventsConfig | None = None + weight_transfer_backend: str | None = None + enable_sleep_mode: bool = False + supports_draft_weight_updates: bool = False class EngineCoreRequest( diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 4e7f25aee5f6..480963cb2770 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -1634,6 +1634,13 @@ def _make_ready_response(self) -> EngineCoreReadyResponse: max_num_batched_tokens=scheduler_config.max_num_batched_tokens, instance_id=self.vllm_config.instance_id, kv_events_config=self.scheduler.get_kv_event_publisher_config(), + weight_transfer_backend=( + self.vllm_config.weight_transfer_config.backend + if self.vllm_config.weight_transfer_config is not None + else None + ), + enable_sleep_mode=self.vllm_config.model_config.enable_sleep_mode, + supports_draft_weight_updates=self.use_spec_decode, ) def process_input_sockets( From 4ac45f637c8cd1564b1312f1bd257e79384cc972 Mon Sep 17 00:00:00 2001 From: Connor Carpenter Date: Thu, 6 Aug 2026 13:25:56 -0700 Subject: [PATCH 2/3] test(grpc): focus RL control coverage Signed-off-by: Connor Carpenter --- .../engine-core-client/src/tests/client.rs | 6 + rust/src/server/src/grpc/tests.rs | 147 +----------------- 2 files changed, 14 insertions(+), 139 deletions(-) diff --git a/rust/src/engine-core-client/src/tests/client.rs b/rust/src/engine-core-client/src/tests/client.rs index 3268190785b5..32a36efb3c25 100644 --- a/rust/src/engine-core-client/src/tests/client.rs +++ b/rust/src/engine-core-client/src/tests/client.rs @@ -2681,6 +2681,12 @@ fn python_msgpack_fixtures_match_rust_encoding() { let ready_response: EngineCoreReadyResponse = rmp_serde::from_slice(&hex::decode(ready_response_hex).unwrap()).unwrap(); + assert_eq!( + ready_response.weight_transfer_backend.as_deref(), + Some("nccl") + ); + assert!(ready_response.enable_sleep_mode); + assert!(ready_response.supports_draft_weight_updates); let kv_events_config = ready_response.kv_events_config.expect("KV events config should decode"); assert!(kv_events_config.enable_kv_cache_events); assert_eq!(kv_events_config.publisher, "zmq"); diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index ef06a976ca36..465887ebfb44 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -1,7 +1,6 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright contributors to the vLLM project -use std::collections::BTreeMap; use std::fs; use std::future::Future; use std::io; @@ -201,16 +200,6 @@ where .await; } -fn collective_args(method: &str, kwargs: BTreeMap) -> Value { - rmpv::ext::to_value(( - method, - Option::::None, - Vec::::new(), - kwargs, - )) - .expect("encode collective arguments") -} - #[derive(Clone, Debug)] struct FakeTextBackend; @@ -1597,134 +1586,6 @@ async fn control_reports_server_and_model_info() { server_task.abort(); } -#[tokio::test(flavor = "multi_thread", worker_threads = 2)] -#[serial] -async fn control_executes_safe_weight_update_lifecycle() { - let mut ready = default_ready_response(); - ready.weight_transfer_backend = Some("nccl".to_string()); - let (service, engine_task) = setup_control_service_with_engine_script(ready, |dealer, push| { - boxed_test_future(async move { - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "collective_rpc"); - assert_eq!( - args, - collective_args( - "init_weight_transfer_engine", - BTreeMap::from([( - "init_info".to_string(), - serde_json::json!({"master_address": "127.0.0.1"}), - )]), - ) - ); - send_utility_result(push, call_id, vec![()]).await; - - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "pause_scheduler"); - assert_eq!( - args, - Value::Array(vec![Value::from("keep"), Value::from(true)]) - ); - send_utility_result(push, call_id, ()).await; - - for expected_method in [ - "start_weight_update", - "update_weights", - "finish_weight_update", - ] { - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "is_scheduler_paused"); - assert_eq!(args, Value::Array(Vec::new())); - send_utility_result(push, call_id, true).await; - - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "collective_rpc"); - let kwargs = if expected_method == "update_weights" { - BTreeMap::from([( - "update_info".to_string(), - serde_json::json!({"names": ["model.weight"]}), - )]) - } else { - BTreeMap::new() - }; - assert_eq!(args, collective_args(expected_method, kwargs)); - send_utility_result(push, call_id, vec![()]).await; - } - - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "set_weight_version"); - assert_eq!(args, Value::Array(vec![Value::from("step-7")])); - send_utility_result(push, call_id, ()).await; - - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "get_weight_version"); - assert_eq!(args, Value::Array(Vec::new())); - send_utility_result(push, call_id, "step-7").await; - - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "resume_scheduler"); - assert_eq!(args, Value::Array(Vec::new())); - send_utility_result(push, call_id, ()).await; - }) - }) - .await; - - pb::control_server::Control::init_weight_transfer_engine( - &service, - tonic::Request::new(pb::InitWeightTransferEngineRequest { - init_info_json: br#"{"master_address":"127.0.0.1"}"#.to_vec(), - }), - ) - .await - .expect("initialize weight transfer"); - pb::control_server::Control::pause_generation( - &service, - tonic::Request::new(pb::PauseGenerationRequest { - mode: pb::PauseMode::Keep as i32, - clear_cache: None, - }), - ) - .await - .expect("pause generation"); - pb::control_server::Control::start_weight_update( - &service, - tonic::Request::new(pb::StartWeightUpdateRequest {}), - ) - .await - .expect("start weight update"); - pb::control_server::Control::update_weights( - &service, - tonic::Request::new(pb::UpdateWeightsRequest { - update_info_json: br#"{"names":["model.weight"]}"#.to_vec(), - }), - ) - .await - .expect("update weights"); - pb::control_server::Control::finish_weight_update( - &service, - tonic::Request::new(pb::FinishWeightUpdateRequest { - weight_version: Some("step-7".to_string()), - }), - ) - .await - .expect("finish weight update"); - let version = pb::control_server::Control::get_weight_version( - &service, - tonic::Request::new(pb::GetWeightVersionRequest {}), - ) - .await - .expect("get weight version") - .into_inner(); - assert_eq!(version.weight_version, "step-7"); - pb::control_server::Control::resume_generation( - &service, - tonic::Request::new(pb::ResumeGenerationRequest {}), - ) - .await - .expect("resume generation"); - - engine_task.await.expect("mock engine task"); -} - #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn control_rejects_weight_update_while_serving() { @@ -1762,6 +1623,9 @@ async fn control_aggregates_multi_engine_capacity() { ready_0.max_model_len = 8_192; ready_0.num_gpu_blocks = 10; ready_0.data_parallel_size = 2; + ready_0.weight_transfer_backend = Some("nccl".to_string()); + ready_0.enable_sleep_mode = true; + ready_0.supports_draft_weight_updates = true; let mut ready_1 = default_ready_response(); ready_1.max_model_len = 4_096; @@ -1812,6 +1676,11 @@ async fn control_aggregates_multi_engine_capacity() { .into_inner(); assert_eq!(server.max_model_len, 4_096); assert_eq!(server.total_kv_blocks, 30); + let rl = server.rl_capabilities.expect("RL capabilities"); + assert!(!rl.weight_transfer_enabled); + assert!(rl.weight_transfer_backend.is_empty()); + assert!(!rl.sleep_mode_enabled); + assert!(!rl.draft_weight_updates_enabled); drop(engine_tasks); } From bb0fbd7ab940d61b6b5907fe9edba033838a59f5 Mon Sep 17 00:00:00 2001 From: Connor Carpenter Date: Thu, 6 Aug 2026 13:35:47 -0700 Subject: [PATCH 3/3] test(grpc): reuse scripted service setup Signed-off-by: Connor Carpenter --- rust/src/server/src/grpc/tests.rs | 165 +++++++++++++----------------- 1 file changed, 73 insertions(+), 92 deletions(-) diff --git a/rust/src/server/src/grpc/tests.rs b/rust/src/server/src/grpc/tests.rs index 465887ebfb44..9cce45d45a14 100644 --- a/rust/src/server/src/grpc/tests.rs +++ b/rust/src/server/src/grpc/tests.rs @@ -32,16 +32,14 @@ use vllm_engine_core_client::mock_engine::{ default_ready_response, }; use vllm_engine_core_client::protocol::decode_value; -use vllm_engine_core_client::protocol::handshake::KvEventsConfig; +use vllm_engine_core_client::protocol::handshake::{EngineCoreReadyResponse, KvEventsConfig}; use vllm_engine_core_client::protocol::output::{ EngineCoreFinishReason, EngineCoreOutput, EngineCoreOutputs, RequestBatchOutputs, UtilityCallOutput, }; use vllm_engine_core_client::protocol::request::EngineCoreRequest; use vllm_engine_core_client::protocol::utility::{UtilityOutput, UtilityResultEnvelope}; -use vllm_engine_core_client::test_utils::{ - IpcNamespace, spawn_mock_engine_task, spawn_mock_engine_task_with_ready, -}; +use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task_with_ready}; use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId, TransportMode}; use vllm_llm::Llm; use vllm_text::tokenizer::DynTokenizer; @@ -167,39 +165,6 @@ async fn recv_engine_message(dealer: &mut DealerSocket) -> Vec { dealer.recv().await.expect("recv engine message").into_vec() } -async fn recv_utility_call(dealer: &mut DealerSocket) -> (u64, String, Value) { - let frames = recv_engine_message(dealer).await; - assert_eq!(frames[0].as_ref(), &[0x03]); - let payload = decode_value(&frames[1]).expect("decode utility payload"); - let fields = payload.as_array().expect("utility payload array"); - ( - fields[1].as_u64().expect("utility call id"), - fields[2].as_str().expect("utility method").to_owned(), - fields[3].clone(), - ) -} - -async fn send_utility_result(push: &mut PushSocket, call_id: u64, result: T) -where - T: serde::Serialize, -{ - send_outputs( - push, - UtilityCallOutput { - output: UtilityOutput { - call_id: call_id.into(), - failure_message: None, - result: Some(UtilityResultEnvelope::without_type_info( - rmpv::ext::to_value(result).expect("encode utility result"), - )), - }, - ..Default::default() - } - .into(), - ) - .await; -} - #[derive(Clone, Debug)] struct FakeTextBackend; @@ -338,13 +303,10 @@ async fn setup_grpc_service_with_backend( where F: FnOnce(&EngineCoreRequest) + Send + 'static, { - let ipc = IpcNamespace::new().expect("create ipc namespace"); - let handshake_address = ipc.handshake_endpoint(); - let engine_id = engine_id.into(); - - let engine_task = MockEngineTask::new(spawn_mock_engine_task( - handshake_address.clone(), - engine_id.clone(), + setup_grpc_service_with_engine_script( + engine_id, + default_ready_response(), + backend, move |dealer, push| { boxed_test_future(async move { let add = recv_engine_message(dealer).await; @@ -358,46 +320,35 @@ where .await; }) }, - )); - - let client = EngineCoreClient::connect( - EngineCoreClientConfig::new_single(handshake_address) - .with_model_name("test-model") - .with_local_input_output_addresses( - Some(ipc.input_endpoint()), - Some(ipc.output_endpoint()), - ), ) .await - .expect("connect client"); - let engine_health = client.subscribe_health(); - - let chat = ChatLlm::from_shared_backend(Llm::new(client), backend); - let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat)); - ( - InferenceServer::new(InferenceServiceImpl::new(state.clone())) - .max_decoding_message_size(crate::DEFAULT_REQUEST_BODY_LIMIT_BYTES), - ControlServer::new(ControlServiceImpl::new(state)), - engine_health, - engine_task, - ) } -async fn setup_control_service_with_engine_script( - ready: vllm_engine_core_client::protocol::handshake::EngineCoreReadyResponse, +async fn setup_grpc_service_with_engine_script( + engine_id: impl Into, + ready: EngineCoreReadyResponse, + backend: Arc, script: F, -) -> (ControlServiceImpl, MockEngineTask) +) -> ( + InferenceServer, + ControlServer, + tokio::sync::watch::Receiver, + MockEngineTask, +) where F: for<'a> FnOnce(&'a mut DealerSocket, &'a mut PushSocket) -> TestFuture<'a> + Send + 'static, { let ipc = IpcNamespace::new().expect("create ipc namespace"); let handshake_address = ipc.handshake_endpoint(); + let engine_id = engine_id.into(); + let engine_task = MockEngineTask::new(spawn_mock_engine_task_with_ready( handshake_address.clone(), - b"engine-grpc-rl".to_vec(), + engine_id.clone(), ready, script, )); + let client = EngineCoreClient::connect( EngineCoreClientConfig::new_single(handshake_address) .with_model_name("test-model") @@ -408,15 +359,15 @@ where ) .await .expect("connect client"); - let chat = ChatLlm::from_shared_backend( - Llm::new(client), - Arc::new(FakeTextBackend) as Arc, - ); + let engine_health = client.subscribe_health(); + + let chat = ChatLlm::from_shared_backend(Llm::new(client), backend); + let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat)); ( - ControlServiceImpl::new(Arc::new(AppState::new( - vec!["test-model".to_string()], - chat, - ))), + InferenceServer::new(InferenceServiceImpl::new(state.clone())) + .max_decoding_message_size(crate::DEFAULT_REQUEST_BODY_LIMIT_BYTES), + ControlServer::new(ControlServiceImpl::new(state)), + engine_health, engine_task, ) } @@ -1591,27 +1542,57 @@ async fn control_reports_server_and_model_info() { async fn control_rejects_weight_update_while_serving() { let mut ready = default_ready_response(); ready.weight_transfer_backend = Some("nccl".to_string()); - let (service, engine_task) = setup_control_service_with_engine_script(ready, |dealer, push| { - boxed_test_future(async move { - let (call_id, method, args) = recv_utility_call(dealer).await; - assert_eq!(method, "is_scheduler_paused"); - assert_eq!(args, Value::Array(Vec::new())); - send_utility_result(push, call_id, false).await; - }) - }) + let (inference_service, control_service, engine_health, engine_task) = + setup_grpc_service_with_engine_script( + b"engine-grpc-rl".to_vec(), + ready, + Arc::new(FakeTextBackend), + |dealer, push| { + boxed_test_future(async move { + let frames = recv_engine_message(dealer).await; + assert_eq!(frames[0].as_ref(), &[0x03]); + let payload = decode_value(&frames[1]).expect("decode utility payload"); + let fields = payload.as_array().expect("utility payload array"); + let call_id = fields[1].as_u64().expect("utility call id"); + assert_eq!(fields[2].as_str(), Some("is_scheduler_paused")); + assert_eq!(fields[3], Value::Array(Vec::new())); + send_outputs( + push, + UtilityCallOutput { + output: UtilityOutput { + call_id: call_id.into(), + failure_message: None, + result: Some(UtilityResultEnvelope::without_type_info( + Value::Boolean(false), + )), + }, + ..Default::default() + } + .into(), + ) + .await; + }) + }, + ) + .await; + let (channel, server_task) = start_grpc_test_server( + inference_service, + control_service, + engine_health, + tokio_util::sync::CancellationToken::new(), + ) .await; - let error = pb::control_server::Control::update_weights( - &service, - tonic::Request::new(pb::UpdateWeightsRequest { + let error = ControlClient::new(channel) + .update_weights(pb::UpdateWeightsRequest { update_info_json: br#"{"names":["model.weight"]}"#.to_vec(), - }), - ) - .await - .expect_err("unpaused weight update should fail"); + }) + .await + .expect_err("unpaused weight update should fail"); assert_eq!(error.code(), tonic::Code::FailedPrecondition); engine_task.await.expect("mock engine task"); + server_task.abort(); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)]