From 7fd366d3f31b5db68a1f931316614dd142d1e49f Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 16 Sep 2026 04:46:07 +0000 Subject: [PATCH 1/2] feat: add sticky least-loaded session routing Balance new sessions by active-session count while preserving existing affinity, with explicit release and idle expiry. Preserve legacy session identifiers and wire the policy through regular and prefill/decode routing. Squashes vllm-project/router#235 commits a06430d and 6422afa. The implementation was adapted upstream from SumanthRH/router commit b50d926f3cf8868e12e8395fa02ce1aa2b1ae5ae. Signed-off-by: bvolpato --- README.md | 38 ++ benches/request_processing.rs | 2 + py_src/vllm_router/router.py | 1 + py_src/vllm_router/router_args.py | 3 + py_test/unit/test_arg_parser.py | 16 + src/config/types.rs | 9 + src/config/validation.rs | 4 + src/lib.rs | 2 + src/main.rs | 7 +- src/policies/factory.rs | 11 + src/policies/hash_key.rs | 99 +++++ src/policies/mod.rs | 17 + src/policies/registry.rs | 65 +++ src/policies/sticky_least_loaded.rs | 604 ++++++++++++++++++++++++++++ src/protocols/spec.rs | 8 + src/routers/http/router.rs | 19 +- src/routers/http/vllm_pd_router.rs | 123 +++++- src/server.rs | 28 +- tests/api_endpoints_test.rs | 166 ++++++++ tests/benchmark_integration.rs | 2 + tests/test_openai_routing.rs | 2 + 21 files changed, 1206 insertions(+), 20 deletions(-) create mode 100644 src/policies/sticky_least_loaded.rs diff --git a/README.md b/README.md index 1bc06289..7cb4c1a8 100644 --- a/README.md +++ b/README.md @@ -206,6 +206,7 @@ The router supports multiple load balancing policies: | `round_robin` | Sequential distribution across workers | No | General purpose, even distribution | | `random` | Uniform random selection | No | Simple deployments | | `consistent_hash` | Routes same session/user to same worker | Yes | Multi-turn chat, KV cache reuse | +| `sticky_least_loaded` | Assigns new sessions to the worker with the fewest active sessions | Yes | Multi-turn workloads with explicit session lifetimes | | `power_of_two` | Picks least loaded of two random workers | No | Load-sensitive workloads | | `cache_aware` | Optimizes for prefix cache hits | Yes | Repeated prompts, few-shot | @@ -219,6 +220,43 @@ curl -X POST http://router:8000/v1/chat/completions \ For detailed configuration options, hash key priorities, and usage examples, see [Load Balancing Documentation](docs/load_balancing/README.md). +#### Sticky least-loaded sessions + +Use `--policy sticky_least_loaded` to balance **active sessions**, not concurrent +requests, token counts, or GPU utilization. New sessions reserve a worker +atomically; ties use rendezvous hashing. Existing sessions keep their healthy +worker, even if session counts later become uneven. Unavailable workers cause +reassignment on the next request. Prefill and decode pools track sessions separately. + +Session IDs use the same header priority as `consistent_hash`, then the JSON fields +`session_params.session_id`, `user`, `session_id`, and `user_id`. Prefer +`X-Session-ID` for multi-turn work; a per-request ID creates a separate session for +each request. Requests without an explicit identifier do not reserve a session. + +Release a session after its final request completes: + +```bash +curl -X POST 'http://router:8000/finish_session?session_id=my-session-123' +``` + +This endpoint uses the router's configured API-key validation. It releases the ID +across default, per-model, prefill, and decode policies. Repeated or unknown IDs +are harmless. It does not cancel generation; use globally unique IDs and finish +only after all requests for that session have completed. Models using the shared +default policy must include the model in their session ID (for example, +`model-a:conversation-123`), or use separate per-model policy instances. Reusing +one ID across disjoint model pools reassigns its single reservation between pools. + +Idle sessions expire after two hours. Set +`VLLM_ROUTER_SLL_SESSION_EXPIRATION_IN_S` on the router to override the TTL; expired +entries are swept on routing requests at most once per minute. Revisited sessions +are checked for expiry before refreshing their affinity. State is local to one +router process and is lost on restart. Multiple routers need consistent ingress +routing and session release on each router that tracked the session. +The policy keeps one entry per active session and serializes assignment under a +mutex, so callers should release sessions +promptly and choose a TTL appropriate for their workload. + ## Advanced Features ### Kubernetes Service Discovery diff --git a/benches/request_processing.rs b/benches/request_processing.rs index 9c1ed893..96d7a87d 100644 --- a/benches/request_processing.rs +++ b/benches/request_processing.rs @@ -20,6 +20,8 @@ fn default_generate_request() -> GenerateRequest { // VLLM Extensions lora_path: None, session_params: None, + session_id: None, + user_id: None, return_hidden_states: false, rid: None, } diff --git a/py_src/vllm_router/router.py b/py_src/vllm_router/router.py index 5d92b900..00af84ef 100644 --- a/py_src/vllm_router/router.py +++ b/py_src/vllm_router/router.py @@ -15,6 +15,7 @@ def policy_from_str(policy_str: Optional[str]) -> PolicyType: "cache_aware": PolicyType.CacheAware, "power_of_two": PolicyType.PowerOfTwo, "consistent_hash": PolicyType.ConsistentHash, + "sticky_least_loaded": PolicyType.StickyLeastLoaded, } return policy_map[policy_str] diff --git a/py_src/vllm_router/router_args.py b/py_src/vllm_router/router_args.py index 5771bf49..4a9e45d4 100644 --- a/py_src/vllm_router/router_args.py +++ b/py_src/vllm_router/router_args.py @@ -140,6 +140,7 @@ def add_cli_args( "cache_aware", "power_of_two", "consistent_hash", + "sticky_least_loaded", ], help="Load balancing policy to use. In PD mode, this is used for both prefill and decode unless overridden", ) @@ -153,6 +154,7 @@ def add_cli_args( "cache_aware", "power_of_two", "consistent_hash", + "sticky_least_loaded", ], help="Specific policy for prefill nodes in PD mode. If not specified, uses the main policy", ) @@ -166,6 +168,7 @@ def add_cli_args( "cache_aware", "power_of_two", "consistent_hash", + "sticky_least_loaded", ], help="Specific policy for decode nodes in PD mode. If not specified, uses the main policy", ) diff --git a/py_test/unit/test_arg_parser.py b/py_test/unit/test_arg_parser.py index c7a43f92..bb3e519e 100644 --- a/py_test/unit/test_arg_parser.py +++ b/py_test/unit/test_arg_parser.py @@ -409,6 +409,7 @@ def test_valid_policies(self): assert policy_from_str("cache_aware") == PolicyType.CacheAware assert policy_from_str("power_of_two") == PolicyType.PowerOfTwo assert policy_from_str("consistent_hash") == PolicyType.ConsistentHash + assert policy_from_str("sticky_least_loaded") == PolicyType.StickyLeastLoaded def test_invalid_policy(self): """Test conversion of invalid policy string.""" @@ -419,6 +420,21 @@ def test_invalid_policy(self): class TestParseRouterArgs: """Test the parse_router_args function.""" + def test_parse_sticky_least_loaded_for_each_pool(self): + args = parse_router_args( + [ + "--policy", + "sticky_least_loaded", + "--prefill-policy", + "sticky_least_loaded", + "--decode-policy", + "sticky_least_loaded", + ] + ) + assert args.policy == "sticky_least_loaded" + assert args.prefill_policy == "sticky_least_loaded" + assert args.decode_policy == "sticky_least_loaded" + def test_parse_basic_args(self): """Test parsing basic router arguments.""" args = [ diff --git a/src/config/types.rs b/src/config/types.rs index d307295b..0b8c316b 100644 --- a/src/config/types.rs +++ b/src/config/types.rs @@ -247,6 +247,14 @@ pub enum PolicyConfig { virtual_nodes: u32, }, + /// Session-aware routing that combines session stickiness with load + /// balancing: existing sessions stick to their assigned replica, new + /// sessions are routed to the least-loaded replica (ties broken by + /// consistent hashing). Session expiry is configured via the + /// `VLLM_ROUTER_SLL_SESSION_EXPIRATION_IN_S` environment variable. + #[serde(rename = "sticky_least_loaded")] + StickyLeastLoaded, + #[serde(rename = "rendezvous_hash")] RendezvousHash, } @@ -259,6 +267,7 @@ impl PolicyConfig { PolicyConfig::CacheAware { .. } => "cache_aware", PolicyConfig::PowerOfTwo { .. } => "power_of_two", PolicyConfig::ConsistentHash { .. } => "consistent_hash", + PolicyConfig::StickyLeastLoaded => "sticky_least_loaded", PolicyConfig::RendezvousHash => "rendezvous_hash", } } diff --git a/src/config/validation.rs b/src/config/validation.rs index bf678a5b..f5b68c18 100644 --- a/src/config/validation.rs +++ b/src/config/validation.rs @@ -188,6 +188,10 @@ impl ConfigValidator { }); } } + PolicyConfig::StickyLeastLoaded => { + // Session expiration is sourced from the environment; no + // structured config to validate here. + } PolicyConfig::RendezvousHash => { // No specific validation needed } diff --git a/src/lib.rs b/src/lib.rs index 65c8a07a..d0bdf22c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -27,6 +27,7 @@ pub enum PolicyType { CacheAware, PowerOfTwo, ConsistentHash, + StickyLeastLoaded, } #[pyclass] @@ -125,6 +126,7 @@ impl Router { PolicyType::ConsistentHash => ConfigPolicyConfig::ConsistentHash { virtual_nodes: 160, // Default value }, + PolicyType::StickyLeastLoaded => ConfigPolicyConfig::StickyLeastLoaded, } }; diff --git a/src/main.rs b/src/main.rs index 3fba1e36..84198b7e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -107,7 +107,7 @@ struct CliArgs { worker_urls: Vec, /// Load balancing policy to use - #[arg(long, default_value = "cache_aware", value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "rendezvous_hash"])] + #[arg(long, default_value = "cache_aware", value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "sticky_least_loaded", "rendezvous_hash"])] policy: String, /// Enable vLLM PD (Prefill-Decode) disaggregated mode with vLLM-specific two-stage processing @@ -124,11 +124,11 @@ struct CliArgs { decode: Vec, /// Specific policy for prefill nodes in PD mode - #[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "rendezvous_hash"])] + #[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "sticky_least_loaded", "rendezvous_hash"])] prefill_policy: Option, /// Specific policy for decode nodes in PD mode - #[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "rendezvous_hash"])] + #[arg(long, value_parser = ["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash", "sticky_least_loaded", "rendezvous_hash"])] decode_policy: Option, /// Timeout in seconds for worker startup @@ -378,6 +378,7 @@ impl CliArgs { "consistent_hash" => PolicyConfig::ConsistentHash { virtual_nodes: 160, // Default value }, + "sticky_least_loaded" => PolicyConfig::StickyLeastLoaded, "rendezvous_hash" => PolicyConfig::RendezvousHash, _ => PolicyConfig::RoundRobin, // Fallback } diff --git a/src/policies/factory.rs b/src/policies/factory.rs index 08ff2d80..952cce07 100644 --- a/src/policies/factory.rs +++ b/src/policies/factory.rs @@ -3,6 +3,7 @@ use super::{ CacheAwareConfig, CacheAwarePolicy, ConsistentHashPolicy, LoadBalancingPolicy, PowerOfTwoPolicy, RandomPolicy, RendezvousHashPolicy, RoundRobinPolicy, + StickyLeastLoadedPolicy, }; use crate::config::PolicyConfig; use std::sync::Arc; @@ -38,6 +39,7 @@ impl PolicyFactory { // The consistent hash policy uses a hardcoded value for now Arc::new(ConsistentHashPolicy::new()) } + PolicyConfig::StickyLeastLoaded => Arc::new(StickyLeastLoadedPolicy::new()), PolicyConfig::RendezvousHash => Arc::new(RendezvousHashPolicy::new()), } } @@ -50,6 +52,9 @@ impl PolicyFactory { "power_of_two" | "poweroftwo" => Some(Arc::new(PowerOfTwoPolicy::new())), "cache_aware" | "cacheaware" => Some(Arc::new(CacheAwarePolicy::new())), "consistent_hash" | "consistenthash" => Some(Arc::new(ConsistentHashPolicy::new())), + "sticky_least_loaded" | "stickyleastloaded" => { + Some(Arc::new(StickyLeastLoadedPolicy::new())) + } "rendezvous_hash" | "rendezvoushash" => Some(Arc::new(RendezvousHashPolicy::new())), _ => None, } @@ -94,6 +99,10 @@ mod tests { // Test RendezvousHash let policy = PolicyFactory::create_from_config(&PolicyConfig::RendezvousHash); assert_eq!(policy.name(), "rendezvous_hash"); + + // Test StickyLeastLoaded + let policy = PolicyFactory::create_from_config(&PolicyConfig::StickyLeastLoaded); + assert_eq!(policy.name(), "sticky_least_loaded"); } #[test] @@ -108,6 +117,8 @@ mod tests { assert!(PolicyFactory::create_by_name("CacheAware").is_some()); assert!(PolicyFactory::create_by_name("consistent_hash").is_some()); assert!(PolicyFactory::create_by_name("ConsistentHash").is_some()); + assert!(PolicyFactory::create_by_name("sticky_least_loaded").is_some()); + assert!(PolicyFactory::create_by_name("StickyLeastLoaded").is_some()); assert!(PolicyFactory::create_by_name("rendezvous_hash").is_some()); assert!(PolicyFactory::create_by_name("RendezvousHash").is_some()); assert!(PolicyFactory::create_by_name("unknown").is_none()); diff --git a/src/policies/hash_key.rs b/src/policies/hash_key.rs index 08c8d316..cd950250 100644 --- a/src/policies/hash_key.rs +++ b/src/policies/hash_key.rs @@ -51,6 +51,51 @@ pub(crate) fn extract_hash_key( } } +/// Extract a raw session identifier (the *value* only, without any prefix) from +/// HTTP headers or request body. +/// +/// Unlike [`extract_hash_key`], this does NOT fall back to hashing the request +/// body when no explicit session/user identifier is present; it returns `None` +/// instead. This is used by load-balancing policies (e.g. +/// `sticky_least_loaded`) that need a stable, externally-addressable +/// session id so that a matching `finish_session(session_id)` call can later +/// release the session. +/// +/// Lookup order mirrors [`extract_hash_key`]: +/// 1. HTTP headers: x-session-id, x-user-id, x-tenant-id, x-correlation-id, x-request-id, x-trace-id +/// 2. Body: session_params.session_id (nested) +/// 3. Body: user (OpenAI format) +/// 4. Body: session_id (legacy) +/// 5. Body: user_id (legacy) +pub(crate) fn extract_session_id( + request_text: Option<&str>, + headers: Option<&RequestHeaders>, +) -> Option { + if let Some(hdrs) = headers { + for header_name in SESSION_HEADER_NAMES { + if let Some(value) = hdrs.get(*header_name) { + if !value.is_empty() { + return Some(value.clone()); + } + } + } + } + + let body: serde_json::Value = serde_json::from_str(request_text?).ok()?; + let session_id = [ + body.pointer("/session_params/session_id"), + body.get("user"), + body.get("session_id"), + body.get("user_id"), + ] + .into_iter() + .flatten() + .filter_map(serde_json::Value::as_str) + .find(|value| !value.is_empty()) + .map(str::to_owned); + session_id +} + /// Extract hash key from HTTP headers pub(crate) fn extract_hash_key_from_headers(headers: &RequestHeaders) -> Option { for header_name in SESSION_HEADER_NAMES { @@ -476,4 +521,58 @@ mod tests { let text = r#"{"other": "value"}"#; assert_eq!(find_field_start(text, "field"), None); } + + // === extract_session_id tests === + + #[test] + fn test_extract_session_id_from_header() { + let mut headers = HashMap::new(); + headers.insert("x-session-id".to_string(), "traj-42".to_string()); + // Returns the raw value, without any "header:" prefix + assert_eq!( + extract_session_id(None, Some(&headers)), + Some("traj-42".to_string()) + ); + } + + #[test] + fn test_extract_session_id_header_priority_over_body() { + let mut headers = HashMap::new(); + headers.insert("x-session-id".to_string(), "from-header".to_string()); + let body = r#"{"session_id": "from-body"}"#; + assert_eq!( + extract_session_id(Some(body), Some(&headers)), + Some("from-header".to_string()) + ); + } + + #[test] + fn test_extract_session_id_from_body() { + let body = r#"{"session_id": "legacy123", "prompt": "hi"}"#; + assert_eq!( + extract_session_id(Some(body), None), + Some("legacy123".to_string()) + ); + } + + #[test] + fn test_extract_session_id_no_fallback() { + // No explicit session identifier -> None (unlike extract_hash_key) + let body = r#"{"prompt": "hello", "model": "llama"}"#; + assert_eq!(extract_session_id(Some(body), None), None); + assert_eq!(extract_session_id(None, None), None); + } + + #[test] + fn test_extract_session_id_preserves_json_escapes() { + let id = "session-\"quoted\\path"; + let body = serde_json::json!({"session_params": {"session_id": id}}).to_string(); + assert_eq!(extract_session_id(Some(&body), None), Some(id.to_string())); + } + + #[test] + fn test_extract_session_id_ignores_unrelated_nested_fields() { + let body = r#"{"metadata":{"session_id":"not-a-session"},"user":null}"#; + assert_eq!(extract_session_id(Some(body), None), None); + } } diff --git a/src/policies/mod.rs b/src/policies/mod.rs index 59481765..6c0239ef 100644 --- a/src/policies/mod.rs +++ b/src/policies/mod.rs @@ -17,6 +17,7 @@ mod random; mod registry; mod rendezvous_hash; mod round_robin; +mod sticky_least_loaded; pub use cache_aware::CacheAwarePolicy; pub use consistent_hash::ConsistentHashPolicy; @@ -27,6 +28,7 @@ pub use random::RandomPolicy; pub use registry::PolicyRegistry; pub use rendezvous_hash::RendezvousHashPolicy; pub use round_robin::RoundRobinPolicy; +pub use sticky_least_loaded::StickyLeastLoadedPolicy; /// HTTP headers passed to policies for routing decisions /// Key is lowercase header name, value is header value @@ -97,6 +99,15 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug { // Default: no-op for stateless policies } + /// Mark a session (e.g. an RL trajectory) as finished. + /// + /// Session-aware policies (e.g. `sticky_least_loaded`) use this to + /// release the active-session assignment. Stateless policies, and + /// policies that don't track sessions, ignore this. + fn finish_session(&self, _session_id: &str) { + // Default: no-op for policies that don't track sessions + } + /// Get policy name for metrics and debugging fn name(&self) -> &'static str; @@ -105,6 +116,12 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug { false // Default: most policies don't need request text } + /// Whether typed routes should supply the serialized body instead of their + /// prompt-derived routing key, for policies that inspect structured fields. + fn needs_request_body(&self) -> bool { + false + } + /// Check if this policy needs HTTP headers for routing decisions fn needs_headers(&self) -> bool { false // Default: most policies don't need headers diff --git a/src/policies/registry.rs b/src/policies/registry.rs index fb759011..a230bf71 100644 --- a/src/policies/registry.rs +++ b/src/policies/registry.rs @@ -7,6 +7,7 @@ use super::{ CacheAwareConfig, CacheAwarePolicy, ConsistentHashPolicy, LoadBalancingPolicy, PowerOfTwoPolicy, RandomPolicy, RendezvousHashPolicy, RoundRobinPolicy, + StickyLeastLoadedPolicy, }; use crate::config::types::PolicyConfig; use std::collections::HashMap; @@ -172,6 +173,7 @@ impl PolicyRegistry { "random" => Arc::new(RandomPolicy::new()), "cache_aware" => Arc::new(CacheAwarePolicy::new()), "power_of_two" => Arc::new(PowerOfTwoPolicy::new()), + "sticky_least_loaded" => Arc::new(StickyLeastLoadedPolicy::new()), "rendezvous_hash" => Arc::new(RendezvousHashPolicy::new()), _ => { warn!("Unknown policy type '{}', using default", policy_type); @@ -203,6 +205,7 @@ impl PolicyRegistry { } PolicyConfig::PowerOfTwo { .. } => Arc::new(PowerOfTwoPolicy::new()), PolicyConfig::ConsistentHash { .. } => Arc::new(ConsistentHashPolicy::new()), + PolicyConfig::StickyLeastLoaded => Arc::new(StickyLeastLoadedPolicy::new()), PolicyConfig::RendezvousHash => Arc::new(RendezvousHashPolicy::new()), } } @@ -221,6 +224,31 @@ impl PolicyRegistry { self.model_worker_counts.read().unwrap().clone() } + /// Mark a session as finished across all registered policies. + /// + /// Session-aware policies (e.g. `sticky_least_loaded`) will release + /// the active-session assignment; all other policies ignore it. + /// This fans out to the default policy, every per-model policy, and the + /// prefill/decode policies (for PD mode), since we don't know which policy + /// instance is tracking the session. + pub fn finish_session(&self, session_id: &str) { + self.default_policy.finish_session(session_id); + + { + let policies = self.model_policies.read().unwrap(); + for policy in policies.values() { + policy.finish_session(session_id); + } + } + + if let Some(policy) = self.prefill_policy.read().unwrap().as_ref() { + policy.finish_session(session_id); + } + if let Some(policy) = self.decode_policy.read().unwrap().as_ref() { + policy.finish_session(session_id); + } + } + /// Clear all policies (useful for testing) pub fn clear(&self) { let mut policies = self.model_policies.write().unwrap(); @@ -274,6 +302,43 @@ impl std::fmt::Debug for PolicyRegistry { mod tests { use super::*; + #[test] + fn test_finish_session_releases_all_policy_scopes() { + use crate::core::{BasicWorker, Worker, WorkerType}; + let registry = PolicyRegistry::new(PolicyConfig::StickyLeastLoaded); + registry.set_prefill_policy(Arc::new(StickyLeastLoadedPolicy::new())); + registry.set_decode_policy(Arc::new(StickyLeastLoadedPolicy::new())); + let policies = [ + registry.get_default_policy(), + registry.on_worker_added("model", Some("sticky_least_loaded")), + registry.get_prefill_policy(), + registry.get_decode_policy(), + ]; + let workers: Vec> = vec![Arc::new(BasicWorker::new( + "http://worker:8000".to_string(), + WorkerType::Regular, + ))]; + let headers = HashMap::from([("x-session-id".to_string(), "session".to_string())]); + for policy in &policies { + assert_eq!( + policy.select_worker_with_headers(&workers, None, Some(&headers)), + Some(0) + ); + } + registry.finish_session("session"); + registry.finish_session("session"); + for policy in policies { + assert_eq!( + policy + .as_any() + .downcast_ref::() + .unwrap() + .active_session_count(), + 0 + ); + } + } + #[test] fn test_policy_registry_basic() { let registry = PolicyRegistry::new(PolicyConfig::RoundRobin); diff --git a/src/policies/sticky_least_loaded.rs b/src/policies/sticky_least_loaded.rs new file mode 100644 index 00000000..5a727002 --- /dev/null +++ b/src/policies/sticky_least_loaded.rs @@ -0,0 +1,604 @@ +//! Sticky least-loaded session routing policy +//! +//! Adapted from SumanthRH/router commit b50d926f3cf8868e12e8395fa02ce1aa2b1ae5ae +//! (Apache-2.0), with typed-body routing and per-entry expiration fixes. +//! +//! This policy combines session stickiness (so all requests for a given +//! trajectory/session reuse the same vLLM replica and maximize KV cache reuse) +//! with load balancing across replicas (so *new* sessions are assigned to the +//! replica with the fewest active sessions). +//! +//! Behavior: +//! - The router bookkeeps the set of active sessions per replica. +//! - For an existing session: requests with the same session id are routed to +//! the replica that session was originally assigned to. +//! - For a new session: the session is assigned to the healthy replica with the +//! least number of active sessions. Ties are broken deterministically using +//! consistent hashing (rendezvous / highest-random-weight) so that, given the +//! same set of least-loaded replicas, a session is always assigned to the same +//! replica. +//! - Sessions expire after a configurable TTL (default 2 hours, overridable via +//! the `VLLM_ROUTER_SLL_SESSION_EXPIRATION_IN_S` environment variable). This +//! protects against leaked sessions when a `finish_session` call is never +//! received (e.g. a cancelled trajectory). +//! - Sessions can be released explicitly via [`finish_session`]. +//! +//! Requests without any session identifier are still load-balanced to the +//! least-loaded replica, but are *not* recorded as active sessions. + +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +use tracing::{debug, info, warn}; + +use super::get_healthy_worker_indices; +use super::hash_key; +use super::ConsistentHashPolicy; +use super::LoadBalancingPolicy; +use super::RequestHeaders; +use crate::core::Worker; +use crate::metrics::RouterMetrics; + +/// Default idle-session expiration (2 hours). +pub const DEFAULT_SESSION_EXPIRATION_SECS: u64 = 7200; + +/// Environment variable to override the session expiration (in seconds). +pub const SESSION_EXPIRATION_ENV: &str = "VLLM_ROUTER_SLL_SESSION_EXPIRATION_IN_S"; + +/// Minimum interval between full expiration sweeps to bound per-request cost. +const SWEEP_INTERVAL_SECS: u64 = 60; + +/// Bookkeeping for a single active session. +#[derive(Debug)] +struct SessionEntry { + worker_url: String, + last_access: Instant, +} + +/// Internal mutable state guarded by a single mutex so that "find least-loaded +/// replica + record session" is atomic across concurrent new-session requests. +#[derive(Debug)] +struct SllState { + /// session id -> assigned replica + sessions: HashMap, + /// replica url -> number of active sessions + active_counts: HashMap, + /// Timestamp of the last expiration sweep. + last_sweep: Instant, +} + +/// Sticky least-loaded routing policy. +#[derive(Debug)] +pub struct StickyLeastLoadedPolicy { + state: Mutex, + session_expiration: Duration, +} + +impl StickyLeastLoadedPolicy { + /// Create a new policy, reading the session expiration from the environment + /// (falling back to [`DEFAULT_SESSION_EXPIRATION_SECS`]). + pub fn new() -> Self { + Self::with_expiration_secs(Self::expiration_secs_from_env()) + } + + /// Create a new policy with an explicit session expiration (in seconds). + pub fn with_expiration_secs(secs: u64) -> Self { + Self { + state: Mutex::new(SllState { + sessions: HashMap::new(), + active_counts: HashMap::new(), + last_sweep: Instant::now(), + }), + session_expiration: Duration::from_secs(secs), + } + } + + fn expiration_secs_from_env() -> u64 { + match std::env::var(SESSION_EXPIRATION_ENV) { + Ok(val) => match val.trim().parse::() { + Ok(secs) => secs, + Err(_) => { + warn!( + "Invalid {} value '{}', using default {}s", + SESSION_EXPIRATION_ENV, val, DEFAULT_SESSION_EXPIRATION_SECS + ); + DEFAULT_SESSION_EXPIRATION_SECS + } + }, + Err(_) => DEFAULT_SESSION_EXPIRATION_SECS, + } + } + + /// Decrement the active-session count for a replica, removing the entry when + /// it reaches zero. + fn decrement_count(counts: &mut HashMap, worker_url: &str) { + if let Some(count) = counts.get_mut(worker_url) { + *count = count.saturating_sub(1); + if *count == 0 { + counts.remove(worker_url); + } + } + } + + /// Remove sessions that have not been accessed within the expiration window. + /// Runs at most once per [`SWEEP_INTERVAL_SECS`] to bound cost. + fn sweep_expired(state: &mut SllState, expiration: Duration, now: Instant) { + if now.duration_since(state.last_sweep) < Duration::from_secs(SWEEP_INTERVAL_SECS) { + return; + } + + let expired: Vec = state + .sessions + .iter() + .filter(|(_, entry)| now.duration_since(entry.last_access) >= expiration) + .map(|(key, _)| key.clone()) + .collect(); + + for key in expired { + if let Some(entry) = state.sessions.remove(&key) { + Self::decrement_count(&mut state.active_counts, &entry.worker_url); + debug!("SLL: expired session '{}' on '{}'", key, entry.worker_url); + } + } + + state.last_sweep = now; + } + + /// Pick the least-loaded replica among `candidates`, breaking ties with + /// consistent (rendezvous) hashing of `tie_break_key`. + /// + /// `candidates` is a slice of `(worker_index, worker_url)` for healthy + /// replicas. Returns the chosen `worker_index`. + fn select_least_loaded( + counts: &HashMap, + candidates: &[(usize, String)], + tie_break_key: &str, + ) -> usize { + let min_count = candidates + .iter() + .map(|(_, url)| counts.get(url).copied().unwrap_or(0)) + .min() + .unwrap_or(0); + + // Among replicas at the minimum load, deterministically pick one using + // rendezvous hashing (highest hash wins). + let mut best: Option<(usize, u64)> = None; + for (idx, url) in candidates { + let count = counts.get(url).copied().unwrap_or(0); + if count != min_count { + continue; + } + let weight = ConsistentHashPolicy::fbi_hash(&format!("{}:{}", tie_break_key, url)); + match best { + Some((_, best_weight)) if best_weight >= weight => {} + _ => best = Some((*idx, weight)), + } + } + + // `candidates` is non-empty (callers ensure healthy workers exist), so + // `best` is always populated. + best.map(|(idx, _)| idx).unwrap_or(candidates[0].0) + } + + fn record_selection(&self, worker: &Arc) { + worker.increment_processed(); + RouterMetrics::record_processed_request(worker.url()); + RouterMetrics::record_policy_decision(self.name(), worker.url()); + } + + /// Number of currently active (tracked) sessions. Exposed for tests/metrics. + pub fn active_session_count(&self) -> usize { + self.state.lock().unwrap().sessions.len() + } + + /// Number of active sessions assigned to a specific replica. For tests. + pub fn active_count_for(&self, worker_url: &str) -> usize { + self.state + .lock() + .unwrap() + .active_counts + .get(worker_url) + .copied() + .unwrap_or(0) + } +} + +impl LoadBalancingPolicy for StickyLeastLoadedPolicy { + fn select_worker_with_headers( + &self, + workers: &[Arc], + request_text: Option<&str>, + headers: Option<&RequestHeaders>, + ) -> Option { + let healthy_indices = get_healthy_worker_indices(workers); + if healthy_indices.is_empty() { + return None; + } + + let candidates: Vec<(usize, String)> = healthy_indices + .iter() + .map(|&idx| (idx, workers[idx].url().to_string())) + .collect(); + + let session_id = hash_key::extract_session_id(request_text, headers); + let mut state = self.state.lock().unwrap(); + let now = Instant::now(); + Self::sweep_expired(&mut state, self.session_expiration, now); + + // Requests without a session id: load-balance but don't record state. + let session_id = match session_id { + Some(id) => id, + None => { + let tie_break = request_text.unwrap_or(""); + let idx = Self::select_least_loaded(&state.active_counts, &candidates, tie_break); + drop(state); + self.record_selection(&workers[idx]); + debug!("SLL: stateless request routed to '{}'", workers[idx].url()); + return Some(idx); + } + }; + + // Existing session: route to the same replica if it's still healthy. + if let Some((worker_url, last_access)) = state + .sessions + .get(&session_id) + .map(|e| (e.worker_url.clone(), e.last_access)) + { + if let Some((idx, _)) = candidates.iter().find(|(_, url)| { + *url == worker_url && now.duration_since(last_access) < self.session_expiration + }) { + let idx = *idx; + if let Some(entry) = state.sessions.get_mut(&session_id) { + entry.last_access = now; + } + drop(state); + self.record_selection(&workers[idx]); + debug!( + "SLL: session '{}' -> existing replica '{}'", + session_id, worker_url + ); + return Some(idx); + } + + // Expired session or unavailable replica: release the old assignment. + state.sessions.remove(&session_id); + Self::decrement_count(&mut state.active_counts, &worker_url); + debug!( + "SLL: session '{}' replica '{}' unavailable, reassigning", + session_id, worker_url + ); + } + + // New session: assign to least-loaded replica (tie-broken by hashing). + let idx = Self::select_least_loaded(&state.active_counts, &candidates, &session_id); + let worker_url = workers[idx].url().to_string(); + state.sessions.insert( + session_id.clone(), + SessionEntry { + worker_url: worker_url.clone(), + last_access: now, + }, + ); + *state.active_counts.entry(worker_url.clone()).or_insert(0) += 1; + let new_count = state.active_counts[&worker_url]; + drop(state); + + self.record_selection(&workers[idx]); + info!( + "SLL: new session '{}' -> replica '{}' (active sessions on replica: {})", + session_id, worker_url, new_count + ); + Some(idx) + } + + fn finish_session(&self, session_id: &str) { + let mut state = self.state.lock().unwrap(); + if let Some(entry) = state.sessions.remove(session_id) { + Self::decrement_count(&mut state.active_counts, &entry.worker_url); + info!( + "SLL: finished session '{}' on replica '{}'", + session_id, entry.worker_url + ); + } else { + debug!("SLL: finish_session for unknown session '{}'", session_id); + } + } + + fn name(&self) -> &'static str { + "sticky_least_loaded" + } + + fn needs_request_text(&self) -> bool { + true + } + + fn needs_request_body(&self) -> bool { + true + } + + fn needs_headers(&self) -> bool { + true + } + + fn reset(&self) { + let mut state = self.state.lock().unwrap(); + state.sessions.clear(); + state.active_counts.clear(); + state.last_sweep = Instant::now(); + info!("SLL: policy reset - all sessions cleared"); + } + + fn as_any(&self) -> &dyn std::any::Any { + self + } +} + +impl Default for StickyLeastLoadedPolicy { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::core::{BasicWorker, WorkerType}; + use std::collections::HashMap as StdHashMap; + + fn make_workers(urls: &[&str]) -> Vec> { + urls.iter() + .map(|url| { + Arc::new(BasicWorker::new(url.to_string(), WorkerType::Regular)) as Arc + }) + .collect() + } + + fn header_with_session(session_id: &str) -> RequestHeaders { + let mut headers: RequestHeaders = StdHashMap::new(); + headers.insert("x-session-id".to_string(), session_id.to_string()); + headers + } + + #[test] + fn test_same_session_sticky() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + let headers = header_with_session("sess-1"); + + let idx1 = policy.select_worker_with_headers(&workers, None, Some(&headers)); + let idx2 = policy.select_worker_with_headers(&workers, None, Some(&headers)); + let idx3 = policy.select_worker_with_headers(&workers, None, Some(&headers)); + + assert!(idx1.is_some()); + assert_eq!(idx1, idx2); + assert_eq!(idx2, idx3); + + // Only one active session should be tracked across repeated requests. + assert_eq!(policy.active_session_count(), 1); + } + + #[test] + fn test_new_sessions_balance_across_replicas() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + + // Assign three distinct sessions; each should land on a distinct replica + // because new sessions go to the least-loaded replica. + for i in 0..3 { + let headers = header_with_session(&format!("sess-{}", i)); + policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + } + + assert_eq!(policy.active_session_count(), 3); + for w in &workers { + assert_eq!( + policy.active_count_for(w.url()), + 1, + "expected each replica to hold exactly one session" + ); + } + } + + #[test] + fn test_finish_session_releases_capacity() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + + let headers = header_with_session("sess-x"); + let idx = policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + let assigned_url = workers[idx].url().to_string(); + assert_eq!(policy.active_count_for(&assigned_url), 1); + + policy.finish_session("sess-x"); + assert_eq!(policy.active_session_count(), 0); + assert_eq!(policy.active_count_for(&assigned_url), 0); + } + + #[test] + fn test_finish_unknown_session_is_noop() { + let policy = StickyLeastLoadedPolicy::new(); + // Should not panic and should leave state empty. + policy.finish_session("does-not-exist"); + assert_eq!(policy.active_session_count(), 0); + } + + #[test] + fn test_session_id_from_body() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let body = r#"{"session_id": "body-sess", "prompt": "hi"}"#; + + let idx1 = policy.select_worker_with_headers(&workers, Some(body), None); + let idx2 = policy.select_worker_with_headers(&workers, Some(body), None); + assert_eq!(idx1, idx2); + assert_eq!(policy.active_session_count(), 1); + + policy.finish_session("body-sess"); + assert_eq!(policy.active_session_count(), 0); + } + + #[test] + fn test_stateless_requests_not_tracked() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let body = r#"{"prompt": "no session here"}"#; + + let idx = policy.select_worker_with_headers(&workers, Some(body), None); + assert!(idx.is_some()); + // No session identifier -> nothing tracked. + assert_eq!(policy.active_session_count(), 0); + } + + #[test] + fn test_reassign_when_replica_unhealthy() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let headers = header_with_session("sess-move"); + + let idx = policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + let original_url = workers[idx].url().to_string(); + + // Mark the assigned replica unhealthy; the session must be reassigned to + // the remaining healthy replica. + workers[idx].set_healthy(false); + let new_idx = policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + assert_ne!(workers[new_idx].url(), original_url); + assert!(workers[new_idx].is_healthy()); + + // Old replica should no longer hold the session. + assert_eq!(policy.active_count_for(&original_url), 0); + assert_eq!(policy.active_count_for(workers[new_idx].url()), 1); + } + + #[test] + fn test_expired_session_is_swept() { + // Zero expiration => any session is immediately eligible for expiry, but + // the sweep only runs every SWEEP_INTERVAL_SECS. Force the sweep by + // back-dating last_sweep. + let policy = StickyLeastLoadedPolicy::with_expiration_secs(0); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let headers = header_with_session("sess-ttl"); + + policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + assert_eq!(policy.active_session_count(), 1); + + { + let mut state = policy.state.lock().unwrap(); + state.last_sweep = Instant::now() - Duration::from_secs(SWEEP_INTERVAL_SECS + 1); + } + + // A subsequent request triggers the sweep, removing the expired session + // and creating a fresh one. + let other = header_with_session("sess-other"); + policy + .select_worker_with_headers(&workers, None, Some(&other)) + .unwrap(); + // The expired session was swept; only the freshly created one remains. + assert_eq!(policy.active_session_count(), 1); + assert_eq!( + policy.active_count_for("http://w1:8000") + policy.active_count_for("http://w2:8000"), + 1 + ); + } + + #[test] + fn test_concurrent_new_sessions_reserve_balanced_capacity() { + let policy = Arc::new(StickyLeastLoadedPolicy::new()); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + let barrier = std::sync::Barrier::new(60); + std::thread::scope(|scope| { + for i in 0..60 { + let policy = &policy; + let workers = &workers; + let barrier = &barrier; + scope.spawn(move || { + let headers = header_with_session(&format!("session-{i}")); + barrier.wait(); + let first = policy.select_worker_with_headers(workers, None, Some(&headers)); + assert_eq!( + first, + policy.select_worker_with_headers(workers, None, Some(&headers)) + ); + }); + } + }); + assert_eq!(policy.active_session_count(), 60); + for worker in &workers { + assert_eq!(policy.active_count_for(worker.url()), 60 / workers.len()); + } + } + + #[test] + fn test_expired_session_rebinds_before_next_sweep() { + let policy = StickyLeastLoadedPolicy::with_expiration_secs(1); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + for session in ["expired", "active"] { + policy.select_worker_with_headers( + &workers[..1], + None, + Some(&header_with_session(session)), + ); + } + policy + .state + .lock() + .unwrap() + .sessions + .get_mut("expired") + .unwrap() + .last_access = Instant::now() - Duration::from_secs(2); + + assert_eq!( + policy.select_worker_with_headers( + &workers, + None, + Some(&header_with_session("expired")) + ), + Some(1), + "expired affinity must not be refreshed before the global sweep" + ); + assert_eq!(policy.active_count_for(workers[0].url()), 1); + assert_eq!(policy.active_count_for(workers[1].url()), 1); + } + + #[test] + fn test_repeated_finish_preserves_other_sessions() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000"]); + for session in ["finished", "active"] { + policy.select_worker_with_headers(&workers, None, Some(&header_with_session(session))); + } + policy.finish_session("finished"); + policy.finish_session("finished"); + assert_eq!(policy.active_session_count(), 1); + assert_eq!(policy.active_count_for(workers[0].url()), 1); + } + + #[test] + fn test_removed_worker_rebinds_session() { + let policy = StickyLeastLoadedPolicy::new(); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let headers = header_with_session("session"); + let original = policy + .select_worker_with_headers(&workers, None, Some(&headers)) + .unwrap(); + let remaining = vec![workers[1 - original].clone()]; + assert_eq!( + policy.select_worker_with_headers(&remaining, None, Some(&headers)), + Some(0) + ); + assert_eq!(policy.active_count_for(workers[original].url()), 0); + assert_eq!(policy.active_count_for(remaining[0].url()), 1); + } +} diff --git a/src/protocols/spec.rs b/src/protocols/spec.rs index 3ed54dc8..39003d26 100644 --- a/src/protocols/spec.rs +++ b/src/protocols/spec.rs @@ -1921,6 +1921,14 @@ pub struct GenerateRequest { #[serde(skip_serializing_if = "Option::is_none")] pub session_params: Option>, + /// Legacy session identifier + #[serde(skip_serializing_if = "Option::is_none")] + pub session_id: Option, + + /// Legacy user identifier + #[serde(skip_serializing_if = "Option::is_none")] + pub user_id: Option, + /// Return model hidden states #[serde(default)] pub return_hidden_states: bool, diff --git a/src/routers/http/router.rs b/src/routers/http/router.rs index d39d3326..c6bd7e2a 100644 --- a/src/routers/http/router.rs +++ b/src/routers/http/router.rs @@ -622,7 +622,24 @@ impl Router { ) -> Response { let start = Instant::now(); let is_stream = typed_req.is_stream(); - let text = typed_req.extract_text_for_routing(); + let policy = match model_id { + Some(model) => self.policy_registry.get_policy_or_default(model), + None => self.policy_registry.get_default_policy(), + }; + let text = if policy.needs_request_body() { + match serde_json::to_string(typed_req) { + Ok(body) => body, + Err(error) => { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Serialization error: {error}"), + ) + .into_response(); + } + } + } else { + typed_req.extract_text_for_routing() + }; // Fall back to the body's `model` field when the caller doesn't pass one, but // only use it as a routing filter when the registry has already indexed that diff --git a/src/routers/http/vllm_pd_router.rs b/src/routers/http/vllm_pd_router.rs index 4a449224..68c3f9c8 100644 --- a/src/routers/http/vllm_pd_router.rs +++ b/src/routers/http/vllm_pd_router.rs @@ -614,6 +614,7 @@ impl VllmPDRouter { instances: &[(String, String)], is_prefill: bool, request_text: Option<&str>, + request_headers: Option<&HashMap>, ) -> Option { if instances.is_empty() { return None; @@ -630,7 +631,7 @@ impl VllmPDRouter { }; // Use policy to select worker - policy.select_worker(&workers, request_text) + policy.select_worker_with_headers(&workers, request_text, request_headers) } /// Process vLLM request using pure service discovery @@ -673,22 +674,40 @@ impl VllmPDRouter { // Use policy-based load balancing to select prefill and decode workers let request_text = serde_json::to_string(&request_json).ok(); let request_str = request_text.as_deref(); + let request_headers: Option> = headers.map(|h| { + h.iter() + .filter_map(|(name, value)| { + value + .to_str() + .ok() + .map(|v| (name.as_str().to_lowercase(), v.to_string())) + }) + .collect() + }); - let prefill_idx = - match self.select_worker_with_policy(&prefill_instances, true, request_str) { - Some(idx) => idx, - None => { - RouterMetrics::record_pd_error("server_selection"); - return ( - axum::http::StatusCode::SERVICE_UNAVAILABLE, - "Prefill policy failed to select a worker".to_string(), - ) - .into_response(); - } - }; + let prefill_idx = match self.select_worker_with_policy( + &prefill_instances, + true, + request_str, + request_headers.as_ref(), + ) { + Some(idx) => idx, + None => { + RouterMetrics::record_pd_error("server_selection"); + return ( + axum::http::StatusCode::SERVICE_UNAVAILABLE, + "Prefill policy failed to select a worker".to_string(), + ) + .into_response(); + } + }; - let decode_idx = match self.select_worker_with_policy(&decode_instances, false, request_str) - { + let decode_idx = match self.select_worker_with_policy( + &decode_instances, + false, + request_str, + request_headers.as_ref(), + ) { Some(idx) => idx, None => { RouterMetrics::record_pd_error("server_selection"); @@ -2552,6 +2571,80 @@ mod tests { use super::*; use serde_json::json; + #[tokio::test] + async fn test_discovered_pd_preserves_header_session_across_turns() { + use crate::config::RouterConfig; + use crate::policies::StickyLeastLoadedPolicy; + use crate::server::AppContext; + + let context = Arc::new( + AppContext::new( + RouterConfig::default(), + reqwest::Client::new(), + 10, + None, + vec![], + ) + .unwrap(), + ); + context + .policy_registry + .set_prefill_policy(Arc::new(StickyLeastLoadedPolicy::new())); + context + .policy_registry + .set_decode_policy(Arc::new(StickyLeastLoadedPolicy::new())); + let router = VllmPDRouter::new(vec![], vec![], None, &context) + .await + .unwrap(); + let mut servers = Vec::new(); + for service_type in [ServiceType::Prefill, ServiceType::Decode] { + for _ in 0..2 { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap().to_string(); + router.service_registry.register_service( + address.clone(), + address, + service_type.clone(), + ); + let app = axum::Router::new().route( + "/v1/completions", + axum::routing::post(|| async { + axum::Json(json!({"choices": [], "kv_transfer_params": {}})) + }), + ); + servers.push(tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + })); + } + } + + let mut headers = HeaderMap::new(); + headers.insert("x-session-id", "multi-turn-session".parse().unwrap()); + for prompt in ["first turn", "different second turn"] { + let response = router + .process_vllm_request( + json!({"model": "test", "prompt": prompt, "max_tokens": 2}), + "/v1/completions", + Some(&headers), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + } + for policy in [ + context.policy_registry.get_prefill_policy(), + context.policy_registry.get_decode_policy(), + ] { + let policy = policy + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(policy.active_session_count(), 1); + } + for server in servers { + server.abort(); + } + } + #[test] fn test_discovery_health_requires_prefill_and_decode_workers() { assert!(!discovery_is_ready(0, 0)); diff --git a/src/server.rs b/src/server.rs index 9d1276c9..e90b73e7 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1009,6 +1009,31 @@ async fn get_loads(State(state): State>, headers: http::HeaderMap) state.router.get_worker_loads().await } +#[derive(Deserialize)] +struct FinishSessionQuery { + session_id: String, +} + +/// POST /finish_session?session_id= +/// +/// Marks a session (e.g. an RL trajectory) as finished so that session-aware +/// routing policies (e.g. `sticky_least_loaded`) can release the +/// active-session assignment. For policies that don't track +/// sessions, this is a no-op. Unknown session ids are ignored. +async fn finish_session( + State(state): State>, + Query(FinishSessionQuery { session_id }): Query, + headers: http::HeaderMap, +) -> Response { + if let Err(response) = authorize_request(&state, &headers).await { + return response; + } + + state.context.policy_registry.finish_session(&session_id); + + Json(json!({ "status": "ok", "session_id": session_id })).into_response() +} + // ---------- Worker management endpoints (RESTful) ---------- /// POST /workers - Add a new worker with full configuration @@ -1246,7 +1271,8 @@ pub fn build_app_with_request_tracing( .route("/remove_worker", post(remove_worker)) .route("/list_workers", get(list_workers)) .route("/flush_cache", post(flush_cache)) - .route("/get_loads", get(get_loads)); + .route("/get_loads", get(get_loads)) + .route("/finish_session", post(finish_session)); // Worker management routes let worker_routes = Router::new() diff --git a/tests/api_endpoints_test.rs b/tests/api_endpoints_test.rs index 51c671c2..7c3f8ee3 100644 --- a/tests/api_endpoints_test.rs +++ b/tests/api_endpoints_test.rs @@ -921,6 +921,172 @@ mod worker_management_tests { #[cfg(test)] mod router_policy_tests { use super::*; + use vllm_router_rs::policies::StickyLeastLoadedPolicy; + use vllm_router_rs::server::{build_app, AppContext, AppState}; + + #[tokio::test] + async fn test_sticky_least_loaded_typed_sessions_and_release() { + let mut workers = Vec::new(); + let mut urls = Vec::new(); + for _ in 0..2 { + let mut worker = MockWorker::new(MockWorkerConfig { + port: 0, + worker_type: WorkerType::Regular, + health_status: HealthStatus::Healthy, + response_delay_ms: 0, + fail_rate: 0.0, + }); + urls.push(worker.start().await.unwrap()); + workers.push(worker); + } + let config = RouterConfig { + mode: RoutingMode::Regular { + worker_urls: urls.clone(), + }, + policy: PolicyConfig::StickyLeastLoaded, + api_key_validation_urls: vec!["http://127.0.0.1:1".to_string()], + ..RouterConfig::default() + }; + let context: Arc = common::create_test_context(config); + context + .api_key_cache + .write() + .await + .insert("test-token".to_string(), true); + let router = Arc::from(RouterFactory::create_router(&context).await.unwrap()); + let app = build_app( + Arc::new(AppState { + router, + context: context.clone(), + concurrency_queue_tx: None, + router_manager: None, + }), + 1024 * 1024, + vec![], + vec![], + true, + ); + let policy = context.policy_registry.get_default_policy(); + let policy = policy + .as_any() + .downcast_ref::() + .unwrap(); + + for session in ["body-a", "body-b", "body-a"] { + let request = Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header(CONTENT_TYPE, "application/json") + .header("authorization", "Bearer test-token") + .body(Body::from( + json!({ + "messages": [{"role": "user", "content": "hello"}], + "session_params": {"session_id": session} + }) + .to_string(), + )) + .unwrap(); + assert_eq!( + app.clone().oneshot(request).await.unwrap().status(), + StatusCode::OK + ); + } + assert_eq!(policy.active_session_count(), 2); + for url in &urls { + assert_eq!(policy.active_count_for(url), 1); + } + + for (uri, body, active_sessions) in [ + ( + "/generate", + json!({"prompt": "hello", "session_id": "generate-session"}), + 3, + ), + ( + "/generate", + json!({"prompt": "hello", "user_id": "generate-user"}), + 4, + ), + ( + "/v1/responses", + json!({"input": "hello", "session_id": "responses-session"}), + 5, + ), + ( + "/v1/responses", + json!({"input": "hello", "user_id": "responses-user"}), + 6, + ), + ] { + let request = Request::builder() + .method("POST") + .uri(uri) + .header(CONTENT_TYPE, "application/json") + .header("authorization", "Bearer test-token") + .body(Body::from(body.to_string())) + .unwrap(); + assert_eq!( + app.clone().oneshot(request).await.unwrap().status(), + StatusCode::OK + ); + assert_eq!(policy.active_session_count(), active_sessions); + } + + let request = Request::builder() + .method("POST") + .uri("/v1/completions") + .header(CONTENT_TYPE, "application/json") + .header("authorization", "Bearer test-token") + .header("x-session-id", "header-session") + .body(Body::from( + json!({"prompt": "hello", "session_id": "ignored-body"}).to_string(), + )) + .unwrap(); + assert_eq!( + app.clone().oneshot(request).await.unwrap().status(), + StatusCode::OK + ); + assert_eq!(policy.active_session_count(), 7); + + let unauthorized = Request::builder() + .method("POST") + .uri("/finish_session?session_id=header-session") + .body(Body::empty()) + .unwrap(); + assert_eq!( + app.clone().oneshot(unauthorized).await.unwrap().status(), + StatusCode::UNAUTHORIZED + ); + assert_eq!(policy.active_session_count(), 7); + + for (session, remaining) in [ + ("ignored-body", 7), + ("header-session", 6), + ("header-session", 6), + ("generate-session", 5), + ("generate-user", 4), + ("responses-session", 3), + ("responses-user", 2), + ("body-a", 1), + ("body-b", 0), + ] { + let request = Request::builder() + .method("POST") + .uri(format!("/finish_session?session_id={session}")) + .header("authorization", "Bearer test-token") + .body(Body::empty()) + .unwrap(); + assert_eq!( + app.clone().oneshot(request).await.unwrap().status(), + StatusCode::OK + ); + assert_eq!(policy.active_session_count(), remaining); + } + assert_eq!(policy.active_session_count(), 0); + for worker in &mut workers { + worker.stop().await; + } + } #[tokio::test] async fn test_random_policy() { diff --git a/tests/benchmark_integration.rs b/tests/benchmark_integration.rs index 80b42678..a046bfe3 100644 --- a/tests/benchmark_integration.rs +++ b/tests/benchmark_integration.rs @@ -23,6 +23,8 @@ fn default_generate_request() -> GenerateRequest { // vLLM Extensions lora_path: None, session_params: None, + session_id: None, + user_id: None, return_hidden_states: false, rid: None, } diff --git a/tests/test_openai_routing.rs b/tests/test_openai_routing.rs index a2990da5..8ddf824b 100644 --- a/tests/test_openai_routing.rs +++ b/tests/test_openai_routing.rs @@ -194,6 +194,8 @@ async fn test_unsupported_endpoints() { return_logprob: false, lora_path: None, session_params: None, + session_id: None, + user_id: None, return_hidden_states: false, rid: None, }; From f64492fa6d3065e6bbae522be6de6fcce9b36197 Mon Sep 17 00:00:00 2001 From: Mika Senghaas <70984473+mikasenghaas@users.noreply.github.com> Date: Wed, 16 Sep 2026 04:46:29 +0000 Subject: [PATCH 2/2] chore: integrate sticky routing into fork Pass the fork-specific optional run_id argument in the upstream prefill/decode regression test. Upgrade Debian runtime packages so the container image includes current security fixes and passes the Trivy scan. Bound sticky session state with graceful stateless fallback and capacity-triggered expiry, balance stateless requests across tied replicas, keep per-request headers from masking stable session IDs, resolve model-specific routing requirements before serializing typed requests, and preserve supported legacy session fields. --- Dockerfile.router | 5 +- README.md | 19 ++- benches/request_processing.rs | 1 + src/policies/hash_key.rs | 32 ++++- src/policies/sticky_least_loaded.rs | 189 ++++++++++++++++++++++++++-- src/protocols/spec.rs | 12 ++ src/routers/http/router.rs | 17 ++- src/routers/http/vllm_pd_router.rs | 1 + tests/benchmark_integration.rs | 3 + tests/responses_api_test.rs | 8 ++ tests/test_openai_routing.rs | 1 + 11 files changed, 254 insertions(+), 34 deletions(-) diff --git a/Dockerfile.router b/Dockerfile.router index 9e2ef4e4..6af683ed 100644 --- a/Dockerfile.router +++ b/Dockerfile.router @@ -12,7 +12,10 @@ RUN cargo build --release --bin vllm-router # Stage 2: Minimal runtime image FROM debian:bookworm-slim -RUN apt-get update && apt-get install -y ca-certificates && rm -rf /var/lib/apt/lists/* +RUN apt-get update \ + && apt-get upgrade -y \ + && apt-get install -y --no-install-recommends ca-certificates \ + && rm -rf /var/lib/apt/lists/* COPY --from=builder /build/target/release/vllm-router /usr/local/bin/vllm-router diff --git a/README.md b/README.md index 7cb4c1a8..c3999dba 100644 --- a/README.md +++ b/README.md @@ -228,10 +228,11 @@ atomically; ties use rendezvous hashing. Existing sessions keep their healthy worker, even if session counts later become uneven. Unavailable workers cause reassignment on the next request. Prefill and decode pools track sessions separately. -Session IDs use the same header priority as `consistent_hash`, then the JSON fields -`session_params.session_id`, `user`, `session_id`, and `user_id`. Prefer -`X-Session-ID` for multi-turn work; a per-request ID creates a separate session for -each request. Requests without an explicit identifier do not reserve a session. +Session IDs prefer stable headers (`x-session-id`, `x-user-id`, `x-tenant-id`, +and `x-correlation-id`), then the JSON fields `session_params.session_id`, `user`, +`session_id`, and `user_id`. Per-request `x-request-id` and `x-trace-id` headers +remain available to `consistent_hash` but do not reserve sticky sessions. Requests +without an explicit session identifier do not reserve a session. Release a session after its final request completes: @@ -254,8 +255,14 @@ are checked for expiry before refreshing their affinity. State is local to one router process and is lost on restart. Multiple routers need consistent ingress routing and session release on each router that tracked the session. The policy keeps one entry per active session and serializes assignment under a -mutex, so callers should release sessions -promptly and choose a TTL appropriate for their workload. +mutex, so callers should release sessions promptly and choose a TTL appropriate +for their workload. + +Each policy instance retains at most 100,000 sessions and accepts session IDs up +to 256 bytes. Override these limits with `VLLM_ROUTER_SLL_MAX_SESSIONS` and +`VLLM_ROUTER_SLL_MAX_SESSION_ID_BYTES`. Requests that introduce a session beyond +either limit continue to route least-loaded without affinity; existing tracked +sessions remain sticky. ## Advanced Features diff --git a/benches/request_processing.rs b/benches/request_processing.rs index 96d7a87d..928a2598 100644 --- a/benches/request_processing.rs +++ b/benches/request_processing.rs @@ -20,6 +20,7 @@ fn default_generate_request() -> GenerateRequest { // VLLM Extensions lora_path: None, session_params: None, + user: None, session_id: None, user_id: None, return_hidden_states: false, diff --git a/src/policies/hash_key.rs b/src/policies/hash_key.rs index cd950250..e3dcb138 100644 --- a/src/policies/hash_key.rs +++ b/src/policies/hash_key.rs @@ -7,8 +7,8 @@ use super::RequestHeaders; use crate::policies::ConsistentHashPolicy; use tracing::debug; -/// HTTP header names to check for session ID (case-insensitive, checked in order) -pub(crate) const SESSION_HEADER_NAMES: &[&str] = &[ +/// HTTP header names used by consistent hashing (case-insensitive, checked in order). +pub(crate) const HASH_HEADER_NAMES: &[&str] = &[ "x-session-id", "x-user-id", "x-tenant-id", @@ -17,6 +17,14 @@ pub(crate) const SESSION_HEADER_NAMES: &[&str] = &[ "x-trace-id", ]; +/// Headers whose values are expected to remain stable for a complete session. +const SESSION_HEADER_NAMES: &[&str] = &[ + "x-session-id", + "x-user-id", + "x-tenant-id", + "x-correlation-id", +]; + /// Extract hash key with priority: HTTP headers > body fields > request content hash /// /// Priority order: @@ -61,8 +69,8 @@ pub(crate) fn extract_hash_key( /// session id so that a matching `finish_session(session_id)` call can later /// release the session. /// -/// Lookup order mirrors [`extract_hash_key`]: -/// 1. HTTP headers: x-session-id, x-user-id, x-tenant-id, x-correlation-id, x-request-id, x-trace-id +/// Lookup order: +/// 1. Stable HTTP headers: x-session-id, x-user-id, x-tenant-id, x-correlation-id /// 2. Body: session_params.session_id (nested) /// 3. Body: user (OpenAI format) /// 4. Body: session_id (legacy) @@ -98,7 +106,7 @@ pub(crate) fn extract_session_id( /// Extract hash key from HTTP headers pub(crate) fn extract_hash_key_from_headers(headers: &RequestHeaders) -> Option { - for header_name in SESSION_HEADER_NAMES { + for header_name in HASH_HEADER_NAMES { if let Some(value) = headers.get(*header_name) { if !value.is_empty() { debug!( @@ -546,6 +554,20 @@ mod tests { ); } + #[test] + fn test_extract_session_id_body_precedes_request_headers() { + let mut headers = HashMap::new(); + headers.insert("x-request-id".to_string(), "request-1".to_string()); + headers.insert("x-trace-id".to_string(), "trace-1".to_string()); + let body = r#"{"session_id": "body-session"}"#; + + assert_eq!( + extract_session_id(Some(body), Some(&headers)), + Some("body-session".to_string()) + ); + assert_eq!(extract_session_id(None, Some(&headers)), None); + } + #[test] fn test_extract_session_id_from_body() { let body = r#"{"session_id": "legacy123", "prompt": "hi"}"#; diff --git a/src/policies/sticky_least_loaded.rs b/src/policies/sticky_least_loaded.rs index 5a727002..fd80bb71 100644 --- a/src/policies/sticky_least_loaded.rs +++ b/src/policies/sticky_least_loaded.rs @@ -46,6 +46,18 @@ pub const DEFAULT_SESSION_EXPIRATION_SECS: u64 = 7200; /// Environment variable to override the session expiration (in seconds). pub const SESSION_EXPIRATION_ENV: &str = "VLLM_ROUTER_SLL_SESSION_EXPIRATION_IN_S"; +/// Default maximum number of sessions retained by one policy instance. +pub const DEFAULT_MAX_SESSIONS: usize = 100_000; + +/// Environment variable to override the maximum retained session count. +pub const MAX_SESSIONS_ENV: &str = "VLLM_ROUTER_SLL_MAX_SESSIONS"; + +/// Default maximum session identifier size in bytes. +pub const DEFAULT_MAX_SESSION_ID_BYTES: usize = 256; + +/// Environment variable to override the maximum session identifier size. +pub const MAX_SESSION_ID_BYTES_ENV: &str = "VLLM_ROUTER_SLL_MAX_SESSION_ID_BYTES"; + /// Minimum interval between full expiration sweeps to bound per-request cost. const SWEEP_INTERVAL_SECS: u64 = 60; @@ -64,6 +76,8 @@ struct SllState { sessions: HashMap, /// replica url -> number of active sessions active_counts: HashMap, + /// Round-robin cursor for requests without a session identifier. + stateless_cursor: usize, /// Timestamp of the last expiration sweep. last_sweep: Instant, } @@ -73,24 +87,37 @@ struct SllState { pub struct StickyLeastLoadedPolicy { state: Mutex, session_expiration: Duration, + max_sessions: usize, + max_session_id_bytes: usize, } impl StickyLeastLoadedPolicy { /// Create a new policy, reading the session expiration from the environment /// (falling back to [`DEFAULT_SESSION_EXPIRATION_SECS`]). pub fn new() -> Self { - Self::with_expiration_secs(Self::expiration_secs_from_env()) + Self::with_limits( + Self::expiration_secs_from_env(), + Self::usize_from_env(MAX_SESSIONS_ENV, DEFAULT_MAX_SESSIONS), + Self::usize_from_env(MAX_SESSION_ID_BYTES_ENV, DEFAULT_MAX_SESSION_ID_BYTES), + ) } /// Create a new policy with an explicit session expiration (in seconds). pub fn with_expiration_secs(secs: u64) -> Self { + Self::with_limits(secs, DEFAULT_MAX_SESSIONS, DEFAULT_MAX_SESSION_ID_BYTES) + } + + fn with_limits(expiration_secs: u64, max_sessions: usize, max_session_id_bytes: usize) -> Self { Self { state: Mutex::new(SllState { sessions: HashMap::new(), active_counts: HashMap::new(), + stateless_cursor: 0, last_sweep: Instant::now(), }), - session_expiration: Duration::from_secs(secs), + session_expiration: Duration::from_secs(expiration_secs), + max_sessions, + max_session_id_bytes, } } @@ -110,6 +137,22 @@ impl StickyLeastLoadedPolicy { } } + fn usize_from_env(name: &str, default: usize) -> usize { + match std::env::var(name) { + Ok(value) => match value.trim().parse::() { + Ok(parsed) => parsed, + Err(_) => { + warn!( + "Invalid {} value '{}', using default {}", + name, value, default + ); + default + } + }, + Err(_) => default, + } + } + /// Decrement the active-session count for a replica, removing the entry when /// it reaches zero. fn decrement_count(counts: &mut HashMap, worker_url: &str) { @@ -122,12 +165,7 @@ impl StickyLeastLoadedPolicy { } /// Remove sessions that have not been accessed within the expiration window. - /// Runs at most once per [`SWEEP_INTERVAL_SECS`] to bound cost. fn sweep_expired(state: &mut SllState, expiration: Duration, now: Instant) { - if now.duration_since(state.last_sweep) < Duration::from_secs(SWEEP_INTERVAL_SECS) { - return; - } - let expired: Vec = state .sessions .iter() @@ -145,6 +183,14 @@ impl StickyLeastLoadedPolicy { state.last_sweep = now; } + /// Sweep at most once per [`SWEEP_INTERVAL_SECS`] during normal routing. + fn sweep_expired_if_due(state: &mut SllState, expiration: Duration, now: Instant) { + if now.duration_since(state.last_sweep) < Duration::from_secs(SWEEP_INTERVAL_SECS) { + return; + } + Self::sweep_expired(state, expiration, now); + } + /// Pick the least-loaded replica among `candidates`, breaking ties with /// consistent (rendezvous) hashing of `tie_break_key`. /// @@ -181,6 +227,39 @@ impl StickyLeastLoadedPolicy { best.map(|(idx, _)| idx).unwrap_or(candidates[0].0) } + fn select_least_loaded_round_robin( + counts: &HashMap, + candidates: &[(usize, String)], + offset: usize, + ) -> usize { + let min_count = candidates + .iter() + .map(|(_, url)| counts.get(url).copied().unwrap_or(0)) + .min() + .unwrap_or(0); + let tied_count = candidates + .iter() + .filter(|(_, url)| counts.get(url).copied().unwrap_or(0) == min_count) + .count(); + let target = offset % tied_count; + candidates + .iter() + .filter(|(_, url)| counts.get(url).copied().unwrap_or(0) == min_count) + .nth(target) + .map(|(idx, _)| *idx) + .unwrap_or(candidates[0].0) + } + + fn select_stateless(state: &mut SllState, candidates: &[(usize, String)]) -> usize { + let idx = Self::select_least_loaded_round_robin( + &state.active_counts, + candidates, + state.stateless_cursor, + ); + state.stateless_cursor = state.stateless_cursor.wrapping_add(1); + idx + } + fn record_selection(&self, worker: &Arc) { worker.increment_processed(); RouterMetrics::record_processed_request(worker.url()); @@ -221,17 +300,26 @@ impl LoadBalancingPolicy for StickyLeastLoadedPolicy { .map(|&idx| (idx, workers[idx].url().to_string())) .collect(); - let session_id = hash_key::extract_session_id(request_text, headers); + let session_id = match hash_key::extract_session_id(request_text, headers) { + Some(id) if id.len() > self.max_session_id_bytes => { + debug!( + "SLL: session identifier exceeds {} bytes; routing without affinity", + self.max_session_id_bytes + ); + None + } + session_id => session_id, + }; + let mut state = self.state.lock().unwrap(); let now = Instant::now(); - Self::sweep_expired(&mut state, self.session_expiration, now); + Self::sweep_expired_if_due(&mut state, self.session_expiration, now); // Requests without a session id: load-balance but don't record state. let session_id = match session_id { Some(id) => id, None => { - let tie_break = request_text.unwrap_or(""); - let idx = Self::select_least_loaded(&state.active_counts, &candidates, tie_break); + let idx = Self::select_stateless(&mut state, &candidates); drop(state); self.record_selection(&workers[idx]); debug!("SLL: stateless request routed to '{}'", workers[idx].url()); @@ -271,6 +359,21 @@ impl LoadBalancingPolicy for StickyLeastLoadedPolicy { } // New session: assign to least-loaded replica (tie-broken by hashing). + if state.sessions.len() >= self.max_sessions { + Self::sweep_expired(&mut state, self.session_expiration, now); + } + if state.sessions.len() >= self.max_sessions { + let idx = Self::select_stateless(&mut state, &candidates); + drop(state); + self.record_selection(&workers[idx]); + debug!( + "SLL: {}-session limit reached; request routed without affinity to '{}'", + self.max_sessions, + workers[idx].url() + ); + return Some(idx); + } + let idx = Self::select_least_loaded(&state.active_counts, &candidates, &session_id); let worker_url = workers[idx].url().to_string(); state.sessions.insert( @@ -325,6 +428,7 @@ impl LoadBalancingPolicy for StickyLeastLoadedPolicy { let mut state = self.state.lock().unwrap(); state.sessions.clear(); state.active_counts.clear(); + state.stateless_cursor = 0; state.last_sweep = Instant::now(); info!("SLL: policy reset - all sessions cleared"); } @@ -448,12 +552,50 @@ mod tests { let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); let body = r#"{"prompt": "no session here"}"#; - let idx = policy.select_worker_with_headers(&workers, Some(body), None); - assert!(idx.is_some()); + let routed: Vec<_> = (0..4) + .map(|_| policy.select_worker_with_headers(&workers, Some(body), None)) + .collect(); + assert_eq!(routed, vec![Some(0), Some(1), Some(0), Some(1)]); // No session identifier -> nothing tracked. assert_eq!(policy.active_session_count(), 0); } + #[test] + fn test_routes_oversized_session_id_without_affinity() { + let policy = StickyLeastLoadedPolicy::with_limits(7200, 10, 8); + let workers = make_workers(&["http://w1:8000"]); + let headers = header_with_session("ninebytes"); + + assert_eq!( + policy.select_worker_with_headers(&workers, None, Some(&headers)), + Some(0) + ); + assert_eq!(policy.active_session_count(), 0); + } + + #[test] + fn test_routes_new_session_without_affinity_at_capacity() { + let policy = StickyLeastLoadedPolicy::with_limits(7200, 1, 256); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + let admitted = header_with_session("admitted"); + let untracked = header_with_session("untracked"); + + let original = policy.select_worker_with_headers(&workers, None, Some(&admitted)); + assert!(original.is_some()); + let first = policy.select_worker_with_headers(&workers, None, Some(&untracked)); + let second = policy.select_worker_with_headers(&workers, None, Some(&untracked)); + assert!(first.is_some()); + assert!(second.is_some()); + assert_ne!(first, second); + assert_ne!(first, original); + assert_ne!(second, original); + assert_eq!( + policy.select_worker_with_headers(&workers, None, Some(&admitted)), + original + ); + assert_eq!(policy.active_session_count(), 1); + } + #[test] fn test_reassign_when_replica_unhealthy() { let policy = StickyLeastLoadedPolicy::new(); @@ -512,6 +654,27 @@ mod tests { ); } + #[test] + fn test_capacity_check_sweeps_expired_sessions() { + let policy = StickyLeastLoadedPolicy::with_limits(0, 1, 256); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + let expired = header_with_session("expired"); + let fresh = header_with_session("fresh"); + + policy + .select_worker_with_headers(&workers, None, Some(&expired)) + .unwrap(); + assert_eq!(policy.active_session_count(), 1); + + policy + .select_worker_with_headers(&workers, None, Some(&fresh)) + .unwrap(); + assert_eq!(policy.active_session_count(), 1); + + policy.finish_session("fresh"); + assert_eq!(policy.active_session_count(), 0); + } + #[test] fn test_concurrent_new_sessions_reserve_balanced_capacity() { let policy = Arc::new(StickyLeastLoadedPolicy::new()); diff --git a/src/protocols/spec.rs b/src/protocols/spec.rs index 39003d26..5f729f98 100644 --- a/src/protocols/spec.rs +++ b/src/protocols/spec.rs @@ -1183,6 +1183,14 @@ pub struct ResponsesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub user: Option, + /// Legacy session identifier + #[serde(skip_serializing_if = "Option::is_none")] + pub session_id: Option, + + /// Legacy user identifier + #[serde(skip_serializing_if = "Option::is_none")] + pub user_id: Option, + // ============= VLLM Extensions ============= /// Request ID #[serde(default = "generate_request_id")] @@ -1921,6 +1929,10 @@ pub struct GenerateRequest { #[serde(skip_serializing_if = "Option::is_none")] pub session_params: Option>, + /// OpenAI user identifier + #[serde(skip_serializing_if = "Option::is_none")] + pub user: Option, + /// Legacy session identifier #[serde(skip_serializing_if = "Option::is_none")] pub session_id: Option, diff --git a/src/routers/http/router.rs b/src/routers/http/router.rs index c6bd7e2a..71635b46 100644 --- a/src/routers/http/router.rs +++ b/src/routers/http/router.rs @@ -622,7 +622,14 @@ impl Router { ) -> Response { let start = Instant::now(); let is_stream = typed_req.is_stream(); - let policy = match model_id { + // Fall back to the body's `model` field when the caller doesn't pass one, but + // only use it as a routing filter when the registry has already indexed that + // model. This keeps compatibility for generic upstream model validation while + // still preventing known LoRA requests from being sent to workers that have not + // loaded the adapter. Run-scoped requests keep the body model as a hard filter. + let effective_model_id = Self::normalize_model_id(model_id) + .or_else(|| self.resolve_body_model_filter(route, typed_req.get_model(), run_id)); + let policy = match effective_model_id { Some(model) => self.policy_registry.get_policy_or_default(model), None => self.policy_registry.get_default_policy(), }; @@ -641,14 +648,6 @@ impl Router { typed_req.extract_text_for_routing() }; - // Fall back to the body's `model` field when the caller doesn't pass one, but - // only use it as a routing filter when the registry has already indexed that - // model. This keeps compatibility for generic upstream model validation while - // still preventing known LoRA requests from being sent to workers that have not - // loaded the adapter. Run-scoped requests keep the body model as a hard filter. - let effective_model_id = Self::normalize_model_id(model_id) - .or_else(|| self.resolve_body_model_filter(route, typed_req.get_model(), run_id)); - let response = RetryExecutor::execute_response_with_retry( &self.retry_config, // operation per attempt diff --git a/src/routers/http/vllm_pd_router.rs b/src/routers/http/vllm_pd_router.rs index 68c3f9c8..2eaf6ccd 100644 --- a/src/routers/http/vllm_pd_router.rs +++ b/src/routers/http/vllm_pd_router.rs @@ -2626,6 +2626,7 @@ mod tests { json!({"model": "test", "prompt": prompt, "max_tokens": 2}), "/v1/completions", Some(&headers), + None, ) .await; assert_eq!(response.status(), StatusCode::OK); diff --git a/tests/benchmark_integration.rs b/tests/benchmark_integration.rs index a046bfe3..6ac09493 100644 --- a/tests/benchmark_integration.rs +++ b/tests/benchmark_integration.rs @@ -23,6 +23,7 @@ fn default_generate_request() -> GenerateRequest { // vLLM Extensions lora_path: None, session_params: None, + user: None, session_id: None, user_id: None, return_hidden_states: false, @@ -215,6 +216,7 @@ fn test_benchmark_serialization_roundtrip() { let generate_req = GenerateRequest { text: Some("Test prompt".to_string()), + user: Some("test-user".to_string()), ..default_generate_request() }; @@ -226,6 +228,7 @@ fn test_benchmark_serialization_roundtrip() { assert_eq!(generate_req.text, deserialized.text); assert_eq!(generate_req.stream, deserialized.stream); assert_eq!(generate_req.return_logprob, deserialized.return_logprob); + assert_eq!(generate_req.user, deserialized.user); } #[test] diff --git a/tests/responses_api_test.rs b/tests/responses_api_test.rs index 366d53f5..aee78508 100644 --- a/tests/responses_api_test.rs +++ b/tests/responses_api_test.rs @@ -34,6 +34,8 @@ fn test_responses_request_creation() { top_p: Some(0.9), truncation: Truncation::Disabled, user: Some("test-user".to_string()), + session_id: None, + user_id: None, request_id: "resp_test123".to_string(), priority: 0, frequency_penalty: 0.0, @@ -75,6 +77,8 @@ fn test_sampling_params_conversion() { top_p: Some(0.95), truncation: Truncation::Auto, user: None, + session_id: None, + user_id: None, request_id: "resp_test456".to_string(), priority: 0, frequency_penalty: 0.1, @@ -187,6 +191,8 @@ fn test_json_serialization() { top_p: Some(0.8), truncation: Truncation::Auto, user: Some("test_user".to_string()), + session_id: Some("test-session".to_string()), + user_id: Some("test-user-id".to_string()), request_id: "resp_comprehensive_test".to_string(), priority: 1, frequency_penalty: 0.3, @@ -207,4 +213,6 @@ fn test_json_serialization() { assert!(parsed.background); assert!(parsed.stream); assert_eq!(parsed.tools.len(), 1); + assert_eq!(parsed.session_id.as_deref(), Some("test-session")); + assert_eq!(parsed.user_id.as_deref(), Some("test-user-id")); } diff --git a/tests/test_openai_routing.rs b/tests/test_openai_routing.rs index 8ddf824b..8596b059 100644 --- a/tests/test_openai_routing.rs +++ b/tests/test_openai_routing.rs @@ -194,6 +194,7 @@ async fn test_unsupported_endpoints() { return_logprob: false, lora_path: None, session_params: None, + user: None, session_id: None, user_id: None, return_hidden_states: false,