refactor(core): remove get_by_model_fast alias in worker_registry (#16313)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user