diff --git a/model_gateway/src/policies/consistent_hashing.rs b/model_gateway/src/policies/consistent_hashing.rs index c8eecf2b58..a98447c80d 100644 --- a/model_gateway/src/policies/consistent_hashing.rs +++ b/model_gateway/src/policies/consistent_hashing.rs @@ -123,10 +123,13 @@ impl ConsistentHashingPolicy { } let target_worker = extract_target_worker(info.headers); - // Both sides apply the hint caps: an over-cap or non-UTF-8 key must not - // influence placement on any path. + // The rid-derived session key (populated only under the routing-key + // override, already capped and lineage-stripped) outranks the header. + // Both header sides apply the hint caps: an over-cap or non-UTF-8 key + // must not influence placement on any path. let routing_key = info - .routing_key + .rid_key + .or(info.routing_key) .or_else(|| extract_routing_key_hint(info.headers)); // Priority 1: X-SMG-Target-Worker - direct routing by worker index @@ -659,6 +662,46 @@ mod tests { assert_eq!(branch, Branch::RoutingKeyHit); } + #[test] + fn test_rid_key_outranks_routing_key_hint() { + let policy = ConsistentHashingPolicy::new(); + let workers = create_workers(&[ + "http://w1:8000", + "http://w2:8000", + "http://w3:8000", + "http://w4:8000", + ]); + let ring = Arc::new(HashRing::new(workers.iter().map(|w| w.url()))); + + let select_by_key = |key: &str| { + let info = SelectWorkerInfo { + routing_key: Some(key), + hash_ring: Some(ring.clone()), + ..Default::default() + }; + policy.select_worker_impl(&workers, &info).0.unwrap() + }; + + // Find a header key that lands on a different worker than the rid. + let rid_idx = select_by_key("conv"); + let other_key = (0..64) + .map(|i| format!("key-{i}")) + .find(|key| select_by_key(key) != rid_idx) + .expect("some key must land on a different worker"); + + let headers = headers_with_routing_key(&other_key); + let info = SelectWorkerInfo { + rid_key: Some("conv"), + routing_key: Some(&other_key), + headers: Some(&headers), + hash_ring: Some(ring.clone()), + ..Default::default() + }; + let (result, branch) = policy.select_worker_impl(&workers, &info); + assert_eq!(result, Some(rid_idx), "the rid-derived key must win"); + assert_eq!(branch, Branch::RoutingKeyHit); + } + #[test] fn test_empty_routing_key_uses_fallback() { let policy = ConsistentHashingPolicy::new(); diff --git a/model_gateway/src/policies/manual.rs b/model_gateway/src/policies/manual.rs index 1816222cc9..d9ddb01e07 100644 --- a/model_gateway/src/policies/manual.rs +++ b/model_gateway/src/policies/manual.rs @@ -10,7 +10,9 @@ //! Use this when you need stronger stickiness guarantees than consistent hashing, //! for example with stateful chat sessions where context is stored on the worker. //! -//! ## Header +//! ## Key sources +//! - `SelectWorkerInfo::rid_key`: the session key derived from the request +//! body's `rid` (populated under `--routing-key-override`); outranks the header //! - `X-SMG-Routing-Key`: The routing key for sticky session routing use std::{sync::Arc, time::Instant}; @@ -286,7 +288,12 @@ impl ManualPolicy { return (None, ExecutionBranch::NoHealthyWorkers); } - if let Some(routing_id) = extract_routing_key(info.headers) { + // The rid-derived session key outranks the routing-key header. The + // registry populates it only under the routing-key override, already + // lineage-stripped, so a proxy that rotates the header per request + // cannot break a conversation's pin. + let routing_id = info.rid_key.or_else(|| extract_routing_key(info.headers)); + if let Some(routing_id) = routing_id { // 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. @@ -451,6 +458,64 @@ mod tests { } } + #[test] + fn test_manual_rid_key_outranks_header_key() { + let policy = ManualPolicy::new(); + let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + + // One conversation (rid-derived key) whose per-request header key + // rotates: the pin must follow the rid, not the header. + let header_a = headers_with_routing_key("key_a"); + let header_b = headers_with_routing_key("key_b"); + let turn1 = SelectWorkerInfo { + headers: Some(&header_a), + rid_key: Some("conv"), + ..Default::default() + }; + let turn2 = SelectWorkerInfo { + headers: Some(&header_b), + rid_key: Some("conv"), + ..Default::default() + }; + + let (first, branch) = policy.select_worker_impl(&workers, &turn1); + assert_eq!(branch, ExecutionBranch::Vacant); + let (second, branch) = policy.select_worker_impl(&workers, &turn2); + assert_eq!(second, first, "rid-derived key must outrank the header key"); + assert_eq!(branch, ExecutionBranch::OccupiedHit); + + // The pin lives under the rid-derived key, not under either header. + assert_eq!(policy.routing_map.len(), 1); + assert!(policy.routing_map.contains_key(&RoutingId::new("conv"))); + } + + #[test] + fn test_manual_rid_key_namespaces_per_leg() { + let policy = ManualPolicy::new(); + let workers = create_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + + let prefill = SelectWorkerInfo { + rid_key: Some("conv"), + leg: WorkerLeg::Prefill, + ..Default::default() + }; + let decode = SelectWorkerInfo { + rid_key: Some("conv"), + leg: WorkerLeg::Decode, + ..Default::default() + }; + + let (p1, _) = policy.select_worker_impl(&workers, &prefill); + let (d1, _) = policy.select_worker_impl(&workers, &decode); + + // Same rid 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(); diff --git a/model_gateway/src/policies/registry.rs b/model_gateway/src/policies/registry.rs index 9fd45004a7..d36b8d58b7 100644 --- a/model_gateway/src/policies/registry.rs +++ b/model_gateway/src/policies/registry.rs @@ -250,8 +250,9 @@ impl PolicyRegistry { /// Select a worker, applying the sticky routing-key override when it is /// enabled, the request carries a key from the configured source, and the /// configured policy does not already honor the key (`manual` / - /// `consistent_hashing`). Otherwise delegates to `policy`. `policy.name()` - /// stays the real policy (for metrics). + /// `consistent_hashing`, which read `rid_key` and the header themselves + /// with the same rid-first precedence). Otherwise delegates to `policy`. + /// `policy.name()` stays the real policy (for metrics). pub fn select_worker( &self, policy: &Arc, @@ -355,7 +356,8 @@ impl PolicyRegistry { finish(Some(idx), branch) } - /// Policies that already honor the routing key keep their own handling; all + /// Policies that already honor the routing key keep their own handling + /// (they consult `rid_key` before the header, like the override does); all /// others (cache_aware, least_load, prefix_hash, ...) get the sticky override. fn routing_key_override_applies(name: &str) -> bool { !matches!(name, "manual" | "consistent_hashing") @@ -985,7 +987,7 @@ mod tests { use super::*; use crate::{ policies::{CacheAwareConfig, LeastLoadPolicy, SelectWorkerInfo}, - worker::{BasicWorkerBuilder, Worker, WorkerLoadGuard, WorkerType}, + worker::{BasicWorkerBuilder, HashRing, Worker, WorkerLoadGuard, WorkerType}, }; fn no_health_check() -> HealthCheckConfig { @@ -1056,6 +1058,72 @@ mod tests { } } + /// `--routing-key-override` means the same thing under every policy: + /// the body rid outranks the routing-key header. Key-native policies + /// skip the sticky override, so they must honor `rid_key` themselves. + #[test] + fn rid_key_outranks_header_under_key_native_policies() { + for config in [ + PolicyConfig::Manual { + eviction_interval_secs: 60, + max_idle_secs: 3600, + assignment_mode: ManualAssignmentMode::Random, + }, + PolicyConfig::ConsistentHashing, + ] { + let reg = PolicyRegistry::with_override( + config, + RoutingKeyOverrideConfig { + enabled: true, + ..Default::default() + }, + ); + let policy = reg.get_default_policy(); + let name = policy.name(); + assert!(!PolicyRegistry::routing_key_override_applies(name)); + let workers = vec![ + worker("http://w1", WorkerType::Regular), + worker("http://w2", WorkerType::Regular), + worker("http://w3", WorkerType::Regular), + worker("http://w4", WorkerType::Regular), + ]; + let hash_ring = Some(Arc::new(HashRing::new(workers.iter().map(|w| w.url())))); + + let rid_key = reg.derive_rid_key(Some("conv_t1")); + assert_eq!(rid_key, Some("conv")); + let pinned = reg + .select_worker( + &policy, + &workers, + &SelectWorkerInfo { + rid_key, + hash_ring: hash_ring.clone(), + ..Default::default() + }, + ) + .unwrap(); + + // Later turns of the conversation with rotating header keys stay + // on the rid's worker, whatever the header would have picked. + for (turn, key) in (2..).zip(["key_a", "key_b", "key_c", "key_d"]) { + let headers = headers_with_key(key); + let rid = format!("conv_t{turn}"); + let info = SelectWorkerInfo { + headers: Some(&headers), + routing_key: reg.resolve_routing_key(Some(&headers)), + rid_key: reg.derive_rid_key(Some(&rid)), + hash_ring: hash_ring.clone(), + ..Default::default() + }; + assert_eq!( + reg.select_worker(&policy, &workers, &info), + Some(pinned), + "{name}: body rid must outrank header {key}" + ); + } + } + } + #[test] fn override_without_key_uses_configured_policy() { let reg = PolicyRegistry::with_override(