Support min num routing keys in key-based load balancing policy (#16564)
This commit is contained in:
@@ -299,10 +299,12 @@ async fn generate_handler(
|
||||
Json(payload): Json<serde_json::Value>,
|
||||
) -> Response {
|
||||
let config = config.read().await;
|
||||
let worker_id = format!("worker-{}", config.port);
|
||||
|
||||
if should_fail(&config).await {
|
||||
return (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
[("x-worker-id", worker_id)],
|
||||
Json(json!({
|
||||
"error": "Random failure for testing"
|
||||
})),
|
||||
@@ -373,28 +375,33 @@ async fn generate_handler(
|
||||
|
||||
let stream = stream::iter(events);
|
||||
|
||||
Sse::new(stream)
|
||||
.keep_alive(KeepAlive::default())
|
||||
(
|
||||
[("x-worker-id", worker_id)],
|
||||
Sse::new(stream).keep_alive(KeepAlive::default()),
|
||||
)
|
||||
.into_response()
|
||||
} else {
|
||||
Json(json!({
|
||||
"text": "This is a mock response.",
|
||||
"meta_info": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"completion_tokens_wo_jump_forward": 5,
|
||||
"input_token_logprobs": null,
|
||||
"output_token_logprobs": null,
|
||||
"first_token_latency": config.response_delay_ms as f64 / 1000.0,
|
||||
"time_to_first_token": config.response_delay_ms as f64 / 1000.0,
|
||||
"time_per_output_token": 0.01,
|
||||
"finish_reason": {
|
||||
"type": "stop",
|
||||
"reason": "length"
|
||||
(
|
||||
[("x-worker-id", worker_id)],
|
||||
Json(json!({
|
||||
"text": "This is a mock response.",
|
||||
"meta_info": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"completion_tokens_wo_jump_forward": 5,
|
||||
"input_token_logprobs": null,
|
||||
"output_token_logprobs": null,
|
||||
"first_token_latency": config.response_delay_ms as f64 / 1000.0,
|
||||
"time_to_first_token": config.response_delay_ms as f64 / 1000.0,
|
||||
"time_per_output_token": 0.01,
|
||||
"finish_reason": {
|
||||
"type": "stop",
|
||||
"reason": "length"
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
.into_response()
|
||||
})),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -3,7 +3,9 @@
|
||||
//! Provides pre-configured RouterConfig and MockWorkerConfig builders
|
||||
//! for common test scenarios.
|
||||
|
||||
use smg::config::{CircuitBreakerConfig, PolicyConfig, RetryConfig, RouterConfig};
|
||||
use smg::config::{
|
||||
CircuitBreakerConfig, ManualAssignmentMode, PolicyConfig, RetryConfig, RouterConfig,
|
||||
};
|
||||
|
||||
use super::mock_worker::{HealthStatus, MockWorkerConfig, WorkerType};
|
||||
|
||||
@@ -94,12 +96,22 @@ impl TestRouterConfig {
|
||||
|
||||
/// Create a manual routing config (for sticky routing tests)
|
||||
pub fn manual(port: u16) -> RouterConfig {
|
||||
Self::manual_with_mode(port, ManualAssignmentMode::Random)
|
||||
}
|
||||
|
||||
/// Create a manual routing config with min_group assignment mode
|
||||
pub fn manual_min_group(port: u16) -> RouterConfig {
|
||||
Self::manual_with_mode(port, ManualAssignmentMode::MinGroup)
|
||||
}
|
||||
|
||||
/// Create a manual routing config with specified assignment mode
|
||||
pub fn manual_with_mode(port: u16, assignment_mode: ManualAssignmentMode) -> RouterConfig {
|
||||
RouterConfig::builder()
|
||||
.regular_mode(vec![])
|
||||
.policy(PolicyConfig::Manual {
|
||||
eviction_interval_secs: 60,
|
||||
max_idle_secs: 3600,
|
||||
assignment_mode: Default::default(),
|
||||
assignment_mode,
|
||||
})
|
||||
.host(defaults::HOST)
|
||||
.port(port)
|
||||
@@ -262,6 +274,13 @@ impl TestWorkerConfig {
|
||||
}
|
||||
}
|
||||
|
||||
/// Create multiple slow workers with sequential ports
|
||||
pub fn slow_workers(start_port: u16, count: u16, delay_ms: u64) -> Vec<MockWorkerConfig> {
|
||||
(0..count)
|
||||
.map(|i| Self::slow(start_port + i, delay_ms))
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Create a flaky worker config (for retry/fault tolerance tests)
|
||||
pub fn flaky(port: u16, fail_rate: f32) -> MockWorkerConfig {
|
||||
MockWorkerConfig {
|
||||
|
||||
@@ -172,3 +172,175 @@ mod manual_routing_tests {
|
||||
ctx.shutdown().await;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod manual_min_group_tests {
|
||||
use super::*;
|
||||
|
||||
async fn send_request(app: axum::Router, routing_key: &str) -> (String, String) {
|
||||
let payload = json!({
|
||||
"text": format!("Request for {}", routing_key),
|
||||
"stream": false
|
||||
});
|
||||
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/generate")
|
||||
.header(CONTENT_TYPE, "application/json")
|
||||
.header(ROUTING_KEY_HEADER, routing_key)
|
||||
.body(Body::from(serde_json::to_string(&payload).unwrap()))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let worker_id = resp
|
||||
.headers()
|
||||
.get("x-worker-id")
|
||||
.expect("Response should have x-worker-id header")
|
||||
.to_str()
|
||||
.unwrap()
|
||||
.to_string();
|
||||
|
||||
(routing_key.to_string(), worker_id)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_min_group_concurrent_distribution() {
|
||||
let config = TestRouterConfig::manual_min_group(3910);
|
||||
|
||||
let ctx =
|
||||
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29910, 3, 500))
|
||||
.await;
|
||||
|
||||
let app = ctx.create_app().await;
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..9 {
|
||||
let routing_key = format!("key-{}", i);
|
||||
let app_clone = app.clone();
|
||||
let handle = tokio::spawn(async move { send_request(app_clone, &routing_key).await });
|
||||
handles.push(handle);
|
||||
}
|
||||
|
||||
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|r| r.unwrap())
|
||||
.collect();
|
||||
|
||||
let key_to_worker: HashMap<String, String> = results.into_iter().collect();
|
||||
|
||||
let worker_counts: HashMap<String, usize> =
|
||||
key_to_worker.values().fold(HashMap::new(), |mut acc, w| {
|
||||
*acc.entry(w.clone()).or_default() += 1;
|
||||
acc
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
worker_counts.len(),
|
||||
3,
|
||||
"min_group should distribute keys across all 3 workers, got {:?}",
|
||||
worker_counts
|
||||
);
|
||||
for (worker, count) in &worker_counts {
|
||||
assert_eq!(
|
||||
*count, 3,
|
||||
"Worker {} should have exactly 3 keys, got {}. Distribution: {:?}",
|
||||
worker, count, key_to_worker
|
||||
);
|
||||
}
|
||||
|
||||
ctx.shutdown().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_min_group_sticky_routing() {
|
||||
let config = TestRouterConfig::manual_min_group(3911);
|
||||
|
||||
let ctx =
|
||||
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29920, 3, 200))
|
||||
.await;
|
||||
|
||||
let app = ctx.create_app().await;
|
||||
|
||||
let routing_key = "sticky-key-123";
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for _ in 0..5 {
|
||||
let app_clone = app.clone();
|
||||
let key = routing_key.to_string();
|
||||
let handle = tokio::spawn(async move { send_request(app_clone, &key).await });
|
||||
handles.push(handle);
|
||||
}
|
||||
|
||||
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|r| r.unwrap())
|
||||
.collect();
|
||||
|
||||
let workers: Vec<String> = results.into_iter().map(|(_, w)| w).collect();
|
||||
let unique_workers: HashSet<&String> = workers.iter().collect();
|
||||
assert_eq!(
|
||||
unique_workers.len(),
|
||||
1,
|
||||
"All requests with same routing key should route to same worker, got {:?}",
|
||||
unique_workers
|
||||
);
|
||||
|
||||
ctx.shutdown().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_min_group_mixed_concurrent_routing() {
|
||||
let config = TestRouterConfig::manual_min_group(3912);
|
||||
|
||||
let ctx =
|
||||
AppTestContext::new_with_config(config, TestWorkerConfig::slow_workers(29930, 2, 300))
|
||||
.await;
|
||||
|
||||
let app = ctx.create_app().await;
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..4 {
|
||||
let routing_key = format!("key-{}", i);
|
||||
for _ in 0..3 {
|
||||
let app_clone = app.clone();
|
||||
let key = routing_key.clone();
|
||||
let handle = tokio::spawn(async move { send_request(app_clone, &key).await });
|
||||
handles.push(handle);
|
||||
}
|
||||
}
|
||||
|
||||
let results: Vec<(String, String)> = futures_util::future::join_all(handles)
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|r| r.unwrap())
|
||||
.collect();
|
||||
|
||||
let mut key_to_workers: HashMap<String, HashSet<String>> = HashMap::new();
|
||||
for (key, worker) in results {
|
||||
key_to_workers.entry(key).or_default().insert(worker);
|
||||
}
|
||||
|
||||
for (key, workers) in &key_to_workers {
|
||||
assert_eq!(
|
||||
workers.len(),
|
||||
1,
|
||||
"Key {} should route to exactly one worker (sticky), but got {:?}",
|
||||
key,
|
||||
workers
|
||||
);
|
||||
}
|
||||
|
||||
let all_workers: HashSet<String> = key_to_workers.values().flatten().cloned().collect();
|
||||
assert_eq!(
|
||||
all_workers.len(),
|
||||
2,
|
||||
"Keys should be distributed across both workers"
|
||||
);
|
||||
|
||||
ctx.shutdown().await;
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user