[model-gateway] Optimize router selection with lock-free snapshots (#15672)

This commit is contained in:
Praneth Paruchuri
2025-12-23 09:09:30 -08:00
committed by GitHub
parent 705287b2e5
commit 80ae2229d3
2 changed files with 41 additions and 51 deletions
+2
View File
@@ -104,6 +104,7 @@ ndarray = "0.16"
base64 = "0.22" base64 = "0.22"
openai-harmony = { git = "https://github.com/openai/harmony", tag = "v0.0.4" } openai-harmony = { git = "https://github.com/openai/harmony", tag = "v0.0.4" }
openmetrics-parser = "0.4.4" openmetrics-parser = "0.4.4"
arc-swap = "1.7.1"
# gRPC and Protobuf dependencies # gRPC and Protobuf dependencies
tonic = { version = "0.14.2", features = ["gzip", "transport"] } tonic = { version = "0.14.2", features = ["gzip", "transport"] }
@@ -144,6 +145,7 @@ opentelemetry-proto = { version = "0.27", features = ["gen-tonic"] }
tonic-v12 = { version = "0.12.3", package = "tonic" } tonic-v12 = { version = "0.12.3", package = "tonic" }
serial_test = "3.0" serial_test = "3.0"
[[bench]] [[bench]]
name = "request_processing" name = "request_processing"
harness = false harness = false
+28 -40
View File
@@ -6,6 +6,7 @@
use std::sync::Arc; use std::sync::Arc;
use arc_swap::ArcSwap;
use async_trait::async_trait; use async_trait::async_trait;
use axum::{ use axum::{
body::Body, body::Body,
@@ -50,6 +51,7 @@ impl RouterId {
pub struct RouterManager { pub struct RouterManager {
worker_registry: Arc<WorkerRegistry>, worker_registry: Arc<WorkerRegistry>,
routers: Arc<DashMap<RouterId, Arc<dyn RouterTrait>>>, routers: Arc<DashMap<RouterId, Arc<dyn RouterTrait>>>,
routers_snapshot: ArcSwap<Vec<Arc<dyn RouterTrait>>>,
default_router: Arc<std::sync::RwLock<Option<RouterId>>>, default_router: Arc<std::sync::RwLock<Option<RouterId>>>,
enable_igw: bool, enable_igw: bool,
} }
@@ -59,6 +61,7 @@ impl RouterManager {
Self { Self {
worker_registry, worker_registry,
routers: Arc::new(DashMap::new()), routers: Arc::new(DashMap::new()),
routers_snapshot: ArcSwap::from_pointee(Vec::new()),
default_router: Arc::new(std::sync::RwLock::new(None)), default_router: Arc::new(std::sync::RwLock::new(None)),
enable_igw: false, // Will be set properly in from_config enable_igw: false, // Will be set properly in from_config
} }
@@ -164,6 +167,10 @@ impl RouterManager {
pub fn register_router(&self, id: RouterId, router: Arc<dyn RouterTrait>) { pub fn register_router(&self, id: RouterId, router: Arc<dyn RouterTrait>) {
self.routers.insert(id.clone(), router); self.routers.insert(id.clone(), router);
// Update the lock-free snapshot for fast per-request iteration
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();
if default_router.is_none() { if default_router.is_none() {
*default_router = Some(id.clone()); *default_router = Some(id.clone());
@@ -274,19 +281,6 @@ impl RouterManager {
} }
} }
// Multi-router mode logic follows
let _priority_threshold = headers.and_then(|h| {
h.get("x-worker-priority")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u32>().ok())
});
let _max_cost = headers.and_then(|h| {
h.get("x-max-cost")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<f32>().ok())
});
let prefer_pd = headers let prefer_pd = headers
.and_then(|h| { .and_then(|h| {
h.get("x-prefer-pd") h.get("x-prefer-pd")
@@ -295,30 +289,26 @@ impl RouterManager {
}) })
.unwrap_or(false); .unwrap_or(false);
let candidate_routers = if let Some(model) = model_id {
if let Some(router) = self.get_router_for_model(model) {
vec![router]
} else {
Vec::new()
}
} else {
self.routers
.iter()
.map(|entry| entry.value().clone())
.collect::<Vec<_>>()
};
if candidate_routers.is_empty() {
return None;
}
let mut best_router = None;
let mut best_score = 0.0;
// Uses O(1) lookups instead of allocating a full vector of workers via get_all()
let (num_regular_workers, num_pd_workers) = self.worker_registry.get_worker_distribution(); let (num_regular_workers, num_pd_workers) = self.worker_registry.get_worker_distribution();
let mut best_router = None;
let mut best_score = -1.0;
for router in candidate_routers { // Extract router validity check into a closure to reduce redundancy
let is_router_valid =
|is_pd: bool| (is_pd && num_pd_workers > 0) || (!is_pd && num_regular_workers > 0);
if let Some(model) = model_id {
// Efficient Single Lookup for Specific Model
if let Some(router) = self.get_router_for_model(model) {
if is_router_valid(router.is_pd_mode()) {
return Some(router);
}
}
} else {
// ZERO-ALLOCATION Snapshot Iteration (Hot Path Optimization)
// Atomic load avoids heap allocations and DashMap shard locks per-request
let routers_snapshot = self.routers_snapshot.load();
for router in routers_snapshot.iter() {
let mut score = 1.0; let mut score = 1.0;
let is_pd = router.is_pd_mode(); let is_pd = router.is_pd_mode();
@@ -327,17 +317,15 @@ impl RouterManager {
} else if !prefer_pd && !is_pd { } else if !prefer_pd && !is_pd {
score += 1.0; score += 1.0;
} }
// TODO: Once routers expose worker stats, we can evaluate: // TODO: Once routers expose worker stats, we can evaluate:
// - Average worker priority vs priority_threshold // - Average worker priority vs priority_threshold
// - Average worker cost vs max_cost // - Average worker cost vs max_cost
// - Current load and health status // - Current load and health status
let valid_router = (router.is_pd_mode() && num_pd_workers > 0) if score > best_score && is_router_valid(is_pd) {
|| (!router.is_pd_mode() && num_regular_workers > 0);
if score > best_score && valid_router {
best_score = score; best_score = score;
best_router = Some(router); best_router = Some(Arc::clone(router));
}
} }
} }