[model-gateway] Optimize router selection with lock-free snapshots (#15672)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user