From 2f7c629235d90abb1416557829c974a0d883d84c Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Wed, 24 Dec 2025 11:28:53 -0500 Subject: [PATCH] [model-gateway] Fix IGW routing and optimize RouterManager (#15741) --- sgl-model-gateway/src/routers/factory.rs | 4 +- .../src/routers/router_manager.rs | 189 +++++++++++------- 2 files changed, 125 insertions(+), 68 deletions(-) diff --git a/sgl-model-gateway/src/routers/factory.rs b/sgl-model-gateway/src/routers/factory.rs index a4c85d5de..fea79f807 100644 --- a/sgl-model-gateway/src/routers/factory.rs +++ b/sgl-model-gateway/src/routers/factory.rs @@ -121,7 +121,9 @@ impl RouterFactory { /// Workers should be registered via the external worker registration workflow /// before using this router. The workflow discovers models from the provided /// endpoints and creates external workers in the registry. - async fn create_openai_router(ctx: &Arc) -> Result, String> { + pub async fn create_openai_router( + ctx: &Arc, + ) -> Result, String> { let router = OpenAIRouter::new(ctx).await?; Ok(Box::new(router)) } diff --git a/sgl-model-gateway/src/routers/router_manager.rs b/sgl-model-gateway/src/routers/router_manager.rs index 7b0588935..95bc2a6ad 100644 --- a/sgl-model-gateway/src/routers/router_manager.rs +++ b/sgl-model-gateway/src/routers/router_manager.rs @@ -36,18 +36,29 @@ use crate::{ }; #[derive(Debug, Clone, Hash, Eq, PartialEq)] -pub struct RouterId(String); +pub struct RouterId(&'static str); impl RouterId { - pub fn new(id: String) -> Self { + pub const fn new(id: &'static str) -> Self { Self(id) } pub fn as_str(&self) -> &str { - &self.0 + self.0 } } +/// Static router ID constants to avoid heap allocations in hot paths +pub mod router_ids { + use super::RouterId; + + pub const HTTP_REGULAR: RouterId = RouterId::new("http-regular"); + pub const HTTP_PD: RouterId = RouterId::new("http-pd"); + pub const HTTP_OPENAI: RouterId = RouterId::new("http-openai"); + pub const GRPC_REGULAR: RouterId = RouterId::new("grpc-regular"); + pub const GRPC_PD: RouterId = RouterId::new("grpc-pd"); +} + pub struct RouterManager { worker_registry: Arc, routers: Arc>>, @@ -83,10 +94,7 @@ impl RouterManager { match RouterFactory::create_regular_router(app_context).await { Ok(http_regular) => { info!("Created HTTP Regular router"); - manager.register_router( - RouterId::new("http-regular".to_string()), - Arc::from(http_regular), - ); + manager.register_router(router_ids::HTTP_REGULAR, Arc::from(http_regular)); } Err(e) => { warn!("Failed to create HTTP Regular router: {e}"); @@ -97,10 +105,7 @@ impl RouterManager { match RouterFactory::create_grpc_router(app_context).await { Ok(grpc_regular) => { info!("Created gRPC Regular router"); - manager.register_router( - RouterId::new("grpc-regular".to_string()), - Arc::from(grpc_regular), - ); + manager.register_router(router_ids::GRPC_REGULAR, Arc::from(grpc_regular)); } Err(e) => { warn!("Failed to create gRPC Regular router: {e}"); @@ -120,8 +125,7 @@ impl RouterManager { { Ok(http_pd) => { info!("Created HTTP PD router"); - manager - .register_router(RouterId::new("http-pd".to_string()), Arc::from(http_pd)); + manager.register_router(router_ids::HTTP_PD, Arc::from(http_pd)); } Err(e) => { warn!("Failed to create HTTP PD router: {e}"); @@ -139,14 +143,24 @@ impl RouterManager { { Ok(grpc_pd) => { info!("Created gRPC PD router"); - manager - .register_router(RouterId::new("grpc-pd".to_string()), Arc::from(grpc_pd)); + manager.register_router(router_ids::GRPC_PD, Arc::from(grpc_pd)); } Err(e) => { warn!("Failed to create gRPC PD router: {e}"); } } + // Create OpenAI router for external OpenAI-compatible backends + match RouterFactory::create_openai_router(app_context).await { + Ok(openai) => { + info!("Created OpenAI router"); + manager.register_router(router_ids::HTTP_OPENAI, Arc::from(openai)); + } + Err(e) => { + warn!("Failed to create OpenAI router: {e}"); + } + } + info!( "RouterManager initialized with {} routers for multi-router mode", manager.router_count(), @@ -177,24 +191,12 @@ impl RouterManager { connection_mode: &ConnectionMode, ) -> RouterId { match (connection_mode, routing_mode) { - (ConnectionMode::Http, RoutingMode::Regular { .. }) => { - RouterId::new("http-regular".to_string()) - } - (ConnectionMode::Http, RoutingMode::PrefillDecode { .. }) => { - RouterId::new("http-pd".to_string()) - } - (ConnectionMode::Http, RoutingMode::OpenAI { .. }) => { - RouterId::new("http-openai".to_string()) - } - (ConnectionMode::Grpc { .. }, RoutingMode::Regular { .. }) => { - RouterId::new("grpc-regular".to_string()) - } - (ConnectionMode::Grpc { .. }, RoutingMode::PrefillDecode { .. }) => { - RouterId::new("grpc-pd".to_string()) - } - (ConnectionMode::Grpc { .. }, RoutingMode::OpenAI { .. }) => { - RouterId::new("grpc-regular".to_string()) - } + (ConnectionMode::Http, RoutingMode::Regular { .. }) => router_ids::HTTP_REGULAR, + (ConnectionMode::Http, RoutingMode::PrefillDecode { .. }) => router_ids::HTTP_PD, + (ConnectionMode::Http, RoutingMode::OpenAI { .. }) => router_ids::HTTP_OPENAI, + (ConnectionMode::Grpc { .. }, RoutingMode::Regular { .. }) => router_ids::GRPC_REGULAR, + (ConnectionMode::Grpc { .. }, RoutingMode::PrefillDecode { .. }) => router_ids::GRPC_PD, + (ConnectionMode::Grpc { .. }, RoutingMode::OpenAI { .. }) => router_ids::GRPC_REGULAR, } } @@ -205,7 +207,10 @@ impl RouterManager { let new_snapshot: Vec<_> = self.routers.iter().map(|e| e.value().clone()).collect(); self.routers_snapshot.store(Arc::new(new_snapshot)); - let mut default_router = self.default_router.write().unwrap(); + let mut default_router = self + .default_router + .write() + .unwrap_or_else(|e| e.into_inner()); if default_router.is_none() { *default_router = Some(id.clone()); info!("Set default router to {}", id.as_str()); @@ -213,7 +218,10 @@ impl RouterManager { } pub fn set_default_router(&self, id: RouterId) { - let mut default_router = self.default_router.write().unwrap(); + let mut default_router = self + .default_router + .write() + .unwrap_or_else(|e| e.into_inner()); *default_router = Some(id); } @@ -221,7 +229,9 @@ impl RouterManager { self.routers.len() } - /// Resolve model_id for a request, inferring from available workers if not specified + /// Resolve model_id for a request, inferring from available workers if not specified. + /// + /// Behavior in IGW mode (must fail fast if model not resolvable): /// - If model_id is provided, use it directly /// - If not provided and only one model exists, use it as implicit default /// - If not provided and multiple models exist, return error requiring specification @@ -270,9 +280,9 @@ impl RouterManager { pub fn get_router_for_model(&self, model_id: &str) -> Option> { let workers = self.worker_registry.get_by_model(model_id); - // Find the best worker type and derive router ID from it - // Priority: grpc-pd (3) > http-pd (2) > grpc-regular (1) > http-regular (0) - let best_score = workers + // Find the best router ID based on worker capabilities + // Priority: grpc-pd > http-pd > grpc-regular > http-regular + let best_router_id = workers .iter() .map(|w| { let is_pd = matches!( @@ -282,29 +292,26 @@ impl RouterManager { let is_grpc = matches!(w.connection_mode(), ConnectionMode::Grpc { .. }); match (is_grpc, is_pd) { - (true, true) => 3, // grpc-pd (best) - (false, true) => 2, // http-pd - (true, false) => 1, // grpc-regular - (false, false) => 0, // http-regular + (true, true) => (3, &router_ids::GRPC_PD), + (false, true) => (2, &router_ids::HTTP_PD), + (true, false) => (1, &router_ids::GRPC_REGULAR), + (false, false) => (0, &router_ids::HTTP_REGULAR), } }) - .max(); + .max_by_key(|(score, _)| *score) + .map(|(_, id)| id); - if let Some(score) = best_score { - let router_id = match score { - 3 => "grpc-pd", - 2 => "http-pd", - 1 => "grpc-regular", - _ => "http-regular", - }; - - if let Some(router) = self.routers.get(&RouterId::new(router_id.to_string())) { + if let Some(router_id) = best_router_id { + if let Some(router) = self.routers.get(router_id) { return Some(router.clone()); } } // Fallback to default router - let default_router = self.default_router.read().unwrap(); + let default_router = self + .default_router + .read() + .unwrap_or_else(|e| e.into_inner()); if let Some(ref default_id) = *default_router { self.routers.get(default_id).map(|r| r.clone()) } else { @@ -319,7 +326,10 @@ impl RouterManager { ) -> Option> { // In single-router mode (enable_igw=false), always use the default router if !self.enable_igw { - let default_router = self.default_router.read().unwrap(); + let default_router = self + .default_router + .read() + .unwrap_or_else(|e| e.into_inner()); if let Some(ref default_id) = *default_router { debug!( "Single-router mode: using default router {} for model {:?}", @@ -454,7 +464,10 @@ impl RouterTrait for RouterManager { async fn get_model_info(&self, req: Request) -> Response { // Route to default router or first available router let router_id = { - let default_router = self.default_router.read().unwrap(); + let default_router = self + .default_router + .read() + .unwrap_or_else(|e| e.into_inner()); default_router.clone() }; @@ -478,17 +491,23 @@ impl RouterTrait for RouterManager { body: &GenerateRequest, model_id: Option<&str>, ) -> Response { - // Resolve model_id intelligently instead of falling back to "unknown" - let resolved_model_id = match self.resolve_model_id(model_id) { - Ok(id) => id, - Err(err_response) => return *err_response, + // In IGW mode, resolve model_id and fail fast if not resolvable + // In non-IGW mode, pass through to router (router handles validation) + let effective_model_id = if self.enable_igw { + match self.resolve_model_id(model_id) { + Ok(id) => Some(id), + Err(err_response) => return *err_response, + } + } else { + None }; - let router = self.select_router_for_request(headers, Some(&resolved_model_id)); + let router = + self.select_router_for_request(headers, effective_model_id.as_deref().or(model_id)); if let Some(router) = router { router - .route_generate(headers, body, Some(&resolved_model_id)) + .route_generate(headers, body, effective_model_id.as_deref().or(model_id)) .await } else { ( @@ -505,10 +524,26 @@ impl RouterTrait for RouterManager { body: &ChatCompletionRequest, model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, model_id); + // In IGW mode, resolve model_id and fail fast if not resolvable + // In non-IGW mode, pass through to router (router handles validation) + let effective_model_id = if self.enable_igw { + // Use provided model_id or fall back to body.model + let model = model_id.or(Some(&body.model)); + match self.resolve_model_id(model) { + Ok(id) => Some(id), + Err(err_response) => return *err_response, + } + } else { + None + }; + + let router = + self.select_router_for_request(headers, effective_model_id.as_deref().or(model_id)); if let Some(router) = router { - router.route_chat(headers, body, model_id).await + router + .route_chat(headers, body, effective_model_id.as_deref().or(model_id)) + .await } else { ( StatusCode::NOT_FOUND, @@ -524,10 +559,26 @@ impl RouterTrait for RouterManager { body: &CompletionRequest, model_id: Option<&str>, ) -> Response { - let router = self.select_router_for_request(headers, model_id); + // In IGW mode, resolve model_id and fail fast if not resolvable + // In non-IGW mode, pass through to router (router handles validation) + let effective_model_id = if self.enable_igw { + // Use provided model_id or fall back to body.model + let model = model_id.or(Some(&body.model)); + match self.resolve_model_id(model) { + Ok(id) => Some(id), + Err(err_response) => return *err_response, + } + } else { + None + }; + + let router = + self.select_router_for_request(headers, effective_model_id.as_deref().or(model_id)); if let Some(router) = router { - router.route_completion(headers, body, model_id).await + router + .route_completion(headers, body, effective_model_id.as_deref().or(model_id)) + .await } else { ( StatusCode::NOT_FOUND, @@ -679,10 +730,14 @@ impl RouterTrait for RouterManager { impl std::fmt::Debug for RouterManager { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let default_router = self + .default_router + .read() + .unwrap_or_else(|e| e.into_inner()); f.debug_struct("RouterManager") .field("routers_count", &self.routers.len()) .field("workers_count", &self.worker_registry.get_all().len()) - .field("default_router", &*self.default_router.read().unwrap()) + .field("default_router", &*default_router) .finish() } }