diff --git a/sgl-model-gateway/src/core/mod.rs b/sgl-model-gateway/src/core/mod.rs index ada1d2a2e..048262342 100644 --- a/sgl-model-gateway/src/core/mod.rs +++ b/sgl-model-gateway/src/core/mod.rs @@ -9,6 +9,12 @@ //! - Workflow steps for multi-step operations //! - Common utilities +/// Default model identifier used when no model is specified. +/// +/// This constant should be used instead of hardcoded "unknown" strings +/// throughout the codebase for consistency. +pub const UNKNOWN_MODEL_ID: &str = "unknown"; + pub mod circuit_breaker; pub mod error; pub mod job_queue; diff --git a/sgl-model-gateway/src/core/model_card.rs b/sgl-model-gateway/src/core/model_card.rs index b0da07e82..3699121d6 100644 --- a/sgl-model-gateway/src/core/model_card.rs +++ b/sgl-model-gateway/src/core/model_card.rs @@ -9,7 +9,10 @@ use serde::{Deserialize, Serialize}; -use super::model_type::{Endpoint, ModelType}; +use super::{ + model_type::{Endpoint, ModelType}, + UNKNOWN_MODEL_ID, +}; /// Provider type for external API transformations. /// @@ -335,7 +338,7 @@ impl ModelCard { impl Default for ModelCard { fn default() -> Self { - Self::new("unknown") + Self::new(UNKNOWN_MODEL_ID) } } diff --git a/sgl-model-gateway/src/policies/bucket.rs b/sgl-model-gateway/src/policies/bucket.rs index 473a20623..bdf64020b 100644 --- a/sgl-model-gateway/src/policies/bucket.rs +++ b/sgl-model-gateway/src/policies/bucket.rs @@ -10,7 +10,7 @@ use rand::Rng; use tracing::{debug, error, info, warn}; use uuid::Uuid; -use super::{get_healthy_worker_indices, BucketConfig, LoadBalancingPolicy}; +use super::{get_healthy_worker_indices, normalize_model_key, BucketConfig, LoadBalancingPolicy}; use crate::core::Worker; #[derive(Debug)] @@ -79,13 +79,7 @@ impl BucketPolicy { // Group workers by model let mut model_workers: HashMap>> = HashMap::new(); for worker in prefill_workers { - // Use "unknown" for empty model_ids - let model_id = worker.model_id(); - let model_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let model_key = normalize_model_key(worker.model_id()); model_workers .entry(model_key.to_string()) .or_default() @@ -118,12 +112,7 @@ impl BucketPolicy { } pub fn add_prefill_url(&self, worker: &dyn Worker) { - let model_id = worker.model_id(); - let model_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let model_key = normalize_model_key(worker.model_id()); let bucket = self .buckets .entry(model_key.to_string()) @@ -166,12 +155,7 @@ impl BucketPolicy { } pub fn remove_prefill_url(&self, worker: &dyn Worker) { - let model_id = worker.model_id(); - let model_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let model_key = normalize_model_key(worker.model_id()); if let Some(bucket_entry) = self.buckets.get(model_key) { let bucket = bucket_entry.value(); @@ -235,12 +219,7 @@ impl LoadBalancingPolicy for BucketPolicy { // Determine the model for this set of workers (router pre-filters by model) // All workers should be from the same model - let first_model = workers[healthy_indices[0]].model_id(); - let model_key = if first_model.is_empty() { - "unknown" - } else { - first_model - }; + let model_key = normalize_model_key(workers[healthy_indices[0]].model_id()); let bucket = self .buckets diff --git a/sgl-model-gateway/src/policies/cache_aware.rs b/sgl-model-gateway/src/policies/cache_aware.rs index 01b07b063..ba528d8e8 100644 --- a/sgl-model-gateway/src/policies/cache_aware.rs +++ b/sgl-model-gateway/src/policies/cache_aware.rs @@ -72,7 +72,10 @@ use dashmap::DashMap; use rand::Rng; use tracing::debug; -use super::{get_healthy_worker_indices, tree::Tree, CacheAwareConfig, LoadBalancingPolicy}; +use super::{ + get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig, + LoadBalancingPolicy, +}; use crate::core::Worker; /// Cache-aware routing policy @@ -160,13 +163,7 @@ impl CacheAwarePolicy { let mut model_workers: std::collections::HashMap>> = std::collections::HashMap::new(); for worker in workers { - // Use "unknown" for empty model_ids - let model_id = worker.model_id(); - let tree_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let tree_key = normalize_model_key(worker.model_id()); model_workers .entry(tree_key.to_string()) .or_default() @@ -187,14 +184,7 @@ impl CacheAwarePolicy { /// Add a single worker to the tree (incremental update) pub fn add_worker(&self, worker: &dyn Worker) { - // For backward compatibility: if model_id is "unknown" or empty, - // use a default tree. This preserves existing behavior for single-model routers. - let model_id = worker.model_id(); - let tree_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let tree_key = normalize_model_key(worker.model_id()); let tree = self .trees .entry(tree_key.to_string()) @@ -213,13 +203,7 @@ impl CacheAwarePolicy { /// Remove a worker from the tree pub fn remove_worker(&self, worker: &dyn Worker) { - // Use same logic as add_worker for consistency - let model_id = worker.model_id(); - let tree_key = if model_id.is_empty() { - "unknown" - } else { - model_id - }; + let tree_key = normalize_model_key(worker.model_id()); if let Some(tree) = self.trees.get(tree_key) { tree.remove_tenant(worker.url()); } @@ -313,12 +297,7 @@ impl LoadBalancingPolicy for CacheAwarePolicy { // Determine the model for this set of workers (router pre-filters by model) // All workers should be from the same model - let first_model = workers[healthy_indices[0]].model_id(); - let model_id = if first_model.is_empty() { - "unknown" - } else { - first_model - }; + let model_id = normalize_model_key(workers[healthy_indices[0]].model_id()); // Get current load statistics - compute min/max in single pass without allocation let (min_load, max_load) = workers.iter().fold((usize::MAX, 0usize), |(min, max), w| { diff --git a/sgl-model-gateway/src/policies/mod.rs b/sgl-model-gateway/src/policies/mod.rs index f74483d6a..e8ef5fa95 100644 --- a/sgl-model-gateway/src/policies/mod.rs +++ b/sgl-model-gateway/src/policies/mod.rs @@ -122,6 +122,19 @@ pub(crate) fn get_healthy_worker_indices(workers: &[Arc]) -> Vec &str { + if model_id.is_empty() { + crate::core::UNKNOWN_MODEL_ID + } else { + model_id + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/sgl-model-gateway/src/routers/grpc/common/stages/dispatch_metadata.rs b/sgl-model-gateway/src/routers/grpc/common/stages/dispatch_metadata.rs index ffb51f086..d07c54314 100644 --- a/sgl-model-gateway/src/routers/grpc/common/stages/dispatch_metadata.rs +++ b/sgl-model-gateway/src/routers/grpc/common/stages/dispatch_metadata.rs @@ -7,9 +7,12 @@ use axum::response::Response; use tracing::error; use super::PipelineStage; -use crate::routers::{ - error, - grpc::context::{DispatchMetadata, RequestContext, RequestType, WorkerSelection}, +use crate::{ + core::UNKNOWN_MODEL_ID, + routers::{ + error, + grpc::context::{DispatchMetadata, RequestContext, RequestType, WorkerSelection}, + }, }; /// Dispatch metadata stage: Prepare metadata for dispatch @@ -31,11 +34,11 @@ impl PipelineStage for DispatchMetadataStage { RequestType::Chat(req) => req.model.clone(), RequestType::Generate(_req) => { // Generate requests don't have a model field - // Use model_id from input or unknown + // Use model_id from input or UNKNOWN_MODEL_ID ctx.input .model_id .clone() - .unwrap_or_else(|| "unknown".to_string()) + .unwrap_or_else(|| UNKNOWN_MODEL_ID.to_string()) } RequestType::Responses(req) => req.model.clone(), RequestType::Embedding(req) => req.model.clone(), diff --git a/sgl-model-gateway/src/routers/grpc/context.rs b/sgl-model-gateway/src/routers/grpc/context.rs index 23438dd49..a5da72f01 100644 --- a/sgl-model-gateway/src/routers/grpc/context.rs +++ b/sgl-model-gateway/src/routers/grpc/context.rs @@ -22,7 +22,7 @@ use crate::{ responses::ResponsesRequest, }, reasoning_parser::ParserFactory as ReasoningParserFactory, - tokenizer::{stop::StopSequenceDecoder, TokenizerRegistry}, + tokenizer::{stop::StopSequenceDecoder, traits::Tokenizer, TokenizerRegistry}, tool_parser::ParserFactory as ToolParserFactory, }; @@ -66,6 +66,10 @@ pub struct ProcessingState { // Stage 1: Preparation outputs pub preparation: Option, + /// Resolved tokenizer (set once in preparation, reused in response processing) + /// This avoids redundant registry lookups across pipeline stages. + pub tokenizer: Option>, + // Stage 2: Worker selection outputs pub workers: Option, @@ -394,6 +398,14 @@ impl RequestContext { RequestType::Embedding(_) => false, // Embeddings are never streaming } } + + /// Get the cached tokenizer, cloning the Arc (cheap 8-byte clone) + /// + /// Returns None if tokenizer hasn't been resolved yet. + /// The tokenizer is resolved once in the preparation stage and cached for reuse. + pub fn tokenizer_arc(&self) -> Option> { + self.state.tokenizer.clone() + } } impl WorkerSelection { diff --git a/sgl-model-gateway/src/routers/grpc/pipeline.rs b/sgl-model-gateway/src/routers/grpc/pipeline.rs index 894fee4f0..19b390849 100644 --- a/sgl-model-gateway/src/routers/grpc/pipeline.rs +++ b/sgl-model-gateway/src/routers/grpc/pipeline.rs @@ -203,7 +203,6 @@ impl RequestPipeline { pub fn new_embeddings( worker_registry: Arc, policy_registry: Arc, - _tokenizer: Arc, ) -> Self { let stages: Vec> = vec![ Box::new(EmbeddingPreparationStage::new()), diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/preparation.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/preparation.rs index 9166a7a90..ee951ab64 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/preparation.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/preparation.rs @@ -43,23 +43,9 @@ impl ChatPreparationStage { ctx: &mut RequestContext, request: &ChatCompletionRequest, ) -> Result<(), Response> { - // Step 0: Resolve tokenizer from registry - let model_id = ctx.input.model_id.as_deref().unwrap(); - let tokenizer = ctx - .components - .tokenizer_registry - .get(model_id) - .ok_or_else(|| { - error!( - function = "ChatPreparationStage::prepare_chat", - model = %model_id, - "Tokenizer not found for model" - ); - error::internal_error( - "tokenizer_not_found", - format!("Tokenizer not found for model: {}", model_id), - ) - })?; + // Step 0: Resolve tokenizer from registry (cached for reuse in response processing) + let tokenizer = + utils::resolve_tokenizer(ctx, "ChatPreparationStage::prepare_chat").map_err(|e| *e)?; // Step 1: Filter tools if needed let body_ref = utils::filter_chat_request_by_tool_choice(request); diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs index fc17f0a6b..2d3001c17 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs @@ -79,23 +79,17 @@ impl ChatResponseProcessingStage { })? .clone(); - // Get tokenizer in real time - let model_id = ctx.input.model_id.as_deref().unwrap(); - let tokenizer = ctx - .components - .tokenizer_registry - .get(model_id) - .ok_or_else(|| { - error!( - function = "ChatPreparationStage::prepare_chat", - model = model_id, - "Tokenizer not found for model" - ); - error::internal_error( - "tokenizer_not_found", - format!("Tokenizer not found for model: {}", model_id), - ) - })?; + // Get cached tokenizer (resolved once in preparation stage) + let tokenizer = ctx.tokenizer_arc().ok_or_else(|| { + error!( + function = "ChatResponseProcessingStage::process_chat_response", + "Tokenizer not cached in context" + ); + error::internal_error( + "tokenizer_not_cached", + "Tokenizer not cached in context - preparation stage may have been skipped", + ) + })?; if is_streaming { // Streaming: Use StreamingProcessor and return SSE response diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/embedding/preparation.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/embedding/preparation.rs index b3c77a837..fae235a3d 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/embedding/preparation.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/embedding/preparation.rs @@ -11,6 +11,7 @@ use crate::{ grpc::{ common::stages::PipelineStage, context::{PreparationOutput, RequestContext, RequestType}, + utils, }, }, }; @@ -55,10 +56,12 @@ impl PipelineStage for EmbeddingPreparationStage { )); } + // Resolve tokenizer from registry (cached for potential reuse) + let tokenizer = + utils::resolve_tokenizer(ctx, "EmbeddingPreparationStage::execute").map_err(|e| *e)?; + // Tokenize - let token_ids = ctx - .components - .tokenizer + let token_ids = tokenizer .encode(&text) .map_err(|e| { error!( diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/preparation.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/preparation.rs index a9d46d140..d1aea48cb 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/preparation.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/preparation.rs @@ -44,24 +44,9 @@ impl GeneratePreparationStage { ctx: &mut RequestContext, request: &GenerateRequest, ) -> Result<(), Response> { - // Get model_id from context (normalized by router_manager) - let model_id = ctx.input.model_id.as_deref().unwrap(); - - let tokenizer = ctx - .components - .tokenizer_registry - .get(model_id) - .ok_or_else(|| { - error!( - function = "GeneratePreparationStage::execute", - model = %model_id, - "Tokenizer not found for model" - ); - error::internal_error( - "tokenizer_not_found", - format!("Tokenizer not found for model: {}", model_id), - ) - })?; + // Resolve tokenizer from registry (cached for reuse in response processing) + let tokenizer = utils::resolve_tokenizer(ctx, "GeneratePreparationStage::prepare_generate") + .map_err(|e| *e)?; let (original_text, token_ids) = match self.resolve_generate_input(request, &tokenizer) { Ok(res) => res, diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs index e6b6012e3..b214e9cf4 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs @@ -77,24 +77,17 @@ impl GenerateResponseProcessingStage { })? .clone(); - // Get model_id from context - let model_id = ctx.input.model_id.as_deref().unwrap(); - - let tokenizer = ctx - .components - .tokenizer_registry - .get(model_id) - .ok_or_else(|| { - error!( - function = "GeneratePreparationStage::execute", - model = %model_id, - "Tokenizer not found for model" - ); - error::internal_error( - "tokenizer_not_found", - format!("Tokenizer not found for model: {}", model_id), - ) - })?; + // Get cached tokenizer (resolved once in preparation stage) + let tokenizer = ctx.tokenizer_arc().ok_or_else(|| { + error!( + function = "GenerateResponseProcessingStage::process_generate_response", + "Tokenizer not cached in context" + ); + error::internal_error( + "tokenizer_not_cached", + "Tokenizer not cached in context - preparation stage may have been skipped", + ) + })?; if is_streaming { // Streaming: Use StreamingProcessor and return SSE response diff --git a/sgl-model-gateway/src/routers/grpc/router.rs b/sgl-model-gateway/src/routers/grpc/router.rs index 31adff8b6..fc8831f98 100644 --- a/sgl-model-gateway/src/routers/grpc/router.rs +++ b/sgl-model-gateway/src/routers/grpc/router.rs @@ -96,11 +96,8 @@ impl GrpcRouter { ); // Create Embedding pipeline - let embedding_pipeline = RequestPipeline::new_embeddings( - worker_registry.clone(), - _policy_registry.clone(), - tokenizer.clone(), - ); + let embedding_pipeline = + RequestPipeline::new_embeddings(worker_registry.clone(), _policy_registry.clone()); // Extract shared dependencies for responses contexts let mcp_manager = ctx diff --git a/sgl-model-gateway/src/routers/grpc/utils.rs b/sgl-model-gateway/src/routers/grpc/utils.rs index 3f0770ca0..83e3d11b4 100644 --- a/sgl-model-gateway/src/routers/grpc/utils.rs +++ b/sgl-model-gateway/src/routers/grpc/utils.rs @@ -10,6 +10,7 @@ use uuid::Uuid; use super::{ client::GrpcClient, + context::RequestContext, proto_wrapper::{ProtoGenerateComplete, ProtoStream}, ProcessedMessages, }; @@ -42,6 +43,49 @@ use crate::{ }, }; +/// Resolve tokenizer from registry and cache it in request context. +/// +/// This is a helper to avoid duplicating tokenizer resolution logic across +/// preparation stages (chat, generate, embedding). +/// +/// Returns the tokenizer Arc, which is also cached in `ctx.state.tokenizer`. +pub fn resolve_tokenizer( + ctx: &mut RequestContext, + stage_name: &str, +) -> Result, Box> { + let model_id = ctx.input.model_id.as_deref().ok_or_else(|| { + error!( + function = %stage_name, + "model_id not set in request context" + ); + Box::new(error::internal_error( + "model_id_not_set", + "model_id not set in request context - this is a bug in request routing", + )) + })?; + + let tokenizer = ctx + .components + .tokenizer_registry + .get(model_id) + .ok_or_else(|| { + error!( + function = %stage_name, + model = %model_id, + "Tokenizer not found for model" + ); + Box::new(error::internal_error( + "tokenizer_not_found", + format!("Tokenizer not found for model: {}", model_id), + )) + })?; + + // Cache tokenizer in context for reuse in response processing stage + ctx.state.tokenizer = Some(tokenizer.clone()); + + Ok(tokenizer) +} + /// Get gRPC client from worker, returning appropriate error response on failure pub async fn get_grpc_client_from_worker(worker: &Arc) -> Result { // Get cached client from worker (or create one if not cached yet)