From c4edcac6d74e67f9559471a89554f522ee71971c Mon Sep 17 00:00:00 2001 From: Chang Su Date: Fri, 2 Jan 2026 16:04:23 -0800 Subject: [PATCH] refactor(core): remove get_by_model_fast alias in worker_registry (#16313) --- .../local/update_policies_for_worker.rs | 2 +- .../worker/local/update_remaining_policies.rs | 2 +- .../steps/worker/shared/update_policies.rs | 2 +- sgl-model-gateway/src/core/worker_registry.rs | 19 +++++-------------- .../src/routers/grpc/harmony/detector.rs | 2 +- .../src/routers/http/pd_router.rs | 4 ++-- 6 files changed, 11 insertions(+), 20 deletions(-) diff --git a/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs b/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs index 2deeeae4c..7dced12a5 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs @@ -35,7 +35,7 @@ impl StepExecutor for UpdatePoliciesForWorkerStep { ); 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 policy.name() == "cache_aware" && !workers.is_empty() { diff --git a/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs b/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs index a02442c51..147529350 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs @@ -29,7 +29,7 @@ impl StepExecutor for UpdateRemainingPoliciesStep { ); 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 policy.name() == "cache_aware" && !remaining_workers.is_empty() { diff --git a/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs b/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs index 282bd2972..e2733066f 100644 --- a/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs +++ b/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs @@ -38,7 +38,7 @@ impl StepExecutor for UpdatePoliciesStep { .on_worker_added(&model_id, policy_hint); // 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 policy.name() == "cache_aware" { app_context diff --git a/sgl-model-gateway/src/core/worker_registry.rs b/sgl-model-gateway/src/core/worker_registry.rs index 0e59bf041..14b86eb0d 100644 --- a/sgl-model-gateway/src/core/worker_registry.rs +++ b/sgl-model-gateway/src/core/worker_registry.rs @@ -356,12 +356,6 @@ impl WorkerRegistry { .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]> { - self.get_by_model(model_id) - } - /// Get all workers by worker type pub fn get_by_type(&self, worker_type: &WorkerType) -> Vec> { self.type_workers @@ -471,7 +465,7 @@ impl WorkerRegistry { // Start with the most efficient collection based on filters // Use model index when possible as it's O(1) lookup let workers: Vec> = if let Some(model) = model_id { - self.get_by_model_fast(model).to_vec() + self.get_by_model(model).to_vec() } else { self.get_all() }; @@ -764,24 +758,21 @@ mod tests { registry.register(Arc::from(worker2)); 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); let urls: Vec = llama_workers.iter().map(|w| w.url().to_string()).collect(); assert!(urls.contains(&"http://worker1: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[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); - 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"); - 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[0].url(), "http://worker2:8080"); } diff --git a/sgl-model-gateway/src/routers/grpc/harmony/detector.rs b/sgl-model-gateway/src/routers/grpc/harmony/detector.rs index 18fcf4b6e..38fecc045 100644 --- a/sgl-model-gateway/src/routers/grpc/harmony/detector.rs +++ b/sgl-model-gateway/src/routers/grpc/harmony/detector.rs @@ -62,7 +62,7 @@ impl HarmonyDetector { /// the model (e.g., during startup before workers are discovered). pub fn is_harmony_model_in_registry(registry: &WorkerRegistry, model_name: &str) -> bool { // 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() { // No workers found - fall back to string-based detection diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 4408c44e6..fd3d17252 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -710,7 +710,7 @@ impl PDRouter { let prefill_workers = if let Some(model) = effective_model_id { self.worker_registry - .get_by_model_fast(model) + .get_by_model(model) .iter() .filter(|w| matches!(w.worker_type(), WorkerType::Prefill { .. })) .cloned() @@ -721,7 +721,7 @@ impl PDRouter { let decode_workers = if let Some(model) = effective_model_id { self.worker_registry - .get_by_model_fast(model) + .get_by_model(model) .iter() .filter(|w| matches!(w.worker_type(), WorkerType::Decode)) .cloned()