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