[model-gateway] Replace PolicyRegistry RwLock with DashMap for lock-free policy lookups (#15361)

This commit is contained in:
Simo Lin
2025-12-18 00:11:20 -08:00
committed by GitHub
parent ee1ca51de8
commit 70607e55e8
+84 -94
View File
@@ -1,8 +1,6 @@
use std::{ use std::sync::{Arc, OnceLock};
collections::HashMap,
sync::{Arc, RwLock},
};
use dashmap::DashMap;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
/// Policy Registry for managing model-to-policy mappings /// Policy Registry for managing model-to-policy mappings
@@ -20,20 +18,20 @@ use crate::{config::types::PolicyConfig, core::Worker};
/// Registry for managing model-to-policy mappings /// Registry for managing model-to-policy mappings
#[derive(Clone)] #[derive(Clone)]
pub struct PolicyRegistry { pub struct PolicyRegistry {
/// Model ID -> Policy instance mapping /// Model ID -> Policy instance mapping (lock-free reads via DashMap)
model_policies: Arc<RwLock<HashMap<String, Arc<dyn LoadBalancingPolicy>>>>, model_policies: Arc<DashMap<String, Arc<dyn LoadBalancingPolicy>>>,
/// Model ID -> Worker count for cleanup tracking /// Model ID -> Worker count for cleanup tracking (lock-free reads via DashMap)
model_worker_counts: Arc<RwLock<HashMap<String, usize>>>, model_worker_counts: Arc<DashMap<String, usize>>,
/// Default policy instance (cached) /// Default policy instance (cached, immutable after creation)
default_policy: Arc<dyn LoadBalancingPolicy>, default_policy: Arc<dyn LoadBalancingPolicy>,
/// Prefill policy for PD mode /// Prefill policy for PD mode (set once at startup, lock-free reads via OnceLock)
prefill_policy: Arc<RwLock<Option<Arc<dyn LoadBalancingPolicy>>>>, prefill_policy: Arc<OnceLock<Arc<dyn LoadBalancingPolicy>>>,
/// Decode policy for PD mode /// Decode policy for PD mode (set once at startup, lock-free reads via OnceLock)
decode_policy: Arc<RwLock<Option<Arc<dyn LoadBalancingPolicy>>>>, decode_policy: Arc<OnceLock<Arc<dyn LoadBalancingPolicy>>>,
} }
impl PolicyRegistry { impl PolicyRegistry {
@@ -42,11 +40,11 @@ impl PolicyRegistry {
let default_policy = Self::create_policy_from_config(&default_policy_config); let default_policy = Self::create_policy_from_config(&default_policy_config);
Self { Self {
model_policies: Arc::new(RwLock::new(HashMap::new())), model_policies: Arc::new(DashMap::new()),
model_worker_counts: Arc::new(RwLock::new(HashMap::new())), model_worker_counts: Arc::new(DashMap::new()),
default_policy, default_policy,
prefill_policy: Arc::new(RwLock::new(None)), prefill_policy: Arc::new(OnceLock::new()),
decode_policy: Arc::new(RwLock::new(None)), decode_policy: Arc::new(OnceLock::new()),
} }
} }
@@ -57,28 +55,23 @@ impl PolicyRegistry {
model_id: &str, model_id: &str,
policy_hint: Option<&str>, policy_hint: Option<&str>,
) -> Arc<dyn LoadBalancingPolicy> { ) -> Arc<dyn LoadBalancingPolicy> {
// Increment worker count // Increment worker count using DashMap entry API
{ let count = self
let mut counts = self.model_worker_counts.write().unwrap(); .model_worker_counts
*counts.entry(model_id.to_string()).or_insert(0) += 1; .entry(model_id.to_string())
debug!( .and_modify(|c| *c += 1)
"Worker added for model {}, count: {}", .or_insert(1);
model_id, debug!("Worker added for model {}, count: {}", model_id, *count);
counts.get(model_id).unwrap() drop(count); // Release the entry lock
);
}
// Check if model already has a policy // Check if model already has a policy (lock-free read via DashMap)
{ if let Some(existing_policy) = self.model_policies.get(model_id) {
let policies = self.model_policies.read().unwrap();
if let Some(existing_policy) = policies.get(model_id) {
debug!( debug!(
"Model {} already has policy: {}", "Model {} already has policy: {}",
model_id, model_id,
existing_policy.name() existing_policy.name()
); );
return Arc::clone(existing_policy); return Arc::clone(&existing_policy);
}
} }
// New model - determine policy // New model - determine policy
@@ -90,24 +83,26 @@ impl PolicyRegistry {
model_id model_id
); );
// Store policy for this model // Store policy for this model (DashMap handles concurrent inserts)
{ self.model_policies
let mut policies = self.model_policies.write().unwrap(); .insert(model_id.to_string(), Arc::clone(&policy));
policies.insert(model_id.to_string(), Arc::clone(&policy));
}
policy policy
} }
/// Called when a worker is removed /// Called when a worker is removed
pub fn on_worker_removed(&self, model_id: &str) { pub fn on_worker_removed(&self, model_id: &str) {
let should_cleanup = { // Decrement worker count and check if cleanup needed
let mut counts = self.model_worker_counts.write().unwrap(); let should_cleanup = if let Some(mut count_ref) = self.model_worker_counts.get_mut(model_id)
if let Some(count) = counts.get_mut(model_id) { {
*count = count.saturating_sub(1); *count_ref = count_ref.saturating_sub(1);
debug!("Worker removed for model {}, count: {}", model_id, *count); debug!(
if *count == 0 { "Worker removed for model {}, count: {}",
counts.remove(model_id); model_id, *count_ref
);
if *count_ref == 0 {
drop(count_ref); // Release before remove
self.model_worker_counts.remove(model_id);
true true
} else { } else {
false false
@@ -118,27 +113,23 @@ impl PolicyRegistry {
model_id model_id
); );
false false
}
}; };
// Clean up policy if this was the last worker // Clean up policy if this was the last worker
if should_cleanup { if should_cleanup {
let mut policies = self.model_policies.write().unwrap(); if let Some((_, policy)) = self.model_policies.remove(model_id) {
if let Some(policy) = policies.remove(model_id) {
info!( info!(
"Removed policy {} for model {} (last worker removed)", "Removed policy {} for model {} (last worker removed)",
policy.name(), policy.name(),
model_id model_id
); );
// Policy will be dropped here, cleaning up any resources
drop(policy);
} }
} }
} }
/// Get the policy for a model /// Get the policy for a model (lock-free via DashMap)
pub fn get_policy(&self, model_id: &str) -> Option<Arc<dyn LoadBalancingPolicy>> { pub fn get_policy(&self, model_id: &str) -> Option<Arc<dyn LoadBalancingPolicy>> {
self.model_policies.read().unwrap().get(model_id).cloned() self.model_policies.get(model_id).map(|r| Arc::clone(&r))
} }
/// Get the default policy /// Get the default policy
@@ -222,58 +213,58 @@ impl PolicyRegistry {
} }
/// Get current model->policy mappings (for debugging/monitoring) /// Get current model->policy mappings (for debugging/monitoring)
pub fn get_all_mappings(&self) -> HashMap<String, String> { pub fn get_all_mappings(&self) -> std::collections::HashMap<String, String> {
let policies = self.model_policies.read().unwrap(); self.model_policies
policies
.iter() .iter()
.map(|(model, policy)| (model.clone(), policy.name().to_string())) .map(|entry| (entry.key().clone(), entry.value().name().to_string()))
.collect() .collect()
} }
/// Get worker counts per model /// Get worker counts per model
pub fn get_worker_counts(&self) -> HashMap<String, usize> { pub fn get_worker_counts(&self) -> std::collections::HashMap<String, usize> {
self.model_worker_counts.read().unwrap().clone() self.model_worker_counts
.iter()
.map(|entry| (entry.key().clone(), *entry.value()))
.collect()
} }
/// Clear all policies (useful for testing) /// Clear all policies (useful for testing)
pub fn clear(&self) { pub fn clear(&self) {
let mut policies = self.model_policies.write().unwrap(); self.model_policies.clear();
policies.clear(); self.model_worker_counts.clear();
let mut counts = self.model_worker_counts.write().unwrap();
counts.clear();
} }
/// Set the prefill policy for PD mode /// Set the prefill policy for PD mode (lock-free, set once at startup)
pub fn set_prefill_policy(&self, policy: Arc<dyn LoadBalancingPolicy>) { pub fn set_prefill_policy(&self, policy: Arc<dyn LoadBalancingPolicy>) {
let mut prefill_policy = self.prefill_policy.write().unwrap(); // OnceLock::set returns Err if already set, which we ignore since
*prefill_policy = Some(policy); // the policy should only be set once at startup
let _ = self.prefill_policy.set(policy);
} }
/// Set the decode policy for PD mode /// Set the decode policy for PD mode (lock-free, set once at startup)
pub fn set_decode_policy(&self, policy: Arc<dyn LoadBalancingPolicy>) { pub fn set_decode_policy(&self, policy: Arc<dyn LoadBalancingPolicy>) {
let mut decode_policy = self.decode_policy.write().unwrap(); // OnceLock::set returns Err if already set, which we ignore since
*decode_policy = Some(policy); // the policy should only be set once at startup
let _ = self.decode_policy.set(policy);
} }
/// Get the prefill policy for PD mode, or default if not set /// Get the prefill policy for PD mode, or default if not set (lock-free)
pub fn get_prefill_policy(&self) -> Arc<dyn LoadBalancingPolicy> { pub fn get_prefill_policy(&self) -> Arc<dyn LoadBalancingPolicy> {
let prefill_policy = self.prefill_policy.read().unwrap(); self.prefill_policy
prefill_policy .get()
.as_ref()
.map(Arc::clone) .map(Arc::clone)
.unwrap_or_else(|| self.get_default_policy()) .unwrap_or_else(|| self.get_default_policy())
} }
/// Get the decode policy for PD mode, or default if not set /// Get the decode policy for PD mode, or default if not set (lock-free)
pub fn get_decode_policy(&self) -> Arc<dyn LoadBalancingPolicy> { pub fn get_decode_policy(&self) -> Arc<dyn LoadBalancingPolicy> {
let decode_policy = self.decode_policy.read().unwrap(); self.decode_policy
decode_policy .get()
.as_ref()
.map(Arc::clone) .map(Arc::clone)
.unwrap_or_else(|| self.get_default_policy()) .unwrap_or_else(|| self.get_default_policy())
} }
/// Get all PowerOfTwo policies that need load updates /// Get all PowerOfTwo policies that need load updates (lock-free)
pub fn get_all_power_of_two_policies(&self) -> Vec<Arc<dyn LoadBalancingPolicy>> { pub fn get_all_power_of_two_policies(&self) -> Vec<Arc<dyn LoadBalancingPolicy>> {
let mut power_of_two_policies = Vec::new(); let mut power_of_two_policies = Vec::new();
@@ -281,29 +272,27 @@ impl PolicyRegistry {
power_of_two_policies.push(Arc::clone(&self.default_policy)); power_of_two_policies.push(Arc::clone(&self.default_policy));
} }
// Cache prefill and decode policies to avoid double-locking prefill_policy // Get prefill and decode policies (lock-free via OnceLock::get)
let prefill_policy_opt = self.prefill_policy.read().unwrap().clone(); let prefill_policy_opt = self.prefill_policy.get();
let decode_policy_opt = self.decode_policy.read().unwrap().clone(); let decode_policy_opt = self.decode_policy.get();
if let Some(ref policy) = prefill_policy_opt { if let Some(policy) = prefill_policy_opt {
if policy.name() == "power_of_two" && !Arc::ptr_eq(policy, &self.default_policy) { if policy.name() == "power_of_two" && !Arc::ptr_eq(policy, &self.default_policy) {
power_of_two_policies.push(Arc::clone(policy)); power_of_two_policies.push(Arc::clone(policy));
} }
} }
if let Some(ref policy) = decode_policy_opt { if let Some(policy) = decode_policy_opt {
if policy.name() == "power_of_two" if policy.name() == "power_of_two"
&& !Arc::ptr_eq(policy, &self.default_policy) && !Arc::ptr_eq(policy, &self.default_policy)
&& !prefill_policy_opt && !prefill_policy_opt.is_some_and(|p| Arc::ptr_eq(p, policy))
.as_ref()
.is_some_and(|p| Arc::ptr_eq(p, policy))
{ {
power_of_two_policies.push(Arc::clone(policy)); power_of_two_policies.push(Arc::clone(policy));
} }
} }
let model_policies = self.model_policies.read().unwrap(); for entry in self.model_policies.iter() {
for policy in model_policies.values() { let policy = entry.value();
if policy.name() == "power_of_two" { if policy.name() == "power_of_two" {
let already_added = power_of_two_policies.iter().any(|p| Arc::ptr_eq(p, policy)); let already_added = power_of_two_policies.iter().any(|p| Arc::ptr_eq(p, policy));
if !already_added { if !already_added {
@@ -350,14 +339,14 @@ impl PolicyRegistry {
} }
} }
/// Initialize cache-aware policies for PD mode (prefill and decode) /// Initialize cache-aware policies for PD mode (prefill and decode) - lock-free
pub fn init_pd_cache_aware_policies( pub fn init_pd_cache_aware_policies(
&self, &self,
prefill_workers: &[Arc<dyn Worker>], prefill_workers: &[Arc<dyn Worker>],
decode_workers: &[Arc<dyn Worker>], decode_workers: &[Arc<dyn Worker>],
) { ) {
// Initialize prefill policy if it's cache-aware // Initialize prefill policy if it's cache-aware (lock-free via OnceLock::get)
if let Some(prefill_policy) = self.prefill_policy.read().unwrap().as_ref() { if let Some(prefill_policy) = self.prefill_policy.get() {
if prefill_policy.name() == "cache_aware" { if prefill_policy.name() == "cache_aware" {
if let Some(cache_aware) = if let Some(cache_aware) =
prefill_policy.as_any().downcast_ref::<CacheAwarePolicy>() prefill_policy.as_any().downcast_ref::<CacheAwarePolicy>()
@@ -373,8 +362,8 @@ impl PolicyRegistry {
} }
} }
// Initialize decode policy if it's cache-aware // Initialize decode policy if it's cache-aware (lock-free via OnceLock::get)
if let Some(decode_policy) = self.decode_policy.read().unwrap().as_ref() { if let Some(decode_policy) = self.decode_policy.get() {
if decode_policy.name() == "cache_aware" { if decode_policy.name() == "cache_aware" {
if let Some(cache_aware) = decode_policy.as_any().downcast_ref::<CacheAwarePolicy>() if let Some(cache_aware) = decode_policy.as_any().downcast_ref::<CacheAwarePolicy>()
{ {
@@ -390,9 +379,10 @@ impl PolicyRegistry {
} }
} }
/// Initialize bucket policies for PD mode - lock-free
pub fn init_pd_bucket_policies(&self, prefill_workers: &[Arc<dyn Worker>]) { pub fn init_pd_bucket_policies(&self, prefill_workers: &[Arc<dyn Worker>]) {
// Initialize prefill policy if it's bucket // Initialize prefill policy if it's bucket (lock-free via OnceLock::get)
if let Some(prefill_policy) = self.prefill_policy.read().unwrap().as_ref() { if let Some(prefill_policy) = self.prefill_policy.get() {
if prefill_policy.name() == "bucket" { if prefill_policy.name() == "bucket" {
if let Some(bucket) = prefill_policy.as_any().downcast_ref::<BucketPolicy>() { if let Some(bucket) = prefill_policy.as_any().downcast_ref::<BucketPolicy>() {
if !prefill_workers.is_empty() { if !prefill_workers.is_empty() {