[model-gateway] Fix IGW routing and optimize RouterManager (#15741)

This commit is contained in:
Simo Lin
2025-12-24 11:28:53 -05:00
committed by GitHub
parent 2fb3160598
commit 2f7c629235
2 changed files with 125 additions and 68 deletions
+3 -1
View File
@@ -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<AppContext>) -> Result<Box<dyn RouterTrait>, String> {
pub async fn create_openai_router(
ctx: &Arc<AppContext>,
) -> Result<Box<dyn RouterTrait>, String> {
let router = OpenAIRouter::new(ctx).await?;
Ok(Box::new(router))
}
+122 -67
View File
@@ -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<WorkerRegistry>,
routers: Arc<DashMap<RouterId, Arc<dyn RouterTrait>>>,
@@ -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<Arc<dyn RouterTrait>> {
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<Arc<dyn RouterTrait>> {
// 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<Body>) -> 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()
}
}