diff --git a/crates/grpc_client/build.rs b/crates/grpc_client/build.rs index 9809b80b10..8a8730c1af 100644 --- a/crates/grpc_client/build.rs +++ b/crates/grpc_client/build.rs @@ -2,6 +2,7 @@ fn main() -> Result<(), Box> { // Rebuild triggers println!("cargo:rerun-if-changed=proto/common.proto"); println!("cargo:rerun-if-changed=proto/sglang_scheduler.proto"); + println!("cargo:rerun-if-changed=proto/tokenspeed_scheduler.proto"); println!("cargo:rerun-if-changed=proto/vllm_engine.proto"); println!("cargo:rerun-if-changed=proto/trtllm_service.proto"); println!("cargo:rerun-if-changed=proto/mlx_engine.proto"); @@ -19,8 +20,8 @@ fn main() -> Result<(), Box> { .build_client(true) .extern_path(".smg.grpc.common", "crate::common_proto") .type_attribute("GetModelInfoResponse", "#[derive(serde::Serialize)]") - // vllm + trtllm ServerInfo have only primitive fields. - // sglang's contains prost_types::{Struct,Timestamp} so it's handled separately. + // Some ServerInfo protos contain prost_types::{Struct, Timestamp}; + // those are handled separately at the wrapper layer. .type_attribute( "vllm.grpc.engine.GetServerInfoResponse", "#[derive(serde::Serialize)]", @@ -40,6 +41,7 @@ fn main() -> Result<(), Box> { "proto/vllm_engine.proto", "proto/trtllm_service.proto", "proto/mlx_engine.proto", + "proto/tokenspeed_scheduler.proto", ], &["proto"], )?; diff --git a/crates/grpc_client/proto/tokenspeed_scheduler.proto b/crates/grpc_client/proto/tokenspeed_scheduler.proto new file mode 100644 index 0000000000..5be4d5b6b6 --- /dev/null +++ b/crates/grpc_client/proto/tokenspeed_scheduler.proto @@ -0,0 +1,268 @@ +syntax = "proto3"; + +package tokenspeed.grpc.scheduler; + +import "google/protobuf/timestamp.proto"; +import "google/protobuf/struct.proto"; + +// TokenSpeed scheduler gRPC service. Fully self-contained wire definition. +// Trimmed to text-generation only (no embed, no multimodal, no +// PD-disaggregated, no LoRA, no hidden-state forwarding). +service TokenSpeedScheduler { + rpc Generate(GenerateRequest) returns (stream GenerateResponse); + rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse); + rpc Abort(AbortRequest) returns (AbortResponse); + rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse); + rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse); + rpc GetLoads(GetLoadsRequest) returns (GetLoadsResponse); +} + +// ===================== +// Sampling +// ===================== + +// Sampling scalars are `optional` so the servicer can distinguish +// "client set 0" from "client unset" via `HasField()`. `min_new_tokens` +// stays non-optional because 0 is its semantic "no minimum" sentinel. +message SamplingParams { + optional float temperature = 1; + optional float top_p = 2; + optional int32 top_k = 3; + optional float min_p = 4; + optional float frequency_penalty = 5; + optional float presence_penalty = 6; + optional float repetition_penalty = 7; + + optional uint32 max_new_tokens = 8; + uint32 min_new_tokens = 9; + + repeated string stop = 10; + repeated uint32 stop_token_ids = 11; + bool ignore_eos = 12; + + bool skip_special_tokens = 13; + bool spaces_between_special_tokens = 14; + + // Number of samples (n in OpenAI API). + uint32 n = 15; + + // Per-token logit bias. + map logit_bias = 16; + + // Structured generation. Currently xfailed in e2e (tokenspeed#361). + oneof constraint { + string regex = 17; + string json_schema = 18; + string ebnf_grammar = 19; + string structural_tag = 20; + } + + // Keep the trailing matched stop token in `output_ids`. + bool no_stop_trim = 22; + + // Escape hatch for backend-specific knobs without bumping the proto. + google.protobuf.Struct custom_params = 21; +} + +// ===================== +// Generate +// ===================== + +message GenerateRequest { + string request_id = 1; + + // Tokenized input (router does its own tokenization). + TokenizedInput tokenized = 2; + + SamplingParams sampling_params = 3; + + // Logprob options. + bool return_logprob = 4; + // Optional: unset = no input logprobs; explicit 0 = "from start". + optional int32 logprob_start_len = 5; + int32 top_logprobs_num = 6; + repeated uint32 token_ids_logprob = 7; + + // Whether the client wants stream chunks (otherwise: complete-only). + bool stream = 8; +} + +message TokenizedInput { + repeated uint32 input_ids = 1; + string original_text = 2; // cosmetic, for worker logs +} + +message GenerateResponse { + string request_id = 1; + + oneof response { + GenerateStreamChunk chunk = 2; + GenerateComplete complete = 3; + } +} + +message GenerateStreamChunk { + // Generated tokens since the previous chunk. + repeated uint32 token_ids = 1; + + uint32 prompt_tokens = 2; + uint32 completion_tokens = 3; + uint32 cached_tokens = 4; + + OutputLogProbs output_logprobs = 5; + + // For ordering when n>1. + uint32 index = 6; +} + +message GenerateComplete { + repeated uint32 output_ids = 1; + + // OpenAI-compatible: "stop", "length", "abort", "tool_calls". + string finish_reason = 2; + + uint32 prompt_tokens = 3; + uint32 completion_tokens = 4; + uint32 cached_tokens = 5; + + OutputLogProbs output_logprobs = 6; + + // Which stop matched (for clients that care which `stop` triggered). + oneof matched_stop { + uint32 matched_token_id = 7; + string matched_stop_str = 8; + } + + uint32 index = 9; +} + +message OutputLogProbs { + repeated float token_logprobs = 1; + repeated uint32 token_ids = 2; + repeated TopLogProbs top_logprobs = 3; +} + +message TopLogProbs { + repeated float values = 1; + repeated uint32 token_ids = 2; +} + +// ===================== +// Management +// ===================== + +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 & Server Info +// ===================== + +message GetModelInfoRequest {} +message GetModelInfoResponse { + string model_path = 1; + string tokenizer_path = 2; + string served_model_name = 3; + string model_type = 4; + repeated string architectures = 5; + + int32 max_context_length = 6; + int32 max_req_input_len = 7; + int32 vocab_size = 8; + + repeated int32 eos_token_ids = 9; + int32 pad_token_id = 10; + int32 bos_token_id = 11; + + string weight_version = 12; + // Backend-reported sampling defaults (JSON string, e.g. + // ``{"temperature": 0.6, "top_p": 0.9}``). Empty = no overrides. Surfaced + // to the router via the ``default_sampling_params_json`` worker label. + string default_sampling_params_json = 13; +} + +message GetServerInfoRequest {} +message GetServerInfoResponse { + google.protobuf.Struct server_args = 1; + google.protobuf.Struct scheduler_info = 2; + + int32 active_requests = 3; + bool is_paused = 4; + double uptime_seconds = 5; + int32 max_total_num_tokens = 6; + + string tokenspeed_version = 7; + google.protobuf.Timestamp start_time = 8; +} + +// ===================== +// Loads +// ===================== + +message GetLoadsRequest { + optional int32 dp_rank = 1; + // Sections: "core" (default), "memory", "queues". Pass "all" for everything. + repeated string include = 2; +} + +message GetLoadsResponse { + string timestamp = 1; + string version = 2; + int32 dp_rank_count = 3; + repeated SchedulerLoad loads = 4; + AggregateMetrics aggregate = 5; +} + +message SchedulerLoad { + int32 dp_rank = 1; + + int32 num_running_reqs = 2; + int32 num_waiting_reqs = 3; + int32 num_total_reqs = 4; + int32 num_used_tokens = 5; + int32 max_total_num_tokens = 6; + int32 max_running_requests = 7; + + double token_usage = 8; + double gen_throughput = 9; + double cache_hit_rate = 10; + double utilization = 11; + + optional MemoryMetrics memory = 12; + optional QueueMetrics queues = 13; +} + +message MemoryMetrics { + double weight_gb = 1; + double kv_cache_gb = 2; + double graph_gb = 3; + int32 token_capacity = 4; +} + +message QueueMetrics { + int32 waiting = 1; + int32 grammar = 2; + int32 paused = 3; + int32 retracted = 4; +} + +message AggregateMetrics { + int32 total_running_reqs = 1; + int32 total_waiting_reqs = 2; + int32 total_reqs = 3; + double avg_token_usage = 4; + double avg_throughput = 5; + double avg_utilization = 6; +} diff --git a/crates/grpc_client/python/smg_grpc_proto/__init__.py b/crates/grpc_client/python/smg_grpc_proto/__init__.py index 6a19b4aead..f7eac4a3e2 100644 --- a/crates/grpc_client/python/smg_grpc_proto/__init__.py +++ b/crates/grpc_client/python/smg_grpc_proto/__init__.py @@ -1,4 +1,4 @@ -"""SMG gRPC Proto - Protocol definitions for SGLang, vLLM, TRT-LLM, and MLX.""" +"""SMG gRPC Proto - Protocol definitions for SGLang, TokenSpeed, vLLM, TRT-LLM, and MLX.""" from importlib.metadata import version @@ -14,6 +14,8 @@ sglang_encoder_pb2_grpc, sglang_scheduler_pb2, sglang_scheduler_pb2_grpc, + tokenspeed_scheduler_pb2, + tokenspeed_scheduler_pb2_grpc, trtllm_service_pb2, trtllm_service_pb2_grpc, vllm_engine_pb2, @@ -25,6 +27,8 @@ "sglang_scheduler_pb2_grpc", "sglang_encoder_pb2", "sglang_encoder_pb2_grpc", + "tokenspeed_scheduler_pb2", + "tokenspeed_scheduler_pb2_grpc", "vllm_engine_pb2", "vllm_engine_pb2_grpc", "trtllm_service_pb2", diff --git a/crates/grpc_client/src/lib.rs b/crates/grpc_client/src/lib.rs index 5b349faa47..a9f5474b32 100644 --- a/crates/grpc_client/src/lib.rs +++ b/crates/grpc_client/src/lib.rs @@ -1,7 +1,8 @@ -//! gRPC clients for SGLang, vLLM, TensorRT-LLM, and MLX backends +//! gRPC clients for SGLang, vLLM, TensorRT-LLM, MLX, and TokenSpeed backends. //! //! This crate provides gRPC client implementations for communicating with -//! SGLang scheduler, vLLM engine, TensorRT-LLM engine, and MLX engine backends. +//! the SGLang scheduler, vLLM engine, TensorRT-LLM engine, MLX engine, and +//! TokenSpeed scheduler backends. pub mod common_proto { #![allow(clippy::all, clippy::absolute_paths, unused_qualifications)] @@ -12,6 +13,7 @@ pub mod channel; pub mod mlx_engine; pub mod sglang_scheduler; pub mod tokenizer_bundle; +pub mod tokenspeed_scheduler; pub mod trtllm_service; pub mod vllm_engine; @@ -22,6 +24,7 @@ pub use abort_on_drop::{AbortOnDropClient, AbortOnDropStream}; pub use channel::{connect_channel, normalize_grpc_endpoint}; pub use mlx_engine::{proto as mlx_proto, MlxEngineClient}; pub use sglang_scheduler::{proto as sglang_proto, SglangSchedulerClient}; +pub use tokenspeed_scheduler::{tokenspeed_proto, TokenSpeedSchedulerClient}; use tonic::metadata::MetadataMap; pub use trtllm_service::{proto as trtllm_proto, TrtllmServiceClient}; pub use vllm_engine::{proto as vllm_proto, VllmEngineClient}; diff --git a/crates/grpc_client/src/tokenspeed_scheduler.rs b/crates/grpc_client/src/tokenspeed_scheduler.rs new file mode 100644 index 0000000000..ea194f48cb --- /dev/null +++ b/crates/grpc_client/src/tokenspeed_scheduler.rs @@ -0,0 +1,684 @@ +//! gRPC client for the TokenSpeed scheduler service. +//! +//! Wire types are TokenSpeed-native end-to-end (`tokenspeed_proto::*`). +//! Sampling-params builders are private `Self::build_*` methods that emit +//! `tokenspeed_proto::SamplingParams` directly. Unary RPC responses +//! (`get_model_info`, `get_server_info`, `get_loads`) also surface as +//! native types — the router consumes them through dedicated +//! `ModelInfo::TokenSpeed` / `ServerInfo::TokenSpeed` enum arms. + +use std::{future::Future, pin::Pin, sync::Arc}; + +use openai_protocol::{ + chat::ChatCompletionRequest, + common::{ResponseFormat, StringOrArray}, + completion::CompletionRequest, + generate::GenerateRequest, + messages::CreateMessageRequest, + responses::ResponsesRequest, + sampling_params::SamplingParams as GenerateSamplingParams, +}; +use tonic::{transport::Channel, Request}; +use tracing::{debug, warn}; + +use crate::{AbortOnDropClient, BoxedTraceInjector, NoopTraceInjector}; + +#[expect(clippy::allow_attributes)] +pub mod tokenspeed_proto { + #![allow(clippy::all, clippy::absolute_paths, unused_qualifications)] + tonic::include_proto!("tokenspeed.grpc.scheduler"); +} + +/// Streaming `generate()` response that auto-aborts on drop. Concrete +/// alias for the generic `crate::AbortOnDropStream`. +pub type AbortOnDropStream = + crate::AbortOnDropStream; + +/// gRPC client for the TokenSpeed scheduler. +#[derive(Clone)] +pub struct TokenSpeedSchedulerClient { + client: tokenspeed_proto::token_speed_scheduler_client::TokenSpeedSchedulerClient, + trace_injector: BoxedTraceInjector, +} + +impl AbortOnDropClient for TokenSpeedSchedulerClient { + fn abort_for_drop( + self, + request_id: String, + ) -> Pin> + Send>> { + Box::pin(async move { + self.abort_request(request_id, "Stream dropped".to_string()) + .await + }) + } +} + +impl TokenSpeedSchedulerClient { + pub async fn connect(endpoint: &str) -> Result> { + Self::connect_with_trace_injector(endpoint, Arc::new(NoopTraceInjector)).await + } + + pub async fn connect_with_trace_injector( + endpoint: &str, + trace_injector: BoxedTraceInjector, + ) -> Result> { + debug!("Connecting to TokenSpeed scheduler at {}", endpoint); + let channel = crate::channel::connect_channel(endpoint).await?; + let client = + tokenspeed_proto::token_speed_scheduler_client::TokenSpeedSchedulerClient::new(channel); + + Ok(Self { + client, + trace_injector, + }) + } + + #[must_use] + pub fn with_trace_injector(mut self, trace_injector: BoxedTraceInjector) -> Self { + self.trace_injector = trace_injector; + self + } + + /// Submit a generation request. + pub async fn generate( + &self, + req: tokenspeed_proto::GenerateRequest, + ) -> Result { + let request_id = req.request_id.clone(); + + let mut client = self.client.clone(); + let mut request = Request::new(req); + + if let Err(e) = self.trace_injector.inject(request.metadata_mut()) { + warn!("Failed to inject trace context: {}", e); + } + + let response = client.generate(request).await?; + + Ok(AbortOnDropStream::new( + response.into_inner(), + request_id, + self.clone(), + )) + } + + pub async fn health_check( + &self, + ) -> Result { + debug!("Sending TokenSpeed health check request"); + let request = Request::new(tokenspeed_proto::HealthCheckRequest {}); + let mut client = self.client.clone(); + let response = client.health_check(request).await?; + Ok(response.into_inner()) + } + + pub async fn abort_request( + &self, + request_id: String, + reason: String, + ) -> Result<(), tonic::Status> { + debug!( + "Sending TokenSpeed abort for {} (reason: {})", + request_id, reason + ); + let request = Request::new(tokenspeed_proto::AbortRequest { + request_id: request_id.clone(), + reason, + }); + let mut client = self.client.clone(); + let response = client.abort(request).await?; + debug!( + "TokenSpeed abort response for {}: success={}, message={}", + request_id, + response.get_ref().success, + response.get_ref().message + ); + Ok(()) + } + + pub async fn get_model_info( + &self, + ) -> Result { + let request = Request::new(tokenspeed_proto::GetModelInfoRequest {}); + let mut client = self.client.clone(); + let response = client.get_model_info(request).await?; + Ok(response.into_inner()) + } + + pub async fn get_server_info( + &self, + ) -> Result { + let request = Request::new(tokenspeed_proto::GetServerInfoRequest {}); + let mut client = self.client.clone(); + let response = client.get_server_info(request).await?; + Ok(response.into_inner()) + } + + pub async fn get_loads( + &self, + include: Vec, + ) -> Result { + let request = Request::new(tokenspeed_proto::GetLoadsRequest { + dp_rank: None, + include, + }); + let mut client = self.client.clone(); + let response = client.get_loads(request).await?; + Ok(response.into_inner()) + } + + // ── Request builders ────────────────────────────────────────────── + + #[expect( + clippy::unused_self, + reason = "receiver kept for API parity with the other engine clients" + )] + pub fn build_generate_request_from_chat( + &self, + request_id: String, + body: &ChatCompletionRequest, + processed_text: String, + token_ids: Vec, + tool_call_constraint: Option<(String, String)>, + ) -> Result { + let sampling_params = Self::build_sampling_params_from_chat(body, tool_call_constraint)?; + Ok(tokenspeed_proto::GenerateRequest { + request_id, + tokenized: Some(tokenspeed_proto::TokenizedInput { + original_text: processed_text, + input_ids: token_ids, + }), + sampling_params: Some(sampling_params), + return_logprob: body.logprobs, + logprob_start_len: Some(-1), + top_logprobs_num: body.top_logprobs.unwrap_or(0) as i32, + stream: body.stream, + ..Default::default() + }) + } + + #[expect( + clippy::unused_self, + reason = "receiver kept for API parity with the other engine clients" + )] + pub fn build_plain_generate_request( + &self, + request_id: String, + body: &GenerateRequest, + original_text: Option, + token_ids: Vec, + ) -> Result { + let sampling_params = + Self::build_sampling_params_from_plain(body.sampling_params.as_ref())?; + Ok(tokenspeed_proto::GenerateRequest { + request_id, + tokenized: Some(tokenspeed_proto::TokenizedInput { + original_text: original_text.unwrap_or_default(), + input_ids: token_ids, + }), + sampling_params: Some(sampling_params), + return_logprob: body.return_logprob.unwrap_or(false), + logprob_start_len: Some(body.logprob_start_len.unwrap_or(-1)), + top_logprobs_num: body.top_logprobs_num.unwrap_or(0), + token_ids_logprob: body.token_ids_logprob.clone().unwrap_or_default(), + stream: body.stream, + }) + } + + #[expect( + clippy::unused_self, + reason = "receiver kept for API parity with the other engine clients" + )] + pub fn build_generate_request_from_responses( + &self, + request_id: String, + body: &ResponsesRequest, + processed_text: String, + token_ids: Vec, + constraint: Option<(String, String)>, + ) -> Result { + let sampling_params = Self::build_sampling_params_from_responses(body, constraint)?; + Ok(tokenspeed_proto::GenerateRequest { + request_id, + tokenized: Some(tokenspeed_proto::TokenizedInput { + original_text: processed_text, + input_ids: token_ids, + }), + sampling_params: Some(sampling_params), + stream: body.stream.unwrap_or(false), + ..Default::default() + }) + } + + #[expect( + clippy::unused_self, + reason = "receiver kept for API parity with the other engine clients" + )] + pub fn build_generate_request_from_messages( + &self, + request_id: String, + body: &CreateMessageRequest, + processed_text: String, + token_ids: Vec, + tool_call_constraint: Option<(String, String)>, + ) -> Result { + let sampling_params = + Self::build_sampling_params_from_messages(body, tool_call_constraint)?; + Ok(tokenspeed_proto::GenerateRequest { + request_id, + tokenized: Some(tokenspeed_proto::TokenizedInput { + original_text: processed_text, + input_ids: token_ids, + }), + sampling_params: Some(sampling_params), + stream: body.stream.unwrap_or(false), + ..Default::default() + }) + } + + #[expect( + clippy::unused_self, + reason = "receiver kept for API parity with the other engine clients" + )] + pub fn build_generate_request_from_completion( + &self, + request_id: String, + body: &CompletionRequest, + original_text: String, + token_ids: Vec, + ) -> Result { + let sampling_params = Self::build_sampling_params_from_completion(body)?; + Ok(tokenspeed_proto::GenerateRequest { + request_id, + tokenized: Some(tokenspeed_proto::TokenizedInput { + original_text, + input_ids: token_ids, + }), + sampling_params: Some(sampling_params), + return_logprob: body.logprobs.is_some(), + logprob_start_len: Some(-1), + top_logprobs_num: body.logprobs.unwrap_or(0) as i32, + stream: body.stream, + ..Default::default() + }) + } + + // ── Private sampling-params builders ───────────────────────────── + // + // TokenSpeed declares every sampling scalar as `optional`. Scalars are + // wrapped in `Some(_)` so the wire presence bit is set; + // `helpers::apply_tokenspeed_sampling_defaults` later overwrites them + // with model-published defaults when a worker advertises them. + + fn build_sampling_params_from_chat( + request: &ChatCompletionRequest, + tool_call_constraint: Option<(String, String)>, + ) -> Result { + let stop_sequences = Self::extract_stop_strings(request.stop.as_ref()); + + // Keep skip_special_tokens true; flipping to false on tool calls regresses BFCL. + Ok(tokenspeed_proto::SamplingParams { + temperature: Some(request.temperature.unwrap_or(1.0)), + top_p: Some(request.top_p.unwrap_or(1.0)), + top_k: Some(request.top_k.unwrap_or(-1)), + min_p: Some(request.min_p.unwrap_or(0.0)), + frequency_penalty: Some(request.frequency_penalty.unwrap_or(0.0)), + presence_penalty: Some(request.presence_penalty.unwrap_or(0.0)), + repetition_penalty: Some(request.repetition_penalty.unwrap_or(1.0)), + max_new_tokens: request.max_completion_tokens, + stop: stop_sequences, + stop_token_ids: request.stop_token_ids.clone().unwrap_or_default(), + skip_special_tokens: true, + spaces_between_special_tokens: true, + ignore_eos: request.ignore_eos, + no_stop_trim: request.no_stop_trim, + n: request.n.unwrap_or(1), + constraint: Self::build_constraint_for_chat(request, tool_call_constraint)?, + ..Default::default() + }) + } + + /// Used by Harmony models only. Regular models use the Chat API path. + /// Constraints come from the Harmony preparation stage (`structural_tag`) + /// or tool handling. + fn build_sampling_params_from_responses( + request: &ResponsesRequest, + constraint: Option<(String, String)>, + ) -> Result { + Ok(tokenspeed_proto::SamplingParams { + temperature: Some(request.temperature.unwrap_or(1.0)), + top_p: Some(request.top_p.unwrap_or(1.0)), + top_k: Some(request.top_k), + min_p: Some(request.min_p), + frequency_penalty: Some(request.frequency_penalty.unwrap_or(0.0)), + presence_penalty: Some(request.presence_penalty.unwrap_or(0.0)), + repetition_penalty: Some(request.repetition_penalty), + max_new_tokens: request.max_output_tokens, + stop: vec![], // Does not pass through request.stop yet + 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, + no_stop_trim: false, + n: 1, // Responses API doesn't support n>1 + constraint: Self::build_constraint_from_pair(constraint)?, + ..Default::default() + }) + } + + fn build_sampling_params_from_messages( + request: &CreateMessageRequest, + tool_call_constraint: Option<(String, String)>, + ) -> Result { + let stop_sequences = request.stop_sequences.clone().unwrap_or_default(); + + Ok(tokenspeed_proto::SamplingParams { + temperature: Some(request.temperature.unwrap_or(1.0) as f32), + top_p: Some(request.top_p.unwrap_or(1.0) as f32), + top_k: Some(request.top_k.map(|v| v as i32).unwrap_or(-1)), + min_p: Some(0.0), + frequency_penalty: Some(0.0), + presence_penalty: Some(0.0), + repetition_penalty: Some(1.0), + max_new_tokens: Some(request.max_tokens), + stop: stop_sequences, + stop_token_ids: vec![], + skip_special_tokens: true, + spaces_between_special_tokens: true, + ignore_eos: false, + no_stop_trim: false, + n: 1, + constraint: Self::build_constraint_from_pair(tool_call_constraint)?, + ..Default::default() + }) + } + + fn build_sampling_params_from_completion( + request: &CompletionRequest, + ) -> Result { + let stop_sequences = match &request.stop { + Some(StringOrArray::String(s)) => vec![s.clone()], + Some(StringOrArray::Array(arr)) => arr.clone(), + None => vec![], + }; + + Ok(tokenspeed_proto::SamplingParams { + temperature: Some(request.temperature.unwrap_or(1.0)), + top_p: Some(request.top_p.unwrap_or(1.0)), + top_k: Some(request.top_k.unwrap_or(-1)), + min_p: Some(request.min_p.unwrap_or(0.0)), + frequency_penalty: Some(request.frequency_penalty.unwrap_or(0.0)), + presence_penalty: Some(request.presence_penalty.unwrap_or(0.0)), + repetition_penalty: Some(request.repetition_penalty.unwrap_or(1.0)), + max_new_tokens: request.max_tokens, + min_new_tokens: request.min_tokens.unwrap_or(0), + stop: stop_sequences, + stop_token_ids: request.stop_token_ids.clone().unwrap_or_default(), + skip_special_tokens: request.skip_special_tokens, + spaces_between_special_tokens: true, + ignore_eos: request.ignore_eos, + no_stop_trim: request.no_stop_trim, + n: request.n.unwrap_or(1), + constraint: Self::build_constraint_from_completion(request)?, + ..Default::default() + }) + } + + fn build_sampling_params_from_plain( + params: Option<&GenerateSamplingParams>, + ) -> Result { + let mut sampling = tokenspeed_proto::SamplingParams { + temperature: Some(1.0), + top_p: Some(1.0), + top_k: Some(-1), + repetition_penalty: Some(1.0), + n: 1, + skip_special_tokens: true, + spaces_between_special_tokens: true, + ..Default::default() + }; + + let Some(p) = params else { + return Ok(sampling); + }; + + if let Some(v) = p.temperature { + sampling.temperature = Some(v); + } + if let Some(v) = p.top_p { + sampling.top_p = Some(v); + } + if let Some(v) = p.top_k { + sampling.top_k = Some(v); + } + if let Some(v) = p.frequency_penalty { + sampling.frequency_penalty = Some(v); + } + if let Some(v) = p.presence_penalty { + sampling.presence_penalty = Some(v); + } + if let Some(v) = p.repetition_penalty { + sampling.repetition_penalty = Some(v); + } + if let Some(v) = p.min_p { + sampling.min_p = Some(v); + } + if let Some(v) = p.ignore_eos { + sampling.ignore_eos = v; + } + if let Some(v) = p.skip_special_tokens { + sampling.skip_special_tokens = v; + } + if let Some(v) = p.no_stop_trim { + sampling.no_stop_trim = v; + } + + if let Some(stop) = &p.stop { + match stop { + StringOrArray::String(s) => sampling.stop.push(s.clone()), + StringOrArray::Array(arr) => sampling.stop.extend(arr.clone()), + } + } + if let Some(stop_token_ids) = &p.stop_token_ids { + sampling.stop_token_ids.clone_from(stop_token_ids); + } + + sampling.max_new_tokens = p.max_new_tokens; + if let Some(v) = p.min_new_tokens { + sampling.min_new_tokens = v; + } + if let Some(v) = p.n { + sampling.n = v; + } + + sampling.constraint = Self::build_constraint_from_plain(p)?; + + Ok(sampling) + } + + // ── Constraint helpers ─────────────────────────────────────────── + + fn extract_stop_strings(stop: Option<&StringOrArray>) -> Vec { + match stop { + Some(StringOrArray::String(s)) => vec![s.clone()], + Some(StringOrArray::Array(arr)) => arr.clone(), + None => vec![], + } + } + + fn build_constraint_for_chat( + request: &ChatCompletionRequest, + tool_call_constraint: Option<(String, String)>, + ) -> Result, String> { + let mut constraints = Vec::new(); + + match &request.response_format { + Some(ResponseFormat::JsonObject) => { + 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(tokenspeed_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(tokenspeed_proto::sampling_params::Constraint::JsonSchema( + schema_str, + )); + } + Some(ResponseFormat::Text) | None => {} + } + + if let Some(ebnf) = &request.ebnf { + constraints.push(tokenspeed_proto::sampling_params::Constraint::EbnfGrammar( + ebnf.clone(), + )); + } + if let Some(regex) = &request.regex { + constraints.push(tokenspeed_proto::sampling_params::Constraint::Regex( + regex.clone(), + )); + } + + // response_format wins over tool_call_constraint when both are set. + if let Some((constraint_type, constraint_value)) = tool_call_constraint { + if constraints.is_empty() { + let tool_constraint = + Self::constraint_from_pair(constraint_type, constraint_value)?; + constraints.push(tool_constraint); + } else { + warn!( + "Constrained decoding is not compatible with tool calls, dropping tool constraint" + ); + } + } + + match constraints.len() { + 0 => Ok(None), + 1 => Ok(constraints.pop()), + _ => Err("Multiple constraints are not allowed.".to_string()), + } + } + + fn build_constraint_from_pair( + constraint: Option<(String, String)>, + ) -> Result, String> { + if let Some((constraint_type, constraint_value)) = constraint { + Ok(Some(Self::constraint_from_pair( + constraint_type, + constraint_value, + )?)) + } else { + Ok(None) + } + } + + fn constraint_from_pair( + constraint_type: String, + constraint_value: String, + ) -> Result { + match constraint_type.as_str() { + "structural_tag" => { + Ok(tokenspeed_proto::sampling_params::Constraint::StructuralTag(constraint_value)) + } + "json_schema" => Ok(tokenspeed_proto::sampling_params::Constraint::JsonSchema( + constraint_value, + )), + "ebnf" => Ok(tokenspeed_proto::sampling_params::Constraint::EbnfGrammar( + constraint_value, + )), + "regex" => Ok(tokenspeed_proto::sampling_params::Constraint::Regex( + constraint_value, + )), + _ => Err(format!("Unknown constraint type: {constraint_type}")), + } + } + + fn build_constraint_from_completion( + request: &CompletionRequest, + ) -> Result, String> { + let mut constraints = Vec::new(); + if let Some(json_schema) = &request.json_schema { + constraints.push(tokenspeed_proto::sampling_params::Constraint::JsonSchema( + json_schema.clone(), + )); + } + if let Some(regex) = &request.regex { + constraints.push(tokenspeed_proto::sampling_params::Constraint::Regex( + regex.clone(), + )); + } + if let Some(ebnf) = &request.ebnf { + constraints.push(tokenspeed_proto::sampling_params::Constraint::EbnfGrammar( + ebnf.clone(), + )); + } + + match constraints.len() { + 0 => Ok(None), + 1 => Ok(constraints.pop()), + _ => Err("Multiple structured constraints are not allowed".to_string()), + } + } + + fn build_constraint_from_plain( + params: &GenerateSamplingParams, + ) -> Result, String> { + let mut constraints = Vec::new(); + if let Some(json_schema) = ¶ms.json_schema { + constraints.push(tokenspeed_proto::sampling_params::Constraint::JsonSchema( + json_schema.clone(), + )); + } + if let Some(regex) = ¶ms.regex { + constraints.push(tokenspeed_proto::sampling_params::Constraint::Regex( + regex.clone(), + )); + } + if let Some(ebnf) = ¶ms.ebnf { + constraints.push(tokenspeed_proto::sampling_params::Constraint::EbnfGrammar( + ebnf.clone(), + )); + } + + match constraints.len() { + 0 => Ok(None), + 1 => Ok(constraints.pop()), + _ => Err("Multiple structured constraints are not allowed".to_string()), + } + } +} + +// --------------------------------------------------------------------------- +// Proto → protocol type conversions +// --------------------------------------------------------------------------- + +impl From for openai_protocol::worker::SchedulerLoadSnapshot { + fn from(load: tokenspeed_proto::SchedulerLoad) -> Self { + Self { + dp_rank: load.dp_rank, + num_running_reqs: load.num_running_reqs, + num_waiting_reqs: load.num_waiting_reqs, + num_total_reqs: load.num_total_reqs, + num_used_tokens: load.num_used_tokens, + max_total_num_tokens: load.max_total_num_tokens, + token_usage: load.token_usage, + gen_throughput: load.gen_throughput, + cache_hit_rate: load.cache_hit_rate, + utilization: load.utilization, + max_running_requests: load.max_running_requests, + } + } +} + +impl From for openai_protocol::worker::WorkerLoadResponse { + fn from(resp: tokenspeed_proto::GetLoadsResponse) -> Self { + Self { + timestamp: resp.timestamp, + dp_rank_count: resp.dp_rank_count, + loads: resp.loads.into_iter().map(Into::into).collect(), + } + } +} diff --git a/crates/protocols/src/worker.rs b/crates/protocols/src/worker.rs index 7a99a9f3a5..ea013bc43b 100644 --- a/crates/protocols/src/worker.rs +++ b/crates/protocols/src/worker.rs @@ -204,6 +204,8 @@ pub enum RuntimeType { Trtllm, /// MLX runtime (Apple Silicon). Mlx, + /// TokenSpeed runtime. + TokenSpeed, /// External OpenAI-compatible API (not local inference). External, } @@ -223,6 +225,7 @@ impl std::fmt::Display for RuntimeType { RuntimeType::Vllm => write!(f, "vllm"), RuntimeType::Trtllm => write!(f, "trtllm"), RuntimeType::Mlx => write!(f, "mlx"), + RuntimeType::TokenSpeed => write!(f, "tokenspeed"), RuntimeType::External => write!(f, "external"), } } @@ -242,6 +245,8 @@ impl std::str::FromStr for RuntimeType { Ok(RuntimeType::Trtllm) } else if s.eq_ignore_ascii_case("mlx") { Ok(RuntimeType::Mlx) + } else if s.eq_ignore_ascii_case("tokenspeed") { + Ok(RuntimeType::TokenSpeed) } else if s.eq_ignore_ascii_case("external") { Ok(RuntimeType::External) } else { diff --git a/model_gateway/src/routers/grpc/client.rs b/model_gateway/src/routers/grpc/client.rs index 81cfc8f11f..370ef8103b 100644 --- a/model_gateway/src/routers/grpc/client.rs +++ b/model_gateway/src/routers/grpc/client.rs @@ -8,7 +8,7 @@ use openai_protocol::{ }; use smg_grpc_client::{ tokenizer_bundle, tokenizer_bundle::StreamBundle, MlxEngineClient, SglangSchedulerClient, - TrtllmServiceClient, VllmEngineClient, + TokenSpeedSchedulerClient, TrtllmServiceClient, VllmEngineClient, }; use crate::routers::grpc::{ @@ -23,13 +23,15 @@ pub struct HealthCheckResponse { pub message: String, } -/// Polymorphic gRPC client that wraps SGLang, vLLM, TensorRT-LLM, or MLX +/// Wraps the per-backend gRPC clients. RPCs absent on a backend's wire +/// return `Status::unimplemented`. #[derive(Clone)] pub enum GrpcClient { Sglang(SglangSchedulerClient), Vllm(VllmEngineClient), Trtllm(TrtllmServiceClient), Mlx(MlxEngineClient), + TokenSpeed(TokenSpeedSchedulerClient), } impl GrpcClient { @@ -137,6 +139,32 @@ impl GrpcClient { matches!(self, Self::Mlx(_)) } + #[expect( + clippy::panic, + reason = "typed accessor: caller guarantees variant via is_tokenspeed() check" + )] + pub fn as_tokenspeed(&self) -> &TokenSpeedSchedulerClient { + match self { + Self::TokenSpeed(client) => client, + _ => panic!("Expected TokenSpeed client"), + } + } + + #[expect( + clippy::panic, + reason = "typed accessor: caller guarantees variant via is_tokenspeed() check" + )] + pub fn as_tokenspeed_mut(&mut self) -> &mut TokenSpeedSchedulerClient { + match self { + Self::TokenSpeed(client) => client, + _ => panic!("Expected TokenSpeed client"), + } + } + + pub fn is_tokenspeed(&self) -> bool { + matches!(self, Self::TokenSpeed(_)) + } + pub async fn connect( url: &str, runtime_type: &str, @@ -146,6 +174,9 @@ impl GrpcClient { "vllm" => Ok(Self::Vllm(VllmEngineClient::connect(url).await?)), "trtllm" | "tensorrt-llm" => Ok(Self::Trtllm(TrtllmServiceClient::connect(url).await?)), "mlx" => Ok(Self::Mlx(MlxEngineClient::connect(url).await?)), + "tokenspeed" => Ok(Self::TokenSpeed( + TokenSpeedSchedulerClient::connect(url).await?, + )), _ => Err(format!("Unknown runtime type: {runtime_type}").into()), } } @@ -182,6 +213,13 @@ impl GrpcClient { message: resp.message, }) } + Self::TokenSpeed(client) => { + let resp = client.health_check().await?; + Ok(HealthCheckResponse { + healthy: resp.healthy, + message: resp.message, + }) + } } } @@ -191,24 +229,32 @@ impl GrpcClient { Self::Vllm(client) => Ok(ModelInfo::Vllm(client.get_model_info().await?)), Self::Trtllm(client) => Ok(ModelInfo::Trtllm(client.get_model_info().await?)), Self::Mlx(client) => Ok(ModelInfo::Mlx(client.get_model_info().await?)), + Self::TokenSpeed(client) => Ok(ModelInfo::TokenSpeed(Box::new( + client.get_model_info().await?, + ))), } } /// Get the full load response from the backend. - /// Only supported for SGLang backends. Returns per-DP-rank load metrics. + /// Returns `Unimplemented` for backends without scheduler load metrics. pub async fn get_loads(&self) -> Result { match self { Self::Sglang(client) => { let resp = client.get_loads(vec!["core".to_string()]).await?; Ok(WorkerLoadResponse::from(resp)) } + Self::TokenSpeed(client) => { + let resp = client.get_loads(vec!["core".to_string()]).await?; + Ok(WorkerLoadResponse::from(resp)) + } _ => Err(tonic::Status::unimplemented( "GetLoads RPC not supported for this backend", )), } } - /// Subscribe to KV cache events (all backends). + /// Subscribe to KV cache events. Returns `Unimplemented` on backends + /// without KV-event streaming. pub async fn subscribe_kv_events( &self, start_seq: u64, @@ -220,6 +266,9 @@ impl GrpcClient { Self::Mlx(_) => Err(tonic::Status::unimplemented( "SubscribeKvEvents RPC not supported for MLX backend", )), + Self::TokenSpeed(_) => Err(tonic::Status::unimplemented( + "SubscribeKvEvents RPC not supported for TokenSpeed backend", + )), } } @@ -231,6 +280,9 @@ impl GrpcClient { Self::Vllm(client) => Ok(ServerInfo::Vllm(client.get_server_info().await?)), Self::Trtllm(client) => Ok(ServerInfo::Trtllm(client.get_server_info().await?)), Self::Mlx(client) => Ok(ServerInfo::Mlx(client.get_server_info().await?)), + Self::TokenSpeed(client) => Ok(ServerInfo::TokenSpeed(Box::new( + client.get_server_info().await?, + ))), } } @@ -243,6 +295,11 @@ impl GrpcClient { Self::Vllm(client) => client.get_tokenizer().await, Self::Trtllm(client) => client.get_tokenizer().await, Self::Mlx(client) => client.get_tokenizer().await, + Self::TokenSpeed(_) => { + return Err(Box::new(tonic::Status::unimplemented( + "TokenSpeed backend does not support GetTokenizer RPC", + ))); + } }?; tokenizer_bundle::validate_bundle_sha256(&bundle).map_err(|e| { @@ -280,6 +337,10 @@ impl GrpcClient { let stream = client.generate(*boxed_req).await?; Ok(ProtoStream::Mlx(stream)) } + (Self::TokenSpeed(client), ProtoGenerateRequest::TokenSpeed(boxed_req)) => { + let stream = client.generate(*boxed_req).await?; + Ok(ProtoStream::TokenSpeed(stream)) + } #[expect( clippy::panic, reason = "client and request types are always matched by construction in the pipeline" @@ -301,6 +362,9 @@ impl GrpcClient { let resp = client.embed(*boxed_req).await?; Ok(ProtoEmbedComplete::Vllm(resp)) } + (Self::TokenSpeed(_), _) => Err(tonic::Status::unimplemented( + "TokenSpeed backend does not support embedding", + )), (Self::Mlx(_), _) => Err(tonic::Status::unimplemented( "MLX backend does not support embedding", )), @@ -382,6 +446,19 @@ impl GrpcClient { )?; Ok(ProtoGenerateRequest::Mlx(Box::new(req))) } + Self::TokenSpeed(client) => { + if multimodal_inputs.is_some() { + return Err("TokenSpeed backend does not support multimodal inputs".to_string()); + } + let req = client.build_generate_request_from_chat( + request_id, + body, + processed_text, + token_ids, + tool_constraints, + )?; + Ok(ProtoGenerateRequest::TokenSpeed(Box::new(req))) + } } } @@ -455,6 +532,19 @@ impl GrpcClient { )?; Ok(ProtoGenerateRequest::Mlx(Box::new(req))) } + Self::TokenSpeed(client) => { + if multimodal_inputs.is_some() { + return Err("TokenSpeed backend does not support multimodal inputs".to_string()); + } + let req = client.build_generate_request_from_messages( + request_id, + body, + processed_text, + token_ids, + tool_constraints, + )?; + Ok(ProtoGenerateRequest::TokenSpeed(Box::new(req))) + } } } @@ -502,6 +592,15 @@ impl GrpcClient { )?; Ok(ProtoGenerateRequest::Mlx(Box::new(req))) } + Self::TokenSpeed(client) => { + let req = client.build_generate_request_from_completion( + request_id, + body, + original_text, + token_ids, + )?; + Ok(ProtoGenerateRequest::TokenSpeed(Box::new(req))) + } } } @@ -549,6 +648,15 @@ impl GrpcClient { )?; Ok(ProtoGenerateRequest::Mlx(Box::new(req))) } + Self::TokenSpeed(client) => { + let req = client.build_plain_generate_request( + request_id, + body, + original_text, + token_ids, + )?; + Ok(ProtoGenerateRequest::TokenSpeed(Box::new(req))) + } } } } @@ -562,6 +670,7 @@ pub enum ModelInfo { Vllm(smg_grpc_client::vllm_proto::GetModelInfoResponse), Trtllm(smg_grpc_client::trtllm_proto::GetModelInfoResponse), Mlx(smg_grpc_client::mlx_proto::GetModelInfoResponse), + TokenSpeed(Box), } pub enum ServerInfo { @@ -569,6 +678,7 @@ pub enum ServerInfo { Vllm(smg_grpc_client::vllm_proto::GetServerInfoResponse), Trtllm(smg_grpc_client::trtllm_proto::GetServerInfoResponse), Mlx(smg_grpc_client::mlx_proto::GetServerInfoResponse), + TokenSpeed(Box), } impl ModelInfo { @@ -578,6 +688,7 @@ impl ModelInfo { ModelInfo::Vllm(info) => flat_labels(info), ModelInfo::Trtllm(info) => flat_labels(info), ModelInfo::Mlx(info) => flat_labels(info), + ModelInfo::TokenSpeed(info) => flat_labels(info), } } } @@ -600,6 +711,16 @@ impl ServerInfo { ServerInfo::Vllm(info) => flat_labels(info), ServerInfo::Trtllm(info) => flat_labels(info), ServerInfo::Mlx(info) => flat_labels(info), + ServerInfo::TokenSpeed(info) => { + let mut labels = HashMap::new(); + if let Some(ref args) = info.server_args { + pick_prost_fields(&mut labels, args, TOKENSPEED_GRPC_KEYS); + } + if !info.tokenspeed_version.is_empty() { + labels.insert("version".to_string(), info.tokenspeed_version.clone()); + } + labels + } } } } @@ -622,6 +743,24 @@ const SGLANG_GRPC_KEYS: &[&str] = &[ "weight_version", ]; +/// Keys worth extracting from TokenSpeed gRPC `server_args` (post-rename: bare +/// names, not `_path` variants — TokenSpeed dropped the legacy suffixes). +const TOKENSPEED_GRPC_KEYS: &[&str] = &[ + "model", + "served_model_name", + "tokenizer", + "tp_size", + "dp_size", + "pp_size", + "context_length", + "max_total_tokens", + "max_running_requests", + "load_balance_method", + "is_embedding", + "vocab_size", + "weight_version", +]; + // --------------------------------------------------------------------------- // Label helpers // --------------------------------------------------------------------------- diff --git a/model_gateway/src/routers/grpc/common/stages/helpers.rs b/model_gateway/src/routers/grpc/common/stages/helpers.rs index 65ef6ddedb..f780c4879f 100644 --- a/model_gateway/src/routers/grpc/common/stages/helpers.rs +++ b/model_gateway/src/routers/grpc/common/stages/helpers.rs @@ -6,7 +6,7 @@ use rand::Rng; use smg_grpc_client::{ mlx_proto, sglang_proto::{self, DisaggregatedParams}, - vllm_proto, + tokenspeed_proto, vllm_proto, }; use tracing::{debug, warn}; @@ -156,6 +156,15 @@ pub(crate) fn apply_sampling_defaults_to_generate_request( }; apply_mlx_sampling_defaults(params, defaults, mask); } + ProtoGenerateRequest::TokenSpeed(req) => { + let Some(params) = req.sampling_params.as_mut() else { + warn!( + "Cannot apply sampling defaults to TokenSpeed request without sampling_params" + ); + return; + }; + apply_tokenspeed_sampling_defaults(params, defaults, mask); + } ProtoGenerateRequest::Trtllm(_) => {} } } @@ -218,6 +227,30 @@ optional_temperature_sampling_defaults_fn!( ); optional_temperature_sampling_defaults_fn!(apply_mlx_sampling_defaults, mlx_proto::SamplingParams); +/// TokenSpeed declares every sampling scalar as `optional` so the servicer +/// can distinguish "client set 0" from "client unset". Apply defaults by +/// writing `Some(value)` rather than the bare value. +fn apply_tokenspeed_sampling_defaults( + params: &mut tokenspeed_proto::SamplingParams, + defaults: SamplingDefaults, + mask: SamplingDefaultsMask, +) { + macro_rules! apply_opt { + ($field:ident) => { + if mask.$field { + if let Some(value) = defaults.$field { + params.$field = Some(value); + } + } + }; + } + apply_opt!(temperature); + apply_opt!(top_p); + apply_opt!(top_k); + apply_opt!(min_p); + apply_opt!(repetition_penalty); +} + /// Inject PD bootstrap metadata for SGLang if needed. /// /// SGLang uses DisaggregatedParams with bootstrap host/port/room. diff --git a/model_gateway/src/routers/grpc/common/stages/request_execution.rs b/model_gateway/src/routers/grpc/common/stages/request_execution.rs index 07cf754966..950b6e4dc6 100644 --- a/model_gateway/src/routers/grpc/common/stages/request_execution.rs +++ b/model_gateway/src/routers/grpc/common/stages/request_execution.rs @@ -114,6 +114,7 @@ impl PipelineStage for RequestExecutionStage { } Some(RuntimeType::Trtllm) | Some(RuntimeType::Mlx) + | Some(RuntimeType::TokenSpeed) | Some(RuntimeType::External) | Some(RuntimeType::Unspecified) => { error!( diff --git a/model_gateway/src/routers/grpc/harmony/stages/request_building.rs b/model_gateway/src/routers/grpc/harmony/stages/request_building.rs index d084d66f3c..edce1e2eee 100644 --- a/model_gateway/src/routers/grpc/harmony/stages/request_building.rs +++ b/model_gateway/src/routers/grpc/harmony/stages/request_building.rs @@ -277,6 +277,50 @@ impl PipelineStage for HarmonyRequestBuildingStage { }; ProtoGenerateRequest::Mlx(Box::new(req)) } + GrpcClient::TokenSpeed(tokenspeed_client) => { + let req = match &ctx.input.request_type { + RequestType::Chat(request) => { + let body = modified_request.as_deref().unwrap_or_else(|| request.as_ref()); + tokenspeed_client + .build_generate_request_from_chat( + request_id, + body, + placeholder_processed_text, + token_ids, + tool_constraints, + ) + .map_err(|e| { + error!(function = "HarmonyRequestBuildingStage::execute", error = %e, "Failed to build TokenSpeed generate request"); + error::bad_request("invalid_request_parameters", format!("Invalid request parameters: {e}")) + })? + } + RequestType::Responses(request) => tokenspeed_client + .build_generate_request_from_responses( + request_id, + request.as_ref(), + placeholder_processed_text, + token_ids, + tool_constraints, + ) + .map_err(|e| { + error!(function = "HarmonyRequestBuildingStage::execute", error = %e, "Failed to build TokenSpeed generate request from responses"); + error::bad_request("invalid_request_parameters", format!("Invalid request parameters: {e}")) + })?, + RequestType::Embedding(_) => { + return Err(error::bad_request( + "harmony_embedding_not_supported", + "Embedding requests are not supported with Harmony models".to_string(), + )); + } + _ => { + return Err(error::bad_request( + "unsupported_request_type", + "Unsupported request type for Harmony models".to_string(), + )); + } + }; + ProtoGenerateRequest::TokenSpeed(Box::new(req)) + } }; // Inject Harmony stop token IDs into sampling params for ALL Harmony requests @@ -322,6 +366,15 @@ impl PipelineStage for HarmonyRequestBuildingStage { ); } } + ProtoGenerateRequest::TokenSpeed(req) => { + if let Some(params) = req.sampling_params.as_mut() { + params.stop_token_ids.extend_from_slice(&harmony_stop_ids); + debug!( + stop_token_count = harmony_stop_ids.len(), + "Injected Harmony stop tokens into TokenSpeed sampling params" + ); + } + } } } diff --git a/model_gateway/src/routers/grpc/multimodal.rs b/model_gateway/src/routers/grpc/multimodal.rs index d28c7e3fcc..a30d421354 100644 --- a/model_gateway/src/routers/grpc/multimodal.rs +++ b/model_gateway/src/routers/grpc/multimodal.rs @@ -708,6 +708,9 @@ pub(crate) fn assemble_multimodal_data( GrpcClient::Mlx(_) => unreachable!( "caller rejects multimodal for MLX in build_chat_request/build_messages_request" ), + GrpcClient::TokenSpeed(_) => unreachable!( + "TokenSpeed backend does not support multimodal; preparation stage should reject earlier" + ), } } diff --git a/model_gateway/src/routers/grpc/proto_wrapper.rs b/model_gateway/src/routers/grpc/proto_wrapper.rs index 971ff388b8..48c7aef229 100644 --- a/model_gateway/src/routers/grpc/proto_wrapper.rs +++ b/model_gateway/src/routers/grpc/proto_wrapper.rs @@ -1,7 +1,8 @@ -//! Protocol buffer type wrappers for SGLang, vLLM, and TensorRT-LLM backends +//! Protocol buffer type wrappers for the supported gRPC backends. //! -//! This module provides unified enums that wrap proto types from SGLang, vLLM, and TensorRT-LLM, -//! allowing the router to work with any backend transparently. +//! This module provides unified enums that wrap proto types from each +//! supported backend, allowing the router to work with any backend +//! transparently. use std::collections::HashMap; @@ -11,6 +12,10 @@ use smg_grpc_client::{ mlx_proto::{self as mlx}, sglang_proto::{self as sglang, generate_complete::MatchedStop as SglangMatchedStop}, sglang_scheduler::AbortOnDropStream as SglangStream, + tokenspeed_proto::{ + self as tokenspeed, generate_complete::MatchedStop as TokenSpeedMatchedStop, + }, + tokenspeed_scheduler::AbortOnDropStream as TokenSpeedStream, trtllm_proto::{self as trtllm, generate_complete::MatchedStop as TrtllmMatchedStop}, trtllm_service::AbortOnDropStream as TrtllmStream, vllm_engine::AbortOnDropStream as VllmStream, @@ -280,6 +285,7 @@ pub enum ProtoGenerateRequest { Vllm(Box), Trtllm(Box), Mlx(Box), + TokenSpeed(Box), } impl ProtoGenerateRequest { @@ -355,6 +361,30 @@ impl ProtoGenerateRequest { } } + /// Get TokenSpeed variant (panics if not TokenSpeed) + #[expect( + clippy::panic, + reason = "typed accessor: caller guarantees variant via is_tokenspeed() check" + )] + pub fn as_tokenspeed(&self) -> &tokenspeed::GenerateRequest { + match self { + Self::TokenSpeed(req) => req, + _ => panic!("Expected TokenSpeed GenerateRequest"), + } + } + + /// Get mutable TokenSpeed variant (panics if not TokenSpeed) + #[expect( + clippy::panic, + reason = "typed accessor: caller guarantees variant via is_tokenspeed() check" + )] + pub fn as_tokenspeed_mut(&mut self) -> &mut tokenspeed::GenerateRequest { + match self { + Self::TokenSpeed(req) => req, + _ => panic!("Expected TokenSpeed GenerateRequest"), + } + } + /// Check if this is SGLang pub fn is_sglang(&self) -> bool { matches!(self, Self::Sglang(_)) @@ -370,6 +400,11 @@ impl ProtoGenerateRequest { matches!(self, Self::Trtllm(_)) } + /// Check if this is TokenSpeed + pub fn is_tokenspeed(&self) -> bool { + matches!(self, Self::TokenSpeed(_)) + } + /// Set max_tokens for prefill-only execution (vLLM PD mode). /// The prefill request uses max_tokens=1 to trigger KV cache computation /// without generating unnecessary tokens. @@ -385,7 +420,7 @@ impl ProtoGenerateRequest { }); } } - Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) => { + Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) | Self::TokenSpeed(_) => { tracing::warn!("set_max_tokens_for_prefill called on non-vLLM request, ignoring"); } } @@ -398,6 +433,7 @@ impl ProtoGenerateRequest { Self::Sglang(req) => req.stream = stream, Self::Trtllm(req) => req.streaming = stream, Self::Mlx(req) => req.stream = stream, + Self::TokenSpeed(req) => req.stream = stream, } } @@ -415,7 +451,8 @@ impl ProtoGenerateRequest { match self { Self::Sglang(req) => req.mm_inputs = None, Self::Vllm(req) => req.mm_inputs = None, - Self::Trtllm(_) | Self::Mlx(_) => {} // TRT-LLM and MLX protos have no mm_inputs field + // TRT-LLM, MLX, and TokenSpeed protos have no mm_inputs field + Self::Trtllm(_) | Self::Mlx(_) | Self::TokenSpeed(_) => {} } } @@ -426,6 +463,7 @@ impl ProtoGenerateRequest { Self::Vllm(req) => &req.request_id, Self::Trtllm(req) => &req.request_id, Self::Mlx(req) => &req.request_id, + Self::TokenSpeed(req) => &req.request_id, } } @@ -439,7 +477,7 @@ impl ProtoGenerateRequest { remote_port, }); } - Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) => { + Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) | Self::TokenSpeed(_) => { tracing::warn!("set_kv_transfer_params called on non-vLLM request, ignoring"); } } @@ -452,6 +490,7 @@ pub enum ProtoGenerateResponse { Vllm(Box), Trtllm(Box), Mlx(Box), + TokenSpeed(Box), } impl ProtoGenerateResponse { @@ -496,6 +535,15 @@ impl ProtoGenerateResponse { } None => ProtoResponseVariant::None, }, + Self::TokenSpeed(resp) => match resp.response { + Some(tokenspeed::generate_response::Response::Chunk(chunk)) => { + ProtoResponseVariant::Chunk(ProtoGenerateStreamChunk::TokenSpeed(chunk)) + } + Some(tokenspeed::generate_response::Response::Complete(complete)) => { + ProtoResponseVariant::Complete(ProtoGenerateComplete::TokenSpeed(complete)) + } + None => ProtoResponseVariant::None, + }, } } } @@ -514,6 +562,7 @@ pub enum ProtoGenerateStreamChunk { Vllm(vllm::GenerateStreamChunk), Trtllm(trtllm::GenerateStreamChunk), Mlx(mlx::GenerateStreamChunk), + TokenSpeed(tokenspeed::GenerateStreamChunk), } impl ProtoGenerateStreamChunk { @@ -573,6 +622,11 @@ impl ProtoGenerateStreamChunk { matches!(self, Self::Mlx(_)) } + /// Check if this is TokenSpeed + pub fn is_tokenspeed(&self) -> bool { + matches!(self, Self::TokenSpeed(_)) + } + /// Get token IDs from chunk (common field) pub fn token_ids(&self) -> &[u32] { match self { @@ -580,6 +634,7 @@ impl ProtoGenerateStreamChunk { Self::Vllm(c) => &c.token_ids, Self::Trtllm(c) => &c.token_ids, Self::Mlx(c) => &c.token_ids, + Self::TokenSpeed(c) => &c.token_ids, } } @@ -591,10 +646,11 @@ impl ProtoGenerateStreamChunk { Self::Vllm(c) => c.index, Self::Trtllm(c) => c.sequence_index, Self::Mlx(c) => c.index, + Self::TokenSpeed(c) => c.index, } } - /// Get output logprobs (SGLang, vLLM, TensorRT-LLM, and MLX) + /// Get output logprobs. pub fn output_logprobs(&self) -> Option { match self { Self::Sglang(c) => c @@ -610,6 +666,10 @@ impl ProtoGenerateStreamChunk { .output_logprobs .as_ref() .map(|lp| convert_output_logprobs!(lp)), + Self::TokenSpeed(c) => c + .output_logprobs + .as_ref() + .map(|lp| convert_output_logprobs!(lp)), } } @@ -624,8 +684,8 @@ impl ProtoGenerateStreamChunk { .input_logprobs .as_ref() .map(|lp| convert_input_logprobs!(lp)), - // TRT-LLM and MLX streaming chunks don't have input_logprobs - Self::Trtllm(_) | Self::Mlx(_) => None, + // TRT-LLM, MLX, and TokenSpeed streaming chunks don't have input_logprobs + Self::Trtllm(_) | Self::Mlx(_) | Self::TokenSpeed(_) => None, } } @@ -636,6 +696,7 @@ impl ProtoGenerateStreamChunk { Self::Vllm(c) => c.prompt_tokens, Self::Trtllm(c) => c.prompt_tokens, Self::Mlx(c) => c.prompt_tokens, + Self::TokenSpeed(c) => c.prompt_tokens, } } @@ -646,6 +707,7 @@ impl ProtoGenerateStreamChunk { Self::Vllm(c) => c.completion_tokens, Self::Trtllm(c) => c.completion_tokens, Self::Mlx(c) => c.completion_tokens, + Self::TokenSpeed(c) => c.completion_tokens, } } @@ -656,6 +718,7 @@ impl ProtoGenerateStreamChunk { Self::Vllm(c) => c.cached_tokens, Self::Trtllm(c) => c.cached_tokens, Self::Mlx(c) => c.cached_tokens, + Self::TokenSpeed(c) => c.cached_tokens, } } } @@ -667,6 +730,7 @@ pub enum ProtoGenerateComplete { Vllm(vllm::GenerateComplete), Trtllm(trtllm::GenerateComplete), Mlx(mlx::GenerateComplete), + TokenSpeed(tokenspeed::GenerateComplete), } impl ProtoGenerateComplete { @@ -738,6 +802,11 @@ impl ProtoGenerateComplete { matches!(self, Self::Mlx(_)) } + /// Check if this is TokenSpeed + pub fn is_tokenspeed(&self) -> bool { + matches!(self, Self::TokenSpeed(_)) + } + /// Get token IDs from either backend (output_ids in proto) pub fn token_ids(&self) -> &[u32] { match self { @@ -745,6 +814,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => &c.output_ids, Self::Trtllm(c) => &c.output_token_ids, Self::Mlx(c) => &c.output_ids, + Self::TokenSpeed(c) => &c.output_ids, } } @@ -755,6 +825,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => c.prompt_tokens, Self::Trtllm(c) => c.prompt_tokens, Self::Mlx(c) => c.prompt_tokens, + Self::TokenSpeed(c) => c.prompt_tokens, } } @@ -765,6 +836,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => c.completion_tokens, Self::Trtllm(c) => c.completion_tokens, Self::Mlx(c) => c.completion_tokens, + Self::TokenSpeed(c) => c.completion_tokens, } } @@ -775,6 +847,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => &c.finish_reason, Self::Trtllm(c) => &c.finish_reason, Self::Mlx(c) => &c.finish_reason, + Self::TokenSpeed(c) => &c.finish_reason, } } @@ -786,6 +859,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => c.index, Self::Trtllm(c) => c.sequence_index, Self::Mlx(c) => c.index, + Self::TokenSpeed(c) => c.index, } } @@ -823,6 +897,11 @@ impl ProtoGenerateComplete { Self::Mlx(c) => c .matched_stop_token_id .map(|id| serde_json::Value::Number(id.into())), + Self::TokenSpeed(c) => convert!( + &c.matched_stop, + TokenSpeedMatchedStop::MatchedTokenId, + TokenSpeedMatchedStop::MatchedStopStr + ), } } @@ -833,6 +912,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => &c.output_ids, Self::Trtllm(c) => &c.output_token_ids, Self::Mlx(c) => &c.output_ids, + Self::TokenSpeed(c) => &c.output_ids, } } @@ -843,6 +923,7 @@ impl ProtoGenerateComplete { Self::Vllm(c) => c.cached_tokens, Self::Trtllm(c) => c.cached_tokens, Self::Mlx(c) => c.cached_tokens, + Self::TokenSpeed(c) => c.cached_tokens, } } @@ -881,12 +962,12 @@ impl ProtoGenerateComplete { }) } } - // MLX does not have input_logprobs - Self::Mlx(_) => None, + // MLX and TokenSpeed do not have input_logprobs + Self::Mlx(_) | Self::TokenSpeed(_) => None, } } - /// Get output logprobs (SGLang, vLLM, TensorRT-LLM, and MLX) + /// Get output logprobs. pub fn output_logprobs(&self) -> Option { match self { Self::Sglang(c) => c @@ -902,6 +983,10 @@ impl ProtoGenerateComplete { .output_logprobs .as_ref() .map(|lp| convert_output_logprobs!(lp)), + Self::TokenSpeed(c) => c + .output_logprobs + .as_ref() + .map(|lp| convert_output_logprobs!(lp)), } } @@ -913,17 +998,22 @@ impl ProtoGenerateComplete { .kv_transfer_params .as_ref() .map(|params| (params.remote_host.clone(), params.remote_port)), - Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) => None, + Self::Sglang(_) | Self::Trtllm(_) | Self::Mlx(_) | Self::TokenSpeed(_) => None, } } } -/// Unified stream wrapper +/// Unified stream wrapper. +/// +/// One variant per backend. Each yields its own native proto response shape; +/// the chunk / complete accessors above match on the corresponding +/// `ProtoGenerateStreamChunk` / `ProtoGenerateComplete` arm. pub enum ProtoStream { Sglang(SglangStream), Vllm(VllmStream), Trtllm(TrtllmStream), Mlx(MlxStream), + TokenSpeed(TokenSpeedStream), } impl ProtoStream { @@ -946,6 +1036,10 @@ impl ProtoStream { .next() .await .map(|result| result.map(|r| ProtoGenerateResponse::Mlx(Box::new(r)))), + Self::TokenSpeed(stream) => stream + .next() + .await + .map(|result| result.map(|r| ProtoGenerateResponse::TokenSpeed(Box::new(r)))), } } @@ -956,6 +1050,7 @@ impl ProtoStream { Self::Vllm(stream) => stream.mark_completed(), Self::Trtllm(stream) => stream.mark_completed(), Self::Mlx(stream) => stream.mark_completed(), + Self::TokenSpeed(stream) => stream.mark_completed(), } } } diff --git a/model_gateway/src/routers/grpc/regular/stages/embedding/request_building.rs b/model_gateway/src/routers/grpc/regular/stages/embedding/request_building.rs index b9ee0a8456..c69eef5b93 100644 --- a/model_gateway/src/routers/grpc/regular/stages/embedding/request_building.rs +++ b/model_gateway/src/routers/grpc/regular/stages/embedding/request_building.rs @@ -96,6 +96,16 @@ impl PipelineStage for EmbeddingRequestBuildingStage { "MLX embedding is not supported via gRPC", )); } + GrpcClient::TokenSpeed(_) => { + error!( + function = "EmbeddingRequestBuildingStage::execute", + "TokenSpeed backend does not support embeddings" + ); + return Err(error::not_implemented( + "unsupported_backend", + "TokenSpeed backend does not support embeddings", + )); + } }; ctx.state.proto_request = Some(ProtoRequest::Embed(proto_req)); diff --git a/model_gateway/src/workflow/steps/local/detect_backend.rs b/model_gateway/src/workflow/steps/local/detect_backend.rs index 3295d6f826..235b672af5 100644 --- a/model_gateway/src/workflow/steps/local/detect_backend.rs +++ b/model_gateway/src/workflow/steps/local/detect_backend.rs @@ -1,8 +1,8 @@ //! Backend runtime detection step. //! -//! Detects the runtime type (sglang, vllm, trtllm, mlx) for both HTTP and gRPC workers. +//! Detects the runtime type (sglang, vllm, trtllm, tokenspeed, mlx) for both HTTP and gRPC workers. //! - HTTP: probes `/v1/models` (owned_by field), falls back to unique endpoints. -//! - gRPC: tries sglang → vllm → trtllm → mlx health checks sequentially. +//! - gRPC: tries sglang → vllm → trtllm → tokenspeed → mlx health checks sequentially. use std::time::Duration; @@ -44,7 +44,7 @@ async fn detect_grpc_backend( } // Try each runtime sequentially (most common first), skipping the hint we already tried - for runtime in &["sglang", "vllm", "trtllm", "mlx"] { + for runtime in &["sglang", "vllm", "trtllm", "tokenspeed", "mlx"] { if Some(*runtime) == runtime_hint { continue; } @@ -57,7 +57,7 @@ async fn detect_grpc_backend( } Err(format!( - "gRPC backend detection failed for {url} (tried sglang, vllm, trtllm, mlx)" + "gRPC backend detection failed for {url} (tried sglang, vllm, trtllm, tokenspeed, mlx)" )) } diff --git a/model_gateway/src/workflow/steps/util.rs b/model_gateway/src/workflow/steps/util.rs index d6ae1a19c2..ec72dcf902 100644 --- a/model_gateway/src/workflow/steps/util.rs +++ b/model_gateway/src/workflow/steps/util.rs @@ -120,17 +120,22 @@ pub(crate) async fn do_grpc_health_check( pub(crate) async fn try_grpc_reachable(url: &str, timeout_secs: u64) -> Result<(), String> { let grpc_url = grpc_reachable_url(url)?; - let (sglang, vllm, trtllm, mlx) = tokio::join!( + let (sglang, vllm, trtllm, mlx, tokenspeed) = tokio::join!( do_grpc_health_check(&grpc_url, timeout_secs, "sglang"), do_grpc_health_check(&grpc_url, timeout_secs, "vllm"), do_grpc_health_check(&grpc_url, timeout_secs, "trtllm"), do_grpc_health_check(&grpc_url, timeout_secs, "mlx"), + do_grpc_health_check(&grpc_url, timeout_secs, "tokenspeed"), ); - match (sglang, vllm, trtllm, mlx) { - (Ok(()), _, _, _) | (_, Ok(()), _, _) | (_, _, Ok(()), _) | (_, _, _, Ok(())) => Ok(()), - (Err(e1), Err(e2), Err(e3), Err(e4)) => Err(format!( - "gRPC not reachable (tried sglang, vllm, trtllm, mlx): sglang={e1}, vllm={e2}, trtllm={e3}, mlx={e4}", + match (sglang, vllm, trtllm, mlx, tokenspeed) { + (Ok(()), _, _, _, _) + | (_, Ok(()), _, _, _) + | (_, _, Ok(()), _, _) + | (_, _, _, Ok(()), _) + | (_, _, _, _, Ok(())) => Ok(()), + (Err(e1), Err(e2), Err(e3), Err(e4), Err(e5)) => Err(format!( + "gRPC not reachable (tried sglang, vllm, trtllm, mlx, tokenspeed): sglang={e1}, vllm={e2}, trtllm={e3}, mlx={e4}, tokenspeed={e5}", )), } }