refactor(core): remove get_by_model_fast alias in worker_registry (#16313)

This commit is contained in:
Chang Su
2026-01-02 16:04:23 -08:00
committed by GitHub
parent e93433892b
commit c4edcac6d7
6 changed files with 11 additions and 20 deletions
@@ -35,7 +35,7 @@ impl StepExecutor for UpdatePoliciesForWorkerStep {
); );
for model_id in &affected_models { for model_id in &affected_models {
let workers = app_context.worker_registry.get_by_model_fast(model_id); let workers = app_context.worker_registry.get_by_model(model_id);
if let Some(policy) = app_context.policy_registry.get_policy(model_id) { if let Some(policy) = app_context.policy_registry.get_policy(model_id) {
if policy.name() == "cache_aware" && !workers.is_empty() { if policy.name() == "cache_aware" && !workers.is_empty() {
@@ -29,7 +29,7 @@ impl StepExecutor for UpdateRemainingPoliciesStep {
); );
for model_id in affected_models.iter() { for model_id in affected_models.iter() {
let remaining_workers = app_context.worker_registry.get_by_model_fast(model_id); let remaining_workers = app_context.worker_registry.get_by_model(model_id);
if let Some(policy) = app_context.policy_registry.get_policy(model_id) { if let Some(policy) = app_context.policy_registry.get_policy(model_id) {
if policy.name() == "cache_aware" && !remaining_workers.is_empty() { if policy.name() == "cache_aware" && !remaining_workers.is_empty() {
@@ -38,7 +38,7 @@ impl StepExecutor for UpdatePoliciesStep {
.on_worker_added(&model_id, policy_hint); .on_worker_added(&model_id, policy_hint);
// Initialize cache-aware policy if configured // Initialize cache-aware policy if configured
let all_workers = app_context.worker_registry.get_by_model_fast(&model_id); let all_workers = app_context.worker_registry.get_by_model(&model_id);
if let Some(policy) = app_context.policy_registry.get_policy(&model_id) { if let Some(policy) = app_context.policy_registry.get_policy(&model_id) {
if policy.name() == "cache_aware" { if policy.name() == "cache_aware" {
app_context app_context
+5 -14
View File
@@ -356,12 +356,6 @@ impl WorkerRegistry {
.unwrap_or_else(|| Arc::from(Self::EMPTY_WORKERS)) .unwrap_or_else(|| Arc::from(Self::EMPTY_WORKERS))
} }
/// Alias for get_by_model for backwards compatibility
#[inline]
pub fn get_by_model_fast(&self, model_id: &str) -> Arc<[Arc<dyn Worker>]> {
self.get_by_model(model_id)
}
/// Get all workers by worker type /// Get all workers by worker type
pub fn get_by_type(&self, worker_type: &WorkerType) -> Vec<Arc<dyn Worker>> { pub fn get_by_type(&self, worker_type: &WorkerType) -> Vec<Arc<dyn Worker>> {
self.type_workers self.type_workers
@@ -471,7 +465,7 @@ impl WorkerRegistry {
// Start with the most efficient collection based on filters // Start with the most efficient collection based on filters
// Use model index when possible as it's O(1) lookup // Use model index when possible as it's O(1) lookup
let workers: Vec<Arc<dyn Worker>> = if let Some(model) = model_id { let workers: Vec<Arc<dyn Worker>> = if let Some(model) = model_id {
self.get_by_model_fast(model).to_vec() self.get_by_model(model).to_vec()
} else { } else {
self.get_all() self.get_all()
}; };
@@ -764,24 +758,21 @@ mod tests {
registry.register(Arc::from(worker2)); registry.register(Arc::from(worker2));
registry.register(Arc::from(worker3)); registry.register(Arc::from(worker3));
let llama_workers = registry.get_by_model_fast("llama-3"); let llama_workers = registry.get_by_model("llama-3");
assert_eq!(llama_workers.len(), 2); assert_eq!(llama_workers.len(), 2);
let urls: Vec<String> = llama_workers.iter().map(|w| w.url().to_string()).collect(); let urls: Vec<String> = llama_workers.iter().map(|w| w.url().to_string()).collect();
assert!(urls.contains(&"http://worker1:8080".to_string())); assert!(urls.contains(&"http://worker1:8080".to_string()));
assert!(urls.contains(&"http://worker2:8080".to_string())); assert!(urls.contains(&"http://worker2:8080".to_string()));
let gpt_workers = registry.get_by_model_fast("gpt-4"); let gpt_workers = registry.get_by_model("gpt-4");
assert_eq!(gpt_workers.len(), 1); assert_eq!(gpt_workers.len(), 1);
assert_eq!(gpt_workers[0].url(), "http://worker3:8080"); assert_eq!(gpt_workers[0].url(), "http://worker3:8080");
let unknown_workers = registry.get_by_model_fast("unknown-model"); let unknown_workers = registry.get_by_model("unknown-model");
assert_eq!(unknown_workers.len(), 0); assert_eq!(unknown_workers.len(), 0);
let llama_workers_slow = registry.get_by_model("llama-3");
assert_eq!(llama_workers.len(), llama_workers_slow.len());
registry.remove_by_url("http://worker1:8080"); registry.remove_by_url("http://worker1:8080");
let llama_workers_after = registry.get_by_model_fast("llama-3"); let llama_workers_after = registry.get_by_model("llama-3");
assert_eq!(llama_workers_after.len(), 1); assert_eq!(llama_workers_after.len(), 1);
assert_eq!(llama_workers_after[0].url(), "http://worker2:8080"); assert_eq!(llama_workers_after[0].url(), "http://worker2:8080");
} }
@@ -62,7 +62,7 @@ impl HarmonyDetector {
/// the model (e.g., during startup before workers are discovered). /// the model (e.g., during startup before workers are discovered).
pub fn is_harmony_model_in_registry(registry: &WorkerRegistry, model_name: &str) -> bool { pub fn is_harmony_model_in_registry(registry: &WorkerRegistry, model_name: &str) -> bool {
// Get workers for this model // Get workers for this model
let workers = registry.get_by_model_fast(model_name); let workers = registry.get_by_model(model_name);
if workers.is_empty() { if workers.is_empty() {
// No workers found - fall back to string-based detection // No workers found - fall back to string-based detection
@@ -710,7 +710,7 @@ impl PDRouter {
let prefill_workers = if let Some(model) = effective_model_id { let prefill_workers = if let Some(model) = effective_model_id {
self.worker_registry self.worker_registry
.get_by_model_fast(model) .get_by_model(model)
.iter() .iter()
.filter(|w| matches!(w.worker_type(), WorkerType::Prefill { .. })) .filter(|w| matches!(w.worker_type(), WorkerType::Prefill { .. }))
.cloned() .cloned()
@@ -721,7 +721,7 @@ impl PDRouter {
let decode_workers = if let Some(model) = effective_model_id { let decode_workers = if let Some(model) = effective_model_id {
self.worker_registry self.worker_registry
.get_by_model_fast(model) .get_by_model(model)
.iter() .iter()
.filter(|w| matches!(w.worker_type(), WorkerType::Decode)) .filter(|w| matches!(w.worker_type(), WorkerType::Decode))
.cloned() .cloned()