[model-gateway] Fix IGW routing and optimize RouterManager (#15741)
This commit is contained in:
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user