[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
+96 -106
View File
@@ -1,8 +1,6 @@
use std::{
collections::HashMap,
sync::{Arc, RwLock},
};
use std::sync::{Arc, OnceLock};
use dashmap::DashMap;
use tracing::{debug, info, warn};
/// 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
#[derive(Clone)]
pub struct PolicyRegistry {
/// Model ID -> Policy instance mapping
model_policies: Arc<RwLock<HashMap<String, Arc<dyn LoadBalancingPolicy>>>>,
/// Model ID -> Policy instance mapping (lock-free reads via DashMap)
model_policies: Arc<DashMap<String, Arc<dyn LoadBalancingPolicy>>>,
/// Model ID -> Worker count for cleanup tracking
model_worker_counts: Arc<RwLock<HashMap<String, usize>>>,
/// Model ID -> Worker count for cleanup tracking (lock-free reads via DashMap)
model_worker_counts: Arc<DashMap<String, usize>>,
/// Default policy instance (cached)
/// Default policy instance (cached, immutable after creation)
default_policy: Arc<dyn LoadBalancingPolicy>,
/// Prefill policy for PD mode
prefill_policy: Arc<RwLock<Option<Arc<dyn LoadBalancingPolicy>>>>,
/// Prefill policy for PD mode (set once at startup, lock-free reads via OnceLock)
prefill_policy: Arc<OnceLock<Arc<dyn LoadBalancingPolicy>>>,
/// Decode policy for PD mode
decode_policy: Arc<RwLock<Option<Arc<dyn LoadBalancingPolicy>>>>,
/// Decode policy for PD mode (set once at startup, lock-free reads via OnceLock)
decode_policy: Arc<OnceLock<Arc<dyn LoadBalancingPolicy>>>,
}
impl PolicyRegistry {
@@ -42,11 +40,11 @@ impl PolicyRegistry {
let default_policy = Self::create_policy_from_config(&default_policy_config);
Self {
model_policies: Arc::new(RwLock::new(HashMap::new())),
model_worker_counts: Arc::new(RwLock::new(HashMap::new())),
model_policies: Arc::new(DashMap::new()),
model_worker_counts: Arc::new(DashMap::new()),
default_policy,
prefill_policy: Arc::new(RwLock::new(None)),
decode_policy: Arc::new(RwLock::new(None)),
prefill_policy: Arc::new(OnceLock::new()),
decode_policy: Arc::new(OnceLock::new()),
}
}
@@ -57,28 +55,23 @@ impl PolicyRegistry {
model_id: &str,
policy_hint: Option<&str>,
) -> Arc<dyn LoadBalancingPolicy> {
// Increment worker count
{
let mut counts = self.model_worker_counts.write().unwrap();
*counts.entry(model_id.to_string()).or_insert(0) += 1;
debug!(
"Worker added for model {}, count: {}",
model_id,
counts.get(model_id).unwrap()
);
}
// Increment worker count using DashMap entry API
let count = self
.model_worker_counts
.entry(model_id.to_string())
.and_modify(|c| *c += 1)
.or_insert(1);
debug!("Worker added for model {}, count: {}", model_id, *count);
drop(count); // Release the entry lock
// Check if model already has a policy
{
let policies = self.model_policies.read().unwrap();
if let Some(existing_policy) = policies.get(model_id) {
debug!(
"Model {} already has policy: {}",
model_id,
existing_policy.name()
);
return Arc::clone(existing_policy);
}
// Check if model already has a policy (lock-free read via DashMap)
if let Some(existing_policy) = self.model_policies.get(model_id) {
debug!(
"Model {} already has policy: {}",
model_id,
existing_policy.name()
);
return Arc::clone(&existing_policy);
}
// New model - determine policy
@@ -90,55 +83,53 @@ impl PolicyRegistry {
model_id
);
// Store policy for this model
{
let mut policies = self.model_policies.write().unwrap();
policies.insert(model_id.to_string(), Arc::clone(&policy));
}
// Store policy for this model (DashMap handles concurrent inserts)
self.model_policies
.insert(model_id.to_string(), Arc::clone(&policy));
policy
}
/// Called when a worker is removed
pub fn on_worker_removed(&self, model_id: &str) {
let should_cleanup = {
let mut counts = self.model_worker_counts.write().unwrap();
if let Some(count) = counts.get_mut(model_id) {
*count = count.saturating_sub(1);
debug!("Worker removed for model {}, count: {}", model_id, *count);
if *count == 0 {
counts.remove(model_id);
true
} else {
false
}
// Decrement worker count and check if cleanup needed
let should_cleanup = if let Some(mut count_ref) = self.model_worker_counts.get_mut(model_id)
{
*count_ref = count_ref.saturating_sub(1);
debug!(
"Worker removed for model {}, count: {}",
model_id, *count_ref
);
if *count_ref == 0 {
drop(count_ref); // Release before remove
self.model_worker_counts.remove(model_id);
true
} else {
warn!(
"Attempted to remove worker for model {} with no registered workers",
model_id
);
false
}
} else {
warn!(
"Attempted to remove worker for model {} with no registered workers",
model_id
);
false
};
// Clean up policy if this was the last worker
if should_cleanup {
let mut policies = self.model_policies.write().unwrap();
if let Some(policy) = policies.remove(model_id) {
if let Some((_, policy)) = self.model_policies.remove(model_id) {
info!(
"Removed policy {} for model {} (last worker removed)",
policy.name(),
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>> {
self.model_policies.read().unwrap().get(model_id).cloned()
self.model_policies.get(model_id).map(|r| Arc::clone(&r))
}
/// Get the default policy
@@ -222,58 +213,58 @@ impl PolicyRegistry {
}
/// Get current model->policy mappings (for debugging/monitoring)
pub fn get_all_mappings(&self) -> HashMap<String, String> {
let policies = self.model_policies.read().unwrap();
policies
pub fn get_all_mappings(&self) -> std::collections::HashMap<String, String> {
self.model_policies
.iter()
.map(|(model, policy)| (model.clone(), policy.name().to_string()))
.map(|entry| (entry.key().clone(), entry.value().name().to_string()))
.collect()
}
/// Get worker counts per model
pub fn get_worker_counts(&self) -> HashMap<String, usize> {
self.model_worker_counts.read().unwrap().clone()
pub fn get_worker_counts(&self) -> std::collections::HashMap<String, usize> {
self.model_worker_counts
.iter()
.map(|entry| (entry.key().clone(), *entry.value()))
.collect()
}
/// Clear all policies (useful for testing)
pub fn clear(&self) {
let mut policies = self.model_policies.write().unwrap();
policies.clear();
let mut counts = self.model_worker_counts.write().unwrap();
counts.clear();
self.model_policies.clear();
self.model_worker_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>) {
let mut prefill_policy = self.prefill_policy.write().unwrap();
*prefill_policy = Some(policy);
// OnceLock::set returns Err if already set, which we ignore since
// 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>) {
let mut decode_policy = self.decode_policy.write().unwrap();
*decode_policy = Some(policy);
// OnceLock::set returns Err if already set, which we ignore since
// 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> {
let prefill_policy = self.prefill_policy.read().unwrap();
prefill_policy
.as_ref()
self.prefill_policy
.get()
.map(Arc::clone)
.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> {
let decode_policy = self.decode_policy.read().unwrap();
decode_policy
.as_ref()
self.decode_policy
.get()
.map(Arc::clone)
.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>> {
let mut power_of_two_policies = Vec::new();
@@ -281,29 +272,27 @@ impl PolicyRegistry {
power_of_two_policies.push(Arc::clone(&self.default_policy));
}
// Cache prefill and decode policies to avoid double-locking prefill_policy
let prefill_policy_opt = self.prefill_policy.read().unwrap().clone();
let decode_policy_opt = self.decode_policy.read().unwrap().clone();
// Get prefill and decode policies (lock-free via OnceLock::get)
let prefill_policy_opt = self.prefill_policy.get();
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) {
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"
&& !Arc::ptr_eq(policy, &self.default_policy)
&& !prefill_policy_opt
.as_ref()
.is_some_and(|p| Arc::ptr_eq(p, policy))
&& !prefill_policy_opt.is_some_and(|p| Arc::ptr_eq(p, policy))
{
power_of_two_policies.push(Arc::clone(policy));
}
}
let model_policies = self.model_policies.read().unwrap();
for policy in model_policies.values() {
for entry in self.model_policies.iter() {
let policy = entry.value();
if policy.name() == "power_of_two" {
let already_added = power_of_two_policies.iter().any(|p| Arc::ptr_eq(p, policy));
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(
&self,
prefill_workers: &[Arc<dyn Worker>],
decode_workers: &[Arc<dyn Worker>],
) {
// Initialize prefill policy if it's cache-aware
if let Some(prefill_policy) = self.prefill_policy.read().unwrap().as_ref() {
// Initialize prefill policy if it's cache-aware (lock-free via OnceLock::get)
if let Some(prefill_policy) = self.prefill_policy.get() {
if prefill_policy.name() == "cache_aware" {
if let Some(cache_aware) =
prefill_policy.as_any().downcast_ref::<CacheAwarePolicy>()
@@ -373,8 +362,8 @@ impl PolicyRegistry {
}
}
// Initialize decode policy if it's cache-aware
if let Some(decode_policy) = self.decode_policy.read().unwrap().as_ref() {
// Initialize decode policy if it's cache-aware (lock-free via OnceLock::get)
if let Some(decode_policy) = self.decode_policy.get() {
if decode_policy.name() == "cache_aware" {
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>]) {
// Initialize prefill policy if it's bucket
if let Some(prefill_policy) = self.prefill_policy.read().unwrap().as_ref() {
// Initialize prefill policy if it's bucket (lock-free via OnceLock::get)
if let Some(prefill_policy) = self.prefill_policy.get() {
if prefill_policy.name() == "bucket" {
if let Some(bucket) = prefill_policy.as_any().downcast_ref::<BucketPolicy>() {
if !prefill_workers.is_empty() {