Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 24 additions & 12 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -368,6 +368,7 @@ struct Router {
host: String,
port: u16,
health_check_port: Option<u16>,
routing_key_override: bool,
worker_urls: Vec<String>,
policy: PolicyType,
worker_startup_timeout_secs: u64,
Expand Down Expand Up @@ -504,6 +505,19 @@ impl Router {
})
}

fn parse_assignment_mode(&self) -> Result<config::ManualAssignmentMode, config::ConfigError> {
match self.assignment_mode.as_str() {
"random" => Ok(config::ManualAssignmentMode::Random),
"min_load" => Ok(config::ManualAssignmentMode::MinLoad),
"min_group" => Ok(config::ManualAssignmentMode::MinGroup),
other => Err(config::ConfigError::InvalidValue {
field: "assignment_mode".to_string(),
value: other.to_string(),
reason: "expected 'random', 'min_load', or 'min_group'".to_string(),
}),
}
}

pub fn to_router_config(&self) -> config::ConfigResult<config::RouterConfig> {
use config::{
DiscoveryConfig, MetricsConfig, PolicyConfig as ConfigPolicyConfig, RoutingMode,
Expand Down Expand Up @@ -541,18 +555,7 @@ impl Router {
PolicyType::Manual => ConfigPolicyConfig::Manual {
eviction_interval_secs: self.eviction_interval_secs,
max_idle_secs: self.max_idle_secs,
assignment_mode: match self.assignment_mode.as_str() {
"random" => config::ManualAssignmentMode::Random,
"min_load" => config::ManualAssignmentMode::MinLoad,
"min_group" => config::ManualAssignmentMode::MinGroup,
other => {
return Err(config::ConfigError::InvalidValue {
field: "assignment_mode".to_string(),
value: other.to_string(),
reason: "expected 'random', 'min_load', or 'min_group'".to_string(),
});
}
},
assignment_mode: self.parse_assignment_mode()?,
},
PolicyType::ConsistentHashing => ConfigPolicyConfig::ConsistentHashing,
PolicyType::PrefixHash => ConfigPolicyConfig::PrefixHash {
Expand Down Expand Up @@ -756,6 +759,12 @@ impl Router {
.maybe_storage_hook_wasm_path(self.storage_hook_wasm_path.as_deref())
.enable_wasm(self.enable_wasm)
.dp_aware(self.dp_aware)
.routing_key_override(config::RoutingKeyOverrideConfig {
enabled: self.routing_key_override,
eviction_interval_secs: self.eviction_interval_secs,
max_idle_secs: self.max_idle_secs,
assignment_mode: self.parse_assignment_mode()?,
})
Comment thread
coderabbitai[bot] marked this conversation as resolved.
.retries(!self.disable_retries)
.circuit_breaker(!self.disable_circuit_breaker)
.igw(self.enable_igw)
Expand Down Expand Up @@ -890,6 +899,7 @@ impl Router {
// positional argument keeps its index for callers that construct
// `_Router(...)` positionally. See the struct-field note above.
health_check_port = None,
routing_key_override = false,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -1009,6 +1019,7 @@ impl Router {
// Appended last to match the `#[pyo3(signature)]` order above and
// preserve positional-argument compatibility.
health_check_port: Option<u16>,
routing_key_override: bool,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand All @@ -1028,6 +1039,7 @@ impl Router {
host,
port,
health_check_port,
routing_key_override,
worker_urls,
policy,
worker_startup_timeout_secs,
Expand Down
6 changes: 6 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ class RouterArgs:
max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches
bucket_adjust_interval_secs: int = 5
dp_aware: bool = False
routing_key_override: bool = False
dp_minimum_tokens_scheduler: bool = False
enable_igw: bool = False # Enable IGW (Inter-Gateway) mode for multi-model support
api_key: str | None = None
Expand Down Expand Up @@ -473,6 +474,11 @@ def add_cli_args(
action="store_true",
help="Enable data parallelism aware schedule",
)
routing_group.add_argument(
f"--{prefix}routing-key-override",
action="store_true",
help="Honor X-SMG-Routing-Key for sticky routing on any policy",
)
routing_group.add_argument(
f"--{prefix}dp-minimum-tokens-scheduler",
action="store_true",
Expand Down
5 changes: 4 additions & 1 deletion model_gateway/src/app_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -526,7 +526,10 @@ impl AppContextBuilder {

/// Create policy registry
fn with_policy_registry(mut self, config: &RouterConfig) -> Self {
self.policy_registry = Some(Arc::new(PolicyRegistry::new(config.policy.clone())));
self.policy_registry = Some(Arc::new(PolicyRegistry::with_override(
config.policy.clone(),
config.routing_key_override.clone(),
)));
self
}

Expand Down
8 changes: 7 additions & 1 deletion model_gateway/src/config/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,8 @@ use smg_mcp::McpConfig;
use super::{
CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig,
HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, PostgresConfig, RedisConfig,
RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, TraceConfig,
RetryConfig, RouterConfig, RoutingKeyOverrideConfig, RoutingMode, TokenizerCacheConfig,
TraceConfig,
};
use crate::worker::ConnectionMode;

Expand Down Expand Up @@ -526,6 +527,11 @@ impl RouterConfigBuilder {
self
}

pub fn routing_key_override(mut self, config: RoutingKeyOverrideConfig) -> Self {
self.config.routing_key_override = config;
self
}

/// Inverse of disable_retries field
pub fn retries(mut self, enable: bool) -> Self {
self.config.disable_retries = !enable;
Expand Down
32 changes: 32 additions & 0 deletions model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ pub struct RouterConfig {
#[serde(default)]
pub connection_mode: ConnectionMode,
pub policy: PolicyConfig,
/// Per-request sticky-routing override (honors `X-SMG-Routing-Key`).
#[serde(default)]
pub routing_key_override: RoutingKeyOverrideConfig,
pub host: String,
pub port: u16,
/// Dedicated port for the isolated Kubernetes liveness/readiness/health
Expand Down Expand Up @@ -303,6 +306,34 @@ pub enum ManualAssignmentMode {
MinGroup,
}

/// Per-request sticky-routing override: when `X-SMG-Routing-Key` is present, any
/// eligible policy routes via manual sticky-map semantics. Reuses the manual
/// policy knobs for the sticky map; eviction defaults match the manual policy so
/// config-file users with only `enabled: true` still get TTL eviction (no leak).
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingKeyOverrideConfig {
/// When false, policies are used unchanged.
#[serde(default)]
pub enabled: bool,
#[serde(default = "default_manual_eviction_interval_secs")]
pub eviction_interval_secs: u64,
#[serde(default = "default_manual_max_idle_secs")]
pub max_idle_secs: u64,
#[serde(default)]
pub assignment_mode: ManualAssignmentMode,
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

impl Default for RoutingKeyOverrideConfig {
fn default() -> Self {
Self {
enabled: false,
eviction_interval_secs: default_manual_eviction_interval_secs(),
max_idle_secs: default_manual_max_idle_secs(),
assignment_mode: ManualAssignmentMode::default(),
}
}
}

/// Policy configuration for routing
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
Expand Down Expand Up @@ -659,6 +690,7 @@ impl Default for RouterConfig {
worker_urls: vec![],
},
policy: PolicyConfig::Random,
routing_key_override: RoutingKeyOverrideConfig::default(),
host: "0.0.0.0".to_string(),
port: 3001,
health_check_port: None,
Expand Down
37 changes: 26 additions & 11 deletions model_gateway/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use smg::{
validate_mesh_server_name, CircuitBreakerConfig, ConfigError, ConfigResult,
DiscoveryConfig, HealthCheckConfig, HistoryBackend, ManualAssignmentMode, MetricsConfig,
OracleConfig, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig,
RoutingMode, SchemaConfig, TokenizerCacheConfig, TraceConfig,
RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, TokenizerCacheConfig, TraceConfig,
},
observability::{
metrics::PrometheusConfig,
Expand Down Expand Up @@ -239,6 +239,11 @@ struct CliArgs {
#[arg(long, default_value_t = false, help_heading = "Routing Policy")]
dp_aware: bool,

/// Honor X-SMG-Routing-Key for sticky routing on any policy (reuses the
/// manual eviction/idle/assignment knobs for the sticky map)
#[arg(long, default_value_t = false, help_heading = "Routing Policy")]
routing_key_override: bool,

/// Enable IGW (Inference Gateway) mode for multi-model support
#[arg(long, default_value_t = false, help_heading = "Routing Policy")]
enable_igw: bool,
Expand Down Expand Up @@ -965,10 +970,6 @@ impl CliArgs {
}))
}

#[expect(
clippy::panic,
reason = "unreachable: clap value_parser restricts valid assignment modes"
)]
fn parse_policy(&self, policy_str: &str) -> PolicyConfig {
match policy_str {
"random" => PolicyConfig::Random,
Expand Down Expand Up @@ -1000,17 +1001,25 @@ impl CliArgs {
"manual" => PolicyConfig::Manual {
eviction_interval_secs: self.eviction_interval,
max_idle_secs: self.max_idle_secs,
assignment_mode: match self.assignment_mode.as_str() {
"random" => ManualAssignmentMode::Random,
"min_load" => ManualAssignmentMode::MinLoad,
"min_group" => ManualAssignmentMode::MinGroup,
other => panic!("Unknown assignment mode: {other}"),
},
assignment_mode: Self::parse_assignment_mode(&self.assignment_mode),
},
_ => PolicyConfig::RoundRobin,
}
}

#[expect(
clippy::panic,
reason = "unreachable: clap value_parser restricts valid assignment modes"
)]
fn parse_assignment_mode(mode: &str) -> ManualAssignmentMode {
match mode {
"random" => ManualAssignmentMode::Random,
"min_load" => ManualAssignmentMode::MinLoad,
"min_group" => ManualAssignmentMode::MinGroup,
other => panic!("Unknown assignment mode: {other}"),
}
}

fn load_schema_config(&self) -> ConfigResult<Option<SchemaConfig>> {
match &self.schema_config {
Some(path) => {
Expand Down Expand Up @@ -1340,6 +1349,12 @@ impl CliArgs {
.maybe_tool_call_parser(self.tool_call_parser.as_ref())
.maybe_mcp_config_path(self.mcp_config_path.as_ref())
.dp_aware(self.dp_aware)
.routing_key_override(RoutingKeyOverrideConfig {
enabled: self.routing_key_override,
eviction_interval_secs: self.eviction_interval,
max_idle_secs: self.max_idle_secs,
assignment_mode: Self::parse_assignment_mode(&self.assignment_mode),
})
.retries(!self.disable_retries)
.circuit_breaker(!self.disable_circuit_breaker)
.enable_wasm(self.enable_wasm)
Expand Down
39 changes: 38 additions & 1 deletion model_gateway/src/policies/manual.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ use tracing::info;

use super::{
get_healthy_worker_indices, utils::PeriodicTask, LoadBalancingPolicy, SelectWorkerInfo,
WorkerLeg,
};
use crate::{
config::ManualAssignmentMode, observability::metrics::Metrics,
Expand Down Expand Up @@ -202,7 +203,15 @@ impl ManualPolicy {
}

if let Some(routing_id) = extract_routing_key(info.headers) {
let (idx, branch) = self.select_by_routing_id(workers, routing_id, &healthy_indices);
// Single is the common leg; route on the bare key to skip the
// per-request allocation. PD legs namespace so prefill and decode
// stick independently.
let (idx, branch) = if info.leg == WorkerLeg::Single {
self.select_by_routing_id(workers, routing_id, &healthy_indices)
} else {
let namespaced = format!("{}{}", info.leg.routing_id_prefix(), routing_id);
self.select_by_routing_id(workers, &namespaced, &healthy_indices)
};
return (Some(idx), branch);
}
Comment on lines 205 to 216

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

format!("{}{}", info.leg.routing_id_prefix(), routing_id) allocates a new String on the heap for every request.

Since WorkerLeg::Single is the default and most common case (where routing_id_prefix() is ""), we can completely avoid this heap allocation on the hot path by checking if info.leg == WorkerLeg::Single and using routing_id directly.

        if let Some(routing_id) = extract_routing_key(info.headers) {
            let (idx, branch) = if info.leg == crate::policies::WorkerLeg::Single {
                self.select_by_routing_id(workers, routing_id, &healthy_indices)
            } else {
                let namespaced = format!("{}{}", info.leg.routing_id_prefix(), routing_id);
                self.select_by_routing_id(workers, &namespaced, &healthy_indices)
            };
            return (Some(idx), branch);
        }
References
  1. Avoid heap allocations in hot or periodic paths. Keep allocation-heavy helper methods restricted to cold paths where the allocation overhead is negligible.


Expand Down Expand Up @@ -330,6 +339,34 @@ mod tests {
headers
}

#[test]
fn test_manual_leg_namespaces_sticky_entries() {
let policy = ManualPolicy::new();
let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]);
let headers = headers_with_routing_key("user-123");

let prefill = SelectWorkerInfo {
headers: Some(&headers),
leg: WorkerLeg::Prefill,
..Default::default()
};
let decode = SelectWorkerInfo {
headers: Some(&headers),
leg: WorkerLeg::Decode,
..Default::default()
};

let (p1, _) = policy.select_worker_impl(&workers, &prefill);
let (d1, _) = policy.select_worker_impl(&workers, &decode);

// Same key under two legs -> two independent sticky entries.
assert_eq!(policy.routing_map.len(), 2);
for _ in 0..5 {
assert_eq!(policy.select_worker_impl(&workers, &prefill).0, p1);
assert_eq!(policy.select_worker_impl(&workers, &decode).0, d1);
}
}

#[test]
fn test_manual_consistent_routing() {
let policy = ManualPolicy::new();
Expand Down
34 changes: 34 additions & 0 deletions model_gateway/src/policies/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -187,6 +187,28 @@ pub(crate) fn normalize_model_key(model_id: &str) -> &str {
}
}

/// Which PD leg a selection is for. `Single` is non-PD (the default) and keeps
/// routing-key stickiness byte-identical to pre-leg behavior.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum WorkerLeg {
#[default]
Single,
Prefill,
Decode,
}

impl WorkerLeg {
/// Prefix used to namespace sticky routing IDs per leg. `Single` is empty so
/// non-PD entries are unchanged.
pub fn routing_id_prefix(self) -> &'static str {
match self {
WorkerLeg::Single => "",
WorkerLeg::Prefill => "prefill:",
WorkerLeg::Decode => "decode:",
}
}
}

/// Information passed to policy for worker selection
#[derive(Debug, Clone, Default)]
pub struct SelectWorkerInfo<'a> {
Expand All @@ -203,6 +225,9 @@ pub struct SelectWorkerInfo<'a> {
/// Pre-computed hash ring for O(log n) consistent hashing
/// Built and cached by WorkerRegistry, passed through to avoid per-request rebuilds
pub hash_ring: Option<Arc<HashRing>>,
/// Which PD leg this selection is for (default `Single`); namespaces
/// header-based sticky routing so prefill and decode stick independently.
pub leg: WorkerLeg,
}

#[cfg(test)]
Expand Down Expand Up @@ -291,4 +316,13 @@ mod tests {
);
}
}

#[test]
fn test_select_worker_info_leg_defaults_to_single() {
let info = SelectWorkerInfo::default();
assert_eq!(info.leg, WorkerLeg::Single);
assert_eq!(WorkerLeg::Single.routing_id_prefix(), "");
assert_eq!(WorkerLeg::Prefill.routing_id_prefix(), "prefill:");
assert_eq!(WorkerLeg::Decode.routing_id_prefix(), "decode:");
}
}
Loading
Loading