[model-gateway] extract header extraction in policy and add (#16566)
This commit is contained in:
@@ -18,16 +18,14 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use http::header::HeaderName;
|
||||
use rand::Rng as _;
|
||||
|
||||
use super::{LoadBalancingPolicy, SelectWorkerInfo};
|
||||
use crate::{core::Worker, observability::metrics::Metrics};
|
||||
|
||||
/// Header for direct worker targeting by index (0-based)
|
||||
static HEADER_TARGET_WORKER: HeaderName = HeaderName::from_static("x-smg-target-worker");
|
||||
/// Header for consistent hash routing
|
||||
static HEADER_ROUTING_KEY: HeaderName = HeaderName::from_static("x-smg-routing-key");
|
||||
use crate::{
|
||||
core::Worker,
|
||||
observability::metrics::Metrics,
|
||||
routers::header_utils::{extract_routing_key, extract_target_worker},
|
||||
};
|
||||
|
||||
/// Execution branch for metrics
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
@@ -112,18 +110,8 @@ impl ConsistentHashingPolicy {
|
||||
return (None, Branch::NoHealthyWorkers);
|
||||
}
|
||||
|
||||
// Extract routing headers - to_str() is O(1), just validates ASCII, no allocation
|
||||
let target_worker = info
|
||||
.headers
|
||||
.and_then(|h| h.get(&HEADER_TARGET_WORKER))
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.filter(|s| !s.is_empty());
|
||||
|
||||
let routing_key = info
|
||||
.headers
|
||||
.and_then(|h| h.get(&HEADER_ROUTING_KEY))
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.filter(|s| !s.is_empty());
|
||||
let target_worker = extract_target_worker(info.headers);
|
||||
let routing_key = extract_routing_key(info.headers);
|
||||
|
||||
// Priority 1: X-SMG-Target-Worker - direct routing by worker index
|
||||
// O(1) parse + O(1) bounds check + O(1) health check
|
||||
|
||||
@@ -16,17 +16,15 @@
|
||||
use std::{sync::Arc, time::Instant};
|
||||
|
||||
use dashmap::{mapref::entry::Entry, DashMap};
|
||||
use http::header::HeaderName;
|
||||
use rand::Rng;
|
||||
use tracing::info;
|
||||
|
||||
use super::{
|
||||
get_healthy_worker_indices, utils::PeriodicTask, LoadBalancingPolicy, SelectWorkerInfo,
|
||||
};
|
||||
use crate::{core::Worker, observability::metrics::Metrics};
|
||||
|
||||
/// Header for routing key based sticky sessions
|
||||
static HEADER_ROUTING_KEY: HeaderName = HeaderName::from_static("x-smg-routing-key");
|
||||
use crate::{
|
||||
core::Worker, observability::metrics::Metrics, routers::header_utils::extract_routing_key,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ExecutionBranch {
|
||||
@@ -191,12 +189,7 @@ impl ManualPolicy {
|
||||
return (None, ExecutionBranch::NoHealthyWorkers);
|
||||
}
|
||||
|
||||
// Extract routing key from header
|
||||
let routing_id = info
|
||||
.headers
|
||||
.and_then(|h| h.get(&HEADER_ROUTING_KEY))
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.filter(|s| !s.is_empty());
|
||||
let routing_id = extract_routing_key(info.headers);
|
||||
|
||||
if let Some(routing_id) = routing_id {
|
||||
let (idx, branch) = self.select_by_routing_id(workers, routing_id, &healthy_indices);
|
||||
|
||||
@@ -3,6 +3,25 @@ use axum::{
|
||||
extract::Request,
|
||||
http::{HeaderMap, HeaderValue},
|
||||
};
|
||||
use http::header::HeaderName;
|
||||
|
||||
static HEADER_TARGET_WORKER: HeaderName = HeaderName::from_static("x-smg-target-worker");
|
||||
static HEADER_ROUTING_KEY: HeaderName = HeaderName::from_static("x-smg-routing-key");
|
||||
|
||||
fn extract_header_value<'a>(headers: Option<&'a HeaderMap>, name: &HeaderName) -> Option<&'a str> {
|
||||
headers
|
||||
.and_then(|h| h.get(name))
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.filter(|s| !s.is_empty())
|
||||
}
|
||||
|
||||
pub fn extract_target_worker(headers: Option<&HeaderMap>) -> Option<&str> {
|
||||
extract_header_value(headers, &HEADER_TARGET_WORKER)
|
||||
}
|
||||
|
||||
pub fn extract_routing_key(headers: Option<&HeaderMap>) -> Option<&str> {
|
||||
extract_header_value(headers, &HEADER_ROUTING_KEY)
|
||||
}
|
||||
|
||||
/// Copy request headers to a Vec of name-value string pairs
|
||||
/// Used for forwarding headers to backend workers
|
||||
@@ -203,6 +222,44 @@ pub fn should_forward_request_header(name: &str) -> bool {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_extract_header_value_returns_value() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-smg-routing-key", "test-key".parse().unwrap());
|
||||
assert_eq!(extract_routing_key(Some(&headers)), Some("test-key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_header_value_returns_none_for_missing() {
|
||||
let headers = HeaderMap::new();
|
||||
assert_eq!(extract_routing_key(Some(&headers)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_header_value_returns_none_for_empty() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-smg-routing-key", "".parse().unwrap());
|
||||
assert_eq!(extract_routing_key(Some(&headers)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_header_value_returns_none_for_none_headers() {
|
||||
assert_eq!(extract_routing_key(None), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_target_worker() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-smg-target-worker", "2".parse().unwrap());
|
||||
assert_eq!(extract_target_worker(Some(&headers)), Some("2"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_target_worker_missing() {
|
||||
let headers = HeaderMap::new();
|
||||
assert_eq!(extract_target_worker(Some(&headers)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_forward_request_header_whitelist() {
|
||||
assert!(should_forward_request_header("authorization"));
|
||||
|
||||
Reference in New Issue
Block a user