[router][grpc] Support vllm backend for grpc router (#13120)
This commit is contained in:
+13
-4
@@ -1,18 +1,27 @@
|
|||||||
fn main() -> Result<(), Box<dyn std::error::Error>> {
|
fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||||
// Only regenerate if the proto file changes
|
// Only regenerate if proto files change
|
||||||
println!("cargo:rerun-if-changed=src/proto/sglang_scheduler.proto");
|
println!("cargo:rerun-if-changed=src/proto/sglang_scheduler.proto");
|
||||||
|
println!("cargo:rerun-if-changed=src/proto/vllm_engine.proto");
|
||||||
|
|
||||||
// Configure tonic-prost-build for gRPC code generation
|
// Configure tonic-prost-build for gRPC code generation
|
||||||
tonic_prost_build::configure()
|
tonic_prost_build::configure()
|
||||||
// Generate both client and server code
|
// Generate both client and server code
|
||||||
.build_server(true)
|
.build_server(true)
|
||||||
.build_client(true)
|
.build_client(true)
|
||||||
|
// Add serde Serialize for model info messages (we only need to serialize to labels)
|
||||||
|
.type_attribute("GetModelInfoResponse", "#[derive(serde::Serialize)]")
|
||||||
// Allow proto3 optional fields
|
// Allow proto3 optional fields
|
||||||
.protoc_arg("--experimental_allow_proto3_optional")
|
.protoc_arg("--experimental_allow_proto3_optional")
|
||||||
// Compile the proto file
|
// Compile both proto files
|
||||||
.compile_protos(&["src/proto/sglang_scheduler.proto"], &["src/proto"])?;
|
.compile_protos(
|
||||||
|
&[
|
||||||
|
"src/proto/sglang_scheduler.proto",
|
||||||
|
"src/proto/vllm_engine.proto",
|
||||||
|
],
|
||||||
|
&["src/proto"],
|
||||||
|
)?;
|
||||||
|
|
||||||
println!("cargo:warning=Protobuf compilation completed successfully");
|
println!("cargo:info=Protobuf compilation completed successfully");
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -396,6 +396,7 @@ impl JobQueue {
|
|||||||
model_id: None,
|
model_id: None,
|
||||||
priority: None,
|
priority: None,
|
||||||
cost: None,
|
cost: None,
|
||||||
|
runtime: None,
|
||||||
tokenizer_path: None,
|
tokenizer_path: None,
|
||||||
reasoning_parser: None,
|
reasoning_parser: None,
|
||||||
tool_parser: None,
|
tool_parser: None,
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ pub use job_queue::{Job, JobQueue, JobQueueConfig};
|
|||||||
pub use retry::{is_retryable_status, BackoffCalculator, RetryError, RetryExecutor};
|
pub use retry::{is_retryable_status, BackoffCalculator, RetryError, RetryExecutor};
|
||||||
pub use worker::{
|
pub use worker::{
|
||||||
worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, HealthChecker, HealthConfig,
|
worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, HealthChecker, HealthConfig,
|
||||||
Worker, WorkerFactory, WorkerLoadGuard, WorkerType,
|
RuntimeType, Worker, WorkerFactory, WorkerLoadGuard, WorkerType,
|
||||||
};
|
};
|
||||||
pub use worker_builder::{BasicWorkerBuilder, DPAwareWorkerBuilder};
|
pub use worker_builder::{BasicWorkerBuilder, DPAwareWorkerBuilder};
|
||||||
pub use worker_manager::{LoadMonitor, WorkerManager};
|
pub use worker_manager::{LoadMonitor, WorkerManager};
|
||||||
|
|||||||
@@ -15,9 +15,9 @@ use tokio::{sync::RwLock, time};
|
|||||||
use super::{CircuitBreaker, WorkerError, WorkerResult};
|
use super::{CircuitBreaker, WorkerError, WorkerResult};
|
||||||
use crate::{
|
use crate::{
|
||||||
core::{BasicWorkerBuilder, CircuitState, DPAwareWorkerBuilder},
|
core::{BasicWorkerBuilder, CircuitState, DPAwareWorkerBuilder},
|
||||||
grpc_client::SglangSchedulerClient,
|
|
||||||
metrics::RouterMetrics,
|
metrics::RouterMetrics,
|
||||||
protocols::worker_spec::WorkerInfo,
|
protocols::worker_spec::WorkerInfo,
|
||||||
|
routers::grpc::client::GrpcClient,
|
||||||
};
|
};
|
||||||
|
|
||||||
static WORKER_CLIENT: LazyLock<reqwest::Client> = LazyLock::new(|| {
|
static WORKER_CLIENT: LazyLock<reqwest::Client> = LazyLock::new(|| {
|
||||||
@@ -229,7 +229,7 @@ pub trait Worker: Send + Sync + fmt::Debug {
|
|||||||
|
|
||||||
/// Get or create a gRPC client for this worker
|
/// Get or create a gRPC client for this worker
|
||||||
/// Returns None for HTTP workers, Some(client) for gRPC workers
|
/// Returns None for HTTP workers, Some(client) for gRPC workers
|
||||||
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<SglangSchedulerClient>>>;
|
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>>;
|
||||||
|
|
||||||
/// Reset the gRPC client connection (for reconnection scenarios)
|
/// Reset the gRPC client connection (for reconnection scenarios)
|
||||||
/// No-op for HTTP workers
|
/// No-op for HTTP workers
|
||||||
@@ -282,6 +282,38 @@ impl fmt::Display for ConnectionMode {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Runtime implementation type for gRPC workers
|
||||||
|
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
|
||||||
|
#[serde(rename_all = "lowercase")]
|
||||||
|
pub enum RuntimeType {
|
||||||
|
/// SGLang runtime (default)
|
||||||
|
#[default]
|
||||||
|
Sglang,
|
||||||
|
/// vLLM runtime
|
||||||
|
Vllm,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl fmt::Display for RuntimeType {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
match self {
|
||||||
|
RuntimeType::Sglang => write!(f, "sglang"),
|
||||||
|
RuntimeType::Vllm => write!(f, "vllm"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl std::str::FromStr for RuntimeType {
|
||||||
|
type Err = String;
|
||||||
|
|
||||||
|
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||||
|
match s.to_lowercase().as_str() {
|
||||||
|
"sglang" => Ok(RuntimeType::Sglang),
|
||||||
|
"vllm" => Ok(RuntimeType::Vllm),
|
||||||
|
_ => Err(format!("Unknown runtime type: {}", s)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Worker type classification
|
/// Worker type classification
|
||||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||||
pub enum WorkerType {
|
pub enum WorkerType {
|
||||||
@@ -345,6 +377,8 @@ pub struct WorkerMetadata {
|
|||||||
pub worker_type: WorkerType,
|
pub worker_type: WorkerType,
|
||||||
/// Connection mode
|
/// Connection mode
|
||||||
pub connection_mode: ConnectionMode,
|
pub connection_mode: ConnectionMode,
|
||||||
|
/// Runtime type (for gRPC workers)
|
||||||
|
pub runtime_type: RuntimeType,
|
||||||
/// Additional labels/tags
|
/// Additional labels/tags
|
||||||
pub labels: std::collections::HashMap<String, String>,
|
pub labels: std::collections::HashMap<String, String>,
|
||||||
/// Health check configuration
|
/// Health check configuration
|
||||||
@@ -368,7 +402,7 @@ pub struct BasicWorker {
|
|||||||
pub consecutive_successes: Arc<AtomicUsize>,
|
pub consecutive_successes: Arc<AtomicUsize>,
|
||||||
pub circuit_breaker: CircuitBreaker,
|
pub circuit_breaker: CircuitBreaker,
|
||||||
/// Lazily initialized gRPC client for gRPC workers
|
/// Lazily initialized gRPC client for gRPC workers
|
||||||
pub grpc_client: Arc<RwLock<Option<Arc<SglangSchedulerClient>>>>,
|
pub grpc_client: Arc<RwLock<Option<Arc<GrpcClient>>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl fmt::Debug for BasicWorker {
|
impl fmt::Debug for BasicWorker {
|
||||||
@@ -506,7 +540,7 @@ impl Worker for BasicWorker {
|
|||||||
&self.circuit_breaker
|
&self.circuit_breaker
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<SglangSchedulerClient>>> {
|
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>> {
|
||||||
match self.metadata.connection_mode {
|
match self.metadata.connection_mode {
|
||||||
ConnectionMode::Http => Ok(None),
|
ConnectionMode::Http => Ok(None),
|
||||||
ConnectionMode::Grpc { .. } => {
|
ConnectionMode::Grpc { .. } => {
|
||||||
@@ -523,16 +557,19 @@ impl Worker for BasicWorker {
|
|||||||
return Ok(Some(client.clone()));
|
return Ok(Some(client.clone()));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let runtime_str = self.metadata.runtime_type.to_string();
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Lazily initializing gRPC client for worker: {}",
|
"Lazily initializing gRPC client ({}) for worker: {}",
|
||||||
|
runtime_str,
|
||||||
self.metadata.url
|
self.metadata.url
|
||||||
);
|
);
|
||||||
match SglangSchedulerClient::connect(&self.metadata.url).await {
|
match GrpcClient::connect(&self.metadata.url, &runtime_str).await {
|
||||||
Ok(client) => {
|
Ok(client) => {
|
||||||
let client_arc = Arc::new(client);
|
let client_arc = Arc::new(client);
|
||||||
*client_guard = Some(client_arc.clone());
|
*client_guard = Some(client_arc.clone());
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Successfully connected gRPC client for worker: {}",
|
"Successfully connected gRPC client ({}) for worker: {}",
|
||||||
|
runtime_str,
|
||||||
self.metadata.url
|
self.metadata.url
|
||||||
);
|
);
|
||||||
Ok(Some(client_arc))
|
Ok(Some(client_arc))
|
||||||
@@ -749,7 +786,7 @@ impl Worker for DPAwareWorker {
|
|||||||
format!("{}{}", self.base_url, route)
|
format!("{}{}", self.base_url, route)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<SglangSchedulerClient>>> {
|
async fn get_grpc_client(&self) -> WorkerResult<Option<Arc<GrpcClient>>> {
|
||||||
self.base_worker.get_grpc_client().await
|
self.base_worker.get_grpc_client().await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -927,6 +964,11 @@ pub fn worker_to_info(worker: &Arc<dyn Worker>) -> WorkerInfo {
|
|||||||
_ => None,
|
_ => None,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let runtime_type = match worker.connection_mode() {
|
||||||
|
ConnectionMode::Grpc { .. } => Some(worker.metadata().runtime_type.to_string()),
|
||||||
|
ConnectionMode::Http => None,
|
||||||
|
};
|
||||||
|
|
||||||
WorkerInfo {
|
WorkerInfo {
|
||||||
id: worker.url().to_string(),
|
id: worker.url().to_string(),
|
||||||
url: worker.url().to_string(),
|
url: worker.url().to_string(),
|
||||||
@@ -937,6 +979,7 @@ pub fn worker_to_info(worker: &Arc<dyn Worker>) -> WorkerInfo {
|
|||||||
is_healthy: worker.is_healthy(),
|
is_healthy: worker.is_healthy(),
|
||||||
load: worker.load(),
|
load: worker.load(),
|
||||||
connection_mode: format!("{:?}", worker.connection_mode()),
|
connection_mode: format!("{:?}", worker.connection_mode()),
|
||||||
|
runtime_type,
|
||||||
tokenizer_path: worker.tokenizer_path().map(String::from),
|
tokenizer_path: worker.tokenizer_path().map(String::from),
|
||||||
reasoning_parser: worker.reasoning_parser().map(String::from),
|
reasoning_parser: worker.reasoning_parser().map(String::from),
|
||||||
tool_parser: worker.tool_parser().map(String::from),
|
tool_parser: worker.tool_parser().map(String::from),
|
||||||
|
|||||||
@@ -3,10 +3,11 @@ use std::collections::HashMap;
|
|||||||
use super::{
|
use super::{
|
||||||
circuit_breaker::{CircuitBreaker, CircuitBreakerConfig},
|
circuit_breaker::{CircuitBreaker, CircuitBreakerConfig},
|
||||||
worker::{
|
worker::{
|
||||||
BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, WorkerMetadata, WorkerType,
|
BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, RuntimeType, WorkerMetadata,
|
||||||
|
WorkerType,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
use crate::grpc_client::SglangSchedulerClient;
|
use crate::routers::grpc::client::GrpcClient;
|
||||||
|
|
||||||
/// Builder for creating BasicWorker instances with fluent API
|
/// Builder for creating BasicWorker instances with fluent API
|
||||||
pub struct BasicWorkerBuilder {
|
pub struct BasicWorkerBuilder {
|
||||||
@@ -14,10 +15,11 @@ pub struct BasicWorkerBuilder {
|
|||||||
api_key: Option<String>,
|
api_key: Option<String>,
|
||||||
worker_type: WorkerType,
|
worker_type: WorkerType,
|
||||||
connection_mode: ConnectionMode,
|
connection_mode: ConnectionMode,
|
||||||
|
runtime_type: RuntimeType,
|
||||||
labels: HashMap<String, String>,
|
labels: HashMap<String, String>,
|
||||||
health_config: HealthConfig,
|
health_config: HealthConfig,
|
||||||
circuit_breaker_config: CircuitBreakerConfig,
|
circuit_breaker_config: CircuitBreakerConfig,
|
||||||
grpc_client: Option<SglangSchedulerClient>,
|
grpc_client: Option<GrpcClient>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl BasicWorkerBuilder {
|
impl BasicWorkerBuilder {
|
||||||
@@ -28,6 +30,7 @@ impl BasicWorkerBuilder {
|
|||||||
api_key: None,
|
api_key: None,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
@@ -42,6 +45,7 @@ impl BasicWorkerBuilder {
|
|||||||
api_key: None,
|
api_key: None,
|
||||||
worker_type,
|
worker_type,
|
||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
@@ -67,6 +71,12 @@ impl BasicWorkerBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set the runtime type (SGLang or vLLM)
|
||||||
|
pub fn runtime_type(mut self, runtime_type: RuntimeType) -> Self {
|
||||||
|
self.runtime_type = runtime_type;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Set labels for worker identification
|
/// Set labels for worker identification
|
||||||
pub fn labels(mut self, labels: HashMap<String, String>) -> Self {
|
pub fn labels(mut self, labels: HashMap<String, String>) -> Self {
|
||||||
self.labels = labels;
|
self.labels = labels;
|
||||||
@@ -92,7 +102,7 @@ impl BasicWorkerBuilder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Set gRPC client for gRPC workers
|
/// Set gRPC client for gRPC workers
|
||||||
pub fn grpc_client(mut self, client: SglangSchedulerClient) -> Self {
|
pub fn grpc_client(mut self, client: GrpcClient) -> Self {
|
||||||
self.grpc_client = Some(client);
|
self.grpc_client = Some(client);
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
@@ -139,6 +149,7 @@ impl BasicWorkerBuilder {
|
|||||||
api_key: self.api_key,
|
api_key: self.api_key,
|
||||||
worker_type: self.worker_type,
|
worker_type: self.worker_type,
|
||||||
connection_mode: self.connection_mode,
|
connection_mode: self.connection_mode,
|
||||||
|
runtime_type: self.runtime_type,
|
||||||
labels: self.labels,
|
labels: self.labels,
|
||||||
health_config: self.health_config,
|
health_config: self.health_config,
|
||||||
bootstrap_host,
|
bootstrap_host,
|
||||||
@@ -168,10 +179,11 @@ pub struct DPAwareWorkerBuilder {
|
|||||||
dp_size: usize,
|
dp_size: usize,
|
||||||
worker_type: WorkerType,
|
worker_type: WorkerType,
|
||||||
connection_mode: ConnectionMode,
|
connection_mode: ConnectionMode,
|
||||||
|
runtime_type: RuntimeType,
|
||||||
labels: HashMap<String, String>,
|
labels: HashMap<String, String>,
|
||||||
health_config: HealthConfig,
|
health_config: HealthConfig,
|
||||||
circuit_breaker_config: CircuitBreakerConfig,
|
circuit_breaker_config: CircuitBreakerConfig,
|
||||||
grpc_client: Option<SglangSchedulerClient>,
|
grpc_client: Option<GrpcClient>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl DPAwareWorkerBuilder {
|
impl DPAwareWorkerBuilder {
|
||||||
@@ -184,6 +196,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
dp_size,
|
dp_size,
|
||||||
worker_type: WorkerType::Regular,
|
worker_type: WorkerType::Regular,
|
||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
@@ -205,6 +218,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
dp_size,
|
dp_size,
|
||||||
worker_type,
|
worker_type,
|
||||||
connection_mode: ConnectionMode::Http,
|
connection_mode: ConnectionMode::Http,
|
||||||
|
runtime_type: RuntimeType::default(),
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
health_config: HealthConfig::default(),
|
health_config: HealthConfig::default(),
|
||||||
circuit_breaker_config: CircuitBreakerConfig::default(),
|
circuit_breaker_config: CircuitBreakerConfig::default(),
|
||||||
@@ -230,6 +244,12 @@ impl DPAwareWorkerBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Set the runtime type (SGLang or vLLM)
|
||||||
|
pub fn runtime_type(mut self, runtime_type: RuntimeType) -> Self {
|
||||||
|
self.runtime_type = runtime_type;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Set labels for worker identification
|
/// Set labels for worker identification
|
||||||
pub fn labels(mut self, labels: HashMap<String, String>) -> Self {
|
pub fn labels(mut self, labels: HashMap<String, String>) -> Self {
|
||||||
self.labels = labels;
|
self.labels = labels;
|
||||||
@@ -255,7 +275,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Set gRPC client for gRPC workers
|
/// Set gRPC client for gRPC workers
|
||||||
pub fn grpc_client(mut self, client: SglangSchedulerClient) -> Self {
|
pub fn grpc_client(mut self, client: GrpcClient) -> Self {
|
||||||
self.grpc_client = Some(client);
|
self.grpc_client = Some(client);
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
@@ -266,6 +286,7 @@ impl DPAwareWorkerBuilder {
|
|||||||
let mut builder = BasicWorkerBuilder::new(worker_url)
|
let mut builder = BasicWorkerBuilder::new(worker_url)
|
||||||
.worker_type(self.worker_type)
|
.worker_type(self.worker_type)
|
||||||
.connection_mode(self.connection_mode)
|
.connection_mode(self.connection_mode)
|
||||||
|
.runtime_type(self.runtime_type)
|
||||||
.labels(self.labels)
|
.labels(self.labels)
|
||||||
.health_config(self.health_config)
|
.health_config(self.health_config)
|
||||||
.circuit_breaker_config(self.circuit_breaker_config);
|
.circuit_breaker_config(self.circuit_breaker_config);
|
||||||
|
|||||||
@@ -24,10 +24,10 @@ use crate::{
|
|||||||
app_context::AppContext,
|
app_context::AppContext,
|
||||||
core::{
|
core::{
|
||||||
workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode,
|
workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode,
|
||||||
DPAwareWorkerBuilder, HealthConfig, Worker, WorkerType,
|
DPAwareWorkerBuilder, HealthConfig, RuntimeType, Worker, WorkerType,
|
||||||
},
|
},
|
||||||
grpc_client::SglangSchedulerClient,
|
|
||||||
protocols::worker_spec::WorkerConfigRequest,
|
protocols::worker_spec::WorkerConfigRequest,
|
||||||
|
routers::grpc::client::GrpcClient,
|
||||||
};
|
};
|
||||||
|
|
||||||
// HTTP client for metadata fetching
|
// HTTP client for metadata fetching
|
||||||
@@ -151,15 +151,13 @@ async fn try_http_health_check(
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Helper: Try gRPC health check
|
/// Helper: Perform gRPC health check with runtime type
|
||||||
async fn try_grpc_health_check(url: &str, timeout_secs: u64) -> Result<(), String> {
|
async fn do_grpc_health_check(
|
||||||
let grpc_url = if url.starts_with("grpc://") {
|
grpc_url: &str,
|
||||||
url.to_string()
|
timeout_secs: u64,
|
||||||
} else {
|
runtime_type: &str,
|
||||||
format!("grpc://{}", strip_protocol(url))
|
) -> Result<(), String> {
|
||||||
};
|
let connect_future = GrpcClient::connect(grpc_url, runtime_type);
|
||||||
|
|
||||||
let connect_future = SglangSchedulerClient::connect(&grpc_url);
|
|
||||||
let client = tokio::time::timeout(Duration::from_secs(timeout_secs), connect_future)
|
let client = tokio::time::timeout(Duration::from_secs(timeout_secs), connect_future)
|
||||||
.await
|
.await
|
||||||
.map_err(|_| "gRPC connection timeout".to_string())?
|
.map_err(|_| "gRPC connection timeout".to_string())?
|
||||||
@@ -174,15 +172,48 @@ async fn try_grpc_health_check(url: &str, timeout_secs: u64) -> Result<(), Strin
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Helper: Fetch gRPC metadata
|
/// Helper: Try gRPC health check
|
||||||
async fn fetch_grpc_metadata(url: &str) -> Result<HashMap<String, String>, String> {
|
///
|
||||||
|
/// If runtime_type is specified, uses the appropriate client (SGLang or vLLM).
|
||||||
|
/// If not specified, tries SGLang first, then falls back to vLLM.
|
||||||
|
async fn try_grpc_health_check(
|
||||||
|
url: &str,
|
||||||
|
timeout_secs: u64,
|
||||||
|
runtime_type: Option<&str>,
|
||||||
|
) -> Result<(), String> {
|
||||||
let grpc_url = if url.starts_with("grpc://") {
|
let grpc_url = if url.starts_with("grpc://") {
|
||||||
url.to_string()
|
url.to_string()
|
||||||
} else {
|
} else {
|
||||||
format!("grpc://{}", strip_protocol(url))
|
format!("grpc://{}", strip_protocol(url))
|
||||||
};
|
};
|
||||||
|
|
||||||
let client = SglangSchedulerClient::connect(&grpc_url)
|
match runtime_type {
|
||||||
|
Some(runtime) => do_grpc_health_check(&grpc_url, timeout_secs, runtime).await,
|
||||||
|
None => {
|
||||||
|
// Runtime not specified: Try SGLang first, then vLLM as fallback
|
||||||
|
if let Ok(()) = do_grpc_health_check(&grpc_url, timeout_secs, "sglang").await {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try vLLM as fallback
|
||||||
|
do_grpc_health_check(&grpc_url, timeout_secs, "vllm")
|
||||||
|
.await
|
||||||
|
.map_err(|e| {
|
||||||
|
format!(
|
||||||
|
"gRPC health check failed (tried both SGLang and vLLM): {}",
|
||||||
|
e
|
||||||
|
)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Fetch metadata from gRPC server with runtime type
|
||||||
|
async fn do_fetch_grpc_metadata(
|
||||||
|
grpc_url: &str,
|
||||||
|
runtime_type: &str,
|
||||||
|
) -> Result<HashMap<String, String>, String> {
|
||||||
|
let client = GrpcClient::connect(grpc_url, runtime_type)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to connect to gRPC: {}", e))?;
|
.map_err(|e| format!("Failed to connect to gRPC: {}", e))?;
|
||||||
|
|
||||||
@@ -191,53 +222,47 @@ async fn fetch_grpc_metadata(url: &str) -> Result<HashMap<String, String>, Strin
|
|||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to fetch gRPC metadata: {}", e))?;
|
.map_err(|e| format!("Failed to fetch gRPC metadata: {}", e))?;
|
||||||
|
|
||||||
let mut labels = HashMap::new();
|
Ok(model_info.to_labels())
|
||||||
|
}
|
||||||
|
|
||||||
// Extract all available fields
|
/// Helper: Fetch gRPC metadata
|
||||||
if !model_info.model_path.is_empty() {
|
///
|
||||||
labels.insert("model_path".to_string(), model_info.model_path.clone());
|
/// If runtime_type is specified, uses the appropriate client (SGLang or vLLM).
|
||||||
}
|
/// If not specified, tries SGLang first, then falls back to vLLM.
|
||||||
if !model_info.tokenizer_path.is_empty() {
|
/// Returns (labels, detected_runtime_type)
|
||||||
labels.insert(
|
async fn fetch_grpc_metadata(
|
||||||
"tokenizer_path".to_string(),
|
url: &str,
|
||||||
model_info.tokenizer_path.clone(),
|
runtime_type: Option<&str>,
|
||||||
);
|
) -> Result<(HashMap<String, String>, String), String> {
|
||||||
}
|
let grpc_url = if url.starts_with("grpc://") {
|
||||||
if !model_info.served_model_name.is_empty() {
|
url.to_string()
|
||||||
labels.insert(
|
} else {
|
||||||
"served_model_name".to_string(),
|
format!("grpc://{}", strip_protocol(url))
|
||||||
model_info.served_model_name.clone(),
|
};
|
||||||
);
|
|
||||||
}
|
|
||||||
if !model_info.weight_version.is_empty() {
|
|
||||||
labels.insert(
|
|
||||||
"weight_version".to_string(),
|
|
||||||
model_info.weight_version.clone(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if !model_info.model_type.is_empty() {
|
|
||||||
labels.insert("model_type".to_string(), model_info.model_type.clone());
|
|
||||||
}
|
|
||||||
if model_info.max_context_length > 0 {
|
|
||||||
labels.insert(
|
|
||||||
"max_context_length".to_string(),
|
|
||||||
model_info.max_context_length.to_string(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if model_info.max_req_input_len > 0 {
|
|
||||||
labels.insert(
|
|
||||||
"max_req_input_len".to_string(),
|
|
||||||
model_info.max_req_input_len.to_string(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if model_info.vocab_size > 0 {
|
|
||||||
labels.insert("vocab_size".to_string(), model_info.vocab_size.to_string());
|
|
||||||
}
|
|
||||||
if model_info.is_generation {
|
|
||||||
labels.insert("is_generation".to_string(), "true".to_string());
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(labels)
|
match runtime_type {
|
||||||
|
Some(runtime) => {
|
||||||
|
let labels = do_fetch_grpc_metadata(&grpc_url, runtime).await?;
|
||||||
|
Ok((labels, runtime.to_string()))
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
// Runtime not specified: Try SGLang first, then vLLM as fallback
|
||||||
|
if let Ok(labels) = do_fetch_grpc_metadata(&grpc_url, "sglang").await {
|
||||||
|
return Ok((labels, "sglang".to_string()));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try vLLM as fallback
|
||||||
|
let labels = do_fetch_grpc_metadata(&grpc_url, "vllm")
|
||||||
|
.await
|
||||||
|
.map_err(|e| {
|
||||||
|
format!(
|
||||||
|
"Failed to fetch gRPC metadata (tried both SGLang and vLLM): {}",
|
||||||
|
e
|
||||||
|
)
|
||||||
|
})?;
|
||||||
|
Ok((labels, "vllm".to_string()))
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Step 1: Detect connection mode by probing both HTTP and gRPC
|
/// Step 1: Detect connection mode by probing both HTTP and gRPC
|
||||||
@@ -263,9 +288,10 @@ impl StepExecutor for DetectConnectionModeStep {
|
|||||||
let url = config.url.clone();
|
let url = config.url.clone();
|
||||||
let timeout = config.health_check_timeout_secs;
|
let timeout = config.health_check_timeout_secs;
|
||||||
let client = &app_context.client;
|
let client = &app_context.client;
|
||||||
|
let runtime_type = config.runtime.as_deref();
|
||||||
let (http_result, grpc_result) = tokio::join!(
|
let (http_result, grpc_result) = tokio::join!(
|
||||||
try_http_health_check(&url, timeout, client),
|
try_http_health_check(&url, timeout, client),
|
||||||
try_grpc_health_check(&url, timeout)
|
try_grpc_health_check(&url, timeout, runtime_type)
|
||||||
);
|
);
|
||||||
|
|
||||||
let connection_mode = match (http_result, grpc_result) {
|
let connection_mode = match (http_result, grpc_result) {
|
||||||
@@ -317,7 +343,7 @@ impl StepExecutor for DiscoverMetadataStep {
|
|||||||
config.url, *connection_mode
|
config.url, *connection_mode
|
||||||
);
|
);
|
||||||
|
|
||||||
let discovered_labels = match connection_mode.as_ref() {
|
let (discovered_labels, detected_runtime) = match connection_mode.as_ref() {
|
||||||
ConnectionMode::Http => {
|
ConnectionMode::Http => {
|
||||||
match get_server_info(&config.url, config.api_key.as_deref()).await {
|
match get_server_info(&config.url, config.api_key.as_deref()).await {
|
||||||
Ok(server_info) => {
|
Ok(server_info) => {
|
||||||
@@ -327,16 +353,21 @@ impl StepExecutor for DiscoverMetadataStep {
|
|||||||
labels.insert("model_path".to_string(), model_path);
|
labels.insert("model_path".to_string(), model_path);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(labels)
|
Ok((labels, None))
|
||||||
}
|
}
|
||||||
Err(e) => Err(e),
|
Err(e) => Err(e),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
ConnectionMode::Grpc { .. } => fetch_grpc_metadata(&config.url).await,
|
ConnectionMode::Grpc { .. } => {
|
||||||
|
let runtime_type = config.runtime.as_deref();
|
||||||
|
fetch_grpc_metadata(&config.url, runtime_type)
|
||||||
|
.await
|
||||||
|
.map(|(labels, runtime)| (labels, Some(runtime)))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
.unwrap_or_else(|e| {
|
.unwrap_or_else(|e| {
|
||||||
warn!("Failed to fetch metadata for {}: {}", config.url, e);
|
warn!("Failed to fetch metadata for {}: {}", config.url, e);
|
||||||
HashMap::new()
|
(HashMap::new(), None)
|
||||||
});
|
});
|
||||||
|
|
||||||
debug!(
|
debug!(
|
||||||
@@ -345,8 +376,12 @@ impl StepExecutor for DiscoverMetadataStep {
|
|||||||
config.url
|
config.url
|
||||||
);
|
);
|
||||||
|
|
||||||
// Store discovered labels in context
|
// Store discovered labels and detected runtime in context
|
||||||
context.set("discovered_labels", discovered_labels);
|
context.set("discovered_labels", discovered_labels);
|
||||||
|
if let Some(runtime) = detected_runtime {
|
||||||
|
debug!("Detected runtime type: {}", runtime);
|
||||||
|
context.set("detected_runtime_type", runtime);
|
||||||
|
}
|
||||||
|
|
||||||
Ok(StepResult::Success)
|
Ok(StepResult::Success)
|
||||||
}
|
}
|
||||||
@@ -496,6 +531,26 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
})
|
})
|
||||||
.unwrap_or(WorkerType::Regular);
|
.unwrap_or(WorkerType::Regular);
|
||||||
|
|
||||||
|
// Get detected runtime type (for gRPC workers)
|
||||||
|
let runtime_type = if matches!(connection_mode.as_ref(), ConnectionMode::Grpc { .. }) {
|
||||||
|
// Try to get detected runtime from context, fall back to config, or default to sglang
|
||||||
|
if let Some(detected_runtime) = context.get::<String>("detected_runtime_type") {
|
||||||
|
match detected_runtime.as_str() {
|
||||||
|
"vllm" => RuntimeType::Vllm,
|
||||||
|
_ => RuntimeType::Sglang,
|
||||||
|
}
|
||||||
|
} else if let Some(ref runtime) = config.runtime {
|
||||||
|
match runtime.as_str() {
|
||||||
|
"vllm" => RuntimeType::Vllm,
|
||||||
|
_ => RuntimeType::Sglang,
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
RuntimeType::Sglang
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
RuntimeType::Sglang // Default for HTTP workers
|
||||||
|
};
|
||||||
|
|
||||||
// Build circuit breaker config
|
// Build circuit breaker config
|
||||||
let circuit_breaker_config = {
|
let circuit_breaker_config = {
|
||||||
let cfg = app_context.router_config.effective_circuit_breaker_config();
|
let cfg = app_context.router_config.effective_circuit_breaker_config();
|
||||||
@@ -561,6 +616,7 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
DPAwareWorkerBuilder::new(normalized_url.clone(), rank, dp_info.dp_size)
|
DPAwareWorkerBuilder::new(normalized_url.clone(), rank, dp_info.dp_size)
|
||||||
.worker_type(worker_type.clone())
|
.worker_type(worker_type.clone())
|
||||||
.connection_mode(connection_mode.as_ref().clone())
|
.connection_mode(connection_mode.as_ref().clone())
|
||||||
|
.runtime_type(runtime_type.clone())
|
||||||
.circuit_breaker_config(circuit_breaker_config.clone())
|
.circuit_breaker_config(circuit_breaker_config.clone())
|
||||||
.health_config(health_config.clone());
|
.health_config(health_config.clone());
|
||||||
|
|
||||||
@@ -595,6 +651,7 @@ impl StepExecutor for CreateWorkerStep {
|
|||||||
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
|
let mut builder = BasicWorkerBuilder::new(normalized_url.clone())
|
||||||
.worker_type(worker_type)
|
.worker_type(worker_type)
|
||||||
.connection_mode(connection_mode.as_ref().clone())
|
.connection_mode(connection_mode.as_ref().clone())
|
||||||
|
.runtime_type(runtime_type)
|
||||||
.circuit_breaker_config(circuit_breaker_config)
|
.circuit_breaker_config(circuit_breaker_config)
|
||||||
.health_config(health_config);
|
.health_config(health_config);
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,7 @@
|
|||||||
pub mod sglang_scheduler;
|
pub mod sglang_scheduler;
|
||||||
|
pub mod vllm_engine;
|
||||||
|
|
||||||
pub use sglang_scheduler::{proto, SglangSchedulerClient};
|
// Export both clients
|
||||||
|
// Re-export proto modules with explicit names
|
||||||
|
pub use sglang_scheduler::{proto as sglang_proto, SglangSchedulerClient};
|
||||||
|
pub use vllm_engine::{proto as vllm_proto, VllmEngineClient};
|
||||||
|
|||||||
@@ -0,0 +1,736 @@
|
|||||||
|
use std::{
|
||||||
|
convert::TryFrom,
|
||||||
|
pin::Pin,
|
||||||
|
sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Arc,
|
||||||
|
},
|
||||||
|
task::{Context, Poll},
|
||||||
|
time::Duration,
|
||||||
|
};
|
||||||
|
|
||||||
|
use tonic::{transport::Channel, Request, Streaming};
|
||||||
|
use tracing::{debug, warn};
|
||||||
|
|
||||||
|
use crate::protocols::{
|
||||||
|
chat::ChatCompletionRequest,
|
||||||
|
common::{ResponseFormat, StringOrArray, ToolChoice, ToolChoiceValue},
|
||||||
|
generate::GenerateRequest,
|
||||||
|
responses::ResponsesRequest,
|
||||||
|
sampling_params::SamplingParams as GenerateSamplingParams,
|
||||||
|
};
|
||||||
|
|
||||||
|
// Include the generated protobuf code
|
||||||
|
#[allow(clippy::all)]
|
||||||
|
pub mod proto {
|
||||||
|
#![allow(clippy::all, unused_qualifications)]
|
||||||
|
tonic::include_proto!("vllm.grpc.engine");
|
||||||
|
}
|
||||||
|
|
||||||
|
// The generated module structure depends on the package name in the .proto file
|
||||||
|
// package vllm.grpc.engine; generates a nested module structure
|
||||||
|
|
||||||
|
/// A smart wrapper around Streaming<GenerateResponse> that automatically
|
||||||
|
/// sends abort when dropped (e.g., due to client disconnection or early termination).
|
||||||
|
///
|
||||||
|
/// This leverages Rust's RAII pattern to ensure cleanup happens automatically,
|
||||||
|
/// regardless of how the stream is dropped (panic, early return, client disconnect, etc.).
|
||||||
|
pub struct AbortOnDropStream {
|
||||||
|
inner: Streaming<proto::GenerateResponse>,
|
||||||
|
request_id: String,
|
||||||
|
client: VllmEngineClient,
|
||||||
|
aborted: Arc<AtomicBool>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl AbortOnDropStream {
|
||||||
|
/// Create a new auto-aborting stream wrapper
|
||||||
|
pub fn new(
|
||||||
|
stream: Streaming<proto::GenerateResponse>,
|
||||||
|
request_id: String,
|
||||||
|
client: VllmEngineClient,
|
||||||
|
) -> Self {
|
||||||
|
debug!("Created AbortOnDropStream for request {}", request_id);
|
||||||
|
Self {
|
||||||
|
inner: stream,
|
||||||
|
request_id,
|
||||||
|
client,
|
||||||
|
aborted: Arc::new(AtomicBool::new(false)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Manually mark the request as completed to prevent abort on drop.
|
||||||
|
/// Call this when the request completes successfully to avoid unnecessary abort RPC.
|
||||||
|
pub fn mark_completed(&self) {
|
||||||
|
// Use Release ordering to ensure that this write is visible to other threads
|
||||||
|
// that use Acquire on the same atomic variable
|
||||||
|
self.aborted.store(true, Ordering::Release);
|
||||||
|
debug!("Request {} marked as completed", self.request_id);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for AbortOnDropStream {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
// Atomically check and set the aborted flag using compare_exchange.
|
||||||
|
// If compare_exchange fails, it means the flag was already true (from mark_completed),
|
||||||
|
// so we don't need to send abort. AcqRel is used for success to synchronize with
|
||||||
|
// mark_completed's Release, and Acquire for failure to see writes from mark_completed.
|
||||||
|
if self
|
||||||
|
.aborted
|
||||||
|
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||||
|
.is_err()
|
||||||
|
{
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
let client = self.client.clone();
|
||||||
|
let request_id = self.request_id.clone();
|
||||||
|
|
||||||
|
// Spawn a background task to send abort (since Drop is sync but abort_request is async)
|
||||||
|
tokio::spawn(async move {
|
||||||
|
debug!(
|
||||||
|
"Stream dropped without completion for request {}, sending abort",
|
||||||
|
request_id
|
||||||
|
);
|
||||||
|
// Clone request_id for the error message since abort_request takes ownership
|
||||||
|
let request_id_for_log = request_id.clone();
|
||||||
|
if let Err(e) = client
|
||||||
|
.abort_request(request_id, "Stream dropped".to_string())
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
warn!(
|
||||||
|
"Failed to send abort on drop for request {}: {}",
|
||||||
|
request_id_for_log, e
|
||||||
|
);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Implement Stream trait to make AbortOnDropStream work like the original Streaming
|
||||||
|
impl futures::Stream for AbortOnDropStream {
|
||||||
|
type Item = Result<proto::GenerateResponse, tonic::Status>;
|
||||||
|
|
||||||
|
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||||
|
// Delegate to the inner stream
|
||||||
|
Pin::new(&mut self.inner).poll_next(cx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// gRPC client for vLLM scheduler
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct VllmEngineClient {
|
||||||
|
client: proto::vllm_engine_client::VllmEngineClient<Channel>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl VllmEngineClient {
|
||||||
|
/// Create a new client and connect to the vLLM server
|
||||||
|
pub async fn connect(endpoint: &str) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
debug!("Connecting to vLLM gRPC server at {}", endpoint);
|
||||||
|
|
||||||
|
// Convert grpc:// to http:// for tonic
|
||||||
|
let http_endpoint = if let Some(addr) = endpoint.strip_prefix("grpc://") {
|
||||||
|
format!("http://{}", addr)
|
||||||
|
} else {
|
||||||
|
endpoint.to_string()
|
||||||
|
};
|
||||||
|
|
||||||
|
let channel = Channel::from_shared(http_endpoint)?
|
||||||
|
.http2_keep_alive_interval(Duration::from_secs(30))
|
||||||
|
.keep_alive_timeout(Duration::from_secs(10))
|
||||||
|
.keep_alive_while_idle(true)
|
||||||
|
.tcp_keepalive(Some(Duration::from_secs(60)))
|
||||||
|
.tcp_nodelay(true)
|
||||||
|
.http2_adaptive_window(true)
|
||||||
|
.initial_stream_window_size(Some(16 * 1024 * 1024)) // 16MB
|
||||||
|
.initial_connection_window_size(Some(32 * 1024 * 1024)) // 32MB
|
||||||
|
.connect()
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
let client = proto::vllm_engine_client::VllmEngineClient::new(channel);
|
||||||
|
|
||||||
|
Ok(Self { client })
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Submit a generation request (returns auto-aborting streaming response)
|
||||||
|
///
|
||||||
|
/// The returned stream automatically sends an abort request when dropped,
|
||||||
|
/// ensuring proper cleanup even if the HTTP client disconnects or an error occurs.
|
||||||
|
/// Call `mark_completed()` on the stream after successful completion to prevent
|
||||||
|
/// unnecessary abort RPCs.
|
||||||
|
pub async fn generate(
|
||||||
|
&self,
|
||||||
|
req: proto::GenerateRequest,
|
||||||
|
) -> Result<AbortOnDropStream, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
let request_id = req.request_id.clone();
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let request = Request::new(req);
|
||||||
|
let response = client.generate(request).await?;
|
||||||
|
|
||||||
|
Ok(AbortOnDropStream::new(
|
||||||
|
response.into_inner(),
|
||||||
|
request_id,
|
||||||
|
self.clone(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Perform health check
|
||||||
|
pub async fn health_check(
|
||||||
|
&self,
|
||||||
|
) -> Result<proto::HealthCheckResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
debug!("Sending health check request");
|
||||||
|
// HealthCheckRequest is now empty - server generates its own health check internally
|
||||||
|
let request = Request::new(proto::HealthCheckRequest {});
|
||||||
|
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let response = client.health_check(request).await?;
|
||||||
|
debug!("Health check response received");
|
||||||
|
Ok(response.into_inner())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Abort a request
|
||||||
|
pub async fn abort_request(
|
||||||
|
&self,
|
||||||
|
request_id: String,
|
||||||
|
reason: String,
|
||||||
|
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
debug!(
|
||||||
|
"Sending abort request for {} (reason: {})",
|
||||||
|
request_id, reason
|
||||||
|
);
|
||||||
|
let request = Request::new(proto::AbortRequest {
|
||||||
|
request_id: request_id.clone(),
|
||||||
|
reason,
|
||||||
|
});
|
||||||
|
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let response = client.abort(request).await?;
|
||||||
|
debug!(
|
||||||
|
"Abort response for {}: success={}, message={}",
|
||||||
|
request_id,
|
||||||
|
response.get_ref().success,
|
||||||
|
response.get_ref().message
|
||||||
|
);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get model information
|
||||||
|
pub async fn get_model_info(
|
||||||
|
&self,
|
||||||
|
) -> Result<proto::GetModelInfoResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
debug!("Requesting model info");
|
||||||
|
let request = Request::new(proto::GetModelInfoRequest {});
|
||||||
|
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let response = client.get_model_info(request).await?;
|
||||||
|
debug!("Model info response received");
|
||||||
|
Ok(response.into_inner())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get server information
|
||||||
|
pub async fn get_server_info(
|
||||||
|
&self,
|
||||||
|
) -> Result<proto::GetServerInfoResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
debug!("Requesting server info");
|
||||||
|
let request = Request::new(proto::GetServerInfoRequest {});
|
||||||
|
|
||||||
|
let mut client = self.client.clone();
|
||||||
|
let response = client.get_server_info(request).await?;
|
||||||
|
debug!("Server info response received");
|
||||||
|
Ok(response.into_inner())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build a single vLLM GenerateRequest from OpenAI ChatCompletionRequest
|
||||||
|
pub fn build_generate_request_from_chat(
|
||||||
|
&self,
|
||||||
|
request_id: String,
|
||||||
|
body: &ChatCompletionRequest,
|
||||||
|
processed_text: String,
|
||||||
|
token_ids: Vec<u32>,
|
||||||
|
tool_call_constraint: Option<(String, String)>, // (constraint_type, constraint_value)
|
||||||
|
) -> Result<proto::GenerateRequest, String> {
|
||||||
|
// Build sampling params
|
||||||
|
let sampling_params =
|
||||||
|
self.build_grpc_sampling_params_from_chat(body, tool_call_constraint)?;
|
||||||
|
|
||||||
|
let grpc_request = proto::GenerateRequest {
|
||||||
|
request_id,
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: processed_text,
|
||||||
|
input_ids: token_ids,
|
||||||
|
}),
|
||||||
|
sampling_params: Some(sampling_params),
|
||||||
|
stream: body.stream,
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(grpc_request)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build a basic GenerateRequest from the vLLM spec GenerateRequest
|
||||||
|
pub fn build_plain_generate_request(
|
||||||
|
&self,
|
||||||
|
request_id: String,
|
||||||
|
body: &GenerateRequest,
|
||||||
|
original_text: Option<String>,
|
||||||
|
token_ids: Vec<u32>,
|
||||||
|
) -> Result<proto::GenerateRequest, String> {
|
||||||
|
let sampling_params =
|
||||||
|
Self::build_sampling_params_from_plain(body.sampling_params.as_ref())?;
|
||||||
|
|
||||||
|
let grpc_request = proto::GenerateRequest {
|
||||||
|
request_id,
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: original_text.unwrap_or_default(),
|
||||||
|
input_ids: token_ids,
|
||||||
|
}),
|
||||||
|
sampling_params: Some(sampling_params),
|
||||||
|
stream: body.stream,
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(grpc_request)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build a GenerateRequest from ResponsesRequest (OpenAI Responses API)
|
||||||
|
///
|
||||||
|
/// NOTE: This is used by the Harmony router only. The Regular router uses
|
||||||
|
/// responses_to_chat() conversion and goes through the chat pipeline.
|
||||||
|
pub fn build_generate_request_from_responses(
|
||||||
|
&self,
|
||||||
|
request_id: String,
|
||||||
|
body: &ResponsesRequest,
|
||||||
|
processed_text: String,
|
||||||
|
token_ids: Vec<u32>,
|
||||||
|
harmony_stop_ids: Option<Vec<u32>>,
|
||||||
|
constraint: Option<(String, String)>,
|
||||||
|
) -> Result<proto::GenerateRequest, String> {
|
||||||
|
// Build sampling params from ResponsesRequest
|
||||||
|
let mut sampling_params =
|
||||||
|
self.build_grpc_sampling_params_from_responses(body, constraint)?;
|
||||||
|
|
||||||
|
// Inject Harmony stop token IDs if provided
|
||||||
|
if let Some(stop_ids) = harmony_stop_ids {
|
||||||
|
sampling_params.stop_token_ids = stop_ids;
|
||||||
|
}
|
||||||
|
|
||||||
|
let grpc_request = proto::GenerateRequest {
|
||||||
|
request_id,
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: processed_text,
|
||||||
|
input_ids: token_ids,
|
||||||
|
}),
|
||||||
|
sampling_params: Some(sampling_params),
|
||||||
|
stream: body.stream.unwrap_or(false),
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(grpc_request)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build gRPC SamplingParams from ChatCompletionRequest
|
||||||
|
fn build_grpc_sampling_params_from_chat(
|
||||||
|
&self,
|
||||||
|
request: &ChatCompletionRequest,
|
||||||
|
tool_call_constraint: Option<(String, String)>,
|
||||||
|
) -> Result<proto::SamplingParams, String> {
|
||||||
|
let stop_sequences = self.extract_stop_strings(request);
|
||||||
|
|
||||||
|
let max_tokens = request.max_completion_tokens.map(|v| v as i32);
|
||||||
|
|
||||||
|
// Handle skip_special_tokens: set to false if tools are present and tool_choice is not "none"
|
||||||
|
let skip_special_tokens = if request.tools.is_some() {
|
||||||
|
match &request.tool_choice {
|
||||||
|
Some(ToolChoice::Value(ToolChoiceValue::None)) => request.skip_special_tokens,
|
||||||
|
Some(_) => false, // tool_choice is not "none"
|
||||||
|
None => false, // TODO: this assumes tool_choice defaults to "auto" when tools present
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
request.skip_special_tokens
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(proto::SamplingParams {
|
||||||
|
temperature: request.temperature.unwrap_or(1.0),
|
||||||
|
top_p: request.top_p.unwrap_or(1.0),
|
||||||
|
top_k: request.top_k.unwrap_or(-1),
|
||||||
|
min_p: request.min_p.unwrap_or(0.0),
|
||||||
|
frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
|
||||||
|
presence_penalty: request.presence_penalty.unwrap_or(0.0),
|
||||||
|
repetition_penalty: request.repetition_penalty.unwrap_or(1.0),
|
||||||
|
max_tokens,
|
||||||
|
stop: stop_sequences,
|
||||||
|
stop_token_ids: request.stop_token_ids.clone().unwrap_or_default(),
|
||||||
|
skip_special_tokens,
|
||||||
|
spaces_between_special_tokens: true, // Default from Python SamplingParams
|
||||||
|
ignore_eos: request.ignore_eos,
|
||||||
|
n: request.n.unwrap_or(1) as i32,
|
||||||
|
constraint: self.build_constraint_for_chat(request, tool_call_constraint)?,
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Extract stop strings from request
|
||||||
|
fn extract_stop_strings(&self, request: &ChatCompletionRequest) -> Vec<String> {
|
||||||
|
match &request.stop {
|
||||||
|
Some(StringOrArray::String(s)) => vec![s.clone()],
|
||||||
|
Some(StringOrArray::Array(arr)) => arr.clone(),
|
||||||
|
None => vec![],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build constraint for structured generation
|
||||||
|
fn build_constraint_for_chat(
|
||||||
|
&self,
|
||||||
|
request: &ChatCompletionRequest,
|
||||||
|
tool_call_constraint: Option<(String, String)>,
|
||||||
|
) -> Result<Option<proto::sampling_params::Constraint>, String> {
|
||||||
|
let mut constraints = Vec::new();
|
||||||
|
|
||||||
|
// Handle response_format constraints
|
||||||
|
match &request.response_format {
|
||||||
|
Some(ResponseFormat::JsonObject) => {
|
||||||
|
// json_object mode - constrain to valid JSON object
|
||||||
|
let schema = serde_json::json!({"type": "object"});
|
||||||
|
let schema_str = serde_json::to_string(&schema)
|
||||||
|
.map_err(|e| format!("Failed to serialize JSON schema: {}", e))?;
|
||||||
|
constraints.push(proto::sampling_params::Constraint::JsonSchema(schema_str));
|
||||||
|
}
|
||||||
|
Some(ResponseFormat::JsonSchema { json_schema }) => {
|
||||||
|
let schema_str = serde_json::to_string(&json_schema.schema)
|
||||||
|
.map_err(|e| format!("Failed to serialize JSON schema: {}", e))?;
|
||||||
|
constraints.push(proto::sampling_params::Constraint::JsonSchema(schema_str));
|
||||||
|
}
|
||||||
|
Some(ResponseFormat::Text) | None => {
|
||||||
|
// No constraint for text format
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// vLLM supports: json_schema, regex, grammar, structural_tag, json_object, choice
|
||||||
|
if let Some(ebnf) = &request.ebnf {
|
||||||
|
constraints.push(proto::sampling_params::Constraint::Grammar(ebnf.clone()));
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(regex) = &request.regex {
|
||||||
|
constraints.push(proto::sampling_params::Constraint::Regex(regex.clone()));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle tool call constraint from preparation stage
|
||||||
|
if let Some((constraint_type, constraint_value)) = tool_call_constraint {
|
||||||
|
if !constraints.is_empty() {
|
||||||
|
return Err("Constrained decoding is not compatible with tool calls.".to_string());
|
||||||
|
}
|
||||||
|
let tool_constraint = match constraint_type.as_str() {
|
||||||
|
"structural_tag" => {
|
||||||
|
proto::sampling_params::Constraint::StructuralTag(constraint_value)
|
||||||
|
}
|
||||||
|
"json_schema" => proto::sampling_params::Constraint::JsonSchema(constraint_value),
|
||||||
|
"grammar" | "ebnf" => proto::sampling_params::Constraint::Grammar(constraint_value),
|
||||||
|
"regex" => proto::sampling_params::Constraint::Regex(constraint_value),
|
||||||
|
_ => return Err(format!("Unknown constraint type: {}", constraint_type)),
|
||||||
|
};
|
||||||
|
constraints.push(tool_constraint);
|
||||||
|
}
|
||||||
|
|
||||||
|
match constraints.len() {
|
||||||
|
0 => Ok(None),
|
||||||
|
1 => Ok(constraints.pop()),
|
||||||
|
_ => Err("Multiple constraints are not allowed.".to_string()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build gRPC SamplingParams from ResponsesRequest
|
||||||
|
fn build_grpc_sampling_params_from_responses(
|
||||||
|
&self,
|
||||||
|
request: &ResponsesRequest,
|
||||||
|
constraint: Option<(String, String)>,
|
||||||
|
) -> Result<proto::SamplingParams, String> {
|
||||||
|
// Used by Harmony models only. Regular models use Chat API path.
|
||||||
|
// Constraints come from Harmony preparation stage (structural_tag) or tool handling.
|
||||||
|
|
||||||
|
let max_tokens = request.max_output_tokens.map(|v| v as i32);
|
||||||
|
|
||||||
|
Ok(proto::SamplingParams {
|
||||||
|
temperature: request.temperature.unwrap_or(1.0),
|
||||||
|
top_p: request.top_p.unwrap_or(1.0),
|
||||||
|
top_k: -1, // ResponsesRequest doesn't expose top_k
|
||||||
|
min_p: 0.0, // ResponsesRequest doesn't expose min_p
|
||||||
|
frequency_penalty: 0.0, // ResponsesRequest doesn't expose frequency_penalty
|
||||||
|
presence_penalty: 0.0, // ResponsesRequest doesn't expose presence_penalty
|
||||||
|
repetition_penalty: 1.0, // ResponsesRequest doesn't expose repetition_penalty
|
||||||
|
max_tokens,
|
||||||
|
stop: vec![], // No stop sequences in Responses API
|
||||||
|
stop_token_ids: vec![], // Handled by Harmony stop tokens
|
||||||
|
skip_special_tokens: false, // Keep special tokens for Harmony
|
||||||
|
spaces_between_special_tokens: true,
|
||||||
|
ignore_eos: false,
|
||||||
|
n: 1, // Responses API doesn't support n>1
|
||||||
|
constraint: self.build_constraint_for_responses(constraint)?,
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build constraint for Responses API
|
||||||
|
///
|
||||||
|
/// Handles constraints from Harmony preparation stage (structural_tag for Harmony models,
|
||||||
|
/// structured output via text field, or tool call constraints).
|
||||||
|
///
|
||||||
|
/// Note: Regular gRPC models use Chat API path with response_format, not this function.
|
||||||
|
fn build_constraint_for_responses(
|
||||||
|
&self,
|
||||||
|
constraint: Option<(String, String)>,
|
||||||
|
) -> Result<Option<proto::sampling_params::Constraint>, String> {
|
||||||
|
if let Some((constraint_type, constraint_value)) = constraint {
|
||||||
|
let parsed_constraint = match constraint_type.as_str() {
|
||||||
|
"structural_tag" => {
|
||||||
|
proto::sampling_params::Constraint::StructuralTag(constraint_value)
|
||||||
|
}
|
||||||
|
"json_schema" => proto::sampling_params::Constraint::JsonSchema(constraint_value),
|
||||||
|
"grammar" | "ebnf" => proto::sampling_params::Constraint::Grammar(constraint_value),
|
||||||
|
"regex" => proto::sampling_params::Constraint::Regex(constraint_value),
|
||||||
|
_ => return Err(format!("Unknown constraint type: {}", constraint_type)),
|
||||||
|
};
|
||||||
|
Ok(Some(parsed_constraint))
|
||||||
|
} else {
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_single_constraint_from_plain(
|
||||||
|
params: &GenerateSamplingParams,
|
||||||
|
) -> Result<Option<proto::sampling_params::Constraint>, String> {
|
||||||
|
let mut constraints = Vec::new();
|
||||||
|
if let Some(json_schema) = ¶ms.json_schema {
|
||||||
|
constraints.push(proto::sampling_params::Constraint::JsonSchema(
|
||||||
|
json_schema.clone(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
if let Some(regex) = ¶ms.regex {
|
||||||
|
constraints.push(proto::sampling_params::Constraint::Regex(regex.clone()));
|
||||||
|
}
|
||||||
|
if let Some(ebnf) = ¶ms.ebnf {
|
||||||
|
constraints.push(proto::sampling_params::Constraint::Grammar(ebnf.clone()));
|
||||||
|
}
|
||||||
|
|
||||||
|
match constraints.len() {
|
||||||
|
0 => Ok(None),
|
||||||
|
1 => Ok(constraints.pop()),
|
||||||
|
_ => Err("Multiple structured constraints are not allowed".to_string()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_sampling_params_from_plain(
|
||||||
|
params: Option<&GenerateSamplingParams>,
|
||||||
|
) -> Result<proto::SamplingParams, String> {
|
||||||
|
let mut sampling = proto::SamplingParams {
|
||||||
|
temperature: 1.0,
|
||||||
|
top_p: 1.0,
|
||||||
|
top_k: -1,
|
||||||
|
repetition_penalty: 1.0,
|
||||||
|
n: 1,
|
||||||
|
skip_special_tokens: true,
|
||||||
|
spaces_between_special_tokens: true,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
|
||||||
|
let Some(p) = params else {
|
||||||
|
return Ok(sampling);
|
||||||
|
};
|
||||||
|
|
||||||
|
// Simple field mappings using a macro
|
||||||
|
macro_rules! map_field {
|
||||||
|
($field:ident) => {
|
||||||
|
if let Some(val) = p.$field {
|
||||||
|
sampling.$field = val;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
map_field!(temperature);
|
||||||
|
map_field!(top_p);
|
||||||
|
map_field!(top_k);
|
||||||
|
map_field!(frequency_penalty);
|
||||||
|
map_field!(presence_penalty);
|
||||||
|
map_field!(repetition_penalty);
|
||||||
|
map_field!(min_p);
|
||||||
|
map_field!(ignore_eos);
|
||||||
|
map_field!(skip_special_tokens);
|
||||||
|
// Note: no_stop_trim not supported in vLLM
|
||||||
|
|
||||||
|
// Handle stop sequences
|
||||||
|
if let Some(stop) = &p.stop {
|
||||||
|
match stop {
|
||||||
|
StringOrArray::String(s) => sampling.stop.push(s.clone()),
|
||||||
|
StringOrArray::Array(arr) => sampling.stop.extend(arr.clone()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle stop token IDs
|
||||||
|
if let Some(stop_token_ids) = &p.stop_token_ids {
|
||||||
|
sampling.stop_token_ids = stop_token_ids.clone();
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle max_tokens with conversion (read from internal max_new_tokens)
|
||||||
|
if let Some(max_new_tokens) = p.max_new_tokens {
|
||||||
|
sampling.max_tokens = Some(
|
||||||
|
i32::try_from(max_new_tokens)
|
||||||
|
.map_err(|_| "max_tokens must fit into a 32-bit signed integer".to_string())?,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle min_tokens with conversion (read from internal min_new_tokens)
|
||||||
|
if let Some(min_new_tokens) = p.min_new_tokens {
|
||||||
|
sampling.min_tokens = i32::try_from(min_new_tokens)
|
||||||
|
.map_err(|_| "min_tokens must fit into a 32-bit signed integer".to_string())?;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle n with conversion
|
||||||
|
if let Some(n) = p.n {
|
||||||
|
sampling.n = i32::try_from(n)
|
||||||
|
.map_err(|_| "n must fit into a 32-bit signed integer".to_string())?;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle constraints (exactly one allowed)
|
||||||
|
sampling.constraint = Self::build_single_constraint_from_plain(p)?;
|
||||||
|
|
||||||
|
Ok(sampling)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_proto_types_compilation() {
|
||||||
|
let _health_req = proto::HealthCheckRequest {};
|
||||||
|
// HealthCheckRequest is now empty - no fields to test
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_generate_request_construction() {
|
||||||
|
let sampling_params = proto::SamplingParams {
|
||||||
|
temperature: 0.7,
|
||||||
|
max_tokens: Some(128),
|
||||||
|
top_p: 0.9,
|
||||||
|
top_k: 50,
|
||||||
|
stop: vec!["</s>".to_string()],
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
|
||||||
|
let gen_req = proto::GenerateRequest {
|
||||||
|
request_id: "test-req-123".to_string(),
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: "Hello world".to_string(),
|
||||||
|
input_ids: vec![9906, 1917], // Mock token IDs for "Hello world"
|
||||||
|
}),
|
||||||
|
sampling_params: Some(sampling_params),
|
||||||
|
stream: false,
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(gen_req.request_id, "test-req-123");
|
||||||
|
if let Some(ref tokenized) = &gen_req.tokenized {
|
||||||
|
assert_eq!(tokenized.original_text, "Hello world");
|
||||||
|
}
|
||||||
|
// vLLM: logprobs are in SamplingParams, not GenerateRequest
|
||||||
|
|
||||||
|
let params = gen_req.sampling_params.unwrap();
|
||||||
|
assert_eq!(params.temperature, 0.7);
|
||||||
|
assert_eq!(params.max_tokens, Some(128));
|
||||||
|
assert_eq!(params.stop, vec!["</s>"]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_health_check_request() {
|
||||||
|
let _health_req = proto::HealthCheckRequest {};
|
||||||
|
// HealthCheckRequest is now empty - server generates its own test internally
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_abort_request_construction() {
|
||||||
|
let abort_req = proto::AbortRequest {
|
||||||
|
request_id: "req-456".to_string(),
|
||||||
|
reason: "User canceled".to_string(),
|
||||||
|
};
|
||||||
|
assert_eq!(abort_req.request_id, "req-456");
|
||||||
|
assert_eq!(abort_req.reason, "User canceled");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_sampling_params_defaults() {
|
||||||
|
let params = proto::SamplingParams::default();
|
||||||
|
// Numeric fields have proto defaults (0)
|
||||||
|
assert_eq!(params.temperature, 0.0);
|
||||||
|
assert_eq!(params.top_p, 0.0);
|
||||||
|
assert_eq!(params.top_k, 0);
|
||||||
|
assert_eq!(params.repetition_penalty, 0.0);
|
||||||
|
assert_eq!(params.n, 0);
|
||||||
|
// Bool fields have proto defaults (false)
|
||||||
|
assert!(!params.skip_special_tokens);
|
||||||
|
assert!(!params.spaces_between_special_tokens);
|
||||||
|
assert!(!params.ignore_eos);
|
||||||
|
assert!(!params.include_stop_str_in_output);
|
||||||
|
// Optional int fields should be None
|
||||||
|
assert_eq!(params.max_tokens, None);
|
||||||
|
assert_eq!(params.logprobs, None);
|
||||||
|
// Other non-optional fields
|
||||||
|
assert_eq!(params.min_p, 0.0);
|
||||||
|
assert_eq!(params.frequency_penalty, 0.0);
|
||||||
|
assert_eq!(params.presence_penalty, 0.0);
|
||||||
|
assert!(params.stop.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: MultimodalInputs not in vLLM proto - skip test
|
||||||
|
// vLLM handles multimodal inputs differently than SGLang
|
||||||
|
|
||||||
|
// TODO: SessionParams not in current proto - skip test
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_embed_request() {
|
||||||
|
let embed_req = proto::EmbedRequest {
|
||||||
|
request_id: "embed-req-202".to_string(),
|
||||||
|
tokenized: Some(proto::TokenizedInput {
|
||||||
|
original_text: "This is a test sentence for embedding".to_string(),
|
||||||
|
input_ids: vec![2028, 374, 264, 1296, 11914, 369, 28537], // Mock token IDs
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(embed_req.request_id, "embed-req-202");
|
||||||
|
if let Some(ref tokenized) = &embed_req.tokenized {
|
||||||
|
assert_eq!(
|
||||||
|
tokenized.original_text,
|
||||||
|
"This is a test sentence for embedding"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
// vLLM: no data_parallel_rank or log_metrics in EmbedRequest
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_client_connect_invalid_endpoint() {
|
||||||
|
let result = VllmEngineClient::connect("invalid://endpoint").await;
|
||||||
|
assert!(result.is_err());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_tokenized_input() {
|
||||||
|
let tokenized = proto::TokenizedInput {
|
||||||
|
original_text: "Hello world".to_string(),
|
||||||
|
input_ids: vec![1, 15043, 1917, 2],
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(tokenized.original_text, "Hello world");
|
||||||
|
assert_eq!(tokenized.input_ids, vec![1, 15043, 1917, 2]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_generate_stream_chunk() {
|
||||||
|
let chunk = proto::GenerateStreamChunk {
|
||||||
|
token_ids: vec![1234, 5678],
|
||||||
|
prompt_tokens: 5,
|
||||||
|
completion_tokens: 2,
|
||||||
|
cached_tokens: 3,
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(chunk.token_ids, vec![1234, 5678]);
|
||||||
|
assert_eq!(chunk.prompt_tokens, 5);
|
||||||
|
assert_eq!(chunk.completion_tokens, 2);
|
||||||
|
assert_eq!(chunk.cached_tokens, 3);
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: ModelInfo not in current proto - skip test
|
||||||
|
}
|
||||||
@@ -0,0 +1,218 @@
|
|||||||
|
syntax = "proto3";
|
||||||
|
|
||||||
|
package vllm.grpc.engine;
|
||||||
|
|
||||||
|
// Service definition for vLLM engine communication
|
||||||
|
// This protocol is designed for efficient binary communication between
|
||||||
|
// the Rust router and vLLM Python engine (AsyncLLM).
|
||||||
|
service VllmEngine {
|
||||||
|
// Submit a generation request (supports streaming)
|
||||||
|
rpc Generate(GenerateRequest) returns (stream GenerateResponse);
|
||||||
|
|
||||||
|
// Submit an embedding request
|
||||||
|
rpc Embed(EmbedRequest) returns (EmbedResponse);
|
||||||
|
|
||||||
|
// Health check
|
||||||
|
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
|
||||||
|
|
||||||
|
// Abort a running request
|
||||||
|
rpc Abort(AbortRequest) returns (AbortResponse);
|
||||||
|
|
||||||
|
// Get model information
|
||||||
|
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
|
||||||
|
|
||||||
|
// Get server information
|
||||||
|
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Common Types
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
// Sampling parameters for text generation
|
||||||
|
message SamplingParams {
|
||||||
|
float temperature = 1;
|
||||||
|
float top_p = 2;
|
||||||
|
int32 top_k = 3;
|
||||||
|
float min_p = 4;
|
||||||
|
float frequency_penalty = 5;
|
||||||
|
float presence_penalty = 6;
|
||||||
|
float repetition_penalty = 7;
|
||||||
|
|
||||||
|
optional int32 max_tokens = 8;
|
||||||
|
int32 min_tokens = 9;
|
||||||
|
|
||||||
|
repeated string stop = 10;
|
||||||
|
repeated uint32 stop_token_ids = 11;
|
||||||
|
|
||||||
|
bool skip_special_tokens = 12;
|
||||||
|
bool spaces_between_special_tokens = 13;
|
||||||
|
bool ignore_eos = 14;
|
||||||
|
|
||||||
|
int32 n = 15; // Number of parallel samples
|
||||||
|
|
||||||
|
// Logprobs configuration
|
||||||
|
optional int32 logprobs = 22; // Number of log probabilities per output token (-1 for all)
|
||||||
|
optional int32 prompt_logprobs = 23; // Number of log probabilities per prompt token (-1 for all)
|
||||||
|
|
||||||
|
// Additional vLLM fields
|
||||||
|
optional int32 seed = 24; // Random seed for reproducibility
|
||||||
|
bool include_stop_str_in_output = 25; // Whether to include stop strings in output
|
||||||
|
map<int32, float> logit_bias = 26; // Token ID to bias mapping (-100 to 100)
|
||||||
|
optional int32 truncate_prompt_tokens = 27; // Prompt truncation (-1 for model max)
|
||||||
|
|
||||||
|
// Structured outputs (one of) - matches vLLM's StructuredOutputsParams
|
||||||
|
oneof constraint {
|
||||||
|
string json_schema = 16; // JSON schema for structured output
|
||||||
|
string regex = 17; // Regex pattern
|
||||||
|
string grammar = 18; // Grammar/EBNF for structured output
|
||||||
|
string structural_tag = 19; // Structural tag (e.g., Harmony models)
|
||||||
|
bool json_object = 20; // Force JSON object output
|
||||||
|
ChoiceConstraint choice = 21; // List of allowed choices
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Choice constraint for structured outputs
|
||||||
|
message ChoiceConstraint {
|
||||||
|
repeated string choices = 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pre-tokenized input from Rust router
|
||||||
|
message TokenizedInput {
|
||||||
|
string original_text = 1; // For reference/debugging
|
||||||
|
repeated uint32 input_ids = 2; // Actual token IDs to process
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Generate Request
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
message GenerateRequest {
|
||||||
|
string request_id = 1;
|
||||||
|
|
||||||
|
// Pre-tokenized input (required)
|
||||||
|
TokenizedInput tokenized = 2;
|
||||||
|
|
||||||
|
// Generation parameters (includes logprobs config)
|
||||||
|
SamplingParams sampling_params = 3;
|
||||||
|
|
||||||
|
// Streaming
|
||||||
|
bool stream = 4;
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Generate Response
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
message GenerateResponse {
|
||||||
|
string request_id = 1;
|
||||||
|
|
||||||
|
oneof response {
|
||||||
|
GenerateStreamChunk chunk = 2; // For streaming
|
||||||
|
GenerateComplete complete = 3; // For final/non-streaming
|
||||||
|
GenerateError error = 4; // For errors
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
message GenerateStreamChunk {
|
||||||
|
repeated uint32 token_ids = 1; // Incremental tokens
|
||||||
|
int32 prompt_tokens = 2;
|
||||||
|
int32 completion_tokens = 3;
|
||||||
|
int32 cached_tokens = 4;
|
||||||
|
|
||||||
|
// Logprobs support (TODO: implement in Phase 4)
|
||||||
|
// OutputLogProbs output_logprobs = 5;
|
||||||
|
// InputLogProbs input_logprobs = 6; // Only in first chunk
|
||||||
|
}
|
||||||
|
|
||||||
|
message GenerateComplete {
|
||||||
|
repeated uint32 output_ids = 1; // All output tokens
|
||||||
|
string finish_reason = 2; // "stop", "length", "abort"
|
||||||
|
int32 prompt_tokens = 3;
|
||||||
|
int32 completion_tokens = 4;
|
||||||
|
int32 cached_tokens = 5;
|
||||||
|
|
||||||
|
// Logprobs support (TODO: implement in Phase 4)
|
||||||
|
// OutputLogProbs output_logprobs = 6;
|
||||||
|
// InputLogProbs input_logprobs = 7;
|
||||||
|
}
|
||||||
|
|
||||||
|
message GenerateError {
|
||||||
|
string message = 1;
|
||||||
|
string http_status_code = 2;
|
||||||
|
string details = 3;
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Embedding Request
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
message EmbedRequest {
|
||||||
|
string request_id = 1;
|
||||||
|
TokenizedInput tokenized = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message EmbedResponse {
|
||||||
|
string request_id = 1;
|
||||||
|
|
||||||
|
oneof response {
|
||||||
|
EmbedComplete complete = 2;
|
||||||
|
EmbedError error = 3;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
message EmbedComplete {
|
||||||
|
repeated float embedding = 1;
|
||||||
|
int32 prompt_tokens = 2;
|
||||||
|
int32 embedding_dim = 3;
|
||||||
|
}
|
||||||
|
|
||||||
|
message EmbedError {
|
||||||
|
string message = 1;
|
||||||
|
string code = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Management Operations
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
message HealthCheckRequest {}
|
||||||
|
|
||||||
|
message HealthCheckResponse {
|
||||||
|
bool healthy = 1;
|
||||||
|
string message = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message AbortRequest {
|
||||||
|
string request_id = 1;
|
||||||
|
string reason = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message AbortResponse {
|
||||||
|
bool success = 1;
|
||||||
|
string message = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
// =====================
|
||||||
|
// Model and Server Info
|
||||||
|
// =====================
|
||||||
|
|
||||||
|
message GetModelInfoRequest {}
|
||||||
|
|
||||||
|
message GetModelInfoResponse {
|
||||||
|
string model_path = 1;
|
||||||
|
bool is_generation = 2;
|
||||||
|
int32 max_context_length = 3;
|
||||||
|
int32 vocab_size = 4;
|
||||||
|
bool supports_vision = 5;
|
||||||
|
}
|
||||||
|
|
||||||
|
message GetServerInfoRequest {}
|
||||||
|
|
||||||
|
message GetServerInfoResponse {
|
||||||
|
int32 active_requests = 1;
|
||||||
|
bool is_paused = 2;
|
||||||
|
double last_receive_timestamp = 3;
|
||||||
|
double uptime_seconds = 4;
|
||||||
|
string server_type = 5; // "vllm-grpc"
|
||||||
|
}
|
||||||
@@ -36,6 +36,11 @@ pub struct WorkerConfigRequest {
|
|||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub bootstrap_port: Option<u16>,
|
pub bootstrap_port: Option<u16>,
|
||||||
|
|
||||||
|
/// Runtime type (optional: "sglang", "vllm", default: "sglang")
|
||||||
|
/// Only relevant for gRPC workers
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub runtime: Option<String>,
|
||||||
|
|
||||||
// gRPC-specific configuration (optional, ignored in HTTP mode)
|
// gRPC-specific configuration (optional, ignored in HTTP mode)
|
||||||
/// Tokenizer path for gRPC mode
|
/// Tokenizer path for gRPC mode
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
@@ -133,6 +138,10 @@ pub struct WorkerInfo {
|
|||||||
/// Connection mode (http or grpc)
|
/// Connection mode (http or grpc)
|
||||||
pub connection_mode: String,
|
pub connection_mode: String,
|
||||||
|
|
||||||
|
/// Runtime type (sglang or vllm, for gRPC workers)
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub runtime_type: Option<String>,
|
||||||
|
|
||||||
// gRPC-specific fields (None for HTTP workers)
|
// gRPC-specific fields (None for HTTP workers)
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
pub tokenizer_path: Option<String>,
|
pub tokenizer_path: Option<String>,
|
||||||
|
|||||||
@@ -0,0 +1,177 @@
|
|||||||
|
//! Unified gRPC client wrapper for SGLang and vLLM backends
|
||||||
|
|
||||||
|
use crate::{
|
||||||
|
grpc_client::{SglangSchedulerClient, VllmEngineClient},
|
||||||
|
routers::grpc::proto_wrapper::{ProtoGenerateRequest, ProtoStream},
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Health check response (common across backends)
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct HealthCheckResponse {
|
||||||
|
pub healthy: bool,
|
||||||
|
pub message: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Polymorphic gRPC client that wraps either SGLang or vLLM
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum GrpcClient {
|
||||||
|
Sglang(SglangSchedulerClient),
|
||||||
|
Vllm(VllmEngineClient),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl GrpcClient {
|
||||||
|
/// Get reference to SGLang client (panics if vLLM)
|
||||||
|
pub fn as_sglang(&self) -> &SglangSchedulerClient {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(client) => client,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang client, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable reference to SGLang client (panics if vLLM)
|
||||||
|
pub fn as_sglang_mut(&mut self) -> &mut SglangSchedulerClient {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(client) => client,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang client, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get reference to vLLM client (panics if SGLang)
|
||||||
|
pub fn as_vllm(&self) -> &VllmEngineClient {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(client) => client,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM client, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable reference to vLLM client (panics if SGLang)
|
||||||
|
pub fn as_vllm_mut(&mut self) -> &mut VllmEngineClient {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(client) => client,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM client, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is a SGLang client
|
||||||
|
pub fn is_sglang(&self) -> bool {
|
||||||
|
matches!(self, Self::Sglang(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is a vLLM client
|
||||||
|
pub fn is_vllm(&self) -> bool {
|
||||||
|
matches!(self, Self::Vllm(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Connect to gRPC server (runtime-aware)
|
||||||
|
pub async fn connect(
|
||||||
|
url: &str,
|
||||||
|
runtime_type: &str,
|
||||||
|
) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
match runtime_type {
|
||||||
|
"sglang" => Ok(Self::Sglang(SglangSchedulerClient::connect(url).await?)),
|
||||||
|
"vllm" => Ok(Self::Vllm(VllmEngineClient::connect(url).await?)),
|
||||||
|
_ => Err(format!("Unknown runtime type: {}", runtime_type).into()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Perform health check (dispatches to appropriate backend)
|
||||||
|
pub async fn health_check(
|
||||||
|
&self,
|
||||||
|
) -> Result<HealthCheckResponse, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(client) => {
|
||||||
|
let resp = client.health_check().await?;
|
||||||
|
Ok(HealthCheckResponse {
|
||||||
|
healthy: resp.healthy,
|
||||||
|
message: resp.message,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
Self::Vllm(client) => {
|
||||||
|
let resp = client.health_check().await?;
|
||||||
|
Ok(HealthCheckResponse {
|
||||||
|
healthy: resp.healthy,
|
||||||
|
message: resp.message,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get model info (returns enum wrapping backend-specific response)
|
||||||
|
pub async fn get_model_info(
|
||||||
|
&self,
|
||||||
|
) -> Result<ModelInfo, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(client) => {
|
||||||
|
let info = client.get_model_info().await?;
|
||||||
|
Ok(ModelInfo::Sglang(info))
|
||||||
|
}
|
||||||
|
Self::Vllm(client) => {
|
||||||
|
let info = client.get_model_info().await?;
|
||||||
|
Ok(ModelInfo::Vllm(info))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Generate streaming response from request
|
||||||
|
///
|
||||||
|
/// Dispatches to the appropriate backend client and wraps the result in ProtoStream
|
||||||
|
pub async fn generate(
|
||||||
|
&mut self,
|
||||||
|
req: ProtoGenerateRequest,
|
||||||
|
) -> Result<ProtoStream, Box<dyn std::error::Error + Send + Sync>> {
|
||||||
|
match (self, req) {
|
||||||
|
(Self::Sglang(client), ProtoGenerateRequest::Sglang(boxed_req)) => {
|
||||||
|
let stream = client.generate(*boxed_req).await?;
|
||||||
|
Ok(ProtoStream::Sglang(stream))
|
||||||
|
}
|
||||||
|
(Self::Vllm(client), ProtoGenerateRequest::Vllm(boxed_req)) => {
|
||||||
|
let stream = client.generate(*boxed_req).await?;
|
||||||
|
Ok(ProtoStream::Vllm(stream))
|
||||||
|
}
|
||||||
|
_ => panic!("Mismatched client and request types"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified ModelInfo wrapper
|
||||||
|
pub enum ModelInfo {
|
||||||
|
Sglang(crate::grpc_client::sglang_proto::GetModelInfoResponse),
|
||||||
|
Vllm(crate::grpc_client::vllm_proto::GetModelInfoResponse),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ModelInfo {
|
||||||
|
/// Convert model info to label map for worker metadata
|
||||||
|
pub fn to_labels(&self) -> std::collections::HashMap<String, String> {
|
||||||
|
let mut labels = std::collections::HashMap::new();
|
||||||
|
|
||||||
|
// Serialize to JSON Value (like pydantic's model_dump)
|
||||||
|
let value = match self {
|
||||||
|
ModelInfo::Sglang(info) => serde_json::to_value(info).ok(),
|
||||||
|
ModelInfo::Vllm(info) => serde_json::to_value(info).ok(),
|
||||||
|
};
|
||||||
|
|
||||||
|
// Convert JSON object to HashMap, filtering out empty/zero/false values
|
||||||
|
if let Some(serde_json::Value::Object(obj)) = value {
|
||||||
|
for (key, val) in obj {
|
||||||
|
match val {
|
||||||
|
// Insert non-empty strings
|
||||||
|
serde_json::Value::String(s) if !s.is_empty() => {
|
||||||
|
labels.insert(key, s);
|
||||||
|
}
|
||||||
|
// Insert positive numbers
|
||||||
|
serde_json::Value::Number(n) if n.as_i64().unwrap_or(0) > 0 => {
|
||||||
|
labels.insert(key, n.to_string());
|
||||||
|
}
|
||||||
|
// Insert true booleans
|
||||||
|
serde_json::Value::Bool(true) => {
|
||||||
|
labels.insert(key, "true".to_string());
|
||||||
|
}
|
||||||
|
// Skip empty strings, zeros, false, nulls, arrays, objects
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
labels
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -5,9 +5,8 @@
|
|||||||
|
|
||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
|
|
||||||
use crate::{
|
use crate::routers::grpc::{
|
||||||
grpc_client::proto,
|
context::ExecutionResult, error, proto_wrapper::ProtoGenerateComplete, utils,
|
||||||
routers::grpc::{context::ExecutionResult, error, utils},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Collect and merge responses from execution result
|
/// Collect and merge responses from execution result
|
||||||
@@ -24,7 +23,7 @@ use crate::{
|
|||||||
pub async fn collect_responses(
|
pub async fn collect_responses(
|
||||||
execution_result: ExecutionResult,
|
execution_result: ExecutionResult,
|
||||||
merge_logprobs: bool,
|
merge_logprobs: bool,
|
||||||
) -> Result<Vec<proto::GenerateComplete>, Response> {
|
) -> Result<Vec<ProtoGenerateComplete>, Response> {
|
||||||
let all_responses = match execution_result {
|
let all_responses = match execution_result {
|
||||||
ExecutionResult::Single { mut stream } => {
|
ExecutionResult::Single { mut stream } => {
|
||||||
let responses = utils::collect_stream_responses(&mut stream, "Single").await?;
|
let responses = utils::collect_stream_responses(&mut stream, "Single").await?;
|
||||||
@@ -68,16 +67,19 @@ pub async fn collect_responses(
|
|||||||
///
|
///
|
||||||
/// Takes input_logprobs from the first prefill response and copies them
|
/// Takes input_logprobs from the first prefill response and copies them
|
||||||
/// into all decode responses. This is used in PD mode when logprobs are requested.
|
/// into all decode responses. This is used in PD mode when logprobs are requested.
|
||||||
|
/// Only works with SGLang (vLLM doesn't support PD mode).
|
||||||
fn merge_prefill_logprobs(
|
fn merge_prefill_logprobs(
|
||||||
prefill_responses: &[proto::GenerateComplete],
|
prefill_responses: &[ProtoGenerateComplete],
|
||||||
decode_responses: &mut [proto::GenerateComplete],
|
decode_responses: &mut [ProtoGenerateComplete],
|
||||||
) {
|
) {
|
||||||
if let Some(prefill_input_logprobs) = prefill_responses
|
// Only SGLang supports PD mode and has input_logprobs
|
||||||
.first()
|
if let Some(ProtoGenerateComplete::Sglang(prefill_first)) = prefill_responses.first() {
|
||||||
.and_then(|r| r.input_logprobs.clone())
|
if let Some(prefill_input_logprobs) = prefill_first.input_logprobs.clone() {
|
||||||
{
|
for response in decode_responses.iter_mut() {
|
||||||
for response in decode_responses.iter_mut() {
|
if let ProtoGenerateComplete::Sglang(decode_resp) = response {
|
||||||
response.input_logprobs = Some(prefill_input_logprobs.clone());
|
decode_resp.input_logprobs = Some(prefill_input_logprobs.clone());
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,7 @@
|
|||||||
//! - Usage calculation from gRPC responses
|
//! - Usage calculation from gRPC responses
|
||||||
//! - ChatCompletionResponse construction
|
//! - ChatCompletionResponse construction
|
||||||
|
|
||||||
use crate::{grpc_client::proto, protocols::common::Usage};
|
use crate::{protocols::common::Usage, routers::grpc::proto_wrapper::ProtoGenerateComplete};
|
||||||
|
|
||||||
/// Build usage information from collected gRPC responses
|
/// Build usage information from collected gRPC responses
|
||||||
///
|
///
|
||||||
@@ -16,9 +16,9 @@ use crate::{grpc_client::proto, protocols::common::Usage};
|
|||||||
///
|
///
|
||||||
/// # Returns
|
/// # Returns
|
||||||
/// Usage object with aggregated token counts
|
/// Usage object with aggregated token counts
|
||||||
pub fn build_usage(responses: &[proto::GenerateComplete]) -> Usage {
|
pub fn build_usage(responses: &[ProtoGenerateComplete]) -> Usage {
|
||||||
let total_prompt_tokens: u32 = responses.iter().map(|r| r.prompt_tokens as u32).sum();
|
let total_prompt_tokens: u32 = responses.iter().map(|r| r.prompt_tokens() as u32).sum();
|
||||||
let total_completion_tokens: u32 = responses.iter().map(|r| r.completion_tokens as u32).sum();
|
let total_completion_tokens: u32 = responses.iter().map(|r| r.completion_tokens() as u32).sum();
|
||||||
|
|
||||||
Usage {
|
Usage {
|
||||||
prompt_tokens: total_prompt_tokens,
|
prompt_tokens: total_prompt_tokens,
|
||||||
|
|||||||
@@ -32,6 +32,19 @@ impl PipelineStage for ClientAcquisitionStage {
|
|||||||
WorkerSelection::Dual { prefill, decode } => {
|
WorkerSelection::Dual { prefill, decode } => {
|
||||||
let prefill_client = utils::get_grpc_client_from_worker(prefill).await?;
|
let prefill_client = utils::get_grpc_client_from_worker(prefill).await?;
|
||||||
let decode_client = utils::get_grpc_client_from_worker(decode).await?;
|
let decode_client = utils::get_grpc_client_from_worker(decode).await?;
|
||||||
|
|
||||||
|
// vLLM does not support dual (PD disaggregated) mode
|
||||||
|
if prefill_client.is_vllm() || decode_client.is_vllm() {
|
||||||
|
error!(
|
||||||
|
function = "ClientAcquisitionStage::execute",
|
||||||
|
"vLLM backend does not support dual (PD disaggregated) mode"
|
||||||
|
);
|
||||||
|
return Err(error::bad_request(
|
||||||
|
"vLLM backend does not support prefill/decode disaggregated mode. \
|
||||||
|
Please use runtime_type: sglang for PD mode, or use a regular (non-PD) worker configuration."
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
ClientSelection::Dual {
|
ClientSelection::Dual {
|
||||||
prefill: prefill_client,
|
prefill: prefill_client,
|
||||||
decode: decode_client,
|
decode: decode_client,
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ impl PipelineStage for DispatchMetadataStage {
|
|||||||
error::internal_error("Proto request not built")
|
error::internal_error("Proto request not built")
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
let request_id = proto_request.request_id.clone();
|
let request_id = proto_request.request_id().to_string();
|
||||||
let model = match &ctx.input.request_type {
|
let model = match &ctx.input.request_type {
|
||||||
RequestType::Chat(req) => req.model.clone(),
|
RequestType::Chat(req) => req.model.clone(),
|
||||||
RequestType::Generate(_req) => {
|
RequestType::Generate(_req) => {
|
||||||
|
|||||||
@@ -2,17 +2,20 @@
|
|||||||
|
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
use proto::DisaggregatedParams;
|
|
||||||
use rand::Rng;
|
use rand::Rng;
|
||||||
use tracing::debug;
|
use tracing::debug;
|
||||||
|
|
||||||
use crate::{core::Worker, grpc_client::proto};
|
use crate::{
|
||||||
|
core::Worker, grpc_client::sglang_proto::DisaggregatedParams,
|
||||||
|
routers::grpc::proto_wrapper::ProtoGenerateRequest,
|
||||||
|
};
|
||||||
|
|
||||||
/// Inject PD bootstrap metadata into a gRPC request
|
/// Inject PD bootstrap metadata into a gRPC request
|
||||||
///
|
///
|
||||||
/// Used by both chat and generate request building stages when in PD mode.
|
/// Used by both chat and generate request building stages when in PD mode.
|
||||||
|
/// Only SGLang supports PD (prefill/decode) disaggregated mode.
|
||||||
pub fn inject_bootstrap_metadata(
|
pub fn inject_bootstrap_metadata(
|
||||||
request: &mut proto::GenerateRequest,
|
request: &mut ProtoGenerateRequest,
|
||||||
prefill_worker: &Arc<dyn Worker>,
|
prefill_worker: &Arc<dyn Worker>,
|
||||||
) {
|
) {
|
||||||
let hostname = prefill_worker.bootstrap_host();
|
let hostname = prefill_worker.bootstrap_host();
|
||||||
@@ -28,8 +31,10 @@ pub fn inject_bootstrap_metadata(
|
|||||||
bootstrap_room: room_id,
|
bootstrap_room: room_id,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Inject metadata directly into request
|
// Inject metadata directly into SGLang request
|
||||||
request.disaggregated_params = Some(disagg_params);
|
// (vLLM doesn't support PD mode, so this will panic if called with vLLM)
|
||||||
|
let sglang_request = request.as_sglang_mut();
|
||||||
|
sglang_request.disaggregated_params = Some(disagg_params);
|
||||||
|
|
||||||
debug!(
|
debug!(
|
||||||
"Injected bootstrap metadata: host={}, port={}, room={}",
|
"Injected bootstrap metadata: host={}, port={}, room={}",
|
||||||
|
|||||||
@@ -5,15 +5,13 @@ use axum::response::Response;
|
|||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use super::PipelineStage;
|
use super::PipelineStage;
|
||||||
use crate::{
|
use crate::routers::grpc::{
|
||||||
grpc_client::{proto, sglang_scheduler::AbortOnDropStream},
|
context::{ClientSelection, ExecutionResult, RequestContext},
|
||||||
routers::grpc::{
|
error,
|
||||||
context::{ClientSelection, ExecutionResult, RequestContext},
|
proto_wrapper::{ProtoGenerateRequest, ProtoStream},
|
||||||
error,
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
type StreamResult = Result<AbortOnDropStream, Box<dyn std::error::Error + Send + Sync>>;
|
type StreamResult = Result<ProtoStream, Box<dyn std::error::Error + Send + Sync>>;
|
||||||
|
|
||||||
/// Request execution stage: Execute gRPC requests (single or dual dispatch)
|
/// Request execution stage: Execute gRPC requests (single or dual dispatch)
|
||||||
pub struct RequestExecutionStage {
|
pub struct RequestExecutionStage {
|
||||||
@@ -72,7 +70,7 @@ impl PipelineStage for RequestExecutionStage {
|
|||||||
impl RequestExecutionStage {
|
impl RequestExecutionStage {
|
||||||
async fn execute_single(
|
async fn execute_single(
|
||||||
&self,
|
&self,
|
||||||
proto_request: proto::GenerateRequest,
|
proto_request: ProtoGenerateRequest,
|
||||||
clients: &mut ClientSelection,
|
clients: &mut ClientSelection,
|
||||||
) -> Result<ExecutionResult, Response> {
|
) -> Result<ExecutionResult, Response> {
|
||||||
let client = clients.single_mut().ok_or_else(|| {
|
let client = clients.single_mut().ok_or_else(|| {
|
||||||
@@ -97,7 +95,7 @@ impl RequestExecutionStage {
|
|||||||
|
|
||||||
async fn execute_dual_dispatch(
|
async fn execute_dual_dispatch(
|
||||||
&self,
|
&self,
|
||||||
proto_request: proto::GenerateRequest,
|
proto_request: ProtoGenerateRequest,
|
||||||
clients: &mut ClientSelection,
|
clients: &mut ClientSelection,
|
||||||
) -> Result<ExecutionResult, Response> {
|
) -> Result<ExecutionResult, Response> {
|
||||||
let (prefill_client, decode_client) = clients.dual_mut().ok_or_else(|| {
|
let (prefill_client, decode_client) = clients.dual_mut().ok_or_else(|| {
|
||||||
@@ -108,7 +106,7 @@ impl RequestExecutionStage {
|
|||||||
error::internal_error("Expected dual clients but got single")
|
error::internal_error("Expected dual clients but got single")
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
let prefill_request = proto_request.clone();
|
let prefill_request = proto_request.clone_inner();
|
||||||
let decode_request = proto_request;
|
let decode_request = proto_request;
|
||||||
|
|
||||||
let (prefill_result, decode_result): (StreamResult, StreamResult) = tokio::join!(
|
let (prefill_result, decode_result): (StreamResult, StreamResult) = tokio::join!(
|
||||||
|
|||||||
@@ -9,9 +9,12 @@ use std::{collections::HashMap, sync::Arc};
|
|||||||
use axum::http::HeaderMap;
|
use axum::http::HeaderMap;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
|
use super::{
|
||||||
|
client::GrpcClient,
|
||||||
|
proto_wrapper::{ProtoGenerateComplete, ProtoGenerateRequest, ProtoStream},
|
||||||
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
core::Worker,
|
core::Worker,
|
||||||
grpc_client::{proto, sglang_scheduler::AbortOnDropStream, SglangSchedulerClient},
|
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
chat::{ChatCompletionRequest, ChatCompletionResponse},
|
||||||
generate::{GenerateRequest, GenerateResponse},
|
generate::{GenerateRequest, GenerateResponse},
|
||||||
@@ -68,7 +71,7 @@ pub struct ProcessingState {
|
|||||||
pub clients: Option<ClientSelection>,
|
pub clients: Option<ClientSelection>,
|
||||||
|
|
||||||
// Stage 4: Request building outputs
|
// Stage 4: Request building outputs
|
||||||
pub proto_request: Option<proto::GenerateRequest>,
|
pub proto_request: Option<ProtoGenerateRequest>,
|
||||||
|
|
||||||
// Stage 5: Dispatch metadata
|
// Stage 5: Dispatch metadata
|
||||||
pub dispatch: Option<DispatchMetadata>,
|
pub dispatch: Option<DispatchMetadata>,
|
||||||
@@ -122,11 +125,11 @@ pub enum WorkerSelection {
|
|||||||
/// Client selection (Step 3)
|
/// Client selection (Step 3)
|
||||||
pub enum ClientSelection {
|
pub enum ClientSelection {
|
||||||
Single {
|
Single {
|
||||||
client: SglangSchedulerClient,
|
client: GrpcClient,
|
||||||
},
|
},
|
||||||
Dual {
|
Dual {
|
||||||
prefill: SglangSchedulerClient,
|
prefill: GrpcClient,
|
||||||
decode: SglangSchedulerClient,
|
decode: GrpcClient,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,7 +153,7 @@ pub struct ResponseState {
|
|||||||
pub streaming: StreamingState,
|
pub streaming: StreamingState,
|
||||||
|
|
||||||
/// Collected responses (non-streaming)
|
/// Collected responses (non-streaming)
|
||||||
pub collected: Option<Vec<proto::GenerateComplete>>,
|
pub collected: Option<Vec<ProtoGenerateComplete>>,
|
||||||
|
|
||||||
/// Execution result (streams from workers)
|
/// Execution result (streams from workers)
|
||||||
pub execution_result: Option<ExecutionResult>,
|
pub execution_result: Option<ExecutionResult>,
|
||||||
@@ -346,56 +349,56 @@ impl ClientSelection {
|
|||||||
matches!(self, Self::Dual { .. })
|
matches!(self, Self::Dual { .. })
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn single(&self) -> Option<&SglangSchedulerClient> {
|
pub fn single(&self) -> Option<&GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Single { client } => Some(client),
|
Self::Single { client } => Some(client),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn single_mut(&mut self) -> Option<&mut SglangSchedulerClient> {
|
pub fn single_mut(&mut self) -> Option<&mut GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Single { client } => Some(client),
|
Self::Single { client } => Some(client),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn dual(&self) -> Option<(&SglangSchedulerClient, &SglangSchedulerClient)> {
|
pub fn dual(&self) -> Option<(&GrpcClient, &GrpcClient)> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { prefill, decode } => Some((prefill, decode)),
|
Self::Dual { prefill, decode } => Some((prefill, decode)),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn dual_mut(&mut self) -> Option<(&mut SglangSchedulerClient, &mut SglangSchedulerClient)> {
|
pub fn dual_mut(&mut self) -> Option<(&mut GrpcClient, &mut GrpcClient)> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { prefill, decode } => Some((prefill, decode)),
|
Self::Dual { prefill, decode } => Some((prefill, decode)),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn prefill_client(&self) -> Option<&SglangSchedulerClient> {
|
pub fn prefill_client(&self) -> Option<&GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { prefill, .. } => Some(prefill),
|
Self::Dual { prefill, .. } => Some(prefill),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn prefill_client_mut(&mut self) -> Option<&mut SglangSchedulerClient> {
|
pub fn prefill_client_mut(&mut self) -> Option<&mut GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { prefill, .. } => Some(prefill),
|
Self::Dual { prefill, .. } => Some(prefill),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn decode_client(&self) -> Option<&SglangSchedulerClient> {
|
pub fn decode_client(&self) -> Option<&GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { decode, .. } => Some(decode),
|
Self::Dual { decode, .. } => Some(decode),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn decode_client_mut(&mut self) -> Option<&mut SglangSchedulerClient> {
|
pub fn decode_client_mut(&mut self) -> Option<&mut GrpcClient> {
|
||||||
match self {
|
match self {
|
||||||
Self::Dual { decode, .. } => Some(decode),
|
Self::Dual { decode, .. } => Some(decode),
|
||||||
_ => None,
|
_ => None,
|
||||||
@@ -404,14 +407,14 @@ impl ClientSelection {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Result of request execution (streams from workers)
|
/// Result of request execution (streams from workers)
|
||||||
/// Uses AbortOnDropStream to automatically abort on cancellation
|
/// Uses ProtoStream to automatically abort on cancellation
|
||||||
pub enum ExecutionResult {
|
pub enum ExecutionResult {
|
||||||
Single {
|
Single {
|
||||||
stream: AbortOnDropStream,
|
stream: ProtoStream,
|
||||||
},
|
},
|
||||||
Dual {
|
Dual {
|
||||||
prefill: AbortOnDropStream,
|
prefill: ProtoStream,
|
||||||
decode: Box<AbortOnDropStream>,
|
decode: Box<ProtoStream>,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -125,6 +125,29 @@ pub fn failed_dependency(message: impl Into<String>) -> Response {
|
|||||||
.into_response()
|
.into_response()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Create a 501 Not Implemented response
|
||||||
|
///
|
||||||
|
/// Use this for features that are not yet implemented or supported.
|
||||||
|
///
|
||||||
|
/// # Example
|
||||||
|
/// ```ignore
|
||||||
|
/// return Err(not_implemented("vLLM backend integration is in progress"));
|
||||||
|
/// ```
|
||||||
|
pub fn not_implemented(message: impl Into<String>) -> Response {
|
||||||
|
let msg = message.into();
|
||||||
|
(
|
||||||
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
|
Json(json!({
|
||||||
|
"error": {
|
||||||
|
"message": msg,
|
||||||
|
"type": "not_implemented_error",
|
||||||
|
"code": 501
|
||||||
|
}
|
||||||
|
})),
|
||||||
|
)
|
||||||
|
.into_response()
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|||||||
@@ -3,12 +3,11 @@
|
|||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId};
|
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use super::HarmonyParserAdapter;
|
use super::HarmonyParserAdapter;
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::proto,
|
grpc_client::sglang_proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatChoice, ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse},
|
chat::{ChatChoice, ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse},
|
||||||
common::{ToolCall, Usage},
|
common::{ToolCall, Usage},
|
||||||
@@ -53,7 +52,7 @@ impl HarmonyResponseProcessor {
|
|||||||
let mut choices: Vec<ChatChoice> = Vec::new();
|
let mut choices: Vec<ChatChoice> = Vec::new();
|
||||||
for (index, complete) in all_responses.iter().enumerate() {
|
for (index, complete) in all_responses.iter().enumerate() {
|
||||||
// Convert matched_stop from proto to JSON
|
// Convert matched_stop from proto to JSON
|
||||||
let matched_stop = complete.matched_stop.as_ref().map(|m| match m {
|
let matched_stop = complete.matched_stop().map(|m| match m {
|
||||||
MatchedTokenId(id) => {
|
MatchedTokenId(id) => {
|
||||||
serde_json::json!(id)
|
serde_json::json!(id)
|
||||||
}
|
}
|
||||||
@@ -75,8 +74,8 @@ impl HarmonyResponseProcessor {
|
|||||||
// Parse Harmony channels with finish_reason and matched_stop
|
// Parse Harmony channels with finish_reason and matched_stop
|
||||||
let parsed = parser
|
let parsed = parser
|
||||||
.parse_complete(
|
.parse_complete(
|
||||||
&complete.output_ids,
|
complete.output_ids(),
|
||||||
complete.finish_reason.clone(),
|
complete.finish_reason().to_string(),
|
||||||
matched_stop.clone(),
|
matched_stop.clone(),
|
||||||
)
|
)
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
@@ -193,7 +192,7 @@ impl HarmonyResponseProcessor {
|
|||||||
})?;
|
})?;
|
||||||
|
|
||||||
// Convert matched_stop from proto to JSON
|
// Convert matched_stop from proto to JSON
|
||||||
let matched_stop = complete.matched_stop.as_ref().map(|m| match m {
|
let matched_stop = complete.matched_stop().map(|m| match m {
|
||||||
MatchedTokenId(id) => {
|
MatchedTokenId(id) => {
|
||||||
serde_json::json!(id)
|
serde_json::json!(id)
|
||||||
}
|
}
|
||||||
@@ -204,8 +203,8 @@ impl HarmonyResponseProcessor {
|
|||||||
|
|
||||||
let parsed = parser
|
let parsed = parser
|
||||||
.parse_complete(
|
.parse_complete(
|
||||||
&complete.output_ids,
|
complete.output_ids(),
|
||||||
complete.finish_reason.clone(),
|
complete.finish_reason().to_string(),
|
||||||
matched_stop,
|
matched_stop,
|
||||||
)
|
)
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ use crate::routers::grpc::{
|
|||||||
common::stages::{helpers, PipelineStage},
|
common::stages::{helpers, PipelineStage},
|
||||||
context::{ClientSelection, RequestContext, RequestType, WorkerSelection},
|
context::{ClientSelection, RequestContext, RequestType, WorkerSelection},
|
||||||
error,
|
error,
|
||||||
|
proto_wrapper::ProtoGenerateRequest,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Harmony Request Building stage: Convert Harmony tokens to gRPC request
|
/// Harmony Request Building stage: Convert Harmony tokens to gRPC request
|
||||||
@@ -51,6 +52,14 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
ClientSelection::Dual { prefill, .. } => prefill,
|
ClientSelection::Dual { prefill, .. } => prefill,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// Harmony model support not yet implemented for vLLM
|
||||||
|
if builder_client.is_vllm() {
|
||||||
|
return Err(error::not_implemented(
|
||||||
|
"Harmony model support is not yet implemented for vLLM backend. \
|
||||||
|
Please use runtime_type: sglang for Harmony models.",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
// Generate request_id based on request type
|
// Generate request_id based on request type
|
||||||
let request_id = match &ctx.input.request_type {
|
let request_id = match &ctx.input.request_type {
|
||||||
RequestType::Chat(_) => format!("chatcmpl-{}", Uuid::new_v4()),
|
RequestType::Chat(_) => format!("chatcmpl-{}", Uuid::new_v4()),
|
||||||
@@ -69,12 +78,14 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
// Build gRPC request using token_ids directly (Harmony encoding already handled message rendering)
|
// Build gRPC request using token_ids directly (Harmony encoding already handled message rendering)
|
||||||
let placeholder_processed_text = "[harmony]".to_string();
|
let placeholder_processed_text = "[harmony]".to_string();
|
||||||
|
|
||||||
let mut proto_request = match &ctx.input.request_type {
|
// Harmony is SGLang-only, so we can safely unwrap as SGLang
|
||||||
|
let sglang_client = builder_client.as_sglang();
|
||||||
|
let proto_request_inner = match &ctx.input.request_type {
|
||||||
RequestType::Chat(request) => {
|
RequestType::Chat(request) => {
|
||||||
// Use filtered request if present from preparation; otherwise original
|
// Use filtered request if present from preparation; otherwise original
|
||||||
let body = prep.filtered_request.as_ref().unwrap_or(request.as_ref());
|
let body = prep.filtered_request.as_ref().unwrap_or(request.as_ref());
|
||||||
|
|
||||||
builder_client
|
sglang_client
|
||||||
.build_generate_request_from_chat(
|
.build_generate_request_from_chat(
|
||||||
request_id,
|
request_id,
|
||||||
body,
|
body,
|
||||||
@@ -92,7 +103,7 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
error::bad_request(format!("Invalid request parameters: {}", e))
|
error::bad_request(format!("Invalid request parameters: {}", e))
|
||||||
})?
|
})?
|
||||||
}
|
}
|
||||||
RequestType::Responses(request) => builder_client
|
RequestType::Responses(request) => sglang_client
|
||||||
.build_generate_request_from_responses(
|
.build_generate_request_from_responses(
|
||||||
request_id,
|
request_id,
|
||||||
request.as_ref(),
|
request.as_ref(),
|
||||||
@@ -112,11 +123,14 @@ impl PipelineStage for HarmonyRequestBuildingStage {
|
|||||||
_ => unreachable!(),
|
_ => unreachable!(),
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let mut proto_request = ProtoGenerateRequest::Sglang(Box::new(proto_request_inner));
|
||||||
|
|
||||||
// Inject Harmony stop token IDs into sampling params for ALL Harmony requests
|
// Inject Harmony stop token IDs into sampling params for ALL Harmony requests
|
||||||
// These stop tokens (<|return|> and <|call|>) prevent the model from generating
|
// These stop tokens (<|return|> and <|call|>) prevent the model from generating
|
||||||
// malformed Harmony sequences
|
// malformed Harmony sequences
|
||||||
if let Some(harmony_stops) = &prep.harmony_stop_ids {
|
if let Some(harmony_stops) = &prep.harmony_stop_ids {
|
||||||
if let Some(params) = proto_request.sampling_params.as_mut() {
|
let sglang_req = proto_request.as_sglang_mut();
|
||||||
|
if let Some(params) = sglang_req.sampling_params.as_mut() {
|
||||||
params.stop_token_ids.extend_from_slice(harmony_stops);
|
params.stop_token_ids.extend_from_slice(harmony_stops);
|
||||||
debug!(
|
debug!(
|
||||||
stop_token_count = harmony_stops.len(),
|
stop_token_count = harmony_stops.len(),
|
||||||
|
|||||||
@@ -9,20 +9,16 @@ use std::{
|
|||||||
use axum::{body::Body, http::StatusCode, response::Response};
|
use axum::{body::Body, http::StatusCode, response::Response};
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use http::header::{HeaderValue, CONTENT_TYPE};
|
use http::header::{HeaderValue, CONTENT_TYPE};
|
||||||
use proto::{
|
|
||||||
generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
|
||||||
generate_response::Response::{Chunk, Complete},
|
|
||||||
};
|
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use tokio::sync::mpsc;
|
use tokio::sync::mpsc;
|
||||||
use tokio_stream::{wrappers::UnboundedReceiverStream, StreamExt};
|
use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||||
use tracing::{debug, error};
|
use tracing::{debug, error};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
processor::ResponsesIterationResult, types::HarmonyChannelDelta, HarmonyParserAdapter,
|
processor::ResponsesIterationResult, types::HarmonyChannelDelta, HarmonyParserAdapter,
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::{proto, sglang_scheduler::AbortOnDropStream},
|
grpc_client::sglang_proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{
|
chat::{
|
||||||
ChatCompletionRequest, ChatCompletionStreamResponse, ChatMessageDelta, ChatStreamChoice,
|
ChatCompletionRequest, ChatCompletionStreamResponse, ChatMessageDelta, ChatStreamChoice,
|
||||||
@@ -33,6 +29,7 @@ use crate::{
|
|||||||
routers::grpc::{
|
routers::grpc::{
|
||||||
common::responses::streaming::{OutputItemType, ResponseStreamEventEmitter},
|
common::responses::streaming::{OutputItemType, ResponseStreamEventEmitter},
|
||||||
context,
|
context,
|
||||||
|
proto_wrapper::{ProtoResponseVariant, ProtoStream},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -179,7 +176,7 @@ impl HarmonyStreamingProcessor {
|
|||||||
|
|
||||||
/// Process streaming chunks from a single stream
|
/// Process streaming chunks from a single stream
|
||||||
async fn process_single_stream(
|
async fn process_single_stream(
|
||||||
mut grpc_stream: AbortOnDropStream,
|
mut grpc_stream: ProtoStream,
|
||||||
dispatch: context::DispatchMetadata,
|
dispatch: context::DispatchMetadata,
|
||||||
original_request: Arc<ChatCompletionRequest>,
|
original_request: Arc<ChatCompletionRequest>,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
@@ -198,8 +195,9 @@ impl HarmonyStreamingProcessor {
|
|||||||
while let Some(result) = grpc_stream.next().await {
|
while let Some(result) = grpc_stream.next().await {
|
||||||
let response = result.map_err(|e| format!("Stream error: {}", e))?;
|
let response = result.map_err(|e| format!("Stream error: {}", e))?;
|
||||||
|
|
||||||
match response.response {
|
match response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk_wrapper) => {
|
||||||
|
let chunk = chunk_wrapper.as_sglang();
|
||||||
let index = chunk.index;
|
let index = chunk.index;
|
||||||
|
|
||||||
// Initialize parser for this index if needed
|
// Initialize parser for this index if needed
|
||||||
@@ -240,7 +238,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete_wrapper) => {
|
||||||
|
let complete = complete_wrapper.as_sglang();
|
||||||
let index = complete.index;
|
let index = complete.index;
|
||||||
|
|
||||||
// Store final metadata
|
// Store final metadata
|
||||||
@@ -249,10 +248,10 @@ impl HarmonyStreamingProcessor {
|
|||||||
index,
|
index,
|
||||||
complete.matched_stop.as_ref().map(|m| match m {
|
complete.matched_stop.as_ref().map(|m| match m {
|
||||||
MatchedTokenId(id) => {
|
MatchedTokenId(id) => {
|
||||||
serde_json::json!(id)
|
json!(id)
|
||||||
}
|
}
|
||||||
MatchedStopStr(s) => {
|
MatchedStopStr(s) => {
|
||||||
serde_json::json!(s)
|
json!(s)
|
||||||
}
|
}
|
||||||
}),
|
}),
|
||||||
);
|
);
|
||||||
@@ -277,10 +276,10 @@ impl HarmonyStreamingProcessor {
|
|||||||
)?;
|
)?;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(proto::generate_response::Response::Error(err)) => {
|
ProtoResponseVariant::Error(error_wrapper) => {
|
||||||
return Err(format!("Server error: {}", err.message));
|
return Err(format!("Server error: {}", error_wrapper.message()));
|
||||||
}
|
}
|
||||||
None => {}
|
ProtoResponseVariant::None => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -306,8 +305,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
|
|
||||||
/// Process streaming chunks from dual streams (prefill + decode)
|
/// Process streaming chunks from dual streams (prefill + decode)
|
||||||
async fn process_dual_stream(
|
async fn process_dual_stream(
|
||||||
mut prefill_stream: AbortOnDropStream,
|
mut prefill_stream: ProtoStream,
|
||||||
mut decode_stream: AbortOnDropStream,
|
mut decode_stream: ProtoStream,
|
||||||
dispatch: context::DispatchMetadata,
|
dispatch: context::DispatchMetadata,
|
||||||
original_request: Arc<ChatCompletionRequest>,
|
original_request: Arc<ChatCompletionRequest>,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
@@ -318,7 +317,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
while let Some(result) = prefill_stream.next().await {
|
while let Some(result) = prefill_stream.next().await {
|
||||||
let response = result.map_err(|e| format!("Prefill stream error: {}", e))?;
|
let response = result.map_err(|e| format!("Prefill stream error: {}", e))?;
|
||||||
|
|
||||||
if let Some(Complete(complete)) = response.response {
|
if let ProtoResponseVariant::Complete(complete_wrapper) = response.into_response() {
|
||||||
|
let complete = complete_wrapper.as_sglang();
|
||||||
prompt_tokens.insert(complete.index, complete.prompt_tokens as u32);
|
prompt_tokens.insert(complete.index, complete.prompt_tokens as u32);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -335,8 +335,9 @@ impl HarmonyStreamingProcessor {
|
|||||||
while let Some(result) = decode_stream.next().await {
|
while let Some(result) = decode_stream.next().await {
|
||||||
let response = result.map_err(|e| format!("Decode stream error: {}", e))?;
|
let response = result.map_err(|e| format!("Decode stream error: {}", e))?;
|
||||||
|
|
||||||
match response.response {
|
match response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk_wrapper) => {
|
||||||
|
let chunk = chunk_wrapper.as_sglang();
|
||||||
let index = chunk.index;
|
let index = chunk.index;
|
||||||
|
|
||||||
// Initialize parser for this index if needed
|
// Initialize parser for this index if needed
|
||||||
@@ -374,7 +375,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete_wrapper) => {
|
||||||
|
let complete = complete_wrapper.as_sglang();
|
||||||
let index = complete.index;
|
let index = complete.index;
|
||||||
|
|
||||||
finish_reasons.insert(index, Some(complete.finish_reason.clone()));
|
finish_reasons.insert(index, Some(complete.finish_reason.clone()));
|
||||||
@@ -408,10 +410,10 @@ impl HarmonyStreamingProcessor {
|
|||||||
)?;
|
)?;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(proto::generate_response::Response::Error(err)) => {
|
ProtoResponseVariant::Error(error_wrapper) => {
|
||||||
return Err(format!("Server error: {}", err.message));
|
return Err(format!("Server error: {}", error_wrapper.message()));
|
||||||
}
|
}
|
||||||
None => {}
|
ProtoResponseVariant::None => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -596,7 +598,7 @@ impl HarmonyStreamingProcessor {
|
|||||||
|
|
||||||
/// Process streaming chunks from a single stream
|
/// Process streaming chunks from a single stream
|
||||||
async fn process_responses_single_stream_mixed(
|
async fn process_responses_single_stream_mixed(
|
||||||
grpc_stream: AbortOnDropStream,
|
grpc_stream: ProtoStream,
|
||||||
emitter: &mut ResponseStreamEventEmitter,
|
emitter: &mut ResponseStreamEventEmitter,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
mcp_tool_names: &std::collections::HashSet<String>,
|
mcp_tool_names: &std::collections::HashSet<String>,
|
||||||
@@ -607,8 +609,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
|
|
||||||
/// Process streaming chunks from dual streams
|
/// Process streaming chunks from dual streams
|
||||||
async fn process_responses_dual_stream_mixed(
|
async fn process_responses_dual_stream_mixed(
|
||||||
mut prefill_stream: AbortOnDropStream,
|
mut prefill_stream: ProtoStream,
|
||||||
decode_stream: AbortOnDropStream,
|
decode_stream: ProtoStream,
|
||||||
emitter: &mut ResponseStreamEventEmitter,
|
emitter: &mut ResponseStreamEventEmitter,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
mcp_tool_names: &std::collections::HashSet<String>,
|
mcp_tool_names: &std::collections::HashSet<String>,
|
||||||
@@ -638,7 +640,7 @@ impl HarmonyStreamingProcessor {
|
|||||||
/// If mcp_tool_names is Some, determines mode per-tool by checking tool name.
|
/// If mcp_tool_names is Some, determines mode per-tool by checking tool name.
|
||||||
/// If mcp_tool_names is None, uses default MCP mode for all tools.
|
/// If mcp_tool_names is None, uses default MCP mode for all tools.
|
||||||
async fn process_decode_stream_with_tool_lookup(
|
async fn process_decode_stream_with_tool_lookup(
|
||||||
mut decode_stream: AbortOnDropStream,
|
mut decode_stream: ProtoStream,
|
||||||
emitter: &mut ResponseStreamEventEmitter,
|
emitter: &mut ResponseStreamEventEmitter,
|
||||||
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &mpsc::UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
mcp_tool_names: Option<&std::collections::HashSet<String>>,
|
mcp_tool_names: Option<&std::collections::HashSet<String>>,
|
||||||
@@ -674,8 +676,9 @@ impl HarmonyStreamingProcessor {
|
|||||||
chunk_count += 1;
|
chunk_count += 1;
|
||||||
let response = result.map_err(|e| format!("Decode stream error: {}", e))?;
|
let response = result.map_err(|e| format!("Decode stream error: {}", e))?;
|
||||||
|
|
||||||
match response.response {
|
match response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk_wrapper) => {
|
||||||
|
let chunk = chunk_wrapper.as_sglang();
|
||||||
// Parse chunk via Harmony parser
|
// Parse chunk via Harmony parser
|
||||||
let delta_result = parser
|
let delta_result = parser
|
||||||
.parse_chunk(&chunk.token_ids)
|
.parse_chunk(&chunk.token_ids)
|
||||||
@@ -847,7 +850,8 @@ impl HarmonyStreamingProcessor {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete_wrapper) => {
|
||||||
|
let complete = complete_wrapper.as_sglang();
|
||||||
// Store final metadata
|
// Store final metadata
|
||||||
finish_reason = complete.finish_reason.clone();
|
finish_reason = complete.finish_reason.clone();
|
||||||
matched_stop = complete.matched_stop.as_ref().map(|m| match m {
|
matched_stop = complete.matched_stop.as_ref().map(|m| match m {
|
||||||
@@ -956,10 +960,10 @@ impl HarmonyStreamingProcessor {
|
|||||||
emitter.send_event_best_effort(&event, tx);
|
emitter.send_event_best_effort(&event, tx);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(proto::generate_response::Response::Error(err)) => {
|
ProtoResponseVariant::Error(error_wrapper) => {
|
||||||
return Err(format!("Server error: {}", err.message));
|
return Err(format!("Server error: {}", error_wrapper.message()));
|
||||||
}
|
}
|
||||||
None => {}
|
ProtoResponseVariant::None => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,15 @@
|
|||||||
//! gRPC router implementations
|
//! gRPC router implementations
|
||||||
|
|
||||||
use crate::{grpc_client::proto, protocols::common::StringOrArray};
|
use crate::{grpc_client::sglang_proto::MultimodalInputs, protocols::common::StringOrArray};
|
||||||
|
|
||||||
|
pub mod client;
|
||||||
pub mod common;
|
pub mod common;
|
||||||
pub mod context;
|
pub mod context;
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod harmony;
|
pub mod harmony;
|
||||||
pub mod pd_router;
|
pub mod pd_router;
|
||||||
pub mod pipeline;
|
pub mod pipeline;
|
||||||
|
pub mod proto_wrapper;
|
||||||
pub mod regular;
|
pub mod regular;
|
||||||
pub mod router;
|
pub mod router;
|
||||||
pub mod utils;
|
pub mod utils;
|
||||||
@@ -16,6 +18,6 @@ pub mod utils;
|
|||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct ProcessedMessages {
|
pub struct ProcessedMessages {
|
||||||
pub text: String,
|
pub text: String,
|
||||||
pub multimodal_inputs: Option<proto::MultimodalInputs>,
|
pub multimodal_inputs: Option<MultimodalInputs>,
|
||||||
pub stop_sequences: Option<StringOrArray>,
|
pub stop_sequences: Option<StringOrArray>,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,381 @@
|
|||||||
|
//! Protocol buffer type wrappers for SGLang and vLLM backends
|
||||||
|
//!
|
||||||
|
//! This module provides unified enums that wrap proto types from both SGLang and vLLM,
|
||||||
|
//! allowing the router to work with either backend transparently.
|
||||||
|
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
|
use crate::grpc_client::{
|
||||||
|
sglang_proto::{self as sglang, generate_complete::MatchedStop},
|
||||||
|
sglang_scheduler::AbortOnDropStream as SglangStream,
|
||||||
|
vllm_engine::AbortOnDropStream as VllmStream,
|
||||||
|
vllm_proto as vllm,
|
||||||
|
};
|
||||||
|
|
||||||
|
/// Unified GenerateRequest that works with both backends
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoGenerateRequest {
|
||||||
|
Sglang(Box<sglang::GenerateRequest>),
|
||||||
|
Vllm(Box<vllm::GenerateRequest>),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoGenerateRequest {
|
||||||
|
/// Get SGLang variant (panics if vLLM)
|
||||||
|
pub fn as_sglang(&self) -> &sglang::GenerateRequest {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => req,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang GenerateRequest, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable SGLang variant (panics if vLLM)
|
||||||
|
pub fn as_sglang_mut(&mut self) -> &mut sglang::GenerateRequest {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => req,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang GenerateRequest, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get vLLM variant (panics if SGLang)
|
||||||
|
pub fn as_vllm(&self) -> &vllm::GenerateRequest {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(req) => req,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM GenerateRequest, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable vLLM variant (panics if SGLang)
|
||||||
|
pub fn as_vllm_mut(&mut self) -> &mut vllm::GenerateRequest {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(req) => req,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM GenerateRequest, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is SGLang
|
||||||
|
pub fn is_sglang(&self) -> bool {
|
||||||
|
matches!(self, Self::Sglang(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is vLLM
|
||||||
|
pub fn is_vllm(&self) -> bool {
|
||||||
|
matches!(self, Self::Vllm(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Clone the inner request (for passing to generate())
|
||||||
|
pub fn clone_inner(&self) -> Self {
|
||||||
|
self.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get request ID
|
||||||
|
pub fn request_id(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(req) => &req.request_id,
|
||||||
|
Self::Vllm(req) => &req.request_id,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified GenerateResponse from stream
|
||||||
|
pub enum ProtoGenerateResponse {
|
||||||
|
Sglang(sglang::GenerateResponse),
|
||||||
|
Vllm(vllm::GenerateResponse),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoGenerateResponse {
|
||||||
|
/// Get the response variant (chunk, complete, or error)
|
||||||
|
///
|
||||||
|
/// Consumes self to avoid cloning large proto messages in hot streaming path
|
||||||
|
pub fn into_response(self) -> ProtoResponseVariant {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(resp) => match resp.response {
|
||||||
|
Some(sglang::generate_response::Response::Chunk(chunk)) => {
|
||||||
|
ProtoResponseVariant::Chunk(ProtoGenerateStreamChunk::Sglang(chunk))
|
||||||
|
}
|
||||||
|
Some(sglang::generate_response::Response::Complete(complete)) => {
|
||||||
|
ProtoResponseVariant::Complete(ProtoGenerateComplete::Sglang(complete))
|
||||||
|
}
|
||||||
|
Some(sglang::generate_response::Response::Error(error)) => {
|
||||||
|
ProtoResponseVariant::Error(ProtoGenerateError::Sglang(error))
|
||||||
|
}
|
||||||
|
None => ProtoResponseVariant::None,
|
||||||
|
},
|
||||||
|
Self::Vllm(resp) => match resp.response {
|
||||||
|
Some(vllm::generate_response::Response::Chunk(chunk)) => {
|
||||||
|
ProtoResponseVariant::Chunk(ProtoGenerateStreamChunk::Vllm(chunk))
|
||||||
|
}
|
||||||
|
Some(vllm::generate_response::Response::Complete(complete)) => {
|
||||||
|
ProtoResponseVariant::Complete(ProtoGenerateComplete::Vllm(complete))
|
||||||
|
}
|
||||||
|
Some(vllm::generate_response::Response::Error(error)) => {
|
||||||
|
ProtoResponseVariant::Error(ProtoGenerateError::Vllm(error))
|
||||||
|
}
|
||||||
|
None => ProtoResponseVariant::None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Response variant extracted from GenerateResponse
|
||||||
|
pub enum ProtoResponseVariant {
|
||||||
|
Chunk(ProtoGenerateStreamChunk),
|
||||||
|
Complete(ProtoGenerateComplete),
|
||||||
|
Error(ProtoGenerateError),
|
||||||
|
None,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified GenerateStreamChunk
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoGenerateStreamChunk {
|
||||||
|
Sglang(sglang::GenerateStreamChunk),
|
||||||
|
Vllm(vllm::GenerateStreamChunk),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoGenerateStreamChunk {
|
||||||
|
/// Get SGLang variant (panics if vLLM)
|
||||||
|
pub fn as_sglang(&self) -> &sglang::GenerateStreamChunk {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(chunk) => chunk,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang GenerateStreamChunk, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get vLLM variant (panics if SGLang)
|
||||||
|
pub fn as_vllm(&self) -> &vllm::GenerateStreamChunk {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(chunk) => chunk,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM GenerateStreamChunk, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is SGLang
|
||||||
|
pub fn is_sglang(&self) -> bool {
|
||||||
|
matches!(self, Self::Sglang(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is vLLM
|
||||||
|
pub fn is_vllm(&self) -> bool {
|
||||||
|
matches!(self, Self::Vllm(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get token IDs from chunk (common field)
|
||||||
|
pub fn token_ids(&self) -> &[u32] {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => &c.token_ids,
|
||||||
|
Self::Vllm(c) => &c.token_ids,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get index (for n>1 support)
|
||||||
|
/// vLLM doesn't support n>1, so always returns 0
|
||||||
|
pub fn index(&self) -> u32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.index,
|
||||||
|
Self::Vllm(_) => 0, // vLLM doesn't support n>1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get output logprobs (SGLang only, returns None for vLLM)
|
||||||
|
pub fn output_logprobs(&self) -> Option<&sglang::OutputLogProbs> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.output_logprobs.as_ref(),
|
||||||
|
Self::Vllm(_) => None, // TODO: vLLM logprobs mapping
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get prompt tokens (cumulative)
|
||||||
|
pub fn prompt_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.prompt_tokens,
|
||||||
|
Self::Vllm(c) => c.prompt_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get completion tokens (cumulative)
|
||||||
|
pub fn completion_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.completion_tokens,
|
||||||
|
Self::Vllm(c) => c.completion_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get cached tokens (cumulative)
|
||||||
|
pub fn cached_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.cached_tokens,
|
||||||
|
Self::Vllm(c) => c.cached_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified GenerateComplete response
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoGenerateComplete {
|
||||||
|
Sglang(sglang::GenerateComplete),
|
||||||
|
Vllm(vllm::GenerateComplete),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoGenerateComplete {
|
||||||
|
/// Get SGLang variant (panics if vLLM)
|
||||||
|
pub fn as_sglang(&self) -> &sglang::GenerateComplete {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(complete) => complete,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang GenerateComplete, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get mutable SGLang variant (panics if vLLM)
|
||||||
|
pub fn as_sglang_mut(&mut self) -> &mut sglang::GenerateComplete {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(complete) => complete,
|
||||||
|
Self::Vllm(_) => panic!("Expected SGLang GenerateComplete, got vLLM"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get vLLM variant (panics if SGLang)
|
||||||
|
pub fn as_vllm(&self) -> &vllm::GenerateComplete {
|
||||||
|
match self {
|
||||||
|
Self::Vllm(complete) => complete,
|
||||||
|
Self::Sglang(_) => panic!("Expected vLLM GenerateComplete, got SGLang"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is SGLang
|
||||||
|
pub fn is_sglang(&self) -> bool {
|
||||||
|
matches!(self, Self::Sglang(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check if this is vLLM
|
||||||
|
pub fn is_vllm(&self) -> bool {
|
||||||
|
matches!(self, Self::Vllm(_))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get token IDs from either backend (output_ids in proto)
|
||||||
|
pub fn token_ids(&self) -> &[u32] {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => &c.output_ids,
|
||||||
|
Self::Vllm(c) => &c.output_ids,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get prompt tokens
|
||||||
|
pub fn prompt_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.prompt_tokens,
|
||||||
|
Self::Vllm(c) => c.prompt_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get completion tokens
|
||||||
|
pub fn completion_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.completion_tokens,
|
||||||
|
Self::Vllm(c) => c.completion_tokens,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get finish reason
|
||||||
|
pub fn finish_reason(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => &c.finish_reason,
|
||||||
|
Self::Vllm(c) => &c.finish_reason,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get index (for n>1 support)
|
||||||
|
/// vLLM doesn't support n>1, so always returns 0
|
||||||
|
pub fn index(&self) -> u32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.index,
|
||||||
|
Self::Vllm(_) => 0, // vLLM doesn't have index field (n>1 not supported)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get matched stop (SGLang only, returns oneof)
|
||||||
|
/// vLLM doesn't have matched_stop, returns None
|
||||||
|
pub fn matched_stop(&self) -> Option<&MatchedStop> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.matched_stop.as_ref(),
|
||||||
|
Self::Vllm(_) => None, // vLLM doesn't have matched_stop
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get output IDs (decode tokens only)
|
||||||
|
pub fn output_ids(&self) -> &[u32] {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => &c.output_ids,
|
||||||
|
Self::Vllm(c) => &c.output_ids,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get cached tokens
|
||||||
|
pub fn cached_tokens(&self) -> i32 {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.cached_tokens,
|
||||||
|
Self::Vllm(_) => 0, // vLLM doesn't have cached_tokens field
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get input logprobs (SGLang only)
|
||||||
|
pub fn input_logprobs(&self) -> Option<&sglang::InputLogProbs> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.input_logprobs.as_ref(),
|
||||||
|
Self::Vllm(_) => None, // vLLM doesn't have input_logprobs
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Get output logprobs
|
||||||
|
pub fn output_logprobs(&self) -> Option<&sglang::OutputLogProbs> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(c) => c.output_logprobs.as_ref(),
|
||||||
|
Self::Vllm(_) => None, // TODO: vLLM logprobs mapping
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified GenerateError
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum ProtoGenerateError {
|
||||||
|
Sglang(sglang::GenerateError),
|
||||||
|
Vllm(vllm::GenerateError),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoGenerateError {
|
||||||
|
/// Get error message
|
||||||
|
pub fn message(&self) -> &str {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(e) => &e.message,
|
||||||
|
Self::Vllm(e) => &e.message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unified stream wrapper
|
||||||
|
pub enum ProtoStream {
|
||||||
|
Sglang(SglangStream),
|
||||||
|
Vllm(VllmStream),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProtoStream {
|
||||||
|
/// Get next item from stream
|
||||||
|
pub async fn next(&mut self) -> Option<Result<ProtoGenerateResponse, tonic::Status>> {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(stream) => stream
|
||||||
|
.next()
|
||||||
|
.await
|
||||||
|
.map(|result| result.map(ProtoGenerateResponse::Sglang)),
|
||||||
|
Self::Vllm(stream) => stream
|
||||||
|
.next()
|
||||||
|
.await
|
||||||
|
.map(|result| result.map(ProtoGenerateResponse::Vllm)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Mark stream as completed (no abort needed)
|
||||||
|
pub fn mark_completed(&mut self) {
|
||||||
|
match self {
|
||||||
|
Self::Sglang(stream) => stream.mark_completed(),
|
||||||
|
Self::Vllm(stream) => stream.mark_completed(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -5,12 +5,11 @@
|
|||||||
|
|
||||||
use std::{sync::Arc, time::Instant};
|
use std::{sync::Arc, time::Instant};
|
||||||
|
|
||||||
use proto::generate_complete::MatchedStop;
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use tracing::error;
|
use tracing::error;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::proto,
|
grpc_client::sglang_proto::generate_complete::MatchedStop,
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatChoice, ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse},
|
chat::{ChatChoice, ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse},
|
||||||
common::{FunctionCallResponse, ToolCall, ToolChoice, ToolChoiceValue},
|
common::{FunctionCallResponse, ToolCall, ToolChoice, ToolChoiceValue},
|
||||||
@@ -20,7 +19,9 @@ use crate::{
|
|||||||
routers::grpc::{
|
routers::grpc::{
|
||||||
common::{response_collection, response_formatting},
|
common::{response_collection, response_formatting},
|
||||||
context::{DispatchMetadata, ExecutionResult},
|
context::{DispatchMetadata, ExecutionResult},
|
||||||
error, utils,
|
error,
|
||||||
|
proto_wrapper::ProtoGenerateComplete,
|
||||||
|
utils,
|
||||||
},
|
},
|
||||||
tokenizer::{
|
tokenizer::{
|
||||||
stop::{SequenceDecoderOutput, StopSequenceDecoder},
|
stop::{SequenceDecoderOutput, StopSequenceDecoder},
|
||||||
@@ -60,7 +61,7 @@ impl ResponseProcessor {
|
|||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
pub async fn process_single_choice(
|
pub async fn process_single_choice(
|
||||||
&self,
|
&self,
|
||||||
complete: &proto::GenerateComplete,
|
complete: &ProtoGenerateComplete,
|
||||||
index: usize,
|
index: usize,
|
||||||
original_request: &ChatCompletionRequest,
|
original_request: &ChatCompletionRequest,
|
||||||
stop_decoder: &mut StopSequenceDecoder,
|
stop_decoder: &mut StopSequenceDecoder,
|
||||||
@@ -71,7 +72,7 @@ impl ResponseProcessor {
|
|||||||
stop_decoder.reset();
|
stop_decoder.reset();
|
||||||
// Decode tokens
|
// Decode tokens
|
||||||
let outputs = stop_decoder
|
let outputs = stop_decoder
|
||||||
.process_tokens(&complete.output_ids)
|
.process_tokens(complete.output_ids())
|
||||||
.map_err(|e| format!("Failed to process tokens: {}", e))?;
|
.map_err(|e| format!("Failed to process tokens: {}", e))?;
|
||||||
|
|
||||||
// Accumulate text with early breaks
|
// Accumulate text with early breaks
|
||||||
@@ -153,7 +154,7 @@ impl ResponseProcessor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Step 3: Use finish reason directly from proto (already OpenAI-compatible string)
|
// Step 3: Use finish reason directly from proto (already OpenAI-compatible string)
|
||||||
let finish_reason_str = &complete.finish_reason;
|
let finish_reason_str = complete.finish_reason();
|
||||||
|
|
||||||
// Override finish reason if we have tool calls
|
// Override finish reason if we have tool calls
|
||||||
let final_finish_reason_str = if tool_calls.is_some() {
|
let final_finish_reason_str = if tool_calls.is_some() {
|
||||||
@@ -163,7 +164,7 @@ impl ResponseProcessor {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Extract matched_stop information from proto
|
// Extract matched_stop information from proto
|
||||||
let matched_stop = match &complete.matched_stop {
|
let matched_stop = match complete.matched_stop() {
|
||||||
Some(MatchedStop::MatchedTokenId(token_id)) => {
|
Some(MatchedStop::MatchedTokenId(token_id)) => {
|
||||||
Some(Value::Number(serde_json::Number::from(*token_id)))
|
Some(Value::Number(serde_json::Number::from(*token_id)))
|
||||||
}
|
}
|
||||||
@@ -172,7 +173,7 @@ impl ResponseProcessor {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Step 4: Convert output logprobs if present
|
// Step 4: Convert output logprobs if present
|
||||||
let logprobs = if let Some(proto_logprobs) = &complete.output_logprobs {
|
let logprobs = if let Some(proto_logprobs) = complete.output_logprobs() {
|
||||||
match utils::convert_proto_to_openai_logprobs(proto_logprobs, &self.tokenizer) {
|
match utils::convert_proto_to_openai_logprobs(proto_logprobs, &self.tokenizer) {
|
||||||
Ok(logprobs) => Some(logprobs),
|
Ok(logprobs) => Some(logprobs),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -370,11 +371,11 @@ impl ResponseProcessor {
|
|||||||
|
|
||||||
// Process each completion
|
// Process each completion
|
||||||
let mut result_array = Vec::new();
|
let mut result_array = Vec::new();
|
||||||
for mut complete in all_responses {
|
for complete in all_responses {
|
||||||
stop_decoder.reset();
|
stop_decoder.reset();
|
||||||
|
|
||||||
// Process tokens through stop decoder
|
// Process tokens through stop decoder
|
||||||
let outputs = match stop_decoder.process_tokens(&complete.output_ids) {
|
let outputs = match stop_decoder.process_tokens(complete.output_ids()) {
|
||||||
Ok(outputs) => outputs,
|
Ok(outputs) => outputs,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
return Err(error::internal_error(format!(
|
return Err(error::internal_error(format!(
|
||||||
@@ -403,15 +404,15 @@ impl ResponseProcessor {
|
|||||||
decoded_text.push_str(&t);
|
decoded_text.push_str(&t);
|
||||||
}
|
}
|
||||||
|
|
||||||
let output_ids = std::mem::take(&mut complete.output_ids);
|
let output_ids = complete.output_ids().to_vec();
|
||||||
let finish_reason_str = complete.finish_reason.to_string();
|
let finish_reason_str = complete.finish_reason();
|
||||||
|
|
||||||
// Parse finish_reason from string to proper type
|
// Parse finish_reason from string to proper type
|
||||||
let finish_reason =
|
let finish_reason =
|
||||||
utils::parse_finish_reason(&finish_reason_str, complete.completion_tokens);
|
utils::parse_finish_reason(finish_reason_str, complete.completion_tokens());
|
||||||
|
|
||||||
// Handle matched_stop if present
|
// Handle matched_stop if present
|
||||||
let matched_stop = complete.matched_stop.take().map(|matched| match matched {
|
let matched_stop = complete.matched_stop().map(|matched| match matched {
|
||||||
MatchedStop::MatchedTokenId(id) => serde_json::json!(id),
|
MatchedStop::MatchedTokenId(id) => serde_json::json!(id),
|
||||||
MatchedStop::MatchedStopStr(s) => serde_json::json!(s),
|
MatchedStop::MatchedStopStr(s) => serde_json::json!(s),
|
||||||
});
|
});
|
||||||
@@ -419,8 +420,7 @@ impl ResponseProcessor {
|
|||||||
// Extract logprobs if requested (convert proto types to Generate format)
|
// Extract logprobs if requested (convert proto types to Generate format)
|
||||||
let input_token_logprobs = if request_logprobs {
|
let input_token_logprobs = if request_logprobs {
|
||||||
complete
|
complete
|
||||||
.input_logprobs
|
.input_logprobs()
|
||||||
.as_ref()
|
|
||||||
.map(utils::convert_generate_input_logprobs)
|
.map(utils::convert_generate_input_logprobs)
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
@@ -428,8 +428,7 @@ impl ResponseProcessor {
|
|||||||
|
|
||||||
let output_token_logprobs = if request_logprobs {
|
let output_token_logprobs = if request_logprobs {
|
||||||
complete
|
complete
|
||||||
.output_logprobs
|
.output_logprobs()
|
||||||
.as_ref()
|
|
||||||
.map(utils::convert_generate_output_logprobs)
|
.map(utils::convert_generate_output_logprobs)
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
@@ -439,15 +438,15 @@ impl ResponseProcessor {
|
|||||||
let meta_info = GenerateMetaInfo {
|
let meta_info = GenerateMetaInfo {
|
||||||
id: dispatch.request_id.clone(),
|
id: dispatch.request_id.clone(),
|
||||||
finish_reason,
|
finish_reason,
|
||||||
prompt_tokens: complete.prompt_tokens as u32,
|
prompt_tokens: complete.prompt_tokens() as u32,
|
||||||
weight_version: dispatch
|
weight_version: dispatch
|
||||||
.weight_version
|
.weight_version
|
||||||
.clone()
|
.clone()
|
||||||
.unwrap_or_else(|| "default".to_string()),
|
.unwrap_or_else(|| "default".to_string()),
|
||||||
input_token_logprobs,
|
input_token_logprobs,
|
||||||
output_token_logprobs,
|
output_token_logprobs,
|
||||||
completion_tokens: complete.completion_tokens as u32,
|
completion_tokens: complete.completion_tokens() as u32,
|
||||||
cached_tokens: complete.cached_tokens as u32,
|
cached_tokens: complete.cached_tokens() as u32,
|
||||||
e2e_latency: start_time.elapsed().as_secs_f64(),
|
e2e_latency: start_time.elapsed().as_secs_f64(),
|
||||||
matched_stop,
|
matched_stop,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -6,9 +6,11 @@ use tracing::error;
|
|||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::routers::grpc::{
|
use crate::routers::grpc::{
|
||||||
|
client::GrpcClient,
|
||||||
common::stages::{helpers, PipelineStage},
|
common::stages::{helpers, PipelineStage},
|
||||||
context::{ClientSelection, RequestContext, WorkerSelection},
|
context::{ClientSelection, RequestContext, WorkerSelection},
|
||||||
error,
|
error,
|
||||||
|
proto_wrapper::ProtoGenerateRequest,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Chat request building stage
|
/// Chat request building stage
|
||||||
@@ -55,23 +57,44 @@ impl PipelineStage for ChatRequestBuildingStage {
|
|||||||
let request_id = format!("chatcmpl-{}", Uuid::new_v4());
|
let request_id = format!("chatcmpl-{}", Uuid::new_v4());
|
||||||
let body_ref = prep.filtered_request.as_ref().unwrap_or(&chat_request);
|
let body_ref = prep.filtered_request.as_ref().unwrap_or(&chat_request);
|
||||||
|
|
||||||
let mut proto_request = builder_client
|
// Dispatch to the appropriate client based on backend type
|
||||||
.build_generate_request_from_chat(
|
let mut proto_request = match builder_client {
|
||||||
request_id,
|
GrpcClient::Sglang(sglang_client) => {
|
||||||
body_ref,
|
let req = sglang_client
|
||||||
prep.processed_messages.as_ref().unwrap().text.clone(),
|
.build_generate_request_from_chat(
|
||||||
prep.token_ids.clone(),
|
request_id,
|
||||||
prep.processed_messages
|
body_ref,
|
||||||
.as_ref()
|
prep.processed_messages.as_ref().unwrap().text.clone(),
|
||||||
.unwrap()
|
prep.token_ids.clone(),
|
||||||
.multimodal_inputs
|
prep.processed_messages
|
||||||
.clone(),
|
.as_ref()
|
||||||
prep.tool_constraints.clone(),
|
.unwrap()
|
||||||
)
|
.multimodal_inputs
|
||||||
.map_err(|e| {
|
.clone(),
|
||||||
error!(function = "ChatRequestBuildingStage::execute", error = %e, "Failed to build generate request");
|
prep.tool_constraints.clone(),
|
||||||
error::bad_request(format!("Invalid request parameters: {}", e))
|
)
|
||||||
})?;
|
.map_err(|e| {
|
||||||
|
error!(function = "ChatRequestBuildingStage::execute", error = %e, "Failed to build SGLang generate request");
|
||||||
|
error::bad_request(format!("Invalid request parameters: {}", e))
|
||||||
|
})?;
|
||||||
|
ProtoGenerateRequest::Sglang(Box::new(req))
|
||||||
|
}
|
||||||
|
GrpcClient::Vllm(vllm_client) => {
|
||||||
|
let req = vllm_client
|
||||||
|
.build_generate_request_from_chat(
|
||||||
|
request_id,
|
||||||
|
body_ref,
|
||||||
|
prep.processed_messages.as_ref().unwrap().text.clone(),
|
||||||
|
prep.token_ids.clone(),
|
||||||
|
prep.tool_constraints.clone(),
|
||||||
|
)
|
||||||
|
.map_err(|e| {
|
||||||
|
error!(function = "ChatRequestBuildingStage::execute", error = %e, "Failed to build vLLM generate request");
|
||||||
|
error::bad_request(format!("Invalid request parameters: {}", e))
|
||||||
|
})?;
|
||||||
|
ProtoGenerateRequest::Vllm(Box::new(req))
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
// Inject PD metadata if needed
|
// Inject PD metadata if needed
|
||||||
if self.inject_pd_metadata {
|
if self.inject_pd_metadata {
|
||||||
|
|||||||
@@ -6,9 +6,11 @@ use tracing::error;
|
|||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::routers::grpc::{
|
use crate::routers::grpc::{
|
||||||
|
client::GrpcClient,
|
||||||
common::stages::{helpers, PipelineStage},
|
common::stages::{helpers, PipelineStage},
|
||||||
context::{ClientSelection, RequestContext, WorkerSelection},
|
context::{ClientSelection, RequestContext, WorkerSelection},
|
||||||
error,
|
error,
|
||||||
|
proto_wrapper::ProtoGenerateRequest,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Generate request building stage
|
/// Generate request building stage
|
||||||
@@ -57,17 +59,37 @@ impl PipelineStage for GenerateRequestBuildingStage {
|
|||||||
.clone()
|
.clone()
|
||||||
.unwrap_or_else(|| format!("gen-{}", Uuid::new_v4()));
|
.unwrap_or_else(|| format!("gen-{}", Uuid::new_v4()));
|
||||||
|
|
||||||
let mut proto_request = builder_client
|
// Dispatch to the appropriate client based on backend type
|
||||||
.build_plain_generate_request(
|
let mut proto_request = match builder_client {
|
||||||
request_id,
|
GrpcClient::Sglang(sglang_client) => {
|
||||||
&generate_request,
|
let req = sglang_client
|
||||||
prep.original_text.clone(),
|
.build_plain_generate_request(
|
||||||
prep.token_ids.clone(),
|
request_id,
|
||||||
)
|
&generate_request,
|
||||||
.map_err(|e| {
|
prep.original_text.clone(),
|
||||||
error!(function = "GenerateRequestBuildingStage::execute", error = %e, "Failed to build generate request");
|
prep.token_ids.clone(),
|
||||||
error::bad_request(e)
|
)
|
||||||
})?;
|
.map_err(|e| {
|
||||||
|
error!(function = "GenerateRequestBuildingStage::execute", error = %e, "Failed to build SGLang generate request");
|
||||||
|
error::bad_request(e)
|
||||||
|
})?;
|
||||||
|
ProtoGenerateRequest::Sglang(Box::new(req))
|
||||||
|
}
|
||||||
|
GrpcClient::Vllm(vllm_client) => {
|
||||||
|
let req = vllm_client
|
||||||
|
.build_plain_generate_request(
|
||||||
|
request_id,
|
||||||
|
&generate_request,
|
||||||
|
prep.original_text.clone(),
|
||||||
|
prep.token_ids.clone(),
|
||||||
|
)
|
||||||
|
.map_err(|e| {
|
||||||
|
error!(function = "GenerateRequestBuildingStage::execute", error = %e, "Failed to build vLLM generate request");
|
||||||
|
error::bad_request(e)
|
||||||
|
})?;
|
||||||
|
ProtoGenerateRequest::Vllm(Box::new(req))
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
// Inject PD metadata if needed
|
// Inject PD metadata if needed
|
||||||
if self.inject_pd_metadata {
|
if self.inject_pd_metadata {
|
||||||
|
|||||||
@@ -5,11 +5,13 @@
|
|||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
|
use tracing::error;
|
||||||
|
|
||||||
use super::{chat::ChatPreparationStage, generate::GeneratePreparationStage};
|
use super::{chat::ChatPreparationStage, generate::GeneratePreparationStage};
|
||||||
use crate::routers::grpc::{
|
use crate::routers::grpc::{
|
||||||
common::stages::PipelineStage,
|
common::stages::PipelineStage,
|
||||||
context::{RequestContext, RequestType},
|
context::{RequestContext, RequestType},
|
||||||
|
error as grpc_error,
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Preparation stage (delegates to endpoint-specific implementations)
|
/// Preparation stage (delegates to endpoint-specific implementations)
|
||||||
@@ -40,8 +42,13 @@ impl PipelineStage for PreparationStage {
|
|||||||
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
||||||
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
||||||
RequestType::Responses(_) => {
|
RequestType::Responses(_) => {
|
||||||
// Responses API has its own preparation handled elsewhere
|
error!(
|
||||||
Ok(None)
|
function = "PreparationStage::execute",
|
||||||
|
"RequestType::Responses reached regular preparation stage"
|
||||||
|
);
|
||||||
|
Err(grpc_error::internal_error(
|
||||||
|
"RequestType::Responses reached regular preparation stage",
|
||||||
|
))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,15 +2,13 @@
|
|||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use uuid::Uuid;
|
use tracing::error;
|
||||||
|
|
||||||
use super::{chat::ChatRequestBuildingStage, generate::GenerateRequestBuildingStage};
|
use super::{chat::ChatRequestBuildingStage, generate::GenerateRequestBuildingStage};
|
||||||
use crate::{
|
use crate::routers::grpc::{
|
||||||
grpc_client::proto,
|
common::stages::PipelineStage,
|
||||||
routers::grpc::{
|
context::{RequestContext, RequestType},
|
||||||
common::stages::PipelineStage,
|
error as grpc_error,
|
||||||
context::{RequestContext, RequestType},
|
|
||||||
},
|
|
||||||
};
|
};
|
||||||
|
|
||||||
/// Request building stage (delegates to endpoint-specific implementations)
|
/// Request building stage (delegates to endpoint-specific implementations)
|
||||||
@@ -35,15 +33,13 @@ impl PipelineStage for RequestBuildingStage {
|
|||||||
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
RequestType::Chat(_) => self.chat_stage.execute(ctx).await,
|
||||||
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
RequestType::Generate(_) => self.generate_stage.execute(ctx).await,
|
||||||
RequestType::Responses(_request) => {
|
RequestType::Responses(_request) => {
|
||||||
// Responses API builds request during the MCP loop
|
error!(
|
||||||
// For now, create minimal request - responses handler will populate it
|
function = "RequestBuildingStage::execute",
|
||||||
let request_id = format!("resp-{}", Uuid::new_v4());
|
"RequestType::Responses reached regular request building stage"
|
||||||
|
);
|
||||||
ctx.state.proto_request = Some(proto::GenerateRequest {
|
Err(grpc_error::internal_error(
|
||||||
request_id,
|
"RequestType::Responses reached regular request building stage",
|
||||||
..Default::default()
|
))
|
||||||
});
|
|
||||||
Ok(None)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,10 +44,10 @@ impl PipelineStage for ResponseProcessingStage {
|
|||||||
RequestType::Responses(_) => {
|
RequestType::Responses(_) => {
|
||||||
error!(
|
error!(
|
||||||
function = "ResponseProcessingStage::execute",
|
function = "ResponseProcessingStage::execute",
|
||||||
"Responses API not supported in regular pipeline"
|
"RequestType::Responses reached regular response processing stage"
|
||||||
);
|
);
|
||||||
Err(error::bad_request(
|
Err(error::internal_error(
|
||||||
"Responses API processing must be handled by responses handler".to_string(),
|
"RequestType::Responses reached regular response processing stage",
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,17 +7,13 @@ use std::{collections::HashMap, io, sync::Arc, time::Instant};
|
|||||||
use axum::{body::Body, http::StatusCode, response::Response};
|
use axum::{body::Body, http::StatusCode, response::Response};
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use http::header::{HeaderValue, CONTENT_TYPE};
|
use http::header::{HeaderValue, CONTENT_TYPE};
|
||||||
use proto::{
|
|
||||||
generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
|
||||||
generate_response::Response::{Chunk, Complete, Error},
|
|
||||||
};
|
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
use tokio::sync::{mpsc, mpsc::UnboundedSender};
|
use tokio::sync::{mpsc, mpsc::UnboundedSender};
|
||||||
use tokio_stream::{wrappers::UnboundedReceiverStream, StreamExt};
|
use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||||
use tracing::{debug, error, warn};
|
use tracing::{debug, error, warn};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::{proto, sglang_scheduler::AbortOnDropStream},
|
grpc_client::sglang_proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatCompletionStreamResponse},
|
chat::{ChatCompletionRequest, ChatCompletionStreamResponse},
|
||||||
common::{
|
common::{
|
||||||
@@ -27,7 +23,11 @@ use crate::{
|
|||||||
generate::GenerateRequest,
|
generate::GenerateRequest,
|
||||||
},
|
},
|
||||||
reasoning_parser::{ParserFactory as ReasoningParserFactory, ParserResult, ReasoningParser},
|
reasoning_parser::{ParserFactory as ReasoningParserFactory, ParserResult, ReasoningParser},
|
||||||
routers::grpc::{context, utils},
|
routers::grpc::{
|
||||||
|
context,
|
||||||
|
proto_wrapper::{ProtoResponseVariant, ProtoStream},
|
||||||
|
utils,
|
||||||
|
},
|
||||||
tokenizer::{
|
tokenizer::{
|
||||||
stop::{SequenceDecoderOutput, StopSequenceDecoder},
|
stop::{SequenceDecoderOutput, StopSequenceDecoder},
|
||||||
traits::Tokenizer,
|
traits::Tokenizer,
|
||||||
@@ -158,7 +158,7 @@ impl StreamingProcessor {
|
|||||||
/// Process streaming chunks from a single stream (Regular mode)
|
/// Process streaming chunks from a single stream (Regular mode)
|
||||||
pub async fn process_streaming_chunks(
|
pub async fn process_streaming_chunks(
|
||||||
&self,
|
&self,
|
||||||
mut grpc_stream: AbortOnDropStream,
|
mut grpc_stream: ProtoStream,
|
||||||
dispatch: context::DispatchMetadata,
|
dispatch: context::DispatchMetadata,
|
||||||
stop_params: (Option<StringOrArray>, Option<Vec<u32>>, bool, bool),
|
stop_params: (Option<StringOrArray>, Option<Vec<u32>>, bool, bool),
|
||||||
original_request: Arc<ChatCompletionRequest>,
|
original_request: Arc<ChatCompletionRequest>,
|
||||||
@@ -244,9 +244,16 @@ impl StreamingProcessor {
|
|||||||
while let Some(response) = grpc_stream.next().await {
|
while let Some(response) = grpc_stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
||||||
|
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
let index = chunk.index;
|
let index = chunk.index();
|
||||||
|
|
||||||
|
// For vLLM, accumulate completion tokens (vLLM sends deltas)
|
||||||
|
// For SGLang, skip (SGLang sends cumulative values)
|
||||||
|
if chunk.is_vllm() {
|
||||||
|
let tokens_count = completion_tokens.entry(index).or_insert(0);
|
||||||
|
*tokens_count += chunk.token_ids().len() as u32;
|
||||||
|
}
|
||||||
|
|
||||||
// Get or create stop decoder for this index
|
// Get or create stop decoder for this index
|
||||||
let stop_decoder = stop_decoders.entry(index).or_insert_with(|| {
|
let stop_decoder = stop_decoders.entry(index).or_insert_with(|| {
|
||||||
@@ -263,14 +270,14 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
// Process tokens through stop decoder
|
// Process tokens through stop decoder
|
||||||
let (chunk_text, _should_stop) =
|
let (chunk_text, _should_stop) =
|
||||||
Self::process_chunk_tokens(stop_decoder, &chunk.token_ids);
|
Self::process_chunk_tokens(stop_decoder, chunk.token_ids());
|
||||||
|
|
||||||
if chunk_text.is_empty() {
|
if chunk_text.is_empty() {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Process logprobs if present
|
// Process logprobs if present
|
||||||
let choice_logprobs = if let Some(ref proto_logprobs) = chunk.output_logprobs {
|
let choice_logprobs = if let Some(proto_logprobs) = chunk.output_logprobs() {
|
||||||
match utils::convert_proto_to_openai_logprobs(
|
match utils::convert_proto_to_openai_logprobs(
|
||||||
proto_logprobs,
|
proto_logprobs,
|
||||||
&self.tokenizer,
|
&self.tokenizer,
|
||||||
@@ -398,8 +405,8 @@ impl StreamingProcessor {
|
|||||||
.map_err(|_| "Failed to send content chunk".to_string())?;
|
.map_err(|_| "Failed to send content chunk".to_string())?;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete) => {
|
||||||
let index = complete.index;
|
let index = complete.index();
|
||||||
|
|
||||||
// Flush any remaining text for this index's stop_decoder
|
// Flush any remaining text for this index's stop_decoder
|
||||||
if let Some(decoder) = stop_decoders.get_mut(&index) {
|
if let Some(decoder) = stop_decoders.get_mut(&index) {
|
||||||
@@ -428,13 +435,21 @@ impl StreamingProcessor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Store metadata
|
// Store metadata
|
||||||
prompt_tokens.insert(index, complete.prompt_tokens as u32);
|
prompt_tokens.insert(index, complete.prompt_tokens() as u32);
|
||||||
completion_tokens.insert(index, complete.completion_tokens as u32);
|
|
||||||
cached_tokens.insert(index, complete.cached_tokens as u32);
|
// For vLLM, use accumulated count (we tracked deltas)
|
||||||
finish_reasons.insert(index, complete.finish_reason.clone());
|
// For SGLang, use complete value (already cumulative)
|
||||||
|
if complete.is_vllm() {
|
||||||
|
completion_tokens.entry(index).or_insert(0);
|
||||||
|
} else {
|
||||||
|
completion_tokens.insert(index, complete.completion_tokens() as u32);
|
||||||
|
}
|
||||||
|
|
||||||
|
cached_tokens.insert(index, complete.cached_tokens() as u32);
|
||||||
|
finish_reasons.insert(index, complete.finish_reason().to_string());
|
||||||
|
|
||||||
// Extract matched_stop
|
// Extract matched_stop
|
||||||
let matched_stop_value = match &complete.matched_stop {
|
let matched_stop_value = match complete.matched_stop() {
|
||||||
Some(MatchedTokenId(token_id)) => {
|
Some(MatchedTokenId(token_id)) => {
|
||||||
Some(Value::Number(serde_json::Number::from(*token_id)))
|
Some(Value::Number(serde_json::Number::from(*token_id)))
|
||||||
}
|
}
|
||||||
@@ -445,10 +460,10 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
// Don't break - continue reading all Complete messages for n>1
|
// Don't break - continue reading all Complete messages for n>1
|
||||||
}
|
}
|
||||||
Some(Error(error)) => {
|
ProtoResponseVariant::Error(error) => {
|
||||||
return Err(error.message);
|
return Err(error.message().to_string());
|
||||||
}
|
}
|
||||||
None => continue,
|
ProtoResponseVariant::None => continue,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -541,8 +556,8 @@ impl StreamingProcessor {
|
|||||||
/// Process dual streaming chunks (prefill + decode) - PD mode
|
/// Process dual streaming chunks (prefill + decode) - PD mode
|
||||||
pub async fn process_dual_streaming_chunks(
|
pub async fn process_dual_streaming_chunks(
|
||||||
&self,
|
&self,
|
||||||
mut prefill_stream: AbortOnDropStream,
|
mut prefill_stream: ProtoStream,
|
||||||
decode_stream: AbortOnDropStream,
|
decode_stream: ProtoStream,
|
||||||
dispatch: context::DispatchMetadata,
|
dispatch: context::DispatchMetadata,
|
||||||
stop_params: (Option<StringOrArray>, Option<Vec<u32>>, bool, bool),
|
stop_params: (Option<StringOrArray>, Option<Vec<u32>>, bool, bool),
|
||||||
original_request: Arc<ChatCompletionRequest>,
|
original_request: Arc<ChatCompletionRequest>,
|
||||||
@@ -552,14 +567,14 @@ impl StreamingProcessor {
|
|||||||
if original_request.logprobs {
|
if original_request.logprobs {
|
||||||
while let Some(response) = prefill_stream.next().await {
|
while let Some(response) = prefill_stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Complete(_complete)) => {
|
ProtoResponseVariant::Complete(_complete) => {
|
||||||
// Input logprobs collected but not yet used in streaming
|
// Input logprobs collected but not yet used in streaming
|
||||||
// (OpenAI spec doesn't require prompt logprobs in streaming responses)
|
// (OpenAI spec doesn't require prompt logprobs in streaming responses)
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
Some(Error(error)) => {
|
ProtoResponseVariant::Error(error) => {
|
||||||
return Err(format!("Prefill error: {}", error.message));
|
return Err(format!("Prefill error: {}", error.message()));
|
||||||
}
|
}
|
||||||
_ => continue,
|
_ => continue,
|
||||||
}
|
}
|
||||||
@@ -661,7 +676,7 @@ impl StreamingProcessor {
|
|||||||
/// Process streaming chunks for generate endpoint (no tool/reasoning parsing)
|
/// Process streaming chunks for generate endpoint (no tool/reasoning parsing)
|
||||||
async fn process_generate_streaming(
|
async fn process_generate_streaming(
|
||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut stream: AbortOnDropStream,
|
mut stream: ProtoStream,
|
||||||
request_id: String,
|
request_id: String,
|
||||||
weight_version: String,
|
weight_version: String,
|
||||||
_include_logprobs: bool,
|
_include_logprobs: bool,
|
||||||
@@ -676,16 +691,19 @@ impl StreamingProcessor {
|
|||||||
while let Some(response) = stream.next().await {
|
while let Some(response) = stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
||||||
|
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
let index = chunk.index;
|
let index = chunk.index();
|
||||||
|
|
||||||
// Update completion tokens for this index
|
// Both backends send delta token_ids, so accumulate for both
|
||||||
let completion_tokens = completion_tokens_map.entry(index).or_insert(0);
|
let completion_tokens = completion_tokens_map.entry(index).or_insert(0);
|
||||||
*completion_tokens += chunk.token_ids.len() as u32;
|
*completion_tokens += chunk.token_ids().len() as u32;
|
||||||
|
let current_completion_tokens = *completion_tokens;
|
||||||
|
|
||||||
// Decode tokens to text (skip_special_tokens=true to handle newlines correctly)
|
// Decode tokens to text (skip_special_tokens=true to handle newlines correctly)
|
||||||
let chunk_text = tokenizer.decode(&chunk.token_ids, true).unwrap_or_default();
|
let chunk_text = tokenizer
|
||||||
|
.decode(chunk.token_ids(), true)
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
// Accumulate text for this index
|
// Accumulate text for this index
|
||||||
let accumulated_text = accumulated_texts.entry(index).or_default();
|
let accumulated_text = accumulated_texts.entry(index).or_default();
|
||||||
@@ -697,14 +715,14 @@ impl StreamingProcessor {
|
|||||||
// Build streaming response chunk (SGLang format)
|
// Build streaming response chunk (SGLang format)
|
||||||
let chunk_response = serde_json::json!({
|
let chunk_response = serde_json::json!({
|
||||||
"text": accumulated_text.clone(),
|
"text": accumulated_text.clone(),
|
||||||
"output_ids": chunk.token_ids,
|
"output_ids": chunk.token_ids(),
|
||||||
"meta_info": {
|
"meta_info": {
|
||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": null,
|
"finish_reason": null,
|
||||||
"prompt_tokens": chunk.prompt_tokens,
|
"prompt_tokens": chunk.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &weight_version,
|
||||||
"completion_tokens": *completion_tokens,
|
"completion_tokens": current_completion_tokens,
|
||||||
"cached_tokens": chunk.cached_tokens
|
"cached_tokens": chunk.cached_tokens()
|
||||||
},
|
},
|
||||||
"index": index
|
"index": index
|
||||||
});
|
});
|
||||||
@@ -716,8 +734,8 @@ impl StreamingProcessor {
|
|||||||
tx.send(Ok(Bytes::from(sse_chunk)))
|
tx.send(Ok(Bytes::from(sse_chunk)))
|
||||||
.map_err(|_| "Failed to send chunk".to_string())?;
|
.map_err(|_| "Failed to send chunk".to_string())?;
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete) => {
|
||||||
let index = complete.index;
|
let index = complete.index();
|
||||||
let accumulated_text =
|
let accumulated_text =
|
||||||
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
||||||
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
||||||
@@ -727,14 +745,14 @@ impl StreamingProcessor {
|
|||||||
// Send final chunk with finish_reason
|
// Send final chunk with finish_reason
|
||||||
let finish_response = serde_json::json!({
|
let finish_response = serde_json::json!({
|
||||||
"text": accumulated_text,
|
"text": accumulated_text,
|
||||||
"output_ids": complete.output_ids[complete.output_ids.len().saturating_sub(1)..].to_vec(),
|
"output_ids": complete.output_ids()[complete.output_ids().len().saturating_sub(1)..].to_vec(),
|
||||||
"meta_info": {
|
"meta_info": {
|
||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": complete.finish_reason,
|
"finish_reason": complete.finish_reason(),
|
||||||
"prompt_tokens": complete.prompt_tokens,
|
"prompt_tokens": complete.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &weight_version,
|
||||||
"completion_tokens": completion_tokens,
|
"completion_tokens": completion_tokens,
|
||||||
"cached_tokens": complete.cached_tokens,
|
"cached_tokens": complete.cached_tokens(),
|
||||||
"e2e_latency": e2e_latency
|
"e2e_latency": e2e_latency
|
||||||
},
|
},
|
||||||
"index": index
|
"index": index
|
||||||
@@ -749,10 +767,10 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
// Continue to process all completions if n>1
|
// Continue to process all completions if n>1
|
||||||
}
|
}
|
||||||
Some(Error(error)) => {
|
ProtoResponseVariant::Error(error) => {
|
||||||
return Err(error.message);
|
return Err(error.message().to_string());
|
||||||
}
|
}
|
||||||
None => continue,
|
ProtoResponseVariant::None => continue,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -765,8 +783,8 @@ impl StreamingProcessor {
|
|||||||
/// Process dual streaming for generate endpoint (PD mode with logprobs support)
|
/// Process dual streaming for generate endpoint (PD mode with logprobs support)
|
||||||
async fn process_generate_streaming_dual(
|
async fn process_generate_streaming_dual(
|
||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut prefill_stream: AbortOnDropStream,
|
mut prefill_stream: ProtoStream,
|
||||||
decode_stream: AbortOnDropStream,
|
decode_stream: ProtoStream,
|
||||||
request_id: String,
|
request_id: String,
|
||||||
weight_version: String,
|
weight_version: String,
|
||||||
return_logprob: bool,
|
return_logprob: bool,
|
||||||
@@ -777,17 +795,16 @@ impl StreamingProcessor {
|
|||||||
let mut input_logprobs = None;
|
let mut input_logprobs = None;
|
||||||
while let Some(response) = prefill_stream.next().await {
|
while let Some(response) = prefill_stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete) => {
|
||||||
// Extract input_logprobs from prefill Complete message (convert proto to SGLang format)
|
// Extract input_logprobs from prefill Complete message (convert proto to SGLang format)
|
||||||
input_logprobs = complete
|
input_logprobs = complete
|
||||||
.input_logprobs
|
.input_logprobs()
|
||||||
.as_ref()
|
|
||||||
.map(utils::convert_generate_input_logprobs);
|
.map(utils::convert_generate_input_logprobs);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
Some(Error(error)) => {
|
ProtoResponseVariant::Error(error) => {
|
||||||
return Err(format!("Prefill error: {}", error.message));
|
return Err(format!("Prefill error: {}", error.message()));
|
||||||
}
|
}
|
||||||
_ => continue,
|
_ => continue,
|
||||||
}
|
}
|
||||||
@@ -822,7 +839,7 @@ impl StreamingProcessor {
|
|||||||
/// Process generate streaming with optional input_logprobs
|
/// Process generate streaming with optional input_logprobs
|
||||||
async fn process_generate_streaming_with_input_logprobs(
|
async fn process_generate_streaming_with_input_logprobs(
|
||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut stream: AbortOnDropStream,
|
mut stream: ProtoStream,
|
||||||
request_id: String,
|
request_id: String,
|
||||||
weight_version: String,
|
weight_version: String,
|
||||||
_include_logprobs: bool,
|
_include_logprobs: bool,
|
||||||
@@ -840,23 +857,26 @@ impl StreamingProcessor {
|
|||||||
while let Some(response) = stream.next().await {
|
while let Some(response) = stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Stream error: {}", e))?;
|
||||||
|
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Chunk(chunk)) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
let index = chunk.index;
|
let index = chunk.index();
|
||||||
|
|
||||||
// Update completion tokens for this index
|
// Both backends send delta token_ids, so accumulate for both
|
||||||
let completion_tokens = completion_tokens_map.entry(index).or_insert(0);
|
let completion_tokens = completion_tokens_map.entry(index).or_insert(0);
|
||||||
*completion_tokens += chunk.token_ids.len() as u32;
|
*completion_tokens += chunk.token_ids().len() as u32;
|
||||||
|
let current_completion_tokens = *completion_tokens;
|
||||||
|
|
||||||
// Decode tokens to text
|
// Decode tokens to text
|
||||||
let chunk_text = tokenizer.decode(&chunk.token_ids, true).unwrap_or_default();
|
let chunk_text = tokenizer
|
||||||
|
.decode(chunk.token_ids(), true)
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
// Accumulate text for this index
|
// Accumulate text for this index
|
||||||
let accumulated_text = accumulated_texts.entry(index).or_default();
|
let accumulated_text = accumulated_texts.entry(index).or_default();
|
||||||
accumulated_text.push_str(&chunk_text);
|
accumulated_text.push_str(&chunk_text);
|
||||||
|
|
||||||
// Store latest output logprobs (cumulative from proto, convert to SGLang format)
|
// Store latest output logprobs (cumulative from proto, convert to SGLang format)
|
||||||
if let Some(ref output_logprobs) = chunk.output_logprobs {
|
if let Some(output_logprobs) = chunk.output_logprobs() {
|
||||||
let converted = utils::convert_generate_output_logprobs(output_logprobs);
|
let converted = utils::convert_generate_output_logprobs(output_logprobs);
|
||||||
accumulated_output_logprobs.insert(index, Some(converted));
|
accumulated_output_logprobs.insert(index, Some(converted));
|
||||||
}
|
}
|
||||||
@@ -869,18 +889,18 @@ impl StreamingProcessor {
|
|||||||
.get(&index)
|
.get(&index)
|
||||||
.and_then(|o| o.as_ref());
|
.and_then(|o| o.as_ref());
|
||||||
|
|
||||||
let chunk_response = serde_json::json!({
|
let chunk_response = json!({
|
||||||
"text": accumulated_text.clone(),
|
"text": accumulated_text.clone(),
|
||||||
"output_ids": chunk.token_ids,
|
"output_ids": chunk.token_ids(),
|
||||||
"meta_info": {
|
"meta_info": {
|
||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": null,
|
"finish_reason": null,
|
||||||
"prompt_tokens": chunk.prompt_tokens,
|
"prompt_tokens": chunk.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &weight_version,
|
||||||
"input_token_logprobs": input_token_logprobs.as_ref(),
|
"input_token_logprobs": input_token_logprobs.as_ref(),
|
||||||
"output_token_logprobs": current_output_logprobs,
|
"output_token_logprobs": current_output_logprobs,
|
||||||
"completion_tokens": *completion_tokens,
|
"completion_tokens": current_completion_tokens,
|
||||||
"cached_tokens": chunk.cached_tokens
|
"cached_tokens": chunk.cached_tokens()
|
||||||
},
|
},
|
||||||
"index": index
|
"index": index
|
||||||
});
|
});
|
||||||
@@ -892,11 +912,14 @@ impl StreamingProcessor {
|
|||||||
tx.send(Ok(Bytes::from(sse_chunk)))
|
tx.send(Ok(Bytes::from(sse_chunk)))
|
||||||
.map_err(|_| "Failed to send chunk".to_string())?;
|
.map_err(|_| "Failed to send chunk".to_string())?;
|
||||||
}
|
}
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete) => {
|
||||||
let index = complete.index;
|
let index = complete.index();
|
||||||
let accumulated_text =
|
let accumulated_text =
|
||||||
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
||||||
|
|
||||||
|
// Use accumulated count (we tracked deltas from both backends)
|
||||||
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
||||||
|
|
||||||
let final_output_logprobs = accumulated_output_logprobs
|
let final_output_logprobs = accumulated_output_logprobs
|
||||||
.get(&index)
|
.get(&index)
|
||||||
.and_then(|o| o.as_ref());
|
.and_then(|o| o.as_ref());
|
||||||
@@ -905,23 +928,23 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
// Parse finish_reason
|
// Parse finish_reason
|
||||||
let finish_reason = utils::parse_finish_reason(
|
let finish_reason = utils::parse_finish_reason(
|
||||||
&complete.finish_reason,
|
complete.finish_reason(),
|
||||||
complete.completion_tokens,
|
complete.completion_tokens(),
|
||||||
);
|
);
|
||||||
|
|
||||||
// Send final chunk with finish_reason
|
// Send final chunk with finish_reason
|
||||||
let finish_response = json!({
|
let finish_response = json!({
|
||||||
"text": accumulated_text,
|
"text": accumulated_text,
|
||||||
"output_ids": complete.output_ids[complete.output_ids.len().saturating_sub(1)..].to_vec(),
|
"output_ids": complete.output_ids()[complete.output_ids().len().saturating_sub(1)..].to_vec(),
|
||||||
"meta_info": {
|
"meta_info": {
|
||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": finish_reason,
|
"finish_reason": finish_reason,
|
||||||
"prompt_tokens": complete.prompt_tokens,
|
"prompt_tokens": complete.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &weight_version,
|
||||||
"input_token_logprobs": input_token_logprobs.as_ref(),
|
"input_token_logprobs": input_token_logprobs.as_ref(),
|
||||||
"output_token_logprobs": final_output_logprobs,
|
"output_token_logprobs": final_output_logprobs,
|
||||||
"completion_tokens": completion_tokens,
|
"completion_tokens": completion_tokens,
|
||||||
"cached_tokens": complete.cached_tokens,
|
"cached_tokens": complete.cached_tokens(),
|
||||||
"e2e_latency": e2e_latency
|
"e2e_latency": e2e_latency
|
||||||
},
|
},
|
||||||
"index": index
|
"index": index
|
||||||
@@ -936,10 +959,10 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
// Continue to process all completions if n>1
|
// Continue to process all completions if n>1
|
||||||
}
|
}
|
||||||
Some(Error(error)) => {
|
ProtoResponseVariant::Error(error) => {
|
||||||
return Err(error.message);
|
return Err(error.message().to_string());
|
||||||
}
|
}
|
||||||
None => continue,
|
ProtoResponseVariant::None => continue,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,15 +3,19 @@
|
|||||||
use std::{collections::HashMap, sync::Arc};
|
use std::{collections::HashMap, sync::Arc};
|
||||||
|
|
||||||
use axum::response::Response;
|
use axum::response::Response;
|
||||||
use futures::StreamExt;
|
|
||||||
use serde_json::{json, Map, Value};
|
use serde_json::{json, Map, Value};
|
||||||
use tracing::{error, warn};
|
use tracing::{error, warn};
|
||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use super::{error, ProcessedMessages};
|
use super::{
|
||||||
|
client::GrpcClient,
|
||||||
|
error,
|
||||||
|
proto_wrapper::{ProtoGenerateComplete, ProtoStream},
|
||||||
|
ProcessedMessages,
|
||||||
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
core::Worker,
|
core::Worker,
|
||||||
grpc_client::{proto, sglang_scheduler::AbortOnDropStream, SglangSchedulerClient},
|
grpc_client::sglang_proto::{InputLogProbs, OutputLogProbs},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatMessage},
|
chat::{ChatCompletionRequest, ChatMessage},
|
||||||
common::{
|
common::{
|
||||||
@@ -24,6 +28,7 @@ use crate::{
|
|||||||
ParserFactory as ReasoningParserFactory, PooledParser as ReasoningPooledParser,
|
ParserFactory as ReasoningParserFactory, PooledParser as ReasoningPooledParser,
|
||||||
ReasoningParser,
|
ReasoningParser,
|
||||||
},
|
},
|
||||||
|
routers::grpc::proto_wrapper::ProtoResponseVariant,
|
||||||
tokenizer::{
|
tokenizer::{
|
||||||
cache::CachedTokenizer,
|
cache::CachedTokenizer,
|
||||||
chat_template::{ChatTemplateContentFormat, ChatTemplateParams},
|
chat_template::{ChatTemplateContentFormat, ChatTemplateParams},
|
||||||
@@ -37,18 +42,24 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
/// Get gRPC client from worker, returning appropriate error response on failure
|
/// Get gRPC client from worker, returning appropriate error response on failure
|
||||||
pub async fn get_grpc_client_from_worker(
|
pub async fn get_grpc_client_from_worker(worker: &Arc<dyn Worker>) -> Result<GrpcClient, Response> {
|
||||||
worker: &Arc<dyn Worker>,
|
// Get cached client from worker (or create one if not cached yet)
|
||||||
) -> Result<SglangSchedulerClient, Response> {
|
|
||||||
let client_arc = worker
|
let client_arc = worker
|
||||||
.get_grpc_client()
|
.get_grpc_client()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
error!(function = "get_grpc_client_from_worker", error = %e, "Failed to get gRPC client");
|
error!(
|
||||||
|
function = "get_grpc_client_from_worker",
|
||||||
|
error = %e,
|
||||||
|
"Failed to get gRPC client from worker"
|
||||||
|
);
|
||||||
error::internal_error(format!("Failed to get gRPC client: {}", e))
|
error::internal_error(format!("Failed to get gRPC client: {}", e))
|
||||||
})?
|
})?
|
||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
error!(function = "get_grpc_client_from_worker", "Selected worker not configured for gRPC");
|
error!(
|
||||||
|
function = "get_grpc_client_from_worker",
|
||||||
|
"Selected worker not configured for gRPC"
|
||||||
|
);
|
||||||
error::internal_error("Selected worker is not configured for gRPC")
|
error::internal_error("Selected worker is not configured for gRPC")
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
@@ -587,32 +598,31 @@ pub fn parse_json_schema_response(
|
|||||||
/// * `Ok(Vec<GenerateComplete>)` - All complete responses collected from the stream
|
/// * `Ok(Vec<GenerateComplete>)` - All complete responses collected from the stream
|
||||||
/// * `Err(Response)` - Error response if the stream fails or returns an error
|
/// * `Err(Response)` - Error response if the stream fails or returns an error
|
||||||
pub async fn collect_stream_responses(
|
pub async fn collect_stream_responses(
|
||||||
stream: &mut AbortOnDropStream,
|
stream: &mut ProtoStream,
|
||||||
worker_name: &str,
|
worker_name: &str,
|
||||||
) -> Result<Vec<proto::GenerateComplete>, Response> {
|
) -> Result<Vec<ProtoGenerateComplete>, Response> {
|
||||||
use proto::generate_response::Response::*;
|
|
||||||
|
|
||||||
let mut all_responses = Vec::new();
|
let mut all_responses = Vec::new();
|
||||||
|
|
||||||
while let Some(response) = stream.next().await {
|
while let Some(response) = stream.next().await {
|
||||||
match response {
|
match response {
|
||||||
Ok(gen_response) => {
|
Ok(gen_response) => {
|
||||||
match gen_response.response {
|
match gen_response.into_response() {
|
||||||
Some(Complete(complete)) => {
|
ProtoResponseVariant::Complete(complete) => {
|
||||||
all_responses.push(complete);
|
all_responses.push(complete);
|
||||||
}
|
}
|
||||||
Some(Error(err)) => {
|
ProtoResponseVariant::Error(err) => {
|
||||||
error!(function = "collect_stream_responses", worker = %worker_name, error = %err.message, "Worker generation error");
|
error!(function = "collect_stream_responses", worker = %worker_name, error = %err.message(), "Worker generation error");
|
||||||
// Don't mark as completed - let Drop send abort for error cases
|
// Don't mark as completed - let Drop send abort for error cases
|
||||||
return Err(error::internal_error(format!(
|
return Err(error::internal_error(format!(
|
||||||
"{} generation failed: {}",
|
"{} generation failed: {}",
|
||||||
worker_name, err.message
|
worker_name,
|
||||||
|
err.message()
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
Some(Chunk(_chunk)) => {
|
ProtoResponseVariant::Chunk(_chunk) => {
|
||||||
// Streaming chunk - no action needed
|
// Streaming chunk - no action needed
|
||||||
}
|
}
|
||||||
None => {
|
ProtoResponseVariant::None => {
|
||||||
// Empty response - no action needed
|
// Empty response - no action needed
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -804,12 +814,12 @@ pub fn create_tool_parser(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Convert proto::OutputLogProbs to OpenAI ChatLogProbs format
|
/// Convert OutputLogProbs to OpenAI ChatLogProbs format
|
||||||
///
|
///
|
||||||
/// This function decodes token IDs using the tokenizer and builds the logprobs structure
|
/// This function decodes token IDs using the tokenizer and builds the logprobs structure
|
||||||
/// expected by the OpenAI API format.
|
/// expected by the OpenAI API format.
|
||||||
pub fn convert_proto_to_openai_logprobs(
|
pub fn convert_proto_to_openai_logprobs(
|
||||||
proto_logprobs: &proto::OutputLogProbs,
|
proto_logprobs: &OutputLogProbs,
|
||||||
tokenizer: &Arc<dyn Tokenizer>,
|
tokenizer: &Arc<dyn Tokenizer>,
|
||||||
) -> Result<ChatLogProbs, String> {
|
) -> Result<ChatLogProbs, String> {
|
||||||
let mut content_items = Vec::new();
|
let mut content_items = Vec::new();
|
||||||
@@ -877,13 +887,11 @@ pub fn convert_proto_to_openai_logprobs(
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Convert proto::OutputLogProbs to Generate format Vec<Vec<Option<f64>>>
|
/// Convert OutputLogProbs to Generate format Vec<Vec<Option<f64>>>
|
||||||
///
|
///
|
||||||
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
||||||
/// Each inner vec contains [logprob (f64), token_id (i32), ...]
|
/// Each inner vec contains [logprob (f64), token_id (i32), ...]
|
||||||
pub fn convert_generate_output_logprobs(
|
pub fn convert_generate_output_logprobs(proto_logprobs: &OutputLogProbs) -> Vec<Vec<Option<f64>>> {
|
||||||
proto_logprobs: &proto::OutputLogProbs,
|
|
||||||
) -> Vec<Vec<Option<f64>>> {
|
|
||||||
proto_logprobs
|
proto_logprobs
|
||||||
.token_logprobs
|
.token_logprobs
|
||||||
.iter()
|
.iter()
|
||||||
@@ -892,13 +900,11 @@ pub fn convert_generate_output_logprobs(
|
|||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Convert proto::InputLogProbs to Generate format Vec<Vec<Option<f64>>>
|
/// Convert InputLogProbs to Generate format Vec<Vec<Option<f64>>>
|
||||||
///
|
///
|
||||||
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
/// Generate format: [[logprob, token_id, ...], [logprob, token_id, ...], ...]
|
||||||
/// First token has null logprob: [[null, token_id], [logprob, token_id], ...]
|
/// First token has null logprob: [[null, token_id], [logprob, token_id], ...]
|
||||||
pub fn convert_generate_input_logprobs(
|
pub fn convert_generate_input_logprobs(proto_logprobs: &InputLogProbs) -> Vec<Vec<Option<f64>>> {
|
||||||
proto_logprobs: &proto::InputLogProbs,
|
|
||||||
) -> Vec<Vec<Option<f64>>> {
|
|
||||||
proto_logprobs
|
proto_logprobs
|
||||||
.token_logprobs
|
.token_logprobs
|
||||||
.iter()
|
.iter()
|
||||||
|
|||||||
@@ -530,6 +530,7 @@ async fn get_worker(State(state): State<Arc<AppState>>, Path(url): Path<String>)
|
|||||||
is_healthy: false,
|
is_healthy: false,
|
||||||
load: 0,
|
load: 0,
|
||||||
connection_mode: "unknown".to_string(),
|
connection_mode: "unknown".to_string(),
|
||||||
|
runtime_type: None,
|
||||||
tokenizer_path: None,
|
tokenizer_path: None,
|
||||||
reasoning_parser: None,
|
reasoning_parser: None,
|
||||||
tool_parser: None,
|
tool_parser: None,
|
||||||
|
|||||||
@@ -376,6 +376,7 @@ async fn handle_pod_event(
|
|||||||
worker_type,
|
worker_type,
|
||||||
priority: None,
|
priority: None,
|
||||||
cost: None,
|
cost: None,
|
||||||
|
runtime: None,
|
||||||
labels: HashMap::new(),
|
labels: HashMap::new(),
|
||||||
bootstrap_port,
|
bootstrap_port,
|
||||||
tokenizer_path: None,
|
tokenizer_path: None,
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ async fn test_policy_registry_with_router_manager() {
|
|||||||
reasoning_parser: None,
|
reasoning_parser: None,
|
||||||
tool_parser: None,
|
tool_parser: None,
|
||||||
chat_template: None,
|
chat_template: None,
|
||||||
|
runtime: None,
|
||||||
health_check_timeout_secs: 30,
|
health_check_timeout_secs: 30,
|
||||||
health_check_interval_secs: 60,
|
health_check_interval_secs: 60,
|
||||||
health_success_threshold: 2,
|
health_success_threshold: 2,
|
||||||
@@ -67,6 +68,7 @@ async fn test_policy_registry_with_router_manager() {
|
|||||||
reasoning_parser: None,
|
reasoning_parser: None,
|
||||||
tool_parser: None,
|
tool_parser: None,
|
||||||
chat_template: None,
|
chat_template: None,
|
||||||
|
runtime: None,
|
||||||
health_check_timeout_secs: 30,
|
health_check_timeout_secs: 30,
|
||||||
health_check_interval_secs: 60,
|
health_check_interval_secs: 60,
|
||||||
health_success_threshold: 2,
|
health_success_threshold: 2,
|
||||||
@@ -93,6 +95,7 @@ async fn test_policy_registry_with_router_manager() {
|
|||||||
tokenizer_path: None,
|
tokenizer_path: None,
|
||||||
reasoning_parser: None,
|
reasoning_parser: None,
|
||||||
tool_parser: None,
|
tool_parser: None,
|
||||||
|
runtime: None,
|
||||||
chat_template: None,
|
chat_template: None,
|
||||||
health_check_timeout_secs: 30,
|
health_check_timeout_secs: 30,
|
||||||
health_check_interval_secs: 60,
|
health_check_interval_secs: 60,
|
||||||
|
|||||||
Reference in New Issue
Block a user