From 8b6fb5a8341d6d1033142d48f85b45bd51b6cd56 Mon Sep 17 00:00:00 2001 From: David <12414531+DavidBellamy@users.noreply.github.com> Date: Tue, 11 Aug 2026 03:42:40 -0700 Subject: [PATCH 1/2] fix(admission): expose bounded distribution headroom Signed-off-by: David <12414531+DavidBellamy@users.noreply.github.com> --- bindings/python/src/lib.rs | 13 + bindings/python/src/smg/router_args.py | 17 + bindings/python/tests/test_arg_parser.py | 12 + model_gateway/src/config/types.rs | 10 + model_gateway/src/config/validation.rs | 87 ++ model_gateway/src/main.rs | 33 +- model_gateway/src/policies/cache_aware.rs | 378 +++++- model_gateway/src/policies/mod.rs | 1 + model_gateway/src/policies/registry.rs | 19 +- .../src/routers/grpc/adaptive_admission.rs | 1045 ++++++++++++++++- .../grpc/common/stages/adaptive_admission.rs | 53 +- .../grpc/common/stages/request_execution.rs | 9 +- .../grpc/common/stages/worker_selection.rs | 265 +++-- model_gateway/src/routers/grpc/context.rs | 53 +- 14 files changed, 1891 insertions(+), 104 deletions(-) diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 0c29c9cfa..c4130af88 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -516,6 +516,8 @@ struct Router { capacity_credit_terminal_retention_secs: u64, capacity_credit_required: bool, priority_scheduler_adaptive_capacity: bool, + adaptive_admission_distribution_headroom_partitions: Vec, + adaptive_admission_distribution_headroom_partition_seed_cap: u32, } impl Router { @@ -831,6 +833,11 @@ impl Router { feedback_max_token_usage: self.adaptive_admission_feedback_max_token_usage, feedback_throughput_improvement_ratio: self .adaptive_admission_feedback_throughput_improvement_ratio, + distribution_headroom_partitions: self + .adaptive_admission_distribution_headroom_partitions + .clone(), + distribution_headroom_partition_seed_cap: self + .adaptive_admission_distribution_headroom_partition_seed_cap, }) .cors_allowed_origins(self.cors_allowed_origins.clone()) .retry_config(config::RetryConfig { @@ -1070,6 +1077,8 @@ impl Router { capacity_credit_terminal_retention_secs = 600, capacity_credit_required = false, priority_scheduler_adaptive_capacity = false, + adaptive_admission_distribution_headroom_partitions = vec![], + adaptive_admission_distribution_headroom_partition_seed_cap = 0, ))] #[expect(clippy::too_many_arguments)] #[expect( @@ -1225,6 +1234,8 @@ impl Router { capacity_credit_terminal_retention_secs: u64, capacity_credit_required: bool, priority_scheduler_adaptive_capacity: bool, + adaptive_admission_distribution_headroom_partitions: Vec, + adaptive_admission_distribution_headroom_partition_seed_cap: u32, ) -> PyResult { let mut all_urls = worker_urls.clone(); @@ -1394,6 +1405,8 @@ impl Router { capacity_credit_terminal_retention_secs, capacity_credit_required, priority_scheduler_adaptive_capacity, + adaptive_admission_distribution_headroom_partitions, + adaptive_admission_distribution_headroom_partition_seed_cap, }) } diff --git a/bindings/python/src/smg/router_args.py b/bindings/python/src/smg/router_args.py index 195ce0b85..acf6c8a12 100644 --- a/bindings/python/src/smg/router_args.py +++ b/bindings/python/src/smg/router_args.py @@ -143,6 +143,10 @@ class RouterArgs: adaptive_admission_feedback_max_waiting_requests_per_healthy_replica: int = 2 adaptive_admission_feedback_max_token_usage: float = 0.9 adaptive_admission_feedback_throughput_improvement_ratio: float = 0.02 + adaptive_admission_distribution_headroom_partitions: list[str] = dataclasses.field( + default_factory=list + ) + adaptive_admission_distribution_headroom_partition_seed_cap: int = 0 # Token bucket refill rate (tokens per second). If not set, defaults to max_concurrent_requests rate_limit_tokens_per_second: int | None = None # Cluster-wide requests-per-second ceiling. Requires mesh and the same value on every gateway. @@ -1022,6 +1026,19 @@ def add_cli_args( default=RouterArgs.adaptive_admission_feedback_throughput_improvement_ratio, help="Relative throughput gain required to raise the learned concurrency knee", ) + adaptive_admission_group.add_argument( + f"--{prefix}adaptive-admission-distribution-headroom-partitions", + type=str, + nargs="+", + default=[], + help="Exact admission partitions allowed to expose clean-worker headroom", + ) + adaptive_admission_group.add_argument( + f"--{prefix}adaptive-admission-distribution-headroom-partition-seed-cap", + type=int, + default=RouterArgs.adaptive_admission_distribution_headroom_partition_seed_cap, + help="Hard active distribution seed-lease cap per allowlisted partition", + ) # Retry configuration retry_group.add_argument( diff --git a/bindings/python/tests/test_arg_parser.py b/bindings/python/tests/test_arg_parser.py index 70ce03b30..84ab796b6 100644 --- a/bindings/python/tests/test_arg_parser.py +++ b/bindings/python/tests/test_arg_parser.py @@ -72,6 +72,8 @@ def test_default_values(self): assert args.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica == 2 assert args.adaptive_admission_feedback_max_token_usage == 0.9 assert args.adaptive_admission_feedback_throughput_improvement_ratio == 0.02 + assert args.adaptive_admission_distribution_headroom_partitions == [] + assert args.adaptive_admission_distribution_headroom_partition_seed_cap == 0 def test_parse_priority_scheduler_options(self): args = parse_router_args( @@ -146,6 +148,11 @@ def test_parse_adaptive_admission_options(self): "0.85", "--adaptive-admission-feedback-throughput-improvement-ratio", "0.03", + "--adaptive-admission-distribution-headroom-partitions", + "k3-prod", + "k3-canary", + "--adaptive-admission-distribution-headroom-partition-seed-cap", + "2", ] ) @@ -162,6 +169,11 @@ def test_parse_adaptive_admission_options(self): assert args.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica == 4 assert args.adaptive_admission_feedback_max_token_usage == 0.85 assert args.adaptive_admission_feedback_throughput_improvement_ratio == 0.03 + assert args.adaptive_admission_distribution_headroom_partitions == [ + "k3-prod", + "k3-canary", + ] + assert args.adaptive_admission_distribution_headroom_partition_seed_cap == 2 def test_parse_selector_valid(self): """Test parsing valid selector arguments.""" diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index 1e2214c6c..85b2939ab 100644 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -150,6 +150,14 @@ pub struct AdaptiveAdmissionConfig { /// concurrency knee upward. Near-equal throughput may move it downward. #[serde(default = "default_feedback_throughput_improvement_ratio")] pub feedback_throughput_improvement_ratio: f64, + /// Exact admission-partition names allowed to expose clean-worker + /// distribution headroom. An empty allowlist keeps the feature disabled. + #[serde(default)] + pub distribution_headroom_partitions: Vec, + /// Hard process-local ceiling on active distribution-headroom seed leases + /// in any allowlisted partition. Zero keeps the feature disabled. + #[serde(default)] + pub distribution_headroom_partition_seed_cap: u32, } impl Default for AdaptiveAdmissionConfig { @@ -169,6 +177,8 @@ impl Default for AdaptiveAdmissionConfig { default_feedback_max_waiting_requests_per_healthy_replica(), feedback_max_token_usage: default_feedback_max_token_usage(), feedback_throughput_improvement_ratio: default_feedback_throughput_improvement_ratio(), + distribution_headroom_partitions: Vec::new(), + distribution_headroom_partition_seed_cap: 0, } } } diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index 31dfe240e..01536b4ce 100755 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -772,6 +772,17 @@ impl ConfigValidator { .to_string(), }); } + if config.priority_scheduler_adaptive_capacity + && config + .adaptive_admission + .distribution_headroom_partition_seed_cap + > 0 + { + return Err(ConfigError::ValidationFailed { + reason: "distribution headroom requires a static scheduler ceiling; adaptive scheduler capacity may collapse before target selection" + .to_string(), + }); + } if config.worker_startup_timeout_secs == 0 { return Err(ConfigError::InvalidValue { @@ -865,6 +876,40 @@ impl ConfigValidator { reason: "Must be finite and in [0, 1]".to_string(), }); } + let mut distribution_partitions = std::collections::HashSet::new(); + for partition in &adaptive.distribution_headroom_partitions { + if partition.is_empty() || partition.trim() != partition || partition.len() > 128 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.distribution_headroom_partitions".to_string(), + value: partition.clone(), + reason: "Partition names must be non-empty, unpadded, and at most 128 bytes" + .to_string(), + }); + } + if !distribution_partitions.insert(partition) { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.distribution_headroom_partitions".to_string(), + value: partition.clone(), + reason: "Partition names must be unique exact matches".to_string(), + }); + } + } + if adaptive.distribution_headroom_partition_seed_cap > 0 { + if adaptive.distribution_headroom_partitions.is_empty() { + return Err(ConfigError::ValidationFailed { + reason: "distribution headroom seed capacity requires a non-empty exact partition allowlist" + .to_string(), + }); + } + if adaptive.mode != AdaptiveAdmissionMode::Enforce + || adaptive.strategy != AdaptiveAdmissionStrategy::EngineFeedback + { + return Err(ConfigError::ValidationFailed { + reason: "distribution headroom seed capacity requires adaptive admission mode=enforce and strategy=engine_feedback" + .to_string(), + }); + } + } Ok(()) } @@ -2045,9 +2090,51 @@ mod tests { assert!(config.capacity_credit_generation.is_none()); assert!(!config.capacity_credit_required); assert!(!config.priority_scheduler_adaptive_capacity); + assert!(config + .adaptive_admission + .distribution_headroom_partitions + .is_empty()); + assert_eq!( + config + .adaptive_admission + .distribution_headroom_partition_seed_cap, + 0 + ); assert!(ConfigValidator::validate(&config).is_ok()); } + #[test] + fn distribution_headroom_requires_exact_allowlist_and_enforced_feedback() { + let mut config = RouterConfig::default(); + config.adaptive_admission.distribution_headroom_partitions = vec!["k3".to_string()]; + assert!(ConfigValidator::validate(&config).is_ok()); + + config + .adaptive_admission + .distribution_headroom_partition_seed_cap = 2; + assert!(ConfigValidator::validate(&config).is_err()); + config.adaptive_admission.mode = AdaptiveAdmissionMode::Enforce; + config.adaptive_admission.strategy = AdaptiveAdmissionStrategy::EngineFeedback; + assert!(ConfigValidator::validate(&config).is_ok()); + + config.priority_scheduler_enabled = true; + config.priority_scheduler_adaptive_capacity = true; + assert!(ConfigValidator::validate(&config).is_err()); + config.priority_scheduler_adaptive_capacity = false; + assert!(ConfigValidator::validate(&config).is_ok()); + + config.adaptive_admission.distribution_headroom_partitions = vec![" k3".to_string()]; + assert!(ConfigValidator::validate(&config).is_err()); + config.adaptive_admission.distribution_headroom_partitions = + vec!["k3".to_string(), "k3".to_string()]; + assert!(ConfigValidator::validate(&config).is_err()); + config + .adaptive_admission + .distribution_headroom_partitions + .clear(); + assert!(ConfigValidator::validate(&config).is_err()); + } + #[test] fn capacity_credit_generation_requires_scheduler_service_key_and_preferred_identity() { let mut config = RouterConfig { diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 86c43170d..edf0a181d 100755 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -631,6 +631,21 @@ struct CliArgs { #[arg(long, default_value_t = 0.02, help_heading = "Adaptive Admission")] adaptive_admission_feedback_throughput_improvement_ratio: f64, + /// Exact admission partitions allowed to expose clean-worker distribution + /// headroom. Empty by default. + #[arg( + long, + value_delimiter = ',', + num_args = 1.., + help_heading = "Adaptive Admission" + )] + adaptive_admission_distribution_headroom_partitions: Vec, + + /// Hard process-local active seed-lease cap per allowlisted partition. + /// Zero disables distribution headroom. + #[arg(long, default_value_t = 0, help_heading = "Adaptive Admission")] + adaptive_admission_distribution_headroom_partition_seed_cap: u32, + // ==================== Tenant Rate Limit ==================== /// Enable per-tenant LLM token/request rate limiting. When unset /// (default), no rate limiter is constructed. @@ -1654,6 +1669,11 @@ impl CliArgs { feedback_max_token_usage: self.adaptive_admission_feedback_max_token_usage, feedback_throughput_improvement_ratio: self .adaptive_admission_feedback_throughput_improvement_ratio, + distribution_headroom_partitions: self + .adaptive_admission_distribution_headroom_partitions + .clone(), + distribution_headroom_partition_seed_cap: self + .adaptive_admission_distribution_headroom_partition_seed_cap, }) .tenant_rate_limit_enabled(self.tenant_rate_limit_enabled) .tenant_rate_limit_config(self.tenant_rate_limit_config.clone()) @@ -2027,7 +2047,7 @@ mod tests { fn engine_feedback_admission_options_flow_into_router_config() { let cli = cli_args_from(&[ "--adaptive-admission-mode", - "shadow", + "enforce", "--adaptive-admission-strategy", "engine-feedback", "--adaptive-admission-feedback-probe-requests-per-healthy-replica", @@ -2038,11 +2058,15 @@ mod tests { "0.85", "--adaptive-admission-feedback-throughput-improvement-ratio", "0.03", + "--adaptive-admission-distribution-headroom-partitions", + "k3-prod,k3-canary", + "--adaptive-admission-distribution-headroom-partition-seed-cap", + "2", ]); let router_config = cli.to_router_config(vec![], vec![]).unwrap(); let adaptive = router_config.adaptive_admission; - assert_eq!(adaptive.mode, AdaptiveAdmissionMode::Shadow); + assert_eq!(adaptive.mode, AdaptiveAdmissionMode::Enforce); assert_eq!(adaptive.strategy, AdaptiveAdmissionStrategy::EngineFeedback); assert_eq!(adaptive.feedback_probe_requests_per_healthy_replica, 3); assert_eq!( @@ -2051,6 +2075,11 @@ mod tests { ); assert_eq!(adaptive.feedback_max_token_usage, 0.85); assert_eq!(adaptive.feedback_throughput_improvement_ratio, 0.03); + assert_eq!( + adaptive.distribution_headroom_partitions, + ["k3-prod", "k3-canary"] + ); + assert_eq!(adaptive.distribution_headroom_partition_seed_cap, 2); } /// The multimodal transport flags must reach both `RouterConfig` and the diff --git a/model_gateway/src/policies/cache_aware.rs b/model_gateway/src/policies/cache_aware.rs index 5e6b598fa..a4ef5879c 100644 --- a/model_gateway/src/policies/cache_aware.rs +++ b/model_gateway/src/policies/cache_aware.rs @@ -60,7 +60,7 @@ use std::{ collections::HashMap, sync::{ - atomic::{AtomicBool, Ordering}, + atomic::{AtomicBool, AtomicU64, Ordering}, Arc, }, time::{Duration, Instant}, @@ -115,12 +115,79 @@ struct PrefixReplicationState { provisional_owner: String, } +#[derive(Debug, Clone)] +struct DistributionSeedPrefixState { + lease_id: u64, + active: bool, + last_transition: Instant, +} + +/// One atomic per-prefix expansion claim. Ordinary cache-aware requests never +/// observe or join this provisional target. Dropping the claim releases the +/// active slot while retaining a conservative cooldown marker. +#[derive(Debug)] +struct DistributionSeedPrefixReservation { + state: Arc>, + key: PrefixBudgetKey, + lease_id: u64, +} + +impl Drop for DistributionSeedPrefixReservation { + fn drop(&mut self) { + let Entry::Occupied(mut entry) = self.state.entry(self.key) else { + return; + }; + if entry.get().lease_id != self.lease_id || !entry.get().active { + return; + } + let state = entry.get_mut(); + state.active = false; + state.last_transition = Instant::now(); + } +} + #[derive(Debug)] struct PrefixOwnership { key: PrefixBudgetKey, owners: Vec, } +/// Routing-neutral worker capacity supplied by adaptive admission. Cache-aware +/// routing may use it to choose a clean non-owner, but it cannot mint capacity +/// or authorize dispatch. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SeedWorkerHeadroom { + pub(crate) worker_url: Arc, + pub(crate) worker_revision: u64, + pub(crate) issuable_slots: u16, +} + +/// Opaque cache-policy plan for one owner-pressure recovery dispatch. +/// +/// The target is not inserted into either prefix tree here. Only backend KV +/// events may make it an authoritative owner after the dispatch succeeds. +#[derive(Debug)] +pub(crate) struct OwnerPressureDispatchPlan { + target_worker_url: Arc, + target_worker_revision: u64, + expands_ownership: bool, + _prefix_reservation: Option, +} + +impl OwnerPressureDispatchPlan { + pub(crate) fn target_worker_url(&self) -> &str { + &self.target_worker_url + } + + pub(crate) fn target_worker_revision(&self) -> u64 { + self.target_worker_revision + } + + pub(crate) fn expands_ownership(&self) -> bool { + self.expands_ownership + } +} + #[derive(Debug, Clone)] struct TimedWorkerLoad { response: WorkerLoadResponse, @@ -203,6 +270,11 @@ pub struct CacheAwarePolicy { /// not the authoritative owner catalog, which always records every worker /// reported by the backend event stream. replication_state: Arc>, + /// Exact per-prefix serialization for scheduler-authorized clean-peer + /// seeds. Kept separate from `replication_state` so ordinary requests + /// cannot join a seed before backend KV events establish ownership. + distribution_seed_state: Arc>, + next_distribution_seed_lease_id: AtomicU64, _replication_gc_task: Option, } @@ -228,6 +300,8 @@ impl CacheAwarePolicy { let token_trees = Arc::new(DashMap::>::new()); let hash_index = Arc::new(DashMap::::new()); let replication_state = Arc::new(DashMap::::new()); + let distribution_seed_state = + Arc::new(DashMap::::new()); // Start background eviction thread if configured let eviction_task = if config.eviction_interval_secs > 0 { @@ -310,12 +384,18 @@ impl CacheAwarePolicy { && config.cache_owner_spill_cooldown_secs > 0 { let state = Arc::clone(&replication_state); + let seed_state = Arc::clone(&distribution_seed_state); let cooldown = Duration::from_secs(config.cache_owner_spill_cooldown_secs); let retention = cooldown.saturating_mul(4).max(Duration::from_secs(60)); Some(PeriodicTask::spawn( config.cache_owner_spill_cooldown_secs.max(1), "Prefix replication budget GC", - move || state.retain(|_, entry| entry.last_spill.elapsed() <= retention), + move || { + state.retain(|_, entry| entry.last_spill.elapsed() <= retention); + seed_state.retain(|_, entry| { + entry.active || entry.last_transition.elapsed() <= retention + }); + }, )) } else { None @@ -332,6 +412,8 @@ impl CacheAwarePolicy { populate_hash_index: AtomicBool::new(false), engine_loads: RwLock::new(HashMap::new()), replication_state, + distribution_seed_state, + next_distribution_seed_lease_id: AtomicU64::new(0), _replication_gc_task: replication_gc_task, } } @@ -1393,6 +1475,162 @@ impl CacheAwarePolicy { owners } + fn try_acquire_distribution_seed_prefix( + &self, + key: PrefixBudgetKey, + ) -> Option { + let cooldown = Duration::from_secs(self.config.cache_owner_spill_cooldown_secs); + if cooldown.is_zero() { + return None; + } + let lease_id = self + .next_distribution_seed_lease_id + .fetch_add(1, Ordering::Relaxed) + .wrapping_add(1); + let state = DistributionSeedPrefixState { + lease_id, + active: true, + last_transition: Instant::now(), + }; + match self.distribution_seed_state.entry(key) { + Entry::Occupied(mut entry) => { + if entry.get().active || entry.get().last_transition.elapsed() < cooldown { + return None; + } + entry.insert(state); + } + Entry::Vacant(entry) => { + entry.insert(state); + } + } + Some(DistributionSeedPrefixReservation { + state: Arc::clone(&self.distribution_seed_state), + key, + lease_id, + }) + } + + /// Plan one exact clean-peer dispatch when every matching cache owner has + /// no issuable engine slot. This is a read-only policy decision: admission + /// still requires a scheduler proof plus a separately acquired headroom + /// lease, and worker selection must bind and revalidate the exact target. + pub(crate) fn owner_pressure_dispatch_plan( + &self, + model_id: &str, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + headroom: &[SeedWorkerHeadroom], + ) -> Option { + if !self.config.engine_load { + return None; + } + + let tokens = info.tokens?; + let healthy_indices = super::get_healthy_worker_indices(workers); + // A seed may expand only ownership reported by the backend KV event + // index. The approximate trees and the normal routing path's + // provisional spill entry are deliberately excluded. + let monitor = self.kv_monitor.read(); + let indexer = monitor.as_ref()?.get_indexer(model_id)?; + if indexer.current_size() == 0 { + return None; + } + let block_size = monitor + .as_ref()? + .block_size(model_id) + .unwrap_or(self.config.block_size); + let (known_owners, matched_blocks) = + Self::score_overlap_with_depth(workers, tokens, &healthy_indices, &indexer, block_size); + if matched_blocks == 0 { + return None; + } + let matched_tokens = matched_blocks.saturating_mul(block_size).min(tokens.len()); + let prefix_key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path(model_id), + prefix_hash: kv_index::hash_token_path(&tokens[..matched_tokens]), + kind: PrefixKind::Token, + }; + if known_owners.is_empty() { + return None; + } + + let eligible_owners = + SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &known_owners); + if eligible_owners.is_empty() { + return None; + } + let mut owner_capacity = Vec::with_capacity(eligible_owners.len()); + for idx in eligible_owners { + let capacity = headroom.iter().find(|candidate| { + candidate.worker_url.as_ref() == workers[idx].url() + && candidate.worker_revision == workers[idx].revision() + })?; + owner_capacity.push((idx, capacity.issuable_slots)); + } + owner_capacity.sort_unstable_by(|(left_idx, left_slots), (right_idx, right_slots)| { + right_slots + .cmp(left_slots) + .then_with(|| workers[*left_idx].load().cmp(&workers[*right_idx].load())) + .then_with(|| workers[*left_idx].url().cmp(workers[*right_idx].url())) + }); + if let Some(&(target_idx, _)) = owner_capacity + .first() + .filter(|(_, issuable_slots)| *issuable_slots > 0) + { + return Some(OwnerPressureDispatchPlan { + target_worker_url: Arc::from(workers[target_idx].url()), + target_worker_revision: workers[target_idx].revision(), + expands_ownership: false, + _prefix_reservation: None, + }); + } + + // Expanding ownership is more restrictive than rebalancing across + // existing authoritative owners. It requires the explicit owner cap + // and cooldown, and never overlaps a normal provisional spill. + if self.config.max_cached_owners_per_prefix == 0 + || self.config.cache_owner_spill_cooldown_secs == 0 + || known_owners.len() >= self.config.max_cached_owners_per_prefix + || self + .recent_provisional_owner(prefix_key, workers, &healthy_indices) + .is_some() + { + return None; + } + + let non_owners: Vec<_> = healthy_indices + .into_iter() + .filter(|idx| !known_owners.contains(idx)) + .collect(); + let eligible_non_owners = + SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &non_owners); + let mut clean: Vec<_> = eligible_non_owners + .into_iter() + .filter_map(|idx| { + let capacity = headroom.iter().find(|candidate| { + candidate.worker_url.as_ref() == workers[idx].url() + && candidate.worker_revision == workers[idx].revision() + && candidate.issuable_slots > 0 + })?; + Some((idx, capacity.issuable_slots)) + }) + .collect(); + clean.sort_unstable_by(|(left_idx, left_slots), (right_idx, right_slots)| { + right_slots + .cmp(left_slots) + .then_with(|| workers[*left_idx].load().cmp(&workers[*right_idx].load())) + .then_with(|| workers[*left_idx].url().cmp(workers[*right_idx].url())) + }); + let (target_idx, _) = *clean.first()?; + let prefix_reservation = self.try_acquire_distribution_seed_prefix(prefix_key)?; + Some(OwnerPressureDispatchPlan { + target_worker_url: Arc::from(workers[target_idx].url()), + target_worker_revision: workers[target_idx].revision(), + expands_ownership: true, + _prefix_reservation: Some(prefix_reservation), + }) + } + fn select_from_cached_owners( &self, workers: &[Arc], @@ -3135,6 +3373,142 @@ mod tests { } } + #[test] + fn owner_pressure_seed_uses_only_authoritative_owners_and_does_not_mutate_them() { + let mut policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + engine_load: true, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + + let mut headroom: Vec<_> = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + let info = SelectWorkerInfo { + tokens: Some(&[1, 2, 3, 4, 5, 6, 7, 8]), + ..Default::default() + }; + let plan = policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .expect("one clean authoritative non-owner should be planned"); + assert_eq!(plan.target_worker_url(), workers[1].url()); + assert_eq!(plan.target_worker_revision(), workers[1].revision()); + assert!(plan.expands_ownership()); + + let monitor = policy.kv_monitor.read(); + let owners = CacheAwarePolicy::score_overlap( + &workers, + info.tokens.unwrap(), + &[0, 1, 2], + &monitor.as_ref().unwrap().get_indexer("unknown").unwrap(), + 4, + ); + assert_eq!( + owners, + vec![0], + "planning must not publish a provisional owner" + ); + drop(monitor); + + headroom[0].issuable_slots = 1; + let plan = policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .expect("an existing authoritative owner with headroom should be preferred"); + assert_eq!(plan.target_worker_url(), workers[0].url()); + assert!(!plan.expands_ownership()); + + policy.config.max_cached_owners_per_prefix = 0; + policy.config.cache_owner_spill_cooldown_secs = 0; + assert!(policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .is_some()); + headroom[0].issuable_slots = 0; + assert!(policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .is_none()); + } + + #[test] + fn owner_pressure_seed_serializes_prefix_expansion_without_publishing_it() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + engine_load: true, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + let headroom: Vec<_> = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let info = SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }; + let key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path("unknown"), + prefix_hash: kv_index::hash_token_path(&tokens), + kind: PrefixKind::Token, + }; + + let first = policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .expect("first clean-peer expansion should claim the prefix"); + assert!(first.expands_ownership()); + assert!(policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .is_none()); + assert!( + policy.replication_state.get(&key).is_none(), + "ordinary requests must not see or join a pending distribution seed" + ); + + drop(first); + assert!(policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .is_none()); + policy + .distribution_seed_state + .get_mut(&key) + .unwrap() + .last_transition = Instant::now() - Duration::from_secs(6); + assert!(policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .is_some()); + } + // -- score_overlap unit tests (scoring helper) -- #[test] diff --git a/model_gateway/src/policies/mod.rs b/model_gateway/src/policies/mod.rs index 1cb3fadf9..ab01f1ff9 100644 --- a/model_gateway/src/policies/mod.rs +++ b/model_gateway/src/policies/mod.rs @@ -27,6 +27,7 @@ pub(crate) mod utils; pub use bucket::BucketPolicy; pub use cache_aware::{CacheAwarePolicy, TreeHandle, TreeKind}; +pub(crate) use cache_aware::{OwnerPressureDispatchPlan, SeedWorkerHeadroom}; pub use consistent_hashing::ConsistentHashingPolicy; pub use dp_min_token::MinimumTokensPolicy; pub use factory::PolicyFactory; diff --git a/model_gateway/src/policies/registry.rs b/model_gateway/src/policies/registry.rs index 94daaa27c..6b8b06bf1 100644 --- a/model_gateway/src/policies/registry.rs +++ b/model_gateway/src/policies/registry.rs @@ -15,7 +15,7 @@ use tracing::{debug, info, warn}; /// When the last worker of a model is removed, the policy mapping is cleaned up. use super::{ BucketPolicy, CacheAwarePolicy, DPRankLoadPolicy, LoadBalancingPolicy, ManualConfig, - ManualPolicy, PolicyFactory, SelectWorkerInfo, + ManualPolicy, OwnerPressureDispatchPlan, PolicyFactory, SeedWorkerHeadroom, SelectWorkerInfo, }; use crate::{ config::types::{PolicyConfig, RoutingKeyOverrideConfig}, @@ -320,6 +320,23 @@ impl PolicyRegistry { .unwrap_or_else(|| self.get_default_policy()) } + /// Ask only the active cache-aware policy for a read-only clean-peer plan. + /// Other policies fail closed, so a scheduler proof can never become a + /// policy-agnostic adaptive-admission bypass. + pub(crate) fn owner_pressure_dispatch_plan( + &self, + model_id: &str, + workers: &[Arc], + info: &SelectWorkerInfo<'_>, + headroom: &[SeedWorkerHeadroom], + ) -> Option { + let policy = self.get_policy_or_default(model_id); + policy + .as_any() + .downcast_ref::()? + .owner_pressure_dispatch_plan(model_id, workers, info, headroom) + } + /// Determine policy for a new model fn determine_policy_for_model( &self, diff --git a/model_gateway/src/routers/grpc/adaptive_admission.rs b/model_gateway/src/routers/grpc/adaptive_admission.rs index 49198e345..fd0929020 100644 --- a/model_gateway/src/routers/grpc/adaptive_admission.rs +++ b/model_gateway/src/routers/grpc/adaptive_admission.rs @@ -8,12 +8,12 @@ //! rejecting traffic. use std::{ - collections::HashMap, + collections::{HashMap, HashSet}, sync::{ - atomic::{AtomicBool, Ordering}, + atomic::{AtomicBool, AtomicU64, Ordering}, Arc, Weak, }, - time::Instant, + time::{Duration, Instant}, }; use metrics::{counter, describe_counter, describe_gauge, describe_histogram, gauge, histogram}; @@ -25,10 +25,11 @@ use crate::{ config::{AdaptiveAdmissionConfig, AdaptiveAdmissionMode, AdaptiveAdmissionStrategy}, middleware::scheduler::state::AdaptiveCapacityProvider, observability::metrics::intern_string, - worker::WorkerRegistry, + worker::{registry::WorkerId, Worker, WorkerRegistry}, }; const ADMISSION_PARTITION_LABEL: &str = "admission_partition"; +const DISTRIBUTION_HEADROOM_MAX_TELEMETRY_AGE: Duration = Duration::from_secs(30); const PREDICTIONS_TOTAL: &str = "smg_adaptive_admission_predictions_total"; const PREDICTED_OUTPUT_TOKENS: &str = "smg_adaptive_admission_predicted_output_tokens"; @@ -531,6 +532,192 @@ impl PartitionLoad { } } +/// Strict per-worker engine snapshot used only to account distribution +/// headroom. Unlike the aggregate admission signal, every DP rank must be +/// present and internally valid before this record exists. +#[derive(Debug, Clone)] +struct DistributionWorkerTelemetry { + partition: Arc, + worker_url: Arc, + worker_id: WorkerId, + worker_instance: Arc, + worker_revision: u64, + telemetry_revision: u64, + observed_at: Instant, + running_requests: u64, + waiting_requests: u64, + max_running_requests: u64, + max_pressure: f64, +} + +impl DistributionWorkerTelemetry { + fn from_load( + partition: Arc, + worker_url: Arc, + worker_id: WorkerId, + worker_instance: Arc, + worker_revision: u64, + telemetry_revision: u64, + observed_at: Instant, + load: &WorkerLoadResponse, + ) -> Option { + let rank_count = usize::try_from(load.dp_rank_count).ok()?; + if rank_count == 0 || load.loads.len() != rank_count { + return None; + } + + let mut ranks = HashSet::with_capacity(rank_count); + let mut running_requests = 0_u64; + let mut waiting_requests = 0_u64; + let mut max_running_requests = 0_u64; + let mut max_pressure = 0.0_f64; + for rank in &load.loads { + let dp_rank = usize::try_from(rank.dp_rank).ok()?; + if dp_rank >= rank_count || !ranks.insert(dp_rank) { + return None; + } + if rank.num_running_reqs < 0 + || rank.num_waiting_reqs < 0 + || rank.num_waiting_uncached_tokens < 0 + || rank.num_total_reqs < 0 + || rank.num_used_tokens < 0 + || rank.max_total_num_tokens < 0 + || rank.max_running_requests <= 0 + { + return None; + } + if !rank.token_usage.is_finite() + || !(0.0..=1.0).contains(&rank.token_usage) + || !rank.utilization.is_finite() + || !(0.0..=1.0).contains(&rank.utilization) + { + return None; + } + running_requests = running_requests.checked_add(rank.num_running_reqs as u64)?; + waiting_requests = waiting_requests.checked_add(rank.num_waiting_reqs as u64)?; + max_running_requests = + max_running_requests.checked_add(rank.max_running_requests as u64)?; + max_pressure = max_pressure.max(rank.token_usage.max(rank.utilization)); + } + if ranks.len() != rank_count || max_running_requests == 0 { + return None; + } + + Some(Self { + partition, + worker_url, + worker_id, + worker_instance, + worker_revision, + telemetry_revision, + observed_at, + running_requests, + waiting_requests, + max_running_requests, + max_pressure, + }) + } + + fn is_fresh_at(&self, now: Instant) -> bool { + now.saturating_duration_since(self.observed_at) <= DISTRIBUTION_HEADROOM_MAX_TELEMETRY_AGE + } + + fn is_clean(&self, max_pressure: f64) -> bool { + self.waiting_requests == 0 && self.max_pressure < max_pressure + } + + fn occupied(&self, worker_load: usize) -> u64 { + let engine_occupied = self.running_requests.saturating_add(self.waiting_requests); + engine_occupied.max(u64::try_from(worker_load).unwrap_or(u64::MAX)) + } + + fn issuable_slots(&self, worker_load: usize, active_target_leases: u64) -> u16 { + self.max_running_requests + .saturating_sub(self.occupied(worker_load)) + .saturating_sub(active_target_leases) + .min(u64::from(u16::MAX)) as u16 + } +} + +#[derive(Debug, Clone)] +struct DistributionHeadroomBinding { + partition: Arc, + model: Arc, + worker_url: Arc, + worker_id: WorkerId, + worker_instance: Arc, + worker_revision: u64, + telemetry_revision: u64, +} + +impl PartialEq for DistributionHeadroomBinding { + fn eq(&self, other: &Self) -> bool { + self.partition == other.partition + && self.model == other.model + && self.worker_url == other.worker_url + && self.worker_id == other.worker_id + && Arc::ptr_eq(&self.worker_instance, &other.worker_instance) + && self.worker_revision == other.worker_revision + && self.telemetry_revision == other.telemetry_revision + } +} + +impl Eq for DistributionHeadroomBinding {} + +/// Read-only, unforgeable snapshot of one worker's currently issuable +/// distribution headroom. Acquisition revalidates every private binding. +#[derive(Debug, Clone)] +pub(crate) struct DistributionHeadroomTarget { + binding: DistributionHeadroomBinding, + issuable_slots: u16, +} + +impl DistributionHeadroomTarget { + pub(crate) fn worker_url(&self) -> &str { + &self.binding.worker_url + } + + pub(crate) fn worker_revision(&self) -> u64 { + self.binding.worker_revision + } + + pub(crate) fn issuable_slots(&self) -> u16 { + self.issuable_slots + } +} + +/// One process-local reservation of a clean worker's engine slot. The lease +/// is non-cloneable and releases itself on every early-return path. +#[derive(Debug)] +pub(crate) struct DistributionHeadroomLease { + controller: Weak, + binding: Option, +} + +impl DistributionHeadroomLease { + /// Revalidate the exact worker, registry revision, telemetry revision, + /// pressure, and reserved slot immediately before dispatch. + pub(crate) fn verify(&self) -> bool { + let Some(binding) = &self.binding else { + return false; + }; + self.controller + .upgrade() + .is_some_and(|controller| controller.verify_distribution_headroom(binding)) + } +} + +impl Drop for DistributionHeadroomLease { + fn drop(&mut self) { + let Some(binding) = self.binding.take() else { + return; + }; + if let Some(controller) = self.controller.upgrade() { + controller.release_distribution_headroom(&binding); + } + } +} + #[derive(Debug, Default)] struct WorkState { outstanding_tokens: HashMap, @@ -538,6 +725,8 @@ struct WorkState { loads: HashMap, capacities: HashMap, feedback_estimates: HashMap, + distribution_telemetry: HashMap, + active_distribution_leases: HashMap, } #[derive(Debug, Clone, Copy)] @@ -636,6 +825,7 @@ pub(crate) struct AdaptiveAdmissionController { prediction_samples: Option>, registry: Arc, capacity_revision: watch::Sender, + telemetry_revision: AtomicU64, } impl AdaptiveAdmissionController { @@ -648,6 +838,7 @@ impl AdaptiveAdmissionController { prediction_samples: PredictionSampleState::from_env().map(Mutex::new), registry, capacity_revision, + telemetry_revision: AtomicU64::new(0), }) } @@ -679,7 +870,13 @@ impl AdaptiveAdmissionController { } fn update_loads(&self, loads: &HashMap) { + let observed_at = Instant::now(); + let telemetry_revision = self + .telemetry_revision + .fetch_add(1, Ordering::AcqRel) + .wrapping_add(1); let mut partitions: HashMap = HashMap::new(); + let mut distribution_telemetry = HashMap::new(); for worker in self .registry .get_all() @@ -695,11 +892,33 @@ impl AdaptiveAdmissionController { .filter(|value| !value.trim().is_empty()) .unwrap_or_else(|| worker.model_id()) .to_string(); - let aggregate = partitions.entry(partition).or_default(); + let aggregate = partitions.entry(partition.clone()).or_default(); aggregate.healthy_replicas = aggregate.healthy_replicas.saturating_add(1); let Some(load) = loads.get(worker.url()) else { continue; }; + if worker.is_available() { + if let Some(worker_id) = self.registry.get_id_by_url(worker.url()) { + if self + .registry + .get(&worker_id) + .is_some_and(|current| Arc::ptr_eq(¤t, &worker)) + { + if let Some(telemetry) = DistributionWorkerTelemetry::from_load( + Arc::from(partition.as_str()), + Arc::from(worker.url()), + worker_id, + Arc::clone(&worker), + worker.revision(), + telemetry_revision, + observed_at, + load, + ) { + distribution_telemetry.insert(worker.url().to_string(), telemetry); + } + } + } + } aggregate.observed_replicas = aggregate.observed_replicas.saturating_add(1); aggregate.generation_tokens_per_second += load.total_gen_throughput().max(0.0); aggregate.running_requests += load @@ -733,7 +952,7 @@ impl AdaptiveAdmissionController { } } - let now = Instant::now(); + let now = observed_at; let mut work = self.work.lock(); work.capacities .retain(|partition, _| partitions.contains_key(partition)); @@ -789,6 +1008,7 @@ impl AdaptiveAdmissionController { } } work.loads = partitions; + work.distribution_telemetry = distribution_telemetry; for (partition, load) in &work.loads { let partition_label = intern_string(partition); gauge!(LOAD_COVERAGE, "partition" => Arc::clone(&partition_label)).set(load.coverage()); @@ -823,6 +1043,261 @@ impl AdaptiveAdmissionController { .send_modify(|revision| *revision = revision.wrapping_add(1)); } + fn distribution_headroom_enabled(&self, partition: &str) -> bool { + self.config.mode == AdaptiveAdmissionMode::Enforce + && self.config.strategy == AdaptiveAdmissionStrategy::EngineFeedback + && self.config.distribution_headroom_partition_seed_cap > 0 + && self + .config + .distribution_headroom_partitions + .iter() + .any(|allowed| allowed == partition) + } + + fn current_model_worker( + &self, + binding: &DistributionHeadroomBinding, + ) -> Option> { + // A caller must supply the canonical model selected at request entry. + // Aliases are not accepted as an independently reusable scope. + if self.registry.resolve_model_alias(&binding.model).is_some() { + return None; + } + if self.registry.get_id_by_url(&binding.worker_url).as_ref() != Some(&binding.worker_id) { + return None; + } + let worker = self.registry.get(&binding.worker_id)?; + if worker.url() != binding.worker_url.as_ref() + || worker.revision() != binding.worker_revision + || !worker.is_available() + || !Arc::ptr_eq(&worker, &binding.worker_instance) + { + return None; + } + self.registry + .get_by_model(&binding.model) + .iter() + .any(|candidate| { + Arc::ptr_eq(candidate, &binding.worker_instance) + && candidate.url() == binding.worker_url.as_ref() + && candidate.revision() == binding.worker_revision + && candidate.is_available() + }) + .then_some(worker) + } + + fn telemetry_matches_binding( + telemetry: &DistributionWorkerTelemetry, + binding: &DistributionHeadroomBinding, + now: Instant, + ) -> bool { + telemetry.partition == binding.partition + && telemetry.worker_url == binding.worker_url + && telemetry.worker_id == binding.worker_id + && Arc::ptr_eq(&telemetry.worker_instance, &binding.worker_instance) + && telemetry.worker_revision == binding.worker_revision + && telemetry.telemetry_revision == binding.telemetry_revision + && telemetry.is_fresh_at(now) + } + + fn active_partition_leases(work: &WorkState, partition: &str) -> u64 { + u64::try_from( + work.active_distribution_leases + .values() + .filter(|binding| binding.partition.as_ref() == partition) + .count(), + ) + .unwrap_or(u64::MAX) + } + + /// Return a complete strictly validated worker-capacity snapshot for one + /// canonical model and exact admission partition. Hot, full, and currently + /// leased workers remain visible with zero issuable slots so routing can + /// prove owner saturation without mistaking an omitted owner for capacity. + /// Missing, malformed, stale, or unavailable workers are omitted and thus + /// fail closed. Targets are advisory; acquisition revalidates every + /// private binding under the lease lock. + pub(crate) fn distribution_headroom_snapshot( + &self, + partition: &str, + model: &str, + ) -> Vec { + if !self.distribution_headroom_enabled(partition) + || self.registry.resolve_model_alias(model).is_some() + { + return Vec::new(); + } + let workers = self.registry.get_by_model(model); + if workers.is_empty() { + return Vec::new(); + } + + let now = Instant::now(); + let work = self.work.lock(); + let partition_has_capacity = Self::active_partition_leases(&work, partition) + < u64::from(self.config.distribution_headroom_partition_seed_cap); + let mut targets = Vec::new(); + for worker in workers.iter().filter(|worker| worker.is_available()) { + let Some(telemetry) = work.distribution_telemetry.get(worker.url()) else { + continue; + }; + let Some(worker_id) = self.registry.get_id_by_url(worker.url()) else { + continue; + }; + if worker_id != telemetry.worker_id || !Arc::ptr_eq(worker, &telemetry.worker_instance) + { + continue; + } + let binding = DistributionHeadroomBinding { + partition: Arc::from(partition), + model: Arc::from(model), + worker_url: Arc::from(worker.url()), + worker_id, + worker_instance: Arc::clone(worker), + worker_revision: worker.revision(), + telemetry_revision: telemetry.telemetry_revision, + }; + if !Self::telemetry_matches_binding(telemetry, &binding, now) { + continue; + } + let active_target_leases = + u64::from(work.active_distribution_leases.contains_key(worker.url())); + let raw_issuable = telemetry.issuable_slots(worker.load(), active_target_leases); + let issuable_slots = if partition_has_capacity + && active_target_leases == 0 + && telemetry.is_clean(self.config.feedback_max_token_usage) + { + raw_issuable + } else { + 0 + }; + targets.push(DistributionHeadroomTarget { + binding, + issuable_slots, + }); + } + targets.sort_unstable_by(|left, right| left.worker_url().cmp(right.worker_url())); + targets + } + + /// Acquire one exact clean-worker slot. This is intentionally independent + /// of routing policy and scheduler authorization; callers must layer those + /// separate decisions around this process-local capacity reservation. + pub(crate) fn try_acquire_distribution_headroom( + self: &Arc, + target: &DistributionHeadroomTarget, + ) -> Option { + let binding = &target.binding; + if target.issuable_slots == 0 || !self.distribution_headroom_enabled(&binding.partition) { + return None; + } + + let worker = self.current_model_worker(binding)?; + let now = Instant::now(); + let mut work = self.work.lock(); + if Self::active_partition_leases(&work, &binding.partition) + >= u64::from(self.config.distribution_headroom_partition_seed_cap) + || work + .active_distribution_leases + .contains_key(binding.worker_url.as_ref()) + { + return None; + } + let telemetry = work + .distribution_telemetry + .get(binding.worker_url.as_ref())?; + if !Self::telemetry_matches_binding(telemetry, binding, now) + || !telemetry.is_clean(self.config.feedback_max_token_usage) + || telemetry.issuable_slots(worker.load(), 0) == 0 + { + return None; + } + // Re-read mutable worker state after all telemetry checks. A later + // registry or health transition is caught by the mandatory verify. + if worker.revision() != binding.worker_revision || !worker.is_available() { + return None; + } + work.active_distribution_leases + .insert(binding.worker_url.to_string(), binding.clone()); + drop(work); + self.capacity_revision + .send_modify(|revision| *revision = revision.wrapping_add(1)); + if self.current_model_worker(binding).is_none() { + self.release_distribution_headroom(binding); + return None; + } + Some(DistributionHeadroomLease { + controller: Arc::downgrade(self), + binding: Some(binding.clone()), + }) + } + + fn verify_distribution_headroom(&self, binding: &DistributionHeadroomBinding) -> bool { + if !self.distribution_headroom_enabled(&binding.partition) { + return false; + } + let Some(worker) = self.current_model_worker(binding) else { + return false; + }; + let now = Instant::now(); + let work = self.work.lock(); + if work + .active_distribution_leases + .get(binding.worker_url.as_ref()) + != Some(binding) + { + return false; + } + let active_partition_leases = Self::active_partition_leases(&work, &binding.partition); + if active_partition_leases == 0 + || active_partition_leases + > u64::from(self.config.distribution_headroom_partition_seed_cap) + { + return false; + } + let Some(telemetry) = work.distribution_telemetry.get(binding.worker_url.as_ref()) else { + return false; + }; + if !Self::telemetry_matches_binding(telemetry, binding, now) + || !telemetry.is_clean(self.config.feedback_max_token_usage) + { + return false; + } + let active_target_leases = 1_u64; + let occupied = telemetry.occupied(worker.load()); + if telemetry.max_running_requests.saturating_sub(occupied) < active_target_leases { + return false; + } + // Evaluate the specified remaining-headroom formula even when this + // lease consumes the final slot. Zero remaining slots is still a valid + // reservation for this already-held lease. + let _remaining_issuable = telemetry.issuable_slots(worker.load(), active_target_leases); + drop(work); + worker.revision() == binding.worker_revision + && worker.is_available() + && self.current_model_worker(binding).is_some() + } + + fn release_distribution_headroom(&self, binding: &DistributionHeadroomBinding) { + let mut work = self.work.lock(); + let removed = if work + .active_distribution_leases + .get(binding.worker_url.as_ref()) + == Some(binding) + { + work.active_distribution_leases + .remove(binding.worker_url.as_ref()) + .is_some() + } else { + false + }; + drop(work); + if removed { + self.capacity_revision + .send_modify(|revision| *revision = revision.wrapping_add(1)); + } + } + pub(crate) fn begin( self: &Arc, partition: String, @@ -1090,16 +1565,16 @@ impl AdaptiveCapacityProvider for AdaptiveAdmissionController { &load, work.feedback_estimates.get(partition), ); - if !constraints.telemetry_usable { - return static_capacity; - } - if constraints.pressure_reason.is_some() { - return 0; - } - let Some(running_limit) = constraints.running_limit else { - return static_capacity; + let ordinary_capacity = if !constraints.telemetry_usable { + static_capacity + } else if constraints.pressure_reason.is_some() { + 0 + } else if let Some(running_limit) = constraints.running_limit { + (running_limit.floor().clamp(0.0, f64::from(u16::MAX)) as u16).min(static_capacity) + } else { + static_capacity }; - (running_limit.floor().clamp(0.0, f64::from(u16::MAX)) as u16).min(static_capacity) + ordinary_capacity } fn subscribe_capacity_changes(&self) -> watch::Receiver { @@ -1150,6 +1625,20 @@ impl AdaptiveRequestTracker { .map_or(0, |inner| inner.decision.retry_after_secs) } + /// Exact admission partition bound when this tracker was created. + pub(crate) fn partition(&self) -> Option<&str> { + self.inner.as_ref().map(|inner| inner.partition.as_str()) + } + + /// Machine-stable reason for an enforced, telemetry-backed rejection. + /// Admitted and telemetry-fallback trackers do not report a rejection. + pub(crate) fn rejection_reason(&self) -> Option<&'static str> { + self.inner.as_ref().and_then(|inner| { + (inner.decision.telemetry_usable && !inner.decision.would_admit) + .then_some(inner.decision.reason) + }) + } + pub(crate) fn complete(mut self, observed_output_tokens: u32) { self.resolve(Some(observed_output_tokens)); } @@ -1175,9 +1664,15 @@ impl Drop for AdaptiveRequestTracker { #[cfg(test)] mod tests { - use std::time::Duration; + use std::{sync::Barrier, time::Duration}; + + use openai_protocol::{ + model_card::ModelCard, + worker::{HealthCheckConfig, SchedulerLoadSnapshot, WorkerStatus}, + }; use super::*; + use crate::worker::BasicWorkerBuilder; fn config() -> AdaptiveAdmissionConfig { AdaptiveAdmissionConfig { @@ -1193,9 +1688,74 @@ mod tests { feedback_max_waiting_requests_per_healthy_replica: 2, feedback_max_token_usage: 0.9, feedback_throughput_improvement_ratio: 0.02, + distribution_headroom_partitions: Vec::new(), + distribution_headroom_partition_seed_cap: 0, + } + } + + fn distribution_config(cap: u32) -> AdaptiveAdmissionConfig { + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Enforce, + strategy: AdaptiveAdmissionStrategy::EngineFeedback, + distribution_headroom_partitions: vec!["k3".to_string()], + distribution_headroom_partition_seed_cap: cap, + ..config() } } + fn register_headroom_worker(registry: &Arc, url: &str) -> Arc { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(url) + .model(ModelCard::new("model").with_alias("model-alias")) + .label(ADMISSION_PARTITION_LABEL, "k3") + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); + registry + .register(Arc::clone(&worker)) + .expect("worker registration"); + worker + } + + fn worker_load(ranks: &[(i32, i32, i32, f64, f64, i32)]) -> WorkerLoadResponse { + WorkerLoadResponse { + dp_rank_count: i32::try_from(ranks.len()).expect("test rank count"), + loads: ranks + .iter() + .map( + |&(dp_rank, running, waiting, token_usage, utilization, maximum)| { + SchedulerLoadSnapshot { + dp_rank, + num_running_reqs: running, + num_waiting_reqs: waiting, + num_total_reqs: running.saturating_add(waiting), + token_usage, + utilization, + max_running_requests: maximum, + ..Default::default() + } + }, + ) + .collect(), + ..Default::default() + } + } + + fn update_worker_loads( + controller: &AdaptiveAdmissionController, + loads: impl IntoIterator, + ) { + controller.update_loads( + &loads + .into_iter() + .map(|(url, load)| (url.to_string(), load)) + .collect(), + ); + } + fn features(user: &str, prompt_tokens: u32, maximum: Option) -> PredictionFeatures { PredictionFeatures { model: "model".to_string(), @@ -1463,6 +2023,8 @@ mod tests { let tracker = controller.begin("model".to_string(), features("u", 10, None)); assert!(tracker.should_reject()); + assert_eq!(tracker.partition(), Some("model")); + assert_eq!(tracker.rejection_reason(), Some(expected_reason)); assert_eq!( tracker.inner.as_ref().unwrap().decision.reason, expected_reason @@ -1687,6 +2249,27 @@ mod tests { assert_eq!(controller.effective_capacity("model", 100), 0); } + #[test] + fn engine_feedback_aggregate_pressure_remains_mean_across_workers() { + let mut settings = config(); + settings.mode = AdaptiveAdmissionMode::Enforce; + settings.strategy = AdaptiveAdmissionStrategy::EngineFeedback; + let load = PartitionLoad { + healthy_replicas: 2, + observed_replicas: 2, + max_token_usage: 0.95, + token_usage_sum: 0.95, + max_running_requests: 128, + max_running_observed_replicas: 2, + ..PartitionLoad::default() + }; + + let constraints = engine_feedback_constraints(&settings, &load, None); + assert_eq!(load.mean_token_usage(), 0.475); + assert_eq!(constraints.pressure_reason, None); + assert_eq!(constraints.running_limit, Some(128.0)); + } + #[test] fn capacity_revision_changes_on_load_samples_not_per_request() { let controller = @@ -1701,6 +2284,436 @@ mod tests { assert!(revision.has_changed().unwrap()); } + #[test] + fn distribution_headroom_exposes_idle_peers_behind_one_hot_worker() { + const HOT: &str = "grpc://hot:30000"; + const IDLE_A: &str = "grpc://idle-a:30000"; + const IDLE_B: &str = "grpc://idle-b:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, HOT); + register_headroom_worker(®istry, IDLE_A); + let idle_b = register_headroom_worker(®istry, IDLE_B); + for _ in 0..5 { + idle_b.increment_load(); + } + let controller = AdaptiveAdmissionController::new(distribution_config(2), registry); + update_worker_loads( + &controller, + [ + (HOT, worker_load(&[(0, 38, 64, 0.89, 0.4, 100)])), + (IDLE_A, worker_load(&[(0, 0, 0, 0.1, 0.2, 43)])), + (IDLE_B, worker_load(&[(0, 0, 0, 0.1, 0.2, 43)])), + ], + ); + + assert_eq!(controller.effective_capacity("k3", 512), 109); + let targets = controller.distribution_headroom_snapshot("k3", "model"); + assert_eq!(targets.len(), 3); + assert_eq!(targets[0].worker_url(), HOT); + assert_eq!(targets[0].issuable_slots(), 0); + assert_eq!(targets[1].worker_url(), IDLE_A); + assert_eq!(targets[1].issuable_slots(), 43); + assert_eq!(targets[2].worker_url(), IDLE_B); + assert_eq!(targets[2].issuable_slots(), 38); + } + + #[test] + fn distribution_headroom_uses_max_pressure_across_every_rank() { + const URL: &str = "grpc://dp2:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); + update_worker_loads( + &controller, + [( + URL, + worker_load(&[(0, 0, 0, 0.1, 0.2, 24), (1, 0, 0, 0.2, 0.95, 24)]), + )], + ); + + { + let work = controller.work.lock(); + assert_eq!( + work.distribution_telemetry.get(URL).unwrap().max_pressure, + 0.95 + ); + assert!( + (work.loads.get("k3").unwrap().mean_token_usage() - 0.15).abs() < f64::EPSILON, + "the existing aggregate path must retain mean token usage" + ); + } + let snapshot = controller.distribution_headroom_snapshot("k3", "model"); + assert_eq!(snapshot.len(), 1); + assert_eq!(snapshot[0].issuable_slots(), 0); + + update_worker_loads( + &controller, + [( + URL, + worker_load(&[(0, 0, 0, 0.1, 0.2, 24), (1, 0, 0, 0.2, 0.8, 24)]), + )], + ); + let targets = controller.distribution_headroom_snapshot("k3", "model"); + assert_eq!(targets.len(), 1); + assert_eq!(targets[0].issuable_slots(), 48); + } + + #[test] + fn distribution_headroom_never_raises_aggregate_scheduler_capacity() { + const HOT: &str = "grpc://floor-hot:30000"; + const IDLE_A: &str = "grpc://floor-idle-a:30000"; + const IDLE_B: &str = "grpc://floor-idle-b:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, HOT); + register_headroom_worker(®istry, IDLE_A); + register_headroom_worker(®istry, IDLE_B); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); + update_worker_loads( + &controller, + [ + (HOT, worker_load(&[(0, 38, 64, 0.89, 0.4, 100)])), + (IDLE_A, worker_load(&[(0, 0, 0, 0.1, 0.2, 43)])), + (IDLE_B, worker_load(&[(0, 0, 0, 0.1, 0.2, 43)])), + ], + ); + + assert_eq!(controller.effective_capacity("k3", 512), 0); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .into_iter() + .find(|target| target.issuable_slots() > 0) + .unwrap(); + let lease = controller + .try_acquire_distribution_headroom(&target) + .unwrap(); + assert_eq!( + controller.effective_capacity("k3", 512), + 0, + "a routing lease must never mint an untyped scheduler slot" + ); + drop(lease); + assert_eq!(controller.effective_capacity("k3", 512), 0); + } + + #[test] + fn distribution_headroom_capacity_floor_requires_complete_strict_telemetry() { + const HOT: &str = "grpc://incomplete-hot:30000"; + const MISSING: &str = "grpc://incomplete-missing:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, HOT); + register_headroom_worker(®istry, MISSING); + let mut settings = distribution_config(1); + settings.min_load_coverage = 0.5; + let controller = AdaptiveAdmissionController::new(settings, registry); + update_worker_loads( + &controller, + [(HOT, worker_load(&[(0, 38, 64, 0.89, 0.4, 100)]))], + ); + + assert_eq!( + controller.effective_capacity("k3", 512), + 0, + "missing strict telemetry must not reopen aggregate pressure" + ); + } + + #[test] + fn distribution_headroom_capacity_floor_is_default_off() { + const HOT: &str = "grpc://off-hot:30000"; + const IDLE: &str = "grpc://off-idle:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, HOT); + register_headroom_worker(®istry, IDLE); + let mut settings = distribution_config(1); + settings.distribution_headroom_partitions.clear(); + settings.distribution_headroom_partition_seed_cap = 0; + let controller = AdaptiveAdmissionController::new(settings, registry); + update_worker_loads( + &controller, + [ + (HOT, worker_load(&[(0, 38, 64, 0.89, 0.4, 100)])), + (IDLE, worker_load(&[(0, 0, 0, 0.1, 0.2, 43)])), + ], + ); + + assert_eq!(controller.effective_capacity("k3", 512), 0); + } + + #[test] + fn distribution_headroom_fails_closed_on_missing_malformed_or_stale_telemetry() { + const URL: &str = "grpc://strict:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); + + controller.update_loads(&HashMap::new()); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .is_empty()); + + let malformed = [ + WorkerLoadResponse { + dp_rank_count: 2, + loads: worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]).loads, + ..Default::default() + }, + worker_load(&[(0, 0, 0, 0.1, 0.1, 16), (0, 0, 0, 0.1, 0.1, 16)]), + worker_load(&[(0, -1, 0, 0.1, 0.1, 16)]), + worker_load(&[(0, 0, 0, 0.1, 0.1, 0)]), + worker_load(&[(0, 0, 0, f64::NAN, 0.1, 16)]), + ]; + for load in malformed { + update_worker_loads(&controller, [(URL, load)]); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .is_empty()); + } + + update_worker_loads( + &controller, + [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]))], + ); + controller + .work + .lock() + .distribution_telemetry + .get_mut(URL) + .unwrap() + .observed_at = Instant::now().checked_sub(Duration::from_secs(31)).unwrap(); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .is_empty()); + } + + #[test] + fn distribution_headroom_allowlist_and_cap_are_default_off_and_exact() { + const URL: &str = "grpc://off:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + + for settings in [ + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Enforce, + strategy: AdaptiveAdmissionStrategy::EngineFeedback, + ..config() + }, + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Enforce, + strategy: AdaptiveAdmissionStrategy::EngineFeedback, + distribution_headroom_partitions: vec!["k3".to_string()], + distribution_headroom_partition_seed_cap: 0, + ..config() + }, + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Enforce, + strategy: AdaptiveAdmissionStrategy::EngineFeedback, + distribution_headroom_partitions: vec!["k3-canary".to_string()], + distribution_headroom_partition_seed_cap: 1, + ..config() + }, + ] { + let controller = AdaptiveAdmissionController::new(settings, Arc::clone(®istry)); + update_worker_loads( + &controller, + [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]))], + ); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .is_empty()); + } + } + + #[test] + fn distribution_headroom_concurrently_enforces_target_and_partition_caps() { + const A: &str = "grpc://cap-a:30000"; + const B: &str = "grpc://cap-b:30000"; + const C: &str = "grpc://cap-c:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, A); + register_headroom_worker(®istry, B); + register_headroom_worker(®istry, C); + let controller = AdaptiveAdmissionController::new(distribution_config(2), registry); + update_worker_loads( + &controller, + [ + (A, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)])), + (B, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)])), + (C, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)])), + ], + ); + + let one_target = controller + .distribution_headroom_snapshot("k3", "model") + .into_iter() + .next() + .unwrap(); + let barrier = Arc::new(Barrier::new(8)); + let claim_handles: Vec<_> = (0..8) + .map(|_| { + let controller = Arc::clone(&controller); + let target = one_target.clone(); + let barrier = Arc::clone(&barrier); + std::thread::spawn(move || { + barrier.wait(); + controller.try_acquire_distribution_headroom(&target) + }) + }) + .collect(); + let claims: Vec<_> = claim_handles + .into_iter() + .map(|handle| handle.join().unwrap()) + .flatten() + .collect(); + assert_eq!(claims.len(), 1); + drop(claims); + + let targets = controller.distribution_headroom_snapshot("k3", "model"); + let barrier = Arc::new(Barrier::new(targets.len())); + let lease_handles: Vec<_> = targets + .into_iter() + .map(|target| { + let controller = Arc::clone(&controller); + let barrier = Arc::clone(&barrier); + std::thread::spawn(move || { + barrier.wait(); + controller.try_acquire_distribution_headroom(&target) + }) + }) + .collect(); + let leases: Vec<_> = lease_handles + .into_iter() + .map(|handle| handle.join().unwrap()) + .flatten() + .collect(); + assert_eq!(leases.len(), 2); + assert_eq!( + AdaptiveAdmissionController::active_partition_leases(&controller.work.lock(), "k3"), + 2 + ); + } + + #[test] + fn distribution_headroom_lease_drop_refunds_capacity() { + const URL: &str = "grpc://refund:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); + update_worker_loads(&controller, [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 1)]))]); + + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + let lease = controller + .try_acquire_distribution_headroom(&target) + .unwrap(); + assert!(lease.verify()); + let snapshot = controller.distribution_headroom_snapshot("k3", "model"); + assert_eq!(snapshot.len(), 1); + assert_eq!(snapshot[0].issuable_slots(), 0); + drop(lease); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + assert!(controller + .try_acquire_distribution_headroom(&target) + .is_some()); + } + + #[test] + fn distribution_headroom_revision_health_and_model_changes_fail_closed() { + const URL: &str = "grpc://revision:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + let worker_id = registry.get_id_by_url(URL).unwrap(); + let controller = + AdaptiveAdmissionController::new(distribution_config(1), Arc::clone(®istry)); + let load = worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]); + update_worker_loads(&controller, [(URL, load.clone())]); + + assert!(controller + .distribution_headroom_snapshot("k3", "model-alias") + .is_empty()); + let stale_target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + let replacement: Arc = Arc::new( + BasicWorkerBuilder::new(URL) + .model(ModelCard::new("model").with_alias("model-alias")) + .label(ADMISSION_PARTITION_LABEL, "k3") + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); + assert!(registry.replace(&worker_id, replacement)); + assert!(controller + .try_acquire_distribution_headroom(&stale_target) + .is_none()); + + update_worker_loads(&controller, [(URL, load.clone())]); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + let lease = controller + .try_acquire_distribution_headroom(&target) + .unwrap(); + assert!(lease.verify()); + let second_replacement: Arc = Arc::new( + BasicWorkerBuilder::new(URL) + .model(ModelCard::new("model").with_alias("model-alias")) + .label(ADMISSION_PARTITION_LABEL, "k3") + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); + assert!(registry.replace(&worker_id, second_replacement)); + assert!( + !lease.verify(), + "same-ID replacement must invalidate exact worker-instance identity" + ); + drop(lease); + + update_worker_loads(&controller, [(URL, load.clone())]); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + let lease = controller + .try_acquire_distribution_headroom(&target) + .unwrap(); + assert!(lease.verify()); + update_worker_loads(&controller, [(URL, load)]); + assert!( + !lease.verify(), + "telemetry revision changes must invalidate" + ); + drop(lease); + + update_worker_loads( + &controller, + [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]))], + ); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + registry + .get_by_url(URL) + .unwrap() + .set_status(WorkerStatus::NotReady); + assert!(controller + .try_acquire_distribution_headroom(&target) + .is_none()); + assert!(!registry.get_by_url(URL).unwrap().is_available()); + } + #[test] fn feedback_knee_moves_down_on_same_throughput_plateau() { let start = Instant::now(); diff --git a/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs index 4a55fc6c5..7680b8606 100644 --- a/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs +++ b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs @@ -8,7 +8,9 @@ use axum::{ use super::PipelineStage; use crate::{ - middleware::scheduler::{LocalAdaptiveRejection, ADMISSION_PARTITION_HEADER}, + middleware::scheduler::{ + LocalAdaptiveRejection, SchedulerAdmissionProof, ADMISSION_PARTITION_HEADER, + }, routers::{ error, grpc::{ @@ -16,7 +18,7 @@ use crate::{ PredictionFeatures, FLAG_MULTIPLE_COMPLETIONS, FLAG_REASONING, FLAG_STREAMING, FLAG_STRUCTURED_OUTPUT, FLAG_TOOLS, }, - context::{RequestContext, RequestType}, + context::{PendingDistributionSeed, RequestContext, RequestType}, }, }, }; @@ -162,6 +164,13 @@ impl GenerationShape { } } +fn is_single_distribution_seed_sample( + shape: GenerationShape, + backend_request_count: usize, +) -> bool { + backend_request_count == 1 && shape.flags & FLAG_MULTIPLE_COMPLETIONS == 0 +} + fn multiplied_limit(per_completion: Option, multiplicity: u32) -> Option { per_completion.map(|limit| limit.saturating_mul(multiplicity)) } @@ -177,7 +186,7 @@ fn trusted_header(ctx: &RequestContext, name: &str) -> Option { .map(str::to_string) } -fn rejection_response(retry_after_secs: u32) -> Response { +pub(crate) fn rejection_response(retry_after_secs: u32) -> Response { let mut response = error::create_error( StatusCode::TOO_MANY_REQUESTS, "adaptive_admission_saturated", @@ -229,6 +238,29 @@ impl PipelineStage for AdaptiveAdmissionStage { }, ); if tracker.should_reject() { + let can_try_distribution_seed = + matches!( + tracker.rejection_reason(), + Some("engine_waiting" | "running_limit") + ) && ctx.state.preparation.as_ref().is_some_and(|preparation| { + is_single_distribution_seed_sample(shape, preparation.backend_request_count()) + }) && ctx + .input + .tenant_request_meta + .as_ref() + .and_then(|meta| meta.extension::()) + .is_some(); + if can_try_distribution_seed { + let Some(partition) = tracker.partition().map(str::to_string) else { + return Err(rejection_response(tracker.retry_after_secs())); + }; + ctx.state.pending_distribution_seed = Some(PendingDistributionSeed { + partition, + retry_after_secs: tracker.retry_after_secs(), + }); + ctx.state.adaptive_request = Some(tracker); + return Ok(None); + } return Err(rejection_response(tracker.retry_after_secs())); } ctx.state.adaptive_request = Some(tracker); @@ -242,8 +274,23 @@ impl PipelineStage for AdaptiveAdmissionStage { #[cfg(test)] mod tests { + use std::sync::Arc; + use super::*; + #[test] + fn multiple_generate_samples_cannot_enter_distribution_seed_path() { + let request = serde_json::from_value(serde_json::json!({ + "text": "hello", + "sampling_params": { "n": 2 } + })) + .unwrap(); + let shape = + GenerationShape::for_request(&RequestType::Generate(Arc::new(request))).unwrap(); + assert_ne!(shape.flags & FLAG_MULTIPLE_COMPLETIONS, 0); + assert!(!is_single_distribution_seed_sample(shape, 1)); + } + #[test] fn local_rejection_carries_internal_scheduler_marker() { let response = rejection_response(3); diff --git a/model_gateway/src/routers/grpc/common/stages/request_execution.rs b/model_gateway/src/routers/grpc/common/stages/request_execution.rs index 8127a81c1..71c59ee31 100644 --- a/model_gateway/src/routers/grpc/common/stages/request_execution.rs +++ b/model_gateway/src/routers/grpc/common/stages/request_execution.rs @@ -7,7 +7,7 @@ use axum::response::Response; use futures::future::{join_all, try_join_all}; use tracing::{debug, error, info_span, Instrument}; -use super::PipelineStage; +use super::{adaptive_admission::rejection_response, PipelineStage}; use crate::{ observability::metrics::{metrics_labels, Metrics}, routers::{ @@ -130,6 +130,11 @@ impl RequestExecutionStage { #[async_trait] impl PipelineStage for RequestExecutionStage { async fn execute(&self, ctx: &mut RequestContext) -> Result, Response> { + if let Some(seed) = ctx.state.distribution_seed_guard.as_ref() { + if !seed.headroom.verify() { + return Err(rejection_response(seed.retry_after_secs)); + } + } let execution_plan = ctx.state.execution_plan.take().ok_or_else(|| { error!( function = "RequestExecutionStage::execute", @@ -172,11 +177,13 @@ impl PipelineStage for RequestExecutionStage { _ => 1, }; let policy_reservation = ctx.state.policy_reservation.take(); + let distribution_seed = ctx.state.distribution_seed_guard.take(); ctx.state.load_guards = Some(LoadGuards::scaled( workers, ctx.input.headers.as_ref(), sub_requests, policy_reservation, + distribution_seed, )); // Extract dispatch metadata for tracing span diff --git a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs index 176655067..4da3223b2 100644 --- a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs +++ b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs @@ -9,15 +9,20 @@ use axum::{ }; use tracing::{error, warn}; -use super::PipelineStage; +use super::{adaptive_admission::rejection_response, PipelineStage}; use crate::{ observability::metrics::{metrics_labels, Metrics}, - policies::{LoadBalancingPolicy, PolicyRegistry, SelectWorkerInfo, WorkerLeg}, + policies::{ + LoadBalancingPolicy, PolicyRegistry, SeedWorkerHeadroom, SelectWorkerInfo, WorkerLeg, + }, routers::{ common::header_utils::worker_url_is_allowed, error, grpc::{ - context::{EncodeWorkerAssignment, PolicyReservation, RequestContext, WorkerSelection}, + context::{ + DistributionSeedDispatchGuard, EncodeWorkerAssignment, PendingDistributionSeed, + PolicyReservation, RequestContext, WorkerSelection, + }, multimodal, }, }, @@ -71,6 +76,7 @@ impl WorkerSelectionStage { #[async_trait] impl PipelineStage for WorkerSelectionStage { async fn execute(&self, ctx: &mut RequestContext) -> Result, Response> { + let pending_distribution_seed = ctx.state.pending_distribution_seed.take(); let prep = ctx.state.preparation.as_ref().ok_or_else(|| { error!( function = "WorkerSelectionStage::execute", @@ -93,85 +99,98 @@ impl PipelineStage for WorkerSelectionStage { let headers = ctx.input.headers.as_ref(); let model_id = ctx.input.model_id.as_str(); - let workers = match self.mode { - WorkerSelectionMode::Regular => { - match self.select_single_worker(model_id, text, tokens, headers) { - Some((worker, reservation)) => { - ctx.state.policy_reservation = reservation; - WorkerSelection::Single { worker } - } - None => { - error!( - function = "WorkerSelectionStage::execute", - mode = "Regular", - model_id = %model_id, - "No available workers for model" - ); - return Err(error::model_not_found(model_id)); - } - } + let workers = if let Some(pending) = pending_distribution_seed { + if self.mode != WorkerSelectionMode::Regular { + return Err(rejection_response(pending.retry_after_secs)); } - WorkerSelectionMode::PrefillDecode => { - match self.select_pd_pair(model_id, text, tokens, headers) { - Some((prefill, decode, runtime_type)) => WorkerSelection::Disaggregated { - encode_assignments: None, - prefill, - decode, - runtime_type, - }, - None => { - error!( - function = "WorkerSelectionStage::execute", - mode = "PrefillDecode", - model_id = %model_id, - "No available PD worker pairs for model" - ); - return Err(error::model_not_found(model_id)); + let Some((worker, guard)) = + self.select_distribution_seed(ctx, &pending, model_id, text, tokens, headers) + else { + return Err(rejection_response(pending.retry_after_secs)); + }; + ctx.state.distribution_seed_guard = Some(guard); + WorkerSelection::Single { worker } + } else { + match self.mode { + WorkerSelectionMode::Regular => { + match self.select_single_worker(model_id, text, tokens, headers) { + Some((worker, reservation)) => { + ctx.state.policy_reservation = reservation; + WorkerSelection::Single { worker } + } + None => { + error!( + function = "WorkerSelectionStage::execute", + mode = "Regular", + model_id = %model_id, + "No available workers for model" + ); + return Err(error::model_not_found(model_id)); + } } } - } - WorkerSelectionMode::EncodePrefillDecode => { - let encode_item_hashes = match encode_item_hashes(intermediate) { - Ok(hashes) => hashes, - Err(err) => { - error!( - function = "WorkerSelectionStage::execute", - error = %err, - "Failed to derive encode item routing hashes" - ); - return Err(error::internal_error( - "encode_routing_hash_failed", - format!("Failed to derive encode routing hashes: {err}"), - )); - } - }; - match self.select_encode_prefill_decode_workers( - model_id, - text, - tokens, - headers, - &encode_item_hashes, - ) { - Some((encode_assignments, prefill, decode, runtime_type)) => { - WorkerSelection::Disaggregated { - encode_assignments: if encode_assignments.is_empty() { - None - } else { - Some(encode_assignments) - }, + WorkerSelectionMode::PrefillDecode => { + match self.select_pd_pair(model_id, text, tokens, headers) { + Some((prefill, decode, runtime_type)) => WorkerSelection::Disaggregated { + encode_assignments: None, prefill, decode, runtime_type, + }, + None => { + error!( + function = "WorkerSelectionStage::execute", + mode = "PrefillDecode", + model_id = %model_id, + "No available PD worker pairs for model" + ); + return Err(error::model_not_found(model_id)); } } - None => { - error!( - function = "WorkerSelectionStage::execute", - mode = "EncodePrefillDecode", - model_id = %model_id, - "No available encode/prefill/decode worker set for model" - ); - return Err(error::model_not_found(model_id)); + } + WorkerSelectionMode::EncodePrefillDecode => { + let encode_item_hashes = match encode_item_hashes(intermediate) { + Ok(hashes) => hashes, + Err(err) => { + error!( + function = "WorkerSelectionStage::execute", + error = %err, + "Failed to derive encode item routing hashes" + ); + return Err(error::internal_error( + "encode_routing_hash_failed", + format!("Failed to derive encode routing hashes: {err}"), + )); + } + }; + match self.select_encode_prefill_decode_workers( + model_id, + text, + tokens, + headers, + &encode_item_hashes, + ) { + Some((encode_assignments, prefill, decode, runtime_type)) => { + WorkerSelection::Disaggregated { + encode_assignments: if encode_assignments.is_empty() { + None + } else { + Some(encode_assignments) + }, + prefill, + decode, + runtime_type, + } + } + None => { + error!( + function = "WorkerSelectionStage::execute", + mode = "EncodePrefillDecode", + model_id = %model_id, + "No available encode/prefill/decode worker set for model" + ); + return Err(error::model_not_found(model_id)); + } } } } @@ -222,6 +241,104 @@ fn worker_is_available_for_request(worker: &dyn Worker, headers: Option<&HeaderM } impl WorkerSelectionStage { + fn select_distribution_seed( + &self, + ctx: &RequestContext, + pending: &PendingDistributionSeed, + model_id: &str, + text: Option<&str>, + tokens: Option<&[u32]>, + headers: Option<&HeaderMap>, + ) -> Option<(Arc, DistributionSeedDispatchGuard)> { + let controller = ctx.components.adaptive_admission.as_ref()?; + let targets = controller.distribution_headroom_snapshot(&pending.partition, model_id); + if targets.is_empty() { + return None; + } + + let workers: Vec<_> = self + .worker_registry + .get_workers_filtered( + Some(model_id), + Some(WorkerType::Regular), + Some(ConnectionMode::Grpc), + None, + false, + ) + .into_iter() + .filter(|worker| worker_is_available_for_request(worker.as_ref(), headers)) + .collect(); + if workers.len() < 2 { + return None; + } + + let info = SelectWorkerInfo { + request_text: text, + tokens, + headers, + hash_ring: self.worker_registry.get_hash_ring(model_id), + max_output_tokens: None, + reserve_work: false, + leg: WorkerLeg::Single, + }; + let capacity: Vec<_> = targets + .iter() + .map(|target| SeedWorkerHeadroom { + worker_url: Arc::from(target.worker_url()), + worker_revision: target.worker_revision(), + issuable_slots: target.issuable_slots(), + }) + .collect(); + let plan = self + .policy_registry + .owner_pressure_dispatch_plan(model_id, &workers, &info, &capacity)?; + let target = targets.iter().find(|target| { + target.worker_url() == plan.target_worker_url() + && target.worker_revision() == plan.target_worker_revision() + })?; + let lease = controller.try_acquire_distribution_headroom(target)?; + + let meta = ctx.input.tenant_request_meta.as_ref()?; + let proof = meta.extension::()?; + let claimed = proof + .try_claim( + meta.tenant_key(), + meta.request_charge_id(), + model_id, + &pending.partition, + ) + .ok()?; + + let selected = workers.into_iter().find(|worker| { + worker.url() == plan.target_worker_url() + && worker.revision() == plan.target_worker_revision() + && worker_is_available_for_request(worker.as_ref(), headers) + })?; + if !lease.verify() { + return None; + } + + Metrics::record_worker_selection( + metrics_labels::WORKER_REGULAR, + metrics_labels::CONNECTION_GRPC, + model_id, + if plan.expands_ownership() { + "cache_aware_distribution_seed" + } else { + "cache_aware_owner_rebalance" + }, + ); + Some(( + selected, + DistributionSeedDispatchGuard { + headroom: lease, + _scheduler_proof: claimed, + _policy_plan: plan, + retry_after_secs: pending.retry_after_secs, + }, + )) + } + fn select_single_worker( &self, model_id: &str, diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index ab19bb7e5..c77060c1c 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -22,7 +22,9 @@ use tool_parser::ParserFactory as ToolParserFactory; use tracing::debug; use super::{ - adaptive_admission::{AdaptiveAdmissionController, AdaptiveRequestTracker}, + adaptive_admission::{ + AdaptiveAdmissionController, AdaptiveRequestTracker, DistributionHeadroomLease, + }, client::GrpcClient, common::stages::encode::EncodeDispatchPlan, multimodal::{MultimodalComponents, MultimodalIntermediate}, @@ -32,8 +34,8 @@ use super::{ }, }; use crate::{ - middleware::TenantRequestMeta, - policies::LoadBalancingPolicy, + middleware::{scheduler::ClaimedSchedulerAdmissionProof, TenantRequestMeta}, + policies::{LoadBalancingPolicy, OwnerPressureDispatchPlan}, worker::{RuntimeType, Worker, WorkerLoadGuard, WorkerRegistry}, }; @@ -176,6 +178,16 @@ pub(crate) struct ProcessingState { /// releases predicted outstanding work without training on a partial run. pub adaptive_request: Option, + /// An aggregate adaptive rejection that may proceed only through the + /// cache-aware, scheduler-authorized clean-peer recovery path in worker + /// selection. No ordinary policy selection is allowed while this is set. + pub pending_distribution_seed: Option, + + /// Non-cloneable authorization and target-capacity guards for an exact + /// clean-peer dispatch. Moved into the request-lifetime load guard before + /// backend execution; every earlier failure releases both by `Drop`. + pub distribution_seed_guard: Option, + // Stage 2: Worker selection outputs pub workers: Option, @@ -200,6 +212,20 @@ pub(crate) struct ProcessingState { pub response: ResponseState, } +#[derive(Debug)] +pub(crate) struct PendingDistributionSeed { + pub(crate) partition: String, + pub(crate) retry_after_secs: u32, +} + +#[derive(Debug)] +pub(crate) struct DistributionSeedDispatchGuard { + pub(crate) headroom: DistributionHeadroomLease, + pub(crate) _scheduler_proof: ClaimedSchedulerAdmissionProof, + pub(crate) _policy_plan: OwnerPressureDispatchPlan, + pub(crate) retry_after_secs: u32, +} + /// Per-item bootstrap rendezvous info for prefill, plus the dispatch plan that /// fans out to encode workers. /// @@ -337,6 +363,16 @@ pub(crate) struct CompletionItem { } impl PreparationOutput { + /// Number of independent backend requests represented by this prepared + /// input. A single scheduler proof and headroom lease may never authorize + /// a batched completion fan-out. + pub(crate) fn backend_request_count(&self) -> usize { + match self { + Self::Completion { items, .. } => items.len(), + _ => 1, + } + } + /// Total prompt tokens represented by this request. Completion batches /// sum every prompt because they fan out into independent backend work. pub fn total_token_count(&self) -> u32 { @@ -468,6 +504,7 @@ pub(crate) enum LoadGuards { Single { _guard: WorkerLoadGuard, _policy_reservation: Option, + _distribution_seed: Option, }, /// Disaggregated guards cover the prefill+decode pair. EPD encode workers are /// assigned per item; their fire-and-supervise RPCs do not hold load guards. @@ -480,6 +517,7 @@ pub(crate) enum LoadGuards { Batch { _guards: Vec, _policy_reservation: Option, + _distribution_seed: Option, }, } @@ -488,16 +526,19 @@ impl LoadGuards { selection: &WorkerSelection, headers: Option<&HeaderMap>, policy_reservation: Option, + distribution_seed: Option, ) -> Self { match selection { WorkerSelection::Single { worker } => LoadGuards::Single { _guard: WorkerLoadGuard::new(worker.clone(), headers), _policy_reservation: policy_reservation, + _distribution_seed: distribution_seed, }, WorkerSelection::Disaggregated { prefill, decode, .. } => { debug_assert!(policy_reservation.is_none()); + debug_assert!(distribution_seed.is_none()); LoadGuards::Disaggregated { _prefill: WorkerLoadGuard::new(prefill.clone(), headers), _decode: WorkerLoadGuard::new(decode.clone(), headers), @@ -512,15 +553,17 @@ impl LoadGuards { headers: Option<&HeaderMap>, count: usize, policy_reservation: Option, + distribution_seed: Option, ) -> Self { if count <= 1 { - Self::new(selection, headers, policy_reservation) + Self::new(selection, headers, policy_reservation, distribution_seed) } else { Self::Batch { _guards: (0..count) - .map(|_| Self::new(selection, headers, None)) + .map(|_| Self::new(selection, headers, None, None)) .collect(), _policy_reservation: policy_reservation, + _distribution_seed: distribution_seed, } } } From aa51d482c9ca8af8c09967f98378260545c214cc Mon Sep 17 00:00:00 2001 From: David <12414531+DavidBellamy@users.noreply.github.com> Date: Tue, 11 Aug 2026 03:42:51 -0700 Subject: [PATCH 2/2] fix(router): bind distribution seeds fail closed Signed-off-by: David <12414531+DavidBellamy@users.noreply.github.com> --- model_gateway/src/config/validation.rs | 21 +- model_gateway/src/policies/cache_aware.rs | 1400 +++++++++++++++-- model_gateway/src/policies/mod.rs | 5 +- .../src/routers/grpc/adaptive_admission.rs | 436 ++++- .../grpc/common/stages/adaptive_admission.rs | 128 +- .../common/stages/distribution_seed_commit.rs | 290 ++++ .../src/routers/grpc/common/stages/mod.rs | 2 + .../grpc/common/stages/request_execution.rs | 4 + .../grpc/common/stages/worker_selection.rs | 498 +++++- model_gateway/src/routers/grpc/context.rs | 73 +- model_gateway/src/routers/grpc/pipeline.rs | 15 + 11 files changed, 2648 insertions(+), 224 deletions(-) create mode 100644 model_gateway/src/routers/grpc/common/stages/distribution_seed_commit.rs diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index 01536b4ce..a4b3663a3 100755 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -895,6 +895,17 @@ impl ConfigValidator { } } if adaptive.distribution_headroom_partition_seed_cap > 0 { + if adaptive.distribution_headroom_partition_seed_cap != 1 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.distribution_headroom_partition_seed_cap" + .to_string(), + value: adaptive + .distribution_headroom_partition_seed_cap + .to_string(), + reason: "Only a single in-flight seed per partition is currently supported" + .to_string(), + }); + } if adaptive.distribution_headroom_partitions.is_empty() { return Err(ConfigError::ValidationFailed { reason: "distribution headroom seed capacity requires a non-empty exact partition allowlist" @@ -2111,12 +2122,20 @@ mod tests { config .adaptive_admission - .distribution_headroom_partition_seed_cap = 2; + .distribution_headroom_partition_seed_cap = 1; assert!(ConfigValidator::validate(&config).is_err()); config.adaptive_admission.mode = AdaptiveAdmissionMode::Enforce; config.adaptive_admission.strategy = AdaptiveAdmissionStrategy::EngineFeedback; assert!(ConfigValidator::validate(&config).is_ok()); + config + .adaptive_admission + .distribution_headroom_partition_seed_cap = 2; + assert!(ConfigValidator::validate(&config).is_err()); + config + .adaptive_admission + .distribution_headroom_partition_seed_cap = 1; + config.priority_scheduler_enabled = true; config.priority_scheduler_adaptive_capacity = true; assert!(ConfigValidator::validate(&config).is_err()); diff --git a/model_gateway/src/policies/cache_aware.rs b/model_gateway/src/policies/cache_aware.rs index a4ef5879c..945bd5a72 100644 --- a/model_gateway/src/policies/cache_aware.rs +++ b/model_gateway/src/policies/cache_aware.rs @@ -69,7 +69,7 @@ use std::{ use dashmap::{mapref::entry::Entry, DashMap}; use kv_index::{compute_request_content_hashes, PositionalIndexer, TokenTree, Tree}; use openai_protocol::worker::WorkerLoadResponse; -use parking_lot::RwLock; +use parking_lot::{Mutex, RwLock}; use serde::{Deserialize, Serialize}; use tokio::sync::watch; use tracing::{debug, warn}; @@ -87,6 +87,7 @@ use crate::{ const ENGINE_LOAD_MAX_AGE: Duration = Duration::from_secs(30); const ENGINE_PRESSURE_SLACK: f64 = 0.10; const ENGINE_PRESSURE_HIGH_WATERMARK: f64 = 0.90; +const MAX_DISTRIBUTION_PROTECTIONS: usize = 64; #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ImbalanceReason { @@ -96,6 +97,31 @@ enum ImbalanceReason { RequestCount, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum CachedOwnerDecision { + Selected(usize), + Blocked, + NoMatch, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum CacheAwareSelection { + Selected(usize), + Blocked, + Unavailable, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct DistributionProtectionSnapshot { + seeds: Vec<(DistributionSeedPhase, Arc, u64)>, +} + +impl DistributionProtectionSnapshot { + pub(crate) fn is_empty(&self) -> bool { + self.seeds.is_empty() + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] enum PrefixKind { String, @@ -109,17 +135,31 @@ struct PrefixBudgetKey { kind: PrefixKind, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +enum DistributionSeedPhase { + Active, + SuccessCooldown, + FailureQuarantine, +} + #[derive(Debug, Clone)] -struct PrefixReplicationState { - last_spill: Instant, - provisional_owner: String, +enum PrefixReplicationKind { + Ordinary { + provisional_owner: String, + }, + DistributionSeed { + lease_id: u64, + phase: DistributionSeedPhase, + target_worker_url: Arc, + target_worker_revision: u64, + prefix_tokens: usize, + }, } #[derive(Debug, Clone)] -struct DistributionSeedPrefixState { - lease_id: u64, - active: bool, +struct PrefixReplicationState { last_transition: Instant, + kind: PrefixReplicationKind, } /// One atomic per-prefix expansion claim. Ordinary cache-aware requests never @@ -127,9 +167,11 @@ struct DistributionSeedPrefixState { /// active slot while retaining a conservative cooldown marker. #[derive(Debug)] struct DistributionSeedPrefixReservation { - state: Arc>, + state: Arc>, + protections: Arc>, key: PrefixBudgetKey, lease_id: u64, + completion: DistributionSeedCompletion, } impl Drop for DistributionSeedPrefixReservation { @@ -137,12 +179,67 @@ impl Drop for DistributionSeedPrefixReservation { let Entry::Occupied(mut entry) = self.state.entry(self.key) else { return; }; - if entry.get().lease_id != self.lease_id || !entry.get().active { + let PrefixReplicationKind::DistributionSeed { + lease_id, + phase, + target_worker_url, + target_worker_revision, + prefix_tokens, + } = &entry.get().kind + else { + return; + }; + if *lease_id != self.lease_id || !matches!(phase, DistributionSeedPhase::Active) { return; } + let target_worker_url = Arc::clone(target_worker_url); + let target_worker_revision = *target_worker_revision; + let prefix_tokens = *prefix_tokens; + let phase = if self.completion.is_committed() { + DistributionSeedPhase::SuccessCooldown + } else { + DistributionSeedPhase::FailureQuarantine + }; + + let transitioned_at = Instant::now(); + let kind = PrefixReplicationKind::DistributionSeed { + lease_id: self.lease_id, + phase, + target_worker_url, + target_worker_revision, + prefix_tokens, + }; let state = entry.get_mut(); - state.active = false; - state.last_transition = Instant::now(); + state.last_transition = transitioned_at; + state.kind = kind.clone(); + if let Some(mut protection) = self.protections.get_mut(&self.key) { + protection.last_transition = transitioned_at; + protection.kind = kind; + } + } +} + +/// Shared one-shot success signal between the terminal commit stage and the +/// cache-policy seed reservation. Only a fully processed backend completion +/// may publish a seeded worker as ordinary cache ownership. +#[derive(Debug, Clone)] +pub(crate) struct DistributionSeedCompletion { + committed: Arc, +} + +impl DistributionSeedCompletion { + fn new() -> Self { + Self { + committed: Arc::new(AtomicBool::new(false)), + } + } + + fn commit(&self) { + self.committed.store(true, Ordering::Release); + } + + fn is_committed(&self) -> bool { + self.committed.load(Ordering::Acquire) } } @@ -171,6 +268,7 @@ pub(crate) struct OwnerPressureDispatchPlan { target_worker_url: Arc, target_worker_revision: u64, expands_ownership: bool, + completion: DistributionSeedCompletion, _prefix_reservation: Option, } @@ -186,6 +284,10 @@ impl OwnerPressureDispatchPlan { pub(crate) fn expands_ownership(&self) -> bool { self.expands_ownership } + + pub(crate) fn commit_success(&mut self) { + self.completion.commit(); + } } #[derive(Debug, Clone)] @@ -268,12 +370,14 @@ pub struct CacheAwarePolicy { engine_loads: RwLock>, /// Per-prefix replication throttle. This controls creation of new owners, /// not the authoritative owner catalog, which always records every worker - /// reported by the backend event stream. + /// reported by the backend event stream. Ordinary provisional owners and + /// scheduler-authorized distribution seeds share this map so they cannot + /// independently expand the same prefix. replication_state: Arc>, - /// Exact per-prefix serialization for scheduler-authorized clean-peer - /// seeds. Kept separate from `replication_state` so ordinary requests - /// cannot join a seed before backend KV events establish ownership. - distribution_seed_state: Arc>, + /// Seed-only descendant protections. This stays small and bounded so the + /// request hot path never scans the ordinary spill history. + distribution_protections: Arc>, + distribution_protection_lock: Mutex<()>, next_distribution_seed_lease_id: AtomicU64, _replication_gc_task: Option, } @@ -300,8 +404,8 @@ impl CacheAwarePolicy { let token_trees = Arc::new(DashMap::>::new()); let hash_index = Arc::new(DashMap::::new()); let replication_state = Arc::new(DashMap::::new()); - let distribution_seed_state = - Arc::new(DashMap::::new()); + let distribution_protections = + Arc::new(DashMap::::new()); // Start background eviction thread if configured let eviction_task = if config.eviction_interval_secs > 0 { @@ -384,16 +488,30 @@ impl CacheAwarePolicy { && config.cache_owner_spill_cooldown_secs > 0 { let state = Arc::clone(&replication_state); - let seed_state = Arc::clone(&distribution_seed_state); + let protections = Arc::clone(&distribution_protections); let cooldown = Duration::from_secs(config.cache_owner_spill_cooldown_secs); let retention = cooldown.saturating_mul(4).max(Duration::from_secs(60)); Some(PeriodicTask::spawn( config.cache_owner_spill_cooldown_secs.max(1), "Prefix replication budget GC", move || { - state.retain(|_, entry| entry.last_spill.elapsed() <= retention); - seed_state.retain(|_, entry| { - entry.active || entry.last_transition.elapsed() <= retention + state.retain(|_, entry| { + matches!( + &entry.kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } + ) || entry.last_transition.elapsed() <= retention + }); + protections.retain(|_, entry| { + matches!( + &entry.kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } + ) || entry.last_transition.elapsed() <= cooldown }); }, )) @@ -412,7 +530,8 @@ impl CacheAwarePolicy { populate_hash_index: AtomicBool::new(false), engine_loads: RwLock::new(HashMap::new()), replication_state, - distribution_seed_state, + distribution_protections, + distribution_protection_lock: Mutex::new(()), next_distribution_seed_lease_id: AtomicU64::new(0), _replication_gc_task: replication_gc_task, } @@ -1101,63 +1220,110 @@ impl TreeHandle for CacheAwarePolicy { } } -impl LoadBalancingPolicy for CacheAwarePolicy { - fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { +impl CacheAwarePolicy { + pub(crate) fn select_worker_decision( + &self, + workers: &[Arc], + info: &SelectWorkerInfo, + ) -> CacheAwareSelection { + self.select_worker_decision_with_active_targets(workers, info, &[]) + } + + pub(crate) fn select_worker_decision_with_active_targets( + &self, + workers: &[Arc], + info: &SelectWorkerInfo, + active_distribution_targets: &[(Arc, u64)], + ) -> CacheAwareSelection { + let model_id = workers + .first() + .map_or("", |worker| normalize_model_key(worker.model_id())); + let protections = self.distribution_protection_snapshot(model_id, info.tokens); + self.select_worker_decision_with_distribution_state( + workers, + info, + active_distribution_targets, + &protections, + ) + } + + pub(crate) fn select_worker_decision_with_distribution_state( + &self, + workers: &[Arc], + info: &SelectWorkerInfo, + active_distribution_targets: &[(Arc, u64)], + protections: &DistributionProtectionSnapshot, + ) -> CacheAwareSelection { let request_text = info.request_text; let request_tokens = info.tokens; // Single O(workers) gather: read each worker once via routing_state() // (status + load + processed under one ArcSwap guard), replacing the // former separate passes whose per-worker guard traffic dominated routing - // CPU at scale. Collects healthy indices and load min/max; cache-owner - // lookup is a hash-free scan over healthy_indices. + // CPU at scale. Cache-owner lookup is a hash-free scan over all healthy + // workers, while pressure/fallback considers only header-eligible + // workers so an excluded idle peer cannot suppress an allowed peer. let mut healthy_indices: Vec = Vec::with_capacity(workers.len()); - let mut min_load = usize::MAX; - let mut max_load = 0usize; for (idx, worker) in workers.iter().enumerate() { let state = worker.routing_state(); if state.healthy && state.can_execute { healthy_indices.push(idx); - min_load = min_load.min(state.load); - max_load = max_load.max(state.load); } } if healthy_indices.is_empty() { - return None; + return CacheAwareSelection::Unavailable; + } + let selectable_indices = + SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &healthy_indices); + if selectable_indices.is_empty() { + return CacheAwareSelection::Unavailable; + } + let mut min_load = usize::MAX; + let mut max_load = 0usize; + for &idx in &selectable_indices { + let state = workers[idx].routing_state(); + min_load = min_load.min(state.load); + max_load = max_load.max(state.load); } let min_load = if min_load == usize::MAX { 0 } else { min_load }; // The router pre-filters workers by model, so any healthy worker gives // us the model key before engine pressure narrows the candidate set. let model_id = normalize_model_key(workers[healthy_indices[0]].model_id()); - let pressure_plan = self.engine_pressure_plan(workers, &healthy_indices); + let pressure_plan = self.engine_pressure_plan(workers, &selectable_indices); // Prefix ownership is evaluated before fleet-wide imbalance. This is // the critical locality invariant: a hot unrelated worker cannot make // a usable cached owner disappear. Only pressure on every matching // owner may create one controlled additional owner. - if let Some(selected_idx) = self.select_cached_owner_or_pressure_spill( + match self.select_cached_owner_or_pressure_spill( workers, info, &healthy_indices, model_id, pressure_plan.as_ref(), + active_distribution_targets, + protections, ) { - let result = if pressure_plan.is_some() { - "cached_owner_scoped" - } else { - "telemetry_fallback" - }; - if self.config.engine_load { - Metrics::record_cache_aware_engine_decision(result); + CachedOwnerDecision::Selected(selected_idx) => { + let result = if pressure_plan.is_some() { + "cached_owner_scoped" + } else { + "telemetry_fallback" + }; + if self.config.engine_load { + Metrics::record_cache_aware_engine_decision(result); + } + return CacheAwareSelection::Selected(selected_idx); } - return Some(selected_idx); + CachedOwnerDecision::Blocked => return CacheAwareSelection::Blocked, + CachedOwnerDecision::NoMatch => {} } let selection_indices = match pressure_plan.as_ref() { Some(plan) => &plan.allowed_indices, - _ => &healthy_indices, + _ => &selectable_indices, }; // Engine pressure is the outer safety filter. When it narrows the @@ -1216,7 +1382,10 @@ impl LoadBalancingPolicy for CacheAwarePolicy { info, model_id, ) - }?; + }; + let Some(selected_idx) = selected_idx else { + return CacheAwareSelection::Unavailable; + }; if let Some(plan) = pressure_plan.as_ref() { let result = if plan.allowed_indices.len() == healthy_indices.len() { @@ -1237,7 +1406,16 @@ impl LoadBalancingPolicy for CacheAwarePolicy { Metrics::record_cache_aware_engine_decision("telemetry_fallback"); } - Some(selected_idx) + CacheAwareSelection::Selected(selected_idx) + } +} + +impl LoadBalancingPolicy for CacheAwarePolicy { + fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { + match self.select_worker_decision(workers, info) { + CacheAwareSelection::Selected(index) => Some(index), + CacheAwareSelection::Blocked | CacheAwareSelection::Unavailable => None, + } } fn update_loads(&self, loads: &HashMap) { @@ -1478,35 +1656,94 @@ impl CacheAwarePolicy { fn try_acquire_distribution_seed_prefix( &self, key: PrefixBudgetKey, + target_worker_url: Arc, + target_worker_revision: u64, + prefix_tokens: usize, + completion: DistributionSeedCompletion, ) -> Option { let cooldown = Duration::from_secs(self.config.cache_owner_spill_cooldown_secs); if cooldown.is_zero() { return None; } + let _protection_guard = self.distribution_protection_lock.lock(); + self.distribution_protections.retain(|_, entry| { + matches!( + &entry.kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } + ) || entry.last_transition.elapsed() < cooldown + }); + if self.distribution_protections.contains_key(&key) + || self.distribution_protections.len() >= MAX_DISTRIBUTION_PROTECTIONS + { + return None; + } let lease_id = self .next_distribution_seed_lease_id .fetch_add(1, Ordering::Relaxed) .wrapping_add(1); - let state = DistributionSeedPrefixState { - lease_id, - active: true, + let state = PrefixReplicationState { last_transition: Instant::now(), + kind: PrefixReplicationKind::DistributionSeed { + lease_id, + phase: DistributionSeedPhase::Active, + target_worker_url, + target_worker_revision, + prefix_tokens, + }, }; - match self.distribution_seed_state.entry(key) { + // Publish the conservative descendant protection before claiming the + // exact primary key. A reader may transiently over-block, but can never + // miss an in-progress seed and open a deeper owner. Roll back by exact + // lease id if the ordinary/shared primary claim wins. + self.distribution_protections.insert(key, state.clone()); + let acquired = match self.replication_state.entry(key) { Entry::Occupied(mut entry) => { - if entry.get().active || entry.get().last_transition.elapsed() < cooldown { - return None; + let active_distribution_seed = matches!( + &entry.get().kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } + ); + if active_distribution_seed || entry.get().last_transition.elapsed() < cooldown { + false + } else { + entry.insert(state.clone()); + true } - entry.insert(state); } Entry::Vacant(entry) => { - entry.insert(state); + entry.insert(state.clone()); + true + } + }; + if !acquired { + let owns_protection = self + .distribution_protections + .get(&key) + .is_some_and(|entry| { + matches!( + &entry.kind, + PrefixReplicationKind::DistributionSeed { + lease_id: current, + .. + } if *current == lease_id + ) + }); + if owns_protection { + self.distribution_protections.remove(&key); } + return None; } Some(DistributionSeedPrefixReservation { - state: Arc::clone(&self.distribution_seed_state), + state: Arc::clone(&self.replication_state), + protections: Arc::clone(&self.distribution_protections), key, lease_id, + completion, }) } @@ -1521,11 +1758,10 @@ impl CacheAwarePolicy { info: &SelectWorkerInfo<'_>, headroom: &[SeedWorkerHeadroom], ) -> Option { - if !self.config.engine_load { + let tokens = info.tokens?; + if self.requires_distribution_serialization(model_id, Some(tokens)) { return None; } - - let tokens = info.tokens?; let healthy_indices = super::get_healthy_worker_indices(workers); // A seed may expand only ownership reported by the backend KV event // index. The approximate trees and the normal routing path's @@ -1577,10 +1813,12 @@ impl CacheAwarePolicy { .first() .filter(|(_, issuable_slots)| *issuable_slots > 0) { + let completion = DistributionSeedCompletion::new(); return Some(OwnerPressureDispatchPlan { target_worker_url: Arc::from(workers[target_idx].url()), target_worker_revision: workers[target_idx].revision(), expands_ownership: false, + completion, _prefix_reservation: None, }); } @@ -1622,11 +1860,20 @@ impl CacheAwarePolicy { .then_with(|| workers[*left_idx].url().cmp(workers[*right_idx].url())) }); let (target_idx, _) = *clean.first()?; - let prefix_reservation = self.try_acquire_distribution_seed_prefix(prefix_key)?; + let target_worker_url: Arc = Arc::from(workers[target_idx].url()); + let completion = DistributionSeedCompletion::new(); + let prefix_reservation = self.try_acquire_distribution_seed_prefix( + prefix_key, + Arc::clone(&target_worker_url), + workers[target_idx].revision(), + matched_tokens, + completion.clone(), + )?; Some(OwnerPressureDispatchPlan { - target_worker_url: Arc::from(workers[target_idx].url()), + target_worker_url, target_worker_revision: workers[target_idx].revision(), expands_ownership: true, + completion, _prefix_reservation: Some(prefix_reservation), }) } @@ -1655,6 +1902,76 @@ impl CacheAwarePolicy { } } + /// Active seeds and recently failed seeds protect only descendants of the + /// exact token prefix they claimed. This closes transient owner-loss and + /// deeper-key races without stalling unrelated cold prefixes in the same + /// model. + pub(crate) fn distribution_protection_snapshot( + &self, + model_id: &str, + tokens: Option<&[u32]>, + ) -> DistributionProtectionSnapshot { + let Some(tokens) = tokens else { + return DistributionProtectionSnapshot::default(); + }; + let model_hash = kv_index::hash_node_path(model_id); + let cooldown = Duration::from_secs(self.config.cache_owner_spill_cooldown_secs); + let mut targets = Vec::new(); + for entry in self.distribution_protections.iter() { + let key = *entry.key(); + if key.kind != PrefixKind::Token || key.model_hash != model_hash { + continue; + } + let PrefixReplicationKind::DistributionSeed { + phase, + target_worker_url, + target_worker_revision, + prefix_tokens, + .. + } = &entry.kind + else { + continue; + }; + if (!matches!(phase, DistributionSeedPhase::Active) + && (cooldown.is_zero() || entry.last_transition.elapsed() >= cooldown)) + || *prefix_tokens == 0 + || *prefix_tokens > tokens.len() + || kv_index::hash_token_path(&tokens[..*prefix_tokens]) != key.prefix_hash + { + continue; + } + targets.push(( + *phase, + Arc::clone(target_worker_url), + *target_worker_revision, + )); + } + // DashMap iteration order is intentionally unspecified. Preserve every + // matching prefix protection, including multiple phases for the same + // target, then sort so the selection-time snapshot comparison is + // deterministic. In particular, a FailureQuarantine may never be + // discarded just because a SuccessCooldown for the same worker was + // observed first. + targets.sort_unstable_by(|left, right| { + left.1 + .as_ref() + .cmp(right.1.as_ref()) + .then_with(|| left.2.cmp(&right.2)) + .then_with(|| left.0.cmp(&right.0)) + }); + DistributionProtectionSnapshot { seeds: targets } + } + + pub(crate) fn requires_distribution_serialization( + &self, + model_id: &str, + tokens: Option<&[u32]>, + ) -> bool { + !self + .distribution_protection_snapshot(model_id, tokens) + .is_empty() + } + /// Prefer the longest-prefix owners before applying fleet-wide imbalance. /// /// Existing owners are balanced by owner-local pressure and atomic reserved @@ -1662,6 +1979,10 @@ impl CacheAwarePolicy { /// pressure high-water mark and the configured replication ceiling has not /// been reached. Concurrent spill requests coalesce on one provisional /// destination until backend cache events catch up. + #[expect( + clippy::too_many_arguments, + reason = "keeps cache ownership, pressure, and distribution snapshots explicit" + )] fn select_cached_owner_or_pressure_spill( &self, workers: &[Arc], @@ -1669,22 +1990,112 @@ impl CacheAwarePolicy { healthy_indices: &[usize], model_id: &str, pressure_plan: Option<&EnginePressurePlan>, - ) -> Option { + active_distribution_targets: &[(Arc, u64)], + protections: &DistributionProtectionSnapshot, + ) -> CachedOwnerDecision { if self.config.max_cached_owners_per_prefix == 0 { - return None; + return CachedOwnerDecision::NoMatch; } let ownership = self.prefix_ownership(workers, info, healthy_indices, model_id); let known_owners = self.owners_with_provisional(&ownership, workers, healthy_indices); let owner_count = known_owners.len(); + let protected_seeds = &protections.seeds; + if protected_seeds + .iter() + .any(|(phase, _, _)| matches!(phase, DistributionSeedPhase::FailureQuarantine)) + { + Metrics::record_cache_aware_replication_decision( + model_id, + "failed_distribution_seed_block", + owner_count, + ); + return CachedOwnerDecision::Blocked; + } if owner_count == 0 { - return None; + return if protected_seeds.is_empty() { + CachedOwnerDecision::NoMatch + } else { + Metrics::record_cache_aware_replication_decision( + model_id, + "distribution_seed_missing_owner_block", + 0, + ); + CachedOwnerDecision::Blocked + }; } - let eligible_owners = + let mut eligible_owners = SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &known_owners); + eligible_owners.retain(|idx| { + !protected_seeds.iter().any(|(phase, target, revision)| { + matches!(phase, DistributionSeedPhase::Active) + && target.as_ref() == workers[*idx].url() + && *revision == workers[*idx].revision() + }) + }); + if !active_distribution_targets.is_empty() + || protected_seeds + .iter() + .any(|(phase, _, _)| matches!(phase, DistributionSeedPhase::Active)) + { + Metrics::record_cache_aware_replication_decision( + model_id, + "active_distribution_seed_hold", + owner_count, + ); + return self + .select_from_cached_owners(workers, info, &eligible_owners, model_id) + .map_or(CachedOwnerDecision::Blocked, CachedOwnerDecision::Selected); + } + if protected_seeds + .iter() + .any(|(phase, _, _)| matches!(phase, DistributionSeedPhase::SuccessCooldown)) + { + Metrics::record_cache_aware_replication_decision( + model_id, + "successful_distribution_seed_cooldown_hold", + owner_count, + ); + return self + .select_from_cached_owners(workers, info, &eligible_owners, model_id) + .map_or(CachedOwnerDecision::Blocked, CachedOwnerDecision::Selected); + } if eligible_owners.is_empty() { - return None; + let candidate_indices = pressure_plan + .map(|plan| plan.allowed_indices.as_slice()) + .unwrap_or(healthy_indices); + let spill_candidates: Vec<_> = candidate_indices + .iter() + .copied() + .filter(|idx| !known_owners.contains(idx)) + .collect(); + let spill_candidates = + SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &spill_candidates); + let Some((selected, claimed_new_owner)) = self.select_or_join_replication_spill( + ownership.key, + workers, + info, + &spill_candidates, + model_id, + ) else { + Metrics::record_cache_aware_replication_decision( + model_id, + "seed_cooldown_excluded_owner_block", + owner_count, + ); + return CachedOwnerDecision::Blocked; + }; + Metrics::record_cache_aware_replication_decision( + model_id, + if claimed_new_owner { + "excluded_owner_spill" + } else { + "excluded_owner_cooldown_hold" + }, + owner_count, + ); + return CachedOwnerDecision::Selected(selected); } let suitable_owners = @@ -1695,13 +2106,17 @@ impl CacheAwarePolicy { "cached_owner_hold", owner_count, ); - return self.select_from_cached_owners(workers, info, &suitable_owners, model_id); + return self + .select_from_cached_owners(workers, info, &suitable_owners, model_id) + .map_or(CachedOwnerDecision::NoMatch, CachedOwnerDecision::Selected); } let Some(plan) = pressure_plan else { // `suitable_cached_owners` only returns empty with a complete, // fresh pressure plan, but keep this fail-open guard explicit. - return self.select_from_cached_owners(workers, info, &eligible_owners, model_id); + return self + .select_from_cached_owners(workers, info, &eligible_owners, model_id) + .map_or(CachedOwnerDecision::NoMatch, CachedOwnerDecision::Selected); }; if owner_count < self.config.max_cached_owners_per_prefix { @@ -1717,20 +2132,29 @@ impl CacheAwarePolicy { let spill_candidates = SizeAwarePowerOfTwoPolicy::eligible_candidates(workers, info, &spill_candidates); if !spill_candidates.is_empty() { - let (selected, claimed_new_owner) = self.select_or_join_replication_spill( + let Some((selected, claimed_new_owner)) = self.select_or_join_replication_spill( ownership.key, workers, info, &spill_candidates, model_id, - )?; + ) else { + Metrics::record_cache_aware_replication_decision( + model_id, + "seed_cooldown_owner_hold", + owner_count, + ); + return self + .select_from_cached_owners(workers, info, &eligible_owners, model_id) + .map_or(CachedOwnerDecision::Blocked, CachedOwnerDecision::Selected); + }; let result = if claimed_new_owner { "owner_pressure_spill" } else { "spill_cooldown_hold" }; Metrics::record_cache_aware_replication_decision(model_id, result, owner_count); - return Some(selected); + return CachedOwnerDecision::Selected(selected); } } @@ -1738,10 +2162,13 @@ impl CacheAwarePolicy { // high-water mark, preserve affinity on the least-pressured owner. The // admission controller remains responsible for stopping new dispatch // when there is no safe capacity anywhere in the model pool. - let best_pressure = eligible_owners + let Some(best_pressure) = eligible_owners .iter() .map(|idx| plan.pressure_by_index[idx].pressure) - .min_by(f64::total_cmp)?; + .min_by(f64::total_cmp) + else { + return CachedOwnerDecision::NoMatch; + }; let least_pressured: Vec = eligible_owners .iter() .copied() @@ -1756,6 +2183,7 @@ impl CacheAwarePolicy { }; Metrics::record_cache_aware_replication_decision(model_id, result, owner_count); self.select_from_cached_owners(workers, info, &least_pressured, model_id) + .map_or(CachedOwnerDecision::NoMatch, CachedOwnerDecision::Selected) } fn recent_provisional_owner( @@ -1769,13 +2197,16 @@ impl CacheAwarePolicy { return None; } let state = self.replication_state.get(&key)?; - if state.last_spill.elapsed() >= cooldown { + if state.last_transition.elapsed() >= cooldown { return None; } + let PrefixReplicationKind::Ordinary { provisional_owner } = &state.kind else { + return None; + }; candidate_indices .iter() .copied() - .find(|&idx| workers[idx].url() == state.provisional_owner) + .find(|&idx| workers[idx].url() == provisional_owner) } /// Atomically claim the next replication slot for a prefix. The first @@ -1804,22 +2235,33 @@ impl CacheAwarePolicy { match self.replication_state.entry(key) { Entry::Occupied(mut entry) => { - let active = entry.get().last_spill.elapsed() < cooldown; - if active { - let provisional_owner = entry.get().provisional_owner.as_str(); - if let Some(idx) = candidate_indices - .iter() - .copied() - .find(|&idx| workers[idx].url() == provisional_owner) - { - drop(entry); - let selected = self.fallback.select_least_loaded_from_candidates( - workers, - info, - &[idx], - )?; - return Some((selected, false)); + let within_cooldown = entry.get().last_transition.elapsed() < cooldown; + match &entry.get().kind { + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } => { + return None; + } + PrefixReplicationKind::DistributionSeed { .. } if within_cooldown => { + return None; + } + PrefixReplicationKind::Ordinary { provisional_owner } if within_cooldown => { + if let Some(idx) = candidate_indices + .iter() + .copied() + .find(|&idx| workers[idx].url() == provisional_owner) + { + drop(entry); + let selected = self.fallback.select_least_loaded_from_candidates( + workers, + info, + &[idx], + )?; + return Some((selected, false)); + } } + _ => {} } // An expired claim, or an active provisional owner excluded by @@ -1827,8 +2269,10 @@ impl CacheAwarePolicy { let selected = self.select_worker_fallback(workers, info, candidate_indices, model_id)?; entry.insert(PrefixReplicationState { - last_spill: Instant::now(), - provisional_owner: workers[selected].url().to_string(), + last_transition: Instant::now(), + kind: PrefixReplicationKind::Ordinary { + provisional_owner: workers[selected].url().to_string(), + }, }); Some((selected, true)) } @@ -1836,8 +2280,10 @@ impl CacheAwarePolicy { let selected = self.select_worker_fallback(workers, info, candidate_indices, model_id)?; entry.insert(PrefixReplicationState { - last_spill: Instant::now(), - provisional_owner: workers[selected].url().to_string(), + last_transition: Instant::now(), + kind: PrefixReplicationKind::Ordinary { + provisional_owner: workers[selected].url().to_string(), + }, }); Some((selected, true)) } @@ -3373,13 +3819,50 @@ mod tests { } } + type DistributionLifecycleFixture = ( + CacheAwarePolicy, + Vec>, + Arc, + Vec, + ); + + fn distribution_lifecycle_fixture() -> DistributionLifecycleFixture { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); + policy.set_kv_event_monitor(Some(monitor)); + let headroom = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + (policy, workers, indexer, headroom) + } + #[test] - fn owner_pressure_seed_uses_only_authoritative_owners_and_does_not_mutate_them() { - let mut policy = CacheAwarePolicy::with_config(CacheAwareConfig { + fn owner_pressure_seed_does_not_require_global_engine_load_or_mutate_owners() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { cache_threshold: 0.0, eviction_interval_secs: 0, block_size: 4, - engine_load: true, + engine_load: false, max_cached_owners_per_prefix: 8, cache_owner_spill_cooldown_secs: 5, ..Default::default() @@ -3389,10 +3872,12 @@ mod tests { let monitor = Arc::new(KvEventMonitor::new(Some(4))); let indexer = setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); - monitor.indexers.insert("unknown".to_string(), indexer); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); policy.set_kv_event_monitor(Some(monitor)); - let mut headroom: Vec<_> = workers + let headroom: Vec<_> = workers .iter() .enumerate() .map(|(index, worker)| SeedWorkerHeadroom { @@ -3427,6 +3912,11 @@ mod tests { ); drop(monitor); + // Use a fresh prefix state for the existing-owner case. The clean-peer + // plan above intentionally keeps this exact prefix protected until its + // request completes and then through the configured cooldown. + let (mut policy, workers, _, mut headroom) = distribution_lifecycle_fixture(); + policy.config.engine_load = false; headroom[0].issuable_slots = 1; let plan = policy .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) @@ -3446,13 +3936,12 @@ mod tests { } #[test] - fn owner_pressure_seed_serializes_prefix_expansion_without_publishing_it() { + fn excluded_authoritative_owner_still_counts_toward_seed_owner_ceiling() { let policy = CacheAwarePolicy::with_config(CacheAwareConfig { cache_threshold: 0.0, eviction_interval_secs: 0, block_size: 4, - engine_load: true, - max_cached_owners_per_prefix: 8, + max_cached_owners_per_prefix: 2, cache_owner_spill_cooldown_secs: 5, ..Default::default() }); @@ -3461,46 +3950,166 @@ mod tests { let monitor = Arc::new(KvEventMonitor::new(Some(4))); let indexer = setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); - monitor.indexers.insert("unknown".to_string(), indexer); + let second_owner = indexer.intern_worker(workers[1].url()).unwrap(); + indexer + .apply_stored( + second_owner, + &[ + StoredBlock { + seq_hash: SequenceHash(1), + content_hash: compute_content_hash(&[1, 2, 3, 4]), + }, + StoredBlock { + seq_hash: SequenceHash(2), + content_hash: compute_content_hash(&[5, 6, 7, 8]), + }, + ], + None, + &mut WorkerBlockMap::default(), + ) + .unwrap(); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); policy.set_kv_event_monitor(Some(monitor)); - let headroom: Vec<_> = workers + let headroom = workers .iter() .enumerate() .map(|(index, worker)| SeedWorkerHeadroom { worker_url: Arc::from(worker.url()), worker_revision: worker.revision(), - issuable_slots: if index == 0 { 0 } else { 38 }, + issuable_slots: if index == 2 { 38 } else { 0 }, }) - .collect(); + .collect::>(); + let mut headers = http::HeaderMap::new(); + headers.insert( + "x-smg-excluded-worker-urls", + workers[0].url().parse().unwrap(), + ); let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; - let info = SelectWorkerInfo { - tokens: Some(&tokens), - ..Default::default() - }; - let key = PrefixBudgetKey { - model_hash: kv_index::hash_node_path("unknown"), - prefix_hash: kv_index::hash_token_path(&tokens), - kind: PrefixKind::Token, - }; - - let first = policy - .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) - .expect("first clean-peer expansion should claim the prefix"); + assert!(policy + .owner_pressure_dispatch_plan( + "unknown", + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + headers: Some(&headers), + ..Default::default() + }, + &headroom, + ) + .is_none()); + } + + #[test] + fn owner_pressure_seed_serializes_prefix_expansion_without_publishing_it() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + engine_load: true, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); + policy.set_kv_event_monitor(Some(monitor)); + policy.update_loads(&HashMap::from([ + (workers[0].url().to_string(), engine_load(0.99, 0.99, 64)), + (workers[1].url().to_string(), engine_load(0.10, 0.10, 0)), + (workers[2].url().to_string(), engine_load(0.10, 0.10, 0)), + ])); + let headroom: Vec<_> = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let info = SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }; + let key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path("unknown"), + prefix_hash: kv_index::hash_token_path(&tokens), + kind: PrefixKind::Token, + }; + + let first = policy + .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) + .expect("first clean-peer expansion should claim the prefix"); assert!(first.expands_ownership()); assert!(policy .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) .is_none()); - assert!( - policy.replication_state.get(&key).is_none(), - "ordinary requests must not see or join a pending distribution seed" + { + let state = policy.replication_state.get(&key).unwrap(); + assert!(matches!( + &state.kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::Active, + .. + } + )); + } + assert!(policy + .recent_provisional_owner(key, &workers, &[0, 1, 2]) + .is_none()); + assert!(policy + .select_or_join_replication_spill(key, &workers, &info, &[1, 2], "unknown") + .is_none()); + + let selected = policy + .select_worker(&workers, &info) + .expect("ordinary traffic should hold on an authoritative owner"); + assert_eq!(selected, 0); + assert_eq!( + CacheAwarePolicy::score_overlap(&workers, &tokens, &[0, 1, 2], &indexer, 4), + vec![0], + "a pending seed must not publish a second cache owner" + ); + + let mut headers = http::HeaderMap::new(); + headers.insert( + "x-smg-excluded-worker-urls", + workers[0].url().parse().unwrap(), ); + assert!(policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + headers: Some(&headers), + ..Default::default() + }, + ) + .is_none()); drop(first); assert!(policy .owner_pressure_dispatch_plan("unknown", &workers, &info, &headroom) .is_none()); + assert!(policy + .select_or_join_replication_spill(key, &workers, &info, &[1, 2], "unknown") + .is_none()); + policy + .replication_state + .get_mut(&key) + .unwrap() + .last_transition = Instant::now() - Duration::from_secs(6); policy - .distribution_seed_state + .distribution_protections .get_mut(&key) .unwrap() .last_transition = Instant::now() - Duration::from_secs(6); @@ -3509,6 +4118,555 @@ mod tests { .is_some()); } + #[test] + fn active_seed_target_blocks_deeper_prefix_key_drift() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); + policy.set_kv_event_monitor(Some(monitor)); + let headroom: Vec<_> = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + let source_tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let plan = policy + .owner_pressure_dispatch_plan( + "unknown", + &workers, + &SelectWorkerInfo { + tokens: Some(&source_tokens), + ..Default::default() + }, + &headroom, + ) + .expect("initial seed should claim the source prefix"); + assert_eq!(plan.target_worker_url(), workers[1].url()); + + let deeper_tokens = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]; + let target_id = indexer.intern_worker(workers[1].url()).unwrap(); + let blocks: Vec<_> = deeper_tokens + .chunks(4) + .enumerate() + .map(|(index, tokens)| StoredBlock { + seq_hash: SequenceHash(index as u64 + 1), + content_hash: compute_content_hash(tokens), + }) + .collect(); + indexer + .apply_stored(target_id, &blocks, None, &mut WorkerBlockMap::default()) + .unwrap(); + let mut headers = http::HeaderMap::new(); + headers.insert( + "x-smg-excluded-worker-urls", + workers[1].url().parse().unwrap(), + ); + let active_targets = vec![(Arc::from(workers[1].url()), workers[1].revision())]; + assert_eq!( + policy.select_worker_decision_with_active_targets( + &workers, + &SelectWorkerInfo { + tokens: Some(&deeper_tokens), + headers: Some(&headers), + ..Default::default() + }, + &active_targets, + ), + CacheAwareSelection::Blocked + ); + let deeper_key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path("unknown"), + prefix_hash: kv_index::hash_token_path(&deeper_tokens), + kind: PrefixKind::Token, + }; + assert!( + !policy.replication_state.contains_key(&deeper_key), + "a deeper seed-owned prefix must not create a third provisional owner" + ); + } + + #[test] + fn terminal_seed_result_controls_descendant_quarantine_and_success_cooldown() { + let source_tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let deeper_tokens = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]; + let source_key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path("unknown"), + prefix_hash: kv_index::hash_token_path(&source_tokens), + kind: PrefixKind::Token, + }; + + let (failed_policy, failed_workers, failed_indexer, failed_headroom) = + distribution_lifecycle_fixture(); + let failed_plan = failed_policy + .owner_pressure_dispatch_plan( + "unknown", + &failed_workers, + &SelectWorkerInfo { + tokens: Some(&source_tokens), + ..Default::default() + }, + &failed_headroom, + ) + .expect("failed seed fixture should acquire a clean target"); + let failed_target = failed_plan.target_worker_url().to_string(); + let failed_monitor = failed_policy.kv_monitor.read().as_ref().cloned().unwrap(); + failed_monitor + .indexers + .insert("unknown".to_string(), Arc::new(PositionalIndexer::new(4))); + assert_eq!( + failed_policy.select_worker_decision( + &failed_workers, + &SelectWorkerInfo { + tokens: Some(&deeper_tokens), + ..Default::default() + }, + ), + CacheAwareSelection::Blocked, + "transient owner loss must not escape an active ancestor seed" + ); + failed_monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&failed_indexer)); + let failed_target_id = failed_indexer.intern_worker(&failed_target).unwrap(); + let deeper_blocks: Vec<_> = deeper_tokens + .chunks(4) + .enumerate() + .map(|(index, tokens)| StoredBlock { + seq_hash: SequenceHash(index as u64 + 1), + content_hash: compute_content_hash(tokens), + }) + .collect(); + failed_indexer + .apply_stored( + failed_target_id, + &deeper_blocks, + None, + &mut WorkerBlockMap::default(), + ) + .unwrap(); + drop(failed_plan); + assert!(matches!( + &failed_policy + .distribution_protections + .get(&source_key) + .unwrap() + .kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::FailureQuarantine, + .. + } + )); + let deeper_info = SelectWorkerInfo { + tokens: Some(&deeper_tokens), + ..Default::default() + }; + assert_eq!( + failed_policy.select_worker_decision(&failed_workers, &deeper_info), + CacheAwareSelection::Blocked, + "an early KV event from a failed seed must not publish its target" + ); + assert!(failed_policy + .owner_pressure_dispatch_plan( + "unknown", + &failed_workers, + &deeper_info, + &failed_headroom, + ) + .is_none()); + assert!(matches!( + failed_policy.select_worker_decision( + &failed_workers, + &SelectWorkerInfo { + tokens: Some(&[91, 92, 93, 94]), + ..Default::default() + }, + ), + CacheAwareSelection::Selected(_) + )); + + let (success_policy, success_workers, success_indexer, success_headroom) = + distribution_lifecycle_fixture(); + let mut success_plan = success_policy + .owner_pressure_dispatch_plan( + "unknown", + &success_workers, + &SelectWorkerInfo { + tokens: Some(&source_tokens), + ..Default::default() + }, + &success_headroom, + ) + .expect("successful seed fixture should acquire a clean target"); + let success_target = success_plan.target_worker_url().to_string(); + let success_target_index = success_workers + .iter() + .position(|worker| worker.url() == success_target) + .unwrap(); + let success_target_id = success_indexer.intern_worker(&success_target).unwrap(); + success_indexer + .apply_stored( + success_target_id, + &deeper_blocks, + None, + &mut WorkerBlockMap::default(), + ) + .unwrap(); + success_plan.commit_success(); + success_plan.commit_success(); + drop(success_plan); + assert!(matches!( + &success_policy + .distribution_protections + .get(&source_key) + .unwrap() + .kind, + PrefixReplicationKind::DistributionSeed { + phase: DistributionSeedPhase::SuccessCooldown, + .. + } + )); + assert_eq!( + success_policy.select_worker_decision(&success_workers, &deeper_info), + CacheAwareSelection::Selected(success_target_index), + "terminal success may expose the target only through authoritative KV events" + ); + assert!(success_policy + .owner_pressure_dispatch_plan( + "unknown", + &success_workers, + &deeper_info, + &success_headroom, + ) + .is_none()); + } + + #[test] + fn failure_quarantine_dominates_overlapping_success_for_same_target() { + let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let (policy, workers, _, _) = distribution_lifecycle_fixture(); + let model_hash = kv_index::hash_node_path("unknown"); + let target_worker_url: Arc = Arc::from(workers[0].url()); + let target_worker_revision = workers[0].revision(); + for (index, (prefix_tokens, phase)) in [ + (4, DistributionSeedPhase::SuccessCooldown), + (8, DistributionSeedPhase::FailureQuarantine), + ] + .into_iter() + .enumerate() + { + policy.distribution_protections.insert( + PrefixBudgetKey { + model_hash, + prefix_hash: kv_index::hash_token_path(&tokens[..prefix_tokens]), + kind: PrefixKind::Token, + }, + PrefixReplicationState { + last_transition: Instant::now(), + kind: PrefixReplicationKind::DistributionSeed { + lease_id: index as u64 + 1, + phase, + target_worker_url: Arc::clone(&target_worker_url), + target_worker_revision, + prefix_tokens, + }, + }, + ); + } + + let snapshot = policy.distribution_protection_snapshot("unknown", Some(&tokens)); + assert_eq!(snapshot.seeds.len(), 2); + assert!(snapshot + .seeds + .iter() + .any(|(phase, _, _)| matches!(phase, DistributionSeedPhase::FailureQuarantine))); + assert_eq!( + policy.select_worker_decision_with_distribution_state( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }, + &[], + &snapshot, + ), + CacheAwareSelection::Blocked, + "a failed descendant must remain quarantined even when the same target has a successful ancestor" + ); + } + + #[test] + fn active_seed_blocks_ordinary_spill_when_existing_owner_moves_deeper() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + cache_threshold: 0.0, + eviction_interval_secs: 0, + block_size: 4, + engine_load: true, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks(workers[0].url(), &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); + policy.set_kv_event_monitor(Some(monitor)); + policy.update_loads(&HashMap::from([ + (workers[0].url().to_string(), engine_load(0.99, 0.99, 64)), + (workers[1].url().to_string(), engine_load(0.1, 0.1, 0)), + (workers[2].url().to_string(), engine_load(0.1, 0.1, 0)), + ])); + let headroom: Vec<_> = workers + .iter() + .enumerate() + .map(|(index, worker)| SeedWorkerHeadroom { + worker_url: Arc::from(worker.url()), + worker_revision: worker.revision(), + issuable_slots: if index == 0 { 0 } else { 38 }, + }) + .collect(); + let source_tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let plan = policy + .owner_pressure_dispatch_plan( + "unknown", + &workers, + &SelectWorkerInfo { + tokens: Some(&source_tokens), + ..Default::default() + }, + &headroom, + ) + .expect("initial seed should claim the source prefix"); + assert_eq!(plan.target_worker_url(), workers[1].url()); + + let deeper_tokens = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]; + let owner_id = indexer.intern_worker(workers[0].url()).unwrap(); + let blocks: Vec<_> = deeper_tokens + .chunks(4) + .enumerate() + .map(|(index, tokens)| StoredBlock { + seq_hash: SequenceHash(index as u64 + 1), + content_hash: compute_content_hash(tokens), + }) + .collect(); + indexer + .apply_stored(owner_id, &blocks, None, &mut WorkerBlockMap::default()) + .unwrap(); + let mut headers = http::HeaderMap::new(); + headers.insert( + "x-smg-excluded-worker-urls", + workers[1].url().parse().unwrap(), + ); + let active_targets = vec![(Arc::from(workers[1].url()), workers[1].revision())]; + assert_eq!( + policy.select_worker_decision_with_active_targets( + &workers, + &SelectWorkerInfo { + tokens: Some(&deeper_tokens), + headers: Some(&headers), + ..Default::default() + }, + &active_targets, + ), + CacheAwareSelection::Selected(0), + "ordinary traffic must hold on the existing owner while a seed is active" + ); + let deeper_key = PrefixBudgetKey { + model_hash: kv_index::hash_node_path("unknown"), + prefix_hash: kv_index::hash_token_path(&deeper_tokens), + kind: PrefixKind::Token, + }; + assert!(!policy.replication_state.contains_key(&deeper_key)); + } + + #[test] + fn excluded_idle_worker_does_not_remove_the_only_selectable_peer() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + eviction_interval_secs: 0, + engine_load: true, + ..Default::default() + }); + let workers = make_workers(&["http://idle:8000", "http://allowed:8000"]); + policy.init_workers(&workers); + policy.update_loads(&HashMap::from([ + (workers[0].url().to_string(), engine_load(0.1, 0.1, 0)), + (workers[1].url().to_string(), engine_load(0.8, 0.8, 0)), + ])); + let mut headers = http::HeaderMap::new(); + headers.insert( + "x-smg-excluded-worker-urls", + workers[0].url().parse().unwrap(), + ); + + assert_eq!( + policy.select_worker_decision( + &workers, + &SelectWorkerInfo { + tokens: Some(&[1, 2, 3, 4]), + headers: Some(&headers), + ..Default::default() + }, + ), + CacheAwareSelection::Selected(1) + ); + } + + #[test] + fn ordinary_spill_blocks_distribution_seed_until_cooldown_expires() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + eviction_interval_secs: 0, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + let key = PrefixBudgetKey { + model_hash: 7, + prefix_hash: 11, + kind: PrefixKind::Token, + }; + let info = SelectWorkerInfo { + tokens: Some(&[1, 2, 3, 4]), + ..Default::default() + }; + + assert!(policy + .select_or_join_replication_spill(key, &workers, &info, &[0, 1], "unknown") + .is_some()); + assert!(policy + .try_acquire_distribution_seed_prefix( + key, + Arc::from(workers[1].url()), + workers[1].revision(), + 4, + DistributionSeedCompletion::new(), + ) + .is_none()); + { + let state = policy.replication_state.get(&key).unwrap(); + assert!(matches!( + &state.kind, + PrefixReplicationKind::Ordinary { .. } + )); + } + + policy + .replication_state + .get_mut(&key) + .unwrap() + .last_transition = Instant::now() - Duration::from_secs(6); + assert!(policy + .try_acquire_distribution_seed_prefix( + key, + Arc::from(workers[1].url()), + workers[1].revision(), + 4, + DistributionSeedCompletion::new(), + ) + .is_some()); + } + + #[test] + fn ordinary_spill_and_distribution_seed_race_has_one_winner() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + eviction_interval_secs: 0, + max_cached_owners_per_prefix: 2, + cache_owner_spill_cooldown_secs: 60, + ..Default::default() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + let info = SelectWorkerInfo { + tokens: Some(&[1, 2, 3, 4]), + ..Default::default() + }; + let model_id = normalize_model_key(workers[0].model_id()); + + for prefix_hash in 0..32 { + let key = PrefixBudgetKey { + model_hash: 13, + prefix_hash, + kind: PrefixKind::Token, + }; + let start = Arc::new(std::sync::Barrier::new(2)); + let (distribution_won, ordinary_won) = std::thread::scope(|scope| { + let distribution_start = Arc::clone(&start); + let distribution_policy = &policy; + let distribution_target = Arc::from(workers[1].url()); + let distribution_revision = workers[1].revision(); + let distribution = scope.spawn(move || { + distribution_start.wait(); + distribution_policy + .try_acquire_distribution_seed_prefix( + key, + distribution_target, + distribution_revision, + 4, + DistributionSeedCompletion::new(), + ) + .is_some() + }); + let ordinary_start = Arc::clone(&start); + let ordinary_policy = &policy; + let ordinary_workers = &workers; + let ordinary_info = &info; + let ordinary = scope.spawn(move || { + ordinary_start.wait(); + ordinary_policy + .select_or_join_replication_spill( + key, + ordinary_workers, + ordinary_info, + &[1], + model_id, + ) + .is_some() + }); + (distribution.join().unwrap(), ordinary.join().unwrap()) + }); + + assert_ne!( + distribution_won, ordinary_won, + "exactly one expansion path must claim a prefix" + ); + let owners = policy.owners_with_provisional( + &PrefixOwnership { + key, + owners: vec![0], + }, + &workers, + &[0, 1], + ); + assert!( + owners.len() <= 2, + "the shared claim must preserve the per-prefix owner ceiling" + ); + } + + assert_eq!(policy.replication_state.len(), 32); + } + // -- score_overlap unit tests (scoring helper) -- #[test] diff --git a/model_gateway/src/policies/mod.rs b/model_gateway/src/policies/mod.rs index ab01f1ff9..756307813 100644 --- a/model_gateway/src/policies/mod.rs +++ b/model_gateway/src/policies/mod.rs @@ -27,7 +27,10 @@ pub(crate) mod utils; pub use bucket::BucketPolicy; pub use cache_aware::{CacheAwarePolicy, TreeHandle, TreeKind}; -pub(crate) use cache_aware::{OwnerPressureDispatchPlan, SeedWorkerHeadroom}; +pub(crate) use cache_aware::{ + CacheAwareSelection, DistributionProtectionSnapshot, OwnerPressureDispatchPlan, + SeedWorkerHeadroom, +}; pub use consistent_hashing::ConsistentHashingPolicy; pub use dp_min_token::MinimumTokensPolicy; pub use factory::PolicyFactory; diff --git a/model_gateway/src/routers/grpc/adaptive_admission.rs b/model_gateway/src/routers/grpc/adaptive_admission.rs index fd0929020..6ddd26c45 100644 --- a/model_gateway/src/routers/grpc/adaptive_admission.rs +++ b/model_gateway/src/routers/grpc/adaptive_admission.rs @@ -18,7 +18,7 @@ use std::{ use metrics::{counter, describe_counter, describe_gauge, describe_histogram, gauge, histogram}; use openai_protocol::worker::WorkerLoadResponse; -use parking_lot::Mutex; +use parking_lot::{Mutex, MutexGuard}; use tokio::sync::watch; use crate::{ @@ -543,6 +543,7 @@ struct DistributionWorkerTelemetry { worker_instance: Arc, worker_revision: u64, telemetry_revision: u64, + source_timestamp: Arc, observed_at: Instant, running_requests: u64, waiting_requests: u64, @@ -551,6 +552,10 @@ struct DistributionWorkerTelemetry { } impl DistributionWorkerTelemetry { + #[expect( + clippy::too_many_arguments, + reason = "constructs one immutable telemetry binding from registry and load identities" + )] fn from_load( partition: Arc, worker_url: Arc, @@ -561,6 +566,10 @@ impl DistributionWorkerTelemetry { observed_at: Instant, load: &WorkerLoadResponse, ) -> Option { + let source_timestamp = load.timestamp.trim(); + if source_timestamp.is_empty() { + return None; + } let rank_count = usize::try_from(load.dp_rank_count).ok()?; if rank_count == 0 || load.loads.len() != rank_count { return None; @@ -593,6 +602,9 @@ impl DistributionWorkerTelemetry { { return None; } + if rank.num_total_reqs != rank.num_running_reqs.checked_add(rank.num_waiting_reqs)? { + return None; + } running_requests = running_requests.checked_add(rank.num_running_reqs as u64)?; waiting_requests = waiting_requests.checked_add(rank.num_waiting_reqs as u64)?; max_running_requests = @@ -610,6 +622,7 @@ impl DistributionWorkerTelemetry { worker_instance, worker_revision, telemetry_revision, + source_timestamp: Arc::from(source_timestamp), observed_at, running_requests, waiting_requests, @@ -684,6 +697,51 @@ impl DistributionHeadroomTarget { pub(crate) fn issuable_slots(&self) -> u16 { self.issuable_slots } + + pub(crate) fn matches_worker(&self, worker: &Arc) -> bool { + self.binding.worker_url.as_ref() == worker.url() + && self.binding.worker_revision == worker.revision() + && Arc::ptr_eq(&self.binding.worker_instance, worker) + } +} + +/// Serializes the synchronous ordinary and distribution-seed selection +/// critical sections. It is deliberately short-lived and never crosses an +/// await point. +pub(crate) struct DistributionSelectionGuard<'a> { + controller: &'a AdaptiveAdmissionController, + _guard: MutexGuard<'a, ()>, +} + +#[derive(Debug)] +struct OrdinaryPredispatchBinding { + partition: Arc, + model: Arc, + worker_id: WorkerId, + worker_instance: Arc, + worker_revision: u64, + claims: u64, +} + +/// Exact ordinary-route claim held from worker selection until +/// `WorkerLoadGuard` increments the selected worker's live load. +#[derive(Debug)] +pub(crate) struct OrdinaryDistributionDispatchGuard { + controller: Weak, + partition: Arc, + model: Arc, + worker_url: Arc, + worker_id: WorkerId, + worker_instance: Arc, + worker_revision: u64, +} + +impl Drop for OrdinaryDistributionDispatchGuard { + fn drop(&mut self) { + if let Some(controller) = self.controller.upgrade() { + controller.release_ordinary_distribution_target(self); + } + } } /// One process-local reservation of a clean worker's engine slot. The lease @@ -727,6 +785,7 @@ struct WorkState { feedback_estimates: HashMap, distribution_telemetry: HashMap, active_distribution_leases: HashMap, + ordinary_predispatch_targets: HashMap, } #[derive(Debug, Clone, Copy)] @@ -826,6 +885,7 @@ pub(crate) struct AdaptiveAdmissionController { registry: Arc, capacity_revision: watch::Sender, telemetry_revision: AtomicU64, + distribution_selection: Mutex<()>, } impl AdaptiveAdmissionController { @@ -839,9 +899,17 @@ impl AdaptiveAdmissionController { registry, capacity_revision, telemetry_revision: AtomicU64::new(0), + distribution_selection: Mutex::new(()), }) } + pub(crate) fn lock_distribution_selection(&self) -> DistributionSelectionGuard<'_> { + DistributionSelectionGuard { + controller: self, + _guard: self.distribution_selection.lock(), + } + } + pub(crate) fn mode(&self) -> AdaptiveAdmissionMode { self.config.mode } @@ -954,6 +1022,23 @@ impl AdaptiveAdmissionController { let now = observed_at; let mut work = self.work.lock(); + // The WorkerMonitor watch snapshot merges independently polled model + // groups. An unrelated fast poll therefore republishes unchanged K3 + // entries. Preserve the original observation time and binding revision + // when the backend's per-response timestamp did not advance, so those + // republishes cannot keep stale distribution telemetry fresh. + for (url, telemetry) in &mut distribution_telemetry { + if let Some(previous) = work.distribution_telemetry.get(url) { + if telemetry.source_timestamp == previous.source_timestamp + && telemetry.worker_id == previous.worker_id + && Arc::ptr_eq(&telemetry.worker_instance, &previous.worker_instance) + && telemetry.worker_revision == previous.worker_revision + { + telemetry.observed_at = previous.observed_at; + telemetry.telemetry_revision = previous.telemetry_revision; + } + } + } work.capacities .retain(|partition, _| partitions.contains_key(partition)); work.feedback_estimates @@ -1043,10 +1128,10 @@ impl AdaptiveAdmissionController { .send_modify(|revision| *revision = revision.wrapping_add(1)); } - fn distribution_headroom_enabled(&self, partition: &str) -> bool { + pub(crate) fn distribution_headroom_enabled(&self, partition: &str) -> bool { self.config.mode == AdaptiveAdmissionMode::Enforce && self.config.strategy == AdaptiveAdmissionStrategy::EngineFeedback - && self.config.distribution_headroom_partition_seed_cap > 0 + && self.config.distribution_headroom_partition_seed_cap == 1 && self .config .distribution_headroom_partitions @@ -1054,6 +1139,144 @@ impl AdaptiveAdmissionController { .any(|allowed| allowed == partition) } + /// Exact targets currently reserved by in-flight distribution seeds for + /// this model/partition. Ordinary routing excludes these workers until the + /// request-lifetime lease is released, even if KV store events arrive + /// before the seed response completes. + pub(crate) fn active_distribution_targets( + &self, + partition: &str, + model: &str, + ) -> Vec<(Arc, u64)> { + let work = self.work.lock(); + work.active_distribution_leases + .values() + .filter(|binding| { + binding.partition.as_ref() == partition && binding.model.as_ref() == model + }) + .map(|binding| (Arc::clone(&binding.worker_url), binding.worker_revision)) + .collect() + } + + pub(crate) fn distribution_target_is_active( + &self, + partition: &str, + model: &str, + worker_url: &str, + worker_revision: u64, + ) -> bool { + self.work + .lock() + .active_distribution_leases + .get(worker_url) + .is_some_and(|binding| { + binding.partition.as_ref() == partition + && binding.model.as_ref() == model + && binding.worker_revision == worker_revision + }) + } + + pub(crate) fn try_reserve_ordinary_distribution_target( + self: &Arc, + selection: &DistributionSelectionGuard<'_>, + partition: &str, + model: &str, + worker: &Arc, + ) -> Option { + if !std::ptr::eq(self.as_ref(), selection.controller) + || !self.distribution_headroom_enabled(partition) + || self.registry.resolve_model_alias(model).is_some() + { + return None; + } + let worker_id = self.registry.get_id_by_url(worker.url())?; + let current = self.registry.get(&worker_id)?; + let worker_partition = current + .metadata() + .spec + .labels + .get(ADMISSION_PARTITION_LABEL) + .map(String::as_str) + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| current.model_id()); + if !Arc::ptr_eq(¤t, worker) + || current.revision() != worker.revision() + || !current.is_available() + || worker_partition != partition + || !self + .registry + .get_by_model(model) + .iter() + .any(|candidate| Arc::ptr_eq(candidate, worker)) + { + return None; + } + + let mut work = self.work.lock(); + if work.active_distribution_leases.contains_key(worker.url()) { + return None; + } + if let Some(binding) = work.ordinary_predispatch_targets.get_mut(worker.url()) { + if binding.partition.as_ref() != partition + || binding.model.as_ref() != model + || binding.worker_id != worker_id + || !Arc::ptr_eq(&binding.worker_instance, worker) + || binding.worker_revision != worker.revision() + { + return None; + } + binding.claims = binding.claims.checked_add(1)?; + } else { + work.ordinary_predispatch_targets.insert( + worker.url().to_string(), + OrdinaryPredispatchBinding { + partition: Arc::from(partition), + model: Arc::from(model), + worker_id: worker_id.clone(), + worker_instance: Arc::clone(worker), + worker_revision: worker.revision(), + claims: 1, + }, + ); + } + Some(OrdinaryDistributionDispatchGuard { + controller: Arc::downgrade(self), + partition: Arc::from(partition), + model: Arc::from(model), + worker_url: Arc::from(worker.url()), + worker_id, + worker_instance: Arc::clone(worker), + worker_revision: worker.revision(), + }) + } + + fn release_ordinary_distribution_target(&self, guard: &OrdinaryDistributionDispatchGuard) { + let mut work = self.work.lock(); + let remove = { + let Some(binding) = work + .ordinary_predispatch_targets + .get_mut(guard.worker_url.as_ref()) + else { + return; + }; + if binding.partition != guard.partition + || binding.model != guard.model + || binding.worker_id != guard.worker_id + || !Arc::ptr_eq(&binding.worker_instance, &guard.worker_instance) + || binding.worker_revision != guard.worker_revision + || binding.claims == 0 + { + return; + } + binding.claims -= 1; + binding.claims == 0 + }; + if remove { + work.ordinary_predispatch_targets + .remove(guard.worker_url.as_ref()); + } + } + fn current_model_worker( &self, binding: &DistributionHeadroomBinding, @@ -1162,9 +1385,17 @@ impl AdaptiveAdmissionController { } let active_target_leases = u64::from(work.active_distribution_leases.contains_key(worker.url())); - let raw_issuable = telemetry.issuable_slots(worker.load(), active_target_leases); + let ordinary_predispatch = work + .ordinary_predispatch_targets + .get(worker.url()) + .map_or(0, |binding| binding.claims); + let raw_issuable = telemetry.issuable_slots( + worker.load(), + active_target_leases.saturating_add(ordinary_predispatch), + ); let issuable_slots = if partition_has_capacity && active_target_leases == 0 + && ordinary_predispatch == 0 && telemetry.is_clean(self.config.feedback_max_token_usage) { raw_issuable @@ -1185,10 +1416,14 @@ impl AdaptiveAdmissionController { /// separate decisions around this process-local capacity reservation. pub(crate) fn try_acquire_distribution_headroom( self: &Arc, + selection: &DistributionSelectionGuard<'_>, target: &DistributionHeadroomTarget, ) -> Option { let binding = &target.binding; - if target.issuable_slots == 0 || !self.distribution_headroom_enabled(&binding.partition) { + if !std::ptr::eq(self.as_ref(), selection.controller) + || target.issuable_slots == 0 + || !self.distribution_headroom_enabled(&binding.partition) + { return None; } @@ -1200,6 +1435,9 @@ impl AdaptiveAdmissionController { || work .active_distribution_leases .contains_key(binding.worker_url.as_ref()) + || work + .ordinary_predispatch_targets + .contains_key(binding.worker_url.as_ref()) { return None; } @@ -1565,7 +1803,7 @@ impl AdaptiveCapacityProvider for AdaptiveAdmissionController { &load, work.feedback_estimates.get(partition), ); - let ordinary_capacity = if !constraints.telemetry_usable { + if !constraints.telemetry_usable { static_capacity } else if constraints.pressure_reason.is_some() { 0 @@ -1573,8 +1811,7 @@ impl AdaptiveCapacityProvider for AdaptiveAdmissionController { (running_limit.floor().clamp(0.0, f64::from(u16::MAX)) as u16).min(static_capacity) } else { static_capacity - }; - ordinary_capacity + } } fn subscribe_capacity_changes(&self) -> watch::Receiver { @@ -1672,7 +1909,7 @@ mod tests { }; use super::*; - use crate::worker::BasicWorkerBuilder; + use crate::worker::{BasicWorkerBuilder, WorkerLoadGuard}; fn config() -> AdaptiveAdmissionConfig { AdaptiveAdmissionConfig { @@ -1748,14 +1985,31 @@ mod tests { controller: &AdaptiveAdmissionController, loads: impl IntoIterator, ) { + static TEST_LOAD_REVISION: AtomicU64 = AtomicU64::new(0); controller.update_loads( &loads .into_iter() - .map(|(url, load)| (url.to_string(), load)) + .map(|(url, mut load)| { + if load.timestamp.is_empty() { + load.timestamp = format!( + "test-{}", + TEST_LOAD_REVISION.fetch_add(1, Ordering::Relaxed) + ); + } + (url.to_string(), load) + }) .collect(), ); } + fn acquire_distribution_headroom( + controller: &Arc, + target: &DistributionHeadroomTarget, + ) -> Option { + let selection = controller.lock_distribution_selection(); + controller.try_acquire_distribution_headroom(&selection, target) + } + fn features(user: &str, prompt_tokens: u32, maximum: Option) -> PredictionFeatures { PredictionFeatures { model: "model".to_string(), @@ -2296,7 +2550,7 @@ mod tests { for _ in 0..5 { idle_b.increment_load(); } - let controller = AdaptiveAdmissionController::new(distribution_config(2), registry); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); update_worker_loads( &controller, [ @@ -2306,7 +2560,11 @@ mod tests { ], ); - assert_eq!(controller.effective_capacity("k3", 512), 109); + assert_eq!( + controller.effective_capacity("k3", 512), + 0, + "routing headroom must not mint untyped scheduler capacity" + ); let targets = controller.distribution_headroom_snapshot("k3", "model"); assert_eq!(targets.len(), 3); assert_eq!(targets[0].worker_url(), HOT); @@ -2383,9 +2641,7 @@ mod tests { .into_iter() .find(|target| target.issuable_slots() > 0) .unwrap(); - let lease = controller - .try_acquire_distribution_headroom(&target) - .unwrap(); + let lease = acquire_distribution_headroom(&controller, &target).unwrap(); assert_eq!( controller.effective_capacity("k3", 512), 0, @@ -2396,7 +2652,7 @@ mod tests { } #[test] - fn distribution_headroom_capacity_floor_requires_complete_strict_telemetry() { + fn distribution_headroom_does_not_reopen_capacity_with_incomplete_telemetry() { const HOT: &str = "grpc://incomplete-hot:30000"; const MISSING: &str = "grpc://incomplete-missing:30000"; let registry = Arc::new(WorkerRegistry::new()); @@ -2418,7 +2674,7 @@ mod tests { } #[test] - fn distribution_headroom_capacity_floor_is_default_off() { + fn distribution_headroom_does_not_reopen_capacity_when_disabled() { const HOT: &str = "grpc://off-hot:30000"; const IDLE: &str = "grpc://off-idle:30000"; let registry = Arc::new(WorkerRegistry::new()); @@ -2451,6 +2707,8 @@ mod tests { .distribution_headroom_snapshot("k3", "model") .is_empty()); + let mut inconsistent_total = worker_load(&[(0, 1, 1, 0.1, 0.1, 16)]); + inconsistent_total.loads[0].num_total_reqs = 1; let malformed = [ WorkerLoadResponse { dp_rank_count: 2, @@ -2461,6 +2719,7 @@ mod tests { worker_load(&[(0, -1, 0, 0.1, 0.1, 16)]), worker_load(&[(0, 0, 0, 0.1, 0.1, 0)]), worker_load(&[(0, 0, 0, f64::NAN, 0.1, 16)]), + inconsistent_total, ]; for load in malformed { update_worker_loads(&controller, [(URL, load)]); @@ -2483,6 +2742,25 @@ mod tests { assert!(controller .distribution_headroom_snapshot("k3", "model") .is_empty()); + + let mut unchanged = worker_load(&[(0, 0, 0, 0.1, 0.1, 16)]); + unchanged.timestamp = "fixed-backend-sample".to_string(); + let unchanged_map = HashMap::from([(URL.to_string(), unchanged)]); + controller.update_loads(&unchanged_map); + controller + .work + .lock() + .distribution_telemetry + .get_mut(URL) + .unwrap() + .observed_at = Instant::now().checked_sub(Duration::from_secs(31)).unwrap(); + controller.update_loads(&unchanged_map); + assert!( + controller + .distribution_headroom_snapshot("k3", "model") + .is_empty(), + "an unrelated watch update must not refresh an unchanged backend sample" + ); } #[test] @@ -2532,7 +2810,7 @@ mod tests { register_headroom_worker(®istry, A); register_headroom_worker(®istry, B); register_headroom_worker(®istry, C); - let controller = AdaptiveAdmissionController::new(distribution_config(2), registry); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); update_worker_loads( &controller, [ @@ -2555,14 +2833,13 @@ mod tests { let barrier = Arc::clone(&barrier); std::thread::spawn(move || { barrier.wait(); - controller.try_acquire_distribution_headroom(&target) + acquire_distribution_headroom(&controller, &target) }) }) .collect(); let claims: Vec<_> = claim_handles .into_iter() - .map(|handle| handle.join().unwrap()) - .flatten() + .filter_map(|handle| handle.join().unwrap()) .collect(); assert_eq!(claims.len(), 1); drop(claims); @@ -2576,19 +2853,18 @@ mod tests { let barrier = Arc::clone(&barrier); std::thread::spawn(move || { barrier.wait(); - controller.try_acquire_distribution_headroom(&target) + acquire_distribution_headroom(&controller, &target) }) }) .collect(); let leases: Vec<_> = lease_handles .into_iter() - .map(|handle| handle.join().unwrap()) - .flatten() + .filter_map(|handle| handle.join().unwrap()) .collect(); - assert_eq!(leases.len(), 2); + assert_eq!(leases.len(), 1); assert_eq!( AdaptiveAdmissionController::active_partition_leases(&controller.work.lock(), "k3"), - 2 + 1 ); } @@ -2604,28 +2880,106 @@ mod tests { .distribution_headroom_snapshot("k3", "model") .pop() .unwrap(); - let lease = controller - .try_acquire_distribution_headroom(&target) - .unwrap(); + let lease = acquire_distribution_headroom(&controller, &target).unwrap(); assert!(lease.verify()); + assert_eq!( + controller.active_distribution_targets("k3", "model"), + vec![(Arc::from(URL), target.worker_revision())] + ); + assert!(controller.distribution_target_is_active( + "k3", + "model", + URL, + target.worker_revision() + )); let snapshot = controller.distribution_headroom_snapshot("k3", "model"); assert_eq!(snapshot.len(), 1); assert_eq!(snapshot[0].issuable_slots(), 0); drop(lease); + assert!(controller + .active_distribution_targets("k3", "model") + .is_empty()); + assert!(!controller.distribution_target_is_active( + "k3", + "model", + URL, + target.worker_revision() + )); + let target = controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap(); + assert!(acquire_distribution_headroom(&controller, &target).is_some()); + } + + #[test] + fn ordinary_predispatch_claim_hands_off_to_worker_load_without_slot_gap() { + const URL: &str = "grpc://predispatch:30000"; + let registry = Arc::new(WorkerRegistry::new()); + let worker = register_headroom_worker(®istry, URL); + let controller = AdaptiveAdmissionController::new(distribution_config(1), registry); + update_worker_loads(&controller, [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 1)]))]); + let target = controller .distribution_headroom_snapshot("k3", "model") .pop() .unwrap(); + let selection = controller.lock_distribution_selection(); + let ordinary = controller + .try_reserve_ordinary_distribution_target(&selection, "k3", "model", &worker) + .unwrap(); + drop(selection); + assert_eq!( + controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap() + .issuable_slots(), + 0 + ); + let seed_selection = controller.lock_distribution_selection(); assert!(controller - .try_acquire_distribution_headroom(&target) - .is_some()); + .try_acquire_distribution_headroom(&seed_selection, &target) + .is_none()); + drop(seed_selection); + + let load_guard = WorkerLoadGuard::new(Arc::clone(&worker), None); + drop(ordinary); + assert_eq!( + controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .unwrap() + .issuable_slots(), + 0, + "live worker load must cover the exact handoff after claim release" + ); + drop(load_guard); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .pop() + .is_some_and(|candidate| candidate.issuable_slots() == 1)); + } + + #[test] + fn distribution_headroom_runtime_rejects_seed_caps_above_one() { + const URL: &str = "grpc://unsupported-cap:30000"; + let registry = Arc::new(WorkerRegistry::new()); + register_headroom_worker(®istry, URL); + let controller = AdaptiveAdmissionController::new(distribution_config(2), registry); + update_worker_loads(&controller, [(URL, worker_load(&[(0, 0, 0, 0.1, 0.1, 8)]))]); + + assert!(!controller.distribution_headroom_enabled("k3")); + assert!(controller + .distribution_headroom_snapshot("k3", "model") + .is_empty()); } #[test] fn distribution_headroom_revision_health_and_model_changes_fail_closed() { const URL: &str = "grpc://revision:30000"; let registry = Arc::new(WorkerRegistry::new()); - register_headroom_worker(®istry, URL); + let original = register_headroom_worker(®istry, URL); let worker_id = registry.get_id_by_url(URL).unwrap(); let controller = AdaptiveAdmissionController::new(distribution_config(1), Arc::clone(®istry)); @@ -2639,6 +2993,7 @@ mod tests { .distribution_headroom_snapshot("k3", "model") .pop() .unwrap(); + assert!(stale_target.matches_worker(&original)); let replacement: Arc = Arc::new( BasicWorkerBuilder::new(URL) .model(ModelCard::new("model").with_alias("model-alias")) @@ -2649,19 +3004,16 @@ mod tests { }) .build(), ); + assert!(!stale_target.matches_worker(&replacement)); assert!(registry.replace(&worker_id, replacement)); - assert!(controller - .try_acquire_distribution_headroom(&stale_target) - .is_none()); + assert!(acquire_distribution_headroom(&controller, &stale_target).is_none()); update_worker_loads(&controller, [(URL, load.clone())]); let target = controller .distribution_headroom_snapshot("k3", "model") .pop() .unwrap(); - let lease = controller - .try_acquire_distribution_headroom(&target) - .unwrap(); + let lease = acquire_distribution_headroom(&controller, &target).unwrap(); assert!(lease.verify()); let second_replacement: Arc = Arc::new( BasicWorkerBuilder::new(URL) @@ -2685,9 +3037,7 @@ mod tests { .distribution_headroom_snapshot("k3", "model") .pop() .unwrap(); - let lease = controller - .try_acquire_distribution_headroom(&target) - .unwrap(); + let lease = acquire_distribution_headroom(&controller, &target).unwrap(); assert!(lease.verify()); update_worker_loads(&controller, [(URL, load)]); assert!( @@ -2708,9 +3058,7 @@ mod tests { .get_by_url(URL) .unwrap() .set_status(WorkerStatus::NotReady); - assert!(controller - .try_acquire_distribution_headroom(&target) - .is_none()); + assert!(acquire_distribution_headroom(&controller, &target).is_none()); assert!(!registry.get_by_url(URL).unwrap().is_available()); } diff --git a/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs index 7680b8606..1cf9328d4 100644 --- a/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs +++ b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs @@ -168,7 +168,7 @@ fn is_single_distribution_seed_sample( shape: GenerationShape, backend_request_count: usize, ) -> bool { - backend_request_count == 1 && shape.flags & FLAG_MULTIPLE_COMPLETIONS == 0 + backend_request_count == 1 && shape.flags & (FLAG_MULTIPLE_COMPLETIONS | FLAG_STREAMING) == 0 } fn multiplied_limit(per_completion: Option, multiplicity: u32) -> Option { @@ -276,19 +276,129 @@ impl PipelineStage for AdaptiveAdmissionStage { mod tests { use std::sync::Arc; + use openai_protocol::{ + chat::ChatCompletionRequest, completion::CompletionRequest, generate::GenerateRequest, + messages::CreateMessageRequest, responses::ResponsesRequest, + }; + use super::*; - #[test] - fn multiple_generate_samples_cannot_enter_distribution_seed_path() { - let request = serde_json::from_value(serde_json::json!({ + fn chat_shape(n: u32) -> GenerationShape { + GenerationShape::for_request(&RequestType::Chat(Arc::new(ChatCompletionRequest { + n: Some(n), + ..Default::default() + }))) + .unwrap() + } + + fn generate_shape(n: u32) -> GenerationShape { + let request: GenerateRequest = serde_json::from_value(serde_json::json!({ "text": "hello", - "sampling_params": { "n": 2 } + "sampling_params": { "n": n } })) .unwrap(); - let shape = - GenerationShape::for_request(&RequestType::Generate(Arc::new(request))).unwrap(); - assert_ne!(shape.flags & FLAG_MULTIPLE_COMPLETIONS, 0); - assert!(!is_single_distribution_seed_sample(shape, 1)); + GenerationShape::for_request(&RequestType::Generate(Arc::new(request))).unwrap() + } + + fn completion_shape(n: Option, best_of: Option) -> GenerationShape { + let mut value = serde_json::json!({ + "model": "unknown", + "prompt": "hello" + }); + if let Some(n) = n { + value["n"] = serde_json::json!(n); + } + if let Some(best_of) = best_of { + value["best_of"] = serde_json::json!(best_of); + } + let request: CompletionRequest = serde_json::from_value(value).unwrap(); + GenerationShape::for_request(&RequestType::Completion(Arc::new(request))).unwrap() + } + + #[test] + fn only_scalar_generation_can_enter_distribution_seed_path() { + let mut streaming_chat = chat_shape(1); + streaming_chat.flags |= FLAG_STREAMING; + let cases = [ + ("scalar chat", chat_shape(1), 1, true), + ("streaming chat", streaming_chat, 1, false), + ("chat n", chat_shape(2), 1, false), + ("generate n", generate_shape(2), 1, false), + ("completion n", completion_shape(Some(2), None), 1, false), + ( + "completion best_of", + completion_shape(Some(1), Some(2)), + 1, + false, + ), + ( + "multi-prompt completion fanout", + completion_shape(Some(1), None), + 2, + false, + ), + ]; + + for (name, shape, backend_request_count, expected) in cases { + assert_eq!( + is_single_distribution_seed_sample(shape, backend_request_count), + expected, + "{name}" + ); + } + } + + #[test] + fn every_streaming_generation_endpoint_is_excluded_from_distribution_seeding() { + let requests = [ + RequestType::Chat(Arc::new(ChatCompletionRequest { + stream: true, + ..Default::default() + })), + RequestType::Generate(Arc::new( + serde_json::from_value::(serde_json::json!({ + "model": "unknown", + "text": "hello", + "stream": true, + })) + .unwrap(), + )), + RequestType::Completion(Arc::new( + serde_json::from_value::(serde_json::json!({ + "model": "unknown", + "prompt": "hello", + "stream": true, + })) + .unwrap(), + )), + RequestType::Messages(Arc::new( + serde_json::from_value::(serde_json::json!({ + "model": "unknown", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 16, + "stream": true, + })) + .unwrap(), + )), + RequestType::Responses(Arc::new( + serde_json::from_value::(serde_json::json!({ + "model": "unknown", + "input": "hello", + "stream": true, + })) + .unwrap(), + )), + ]; + + for request in requests { + let shape = GenerationShape::for_request(&request).unwrap(); + assert_ne!(shape.flags & FLAG_STREAMING, 0, "{}", shape.endpoint); + assert!( + !is_single_distribution_seed_sample(shape, 1), + "{} streaming request reached the seed path", + shape.endpoint + ); + } } #[test] diff --git a/model_gateway/src/routers/grpc/common/stages/distribution_seed_commit.rs b/model_gateway/src/routers/grpc/common/stages/distribution_seed_commit.rs new file mode 100644 index 000000000..a325552bb --- /dev/null +++ b/model_gateway/src/routers/grpc/common/stages/distribution_seed_commit.rs @@ -0,0 +1,290 @@ +//! Terminal-success commit for cache-aware distribution seeds. + +use async_trait::async_trait; +use axum::response::Response; + +use super::{adaptive_admission::rejection_response, PipelineStage}; +use crate::routers::{ + error, + grpc::context::{FinalResponse, RequestContext, RequestType}, +}; + +pub(crate) struct DistributionSeedCommitStage; + +fn has_matching_terminal_response( + request_type: &RequestType, + final_response: Option<&FinalResponse>, + has_responses_iteration_result: bool, +) -> bool { + match request_type { + RequestType::Chat(_) => matches!(final_response, Some(FinalResponse::Chat(_))), + RequestType::Generate(_) => matches!(final_response, Some(FinalResponse::Generate(_))), + RequestType::Completion(_) => { + matches!(final_response, Some(FinalResponse::Completion(_))) + } + RequestType::Messages(_) => matches!(final_response, Some(FinalResponse::Messages(_))), + RequestType::Responses(_) => has_responses_iteration_result, + RequestType::Embedding(_) | RequestType::Classify(_) => false, + } +} + +#[async_trait] +impl PipelineStage for DistributionSeedCommitStage { + async fn execute(&self, ctx: &mut RequestContext) -> Result, Response> { + let has_seed = ctx + .state + .load_guards + .as_ref() + .is_some_and(|guards| guards.has_distribution_seed()); + if !has_seed { + return Ok(None); + } + if ctx.is_streaming() { + return Err(rejection_response(1)); + } + if !has_matching_terminal_response( + &ctx.input.request_type, + ctx.state.response.final_response.as_ref(), + ctx.state.response.responses_iteration_result.is_some(), + ) { + return Err(error::internal_error( + "distribution_seed_terminal_response_missing", + "Distribution seed did not produce the expected terminal response", + )); + } + if let Some(guards) = ctx.state.load_guards.as_mut() { + guards.commit_distribution_seed_success(); + } + Ok(None) + } + + fn name(&self) -> &'static str { + "DistributionSeedCommit" + } +} + +#[cfg(test)] +mod tests { + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; + + use axum::http::StatusCode; + use llm_tokenizer::TokenizerRegistry; + use openai_protocol::{ + chat::{ChatCompletionRequest, ChatCompletionResponse}, + completion::{CompletionRequest, CompletionResponse}, + generate::GenerateRequest, + messages::{CreateMessageRequest, Message}, + responses::ResponsesRequest, + }; + use reasoning_parser::ParserFactory as ReasoningParserFactory; + use serde_json::json; + use tool_parser::ParserFactory as ToolParserFactory; + + use super::*; + use crate::{ + routers::grpc::context::{LoadGuards, SharedComponents}, + worker::WorkerRegistry, + }; + + fn components() -> Arc { + Arc::new(SharedComponents { + tokenizer_registry: Arc::new(TokenizerRegistry::new()), + worker_registry: Arc::new(WorkerRegistry::new()), + tool_parser_factory: ToolParserFactory::default(), + reasoning_parser_factory: ReasoningParserFactory::default(), + configured_tool_parser: None, + configured_reasoning_parser: None, + multimodal: None, + adaptive_admission: None, + }) + } + + fn generate_context(stream: bool) -> RequestContext { + let request: GenerateRequest = serde_json::from_value(json!({ + "model": "kimi-k3", + "text": "hello", + "stream": stream, + })) + .expect("generate request"); + RequestContext::for_generate(Arc::new(request), None, "kimi-k3".to_string(), components()) + } + + fn test_seed(ctx: &mut RequestContext) -> Arc { + let committed = Arc::new(AtomicBool::new(false)); + ctx.state.load_guards = Some(LoadGuards::test_distribution_seed(Arc::clone(&committed))); + committed + } + + #[test] + fn terminal_response_match_is_endpoint_exact() { + let chat = RequestType::Chat(Arc::new( + serde_json::from_value::(json!({ + "model": "kimi-k3", + "messages": [{"role": "user", "content": "hello"}], + })) + .expect("chat request"), + )); + let generate = RequestType::Generate(Arc::new( + serde_json::from_value::(json!({ + "model": "kimi-k3", + "text": "hello", + })) + .expect("generate request"), + )); + let completion = RequestType::Completion(Arc::new( + serde_json::from_value::(json!({ + "model": "kimi-k3", + "prompt": "hello", + })) + .expect("completion request"), + )); + let messages = RequestType::Messages(Arc::new( + serde_json::from_value::(json!({ + "model": "kimi-k3", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 16, + })) + .expect("messages request"), + )); + let responses = RequestType::Responses(Arc::new( + serde_json::from_value::(json!({ + "model": "kimi-k3", + "input": "hello", + })) + .expect("responses request"), + )); + + let chat_final = FinalResponse::Chat( + serde_json::from_value::(json!({ + "id": "chat-1", + "object": "chat.completion", + "created": 1, + "model": "kimi-k3", + "choices": [], + })) + .expect("chat response"), + ); + let generate_final = FinalResponse::Generate(Vec::new()); + let completion_final = FinalResponse::Completion( + serde_json::from_value::(json!({ + "id": "completion-1", + "object": "text_completion", + "created": 1, + "model": "kimi-k3", + "choices": [], + })) + .expect("completion response"), + ); + let messages_final = FinalResponse::Messages( + serde_json::from_value::(json!({ + "id": "message-1", + "type": "message", + "role": "assistant", + "content": [], + "model": "kimi-k3", + "stop_reason": "end_turn", + "stop_sequence": null, + "usage": {"input_tokens": 1, "output_tokens": 1}, + })) + .expect("messages response"), + ); + + assert!(has_matching_terminal_response( + &chat, + Some(&chat_final), + false + )); + assert!(has_matching_terminal_response( + &generate, + Some(&generate_final), + false + )); + assert!(has_matching_terminal_response( + &completion, + Some(&completion_final), + false + )); + assert!(has_matching_terminal_response( + &messages, + Some(&messages_final), + false + )); + assert!(has_matching_terminal_response(&responses, None, true)); + + assert!(!has_matching_terminal_response( + &chat, + Some(&generate_final), + false + )); + assert!(!has_matching_terminal_response(&generate, None, false)); + assert!(!has_matching_terminal_response(&responses, None, false)); + } + + #[tokio::test] + async fn no_seed_is_a_noop_without_terminal_state() { + let mut ctx = generate_context(false); + let result = DistributionSeedCommitStage.execute(&mut ctx).await; + assert!(matches!(result, Ok(None))); + } + + #[tokio::test] + async fn matching_non_streaming_terminal_response_commits_seed() { + let mut ctx = generate_context(false); + let committed = test_seed(&mut ctx); + ctx.state.response.final_response = Some(FinalResponse::Generate(Vec::new())); + + let result = DistributionSeedCommitStage.execute(&mut ctx).await; + + assert!(matches!(result, Ok(None))); + assert!(committed.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn missing_or_mismatched_terminal_state_does_not_commit() { + let mut missing = generate_context(false); + let missing_committed = test_seed(&mut missing); + let response = DistributionSeedCommitStage + .execute(&mut missing) + .await + .expect_err("missing terminal response must fail closed"); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + assert!(!missing_committed.load(Ordering::Acquire)); + + let mut mismatched = generate_context(false); + let mismatched_committed = test_seed(&mut mismatched); + mismatched.state.response.final_response = Some(FinalResponse::Completion( + serde_json::from_value::(json!({ + "id": "completion-1", + "object": "text_completion", + "created": 1, + "model": "kimi-k3", + "choices": [], + })) + .expect("completion response"), + )); + let response = DistributionSeedCommitStage + .execute(&mut mismatched) + .await + .expect_err("mismatched terminal response must fail closed"); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + assert!(!mismatched_committed.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn streaming_seed_is_rejected_without_commit() { + let mut ctx = generate_context(true); + let committed = test_seed(&mut ctx); + ctx.state.response.final_response = Some(FinalResponse::Generate(Vec::new())); + + let response = DistributionSeedCommitStage + .execute(&mut ctx) + .await + .expect_err("streaming seeds are unsupported in the first release"); + + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert!(!committed.load(Ordering::Acquire)); + } +} diff --git a/model_gateway/src/routers/grpc/common/stages/mod.rs b/model_gateway/src/routers/grpc/common/stages/mod.rs index d31ba8051..f5172c26f 100644 --- a/model_gateway/src/routers/grpc/common/stages/mod.rs +++ b/model_gateway/src/routers/grpc/common/stages/mod.rs @@ -42,6 +42,7 @@ pub trait PipelineStage: Send + Sync { mod adaptive_admission; mod client_acquisition; mod dispatch_metadata; +mod distribution_seed_commit; pub(crate) mod encode; pub(crate) mod helpers; mod request_execution; @@ -51,6 +52,7 @@ mod worker_selection; pub(crate) use adaptive_admission::AdaptiveAdmissionStage; pub(crate) use client_acquisition::ClientAcquisitionStage; pub(crate) use dispatch_metadata::DispatchMetadataStage; +pub(crate) use distribution_seed_commit::DistributionSeedCommitStage; pub(crate) use encode::EncodeStage; pub(crate) use request_execution::RequestExecutionStage; pub(crate) use worker_selection::{WorkerSelectionMode, WorkerSelectionStage}; diff --git a/model_gateway/src/routers/grpc/common/stages/request_execution.rs b/model_gateway/src/routers/grpc/common/stages/request_execution.rs index 71c59ee31..4b63831ea 100644 --- a/model_gateway/src/routers/grpc/common/stages/request_execution.rs +++ b/model_gateway/src/routers/grpc/common/stages/request_execution.rs @@ -185,6 +185,10 @@ impl PipelineStage for RequestExecutionStage { policy_reservation, distribution_seed, )); + // `LoadGuards::scaled` has now incremented the selected worker's live + // load. Releasing the exact predispatch claim after that handoff makes + // the ordinary-selection-to-seed-acquisition transition gap-free. + drop(ctx.state.ordinary_distribution_guard.take()); // Extract dispatch metadata for tracing span let dispatch = ctx.state.dispatch.as_ref(); diff --git a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs index 4da3223b2..5b82e3f41 100644 --- a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs +++ b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs @@ -11,14 +11,17 @@ use tracing::{error, warn}; use super::{adaptive_admission::rejection_response, PipelineStage}; use crate::{ + middleware::scheduler::SchedulerAdmissionProof, observability::metrics::{metrics_labels, Metrics}, policies::{ - LoadBalancingPolicy, PolicyRegistry, SeedWorkerHeadroom, SelectWorkerInfo, WorkerLeg, + CacheAwarePolicy, CacheAwareSelection, DistributionProtectionSnapshot, LoadBalancingPolicy, + PolicyRegistry, SeedWorkerHeadroom, SelectWorkerInfo, WorkerLeg, }, routers::{ common::header_utils::worker_url_is_allowed, error, grpc::{ + adaptive_admission::{AdaptiveAdmissionController, OrdinaryDistributionDispatchGuard}, context::{ DistributionSeedDispatchGuard, EncodeWorkerAssignment, PendingDistributionSeed, PolicyReservation, RequestContext, WorkerSelection, @@ -33,6 +36,11 @@ use crate::{ /// Result type for PD worker pair selection: (prefill, decode, runtime_type) type PdWorkerPair = (Arc, Arc, RuntimeType); +type SingleWorkerSelection = ( + Arc, + Option, + Option, +); /// Result type for EPD worker selection: (encode assignments, prefill, decode, runtime_type). type EncodePrefillDecodeWorkerSelection = ( @@ -99,6 +107,13 @@ impl PipelineStage for WorkerSelectionStage { let headers = ctx.input.headers.as_ref(); let model_id = ctx.input.model_id.as_str(); + let distribution_scope = ctx.state.adaptive_request.as_ref().and_then(|tracker| { + let partition = tracker.partition()?; + let controller = ctx.components.adaptive_admission.as_ref()?; + controller + .distribution_headroom_enabled(partition) + .then(|| (Arc::clone(controller), partition.to_string())) + }); let workers = if let Some(pending) = pending_distribution_seed { if self.mode != WorkerSelectionMode::Regular { return Err(rejection_response(pending.retry_after_secs)); @@ -113,9 +128,18 @@ impl PipelineStage for WorkerSelectionStage { } else { match self.mode { WorkerSelectionMode::Regular => { - match self.select_single_worker(model_id, text, tokens, headers) { - Some((worker, reservation)) => { + match self.select_single_worker( + model_id, + text, + tokens, + headers, + distribution_scope + .as_ref() + .map(|(controller, partition)| (controller, partition.as_str())), + )? { + Some((worker, reservation, ordinary_distribution_guard)) => { ctx.state.policy_reservation = reservation; + ctx.state.ordinary_distribution_guard = ordinary_distribution_guard; WorkerSelection::Single { worker } } None => { @@ -251,11 +275,6 @@ impl WorkerSelectionStage { headers: Option<&HeaderMap>, ) -> Option<(Arc, DistributionSeedDispatchGuard)> { let controller = ctx.components.adaptive_admission.as_ref()?; - let targets = controller.distribution_headroom_snapshot(&pending.partition, model_id); - if targets.is_empty() { - return None; - } - let workers: Vec<_> = self .worker_registry .get_workers_filtered( @@ -266,7 +285,10 @@ impl WorkerSelectionStage { false, ) .into_iter() - .filter(|worker| worker_is_available_for_request(worker.as_ref(), headers)) + // Header-excluded authoritative owners must remain visible to the + // cache-policy owner ceiling. SelectWorkerInfo applies the same + // headers only to actual target eligibility below. + .filter(|worker| worker.is_available()) .collect(); if workers.len() < 2 { return None; @@ -281,6 +303,25 @@ impl WorkerSelectionStage { reserve_work: false, leg: WorkerLeg::Single, }; + let meta = ctx.input.tenant_request_meta.as_ref()?; + let proof = meta.extension::()?; + let claimed = proof + .try_claim( + meta.tenant_key(), + meta.request_charge_id(), + model_id, + &pending.partition, + ) + .ok()?; + + // Claim scheduler authority before mutating policy cooldown or target + // capacity. The short controller lock then orders this seed selection + // against ordinary cache-aware selection without crossing an await. + let selection_guard = controller.lock_distribution_selection(); + let targets = controller.distribution_headroom_snapshot(&pending.partition, model_id); + if targets.is_empty() { + return None; + } let capacity: Vec<_> = targets .iter() .map(|target| SeedWorkerHeadroom { @@ -296,24 +337,15 @@ impl WorkerSelectionStage { target.worker_url() == plan.target_worker_url() && target.worker_revision() == plan.target_worker_revision() })?; - let lease = controller.try_acquire_distribution_headroom(target)?; - - let meta = ctx.input.tenant_request_meta.as_ref()?; - let proof = meta.extension::()?; - let claimed = proof - .try_claim( - meta.tenant_key(), - meta.request_charge_id(), - model_id, - &pending.partition, - ) - .ok()?; - let selected = workers.into_iter().find(|worker| { worker.url() == plan.target_worker_url() && worker.revision() == plan.target_worker_revision() && worker_is_available_for_request(worker.as_ref(), headers) })?; + if !target.matches_worker(&selected) { + return None; + } + let lease = controller.try_acquire_distribution_headroom(&selection_guard, target)?; if !lease.verify() { return None; } @@ -333,19 +365,24 @@ impl WorkerSelectionStage { DistributionSeedDispatchGuard { headroom: lease, _scheduler_proof: claimed, - _policy_plan: plan, + policy_plan: plan, retry_after_secs: pending.retry_after_secs, }, )) } + #[expect( + clippy::result_large_err, + reason = "the pipeline contract returns an Axum Response on local rejection" + )] fn select_single_worker( &self, model_id: &str, text: Option<&str>, tokens: Option<&[u32]>, headers: Option<&HeaderMap>, - ) -> Option<(Arc, Option)> { + distribution_scope: Option<(&Arc, &str)>, + ) -> Result, Response> { // Treat "unknown" model as wildcard (match any worker) let model_filter = if model_id == UNKNOWN_MODEL_ID { None @@ -362,41 +399,180 @@ impl WorkerSelectionStage { false, // get all workers, we'll filter by is_available() next ); - // Use into_iter() to take ownership of Arcs without cloning (avoids atomic inc/dec) - let available: Vec> = workers - .into_iter() - .filter(|w| worker_is_available_for_request(w.as_ref(), headers)) - .collect(); - - if available.is_empty() { - return None; - } - // Get the appropriate policy for this model let policy = self.policy_registry.get_policy_or_default(model_id); // Get cached hash ring for consistent hashing (O(log n) lookup) let hash_ring = self.worker_registry.get_hash_ring(model_id); + // Distribution-enabled cache-aware routing must see header-excluded + // authoritative owners so its per-prefix claim can serialize any + // replacement. The policy itself enforces header eligibility. Active + // seed targets remain visible for ownership matching but are appended + // to the policy-only exclusion header until their request-lifetime + // headroom lease ends. This covers early/deeper KV store events and the + // routing-key override path without hiding the provisional owner. + // Other models retain the legacy prefilter. + let distribution_cache_aware = distribution_scope.and_then(|(controller, partition)| { + policy + .as_any() + .downcast_ref::() + .map(|cache_aware| (controller, partition, cache_aware)) + }); + let active_targets = + distribution_cache_aware.map_or_else(Vec::new, |(controller, partition, _)| { + controller.active_distribution_targets(partition, model_id) + }); + let distribution_protections = distribution_cache_aware.map_or_else( + DistributionProtectionSnapshot::default, + |(_, _, cache_aware)| cache_aware.distribution_protection_snapshot(model_id, tokens), + ); + let serialize_distribution = + !active_targets.is_empty() || !distribution_protections.is_empty(); + let policy_headers = if active_targets.is_empty() { + None + } else { + let mut merged = headers.cloned().unwrap_or_default(); + let targeted_worker = merged + .get("x-smg-target-worker-url") + .and_then(|value| value.to_str().ok()); + if targeted_worker.is_some_and(|targeted| { + active_targets + .iter() + .any(|(url, _)| url.as_ref() == targeted) + }) { + return Err(rejection_response(1)); + } + let mut excluded = merged + .get("x-smg-excluded-worker-urls") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .to_string(); + for (url, _) in &active_targets { + if !excluded.is_empty() { + excluded.push(','); + } + excluded.push_str(url); + } + let Ok(excluded) = HeaderValue::from_str(&excluded) else { + return Err(rejection_response(1)); + }; + merged.insert("x-smg-excluded-worker-urls", excluded); + Some(merged) + }; + let effective_headers = policy_headers.as_ref().or(headers); let info = SelectWorkerInfo { request_text: text, tokens, - headers, + headers: effective_headers, hash_ring, max_output_tokens: None, reserve_work: true, leg: WorkerLeg::Single, }; + let available: Vec> = workers + .into_iter() + .filter(|worker| { + if serialize_distribution { + worker.is_available() + } else { + worker_is_available_for_request(worker.as_ref(), headers) + } + }) + .collect(); + + if available.is_empty() { + return if serialize_distribution { + Err(rejection_response(1)) + } else { + Ok(None) + }; + } // Select and reserve prompt/output work atomically. The returned guard // releases the reservation when the request finishes or any later // gRPC pipeline stage fails. - let (idx, reservation_cost) = self - .policy_registry - .select_worker_with_reservation(&policy, &available, &info)?; - let selected = available[idx].clone(); + let selection = if serialize_distribution { + let Some((_, _, cache_aware)) = distribution_cache_aware else { + return Err(rejection_response(1)); + }; + match cache_aware.select_worker_decision_with_distribution_state( + &available, + &info, + &active_targets, + &distribution_protections, + ) { + CacheAwareSelection::Selected(index) => { + Some((index, policy.reservation_cost(&info))) + } + CacheAwareSelection::Blocked => return Err(rejection_response(1)), + CacheAwareSelection::Unavailable => return Err(rejection_response(1)), + } + } else { + self.policy_registry + .select_worker_with_reservation(&policy, &available, &info) + }; + let Some((idx, reservation_cost)) = selection else { + return Ok(None); + }; + let Some(selected) = available.get(idx).cloned() else { + return Ok(None); + }; let reservation = reservation_cost .map(|cost| PolicyReservation::new(policy.clone(), selected.url().to_string(), cost)); + if !worker_url_is_allowed(headers, selected.url()) { + return Err(rejection_response(1)); + } + // Cache hashing and ordinary policy selection stay outside the global + // controller lock. Only this final atomic snapshot check plus exact + // predispatch claim is serialized against the rare seed path. + let distribution_selection_guard = distribution_cache_aware + .map(|(controller, _, _)| controller.lock_distribution_selection()); + if let Some((controller, partition)) = distribution_scope { + let protection_snapshot_unchanged = + distribution_cache_aware.is_some_and(|(_, _, cache_aware)| { + cache_aware.distribution_protection_snapshot(model_id, tokens) + == distribution_protections + }); + let current_targets = controller.active_distribution_targets(partition, model_id); + let snapshot_unchanged = current_targets.len() == active_targets.len() + && current_targets.iter().all(|(url, revision)| { + active_targets + .iter() + .any(|(snapshot_url, snapshot_revision)| { + snapshot_url == url && snapshot_revision == revision + }) + }); + if !protection_snapshot_unchanged || !snapshot_unchanged { + return Err(rejection_response(1)); + } + if controller.distribution_target_is_active( + partition, + model_id, + selected.url(), + selected.revision(), + ) { + return Err(rejection_response(1)); + } + } + let ordinary_distribution_guard = + if let (Some((controller, partition, _)), Some(selection_guard)) = ( + distribution_cache_aware, + distribution_selection_guard.as_ref(), + ) { + Some( + controller + .try_reserve_ordinary_distribution_target( + selection_guard, + partition, + model_id, + &selected, + ) + .ok_or_else(|| rejection_response(1))?, + ) + } else { + None + }; // Record worker selection metric Metrics::record_worker_selection( @@ -406,7 +582,7 @@ impl WorkerSelectionStage { policy.name(), ); - Some((selected, reservation)) + Ok(Some((selected, reservation, ordinary_distribution_guard))) } fn select_pd_pair( @@ -795,10 +971,27 @@ fn hex_encode(bytes: &[u8]) -> String { #[cfg(test)] mod tests { - use openai_protocol::worker::HealthCheckConfig; + use std::collections::HashMap; + + use axum::http::{header::RETRY_AFTER, StatusCode}; + use kv_index::{ + compute_content_hash, PositionalIndexer, SequenceHash, StoredBlock, WorkerBlockMap, + }; + use openai_protocol::{ + model_card::ModelCard, + worker::{HealthCheckConfig, SchedulerLoadSnapshot, WorkerLoadResponse}, + }; + use tokio::sync::watch; use super::*; - use crate::worker::BasicWorkerBuilder; + use crate::{ + config::{ + AdaptiveAdmissionConfig, AdaptiveAdmissionMode, AdaptiveAdmissionStrategy, + ManualAssignmentMode, PolicyConfig, RoutingKeyOverrideConfig, + }, + middleware::scheduler::LocalAdaptiveRejection, + worker::{BasicWorkerBuilder, KvEventMonitor}, + }; fn ready_worker(url: &str) -> Arc { Arc::new( @@ -812,6 +1005,62 @@ mod tests { ) } + fn stage_worker(url: &str) -> Arc { + Arc::new( + BasicWorkerBuilder::new(url) + .model(ModelCard::new("kimi-k3")) + .label("admission_partition", "k3") + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ) + } + + fn clean_headroom_load(timestamp: &str) -> WorkerLoadResponse { + WorkerLoadResponse { + timestamp: timestamp.to_string(), + dp_rank_count: 1, + loads: vec![SchedulerLoadSnapshot { + dp_rank: 0, + num_running_reqs: 0, + num_waiting_reqs: 0, + num_total_reqs: 0, + token_usage: 0.1, + utilization: 0.1, + max_running_requests: 38, + ..Default::default() + }], + } + } + + fn assert_local_adaptive_rejection(result: Result, Response>) { + let response = match result { + Err(response) => response, + Ok(_) => panic!( + "active seed B must block locally rather than route to sticky B or uncached C" + ), + }; + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert!( + response + .extensions() + .get::() + .is_some(), + "the rejection must use the scheduler refund marker" + ); + assert_eq!( + response + .headers() + .get(RETRY_AFTER) + .and_then(|value| value.to_str().ok()), + Some("1") + ); + } + #[test] fn trusted_worker_url_constraints_filter_before_policy_selection() { let public = ready_worker("grpc://public:30000"); @@ -842,4 +1091,169 @@ mod tests { Some(&private_headers) )); } + + #[tokio::test] + async fn active_distribution_seed_blocks_sticky_deeper_owner_after_header_filtering() { + let worker_registry = Arc::new(WorkerRegistry::new()); + let worker_a = stage_worker("grpc://worker-a:30000"); + let worker_b = stage_worker("grpc://worker-b:30000"); + let worker_c = stage_worker("grpc://worker-c:30000"); + for worker in [&worker_a, &worker_b, &worker_c] { + worker_registry + .register(Arc::clone(worker)) + .expect("unique worker registration"); + } + + let policy_registry = Arc::new(PolicyRegistry::with_override( + PolicyConfig::CacheAware { + cache_threshold: 0.0, + balance_abs_threshold: 32, + balance_rel_threshold: 1.1, + eviction_interval_secs: 0, + max_tree_size: 10_000, + fallback_output_token_estimate: 4096, + block_size: 4, + engine_load: false, + balance_token_usage_threshold: 1.0, + overload_token_usage_threshold: 1.0, + max_cached_owners_per_prefix: 8, + cache_owner_spill_cooldown_secs: 5, + }, + RoutingKeyOverrideConfig { + enabled: true, + assignment_mode: ManualAssignmentMode::MinLoad, + ..Default::default() + }, + )); + + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + monitor.set_block_size("kimi-k3", 4); + let indexer = Arc::new(PositionalIndexer::new(4)); + let owner_a = indexer.intern_worker(worker_a.url()).unwrap(); + let seed_b = indexer.intern_worker(worker_b.url()).unwrap(); + let prefix_tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let owner_blocks: Vec<_> = prefix_tokens + .chunks(4) + .enumerate() + .map(|(index, tokens)| StoredBlock { + seq_hash: SequenceHash(index as u64 + 1), + content_hash: compute_content_hash(tokens), + }) + .collect(); + indexer + .apply_stored(owner_a, &owner_blocks, None, &mut WorkerBlockMap::default()) + .unwrap(); + monitor + .indexers + .insert("kimi-k3".to_string(), Arc::clone(&indexer)); + policy_registry.set_kv_event_monitor(Some(monitor)); + + let policy = policy_registry.get_policy_or_default("kimi-k3"); + let mut sticky_headers = HeaderMap::new(); + sticky_headers.insert("x-smg-routing-key", "trajectory-1".parse().unwrap()); + let sticky_info = SelectWorkerInfo { + tokens: Some(&prefix_tokens), + headers: Some(&sticky_headers), + ..Default::default() + }; + assert_eq!( + policy_registry.select_worker(&policy, &[Arc::clone(&worker_b)], &sticky_info), + Some(0), + "the routing key should first become sticky to worker B" + ); + assert_eq!( + policy_registry.select_worker( + &policy, + &[ + Arc::clone(&worker_a), + Arc::clone(&worker_b), + Arc::clone(&worker_c), + ], + &sticky_info, + ), + Some(1), + "the routing-key override should still resolve to worker B" + ); + + let controller = AdaptiveAdmissionController::new( + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Enforce, + strategy: AdaptiveAdmissionStrategy::EngineFeedback, + distribution_headroom_partitions: vec!["k3".to_string()], + distribution_headroom_partition_seed_cap: 1, + ..Default::default() + }, + Arc::clone(&worker_registry), + ); + let loads = HashMap::from([ + ( + worker_a.url().to_string(), + clean_headroom_load("stage-worker-a"), + ), + ( + worker_b.url().to_string(), + clean_headroom_load("stage-worker-b"), + ), + ( + worker_c.url().to_string(), + clean_headroom_load("stage-worker-c"), + ), + ]); + let (_load_tx, load_rx) = watch::channel(loads); + controller.start_load_updates(load_rx); + let target_b = controller + .distribution_headroom_snapshot("k3", "kimi-k3") + .into_iter() + .find(|target| target.worker_url() == worker_b.url()) + .expect("worker B should expose clean distribution headroom"); + let seed_lease = { + let selection = controller.lock_distribution_selection(); + controller + .try_acquire_distribution_headroom(&selection, &target_b) + .expect("worker B seed lease") + }; + assert!(seed_lease.verify()); + + let request_tokens = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]; + let mut request_headers = sticky_headers; + request_headers.insert( + "x-smg-excluded-worker-urls", + worker_a.url().parse().unwrap(), + ); + let stage = WorkerSelectionStage::new( + worker_registry, + policy_registry, + WorkerSelectionMode::Regular, + ); + assert_local_adaptive_rejection(stage.select_single_worker( + "kimi-k3", + None, + Some(&request_tokens), + Some(&request_headers), + Some((&controller, "k3")), + )); + + // The seed's store event can make B a deeper authoritative owner while + // its request is still in flight. The stage must continue to reject, + // even though the routing key is sticky to B and A is header-excluded. + let seed_blocks: Vec<_> = request_tokens + .chunks(4) + .enumerate() + .map(|(index, tokens)| StoredBlock { + seq_hash: SequenceHash(index as u64 + 1), + content_hash: compute_content_hash(tokens), + }) + .collect(); + indexer + .apply_stored(seed_b, &seed_blocks, None, &mut WorkerBlockMap::default()) + .unwrap(); + + assert_local_adaptive_rejection(stage.select_single_worker( + "kimi-k3", + None, + Some(&request_tokens), + Some(&request_headers), + Some((&controller, "k3")), + )); + } } diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index c77060c1c..86e6311a7 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -4,6 +4,8 @@ //! eliminating deep parameter passing chains and providing a single source of truth //! for request state. +#[cfg(test)] +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use axum::http::HeaderMap; @@ -24,6 +26,7 @@ use tracing::debug; use super::{ adaptive_admission::{ AdaptiveAdmissionController, AdaptiveRequestTracker, DistributionHeadroomLease, + OrdinaryDistributionDispatchGuard, }, client::GrpcClient, common::stages::encode::EncodeDispatchPlan, @@ -188,6 +191,11 @@ pub(crate) struct ProcessingState { /// backend execution; every earlier failure releases both by `Drop`. pub distribution_seed_guard: Option, + /// Exact ordinary target claim held until request execution installs the + /// worker load guard. This closes the worker-selection-to-dispatch window + /// in which a clean-peer seed might otherwise reserve the same target. + pub ordinary_distribution_guard: Option, + // Stage 2: Worker selection outputs pub workers: Option, @@ -220,12 +228,21 @@ pub(crate) struct PendingDistributionSeed { #[derive(Debug)] pub(crate) struct DistributionSeedDispatchGuard { - pub(crate) headroom: DistributionHeadroomLease, + pub(crate) policy_plan: OwnerPressureDispatchPlan, pub(crate) _scheduler_proof: ClaimedSchedulerAdmissionProof, - pub(crate) _policy_plan: OwnerPressureDispatchPlan, + // Keep this field last. Rust drops fields in declaration order, so the + // prefix plan must enter cooldown before model-wide target suppression is + // released by the headroom lease. + pub(crate) headroom: DistributionHeadroomLease, pub(crate) retry_after_secs: u32, } +impl DistributionSeedDispatchGuard { + fn commit_success(&mut self) { + self.policy_plan.commit_success(); + } +} + /// Per-item bootstrap rendezvous info for prefill, plus the dispatch plan that /// fans out to encode workers. /// @@ -504,7 +521,7 @@ pub(crate) enum LoadGuards { Single { _guard: WorkerLoadGuard, _policy_reservation: Option, - _distribution_seed: Option, + distribution_seed: Option, }, /// Disaggregated guards cover the prefill+decode pair. EPD encode workers are /// assigned per item; their fire-and-supervise RPCs do not hold load guards. @@ -517,11 +534,55 @@ pub(crate) enum LoadGuards { Batch { _guards: Vec, _policy_reservation: Option, - _distribution_seed: Option, + distribution_seed: Option, }, + /// Test-only terminal-success probe for the distribution-seed commit stage. + /// Production construction still requires the real scheduler proof, policy + /// reservation, and adaptive headroom lease. + #[cfg(test)] + TestDistributionSeed { committed: Arc }, } impl LoadGuards { + pub(crate) fn has_distribution_seed(&self) -> bool { + match self { + Self::Single { + distribution_seed, .. + } + | Self::Batch { + distribution_seed, .. + } => distribution_seed.is_some(), + Self::Disaggregated { .. } => false, + #[cfg(test)] + Self::TestDistributionSeed { .. } => true, + } + } + + pub(crate) fn commit_distribution_seed_success(&mut self) { + match self { + Self::Single { + distribution_seed, .. + } + | Self::Batch { + distribution_seed, .. + } => { + if let Some(seed) = distribution_seed { + seed.commit_success(); + } + } + Self::Disaggregated { .. } => {} + #[cfg(test)] + Self::TestDistributionSeed { committed } => { + committed.store(true, Ordering::Release); + } + } + } + + #[cfg(test)] + pub(crate) fn test_distribution_seed(committed: Arc) -> Self { + Self::TestDistributionSeed { committed } + } + pub fn new( selection: &WorkerSelection, headers: Option<&HeaderMap>, @@ -532,7 +593,7 @@ impl LoadGuards { WorkerSelection::Single { worker } => LoadGuards::Single { _guard: WorkerLoadGuard::new(worker.clone(), headers), _policy_reservation: policy_reservation, - _distribution_seed: distribution_seed, + distribution_seed, }, WorkerSelection::Disaggregated { prefill, decode, .. @@ -563,7 +624,7 @@ impl LoadGuards { .map(|_| Self::new(selection, headers, None, None)) .collect(), _policy_reservation: policy_reservation, - _distribution_seed: distribution_seed, + distribution_seed, } } } diff --git a/model_gateway/src/routers/grpc/pipeline.rs b/model_gateway/src/routers/grpc/pipeline.rs index 14655986a..b853de91c 100644 --- a/model_gateway/src/routers/grpc/pipeline.rs +++ b/model_gateway/src/routers/grpc/pipeline.rs @@ -281,6 +281,7 @@ impl RequestPipeline { processor, streaming_processor, )), + Box::new(DistributionSeedCommitStage), ]); stages } @@ -310,6 +311,7 @@ impl RequestPipeline { processor, streaming_processor, )), + Box::new(DistributionSeedCommitStage), ]); stages } @@ -340,6 +342,7 @@ impl RequestPipeline { processor, streaming_processor, )), + Box::new(DistributionSeedCommitStage), ]); stages } @@ -364,6 +367,7 @@ impl RequestPipeline { Box::new(DispatchMetadataStage), Box::new(RequestExecutionStage::new()), Box::new(harmony::stages::HarmonyResponseProcessingStage::new()), + Box::new(DistributionSeedCommitStage), ] } Endpoint::Embeddings => { @@ -1246,6 +1250,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "ChatGenerateResponseProcessingStage", + "DistributionSeedCommitStage", ]), REGULAR, ), @@ -1259,6 +1264,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "ChatGenerateResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1273,6 +1279,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "ChatGenerateResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1286,6 +1293,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "MessageResponseProcessingStage", + "DistributionSeedCommitStage", ]), REGULAR, ), @@ -1299,6 +1307,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "MessageResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1313,6 +1322,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "MessageResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1326,6 +1336,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "CompletionResponseProcessingStage", + "DistributionSeedCommitStage", ]), REGULAR, ), @@ -1339,6 +1350,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "CompletionResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1353,6 +1365,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "CompletionResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ), @@ -1366,6 +1379,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "HarmonyResponseProcessingStage", + "DistributionSeedCommitStage", ]), REGULAR, ), @@ -1379,6 +1393,7 @@ mod build_parity_tests { "DispatchMetadataStage", "RequestExecutionStage", "HarmonyResponseProcessingStage", + "DistributionSeedCommitStage", ]), PD, ),