Tiny refactor select_workers API for future passing more information (#15596)
This commit is contained in:
+1321
-1127
File diff suppressed because it is too large
Load Diff
@@ -74,7 +74,7 @@ use tracing::debug;
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig,
|
get_healthy_worker_indices, normalize_model_key, tree::Tree, CacheAwareConfig,
|
||||||
LoadBalancingPolicy,
|
LoadBalancingPolicy, SelectWorkerInfo,
|
||||||
};
|
};
|
||||||
use crate::core::Worker;
|
use crate::core::Worker;
|
||||||
|
|
||||||
@@ -284,11 +284,8 @@ impl CacheAwarePolicy {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl LoadBalancingPolicy for CacheAwarePolicy {
|
impl LoadBalancingPolicy for CacheAwarePolicy {
|
||||||
fn select_worker(
|
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize> {
|
||||||
&self,
|
let request_text = info.request_text;
|
||||||
workers: &[Arc<dyn Worker>],
|
|
||||||
request_text: Option<&str>,
|
|
||||||
) -> Option<usize> {
|
|
||||||
let healthy_indices = get_healthy_worker_indices(workers);
|
let healthy_indices = get_healthy_worker_indices(workers);
|
||||||
|
|
||||||
if healthy_indices.is_empty() {
|
if healthy_indices.is_empty() {
|
||||||
@@ -458,14 +455,35 @@ mod tests {
|
|||||||
policy.init_workers(&workers);
|
policy.init_workers(&workers);
|
||||||
|
|
||||||
// First request should be distributed
|
// First request should be distributed
|
||||||
let idx1 = policy.select_worker(&workers, Some("hello world")).unwrap();
|
let idx1 = policy
|
||||||
|
.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("hello world"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
// Same request should go to same worker (cache hit)
|
// Same request should go to same worker (cache hit)
|
||||||
let idx2 = policy.select_worker(&workers, Some("hello world")).unwrap();
|
let idx2 = policy
|
||||||
|
.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("hello world"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
assert_eq!(idx1, idx2);
|
assert_eq!(idx1, idx2);
|
||||||
|
|
||||||
// Similar request should also go to same worker
|
// Similar request should also go to same worker
|
||||||
let idx3 = policy.select_worker(&workers, Some("hello")).unwrap();
|
let idx3 = policy
|
||||||
|
.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("hello"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
assert_eq!(idx1, idx3);
|
assert_eq!(idx1, idx3);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -496,8 +514,11 @@ mod tests {
|
|||||||
policy.init_workers(&workers);
|
policy.init_workers(&workers);
|
||||||
|
|
||||||
// Should select worker2 (lower load) despite cache affinity
|
// Should select worker2 (lower load) despite cache affinity
|
||||||
|
let info = SelectWorkerInfo {
|
||||||
|
request_text: Some("test"),
|
||||||
|
};
|
||||||
for _ in 0..5 {
|
for _ in 0..5 {
|
||||||
let idx = policy.select_worker(&workers, Some("test")).unwrap();
|
let idx = policy.select_worker(&workers, &info).unwrap();
|
||||||
assert_eq!(idx, 1); // Should always pick worker2
|
assert_eq!(idx, 1); // Should always pick worker2
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -525,15 +546,32 @@ mod tests {
|
|||||||
policy.init_workers(&workers);
|
policy.init_workers(&workers);
|
||||||
|
|
||||||
// Route some requests
|
// Route some requests
|
||||||
policy.select_worker(&workers, Some("test1"));
|
policy.select_worker(
|
||||||
policy.select_worker(&workers, Some("test2"));
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("test1"),
|
||||||
|
},
|
||||||
|
);
|
||||||
|
policy.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("test2"),
|
||||||
|
},
|
||||||
|
);
|
||||||
|
|
||||||
// Remove a worker
|
// Remove a worker
|
||||||
policy.remove_worker_by_url("http://w1:8000");
|
policy.remove_worker_by_url("http://w1:8000");
|
||||||
workers[0].set_healthy(false);
|
workers[0].set_healthy(false);
|
||||||
|
|
||||||
// All requests should now go to worker2
|
// All requests should now go to worker2
|
||||||
let idx = policy.select_worker(&workers, Some("test1")).unwrap();
|
let idx = policy
|
||||||
|
.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("test1"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
assert_eq!(idx, 1);
|
assert_eq!(idx, 1);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -33,11 +33,11 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug {
|
|||||||
///
|
///
|
||||||
/// This is used for regular routing mode where requests go to a single worker.
|
/// This is used for regular routing mode where requests go to a single worker.
|
||||||
/// Now uses Arc<dyn Worker> for better performance and to avoid unnecessary cloning.
|
/// Now uses Arc<dyn Worker> for better performance and to avoid unnecessary cloning.
|
||||||
fn select_worker(
|
///
|
||||||
&self,
|
/// # Arguments
|
||||||
workers: &[Arc<dyn Worker>],
|
/// * `workers` - Available workers to select from
|
||||||
request_text: Option<&str>,
|
/// * `info` - Additional information for routing decisions
|
||||||
) -> Option<usize>;
|
fn select_worker(&self, workers: &[Arc<dyn Worker>], info: &SelectWorkerInfo) -> Option<usize>;
|
||||||
|
|
||||||
/// Update policy state after request completion
|
/// Update policy state after request completion
|
||||||
///
|
///
|
||||||
@@ -135,6 +135,13 @@ pub(crate) fn normalize_model_key(model_id: &str) -> &str {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Information passed to policy for worker selection
|
||||||
|
#[derive(Debug, Default, Clone)]
|
||||||
|
pub struct SelectWorkerInfo<'a> {
|
||||||
|
/// Request text for cache-aware routing
|
||||||
|
pub request_text: Option<&'a str>,
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ use std::{
|
|||||||
use rand::Rng;
|
use rand::Rng;
|
||||||
use tracing::debug;
|
use tracing::debug;
|
||||||
|
|
||||||
use super::{get_healthy_worker_indices, LoadBalancingPolicy};
|
use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
|
||||||
use crate::core::Worker;
|
use crate::core::Worker;
|
||||||
|
|
||||||
/// Power-of-two choices policy
|
/// Power-of-two choices policy
|
||||||
@@ -33,7 +33,7 @@ impl LoadBalancingPolicy for PowerOfTwoPolicy {
|
|||||||
fn select_worker(
|
fn select_worker(
|
||||||
&self,
|
&self,
|
||||||
workers: &[Arc<dyn Worker>],
|
workers: &[Arc<dyn Worker>],
|
||||||
_request_text: Option<&str>,
|
_info: &SelectWorkerInfo,
|
||||||
) -> Option<usize> {
|
) -> Option<usize> {
|
||||||
let healthy_indices = get_healthy_worker_indices(workers);
|
let healthy_indices = get_healthy_worker_indices(workers);
|
||||||
|
|
||||||
@@ -157,8 +157,9 @@ mod tests {
|
|||||||
|
|
||||||
// Run multiple selections
|
// Run multiple selections
|
||||||
let mut selected_counts = [0; 3];
|
let mut selected_counts = [0; 3];
|
||||||
|
let info = SelectWorkerInfo::default();
|
||||||
for _ in 0..100 {
|
for _ in 0..100 {
|
||||||
if let Some(idx) = policy.select_worker(&workers, None) {
|
if let Some(idx) = policy.select_worker(&workers, &info) {
|
||||||
selected_counts[idx] += 1;
|
selected_counts[idx] += 1;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -192,8 +193,9 @@ mod tests {
|
|||||||
|
|
||||||
// Should prefer worker2 with lower cached load
|
// Should prefer worker2 with lower cached load
|
||||||
let mut w2_selected = 0;
|
let mut w2_selected = 0;
|
||||||
|
let info = SelectWorkerInfo::default();
|
||||||
for _ in 0..50 {
|
for _ in 0..50 {
|
||||||
if let Some(idx) = policy.select_worker(&workers, None) {
|
if let Some(idx) = policy.select_worker(&workers, &info) {
|
||||||
if idx == 1 {
|
if idx == 1 {
|
||||||
w2_selected += 1;
|
w2_selected += 1;
|
||||||
}
|
}
|
||||||
@@ -214,7 +216,10 @@ mod tests {
|
|||||||
)];
|
)];
|
||||||
|
|
||||||
// With single worker, should always select it
|
// With single worker, should always select it
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
assert_eq!(
|
||||||
|
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||||
|
Some(0)
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -251,7 +256,7 @@ mod tests {
|
|||||||
|
|
||||||
// 5. Run selection
|
// 5. Run selection
|
||||||
let selected_idx = policy
|
let selected_idx = policy
|
||||||
.select_worker(&workers, None)
|
.select_worker(&workers, &SelectWorkerInfo::default())
|
||||||
.expect("Should select a worker");
|
.expect("Should select a worker");
|
||||||
|
|
||||||
// 6. Verify the Fix
|
// 6. Verify the Fix
|
||||||
@@ -307,7 +312,9 @@ mod tests {
|
|||||||
loads_1.insert("http://b:8000".to_string(), 100_000);
|
loads_1.insert("http://b:8000".to_string(), 100_000);
|
||||||
policy.update_loads(&loads_1);
|
policy.update_loads(&loads_1);
|
||||||
|
|
||||||
let idx_1 = policy.select_worker(&workers_1, None).unwrap();
|
let idx_1 = policy
|
||||||
|
.select_worker(&workers_1, &SelectWorkerInfo::default())
|
||||||
|
.unwrap();
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
idx_1, 0,
|
idx_1, 0,
|
||||||
"Happy Path Failed: Should select Worker A (fewer tokens) despite higher request count"
|
"Happy Path Failed: Should select Worker A (fewer tokens) despite higher request count"
|
||||||
@@ -326,7 +333,9 @@ mod tests {
|
|||||||
// http://d:8000 is MISSING
|
// http://d:8000 is MISSING
|
||||||
policy.update_loads(&loads_2);
|
policy.update_loads(&loads_2);
|
||||||
|
|
||||||
let idx_2 = policy.select_worker(&workers_2, None).unwrap();
|
let idx_2 = policy
|
||||||
|
.select_worker(&workers_2, &SelectWorkerInfo::default())
|
||||||
|
.unwrap();
|
||||||
assert_eq!(idx_2, 1, "Partial Fail 1 Failed: Should fallback to requests and select Worker B (fewer requests)");
|
assert_eq!(idx_2, 1, "Partial Fail 1 Failed: Should fallback to requests and select Worker B (fewer requests)");
|
||||||
|
|
||||||
// Scenario 3: Partial Failure (Worker A is missing, Worker B has tokens)
|
// Scenario 3: Partial Failure (Worker A is missing, Worker B has tokens)
|
||||||
@@ -342,7 +351,9 @@ mod tests {
|
|||||||
loads_3.insert("http://f:8000".to_string(), 1_000);
|
loads_3.insert("http://f:8000".to_string(), 1_000);
|
||||||
policy.update_loads(&loads_3);
|
policy.update_loads(&loads_3);
|
||||||
|
|
||||||
let idx_3 = policy.select_worker(&workers_3, None).unwrap();
|
let idx_3 = policy
|
||||||
|
.select_worker(&workers_3, &SelectWorkerInfo::default())
|
||||||
|
.unwrap();
|
||||||
assert_eq!(idx_3, 0, "Partial Fail 2 Failed: Should fallback to requests and select Worker A (fewer requests)");
|
assert_eq!(idx_3, 0, "Partial Fail 2 Failed: Should fallback to requests and select Worker A (fewer requests)");
|
||||||
|
|
||||||
// Scenario 4: Total Failure (Both missing)
|
// Scenario 4: Total Failure (Both missing)
|
||||||
@@ -356,7 +367,9 @@ mod tests {
|
|||||||
let loads_4 = HashMap::new();
|
let loads_4 = HashMap::new();
|
||||||
policy.update_loads(&loads_4);
|
policy.update_loads(&loads_4);
|
||||||
|
|
||||||
let idx_4 = policy.select_worker(&workers_4, None).unwrap();
|
let idx_4 = policy
|
||||||
|
.select_worker(&workers_4, &SelectWorkerInfo::default())
|
||||||
|
.unwrap();
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
idx_4, 1,
|
idx_4, 1,
|
||||||
"Total Fail Failed: Should select Worker B based on request count"
|
"Total Fail Failed: Should select Worker B based on request count"
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use rand::Rng;
|
use rand::Rng;
|
||||||
|
|
||||||
use super::{get_healthy_worker_indices, LoadBalancingPolicy};
|
use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
|
||||||
use crate::core::Worker;
|
use crate::core::Worker;
|
||||||
|
|
||||||
/// Random selection policy
|
/// Random selection policy
|
||||||
@@ -23,7 +23,7 @@ impl LoadBalancingPolicy for RandomPolicy {
|
|||||||
fn select_worker(
|
fn select_worker(
|
||||||
&self,
|
&self,
|
||||||
workers: &[Arc<dyn Worker>],
|
workers: &[Arc<dyn Worker>],
|
||||||
_request_text: Option<&str>,
|
_info: &SelectWorkerInfo,
|
||||||
) -> Option<usize> {
|
) -> Option<usize> {
|
||||||
let healthy_indices = get_healthy_worker_indices(workers);
|
let healthy_indices = get_healthy_worker_indices(workers);
|
||||||
|
|
||||||
@@ -76,7 +76,7 @@ mod tests {
|
|||||||
|
|
||||||
let mut counts = HashMap::new();
|
let mut counts = HashMap::new();
|
||||||
for _ in 0..100 {
|
for _ in 0..100 {
|
||||||
if let Some(idx) = policy.select_worker(&workers, None) {
|
if let Some(idx) = policy.select_worker(&workers, &SelectWorkerInfo::default()) {
|
||||||
*counts.entry(idx).or_insert(0) += 1;
|
*counts.entry(idx).or_insert(0) += 1;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -107,7 +107,10 @@ mod tests {
|
|||||||
|
|
||||||
// Should always select the healthy worker (index 1)
|
// Should always select the healthy worker (index 1)
|
||||||
for _ in 0..10 {
|
for _ in 0..10 {
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(1));
|
assert_eq!(
|
||||||
|
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||||
|
Some(1)
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -121,6 +124,9 @@ mod tests {
|
|||||||
)];
|
)];
|
||||||
|
|
||||||
workers[0].set_healthy(false);
|
workers[0].set_healthy(false);
|
||||||
assert_eq!(policy.select_worker(&workers, None), None);
|
assert_eq!(
|
||||||
|
policy.select_worker(&workers, &SelectWorkerInfo::default()),
|
||||||
|
None
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ use std::sync::{
|
|||||||
Arc,
|
Arc,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{get_healthy_worker_indices, LoadBalancingPolicy};
|
use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo};
|
||||||
use crate::core::Worker;
|
use crate::core::Worker;
|
||||||
|
|
||||||
/// Round-robin selection policy
|
/// Round-robin selection policy
|
||||||
@@ -28,7 +28,7 @@ impl LoadBalancingPolicy for RoundRobinPolicy {
|
|||||||
fn select_worker(
|
fn select_worker(
|
||||||
&self,
|
&self,
|
||||||
workers: &[Arc<dyn Worker>],
|
workers: &[Arc<dyn Worker>],
|
||||||
_request_text: Option<&str>,
|
_info: &SelectWorkerInfo,
|
||||||
) -> Option<usize> {
|
) -> Option<usize> {
|
||||||
let healthy_indices = get_healthy_worker_indices(workers);
|
let healthy_indices = get_healthy_worker_indices(workers);
|
||||||
|
|
||||||
@@ -83,11 +83,12 @@ mod tests {
|
|||||||
];
|
];
|
||||||
|
|
||||||
// Should select workers in order: 0, 1, 2, 0, 1, 2, ...
|
// Should select workers in order: 0, 1, 2, 0, 1, 2, ...
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
let info = SelectWorkerInfo::default();
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(1));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(2));
|
assert_eq!(policy.select_worker(&workers, &info), Some(1));
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
assert_eq!(policy.select_worker(&workers, &info), Some(2));
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(1));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
|
assert_eq!(policy.select_worker(&workers, &info), Some(1));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -115,10 +116,11 @@ mod tests {
|
|||||||
workers[1].set_healthy(false);
|
workers[1].set_healthy(false);
|
||||||
|
|
||||||
// Should skip unhealthy worker: 0, 2, 0, 2, ...
|
// Should skip unhealthy worker: 0, 2, 0, 2, ...
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
let info = SelectWorkerInfo::default();
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(2));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
assert_eq!(policy.select_worker(&workers, &info), Some(2));
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(2));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
|
assert_eq!(policy.select_worker(&workers, &info), Some(2));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -138,11 +140,12 @@ mod tests {
|
|||||||
];
|
];
|
||||||
|
|
||||||
// Advance the counter
|
// Advance the counter
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
let info = SelectWorkerInfo::default();
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(1));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
|
assert_eq!(policy.select_worker(&workers, &info), Some(1));
|
||||||
|
|
||||||
// Reset should start from beginning
|
// Reset should start from beginning
|
||||||
policy.reset();
|
policy.reset();
|
||||||
assert_eq!(policy.select_worker(&workers, None), Some(0));
|
assert_eq!(policy.select_worker(&workers, &info), Some(0));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ use super::PipelineStage;
|
|||||||
use crate::{
|
use crate::{
|
||||||
core::{ConnectionMode, Worker, WorkerRegistry, WorkerType},
|
core::{ConnectionMode, Worker, WorkerRegistry, WorkerType},
|
||||||
observability::metrics::{metrics_labels, Metrics},
|
observability::metrics::{metrics_labels, Metrics},
|
||||||
policies::PolicyRegistry,
|
policies::{PolicyRegistry, SelectWorkerInfo},
|
||||||
routers::{
|
routers::{
|
||||||
error,
|
error,
|
||||||
grpc::context::{RequestContext, WorkerSelection},
|
grpc::context::{RequestContext, WorkerSelection},
|
||||||
@@ -146,7 +146,7 @@ impl WorkerSelectionStage {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Select worker using the policy
|
// Select worker using the policy
|
||||||
let idx = policy.select_worker(&available, text)?;
|
let idx = policy.select_worker(&available, &SelectWorkerInfo { request_text: text })?;
|
||||||
let selected = available[idx].clone();
|
let selected = available[idx].clone();
|
||||||
|
|
||||||
// Record worker selection metric
|
// Record worker selection metric
|
||||||
@@ -203,8 +203,9 @@ impl WorkerSelectionStage {
|
|||||||
None => self.policy_registry.get_default_policy(),
|
None => self.policy_registry.get_default_policy(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let prefill_idx = policy.select_worker(&available_prefill, text)?;
|
let info = SelectWorkerInfo { request_text: text };
|
||||||
let decode_idx = policy.select_worker(&available_decode, text)?;
|
let prefill_idx = policy.select_worker(&available_prefill, &info)?;
|
||||||
|
let decode_idx = policy.select_worker(&available_decode, &info)?;
|
||||||
|
|
||||||
let model = model_id.unwrap_or("default");
|
let model = model_id.unwrap_or("default");
|
||||||
let policy_name = policy.name();
|
let policy_name = policy.name();
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ use crate::{
|
|||||||
metrics::{bool_to_static_str, metrics_labels, Metrics},
|
metrics::{bool_to_static_str, metrics_labels, Metrics},
|
||||||
otel_trace::inject_trace_context_http,
|
otel_trace::inject_trace_context_http,
|
||||||
},
|
},
|
||||||
policies::{LoadBalancingPolicy, PolicyRegistry},
|
policies::{LoadBalancingPolicy, PolicyRegistry, SelectWorkerInfo},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatMessage, MessageContent},
|
chat::{ChatCompletionRequest, ChatMessage, MessageContent},
|
||||||
common::{InputIds, StringOrArray},
|
common::{InputIds, StringOrArray},
|
||||||
@@ -784,7 +784,7 @@ impl PDRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let selected_idx = policy
|
let selected_idx = policy
|
||||||
.select_worker(&available_workers, request_text)
|
.select_worker(&available_workers, &SelectWorkerInfo { request_text })
|
||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
format!(
|
format!(
|
||||||
"Policy {} failed to select a {} worker",
|
"Policy {} failed to select a {} worker",
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ use crate::{
|
|||||||
metrics::{bool_to_static_str, metrics_labels, Metrics},
|
metrics::{bool_to_static_str, metrics_labels, Metrics},
|
||||||
otel_trace::inject_trace_context_http,
|
otel_trace::inject_trace_context_http,
|
||||||
},
|
},
|
||||||
policies::PolicyRegistry,
|
policies::{PolicyRegistry, SelectWorkerInfo},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::ChatCompletionRequest,
|
chat::ChatCompletionRequest,
|
||||||
classify::ClassifyRequest,
|
classify::ClassifyRequest,
|
||||||
@@ -168,7 +168,7 @@ impl Router {
|
|||||||
None => self.policy_registry.get_default_policy(),
|
None => self.policy_registry.get_default_policy(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let idx = policy.select_worker(&available, text)?;
|
let idx = policy.select_worker(&available, &SelectWorkerInfo { request_text: text })?;
|
||||||
|
|
||||||
// Record worker selection metric (Layer 3)
|
// Record worker selection metric (Layer 3)
|
||||||
Metrics::record_worker_selection(
|
Metrics::record_worker_selection(
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ use std::{collections::HashMap, sync::Arc};
|
|||||||
|
|
||||||
use sgl_model_gateway::{
|
use sgl_model_gateway::{
|
||||||
core::{BasicWorkerBuilder, Worker, WorkerType},
|
core::{BasicWorkerBuilder, Worker, WorkerType},
|
||||||
policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy},
|
policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy, SelectWorkerInfo},
|
||||||
};
|
};
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -40,7 +40,12 @@ fn test_backward_compatibility_with_empty_model_id() {
|
|||||||
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker1.clone()), Arc::new(worker2.clone())];
|
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker1.clone()), Arc::new(worker2.clone())];
|
||||||
|
|
||||||
// Select worker - should work without errors
|
// Select worker - should work without errors
|
||||||
let selected = policy.select_worker(&workers, Some("test request"));
|
let selected = policy.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("test request"),
|
||||||
|
},
|
||||||
|
);
|
||||||
assert!(selected.is_some(), "Should select a worker");
|
assert!(selected.is_some(), "Should select a worker");
|
||||||
|
|
||||||
// Remove workers - should work without errors
|
// Remove workers - should work without errors
|
||||||
@@ -97,12 +102,15 @@ fn test_mixed_model_ids() {
|
|||||||
|
|
||||||
let default_workers: Vec<Arc<dyn Worker>> =
|
let default_workers: Vec<Arc<dyn Worker>> =
|
||||||
vec![Arc::new(worker1.clone()), Arc::new(worker3.clone())];
|
vec![Arc::new(worker1.clone()), Arc::new(worker3.clone())];
|
||||||
let selected = policy.select_worker(&default_workers, Some("test request"));
|
let info = SelectWorkerInfo {
|
||||||
|
request_text: Some("test request"),
|
||||||
|
};
|
||||||
|
let selected = policy.select_worker(&default_workers, &info);
|
||||||
assert!(selected.is_some(), "Should select from default workers");
|
assert!(selected.is_some(), "Should select from default workers");
|
||||||
|
|
||||||
let llama_workers: Vec<Arc<dyn Worker>> =
|
let llama_workers: Vec<Arc<dyn Worker>> =
|
||||||
vec![Arc::new(worker2.clone()), Arc::new(worker4.clone())];
|
vec![Arc::new(worker2.clone()), Arc::new(worker4.clone())];
|
||||||
let selected = policy.select_worker(&llama_workers, Some("test request"));
|
let selected = policy.select_worker(&llama_workers, &info);
|
||||||
assert!(selected.is_some(), "Should select from llama-3 workers");
|
assert!(selected.is_some(), "Should select from llama-3 workers");
|
||||||
|
|
||||||
let all_workers: Vec<Arc<dyn Worker>> = vec![
|
let all_workers: Vec<Arc<dyn Worker>> = vec![
|
||||||
@@ -111,7 +119,7 @@ fn test_mixed_model_ids() {
|
|||||||
Arc::new(worker3.clone()),
|
Arc::new(worker3.clone()),
|
||||||
Arc::new(worker4.clone()),
|
Arc::new(worker4.clone()),
|
||||||
];
|
];
|
||||||
let selected = policy.select_worker(&all_workers, Some("test request"));
|
let selected = policy.select_worker(&all_workers, &info);
|
||||||
assert!(selected.is_some(), "Should select from all workers");
|
assert!(selected.is_some(), "Should select from all workers");
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -144,6 +152,11 @@ fn test_remove_worker_by_url_backward_compat() {
|
|||||||
policy.remove_worker_by_url("http://worker1:8080");
|
policy.remove_worker_by_url("http://worker1:8080");
|
||||||
|
|
||||||
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker2.clone())];
|
let workers: Vec<Arc<dyn Worker>> = vec![Arc::new(worker2.clone())];
|
||||||
let selected = policy.select_worker(&workers, Some("test"));
|
let selected = policy.select_worker(
|
||||||
|
&workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: Some("test"),
|
||||||
|
},
|
||||||
|
);
|
||||||
assert_eq!(selected, Some(0), "Should only have worker2 left");
|
assert_eq!(selected, Some(0), "Should only have worker2 left");
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user