[model-gateway] update ManualPolicy with header-based routing (#15847)
This commit is contained in:
@@ -337,6 +337,9 @@ pub enum PolicyConfig {
|
|||||||
bucket_adjust_interval_secs: usize,
|
bucket_adjust_interval_secs: usize,
|
||||||
},
|
},
|
||||||
|
|
||||||
|
/// Manual routing policy supporting header-based routing:
|
||||||
|
/// - X-SMG-Target-Worker: Direct routing to a specific worker by URL
|
||||||
|
/// - X-SMG-Routing-Key: Consistent hash routing for session affinity
|
||||||
#[serde(rename = "manual")]
|
#[serde(rename = "manual")]
|
||||||
Manual,
|
Manual,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -147,7 +147,7 @@ impl ConfigValidator {
|
|||||||
|
|
||||||
fn validate_policy(policy: &PolicyConfig) -> ConfigResult<()> {
|
fn validate_policy(policy: &PolicyConfig) -> ConfigResult<()> {
|
||||||
match policy {
|
match policy {
|
||||||
PolicyConfig::Random | PolicyConfig::RoundRobin => {}
|
PolicyConfig::Random | PolicyConfig::RoundRobin | PolicyConfig::Manual => {}
|
||||||
PolicyConfig::CacheAware {
|
PolicyConfig::CacheAware {
|
||||||
cache_threshold,
|
cache_threshold,
|
||||||
balance_abs_threshold: _,
|
balance_abs_threshold: _,
|
||||||
@@ -226,7 +226,6 @@ impl ConfigValidator {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
PolicyConfig::Manual => {}
|
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -136,7 +136,7 @@ struct CliArgs {
|
|||||||
#[arg(long, num_args = 0..)]
|
#[arg(long, num_args = 0..)]
|
||||||
worker_urls: Vec<String>,
|
worker_urls: Vec<String>,
|
||||||
|
|
||||||
#[arg(long, default_value = "cache_aware", value_parser = ["random", "round_robin", "cache_aware", "power_of_two"])]
|
#[arg(long, default_value = "cache_aware", value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "manual"])]
|
||||||
policy: String,
|
policy: String,
|
||||||
|
|
||||||
#[arg(long, default_value_t = false)]
|
#[arg(long, default_value_t = false)]
|
||||||
@@ -145,10 +145,10 @@ struct CliArgs {
|
|||||||
#[arg(long, action = ArgAction::Append)]
|
#[arg(long, action = ArgAction::Append)]
|
||||||
decode: Vec<String>,
|
decode: Vec<String>,
|
||||||
|
|
||||||
#[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two"])]
|
#[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "manual"])]
|
||||||
prefill_policy: Option<String>,
|
prefill_policy: Option<String>,
|
||||||
|
|
||||||
#[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two"])]
|
#[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "manual"])]
|
||||||
decode_policy: Option<String>,
|
decode_policy: Option<String>,
|
||||||
|
|
||||||
#[arg(long, default_value_t = 1800)]
|
#[arg(long, default_value_t = 1800)]
|
||||||
@@ -415,6 +415,7 @@ impl CliArgs {
|
|||||||
"power_of_two" => PolicyConfig::PowerOfTwo {
|
"power_of_two" => PolicyConfig::PowerOfTwo {
|
||||||
load_check_interval_secs: 5,
|
load_check_interval_secs: 5,
|
||||||
},
|
},
|
||||||
|
"manual" => PolicyConfig::Manual,
|
||||||
_ => PolicyConfig::RoundRobin,
|
_ => PolicyConfig::RoundRobin,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -190,10 +190,6 @@ pub fn init_metrics() {
|
|||||||
"smg_worker_errors_total",
|
"smg_worker_errors_total",
|
||||||
"Worker-level errors by worker_type, connection_mode, error_type"
|
"Worker-level errors by worker_type, connection_mode, error_type"
|
||||||
);
|
);
|
||||||
describe_counter!(
|
|
||||||
"smg_worker_manual_policy_branch_total",
|
|
||||||
"Manual policy execution branch by branch type"
|
|
||||||
);
|
|
||||||
|
|
||||||
// Layer 3: Worker resilience metrics (circuit breaker)
|
// Layer 3: Worker resilience metrics (circuit breaker)
|
||||||
describe_gauge!(
|
describe_gauge!(
|
||||||
@@ -805,6 +801,15 @@ impl Metrics {
|
|||||||
.increment(1);
|
.increment(1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Record manual policy execution branch for routing decisions
|
||||||
|
pub fn record_worker_manual_policy_branch(branch: &'static str) {
|
||||||
|
counter!(
|
||||||
|
"smg_manual_policy_branch_total",
|
||||||
|
"branch" => branch
|
||||||
|
)
|
||||||
|
.increment(1);
|
||||||
|
}
|
||||||
|
|
||||||
/// Set running requests per worker
|
/// Set running requests per worker
|
||||||
pub fn set_worker_requests_active(worker: &str, count: usize) {
|
pub fn set_worker_requests_active(worker: &str, count: usize) {
|
||||||
gauge!(
|
gauge!(
|
||||||
@@ -814,15 +819,6 @@ impl Metrics {
|
|||||||
.set(count as f64);
|
.set(count as f64);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Record manual policy execution branch
|
|
||||||
pub fn record_worker_manual_policy_branch(branch: &'static str) {
|
|
||||||
counter!(
|
|
||||||
"smg_worker_manual_policy_branch_total",
|
|
||||||
"branch" => branch
|
|
||||||
)
|
|
||||||
.increment(1);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Set worker health status
|
/// Set worker health status
|
||||||
pub fn set_worker_health(worker_url: &str, healthy: bool) {
|
pub fn set_worker_health(worker_url: &str, healthy: bool) {
|
||||||
gauge!(
|
gauge!(
|
||||||
|
|||||||
@@ -517,16 +517,12 @@ 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"),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
for _ in 0..5 {
|
for _ in 0..5 {
|
||||||
let idx = policy
|
let idx = policy.select_worker(&workers, &info).unwrap();
|
||||||
.select_worker(
|
|
||||||
&workers,
|
|
||||||
&SelectWorkerInfo {
|
|
||||||
request_text: Some("test"),
|
|
||||||
..Default::default()
|
|
||||||
},
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(idx, 1); // Should always pick worker2
|
assert_eq!(idx, 1); // Should always pick worker2
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -96,9 +96,6 @@ mod tests {
|
|||||||
bucket_adjust_interval_secs: 5,
|
bucket_adjust_interval_secs: 5,
|
||||||
});
|
});
|
||||||
assert_eq!(policy.name(), "bucket");
|
assert_eq!(policy.name(), "bucket");
|
||||||
|
|
||||||
let policy = PolicyFactory::create_from_config(&PolicyConfig::Manual);
|
|
||||||
assert_eq!(policy.name(), "manual");
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -113,8 +110,6 @@ mod tests {
|
|||||||
assert!(PolicyFactory::create_by_name("CacheAware").is_some());
|
assert!(PolicyFactory::create_by_name("CacheAware").is_some());
|
||||||
assert!(PolicyFactory::create_by_name("bucket").is_some());
|
assert!(PolicyFactory::create_by_name("bucket").is_some());
|
||||||
assert!(PolicyFactory::create_by_name("Bucket").is_some());
|
assert!(PolicyFactory::create_by_name("Bucket").is_some());
|
||||||
assert!(PolicyFactory::create_by_name("manual").is_some());
|
|
||||||
assert!(PolicyFactory::create_by_name("Manual").is_some());
|
|
||||||
assert!(PolicyFactory::create_by_name("unknown").is_none());
|
assert!(PolicyFactory::create_by_name("unknown").is_none());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -57,11 +57,6 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug {
|
|||||||
false // Default: most policies don't need request text
|
false // Default: most policies don't need request text
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Check if this policy needs routing_id for routing decisions
|
|
||||||
fn needs_routing_id(&self) -> bool {
|
|
||||||
false // Default: most policies don't need routing_id
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Update worker load information
|
/// Update worker load information
|
||||||
///
|
///
|
||||||
/// This is called periodically with current load information for load-aware policies.
|
/// This is called periodically with current load information for load-aware policies.
|
||||||
@@ -147,8 +142,11 @@ pub(crate) fn normalize_model_key(model_id: &str) -> &str {
|
|||||||
pub struct SelectWorkerInfo<'a> {
|
pub struct SelectWorkerInfo<'a> {
|
||||||
/// Request text for cache-aware routing
|
/// Request text for cache-aware routing
|
||||||
pub request_text: Option<&'a str>,
|
pub request_text: Option<&'a str>,
|
||||||
/// Routing ID for manual routing policy (consistent hashing)
|
/// HTTP headers for header-based routing policies
|
||||||
pub routing_id: Option<&'a str>,
|
/// Policies can extract routing information from headers like:
|
||||||
|
/// - X-Target-Worker: Direct routing to a specific worker by URL
|
||||||
|
/// - X-Routing-Key: Consistent hash routing for session affinity
|
||||||
|
pub headers: Option<&'a http::HeaderMap>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
|||||||
@@ -61,18 +61,17 @@ impl PipelineStage for WorkerSelectionStage {
|
|||||||
|
|
||||||
// For Harmony, use selection_text produced during Harmony encoding
|
// For Harmony, use selection_text produced during Harmony encoding
|
||||||
// Otherwise, use original_text from regular preparation
|
// Otherwise, use original_text from regular preparation
|
||||||
let info = SelectWorkerInfo {
|
let text = if prep.harmony_mode {
|
||||||
request_text: if prep.harmony_mode {
|
prep.selection_text.as_deref()
|
||||||
prep.selection_text.as_deref()
|
} else {
|
||||||
} else {
|
prep.original_text.as_deref()
|
||||||
prep.original_text.as_deref()
|
|
||||||
},
|
|
||||||
routing_id: prep.routing_id.as_deref(),
|
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let headers = ctx.input.headers.as_ref();
|
||||||
|
|
||||||
let workers = match self.mode {
|
let workers = match self.mode {
|
||||||
WorkerSelectionMode::Regular => {
|
WorkerSelectionMode::Regular => {
|
||||||
match self.select_single_worker(ctx.input.model_id.as_deref(), &info) {
|
match self.select_single_worker(ctx.input.model_id.as_deref(), text, headers) {
|
||||||
Some(w) => WorkerSelection::Single { worker: w },
|
Some(w) => WorkerSelection::Single { worker: w },
|
||||||
None => {
|
None => {
|
||||||
error!(
|
error!(
|
||||||
@@ -89,7 +88,7 @@ impl PipelineStage for WorkerSelectionStage {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
WorkerSelectionMode::PrefillDecode => {
|
WorkerSelectionMode::PrefillDecode => {
|
||||||
match self.select_pd_pair(ctx.input.model_id.as_deref(), &info) {
|
match self.select_pd_pair(ctx.input.model_id.as_deref(), text, headers) {
|
||||||
Some((prefill, decode)) => WorkerSelection::Dual { prefill, decode },
|
Some((prefill, decode)) => WorkerSelection::Dual { prefill, decode },
|
||||||
None => {
|
None => {
|
||||||
error!(
|
error!(
|
||||||
@@ -123,7 +122,8 @@ impl WorkerSelectionStage {
|
|||||||
fn select_single_worker(
|
fn select_single_worker(
|
||||||
&self,
|
&self,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
info: &SelectWorkerInfo,
|
text: Option<&str>,
|
||||||
|
headers: Option<&http::HeaderMap>,
|
||||||
) -> Option<Arc<dyn Worker>> {
|
) -> Option<Arc<dyn Worker>> {
|
||||||
// Get workers for the specified model, filtered by connection mode
|
// Get workers for the specified model, filtered by connection mode
|
||||||
let workers = self.worker_registry.get_workers_filtered(
|
let workers = self.worker_registry.get_workers_filtered(
|
||||||
@@ -149,7 +149,13 @@ impl WorkerSelectionStage {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Select worker using the policy
|
// Select worker using the policy
|
||||||
let idx = policy.select_worker(&available, info)?;
|
let idx = policy.select_worker(
|
||||||
|
&available,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: text,
|
||||||
|
headers,
|
||||||
|
},
|
||||||
|
)?;
|
||||||
let selected = available[idx].clone();
|
let selected = available[idx].clone();
|
||||||
|
|
||||||
// Record worker selection metric
|
// Record worker selection metric
|
||||||
@@ -166,7 +172,8 @@ impl WorkerSelectionStage {
|
|||||||
fn select_pd_pair(
|
fn select_pd_pair(
|
||||||
&self,
|
&self,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
info: &SelectWorkerInfo,
|
text: Option<&str>,
|
||||||
|
headers: Option<&http::HeaderMap>,
|
||||||
) -> Option<(Arc<dyn Worker>, Arc<dyn Worker>)> {
|
) -> Option<(Arc<dyn Worker>, Arc<dyn Worker>)> {
|
||||||
let all_workers = self.worker_registry.get_workers_filtered(
|
let all_workers = self.worker_registry.get_workers_filtered(
|
||||||
model_id,
|
model_id,
|
||||||
@@ -206,8 +213,12 @@ impl WorkerSelectionStage {
|
|||||||
None => self.policy_registry.get_default_policy(),
|
None => self.policy_registry.get_default_policy(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let prefill_idx = policy.select_worker(&available_prefill, info)?;
|
let info = SelectWorkerInfo {
|
||||||
let decode_idx = policy.select_worker(&available_decode, info)?;
|
request_text: text,
|
||||||
|
headers,
|
||||||
|
};
|
||||||
|
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();
|
||||||
|
|||||||
@@ -94,9 +94,6 @@ pub struct PreparationOutput {
|
|||||||
/// Original text (for chat) or resolved text (for generate)
|
/// Original text (for chat) or resolved text (for generate)
|
||||||
pub original_text: Option<String>,
|
pub original_text: Option<String>,
|
||||||
|
|
||||||
/// Routing ID for manual routing policy
|
|
||||||
pub routing_id: Option<String>,
|
|
||||||
|
|
||||||
/// Tokenized input
|
/// Tokenized input
|
||||||
pub token_ids: Vec<u32>,
|
pub token_ids: Vec<u32>,
|
||||||
|
|
||||||
|
|||||||
@@ -19,7 +19,6 @@ use crate::{
|
|||||||
context::{PreparationOutput, RequestContext, RequestType},
|
context::{PreparationOutput, RequestContext, RequestType},
|
||||||
utils,
|
utils,
|
||||||
},
|
},
|
||||||
header_utils,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -124,7 +123,6 @@ impl HarmonyPreparationStage {
|
|||||||
// Step 4: Store results
|
// Step 4: Store results
|
||||||
ctx.state.preparation = Some(PreparationOutput {
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
original_text: None,
|
original_text: None,
|
||||||
routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
|
|
||||||
token_ids: build_output.input_ids,
|
token_ids: build_output.input_ids,
|
||||||
processed_messages: None,
|
processed_messages: None,
|
||||||
tool_constraints,
|
tool_constraints,
|
||||||
@@ -205,7 +203,6 @@ impl HarmonyPreparationStage {
|
|||||||
// Step 4: Store results with constraint
|
// Step 4: Store results with constraint
|
||||||
ctx.state.preparation = Some(PreparationOutput {
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
original_text: None,
|
original_text: None,
|
||||||
routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
|
|
||||||
token_ids: build_output.input_ids,
|
token_ids: build_output.input_ids,
|
||||||
processed_messages: None,
|
processed_messages: None,
|
||||||
tool_constraints: constraint,
|
tool_constraints: constraint,
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ use crate::{
|
|||||||
context::{PreparationOutput, RequestContext},
|
context::{PreparationOutput, RequestContext},
|
||||||
utils,
|
utils,
|
||||||
},
|
},
|
||||||
header_utils,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -97,7 +96,6 @@ impl ChatPreparationStage {
|
|||||||
// Store results in context
|
// Store results in context
|
||||||
ctx.state.preparation = Some(PreparationOutput {
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
original_text: Some(processed_messages.text.clone()),
|
original_text: Some(processed_messages.text.clone()),
|
||||||
routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
|
|
||||||
token_ids,
|
token_ids,
|
||||||
processed_messages: Some(processed_messages),
|
processed_messages: Some(processed_messages),
|
||||||
tool_constraints: tool_call_constraint,
|
tool_constraints: tool_call_constraint,
|
||||||
|
|||||||
@@ -13,7 +13,6 @@ use crate::{
|
|||||||
context::{PreparationOutput, RequestContext, RequestType},
|
context::{PreparationOutput, RequestContext, RequestType},
|
||||||
utils,
|
utils,
|
||||||
},
|
},
|
||||||
header_utils,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -48,9 +47,8 @@ impl PipelineStage for EmbeddingPreparationStage {
|
|||||||
));
|
));
|
||||||
};
|
};
|
||||||
|
|
||||||
// Extract text from request before borrowing ctx mutably
|
// Extract text from request
|
||||||
let text = request.extract_text_for_routing();
|
let text = request.extract_text_for_routing();
|
||||||
let routing_id = header_utils::extract_routing_id(ctx.input.headers.as_ref());
|
|
||||||
if text.is_empty() {
|
if text.is_empty() {
|
||||||
return Err(error::bad_request(
|
return Err(error::bad_request(
|
||||||
"empty_input",
|
"empty_input",
|
||||||
@@ -79,7 +77,6 @@ impl PipelineStage for EmbeddingPreparationStage {
|
|||||||
// Store preparation output
|
// Store preparation output
|
||||||
ctx.state.preparation = Some(PreparationOutput {
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
original_text: Some(text),
|
original_text: Some(text),
|
||||||
routing_id,
|
|
||||||
token_ids,
|
token_ids,
|
||||||
processed_messages: None,
|
processed_messages: None,
|
||||||
tool_constraints: None,
|
tool_constraints: None,
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ use crate::{
|
|||||||
context::{PreparationOutput, RequestContext},
|
context::{PreparationOutput, RequestContext},
|
||||||
utils,
|
utils,
|
||||||
},
|
},
|
||||||
header_utils,
|
|
||||||
},
|
},
|
||||||
tokenizer::traits::Tokenizer,
|
tokenizer::traits::Tokenizer,
|
||||||
};
|
};
|
||||||
@@ -69,7 +68,6 @@ impl GeneratePreparationStage {
|
|||||||
|
|
||||||
ctx.state.preparation = Some(PreparationOutput {
|
ctx.state.preparation = Some(PreparationOutput {
|
||||||
original_text,
|
original_text,
|
||||||
routing_id: header_utils::extract_routing_id(ctx.input.headers.as_ref()),
|
|
||||||
token_ids,
|
token_ids,
|
||||||
processed_messages: None,
|
processed_messages: None,
|
||||||
tool_constraints: None,
|
tool_constraints: None,
|
||||||
|
|||||||
@@ -157,17 +157,6 @@ pub fn apply_provider_headers(
|
|||||||
req
|
req
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Header name for routing key used by manual routing policy
|
|
||||||
pub const ROUTING_KEY_HEADER: &str = "X-SMG-Routing-Key";
|
|
||||||
|
|
||||||
/// Extract routing ID from HTTP headers for manual routing policy
|
|
||||||
pub fn extract_routing_id(headers: Option<&HeaderMap>) -> Option<String> {
|
|
||||||
headers
|
|
||||||
.and_then(|h| h.get(ROUTING_KEY_HEADER))
|
|
||||||
.and_then(|v| v.to_str().ok())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract auth header with passthrough semantics.
|
/// Extract auth header with passthrough semantics.
|
||||||
///
|
///
|
||||||
/// Passthrough mode: User's Authorization header takes priority.
|
/// Passthrough mode: User's Authorization header takes priority.
|
||||||
@@ -194,31 +183,3 @@ pub fn extract_auth_header(
|
|||||||
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
|
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_routing_id_with_header() {
|
|
||||||
let mut headers = HeaderMap::new();
|
|
||||||
headers.insert(ROUTING_KEY_HEADER, HeaderValue::from_static("user-123"));
|
|
||||||
let result = extract_routing_id(Some(&headers));
|
|
||||||
assert_eq!(result, Some("user-123".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_routing_id_without_header() {
|
|
||||||
let headers = HeaderMap::new();
|
|
||||||
let result = extract_routing_id(Some(&headers));
|
|
||||||
assert_eq!(result, None);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_extract_routing_id_empty_value() {
|
|
||||||
let mut headers = HeaderMap::new();
|
|
||||||
headers.insert(ROUTING_KEY_HEADER, HeaderValue::from_static(""));
|
|
||||||
let result = extract_routing_id(Some(&headers));
|
|
||||||
assert_eq!(result, Some("".to_string()));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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},
|
||||||
@@ -58,8 +58,8 @@ struct PDRequestContext<'a> {
|
|||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
return_logprob: bool,
|
return_logprob: bool,
|
||||||
request_text: Option<String>,
|
request_text: Option<String>,
|
||||||
routing_id: Option<String>,
|
|
||||||
model_id: Option<&'a str>,
|
model_id: Option<&'a str>,
|
||||||
|
headers: Option<HeaderMap>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl PDRouter {
|
impl PDRouter {
|
||||||
@@ -306,8 +306,8 @@ impl PDRouter {
|
|||||||
let (prefill, decode) = match self
|
let (prefill, decode) = match self
|
||||||
.select_pd_pair(
|
.select_pd_pair(
|
||||||
context.request_text.as_deref(),
|
context.request_text.as_deref(),
|
||||||
context.routing_id.as_deref(),
|
|
||||||
context.model_id,
|
context.model_id,
|
||||||
|
context.headers.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
@@ -696,8 +696,8 @@ impl PDRouter {
|
|||||||
async fn select_pd_pair(
|
async fn select_pd_pair(
|
||||||
&self,
|
&self,
|
||||||
request_text: Option<&str>,
|
request_text: Option<&str>,
|
||||||
routing_id: Option<&str>,
|
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
|
headers: Option<&HeaderMap>,
|
||||||
) -> Result<(Arc<dyn Worker>, Arc<dyn Worker>), String> {
|
) -> Result<(Arc<dyn Worker>, Arc<dyn Worker>), String> {
|
||||||
let effective_model_id = if !self.enable_igw { None } else { model_id };
|
let effective_model_id = if !self.enable_igw { None } else { model_id };
|
||||||
|
|
||||||
@@ -731,16 +731,21 @@ impl PDRouter {
|
|||||||
let prefill_policy = self.policy_registry.get_prefill_policy();
|
let prefill_policy = self.policy_registry.get_prefill_policy();
|
||||||
let decode_policy = self.policy_registry.get_decode_policy();
|
let decode_policy = self.policy_registry.get_decode_policy();
|
||||||
|
|
||||||
let info = crate::policies::SelectWorkerInfo {
|
let prefill = Self::pick_worker_by_policy_arc(
|
||||||
|
&prefill_workers,
|
||||||
|
&*prefill_policy,
|
||||||
request_text,
|
request_text,
|
||||||
routing_id,
|
headers,
|
||||||
};
|
"prefill",
|
||||||
|
)?;
|
||||||
|
|
||||||
let prefill =
|
let decode = Self::pick_worker_by_policy_arc(
|
||||||
Self::pick_worker_by_policy_arc(&prefill_workers, &*prefill_policy, &info, "prefill")?;
|
&decode_workers,
|
||||||
|
&*decode_policy,
|
||||||
let decode =
|
request_text,
|
||||||
Self::pick_worker_by_policy_arc(&decode_workers, &*decode_policy, &info, "decode")?;
|
headers,
|
||||||
|
"decode",
|
||||||
|
)?;
|
||||||
|
|
||||||
// Record worker selection metrics (Layer 3)
|
// Record worker selection metrics (Layer 3)
|
||||||
let model = model_id.unwrap_or("default");
|
let model = model_id.unwrap_or("default");
|
||||||
@@ -763,7 +768,8 @@ impl PDRouter {
|
|||||||
fn pick_worker_by_policy_arc(
|
fn pick_worker_by_policy_arc(
|
||||||
workers: &[Arc<dyn Worker>],
|
workers: &[Arc<dyn Worker>],
|
||||||
policy: &dyn LoadBalancingPolicy,
|
policy: &dyn LoadBalancingPolicy,
|
||||||
info: &crate::policies::SelectWorkerInfo,
|
request_text: Option<&str>,
|
||||||
|
headers: Option<&HeaderMap>,
|
||||||
worker_type: &str,
|
worker_type: &str,
|
||||||
) -> Result<Arc<dyn Worker>, String> {
|
) -> Result<Arc<dyn Worker>, String> {
|
||||||
if workers.is_empty() {
|
if workers.is_empty() {
|
||||||
@@ -787,7 +793,13 @@ impl PDRouter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let selected_idx = policy
|
let selected_idx = policy
|
||||||
.select_worker(&available_workers, info)
|
.select_worker(
|
||||||
|
&available_workers,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text,
|
||||||
|
headers,
|
||||||
|
},
|
||||||
|
)
|
||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
format!(
|
format!(
|
||||||
"Policy {} failed to select a {} worker",
|
"Policy {} failed to select a {} worker",
|
||||||
@@ -1245,8 +1257,8 @@ impl RouterTrait for PDRouter {
|
|||||||
is_stream,
|
is_stream,
|
||||||
return_logprob,
|
return_logprob,
|
||||||
request_text,
|
request_text,
|
||||||
routing_id: header_utils::extract_routing_id(headers),
|
|
||||||
model_id,
|
model_id,
|
||||||
|
headers: headers.cloned(),
|
||||||
};
|
};
|
||||||
|
|
||||||
self.execute_dual_dispatch(headers, body, context).await
|
self.execute_dual_dispatch(headers, body, context).await
|
||||||
@@ -1287,8 +1299,8 @@ impl RouterTrait for PDRouter {
|
|||||||
is_stream,
|
is_stream,
|
||||||
return_logprob,
|
return_logprob,
|
||||||
request_text,
|
request_text,
|
||||||
routing_id: header_utils::extract_routing_id(headers),
|
|
||||||
model_id,
|
model_id,
|
||||||
|
headers: headers.cloned(),
|
||||||
};
|
};
|
||||||
|
|
||||||
self.execute_dual_dispatch(headers, body, context).await
|
self.execute_dual_dispatch(headers, body, context).await
|
||||||
@@ -1321,8 +1333,8 @@ impl RouterTrait for PDRouter {
|
|||||||
is_stream,
|
is_stream,
|
||||||
return_logprob,
|
return_logprob,
|
||||||
request_text,
|
request_text,
|
||||||
routing_id: header_utils::extract_routing_id(headers),
|
|
||||||
model_id,
|
model_id,
|
||||||
|
headers: headers.cloned(),
|
||||||
};
|
};
|
||||||
|
|
||||||
self.execute_dual_dispatch(headers, body, context).await
|
self.execute_dual_dispatch(headers, body, context).await
|
||||||
@@ -1334,6 +1346,7 @@ impl RouterTrait for PDRouter {
|
|||||||
body: &RerankRequest,
|
body: &RerankRequest,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
|
// Extract text for cache-aware routing
|
||||||
let req_text = if self.policies_need_request_text() {
|
let req_text = if self.policies_need_request_text() {
|
||||||
Some(body.query.clone())
|
Some(body.query.clone())
|
||||||
} else {
|
} else {
|
||||||
@@ -1346,8 +1359,8 @@ impl RouterTrait for PDRouter {
|
|||||||
is_stream: false,
|
is_stream: false,
|
||||||
return_logprob: false,
|
return_logprob: false,
|
||||||
request_text: req_text,
|
request_text: req_text,
|
||||||
routing_id: header_utils::extract_routing_id(headers),
|
|
||||||
model_id,
|
model_id,
|
||||||
|
headers: headers.cloned(),
|
||||||
};
|
};
|
||||||
|
|
||||||
self.execute_dual_dispatch(headers, body, context).await
|
self.execute_dual_dispatch(headers, body, context).await
|
||||||
|
|||||||
@@ -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,
|
||||||
@@ -39,7 +39,7 @@ use crate::{
|
|||||||
responses::{ResponsesGetParams, ResponsesRequest},
|
responses::{ResponsesGetParams, ResponsesRequest},
|
||||||
},
|
},
|
||||||
routers::{
|
routers::{
|
||||||
error,
|
error::{self, extract_error_code_from_response},
|
||||||
grpc::utils::{error_type_from_status, route_to_endpoint},
|
grpc::utils::{error_type_from_status, route_to_endpoint},
|
||||||
header_utils, RouterTrait,
|
header_utils, RouterTrait,
|
||||||
},
|
},
|
||||||
@@ -140,7 +140,8 @@ impl Router {
|
|||||||
fn select_worker_for_model(
|
fn select_worker_for_model(
|
||||||
&self,
|
&self,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
info: &crate::policies::SelectWorkerInfo,
|
text: Option<&str>,
|
||||||
|
headers: Option<&HeaderMap>,
|
||||||
) -> Option<Arc<dyn Worker>> {
|
) -> Option<Arc<dyn Worker>> {
|
||||||
let effective_model_id = if !self.enable_igw { None } else { model_id };
|
let effective_model_id = if !self.enable_igw { None } else { model_id };
|
||||||
|
|
||||||
@@ -168,7 +169,13 @@ impl Router {
|
|||||||
None => self.policy_registry.get_default_policy(),
|
None => self.policy_registry.get_default_policy(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let idx = policy.select_worker(&available, info)?;
|
let idx = policy.select_worker(
|
||||||
|
&available,
|
||||||
|
&SelectWorkerInfo {
|
||||||
|
request_text: text,
|
||||||
|
headers,
|
||||||
|
},
|
||||||
|
)?;
|
||||||
|
|
||||||
// Record worker selection metric (Layer 3)
|
// Record worker selection metric (Layer 3)
|
||||||
Metrics::record_worker_selection(
|
Metrics::record_worker_selection(
|
||||||
@@ -191,11 +198,6 @@ impl Router {
|
|||||||
let start = Instant::now();
|
let start = Instant::now();
|
||||||
let is_stream = typed_req.is_stream();
|
let is_stream = typed_req.is_stream();
|
||||||
let text = typed_req.extract_text_for_routing();
|
let text = typed_req.extract_text_for_routing();
|
||||||
let routing_id = header_utils::extract_routing_id(headers);
|
|
||||||
let info = crate::policies::SelectWorkerInfo {
|
|
||||||
request_text: Some(&text),
|
|
||||||
routing_id: routing_id.as_deref(),
|
|
||||||
};
|
|
||||||
let model = model_id.unwrap_or("default");
|
let model = model_id.unwrap_or("default");
|
||||||
let endpoint = route_to_endpoint(route);
|
let endpoint = route_to_endpoint(route);
|
||||||
|
|
||||||
@@ -214,7 +216,7 @@ impl Router {
|
|||||||
// operation per attempt
|
// operation per attempt
|
||||||
|_: u32| async {
|
|_: u32| async {
|
||||||
let res = self
|
let res = self
|
||||||
.route_typed_request_once(headers, typed_req, route, model_id, is_stream, &info)
|
.route_typed_request_once(headers, typed_req, route, model_id, is_stream, &text)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
// Need to be outside `route_typed_request_once` because that function has multiple return paths
|
// Need to be outside `route_typed_request_once` because that function has multiple return paths
|
||||||
@@ -272,9 +274,9 @@ impl Router {
|
|||||||
route: &'static str,
|
route: &'static str,
|
||||||
model_id: Option<&str>,
|
model_id: Option<&str>,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
info: &crate::policies::SelectWorkerInfo<'_>,
|
text: &str,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
let worker = match self.select_worker_for_model(model_id, info) {
|
let worker = match self.select_worker_for_model(model_id, Some(text), headers) {
|
||||||
Some(w) => w,
|
Some(w) => w,
|
||||||
None => {
|
None => {
|
||||||
return error::service_unavailable(
|
return error::service_unavailable(
|
||||||
@@ -696,8 +698,6 @@ fn convert_reqwest_error(e: reqwest::Error) -> Response {
|
|||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
|
||||||
use crate::routers::error::extract_error_code_from_response;
|
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
impl RouterTrait for Router {
|
impl RouterTrait for Router {
|
||||||
fn as_any(&self) -> &dyn std::any::Any {
|
fn as_any(&self) -> &dyn std::any::Any {
|
||||||
|
|||||||
@@ -103,24 +103,16 @@ 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(
|
let info = SelectWorkerInfo {
|
||||||
&default_workers,
|
request_text: Some("test request"),
|
||||||
&SelectWorkerInfo {
|
..Default::default()
|
||||||
request_text: Some("test request"),
|
};
|
||||||
..Default::default()
|
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(
|
let selected = policy.select_worker(&llama_workers, &info);
|
||||||
&llama_workers,
|
|
||||||
&SelectWorkerInfo {
|
|
||||||
request_text: Some("test request"),
|
|
||||||
..Default::default()
|
|
||||||
},
|
|
||||||
);
|
|
||||||
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![
|
||||||
@@ -129,13 +121,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(
|
let selected = policy.select_worker(&all_workers, &info);
|
||||||
&all_workers,
|
|
||||||
&SelectWorkerInfo {
|
|
||||||
request_text: Some("test request"),
|
|
||||||
..Default::default()
|
|
||||||
},
|
|
||||||
);
|
|
||||||
assert!(selected.is_some(), "Should select from all workers");
|
assert!(selected.is_some(), "Should select from all workers");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user