[model-gateway] extract header extraction in policy and add (#16566)

This commit is contained in:
fzyzcjy
2026-01-07 12:18:44 +08:00
committed by GitHub
parent 38895a0064
commit 3be1e734ee
3 changed files with 68 additions and 30 deletions

View File

@@ -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

View File

@@ -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);

View File

@@ -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"));