From f1107c21d94f97df3e71a71f8e796f4b14a115f9 Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Wed, 23 Sep 2026 22:30:51 +0300 Subject: [PATCH] Process-wide concurrency + Cancellation ownership fixes --- src/api/runtime_edge.rs | 6 +- src/api/users/view.rs | 2 +- src/conntrack_control/firewall.rs | 29 +- src/ip_tracker.rs | 47 +++- src/ip_tracker/admission.rs | 95 ++++--- src/ip_tracker/cleanup.rs | 193 ++++++++++---- src/ip_tracker/snapshot.rs | 25 +- src/ip_tracker/tests.rs | 83 +++++- src/ip_tracker/tests/cleanup_invariants.rs | 107 ++++++++ src/maestro/generation.rs | 31 ++- src/maestro/generation/lifecycle.rs | 157 +++++++++++ src/maestro/orchestrator.rs | 43 ++- src/maestro/reload_supervisor.rs | 41 ++- src/maestro/runtime_build.rs | 63 ++--- src/maestro/runtime_build_tests.rs | 33 +++ src/maestro/runtime_startup.rs | 10 +- src/maestro/runtime_tasks.rs | 53 ++-- src/metrics/render/users.rs | 2 +- src/metrics/tests.rs | 4 + src/proxy/authenticated.rs | 134 ++++++---- src/proxy/client/authenticated.rs | 10 +- src/proxy/direct_buffer_budget.rs | 248 ++---------------- src/proxy/direct_buffer_budget/controller.rs | 238 +++++++++++++++++ src/proxy/middle_relay.rs | 3 +- src/proxy/middle_relay/quota.rs | 3 +- src/proxy/middle_relay/session.rs | 30 ++- src/proxy/middle_relay/session/children.rs | 65 +++++ src/proxy/mod.rs | 2 + src/proxy/relay/io.rs | 118 ++++----- src/proxy/relay/io/quota.rs | 3 +- src/proxy/shared_state.rs | 15 +- src/proxy/tests/client_security_tests.rs | 140 +++++++++- src/proxy/tests/direct_buffer_budget_tests.rs | 60 +++++ src/proxy/traffic_limiter.rs | 2 +- src/proxy/traffic_limiter/lease.rs | 17 +- src/proxy/traffic_limiter/limiter.rs | 43 ++- src/proxy/traffic_limiter/tests.rs | 73 ++++++ src/proxy/user_connection_authority.rs | 123 +++++++++ src/stats/mod.rs | 27 +- src/stats/tests.rs | 26 ++ src/stats/users.rs | 63 +++-- src/synlimit_control/command.rs | 36 ++- src/transport/middle_proxy/send/selection.rs | 219 +++++++++------- src/web/http/websocket/driver.rs | 81 +++++- src/web/http/websocket/driver/lane.rs | 37 ++- src/web/session/lifecycle.rs | 22 +- src/web/session/uplink_tests.rs | 77 ++++++ 47 files changed, 2177 insertions(+), 762 deletions(-) create mode 100644 src/ip_tracker/tests/cleanup_invariants.rs create mode 100644 src/maestro/generation/lifecycle.rs create mode 100644 src/proxy/direct_buffer_budget/controller.rs create mode 100644 src/proxy/middle_relay/session/children.rs create mode 100644 src/proxy/user_connection_authority.rs diff --git a/src/api/runtime_edge.rs b/src/api/runtime_edge.rs index 639dbe0..0303de2 100644 --- a/src/api/runtime_edge.rs +++ b/src/api/runtime_edge.rs @@ -314,9 +314,9 @@ async fn recompute_connections_payload( let mut active_users = 0usize; for entry in shared.stats.iter_user_stats() { let user_stats = entry.value(); - let current_connections = user_stats - .curr_connects - .load(std::sync::atomic::Ordering::Relaxed); + let current_connections = shared + .stats + .get_process_user_curr_connects(entry.key()); let total_octets = user_stats .octets_from_client .load(std::sync::atomic::Ordering::Relaxed) diff --git a/src/api/users/view.rs b/src/api/users/view.rs index f38ce75..4c89d4d 100644 --- a/src/api/users/view.rs +++ b/src/api/users/view.rs @@ -71,7 +71,7 @@ pub(in crate::api) async fn users_from_config( .filter(|limit| *limit > 0) .or((cfg.access.user_max_unique_ips_global_each > 0) .then_some(cfg.access.user_max_unique_ips_global_each)), - current_connections: stats.get_user_curr_connects(&username), + current_connections: stats.get_process_user_curr_connects(&username), active_unique_ips: active_ip_list.len(), active_unique_ips_list: active_ip_list, recent_unique_ips: recent_ip_list.len(), diff --git a/src/conntrack_control/firewall.rs b/src/conntrack_control/firewall.rs index 5632594..2fc42f7 100644 --- a/src/conntrack_control/firewall.rs +++ b/src/conntrack_control/firewall.rs @@ -1,5 +1,6 @@ use std::collections::BTreeSet; use std::net::IpAddr; +use std::time::Duration; use tokio::io::AsyncWriteExt; use tokio::process::Command; @@ -363,6 +364,7 @@ pub(super) async fn delete_conntrack_entry(event: ConntrackCloseEvent) -> Delete } async fn run_command(binary: &str, args: &[&str], stdin: Option) -> Result<(), String> { + const COMMAND_TIMEOUT: Duration = Duration::from_secs(30); #[cfg(unix)] let Some(command_path) = resolve_trusted_helper(binary) else { return Err(format!("{binary} is not available")); @@ -377,21 +379,26 @@ async fn run_command(binary: &str, args: &[&str], stdin: Option) -> Resu } command.stdout(std::process::Stdio::null()); command.stderr(std::process::Stdio::piped()); + command.kill_on_drop(true); let mut child = command .spawn() .map_err(|error| format!("spawn {binary} failed: {error}"))?; - if let Some(blob) = stdin - && let Some(mut writer) = child.stdin.take() - { - writer - .write_all(blob.as_bytes()) + let output = tokio::time::timeout(COMMAND_TIMEOUT, async move { + if let Some(blob) = stdin + && let Some(mut writer) = child.stdin.take() + { + writer + .write_all(blob.as_bytes()) + .await + .map_err(|error| format!("stdin write {binary} failed: {error}"))?; + } + child + .wait_with_output() .await - .map_err(|error| format!("stdin write {binary} failed: {error}"))?; - } - let output = child - .wait_with_output() - .await - .map_err(|error| format!("wait {binary} failed: {error}"))?; + .map_err(|error| format!("wait {binary} failed: {error}")) + }) + .await + .map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??; if output.status.success() { return Ok(()); } diff --git a/src/ip_tracker.rs b/src/ip_tracker.rs index 990a405..b5dae17 100644 --- a/src/ip_tracker.rs +++ b/src/ip_tracker.rs @@ -38,11 +38,16 @@ struct UserIpShard { #[derive(Debug, Default)] struct CleanupShard { - queue: Mutex>>, + queue: Mutex, } +type CleanupQueue = + HashMap>>; +type CleanupBatch = HashMap<(String, UserIncarnation, IpAddr), usize>; + #[derive(Debug, Clone)] struct UserIpLimitPolicy { + source_generation: u64, max_ips: Arc>, default_max_ips: usize, mode: UserMaxUniqueIpsMode, @@ -52,6 +57,7 @@ struct UserIpLimitPolicy { impl Default for UserIpLimitPolicy { fn default() -> Self { Self { + source_generation: 0, max_ips: Arc::new(HashMap::new()), default_max_ips: 0, mode: UserMaxUniqueIpsMode::ActiveWindow, @@ -70,6 +76,7 @@ pub struct UserIpTracker { recent_cap_rejects: Arc, cleanup_deferred_releases: Arc, limit_policy: Arc>, + policy_update: Arc>, last_compact_epoch_secs: Arc, cleanup_queue_len: Arc, cleanup_shards: Arc>, @@ -121,6 +128,7 @@ impl UserIpTracker { recent_cap_rejects: Arc::new(AtomicU64::new(0)), cleanup_deferred_releases: Arc::new(AtomicU64::new(0)), limit_policy: Arc::new(ArcSwap::from_pointee(UserIpLimitPolicy::default())), + policy_update: Arc::new(Mutex::new(())), last_compact_epoch_secs: Arc::new(AtomicU64::new(0)), cleanup_queue_len: Arc::new(AtomicU64::new(0)), cleanup_shards: Arc::new(cleanup_shards), @@ -196,19 +204,21 @@ impl UserIpTracker { } pub(super) fn pop_one_cleanup( - queue: &mut HashMap<(String, UserIncarnation), HashMap>, + queue: &mut CleanupQueue, ) -> Option<(String, UserIncarnation, IpAddr, usize)> { - let owner = queue.keys().next().cloned()?; - let ip = queue.get(&owner)?.keys().next().copied()?; - let count = queue.get_mut(&owner)?.remove(&ip)?; - let remove_user = queue - .get(&owner) - .map(|user_queue| user_queue.is_empty()) - .unwrap_or(false); - if remove_user { - queue.remove(&owner); + let user = queue.keys().next().cloned()?; + let incarnation = queue.get(&user)?.keys().next().copied()?; + let ip = queue.get(&user)?.get(&incarnation)?.keys().next().copied()?; + let incarnations = queue.get_mut(&user)?; + let ips = incarnations.get_mut(&incarnation)?; + let count = ips.remove(&ip)?; + if ips.is_empty() { + incarnations.remove(&incarnation); } - Some((owner.0, owner.1, ip, count)) + if incarnations.is_empty() { + queue.remove(&user); + } + Some((user, incarnation, ip, count)) } #[cfg(test)] @@ -224,6 +234,19 @@ impl UserIpTracker { #[cfg(not(test))] pub(super) fn observe_cleanup_poison_for_tests(&self) {} + #[cfg(test)] + pub(crate) async fn hold_user_shard_for_tests( + &self, + user: &str, + entered: tokio::sync::oneshot::Sender<()>, + release: tokio::sync::oneshot::Receiver<()>, + ) { + let shard_idx = Self::shard_idx(user); + let _guard = self.shards[shard_idx].write().await; + let _ = entered.send(()); + let _ = release.await; + } + pub(super) fn now_epoch_secs() -> u64 { std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) diff --git a/src/ip_tracker/admission.rs b/src/ip_tracker/admission.rs index 26e477b..8fb39a1 100644 --- a/src/ip_tracker/admission.rs +++ b/src/ip_tracker/admission.rs @@ -2,47 +2,84 @@ use super::*; impl UserIpTracker { pub async fn set_limit_policy(&self, mode: UserMaxUniqueIpsMode, window_secs: u64) { - self.limit_policy.rcu(|current| { - Arc::new(UserIpLimitPolicy { - mode, - window_secs: window_secs.max(1), - ..(**current).clone() - }) + let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| { + self.policy_update.clear_poison(); + poisoned.into_inner() }); + let current = self.limit_policy.load_full(); + self.limit_policy.store(Arc::new(UserIpLimitPolicy { + mode, + window_secs: window_secs.max(1), + ..(*current).clone() + })); } pub async fn set_user_limit(&self, username: &str, max_ips: usize) { - let username = username.to_string(); - self.limit_policy.rcu(|current| { - let mut limits = current.max_ips.as_ref().clone(); - limits.insert(username.clone(), max_ips); - Arc::new(UserIpLimitPolicy { - max_ips: Arc::new(limits), - ..(**current).clone() - }) + let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| { + self.policy_update.clear_poison(); + poisoned.into_inner() }); + let current = self.limit_policy.load_full(); + let mut limits = current.max_ips.as_ref().clone(); + limits.insert(username.to_string(), max_ips); + self.limit_policy.store(Arc::new(UserIpLimitPolicy { + max_ips: Arc::new(limits), + ..(*current).clone() + })); } pub async fn remove_user_limit(&self, username: &str) { - self.limit_policy.rcu(|current| { - let mut limits = current.max_ips.as_ref().clone(); - limits.remove(username); - Arc::new(UserIpLimitPolicy { - max_ips: Arc::new(limits), - ..(**current).clone() - }) + let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| { + self.policy_update.clear_poison(); + poisoned.into_inner() }); + let current = self.limit_policy.load_full(); + let mut limits = current.max_ips.as_ref().clone(); + limits.remove(username); + self.limit_policy.store(Arc::new(UserIpLimitPolicy { + max_ips: Arc::new(limits), + ..(*current).clone() + })); } pub async fn load_limits(&self, default_limit: usize, limits: &HashMap) { - let limits = Arc::new(limits.clone()); - self.limit_policy.rcu(|current| { - Arc::new(UserIpLimitPolicy { - max_ips: Arc::clone(&limits), - default_max_ips: default_limit, - ..(**current).clone() - }) + let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| { + self.policy_update.clear_poison(); + poisoned.into_inner() }); + let current = self.limit_policy.load_full(); + self.limit_policy.store(Arc::new(UserIpLimitPolicy { + max_ips: Arc::new(limits.clone()), + default_max_ips: default_limit, + ..(*current).clone() + })); + } + + /// Atomically publishes one coherent policy from the active runtime generation. + pub(crate) async fn apply_policy_from_source( + &self, + source_generation: u64, + default_limit: usize, + limits: &HashMap, + mode: UserMaxUniqueIpsMode, + window_secs: u64, + ) -> bool { + let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| { + self.policy_update.clear_poison(); + poisoned.into_inner() + }); + let current = self.limit_policy.load_full(); + if source_generation < current.source_generation { + return false; + } + self.limit_policy.store(Arc::new(UserIpLimitPolicy { + source_generation, + max_ips: Arc::new(limits.clone()), + default_max_ips: default_limit, + mode, + window_secs: window_secs.max(1), + })); + true } pub(super) fn prune_recent( @@ -70,7 +107,6 @@ impl UserIpTracker { ip: IpAddr, ) -> Result<(), String> { self.drain_cleanup_for_user(username).await; - self.maybe_compact_empty_users().await; let policy = self.limit_policy.load(); let limit = Self::user_limit(&policy, username); let mode = policy.mode; @@ -218,7 +254,6 @@ impl UserIpTracker { incarnation: UserIncarnation, ip: IpAddr, ) { - self.maybe_compact_empty_users().await; let shard_idx = Self::shard_idx(username); let mut shard = self.shards[shard_idx].write().await; if shard.incarnations.get(username).copied() != Some(incarnation) { diff --git a/src/ip_tracker/cleanup.rs b/src/ip_tracker/cleanup.rs index a884009..f72250a 100644 --- a/src/ip_tracker/cleanup.rs +++ b/src/ip_tracker/cleanup.rs @@ -1,5 +1,68 @@ use super::*; +struct DetachedCleanupBatch<'a> { + tracker: &'a UserIpTracker, + shard_idx: usize, + entries: CleanupBatch, +} + +impl DetachedCleanupBatch<'_> { + fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + fn entries(&self) -> impl Iterator { + self.entries.iter() + } + + fn commit(mut self) { + let committed = self.entries.len(); + self.entries.clear(); + UserIpTracker::decrement_counter(&self.tracker.cleanup_queue_len, committed); + } +} + +impl Drop for DetachedCleanupBatch<'_> { + fn drop(&mut self) { + if self.entries.is_empty() { + return; + } + + let cleanup_shard = &self.tracker.cleanup_shards[self.shard_idx]; + let mut duplicate_entries = 0usize; + let mut restore = |queue: &mut CleanupQueue| { + for ((user, incarnation, ip), count) in self.entries.drain() { + let queued = queue + .entry(user) + .or_default() + .entry(incarnation) + .or_default() + .entry(ip) + .or_insert(0); + if *queued != 0 { + duplicate_entries = duplicate_entries.saturating_add(1); + } + *queued = queued.saturating_add(count); + } + }; + match cleanup_shard.queue.lock() { + Ok(mut queue) => restore(&mut queue), + Err(poisoned) => { + let mut queue = poisoned.into_inner(); + restore(&mut queue); + cleanup_shard.queue.clear_poison(); + tracing::warn!( + "UserIpTracker cleanup_queue lock poisoned while restoring a cancelled cleanup batch" + ); + } + } + UserIpTracker::decrement_counter( + &self.tracker.cleanup_queue_len, + duplicate_entries, + ); + } +} + impl UserIpTracker { /// Queues a deferred active IP cleanup for a later async drain. pub fn enqueue_cleanup(&self, user: String, ip: IpAddr) { @@ -18,8 +81,13 @@ impl UserIpTracker { let cleanup_shard = &self.cleanup_shards[shard_idx]; match cleanup_shard.queue.lock() { Ok(mut queue) => { - let user_queue = queue.entry((user, incarnation)).or_default(); - let count = user_queue.entry(ip).or_insert(0); + let count = queue + .entry(user) + .or_default() + .entry(incarnation) + .or_default() + .entry(ip) + .or_insert(0); if *count == 0 { self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); } @@ -29,8 +97,13 @@ impl UserIpTracker { } Err(poisoned) => { let mut queue = poisoned.into_inner(); - let user_queue = queue.entry((user.clone(), incarnation)).or_default(); - let count = user_queue.entry(ip).or_insert(0); + let count = queue + .entry(user.clone()) + .or_default() + .entry(incarnation) + .or_default() + .entry(ip) + .or_insert(0); if *count == 0 { self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); } @@ -52,6 +125,26 @@ impl UserIpTracker { self.cleanup_queue_len.load(Ordering::Relaxed) as usize } + #[cfg(test)] + pub(crate) fn cleanup_queue_physical_entries_for_tests(&self) -> usize { + self.cleanup_shards + .iter() + .map(|cleanup_shard| { + let count = |queue: &CleanupQueue| { + queue + .values() + .flat_map(HashMap::values) + .map(HashMap::len) + .sum::() + }; + match cleanup_shard.queue.lock() { + Ok(queue) => count(&queue), + Err(poisoned) => count(&poisoned.into_inner()), + } + }) + .sum() + } + #[cfg(test)] pub(crate) fn cleanup_queue_mutex_for_tests( &self, @@ -73,12 +166,13 @@ impl UserIpTracker { return; } let shard_idx = Self::shard_idx(user); + let _drain_guard = self.cleanup_drain_locks[shard_idx].lock().await; let cleanup_shard = &self.cleanup_shards[shard_idx]; let to_remove = match cleanup_shard.queue.lock() { - Ok(mut queue) => drain_user_cleanup(&mut queue, user), + Ok(mut queue) => detach_user_cleanup(self, shard_idx, &mut queue, user), Err(poisoned) => { let mut queue = poisoned.into_inner(); - let drained = drain_user_cleanup(&mut queue, user); + let drained = detach_user_cleanup(self, shard_idx, &mut queue, user); cleanup_shard.queue.clear_poison(); drained } @@ -86,25 +180,24 @@ impl UserIpTracker { if to_remove.is_empty() { return; } - let removed_queue_entries = to_remove - .iter() - .map(|(_, ips)| ips.len()) - .sum::(); - self.cleanup_queue_len - .fetch_sub(removed_queue_entries as u64, Ordering::Relaxed); let mut shard = self.shards[shard_idx].write().await; let mut removed_active_entries = 0usize; - for (incarnation, ips) in to_remove { - if shard.incarnations.get(user).copied() != Some(incarnation) { + for ((queued_user, incarnation, ip), pending_count) in to_remove.entries() { + if shard.incarnations.get(queued_user).copied() != Some(*incarnation) { continue; } - for (ip, pending_count) in ips { - removed_active_entries = removed_active_entries.saturating_add( - Self::apply_active_cleanup(&mut shard.active_ips, user, ip, pending_count), - ); - } + removed_active_entries = removed_active_entries.saturating_add( + Self::apply_active_cleanup( + &mut shard.active_ips, + queued_user, + *ip, + *pending_count, + ), + ); } Self::decrement_counter(&self.active_entry_count, removed_active_entries); + drop(shard); + to_remove.commit(); } pub(super) async fn drain_cleanup_shard(&self, shard_idx: usize) { @@ -119,18 +212,20 @@ impl UserIpTracker { if queue.is_empty() { return; } - let mut drained = - HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); + let mut drained = HashMap::with_capacity(CLEANUP_DRAIN_BATCH_LIMIT); for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { let Some((user, incarnation, ip, count)) = Self::pop_one_cleanup(&mut queue) else { break; }; - self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); drained.insert((user, incarnation, ip), count); } - drained + DetachedCleanupBatch { + tracker: self, + shard_idx, + entries: drained, + } } Err(poisoned) => { let mut queue = poisoned.into_inner(); @@ -138,55 +233,61 @@ impl UserIpTracker { cleanup_shard.queue.clear_poison(); return; } - let mut drained = - HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); + let mut drained = HashMap::with_capacity(CLEANUP_DRAIN_BATCH_LIMIT); for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { let Some((user, incarnation, ip, count)) = Self::pop_one_cleanup(&mut queue) else { break; }; - self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); drained.insert((user, incarnation, ip), count); } cleanup_shard.queue.clear_poison(); - drained + DetachedCleanupBatch { + tracker: self, + shard_idx, + entries: drained, + } } } }; - drop(_drain_guard); if to_remove.is_empty() { return; } let mut shard = self.shards[shard_idx].write().await; let mut removed_active_entries = 0usize; - for ((user, incarnation, ip), pending_count) in to_remove { - if shard.incarnations.get(&user).copied() != Some(incarnation) { + for ((user, incarnation, ip), pending_count) in to_remove.entries() { + if shard.incarnations.get(user).copied() != Some(*incarnation) { continue; } removed_active_entries = removed_active_entries.saturating_add( - Self::apply_active_cleanup(&mut shard.active_ips, &user, ip, pending_count), + Self::apply_active_cleanup(&mut shard.active_ips, user, *ip, *pending_count), ); } Self::decrement_counter(&self.active_entry_count, removed_active_entries); + drop(shard); + to_remove.commit(); } } -fn drain_user_cleanup( - queue: &mut HashMap<(String, UserIncarnation), HashMap>, +fn detach_user_cleanup<'a>( + tracker: &'a UserIpTracker, + shard_idx: usize, + queue: &mut CleanupQueue, user: &str, -) -> Vec<(UserIncarnation, HashMap)> { - let owners = queue - .keys() - .filter(|(queued_user, _)| queued_user == user) - .cloned() - .collect::>(); - owners - .into_iter() - .filter_map(|owner| { - let incarnation = owner.1; - queue.remove(&owner).map(|ips| (incarnation, ips)) - }) - .collect() +) -> DetachedCleanupBatch<'a> { + let mut entries = CleanupBatch::new(); + if let Some(incarnations) = queue.remove(user) { + for (incarnation, ips) in incarnations { + for (ip, count) in ips { + entries.insert((user.to_string(), incarnation, ip), count); + } + } + } + DetachedCleanupBatch { + tracker, + shard_idx, + entries, + } } diff --git a/src/ip_tracker/snapshot.rs b/src/ip_tracker/snapshot.rs index 816e01b..7e89ee1 100644 --- a/src/ip_tracker/snapshot.rs +++ b/src/ip_tracker/snapshot.rs @@ -68,10 +68,13 @@ impl UserIpTracker { } } - pub async fn run_periodic_maintenance(self: Arc) { + pub async fn run_periodic_maintenance(self: Arc, source_generation: u64) { let mut interval = tokio::time::interval(Duration::from_secs(1)); loop { interval.tick().await; + if self.limit_policy.load().source_generation != source_generation { + continue; + } self.drain_cleanup_queue().await; self.maybe_compact_empty_users().await; } @@ -263,6 +266,10 @@ impl UserIpTracker { } pub async fn clear_all(&self) { + let mut cleanup_drain_guards = Vec::with_capacity(USER_IP_TRACKER_SHARDS); + for drain_lock in self.cleanup_drain_locks.iter() { + cleanup_drain_guards.push(drain_lock.lock().await); + } for shard_lock in self.shards.iter() { let mut shard = shard_lock.write().await; shard.active_ips.clear(); @@ -271,16 +278,24 @@ impl UserIpTracker { } self.active_entry_count.store(0, Ordering::Relaxed); self.recent_entry_count.store(0, Ordering::Relaxed); + let mut cleanup_queue_guards = Vec::with_capacity(USER_IP_TRACKER_SHARDS); for cleanup_shard in self.cleanup_shards.iter() { - match cleanup_shard.queue.lock() { - Ok(mut queue) => queue.clear(), + let queue = match cleanup_shard.queue.lock() { + Ok(queue) => queue, Err(poisoned) => { - poisoned.into_inner().clear(); + let queue = poisoned.into_inner(); cleanup_shard.queue.clear_poison(); + queue } - } + }; + cleanup_queue_guards.push(queue); + } + for queue in cleanup_queue_guards.iter_mut() { + queue.clear(); } self.cleanup_queue_len.store(0, Ordering::Relaxed); + drop(cleanup_queue_guards); + drop(cleanup_drain_guards); } pub async fn is_ip_active(&self, username: &str, ip: IpAddr) -> bool { diff --git a/src/ip_tracker/tests.rs b/src/ip_tracker/tests.rs index 13e41ff..dc024e4 100644 --- a/src/ip_tracker/tests.rs +++ b/src/ip_tracker/tests.rs @@ -3,6 +3,8 @@ use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; use std::sync::atomic::AtomicBool; use std::sync::atomic::Ordering; +mod cleanup_invariants; + fn test_ipv4(oct1: u8, oct2: u8, oct3: u8, oct4: u8) -> IpAddr { IpAddr::V4(Ipv4Addr::new(oct1, oct2, oct3, oct4)) } @@ -266,6 +268,82 @@ async fn test_load_limits_replaces_previous_map() { assert_eq!(tracker.get_user_limit("user2").await, Some(5)); } +#[tokio::test] +async fn stale_runtime_cannot_overwrite_newer_ip_policy() { + let tracker = UserIpTracker::new(); + let mut newer = HashMap::new(); + newer.insert("alice".to_string(), 5); + assert!( + tracker + .apply_policy_from_source( + 2, + 7, + &newer, + UserMaxUniqueIpsMode::Combined, + 90, + ) + .await + ); + + let mut stale = HashMap::new(); + stale.insert("alice".to_string(), 1); + assert!( + !tracker + .apply_policy_from_source( + 1, + 1, + &stale, + UserMaxUniqueIpsMode::ActiveWindow, + 1, + ) + .await + ); + + let policy = tracker.limit_policy.load_full(); + assert_eq!(policy.source_generation, 2); + assert_eq!(policy.default_max_ips, 7); + assert_eq!(policy.max_ips["alice"], 5); + assert_eq!(policy.mode, UserMaxUniqueIpsMode::Combined); + assert_eq!(policy.window_secs, 90); +} + +#[tokio::test] +async fn active_runtime_can_publish_coherent_same_generation_ip_policy() { + let tracker = UserIpTracker::new(); + assert!( + tracker + .apply_policy_from_source( + 3, + 1, + &HashMap::new(), + UserMaxUniqueIpsMode::ActiveWindow, + 10, + ) + .await + ); + let mut limits = HashMap::new(); + limits.insert("alice".to_string(), 4); + + assert!( + tracker + .apply_policy_from_source( + 3, + 6, + &limits, + UserMaxUniqueIpsMode::TimeWindow, + 30, + ) + .await + ); + + let policy = tracker.limit_policy.load_full(); + assert_eq!(policy.source_generation, 3); + assert_eq!(policy.default_max_ips, 6); + assert_eq!(policy.max_ips["alice"], 4); + assert_eq!(policy.mode, UserMaxUniqueIpsMode::TimeWindow); + assert_eq!(policy.window_secs, 30); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn concurrent_policy_replacement_never_exposes_partial_limit_map() { const USER_COUNT: usize = 4_096; @@ -437,10 +515,7 @@ async fn test_compact_prunes_stale_recent_entries() { } tracker.last_compact_epoch_secs.store(0, Ordering::Relaxed); - tracker - .check_and_add("trigger-user", test_ipv4(10, 3, 0, 2)) - .await - .unwrap(); + tracker.maybe_compact_empty_users().await; let shard_idx = UserIpTracker::shard_idx(&stale_user); let shard = tracker.shards[shard_idx].read().await; diff --git a/src/ip_tracker/tests/cleanup_invariants.rs b/src/ip_tracker/tests/cleanup_invariants.rs new file mode 100644 index 0000000..a98d659 --- /dev/null +++ b/src/ip_tracker/tests/cleanup_invariants.rs @@ -0,0 +1,107 @@ +use super::super::*; +use std::net::{IpAddr, Ipv4Addr}; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +fn test_ipv4(oct1: u8, oct2: u8, oct3: u8, oct4: u8) -> IpAddr { + IpAddr::V4(Ipv4Addr::new(oct1, oct2, oct3, oct4)) +} + +#[tokio::test] +async fn cancelled_cleanup_drain_restores_detached_batch() { + let tracker = Arc::new(UserIpTracker::new()); + let user = "cancelled-cleanup-user"; + let ip = test_ipv4(10, 2, 1, 1); + tracker.check_and_add(user, ip).await.unwrap(); + tracker.enqueue_cleanup(user.to_string(), ip); + + let (entered_tx, entered_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel(); + let held_tracker = Arc::clone(&tracker); + let held_user = user.to_string(); + let holder = tokio::spawn(async move { + held_tracker + .hold_user_shard_for_tests(&held_user, entered_tx, release_rx) + .await; + }); + entered_rx.await.unwrap(); + + let shard_idx = UserIpTracker::shard_idx(user); + let drain_tracker = Arc::clone(&tracker); + let drain = tokio::spawn(async move { + drain_tracker.drain_cleanup_shard(shard_idx).await; + }); + tokio::time::timeout(Duration::from_secs(1), async { + while tracker.cleanup_queue_physical_entries_for_tests() != 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("cleanup batch must detach before waiting for the IP shard"); + + drain.abort(); + assert!(drain.await.unwrap_err().is_cancelled()); + assert_eq!(tracker.cleanup_queue_len_for_tests(), 1); + assert_eq!(tracker.cleanup_queue_physical_entries_for_tests(), 1); + + let _ = release_tx.send(()); + holder.await.unwrap(); + tracker.drain_cleanup_queue().await; + assert_eq!(tracker.cleanup_queue_len_for_tests(), 0); + assert_eq!(tracker.get_active_ip_count(user).await, 0); +} + +#[test] +fn clear_all_serializes_queue_reset_with_concurrent_enqueue() { + let tracker = Arc::new(UserIpTracker::new()); + let first_shard_user = (0u64..) + .map(|index| format!("clear-race-{index}")) + .find(|user| UserIpTracker::shard_idx(user) == 0) + .unwrap(); + let last_shard = USER_IP_TRACKER_SHARDS - 1; + let last_queue_guard = tracker.cleanup_shards[last_shard].queue.lock().unwrap(); + + let clear_tracker = Arc::clone(&tracker); + let clear = std::thread::spawn(move || { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + runtime.block_on(clear_tracker.clear_all()); + }); + + let wait_deadline = Instant::now() + Duration::from_secs(1); + loop { + match tracker.cleanup_shards[0].queue.try_lock() { + Ok(queue) => drop(queue), + Err(std::sync::TryLockError::WouldBlock) => break, + Err(std::sync::TryLockError::Poisoned(_)) => panic!("cleanup queue lock poisoned"), + } + assert!(Instant::now() < wait_deadline, "clear_all did not reach queue reset"); + std::thread::yield_now(); + } + + let (started_tx, started_rx) = std::sync::mpsc::channel(); + let (completed_tx, completed_rx) = std::sync::mpsc::channel(); + let enqueue_tracker = Arc::clone(&tracker); + let enqueue = std::thread::spawn(move || { + started_tx.send(()).unwrap(); + enqueue_tracker.enqueue_cleanup( + first_shard_user, + test_ipv4(10, 2, 2, 1), + ); + completed_tx.send(()).unwrap(); + }); + started_rx.recv().unwrap(); + assert!(completed_rx + .recv_timeout(Duration::from_millis(50)) + .is_err()); + + drop(last_queue_guard); + clear.join().unwrap(); + completed_rx.recv_timeout(Duration::from_secs(1)).unwrap(); + enqueue.join().unwrap(); + + assert_eq!(tracker.cleanup_queue_len_for_tests(), 1); + assert_eq!(tracker.cleanup_queue_physical_entries_for_tests(), 1); +} diff --git a/src/maestro/generation.rs b/src/maestro/generation.rs index 870b683..10da12f 100644 --- a/src/maestro/generation.rs +++ b/src/maestro/generation.rs @@ -22,6 +22,11 @@ use crate::tls_front::TlsFrontCache; use crate::transport::UpstreamManager; use crate::transport::middle_proxy::MePool; +// Cancellation guards preserve runtime ownership across preparation and drain futures. +mod lifecycle; +pub(crate) use lifecycle::RuntimeTaskScopePreparationGuard; +use lifecycle::SessionDrainCancellationGuard; + const SESSION_STOP_TIMEOUT: Duration = Duration::from_secs(5); const BACKGROUND_STOP_TIMEOUT: Duration = Duration::from_secs(5); const SESSION_ADMISSION_CLOSED: usize = 1 << (usize::BITS - 1); @@ -145,12 +150,17 @@ impl RuntimeTaskScope { self.cancel.clone() } - /// Cancels the scope and waits within the bounded background-task budget. - pub(crate) async fn stop(&self) { + /// Synchronously closes task admission and signals every tracked task. + pub(crate) fn begin_stop(&self) { self.admission.close(); - self.admission.wait_for_registrations().await; self.cancel.cancel(); self.tracker.close(); + } + + /// Cancels the scope and waits within the bounded background-task budget. + pub(crate) async fn stop(&self) { + self.begin_stop(); + self.admission.wait_for_registrations().await; let _ = tokio::time::timeout(BACKGROUND_STOP_TIMEOUT, self.tracker.wait()).await; } } @@ -298,24 +308,31 @@ impl RuntimeGeneration { /// Waits for registered sessions and cancels them when the deadline expires. pub(crate) async fn drain_sessions(&self, timeout: Duration) -> bool { self.stop_accepting_sessions(); + let mut cancellation_guard = SessionDrainCancellationGuard::new(self); self.session_admission.wait_for_registrations().await; self.sessions.close(); if tokio::time::timeout(timeout, self.sessions.wait()) .await .is_ok() { + cancellation_guard.disarm(); return true; } + cancellation_guard.disarm(); self.stop_sessions().await; false } - /// Cancels all sessions and waits within the bounded session-stop budget. - pub(crate) async fn stop_sessions(&self) { + fn begin_stop_sessions(&self) { self.stop_accepting_sessions(); - self.session_admission.wait_for_registrations().await; self.session_cancel.cancel(); self.sessions.close(); + } + + /// Cancels all sessions and waits within the bounded session-stop budget. + pub(crate) async fn stop_sessions(&self) { + self.begin_stop_sessions(); + self.session_admission.wait_for_registrations().await; let _ = tokio::time::timeout(SESSION_STOP_TIMEOUT, self.sessions.wait()).await; } @@ -335,6 +352,8 @@ impl RuntimeGeneration { impl Drop for RuntimeGeneration { fn drop(&mut self) { + self.background_tasks.begin_stop(); + self.begin_stop_sessions(); if let Some(pool) = self.me_pool.as_ref() { pool.begin_shutdown(); } diff --git a/src/maestro/generation/lifecycle.rs b/src/maestro/generation/lifecycle.rs new file mode 100644 index 0000000..8176b84 --- /dev/null +++ b/src/maestro/generation/lifecycle.rs @@ -0,0 +1,157 @@ +use super::*; + +/// Cancels tasks if runtime preparation exits before ownership reaches a generation. +#[must_use = "runtime preparation guards must be disarmed after ownership transfer"] +pub(crate) struct RuntimeTaskScopePreparationGuard { + scope: RuntimeTaskScope, + armed: bool, +} + +impl RuntimeTaskScopePreparationGuard { + /// Arms cancellation for a newly created, not-yet-published task scope. + pub(crate) fn new(scope: RuntimeTaskScope) -> Self { + Self { scope, armed: true } + } + + /// Confirms that a runtime generation now owns the task scope. + pub(crate) fn disarm(mut self) { + self.armed = false; + } +} + +impl Drop for RuntimeTaskScopePreparationGuard { + fn drop(&mut self) { + if self.armed { + self.scope.begin_stop(); + } + } +} + +pub(super) struct SessionDrainCancellationGuard<'a> { + generation: &'a RuntimeGeneration, + armed: bool, +} + +impl<'a> SessionDrainCancellationGuard<'a> { + pub(super) fn new(generation: &'a RuntimeGeneration) -> Self { + Self { + generation, + armed: true, + } + } + + pub(super) fn disarm(&mut self) { + self.armed = false; + } +} + +impl Drop for SessionDrainCancellationGuard<'_> { + fn drop(&mut self) { + if self.armed { + self.generation.begin_stop_sessions(); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct NotifyOnDrop(Arc); + + impl Drop for NotifyOnDrop { + fn drop(&mut self) { + self.0.notify_one(); + } + } + + #[tokio::test] + async fn preparation_guard_drop_cancels_scope_and_rejects_late_spawn() { + let scope = RuntimeTaskScope::new(); + let guard = RuntimeTaskScopePreparationGuard::new(scope.clone()); + drop(guard); + + assert!(scope.cancellation_token().is_cancelled()); + let ran = Arc::new(AtomicUsize::new(0)); + let ran_task = Arc::clone(&ran); + scope.spawn(async move { + ran_task.fetch_add(1, Ordering::AcqRel); + }); + tokio::task::yield_now().await; + assert_eq!(ran.load(Ordering::Acquire), 0); + } + + #[tokio::test] + async fn preparation_guard_disarm_transfers_ownership() { + let scope = RuntimeTaskScope::new(); + RuntimeTaskScopePreparationGuard::new(scope.clone()).disarm(); + + assert!(!scope.cancellation_token().is_cancelled()); + scope.stop().await; + } + + #[tokio::test] + async fn aborted_scope_stop_still_cancels_children() { + let scope = RuntimeTaskScope::new(); + let registration = scope.admission.try_register().unwrap(); + let started = Arc::new(Notify::new()); + let dropped = Arc::new(Notify::new()); + let started_task = Arc::clone(&started); + let dropped_task = Arc::clone(&dropped); + scope.spawn(async move { + let _drop_signal = NotifyOnDrop(dropped_task); + started_task.notify_one(); + std::future::pending::<()>().await; + }); + started.notified().await; + + let stop_scope = scope.clone(); + let stop = tokio::spawn(async move { + stop_scope.stop().await; + }); + scope.cancellation_token().cancelled().await; + stop.abort(); + assert!(stop.await.unwrap_err().is_cancelled()); + drop(registration); + + tokio::time::timeout(Duration::from_secs(1), dropped.notified()) + .await + .expect("scope cancellation must drop the tracked child"); + assert!(scope.admission.try_register().is_none()); + } + + #[tokio::test] + async fn aborted_graceful_drain_forces_session_cancellation() { + let generation = test_runtime_generation(1, ProxyConfig::default()); + let registration = generation.session_admission.try_register().unwrap(); + let started = Arc::new(Notify::new()); + let dropped = Arc::new(Notify::new()); + let started_task = Arc::clone(&started); + let dropped_task = Arc::clone(&dropped); + assert!(generation.spawn_session(async move { + let _drop_signal = NotifyOnDrop(dropped_task); + started_task.notify_one(); + std::future::pending::<()>().await; + })); + started.notified().await; + + let drain_generation = Arc::clone(&generation); + let drain = tokio::spawn(async move { + drain_generation.drain_sessions(Duration::from_secs(60)).await + }); + while generation.session_admission.state.load(Ordering::Acquire) + & SESSION_ADMISSION_CLOSED + == 0 + { + tokio::task::yield_now().await; + } + drain.abort(); + assert!(drain.await.unwrap_err().is_cancelled()); + assert!(generation.session_cancel.is_cancelled()); + drop(registration); + + tokio::time::timeout(Duration::from_secs(1), dropped.notified()) + .await + .expect("aborted graceful drain must cancel existing sessions"); + } +} diff --git a/src/maestro/orchestrator.rs b/src/maestro/orchestrator.rs index fe0344d..b0e832e 100644 --- a/src/maestro/orchestrator.rs +++ b/src/maestro/orchestrator.rs @@ -3,7 +3,7 @@ use std::net::{IpAddr, SocketAddr}; use std::sync::Arc; use arc_swap::ArcSwap; -use tokio::sync::{RwLock, watch}; +use tokio::sync::{RwLock, Semaphore, watch}; use tracing::{error, info}; use crate::api; @@ -12,7 +12,9 @@ use crate::network::probe::{decide_network_capabilities, log_probe_result, run_p use crate::proxy::direct_buffer_budget::{DirectBufferBudget, resolve_direct_buffer_hard_limit}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::ProxySharedState; +use crate::proxy::traffic_limiter::TrafficLimiter; use crate::proxy::user_admission::UserAdmissionAuthority; +use crate::proxy::user_connection_authority::UserConnectionAuthority; use crate::startup::{COMPONENT_API_BOOTSTRAP, COMPONENT_NETWORK_PROBE}; use crate::stats::telemetry::TelemetryPolicy; use crate::stats::{QuotaStore, Stats}; @@ -47,10 +49,16 @@ pub(super) async fn run_telemt_core( } = bootstrap::bootstrap(privilege_drop_requested).await?; let quota_store = Arc::new(QuotaStore::default()); - let stats = Arc::new(Stats::with_quota_store(quota_store.clone())); + let connection_authority = Arc::new(UserConnectionAuthority::default()); + let stats = Arc::new(Stats::with_process_authorities( + quota_store.clone(), + connection_authority, + )); let tls_full_cert_budget = Arc::new(TlsFullCertBudget::new()); let process_control_plane = control_plane::ProcessControlPlane::new(); let runtime_task_scope = generation::RuntimeTaskScope::new(); + let runtime_task_scope_guard = + generation::RuntimeTaskScopePreparationGuard::new(runtime_task_scope.clone()); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry)); let quota_state_path = config.general.quota_state_path.clone(); let quota_state = @@ -72,14 +80,11 @@ pub(super) async fn run_telemt_core( .with_dns_overrides(&config.network.dns_overrides)?, ); let ip_tracker = Arc::new(UserIpTracker::new()); - ip_tracker - .load_limits( + let _ = ip_tracker + .apply_policy_from_source( + 1, config.access.user_max_unique_ips_global_each, &config.access.user_max_unique_ips, - ) - .await; - ip_tracker - .set_limit_policy( config.access.user_max_unique_ips_mode, config.access.user_max_unique_ips_window_secs, ) @@ -102,14 +107,22 @@ pub(super) async fn run_telemt_core( let direct_buffer_hard_limit = resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await; let direct_buffer_budget = DirectBufferBudget::new(direct_buffer_hard_limit); + direct_buffer_budget.activate_controller(1); info!( hard_limit_bytes = direct_buffer_hard_limit, configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes, "Direct relay buffer budget initialized" ); let user_admission = UserAdmissionAuthority::new_with_quota_store(quota_store.clone()); - let shared_state = ProxySharedState::new_with_direct_buffer_budget_and_user_admission( + let traffic_limiter = TrafficLimiter::new(); + let _ = traffic_limiter.apply_policy_from_source( + 1, + config.access.user_rate_limits.clone(), + config.access.cidr_rate_limits.clone(), + ); + let shared_state = ProxySharedState::new_with_process_authorities( direct_buffer_budget.clone(), + traffic_limiter, user_admission, ); let _ = shared_state.activate_user_config_source( @@ -118,10 +131,12 @@ pub(super) async fn run_telemt_core( &config.access.users, &config.access.user_enabled, ); - shared_state.traffic_limiter.apply_policy( - config.access.user_rate_limits.clone(), - config.access.cidr_rate_limits.clone(), - ); + let max_connections_limit = if config.server.max_connections == 0 { + Semaphore::MAX_PERMITS + } else { + config.server.max_connections as usize + }; + let max_connections = Arc::new(Semaphore::new(max_connections_limit)); let web_trace = WebTraceStore::new(config.web.debug.clone(), &config.web.limits); let web_runtime_control = WebRuntimeControl::new(); @@ -303,6 +318,7 @@ pub(super) async fn run_telemt_core( ip_tracker.clone(), shared_state.clone(), direct_buffer_budget, + max_connections, route_runtime.clone(), api_me_pool.clone(), runtime_task_scope.clone(), @@ -333,6 +349,7 @@ pub(super) async fn run_telemt_core( runtime.max_connections, runtime_task_scope, ); + runtime_task_scope_guard.disarm(); let active_runtime = Arc::new(ArcSwap::from(runtime_generation)); let bound = listeners::bind_listeners( &runtime.config, diff --git a/src/maestro/reload_supervisor.rs b/src/maestro/reload_supervisor.rs index 725c329..c1e3998 100644 --- a/src/maestro/reload_supervisor.rs +++ b/src/maestro/reload_supervisor.rs @@ -175,9 +175,14 @@ impl ReloadSupervisor { resolved.effective, &self.config_path, self.quota_store.clone(), + old_runtime.stats.connection_authority(), self.runtime_log_filter.clone(), self.tls_full_cert_budget.clone(), old_runtime.proxy_shared.user_admission(), + old_runtime.ip_tracker.clone(), + old_runtime.proxy_shared.traffic_limiter.clone(), + old_runtime.proxy_shared.direct_buffer_budget.clone(), + old_runtime.max_connections.clone(), ) .await { @@ -300,15 +305,37 @@ impl ReloadSupervisor { } else { None }; + let config = new_runtime.config(); + let _ = new_runtime.proxy_shared.activate_user_config_source( + new_runtime.id, + Some(user_admission_epoch), + &config.access.users, + &config.access.user_enabled, + ); + let _ = new_runtime + .ip_tracker + .apply_policy_from_source( + new_runtime.id, + config.access.user_max_unique_ips_global_each, + &config.access.user_max_unique_ips, + config.access.user_max_unique_ips_mode, + config.access.user_max_unique_ips_window_secs, + ) + .await; + let _ = new_runtime + .proxy_shared + .traffic_limiter + .apply_policy_from_source( + new_runtime.id, + config.access.user_rate_limits.clone(), + config.access.cidr_rate_limits.clone(), + ); + new_runtime + .proxy_shared + .direct_buffer_budget + .activate_controller(new_runtime.id); let replaced = { let listener_manager = self.listener_manager.lock().await; - let config = new_runtime.config(); - let _ = new_runtime.proxy_shared.activate_user_config_source( - new_runtime.id, - Some(user_admission_epoch), - &config.access.users, - &config.access.user_enabled, - ); old_runtime.stop_accepting_sessions(); listener_manager.activate_runtime_generation(new_runtime.clone()) }; diff --git a/src/maestro/runtime_build.rs b/src/maestro/runtime_build.rs index a289896..ea2221a 100644 --- a/src/maestro/runtime_build.rs +++ b/src/maestro/runtime_build.rs @@ -12,11 +12,13 @@ use crate::crypto::SecureRandom; use crate::ip_tracker::UserIpTracker; use crate::network::probe::{decide_network_capabilities, run_probe}; use crate::proxy::direct_buffer_budget::{ - DirectBufferBudget, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller, + DirectBufferBudget, run_direct_buffer_budget_controller, }; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::ProxySharedState; +use crate::proxy::traffic_limiter::TrafficLimiter; use crate::proxy::user_admission::UserAdmissionAuthority; +use crate::proxy::user_connection_authority::UserConnectionAuthority; use crate::startup::StartupTracker; use crate::stats::beobachten::BeobachtenStore; use crate::stats::telemetry::TelemetryPolicy; @@ -27,7 +29,9 @@ use crate::transport::UpstreamManager; use crate::transport::middle_proxy::MePool; use super::admission; -use super::generation::{RuntimeGeneration, RuntimeTaskScope}; +use super::generation::{ + RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard, +}; use super::listeners::listener_rebind_supported; use super::runtime_tasks::RuntimeLogFilter; use super::{me_startup, runtime_tasks, tls_bootstrap}; @@ -49,9 +53,14 @@ pub(crate) async fn prepare_runtime( config: ProxyConfig, config_path: &Path, quota_store: Arc, + connection_authority: Arc, runtime_log_filter: RuntimeLogFilter, tls_full_cert_budget: Arc, user_admission: Arc, + ip_tracker: Arc, + traffic_limiter: Arc, + direct_buffer_budget: Arc, + max_connections: Arc, ) -> Result { let user_admission_epoch = user_admission.epoch(); config @@ -63,7 +72,11 @@ pub(crate) async fn prepare_runtime( .as_secs(); let startup_tracker = Arc::new(StartupTracker::new(started_at_epoch_secs)); let task_scope = RuntimeTaskScope::new(); - let stats = Arc::new(Stats::with_quota_store(quota_store)); + let task_scope_guard = RuntimeTaskScopePreparationGuard::new(task_scope.clone()); + let stats = Arc::new(Stats::with_process_authorities( + quota_store, + connection_authority, + )); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry)); let upstream_manager = Arc::new( @@ -80,31 +93,11 @@ pub(crate) async fn prepare_runtime( .with_dns_overrides(&config.network.dns_overrides) .map_err(|error| format!("DNS override preparation failed: {}", error))?, ); - let ip_tracker = Arc::new(UserIpTracker::new()); - ip_tracker - .load_limits( - config.access.user_max_unique_ips_global_each, - &config.access.user_max_unique_ips, - ) - .await; - ip_tracker - .set_limit_policy( - config.access.user_max_unique_ips_mode, - config.access.user_max_unique_ips_window_secs, - ) - .await; - - let hard_limit = - resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await; - let direct_buffer_budget = DirectBufferBudget::new(hard_limit); - let proxy_shared = ProxySharedState::new_with_direct_buffer_budget_and_user_admission( + let proxy_shared = ProxySharedState::new_with_process_authorities( direct_buffer_budget.clone(), + traffic_limiter, user_admission, ); - proxy_shared.traffic_limiter.apply_policy( - config.access.user_rate_limits.clone(), - config.access.cidr_rate_limits.clone(), - ); let probe = run_probe( &config.network, @@ -184,12 +177,6 @@ pub(crate) async fn prepare_runtime( Duration::from_secs(config.access.replay_window_secs), )); let buffer_pool = Arc::new(BufferPool::with_config(64 * 1024, 4096)); - let max_connections_limit = if config.server.max_connections == 0 { - Semaphore::MAX_PERMITS - } else { - config.server.max_connections as usize - }; - let max_connections = Arc::new(Semaphore::new(max_connections_limit)); let (config_watcher_activation, config_watcher_activation_rx) = watch::channel(false); let watches = runtime_tasks::spawn_runtime_tasks( generation_id, @@ -288,10 +275,12 @@ pub(crate) async fn prepare_runtime( conntrack_scope.cancellation_token(), )); task_scope.spawn(run_direct_buffer_budget_controller( + generation_id, direct_buffer_budget, buffer_pool.clone(), stats.clone(), proxy_shared.clone(), + max_connections.clone(), config.server.max_connections, )); let generation = RuntimeGeneration::new( @@ -313,6 +302,7 @@ pub(crate) async fn prepare_runtime( max_connections, task_scope, ); + task_scope_guard.disarm(); drop(admission_tx); Ok(PreparedRuntime { @@ -414,6 +404,17 @@ pub(crate) fn resolve_reload_config( effective.server.metrics_listen = old.server.metrics_listen.clone(); effective.server.metrics_port = old.server.metrics_port; } + if old.server.max_connections != desired.server.max_connections { + fields.push("server.max_connections".to_string()); + effective.server.max_connections = old.server.max_connections; + } + if old.general.direct_relay_buffer_budget_max_bytes + != desired.general.direct_relay_buffer_budget_max_bytes + { + fields.push("general.direct_relay_buffer_budget_max_bytes".to_string()); + effective.general.direct_relay_buffer_budget_max_bytes = + old.general.direct_relay_buffer_budget_max_bytes; + } if old.general.quota_state_path != desired.general.quota_state_path { fields.push("general.quota_state_path".to_string()); effective.general.quota_state_path = old.general.quota_state_path.clone(); diff --git a/src/maestro/runtime_build_tests.rs b/src/maestro/runtime_build_tests.rs index babaa54..38b7f87 100644 --- a/src/maestro/runtime_build_tests.rs +++ b/src/maestro/runtime_build_tests.rs @@ -94,6 +94,39 @@ fn global_mss_profiles_are_deferred_with_the_listener_socket_group() { assert!(!resolved.runtime_changed); } +#[test] +fn process_wide_connection_and_direct_buffer_envelopes_are_restart_only() { + let old = ProxyConfig::default(); + let mut desired = old.clone(); + desired.server.max_connections = old.server.max_connections.saturating_add(1); + desired.general.direct_relay_buffer_budget_max_bytes = old + .general + .direct_relay_buffer_budget_max_bytes + .saturating_add(4 * 1024); + + let resolved = resolve_reload_config(&old, &desired).unwrap(); + + assert_eq!( + resolved.deferred_process_fields, + vec![ + "server.max_connections".to_string(), + "general.direct_relay_buffer_budget_max_bytes".to_string(), + ] + ); + assert_eq!( + resolved.effective.server.max_connections, + old.server.max_connections + ); + assert_eq!( + resolved + .effective + .general + .direct_relay_buffer_budget_max_bytes, + old.general.direct_relay_buffer_budget_max_bytes + ); + assert!(!resolved.runtime_changed); +} + #[test] fn mixed_reload_retains_process_state_and_applies_runtime_state() { let old = ProxyConfig::default(); diff --git a/src/maestro/runtime_startup.rs b/src/maestro/runtime_startup.rs index e7193d1..d580fb0 100644 --- a/src/maestro/runtime_startup.rs +++ b/src/maestro/runtime_startup.rs @@ -56,6 +56,7 @@ pub(super) async fn prepare_runtime( ip_tracker: Arc, shared_state: Arc, direct_buffer_budget: Arc, + max_connections: Arc, route_runtime: Arc, api_me_pool: Arc>>>, runtime_task_scope: RuntimeTaskScope, @@ -69,13 +70,6 @@ pub(super) async fn prepare_runtime( let beobachten = Arc::new(BeobachtenStore::new()); let rng = Arc::new(SecureRandom::new()); - let max_connections_limit = if config.server.max_connections == 0 { - Semaphore::MAX_PERMITS - } else { - config.server.max_connections as usize - }; - let max_connections = Arc::new(Semaphore::new(max_connections_limit)); - let me2dc_fallback = config.general.me2dc_fallback; let me_init_retry_attempts = config.general.me_init_retry_attempts; if use_middle_proxy && !decision.ipv4_me && !decision.ipv6_me { @@ -346,10 +340,12 @@ pub(super) async fn prepare_runtime( conntrack_scope.cancellation_token(), )); runtime_task_scope.spawn(run_direct_buffer_budget_controller( + 1, direct_buffer_budget, buffer_pool.clone(), stats, shared_state, + max_connections.clone(), config.server.max_connections, )); diff --git a/src/maestro/runtime_tasks.rs b/src/maestro/runtime_tasks.rs index e46e604..891e6d6 100644 --- a/src/maestro/runtime_tasks.rs +++ b/src/maestro/runtime_tasks.rs @@ -138,7 +138,9 @@ pub(crate) async fn spawn_runtime_tasks( let ip_tracker_maintenance = ip_tracker.clone(); task_scope.spawn(async move { - ip_tracker_maintenance.run_periodic_maintenance().await; + ip_tracker_maintenance + .run_periodic_maintenance(generation_id) + .await; }); let detected_ip_v4: Option = probe.detected_ipv4.map(IpAddr::V4); @@ -197,20 +199,7 @@ pub(crate) async fn spawn_runtime_tasks( let ip_tracker_policy = ip_tracker.clone(); let mut config_rx_ip_limits = config_rx.clone(); task_scope.spawn(async move { - let mut prev_limits = config_rx_ip_limits - .borrow() - .access - .user_max_unique_ips - .clone(); - let mut prev_global_each = config_rx_ip_limits - .borrow() - .access - .user_max_unique_ips_global_each; - let mut prev_mode = config_rx_ip_limits.borrow().access.user_max_unique_ips_mode; - let mut prev_window = config_rx_ip_limits - .borrow() - .access - .user_max_unique_ips_window_secs; + let mut previous = config_rx_ip_limits.borrow().access.clone(); loop { if config_rx_ip_limits.changed().await.is_err() { @@ -218,39 +207,28 @@ pub(crate) async fn spawn_runtime_tasks( } let cfg = config_rx_ip_limits.borrow_and_update().clone(); - if prev_limits != cfg.access.user_max_unique_ips - || prev_global_each != cfg.access.user_max_unique_ips_global_each + if previous.user_max_unique_ips != cfg.access.user_max_unique_ips + || previous.user_max_unique_ips_global_each + != cfg.access.user_max_unique_ips_global_each + || previous.user_max_unique_ips_mode != cfg.access.user_max_unique_ips_mode + || previous.user_max_unique_ips_window_secs + != cfg.access.user_max_unique_ips_window_secs { - ip_tracker_policy - .load_limits( + let _ = ip_tracker_policy + .apply_policy_from_source( + generation_id, cfg.access.user_max_unique_ips_global_each, &cfg.access.user_max_unique_ips, - ) - .await; - prev_limits = cfg.access.user_max_unique_ips.clone(); - prev_global_each = cfg.access.user_max_unique_ips_global_each; - } - - if prev_mode != cfg.access.user_max_unique_ips_mode - || prev_window != cfg.access.user_max_unique_ips_window_secs - { - ip_tracker_policy - .set_limit_policy( cfg.access.user_max_unique_ips_mode, cfg.access.user_max_unique_ips_window_secs, ) .await; - prev_mode = cfg.access.user_max_unique_ips_mode; - prev_window = cfg.access.user_max_unique_ips_window_secs; + previous = cfg.access.clone(); } } }); let limiter = shared_state.traffic_limiter.clone(); - limiter.apply_policy( - config.access.user_rate_limits.clone(), - config.access.cidr_rate_limits.clone(), - ); let mut config_rx_rate_limits = config_rx.clone(); task_scope.spawn(async move { let mut prev_user_limits = config_rx_rate_limits @@ -271,7 +249,8 @@ pub(crate) async fn spawn_runtime_tasks( if prev_user_limits != cfg.access.user_rate_limits || prev_cidr_limits != cfg.access.cidr_rate_limits { - limiter.apply_policy( + let _ = limiter.apply_policy_from_source( + generation_id, cfg.access.user_rate_limits.clone(), cfg.access.cidr_rate_limits.clone(), ); diff --git a/src/metrics/render/users.rs b/src/metrics/render/users.rs index bdfd5e8..87897e8 100644 --- a/src/metrics/render/users.rs +++ b/src/metrics/render/users.rs @@ -153,7 +153,7 @@ pub(super) async fn render( out, "telemt_user_connections_current{{user=\"{}\"}} {}", user, - s.curr_connects.load(std::sync::atomic::Ordering::Relaxed) + stats.get_process_user_curr_connects(user) ); let _ = writeln!( out, diff --git a/src/metrics/tests.rs b/src/metrics/tests.rs index f157e3b..9d9d830 100644 --- a/src/metrics/tests.rs +++ b/src/metrics/tests.rs @@ -69,6 +69,10 @@ async fn test_render_metrics_format() { stats.increment_me_endpoint_quarantine_draining_suppressed_total(); stats.increment_user_connects("alice"); stats.increment_user_curr_connects("alice"); + let _connection_permit = stats + .connection_authority() + .try_acquire("alice", None) + .unwrap(); stats.add_user_octets_from("alice", 1024); stats.add_user_octets_to("alice", 2048); stats.increment_user_msgs_from("alice"); diff --git a/src/proxy/authenticated.rs b/src/proxy/authenticated.rs index 571f9a1..b85fd14 100644 --- a/src/proxy/authenticated.rs +++ b/src/proxy/authenticated.rs @@ -15,7 +15,8 @@ use crate::proxy::middle_relay::{handle_via_middle_proxy, handle_via_middle_prox use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState}; use crate::proxy::user_admission::UserIncarnation; -use crate::stats::{Stats, UserQuotaHandle}; +use crate::proxy::user_connection_authority::UserConnectionPermit; +use crate::stats::{Stats, UserConnectionObservation, UserQuotaHandle}; use crate::stream::{BufferPool, CryptoReader, CryptoWriter}; use crate::transport::UpstreamManager; use crate::transport::middle_proxy::MePool; @@ -238,13 +239,63 @@ where /// Owns one authenticated user's connection and source-IP admission slots. pub(crate) struct UserConnectionReservation { stats: Arc, - ip_tracker: Arc, - user: String, - ip: IpAddr, - incarnation: UserIncarnation, quota_handle: UserQuotaHandle, - tracks_ip: bool, - active: bool, + _connection_permit: UserConnectionPermit, + _stats_observation: Option, + ip_permit: Option, + released: bool, +} + +struct UserIpPermit { + tracker: Arc, + owner: Option, +} + +struct UserIpOwner { + user: String, + incarnation: UserIncarnation, + ip: IpAddr, +} + +impl UserIpPermit { + fn new( + tracker: Arc, + user: String, + incarnation: UserIncarnation, + ip: IpAddr, + ) -> Self { + Self { + tracker, + owner: Some(UserIpOwner { + user, + incarnation, + ip, + }), + } + } + + async fn release(mut self) { + let Some(owner) = self.owner.as_ref() else { + return; + }; + self.tracker + .remove_ip_for_incarnation(&owner.user, owner.incarnation, owner.ip) + .await; + self.owner = None; + } +} + +impl Drop for UserIpPermit { + fn drop(&mut self) { + let Some(owner) = self.owner.take() else { + return; + }; + self.tracker.enqueue_cleanup_for_incarnation( + owner.user, + owner.incarnation, + owner.ip, + ); + } } impl UserConnectionReservation { @@ -257,6 +308,11 @@ impl UserConnectionReservation { tracks_ip: bool, ) -> Self { let quota_handle = stats.current_user_quota_handle(&user); + let connection_permit = stats + .connection_authority() + .try_acquire(&user, None) + .expect("unlimited test connection permit must be available"); + let stats_observation = stats.observe_user_current_connection(&user); Self::new_for_incarnation( stats, ip_tracker, @@ -264,6 +320,8 @@ impl UserConnectionReservation { ip, 0, quota_handle, + connection_permit, + stats_observation, tracks_ip, ) } @@ -276,17 +334,20 @@ impl UserConnectionReservation { ip: IpAddr, incarnation: UserIncarnation, quota_handle: UserQuotaHandle, + connection_permit: UserConnectionPermit, + stats_observation: Option, tracks_ip: bool, ) -> Self { + let ip_permit = tracks_ip.then(|| { + UserIpPermit::new(ip_tracker, user, incarnation, ip) + }); Self { stats, - ip_tracker, - user, - ip, - incarnation, quota_handle, - tracks_ip, - active: true, + _connection_permit: connection_permit, + _stats_observation: stats_observation, + ip_permit, + released: false, } } @@ -297,50 +358,24 @@ impl UserConnectionReservation { /// Releases both admission counters through the asynchronous cleanup path. pub(crate) async fn release(mut self) { - if !self.active { - return; + if let Some(ip_permit) = self.ip_permit.take() { + ip_permit.release().await; } - self.active = false; - if self.tracks_ip { - self.ip_tracker - .remove_ip_for_incarnation(&self.user, self.incarnation, self.ip) - .await; - } - self.stats.decrement_user_curr_connects(&self.user); + self.released = true; } /// Defers IP cleanup when admission fails after the asynchronous reservation step. pub(crate) fn release_deferred(mut self) { - if !self.active { - return; - } - self.active = false; - self.stats.decrement_user_curr_connects(&self.user); - if self.tracks_ip { - self.ip_tracker.enqueue_cleanup_for_incarnation( - self.user.clone(), - self.incarnation, - self.ip, - ); - } + self.released = true; } } impl Drop for UserConnectionReservation { fn drop(&mut self) { - if !self.active { + if self.released { return; } - self.active = false; self.stats.increment_session_drop_fallback_total(); - self.stats.decrement_user_curr_connects(&self.user); - if self.tracks_ip { - self.ip_tracker.enqueue_cleanup_for_incarnation( - self.user.clone(), - self.incarnation, - self.ip, - ); - } } } @@ -400,17 +435,20 @@ async fn acquire_user_connection_reservation_for_incarnation( .or((config.access.user_max_tcp_conns_global_each > 0) .then_some(config.access.user_max_tcp_conns_global_each)) .map(|value| value as u64); - if !stats.try_acquire_user_curr_connects(user, limit) { + let Some(connection_permit) = stats + .connection_authority() + .try_acquire(user, limit) + else { return Err(ProxyError::ConnectionLimitExceeded { user: user.to_string(), }); - } + }; + let stats_observation = stats.observe_user_current_connection(user); if let Err(reason) = ip_tracker .check_and_add_for_incarnation(user, incarnation, peer_addr.ip()) .await { - stats.decrement_user_curr_connects(user); warn!( user = %user, ip = %peer_addr.ip(), @@ -429,6 +467,8 @@ async fn acquire_user_connection_reservation_for_incarnation( peer_addr.ip(), incarnation, quota_handle, + connection_permit, + stats_observation, true, )) } diff --git a/src/proxy/client/authenticated.rs b/src/proxy/client/authenticated.rs index b6a0f94..9b12cc1 100644 --- a/src/proxy/client/authenticated.rs +++ b/src/proxy/client/authenticated.rs @@ -155,18 +155,20 @@ impl RunningClientHandler { .or((config.access.user_max_tcp_conns_global_each > 0) .then_some(config.access.user_max_tcp_conns_global_each)) .map(|v| v as u64); - if !stats.try_acquire_user_curr_connects(user, limit) { + let Some(_connection_permit) = stats + .connection_authority() + .try_acquire(user, limit) + else { return Err(ProxyError::ConnectionLimitExceeded { user: user.to_string(), }); - } + }; match ip_tracker.check_and_add(user, peer_addr.ip()).await { Ok(()) => { ip_tracker.remove_ip(user, peer_addr.ip()).await; } Err(reason) => { - stats.decrement_user_curr_connects(user); warn!( user = %user, ip = %peer_addr.ip(), @@ -178,8 +180,6 @@ impl RunningClientHandler { }); } } - - stats.decrement_user_curr_connects(user); Ok(()) } } diff --git a/src/proxy/direct_buffer_budget.rs b/src/proxy/direct_buffer_budget.rs index 0242294..8eb6b14 100644 --- a/src/proxy/direct_buffer_budget.rs +++ b/src/proxy/direct_buffer_budget.rs @@ -2,12 +2,16 @@ use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::Duration; +use parking_lot::{Mutex as ParkingMutex, MutexGuard as ParkingMutexGuard}; use tokio::sync::watch; -use crate::stats::Stats; -use crate::stream::BufferPool; - -use super::shared_state::ProxySharedState; +// Process controller and system-memory sampling remain outside data-plane accounting. +mod controller; +pub(crate) use controller::{ + resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller, +}; +#[cfg(test)] +use controller::connection_fill_pct; /// Accounting granularity for process-wide Direct copy-buffer reservations. pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024; @@ -70,6 +74,8 @@ pub(crate) struct DirectBufferBudget { hard_limit_bytes: u64, target_bytes: AtomicU64, reserved_bytes: AtomicU64, + active_controller_generation: AtomicU64, + controller_update: ParkingMutex<()>, pressure_generation: AtomicU64, pressure_tx: watch::Sender, memory_total_bytes: AtomicU64, @@ -94,6 +100,8 @@ impl DirectBufferBudget { hard_limit_bytes, target_bytes: AtomicU64::new(hard_limit_bytes), reserved_bytes: AtomicU64::new(0), + active_controller_generation: AtomicU64::new(0), + controller_update: ParkingMutex::new(()), pressure_generation: AtomicU64::new(0), pressure_tx, memory_total_bytes: AtomicU64::new(0), @@ -120,6 +128,22 @@ impl DirectBufferBudget { self.pressure_tx.subscribe() } + /// Transfers adaptive-target writes to the active runtime generation. + pub(crate) fn activate_controller(&self, generation: u64) { + let _controller_update = self.controller_update.lock(); + self.active_controller_generation + .fetch_max(generation, Ordering::AcqRel); + } + + fn begin_controller_update( + &self, + generation: u64, + ) -> Option> { + let controller_update = self.controller_update.lock(); + (self.active_controller_generation.load(Ordering::Acquire) == generation) + .then_some(controller_update) + } + /// Reserves bytes against either the adaptive target or the absolute ceiling. pub(crate) fn try_reserve( self: &Arc, @@ -322,222 +346,6 @@ impl Drop for DirectBufferLease { } } -/// Resolves the startup hard ceiling from config, cgroup, and host memory. -pub(crate) async fn resolve_direct_buffer_hard_limit(configured: usize) -> usize { - if configured != 0 { - return align_down(configured); - } - let sample = read_system_memory_sample().await; - if sample.total_bytes == 0 { - return AUTO_HARD_FALLBACK_BYTES; - } - let derived = (sample.total_bytes / 4) - .clamp(AUTO_HARD_MIN_BYTES as u64, AUTO_HARD_MAX_BYTES as u64) - .min(sample.total_bytes); - align_down(derived as usize).max(DIRECT_BUFFER_UNIT_BYTES) -} - -/// Runs the control-plane loop for Direct budget and shared pool pressure. -pub(crate) async fn run_direct_buffer_budget_controller( - budget: Arc, - buffer_pool: Arc, - stats: Arc, - shared: Arc, - max_connections: u32, -) { - let mut interval = tokio::time::interval(CONTROL_INTERVAL); - interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); - let mut healthy_streak = 0u8; - let mut previous_denied = 0u64; - let mut previous_fallback = 0u64; - let mut previous_rejected = 0u64; - let pool_trim_low = buffer_pool - .max_buffers() - .min(BUFFER_POOL_TRIM_LOW_WATERMARK); - let pool_trim_high = buffer_pool - .max_buffers() - .min(BUFFER_POOL_TRIM_HIGH_WATERMARK); - let mut pool_trim_armed = true; - - loop { - interval.tick().await; - let sample = read_system_memory_sample().await; - budget.update_system_sample(sample); - - let snapshot = budget.snapshot(); - let denied_delta = snapshot - .promotion_denied_total - .saturating_sub(previous_denied); - previous_denied = snapshot.promotion_denied_total; - let fallback_delta = snapshot - .minimum_fallback_total - .saturating_sub(previous_fallback); - previous_fallback = snapshot.minimum_fallback_total; - let rejected_delta = snapshot - .admission_rejected_total - .saturating_sub(previous_rejected); - previous_rejected = snapshot.admission_rejected_total; - - let connection_pct = connection_fill_pct(stats.as_ref(), max_connections); - let memory_available_pct = percentage(sample.available_bytes, sample.total_bytes); - let target_utilization_pct = percentage(snapshot.reserved_bytes, snapshot.target_bytes); - let pressure = shared.conntrack_pressure_active() - || connection_pct.is_some_and(|value| value >= 85) - || memory_available_pct.is_some_and(|value| value <= 15) - || target_utilization_pct.is_some_and(|value| value >= 90) - || denied_delta > 0 - || fallback_delta > 0 - || rejected_delta > 0; - - if !pressure { - pool_trim_armed = true; - } else if pool_trim_armed && buffer_pool.pooled() > pool_trim_high { - buffer_pool.trim_to(pool_trim_low); - pool_trim_armed = false; - } - - let pool_snapshot = buffer_pool.stats(); - stats.set_buffer_pool_gauges( - pool_snapshot.pooled, - pool_snapshot.allocated, - pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled), - ); - stats.set_buffer_pool_replaced_nonstandard_total(pool_snapshot.replaced_nonstandard); - - let headroom_target = if sample.total_bytes == 0 { - snapshot.hard_limit_bytes - } else { - snapshot - .reserved_bytes - .saturating_add(sample.available_bytes / 4) - .min(snapshot.hard_limit_bytes) - }; - - if pressure { - healthy_streak = 0; - let reduced = snapshot.target_bytes.saturating_mul(3) / 4; - budget.set_target_bytes(reduced.min(headroom_target)); - continue; - } - - let healthy = memory_available_pct.is_none_or(|value| value >= 30) - && connection_pct.is_none_or(|value| value <= 70); - if !healthy { - healthy_streak = 0; - if headroom_target < snapshot.target_bytes { - budget.set_target_bytes(headroom_target); - } - continue; - } - - healthy_streak = healthy_streak.saturating_add(1); - if healthy_streak >= HEALTHY_RECOVERY_SAMPLES { - healthy_streak = 0; - let increment = (snapshot.target_bytes / 16).max(4 * 1024 * 1024); - budget.set_target_bytes( - snapshot - .target_bytes - .saturating_add(increment) - .min(headroom_target), - ); - } - } -} - -fn connection_fill_pct(stats: &Stats, max_connections: u32) -> Option { - if max_connections == 0 { - return None; - } - Some( - ((stats.get_current_connections_total().saturating_mul(100)) / u64::from(max_connections)) - .min(100) as u8, - ) -} - -fn percentage(value: u64, total: u64) -> Option { - if total == 0 { - return None; - } - Some(((value.saturating_mul(100)) / total).min(100) as u8) -} - -async fn read_system_memory_sample() -> SystemMemorySample { - #[cfg(target_os = "linux")] - { - let meminfo = tokio::fs::read_to_string("/proc/meminfo") - .await - .unwrap_or_default(); - let status = tokio::fs::read_to_string("/proc/self/status") - .await - .unwrap_or_default(); - let host_total = parse_kib_field(&meminfo, "MemTotal:"); - let host_available = parse_kib_field(&meminfo, "MemAvailable:"); - let process_rss = parse_kib_field(&status, "VmRSS:"); - - let cgroup_v2_max = read_cgroup_limit("/sys/fs/cgroup/memory.max").await; - let cgroup_v2_current = read_u64_file("/sys/fs/cgroup/memory.current").await; - let cgroup_v1_max = read_cgroup_limit("/sys/fs/cgroup/memory/memory.limit_in_bytes").await; - let cgroup_v1_current = read_u64_file("/sys/fs/cgroup/memory/memory.usage_in_bytes").await; - let cgroup_max = cgroup_v2_max.or(cgroup_v1_max); - let cgroup_current = cgroup_v2_current.or(cgroup_v1_current); - - let total = match (host_total, cgroup_max) { - (0, Some(limit)) => limit, - (host, Some(limit)) => host.min(limit), - (host, None) => host, - }; - let cgroup_available = cgroup_max - .zip(cgroup_current) - .map(|(limit, current)| limit.saturating_sub(current)); - let available = match (host_available, cgroup_available) { - (0, Some(value)) => value, - (host, Some(value)) => host.min(value), - (host, None) => host, - }; - return SystemMemorySample { - total_bytes: total, - available_bytes: available, - process_rss_bytes: process_rss, - }; - } - #[cfg(not(target_os = "linux"))] - { - SystemMemorySample::default() - } -} - -#[cfg(target_os = "linux")] -async fn read_cgroup_limit(path: &str) -> Option { - let raw = tokio::fs::read_to_string(path).await.ok()?; - let raw = raw.trim(); - if raw == "max" { - return None; - } - let value = raw.parse::().ok()?; - (value < (1u64 << 60)).then_some(value) -} - -#[cfg(target_os = "linux")] -async fn read_u64_file(path: &str) -> Option { - tokio::fs::read_to_string(path) - .await - .ok()? - .trim() - .parse() - .ok() -} - -#[cfg(target_os = "linux")] -fn parse_kib_field(raw: &str, key: &str) -> u64 { - raw.lines() - .find_map(|line| { - let value = line.strip_prefix(key)?.split_whitespace().next()?; - value.parse::().ok() - }) - .unwrap_or(0) - .saturating_mul(1024) -} - fn align_up(bytes: usize) -> usize { bytes .div_ceil(DIRECT_BUFFER_UNIT_BYTES) diff --git a/src/proxy/direct_buffer_budget/controller.rs b/src/proxy/direct_buffer_budget/controller.rs new file mode 100644 index 0000000..c7fe8fe --- /dev/null +++ b/src/proxy/direct_buffer_budget/controller.rs @@ -0,0 +1,238 @@ +use std::sync::Arc; + +use tokio::sync::Semaphore; + +use super::*; +use crate::proxy::shared_state::ProxySharedState; +use crate::stats::Stats; +use crate::stream::BufferPool; + +/// Resolves the startup hard ceiling from config, cgroup, and host memory. +pub(crate) async fn resolve_direct_buffer_hard_limit(configured: usize) -> usize { + if configured != 0 { + return align_down(configured); + } + let sample = read_system_memory_sample().await; + if sample.total_bytes == 0 { + return AUTO_HARD_FALLBACK_BYTES; + } + let derived = (sample.total_bytes / 4) + .clamp(AUTO_HARD_MIN_BYTES as u64, AUTO_HARD_MAX_BYTES as u64) + .min(sample.total_bytes); + align_down(derived as usize).max(DIRECT_BUFFER_UNIT_BYTES) +} + +/// Runs the control-plane loop for Direct budget and shared pool pressure. +pub(crate) async fn run_direct_buffer_budget_controller( + source_generation: u64, + budget: Arc, + buffer_pool: Arc, + stats: Arc, + shared: Arc, + connection_slots: Arc, + max_connections: u32, +) { + let mut interval = tokio::time::interval(CONTROL_INTERVAL); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut healthy_streak = 0u8; + let mut previous_denied = 0u64; + let mut previous_fallback = 0u64; + let mut previous_rejected = 0u64; + let pool_trim_low = buffer_pool + .max_buffers() + .min(BUFFER_POOL_TRIM_LOW_WATERMARK); + let pool_trim_high = buffer_pool + .max_buffers() + .min(BUFFER_POOL_TRIM_HIGH_WATERMARK); + let mut pool_trim_armed = true; + + loop { + interval.tick().await; + if budget.active_controller_generation.load(Ordering::Acquire) != source_generation { + continue; + } + let sample = read_system_memory_sample().await; + let Some(_controller_update) = budget.begin_controller_update(source_generation) else { + continue; + }; + budget.update_system_sample(sample); + + let snapshot = budget.snapshot(); + let denied_delta = snapshot + .promotion_denied_total + .saturating_sub(previous_denied); + previous_denied = snapshot.promotion_denied_total; + let fallback_delta = snapshot + .minimum_fallback_total + .saturating_sub(previous_fallback); + previous_fallback = snapshot.minimum_fallback_total; + let rejected_delta = snapshot + .admission_rejected_total + .saturating_sub(previous_rejected); + previous_rejected = snapshot.admission_rejected_total; + + let connection_pct = connection_fill_pct(connection_slots.as_ref(), max_connections); + let memory_available_pct = percentage(sample.available_bytes, sample.total_bytes); + let target_utilization_pct = percentage(snapshot.reserved_bytes, snapshot.target_bytes); + let pressure = shared.conntrack_pressure_active() + || connection_pct.is_some_and(|value| value >= 85) + || memory_available_pct.is_some_and(|value| value <= 15) + || target_utilization_pct.is_some_and(|value| value >= 90) + || denied_delta > 0 + || fallback_delta > 0 + || rejected_delta > 0; + + if !pressure { + pool_trim_armed = true; + } else if pool_trim_armed && buffer_pool.pooled() > pool_trim_high { + buffer_pool.trim_to(pool_trim_low); + pool_trim_armed = false; + } + + let pool_snapshot = buffer_pool.stats(); + stats.set_buffer_pool_gauges( + pool_snapshot.pooled, + pool_snapshot.allocated, + pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled), + ); + stats.set_buffer_pool_replaced_nonstandard_total(pool_snapshot.replaced_nonstandard); + + let headroom_target = if sample.total_bytes == 0 { + snapshot.hard_limit_bytes + } else { + snapshot + .reserved_bytes + .saturating_add(sample.available_bytes / 4) + .min(snapshot.hard_limit_bytes) + }; + + if pressure { + healthy_streak = 0; + let reduced = snapshot.target_bytes.saturating_mul(3) / 4; + budget.set_target_bytes(reduced.min(headroom_target)); + continue; + } + + let healthy = memory_available_pct.is_none_or(|value| value >= 30) + && connection_pct.is_none_or(|value| value <= 70); + if !healthy { + healthy_streak = 0; + if headroom_target < snapshot.target_bytes { + budget.set_target_bytes(headroom_target); + } + continue; + } + + healthy_streak = healthy_streak.saturating_add(1); + if healthy_streak >= HEALTHY_RECOVERY_SAMPLES { + healthy_streak = 0; + let increment = (snapshot.target_bytes / 16).max(4 * 1024 * 1024); + budget.set_target_bytes( + snapshot + .target_bytes + .saturating_add(increment) + .min(headroom_target), + ); + } + } +} + +pub(super) fn connection_fill_pct( + connection_slots: &Semaphore, + max_connections: u32, +) -> Option { + if max_connections == 0 { + return None; + } + let max_connections = max_connections as usize; + let active = max_connections.saturating_sub( + connection_slots + .available_permits() + .min(max_connections), + ); + Some((active.saturating_mul(100) / max_connections).min(100) as u8) +} + +fn percentage(value: u64, total: u64) -> Option { + if total == 0 { + return None; + } + Some(((value.saturating_mul(100)) / total).min(100) as u8) +} + +async fn read_system_memory_sample() -> SystemMemorySample { + #[cfg(target_os = "linux")] + { + let meminfo = tokio::fs::read_to_string("/proc/meminfo") + .await + .unwrap_or_default(); + let status = tokio::fs::read_to_string("/proc/self/status") + .await + .unwrap_or_default(); + let host_total = parse_kib_field(&meminfo, "MemTotal:"); + let host_available = parse_kib_field(&meminfo, "MemAvailable:"); + let process_rss = parse_kib_field(&status, "VmRSS:"); + + let cgroup_v2_max = read_cgroup_limit("/sys/fs/cgroup/memory.max").await; + let cgroup_v2_current = read_u64_file("/sys/fs/cgroup/memory.current").await; + let cgroup_v1_max = read_cgroup_limit("/sys/fs/cgroup/memory/memory.limit_in_bytes").await; + let cgroup_v1_current = read_u64_file("/sys/fs/cgroup/memory/memory.usage_in_bytes").await; + let cgroup_max = cgroup_v2_max.or(cgroup_v1_max); + let cgroup_current = cgroup_v2_current.or(cgroup_v1_current); + + let total = match (host_total, cgroup_max) { + (0, Some(limit)) => limit, + (host, Some(limit)) => host.min(limit), + (host, None) => host, + }; + let cgroup_available = cgroup_max + .zip(cgroup_current) + .map(|(limit, current)| limit.saturating_sub(current)); + let available = match (host_available, cgroup_available) { + (0, Some(value)) => value, + (host, Some(value)) => host.min(value), + (host, None) => host, + }; + return SystemMemorySample { + total_bytes: total, + available_bytes: available, + process_rss_bytes: process_rss, + }; + } + #[cfg(not(target_os = "linux"))] + { + SystemMemorySample::default() + } +} + +#[cfg(target_os = "linux")] +async fn read_cgroup_limit(path: &str) -> Option { + let raw = tokio::fs::read_to_string(path).await.ok()?; + let raw = raw.trim(); + if raw == "max" { + return None; + } + let value = raw.parse::().ok()?; + (value < (1u64 << 60)).then_some(value) +} + +#[cfg(target_os = "linux")] +async fn read_u64_file(path: &str) -> Option { + tokio::fs::read_to_string(path) + .await + .ok()? + .trim() + .parse() + .ok() +} + +#[cfg(target_os = "linux")] +fn parse_kib_field(raw: &str, key: &str) -> u64 { + raw.lines() + .find_map(|line| { + let value = line.strip_prefix(key)?.split_whitespace().next()?; + value.parse::().ok() + }) + .unwrap_or(0) + .saturating_mul(1024) +} diff --git a/src/proxy/middle_relay.rs b/src/proxy/middle_relay.rs index 9db3730..56a4a5a 100644 --- a/src/proxy/middle_relay.rs +++ b/src/proxy/middle_relay.rs @@ -15,6 +15,7 @@ use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch}; use tokio::time::timeout; use tokio_util::sync::CancellationToken; +use tokio_util::task::AbortOnDropHandle; use tracing::{debug, info, trace, warn}; use crate::config::{ConntrackPressureProfile, ProxyConfig}; @@ -155,7 +156,7 @@ const ME_D2C_FLUSH_BATCH_MAX_FRAMES_MIN: usize = 1; const ME_D2C_FLUSH_BATCH_MAX_BYTES_MIN: usize = 4096; const ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR: usize = 2; const ME_D2C_SINGLE_WRITE_COALESCE_MAX_BYTES: usize = 128 * 1024; -const QUOTA_RESERVE_SPIN_RETRIES: usize = 32; +const QUOTA_RESERVE_ATTEMPTS_PER_ROUND: usize = 4; const QUOTA_RESERVE_BACKOFF_MIN_MS: u64 = 1; const QUOTA_RESERVE_BACKOFF_MAX_MS: u64 = 16; const QUOTA_RESERVE_MAX_BACKOFF_ROUNDS: usize = 16; diff --git a/src/proxy/middle_relay/quota.rs b/src/proxy/middle_relay/quota.rs index 1484f2b..e7a6eb4 100644 --- a/src/proxy/middle_relay/quota.rs +++ b/src/proxy/middle_relay/quota.rs @@ -22,7 +22,7 @@ pub(super) async fn reserve_user_quota_with_yield( let mut backoff_ms = QUOTA_RESERVE_BACKOFF_MIN_MS; let mut backoff_rounds = 0usize; loop { - for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { + for _ in 0..QUOTA_RESERVE_ATTEMPTS_PER_ROUND { match quota_handle.try_reserve(bytes, limit) { Ok(reservation) => return Ok(reservation.commit()), Err(QuotaReserveError::LimitExceeded) => { @@ -30,7 +30,6 @@ pub(super) async fn reserve_user_quota_with_yield( } Err(QuotaReserveError::Contended) => { stats.increment_quota_contention_total(); - std::hint::spin_loop(); } } } diff --git a/src/proxy/middle_relay/session.rs b/src/proxy/middle_relay/session.rs index 89fc26e..3531cb7 100644 --- a/src/proxy/middle_relay/session.rs +++ b/src/proxy/middle_relay/session.rs @@ -2,11 +2,15 @@ use super::*; // Bounded C2ME sender and downstream writer tasks. mod tasks; +// Child-task ownership aborts relay tasks when the parent future is cancelled. +mod children; // Conntrack close classification. mod close_reason; +use children::RelayChildTasks; use close_reason::classify_conntrack_close_reason; use tasks::{run_c2me_sender, run_me_writer}; + struct RelayConnLease { connection: Option, conn_id: u64, @@ -185,7 +189,7 @@ where let c2me_byte_semaphore = Arc::new(Semaphore::new(c2me_byte_budget)); let (c2me_tx, c2me_rx) = mpsc::channel::(c2me_channel_capacity); let me_pool_c2me = me_pool.clone(); - let mut c2me_sender = tokio::spawn(run_c2me_sender( + let c2me_sender = AbortOnDropHandle::new(tokio::spawn(run_c2me_sender( c2me_rx, me_pool_c2me, conn_id, @@ -193,7 +197,7 @@ where peer, translated_local_addr, effective_tag_array, - )); + ))); let (stop_tx, stop_rx) = oneshot::channel::<()>(); let flow_cancel = CancellationToken::new(); @@ -208,7 +212,7 @@ where let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone(); let bytes_me2c_clone = bytes_me2c.clone(); let d2c_flush_policy = MeD2cFlushPolicy::from_config(&config); - let mut me_writer = tokio::spawn(run_me_writer( + let me_writer = AbortOnDropHandle::new(tokio::spawn(run_me_writer( crypto_writer, me_rx_task, stats_clone, @@ -226,7 +230,13 @@ where session_started_at, conn_id, stop_rx, - )); + ))); + let mut child_tasks = RelayChildTasks { + c2me_sender, + me_writer, + flow_cancel: flow_cancel.clone(), + stop_tx: Some(stop_tx), + }; let mut main_result: Result<()> = Ok(()); let mut client_closed = false; @@ -454,28 +464,30 @@ where } drop(c2me_tx); - let c2me_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut c2me_sender).await { + let c2me_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut child_tasks.c2me_sender).await { Ok(joined) => { joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME sender join error: {e}")))) } Err(_) => { stats.increment_me_child_join_timeout_total(); stats.increment_me_child_abort_total(); - c2me_sender.abort(); + child_tasks.c2me_sender.abort(); Err(ProxyError::Proxy("ME sender join timeout".into())) } }; flow_cancel.cancel(); - let _ = stop_tx.send(()); - let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut me_writer).await { + if let Some(stop_tx) = child_tasks.stop_tx.take() { + let _ = stop_tx.send(()); + } + let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut child_tasks.me_writer).await { Ok(joined) => { joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME writer join error: {e}")))) } Err(_) => { stats.increment_me_child_join_timeout_total(); stats.increment_me_child_abort_total(); - me_writer.abort(); + child_tasks.me_writer.abort(); Err(ProxyError::Proxy("ME writer join timeout".into())) } }; diff --git a/src/proxy/middle_relay/session/children.rs b/src/proxy/middle_relay/session/children.rs new file mode 100644 index 0000000..16930a1 --- /dev/null +++ b/src/proxy/middle_relay/session/children.rs @@ -0,0 +1,65 @@ +use super::*; + +pub(super) struct RelayChildTasks { + pub(super) c2me_sender: AbortOnDropHandle>, + pub(super) me_writer: AbortOnDropHandle>, + pub(super) flow_cancel: CancellationToken, + pub(super) stop_tx: Option>, +} + +impl Drop for RelayChildTasks { + fn drop(&mut self) { + self.flow_cancel.cancel(); + if let Some(stop_tx) = self.stop_tx.take() { + let _ = stop_tx.send(()); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + struct DropSignal(Arc); + + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::AcqRel); + } + } + + async fn pending_child(signal: DropSignal) -> Result<()> { + let _signal = signal; + std::future::pending::<()>().await; + Ok(()) + } + + #[tokio::test] + async fn relay_child_scope_drop_aborts_both_children() { + let dropped = Arc::new(AtomicUsize::new(0)); + let flow_cancel = CancellationToken::new(); + let (stop_tx, stop_rx) = oneshot::channel(); + let child_tasks = RelayChildTasks { + c2me_sender: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal( + Arc::clone(&dropped), + )))), + me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal( + Arc::clone(&dropped), + )))), + flow_cancel: flow_cancel.clone(), + stop_tx: Some(stop_tx), + }; + + drop(child_tasks); + assert!(flow_cancel.is_cancelled()); + assert!(stop_rx.await.is_ok()); + tokio::time::timeout(Duration::from_secs(1), async { + while dropped.load(Ordering::Acquire) != 2 { + tokio::task::yield_now().await; + } + }) + .await + .expect("both relay child futures must be dropped after scope cancellation"); + } +} diff --git a/src/proxy/mod.rs b/src/proxy/mod.rs index 7908e12..01b9f08 100644 --- a/src/proxy/mod.rs +++ b/src/proxy/mod.rs @@ -74,6 +74,8 @@ pub mod session_eviction; pub mod shared_state; pub mod traffic_limiter; pub(crate) mod user_admission; +// Process-wide per-user connection admission remains independent from telemetry. +pub(crate) mod user_connection_authority; pub use client::ClientHandler; #[allow(unused_imports)] diff --git a/src/proxy/relay/io.rs b/src/proxy/relay/io.rs index 29171a9..50697ce 100644 --- a/src/proxy/relay/io.rs +++ b/src/proxy/relay/io.rs @@ -16,7 +16,7 @@ mod quota; pub(super) use self::combined::CombinedStream; pub(super) use self::counters::SharedCounters; pub(super) use self::quota::is_quota_io_error; -use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error}; +use self::quota::{QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL, quota_io_error}; pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check}; /// Transparent I/O wrapper that tracks per-user statistics and activity. @@ -218,48 +218,33 @@ impl AsyncRead for StatsIo { let mut quota_reservation = None; let mut read_limit = buf.remaining(); if let Some(limit) = this.quota_limit { - let used_before = this.quota_handle.used(); - let remaining = limit.saturating_sub(used_before); - if remaining == 0 { - this.quota_exceeded.store(true, Ordering::Release); - return Poll::Ready(Err(quota_io_error())); - } - remaining_before = Some(remaining); - read_limit = read_limit.min(remaining as usize); - if read_limit == 0 { - this.quota_exceeded.store(true, Ordering::Release); - return Poll::Ready(Err(quota_io_error())); - } - - let desired = read_limit as u64; - let mut reserve_rounds = 0usize; - while quota_reservation.is_none() { - for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.quota_handle.try_reserve(desired, limit) { - Ok(reservation) => { - quota_reservation = Some(reservation); - break; - } - Err(crate::stats::QuotaReserveError::LimitExceeded) => { - this.quota_exceeded.store(true, Ordering::Release); - return Poll::Ready(Err(quota_io_error())); - } - Err(crate::stats::QuotaReserveError::Contended) => { - this.stats.increment_quota_contention_total(); - } + for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL { + let used_before = this.quota_handle.used(); + let remaining = limit.saturating_sub(used_before); + if remaining == 0 { + this.quota_exceeded.store(true, Ordering::Release); + return Poll::Ready(Err(quota_io_error())); + } + let desired = remaining.min(read_limit as u64); + match this.quota_handle.try_reserve(desired, limit) { + Ok(reservation) => { + remaining_before = Some(remaining); + read_limit = desired as usize; + quota_reservation = Some(reservation); + break; + } + Err(crate::stats::QuotaReserveError::LimitExceeded) + | Err(crate::stats::QuotaReserveError::Contended) => { + this.stats.increment_quota_contention_total(); } } - - if quota_reservation.is_none() { - reserve_rounds = reserve_rounds.saturating_add(1); - if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { - this.stats.increment_quota_contention_timeout_total(); - if this.arm_quota_wait(cx).is_pending() { - return Poll::Pending; - } - reserve_rounds = 0; - } + } + if quota_reservation.is_none() { + this.stats.increment_quota_contention_timeout_total(); + if this.arm_quota_wait(cx).is_ready() { + cx.waker().wake_by_ref(); } + return Poll::Pending; } } @@ -404,8 +389,7 @@ impl AsyncWrite for StatsIo { let mut quota_reservation = None; if let Some(limit) = this.quota_limit { if !write_buf.is_empty() { - let mut reserve_rounds = 0usize; - while quota_reservation.is_none() { + for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL { let used_before = this.quota_handle.used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { @@ -415,36 +399,32 @@ impl AsyncWrite for StatsIo { remaining_before = Some(remaining); let desired = remaining.min(write_buf.len() as u64); - let mut saw_contention = false; - for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.quota_handle.try_reserve(desired, limit) { - Ok(reservation) => { - quota_reservation = Some(reservation); - write_buf = &write_buf[..desired as usize]; - break; - } - Err(crate::stats::QuotaReserveError::LimitExceeded) => { - break; - } - Err(crate::stats::QuotaReserveError::Contended) => { - this.stats.increment_quota_contention_total(); - saw_contention = true; - } + match this.quota_handle.try_reserve(desired, limit) { + Ok(reservation) => { + quota_reservation = Some(reservation); + write_buf = &write_buf[..desired as usize]; + break; + } + Err(crate::stats::QuotaReserveError::LimitExceeded) + | Err(crate::stats::QuotaReserveError::Contended) => { + this.stats.increment_quota_contention_total(); } } - - if quota_reservation.is_none() { - reserve_rounds = reserve_rounds.saturating_add(1); - if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { - this.stats.increment_quota_contention_timeout_total(); - Self::arm_wait(&mut this.quota_wait, false, false); - let _ = - Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up); - return Poll::Pending; - } else if saw_contention { - std::hint::spin_loop(); - } + } + if quota_reservation.is_none() { + this.stats.increment_quota_contention_timeout_total(); + Self::arm_wait(&mut this.quota_wait, false, false); + if Self::poll_wait( + &mut this.quota_wait, + cx, + None, + RateDirection::Up, + ) + .is_ready() + { + cx.waker().wake_by_ref(); } + return Poll::Pending; } } else { let used_before = this.quota_handle.used(); diff --git a/src/proxy/relay/io/quota.rs b/src/proxy/relay/io/quota.rs index b5f87e0..2b9c8c1 100644 --- a/src/proxy/relay/io/quota.rs +++ b/src/proxy/relay/io/quota.rs @@ -27,8 +27,7 @@ const QUOTA_NEAR_LIMIT_BYTES: u64 = 64 * 1024; const QUOTA_LARGE_CHARGE_BYTES: u64 = 16 * 1024; const QUOTA_ADAPTIVE_INTERVAL_MIN_BYTES: u64 = 4 * 1024; const QUOTA_ADAPTIVE_INTERVAL_MAX_BYTES: u64 = 64 * 1024; -pub(super) const QUOTA_RESERVE_SPIN_RETRIES: usize = 64; -pub(super) const QUOTA_RESERVE_MAX_ROUNDS: usize = 8; +pub(super) const QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL: usize = 4; #[inline] pub(in crate::proxy::relay) fn quota_adaptive_interval_bytes(remaining_before: u64) -> u64 { diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index f19f395..a4e15ff 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -123,6 +123,19 @@ impl ProxySharedState { pub(crate) fn new_with_direct_buffer_budget_and_user_admission( direct_buffer_budget: Arc, user_admission: Arc, + ) -> Arc { + Self::new_with_process_authorities( + direct_buffer_budget, + TrafficLimiter::new(), + user_admission, + ) + } + + /// Creates generation-local caches around process-owned data-plane authorities. + pub(crate) fn new_with_process_authorities( + direct_buffer_budget: Arc, + traffic_limiter: Arc, + user_admission: Arc, ) -> Arc { Arc::new(Self { handshake: HandshakeSharedState { @@ -163,7 +176,7 @@ impl ProxySharedState { relay_idle_registry: RelayIdleCandidateRegistry::default(), relay_idle_mark_seq: AtomicU64::new(0), }, - traffic_limiter: TrafficLimiter::new(), + traffic_limiter, direct_buffer_budget, user_admission, conntrack_pressure_active: AtomicBool::new(false), diff --git a/src/proxy/tests/client_security_tests.rs b/src/proxy/tests/client_security_tests.rs index 338129d..9201736 100644 --- a/src/proxy/tests/client_security_tests.rs +++ b/src/proxy/tests/client_security_tests.rs @@ -276,14 +276,13 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() { ip_tracker.set_user_limit(&user, 1).await; ip_tracker.check_and_add(&user, ip).await.unwrap(); - stats.increment_user_curr_connects(&user); - assert_eq!(ip_tracker.get_active_ip_count(&user).await, 1); - assert_eq!(stats.get_user_curr_connects(&user), 1); let reservation = UserConnectionReservation::new(stats.clone(), ip_tracker.clone(), user.clone(), ip, true); + assert_eq!(stats.get_user_curr_connects(&user), 1); + // Drop the reservation synchronously without any tokio::spawn/await yielding! drop(reservation); @@ -304,6 +303,117 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() { assert_eq!(ip_tracker.get_active_ip_count(&user).await, 0); } +#[tokio::test] +async fn cancelled_ip_admission_releases_process_connection_permit() { + let ip_tracker = Arc::new(UserIpTracker::new()); + let stats = Arc::new(Stats::new()); + let user = "cancelled-admission-user"; + let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap(); + let mut config = ProxyConfig::default(); + config + .access + .user_max_tcp_conns + .insert(user.to_string(), 1); + + let (entered_tx, entered_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel(); + let held_tracker = Arc::clone(&ip_tracker); + let held_user = user.to_string(); + let holder = tokio::spawn(async move { + held_tracker + .hold_user_shard_for_tests(&held_user, entered_tx, release_rx) + .await; + }); + entered_rx.await.unwrap(); + + let acquire_stats = Arc::clone(&stats); + let acquire_tracker = Arc::clone(&ip_tracker); + let acquire_config = config.clone(); + let acquire = tokio::spawn(async move { + acquire_user_connection_reservation( + user, + &acquire_config, + acquire_stats, + peer_addr, + acquire_tracker, + ) + .await + }); + + tokio::time::timeout(Duration::from_secs(1), async { + while stats.get_process_user_curr_connects(user) != 1 { + tokio::task::yield_now().await; + } + }) + .await + .expect("connection permit must be acquired before IP admission completes"); + + acquire.abort(); + let acquire_result = acquire.await; + assert!( + acquire_result + .as_ref() + .is_err_and(tokio::task::JoinError::is_cancelled) + ); + assert_eq!(stats.get_process_user_curr_connects(user), 0); + assert_eq!(stats.get_user_curr_connects(user), 0); + assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 0); + + let _ = release_tx.send(()); + holder.await.unwrap(); +} + +#[tokio::test] +async fn cancelled_async_release_preserves_ip_cleanup_ownership() { + let ip_tracker = Arc::new(UserIpTracker::new()); + let stats = Arc::new(Stats::new()); + let user = "cancelled-release-user"; + let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap(); + let mut config = ProxyConfig::default(); + config + .access + .user_max_tcp_conns + .insert(user.to_string(), 1); + + let reservation = acquire_user_connection_reservation( + user, + &config, + Arc::clone(&stats), + peer_addr, + Arc::clone(&ip_tracker), + ) + .await + .unwrap(); + assert_eq!(stats.get_process_user_curr_connects(user), 1); + assert_eq!(ip_tracker.get_active_ip_count(user).await, 1); + + let (entered_tx, entered_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel(); + let held_tracker = Arc::clone(&ip_tracker); + let held_user = user.to_string(); + let holder = tokio::spawn(async move { + held_tracker + .hold_user_shard_for_tests(&held_user, entered_tx, release_rx) + .await; + }); + entered_rx.await.unwrap(); + + let release = tokio::spawn(reservation.release()); + tokio::task::yield_now().await; + release.abort(); + assert!(release.await.unwrap_err().is_cancelled()); + + assert_eq!(stats.get_process_user_curr_connects(user), 0); + assert_eq!(stats.get_user_curr_connects(user), 0); + assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 1); + + let _ = release_tx.send(()); + holder.await.unwrap(); + ip_tracker.drain_cleanup_queue().await; + assert_eq!(ip_tracker.get_active_ip_count(user).await, 0); + assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 0); +} + #[tokio::test] async fn relay_task_abort_releases_user_gate_and_ip_reservation() { let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); @@ -2813,7 +2923,11 @@ async fn tcp_limit_rejection_does_not_reserve_ip_or_trigger_rollback() { .insert("user".to_string(), 1); let stats = Stats::new(); - stats.increment_user_curr_connects("user"); + let _existing_connection = stats + .connection_authority() + .try_acquire("user", Some(1)) + .expect("existing connection must occupy the process admission slot"); + let _existing_observation = stats.observe_user_current_connection("user"); let ip_tracker = UserIpTracker::new(); let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap(); @@ -2853,7 +2967,11 @@ async fn zero_tcp_limit_uses_global_fallback_and_rejects_without_side_effects() config.access.user_max_tcp_conns_global_each = 1; let stats = Stats::new(); - stats.increment_user_curr_connects("user"); + let _existing_connection = stats + .connection_authority() + .try_acquire("user", Some(1)) + .expect("existing connection must occupy the process admission slot"); + let _existing_observation = stats.observe_user_current_connection("user"); let ip_tracker = UserIpTracker::new(); let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap(); @@ -2914,7 +3032,11 @@ async fn global_tcp_fallback_applies_when_per_user_limit_is_missing() { config.access.user_max_tcp_conns_global_each = 1; let stats = Stats::new(); - stats.increment_user_curr_connects("user"); + let _existing_connection = stats + .connection_authority() + .try_acquire("user", Some(1)) + .expect("existing connection must occupy the process admission slot"); + let _existing_observation = stats.observe_user_current_connection("user"); let ip_tracker = UserIpTracker::new(); let peer_addr: SocketAddr = "198.51.100.213:50003".parse().unwrap(); @@ -4024,7 +4146,11 @@ async fn concurrent_limit_rejections_from_mixed_ips_leave_no_ip_footprint() { let config = Arc::new(config); let stats = Arc::new(Stats::new()); - stats.increment_user_curr_connects("user"); + let _existing_connection = stats + .connection_authority() + .try_acquire("user", Some(1)) + .expect("existing connection must occupy the process admission slot"); + let _existing_observation = stats.observe_user_current_connection("user"); let ip_tracker = Arc::new(UserIpTracker::new()); let mut tasks = tokio::task::JoinSet::new(); diff --git a/src/proxy/tests/direct_buffer_budget_tests.rs b/src/proxy/tests/direct_buffer_budget_tests.rs index 796e7fc..e4e663a 100644 --- a/src/proxy/tests/direct_buffer_budget_tests.rs +++ b/src/proxy/tests/direct_buffer_budget_tests.rs @@ -1,4 +1,5 @@ use super::*; +use tokio::sync::Semaphore; #[test] fn lease_drop_releases_the_complete_reservation() { @@ -36,3 +37,62 @@ fn growth_and_shrink_keep_accounting_balanced() { drop(lease); assert_eq!(budget.snapshot().reserved_bytes, 0); } + +#[test] +fn runtime_generations_share_one_absolute_reservation_envelope() { + let first_generation = DirectBufferBudget::new(16 * 1024); + let second_generation = Arc::clone(&first_generation); + let first = first_generation + .try_reserve(12 * 1024, true) + .expect("first generation reservation must fit"); + + assert!(second_generation.try_reserve(8 * 1024, true).is_none()); + assert_eq!(second_generation.snapshot().reserved_bytes, 12 * 1024); + + drop(first); + assert!(second_generation.try_reserve(8 * 1024, true).is_some()); +} + +#[test] +fn stale_runtime_cannot_reclaim_direct_controller_ownership() { + let budget = DirectBufferBudget::new(16 * 1024); + budget.activate_controller(2); + budget.activate_controller(1); + + assert_eq!( + budget.active_controller_generation.load(Ordering::Acquire), + 2 + ); +} + +#[test] +fn controller_handoff_waits_for_inflight_update_and_fences_old_generation() { + let budget = DirectBufferBudget::new(16 * 1024); + budget.activate_controller(1); + let update = budget.begin_controller_update(1).unwrap(); + let (activated_tx, activated_rx) = std::sync::mpsc::channel(); + let next_budget = Arc::clone(&budget); + let activation = std::thread::spawn(move || { + next_budget.activate_controller(2); + activated_tx.send(()).unwrap(); + }); + + assert!(activated_rx + .recv_timeout(Duration::from_millis(50)) + .is_err()); + drop(update); + activated_rx.recv_timeout(Duration::from_secs(1)).unwrap(); + activation.join().unwrap(); + + assert!(budget.begin_controller_update(1).is_none()); + assert!(budget.begin_controller_update(2).is_some()); +} + +#[test] +fn connection_pressure_uses_process_wide_slot_ownership() { + let slots = Arc::new(Semaphore::new(10)); + let _old_generation = Arc::clone(&slots).try_acquire_many_owned(3).unwrap(); + let _new_generation = Arc::clone(&slots).try_acquire_many_owned(2).unwrap(); + + assert_eq!(connection_fill_pct(slots.as_ref(), 10), Some(50)); +} diff --git a/src/proxy/traffic_limiter.rs b/src/proxy/traffic_limiter.rs index 40047af..51e2b6e 100644 --- a/src/proxy/traffic_limiter.rs +++ b/src/proxy/traffic_limiter.rs @@ -142,6 +142,7 @@ enum CidrPolicyMatch<'a> { #[derive(Default)] struct PolicySnapshot { revision: u64, + source_generation: u64, user_limits: HashMap, cidr_rules_v4: Vec, cidr_rules_v6: Vec, @@ -175,7 +176,6 @@ pub struct TrafficLease { pub struct TrafficLimiter { policy: ArcSwap, policy_update: ParkingMutex<()>, - published_revision: AtomicU64, user_buckets: ShardedRegistry, cidr_buckets: ShardedRegistry, user_scope: ScopeMetrics, diff --git a/src/proxy/traffic_limiter/lease.rs b/src/proxy/traffic_limiter/lease.rs index 452cdfa..02ea7f0 100644 --- a/src/proxy/traffic_limiter/lease.rs +++ b/src/proxy/traffic_limiter/lease.rs @@ -2,20 +2,16 @@ use super::*; impl TrafficLease { fn current_binding(&self) -> Arc { - let published_revision = self.limiter.published_revision.load(Ordering::Acquire); + let policy = self.limiter.policy.load(); let current = self.binding.load_full(); - if current.revision == published_revision { + if current.revision == policy.revision { return current; } + drop(policy); - let refresh = self.refresh.lock(); - let published_revision = self.limiter.published_revision.load(Ordering::Acquire); - let current = self.binding.load_full(); - if current.revision == published_revision { - return current; - } - let policy_update = self.limiter.policy_update.lock(); + let _refresh = self.refresh.lock(); let policy = self.limiter.policy.load_full(); + let current = self.binding.load_full(); if current.revision == policy.revision { return current; } @@ -23,9 +19,6 @@ impl TrafficLease { .limiter .build_binding(&self.user, self.client_ip, &policy); self.binding.store(Arc::clone(&next)); - drop(policy_update); - drop(refresh); - self.limiter.maybe_cleanup(); next } diff --git a/src/proxy/traffic_limiter/limiter.rs b/src/proxy/traffic_limiter/limiter.rs index d407602..6c55b89 100644 --- a/src/proxy/traffic_limiter/limiter.rs +++ b/src/proxy/traffic_limiter/limiter.rs @@ -6,7 +6,6 @@ impl TrafficLimiter { Arc::new(Self { policy: ArcSwap::from_pointee(PolicySnapshot::default()), policy_update: ParkingMutex::new(()), - published_revision: AtomicU64::new(0), user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), user_scope: ScopeMetrics::default(), @@ -15,16 +14,31 @@ impl TrafficLimiter { }) } + #[cfg(test)] pub fn apply_policy( &self, user_limits: HashMap, cidr_limits: HashMap, ) { - let policy_update = self.policy_update.lock(); - // Revision wrap could otherwise let an old lease restore stale rates. - let Some(revision) = self.policy.load().revision.checked_add(1) else { - return; - }; + let _ = self.apply_policy_inner(None, user_limits, cidr_limits); + } + + /// Publishes policy only when the source runtime is not older than the active source. + pub(crate) fn apply_policy_from_source( + &self, + source_generation: u64, + user_limits: HashMap, + cidr_limits: HashMap, + ) -> bool { + self.apply_policy_inner(Some(source_generation), user_limits, cidr_limits) + } + + fn apply_policy_inner( + &self, + source_generation: Option, + user_limits: HashMap, + cidr_limits: HashMap, + ) -> bool { let filtered_users = user_limits .into_iter() .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) @@ -77,6 +91,17 @@ impl TrafficLimiter { let cidr_policy_entries = cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len(); + let policy_update = self.policy_update.lock(); + let current = self.policy.load_full(); + if source_generation.is_some_and(|source| source < current.source_generation) { + return false; + } + // Revision wrap could otherwise let an old lease restore stale rates. + let Some(revision) = current.revision.checked_add(1) else { + return false; + }; + let source_generation = source_generation.unwrap_or(current.source_generation); + self.user_scope .policy_entries .store(filtered_users.len() as u64, Ordering::Relaxed); @@ -86,6 +111,7 @@ impl TrafficLimiter { self.policy.store(Arc::new(PolicySnapshot { revision, + source_generation, user_limits: filtered_users, cidr_rules_v4, cidr_rules_v6, @@ -93,10 +119,10 @@ impl TrafficLimiter { cidr_auto_rules_v6, cidr_rule_keys, })); - self.published_revision.store(revision, Ordering::Release); drop(policy_update); self.maybe_cleanup(); + true } pub fn acquire_lease( @@ -104,11 +130,8 @@ impl TrafficLimiter { user: &str, client_ip: IpAddr, ) -> Option> { - let policy_update = self.policy_update.lock(); let policy = self.policy.load_full(); let binding = self.build_binding(user, client_ip, &policy); - drop(policy_update); - self.maybe_cleanup(); Some(Arc::new(TrafficLease { limiter: Arc::clone(self), user: user.to_string(), diff --git a/src/proxy/traffic_limiter/tests.rs b/src/proxy/traffic_limiter/tests.rs index b1a1e2d..975dff8 100644 --- a/src/proxy/traffic_limiter/tests.rs +++ b/src/proxy/traffic_limiter/tests.rs @@ -5,6 +5,79 @@ fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps { RateLimitBps { up_bps, down_bps } } +#[test] +fn stale_runtime_cannot_overwrite_newer_rate_policy() { + let limiter = TrafficLimiter::new(); + let mut newer = HashMap::new(); + newer.insert("alice".to_string(), rate(2_000, 3_000)); + assert!(limiter.apply_policy_from_source(2, newer, HashMap::new())); + + let mut stale = HashMap::new(); + stale.insert("alice".to_string(), rate(1_000, 1_000)); + assert!(!limiter.apply_policy_from_source(1, stale, HashMap::new())); + + let policy = limiter.policy.load_full(); + assert_eq!(policy.source_generation, 2); + assert_eq!(policy.user_limits["alice"].up_bps, 2_000); + assert_eq!(policy.user_limits["alice"].down_bps, 3_000); +} + +#[test] +fn active_runtime_can_publish_same_generation_rate_update() { + let limiter = TrafficLimiter::new(); + assert!(limiter.apply_policy_from_source(4, HashMap::new(), HashMap::new())); + let mut updated = HashMap::new(); + updated.insert("alice".to_string(), rate(4_000, 5_000)); + + assert!(limiter.apply_policy_from_source(4, updated, HashMap::new())); + + let policy = limiter.policy.load_full(); + assert_eq!(policy.source_generation, 4); + assert_eq!(policy.user_limits["alice"].up_bps, 4_000); +} + +#[test] +fn lease_acquisition_and_refresh_do_not_wait_for_policy_publication_lock() { + let limiter = TrafficLimiter::new(); + let mut initial = HashMap::new(); + initial.insert("alice".to_string(), rate(1_000, 1_000)); + limiter.apply_policy(initial, HashMap::new()); + let lease = limiter + .acquire_lease("alice", "203.0.113.7".parse().unwrap()) + .unwrap(); + + let mut updated = HashMap::new(); + updated.insert("alice".to_string(), rate(2_000, 2_000)); + limiter.apply_policy(updated, HashMap::new()); + + let publication = limiter.policy_update.lock(); + let (completed_tx, completed_rx) = std::sync::mpsc::channel(); + let acquire_limiter = Arc::clone(&limiter); + let acquire_tx = completed_tx.clone(); + let acquire = std::thread::spawn(move || { + let _lease = acquire_limiter + .acquire_lease("bob", "203.0.113.8".parse().unwrap()) + .unwrap(); + acquire_tx.send(()).unwrap(); + }); + let refresh = std::thread::spawn(move || { + let _ = lease.try_consume(RateDirection::Up, 1); + completed_tx.send(()).unwrap(); + }); + + let first_completed = completed_rx + .recv_timeout(std::time::Duration::from_secs(1)) + .is_ok(); + let second_completed = completed_rx + .recv_timeout(std::time::Duration::from_secs(1)) + .is_ok(); + drop(publication); + acquire.join().unwrap(); + refresh.join().unwrap(); + + assert!(first_completed && second_completed); +} + #[test] fn explicit_cidr_rule_wins_over_auto_template() { let limiter = TrafficLimiter::new(); diff --git a/src/proxy/user_connection_authority.rs b/src/proxy/user_connection_authority.rs new file mode 100644 index 0000000..373f99d --- /dev/null +++ b/src/proxy/user_connection_authority.rs @@ -0,0 +1,123 @@ +use std::sync::Arc; + +use dashmap::DashMap; +use dashmap::mapref::entry::Entry; + +/// Process-wide per-user connection admission shared by runtime generations. +#[derive(Default)] +pub(crate) struct UserConnectionAuthority { + active: DashMap, +} + +/// Owns one exact connection slot until the authenticated connection exits. +#[must_use = "connection permits must be retained for the connection lifetime"] +pub(crate) struct UserConnectionPermit { + authority: Arc, + user: String, +} + +impl UserConnectionAuthority { + /// Acquires a slot without consulting optional telemetry state. + pub(crate) fn try_acquire( + self: &Arc, + user: &str, + limit: Option, + ) -> Option { + match self.active.entry(user.to_string()) { + Entry::Occupied(mut entry) => { + if limit.is_some_and(|max| *entry.get() >= max) { + return None; + } + let next = entry.get().checked_add(1)?; + *entry.get_mut() = next; + } + Entry::Vacant(entry) => { + if limit == Some(0) { + return None; + } + entry.insert(1); + } + } + Some(UserConnectionPermit { + authority: Arc::clone(self), + user: user.to_string(), + }) + } + + /// Returns the authoritative active connection count for one username. + pub(crate) fn active(&self, user: &str) -> u64 { + self.active.get(user).map(|entry| *entry).unwrap_or(0) + } + + #[cfg(test)] + fn tracked_users(&self) -> usize { + self.active.len() + } +} + +impl Drop for UserConnectionPermit { + fn drop(&mut self) { + let Entry::Occupied(mut entry) = self.authority.active.entry(self.user.clone()) else { + debug_assert!(false, "connection permit owner entry disappeared"); + return; + }; + debug_assert!(*entry.get() > 0, "connection permit counter underflow"); + if *entry.get() <= 1 { + entry.remove(); + } else { + *entry.get_mut() -= 1; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Barrier; + + #[test] + fn permit_drop_releases_and_removes_zero_entry() { + let authority = Arc::new(UserConnectionAuthority::default()); + let permit = authority.try_acquire("alice", Some(1)).unwrap(); + assert_eq!(authority.active("alice"), 1); + assert!(authority.try_acquire("alice", Some(1)).is_none()); + + drop(permit); + + assert_eq!(authority.active("alice"), 0); + assert_eq!(authority.tracked_users(), 0); + } + + #[test] + fn concurrent_acquire_never_exceeds_limit() { + const CONTENDERS: usize = 64; + const LIMIT: u64 = 7; + + let authority = Arc::new(UserConnectionAuthority::default()); + let barrier = Arc::new(Barrier::new(CONTENDERS + 1)); + let (permit_tx, permit_rx) = std::sync::mpsc::channel(); + let mut threads = Vec::with_capacity(CONTENDERS); + for _ in 0..CONTENDERS { + let authority = Arc::clone(&authority); + let barrier = Arc::clone(&barrier); + let permit_tx = permit_tx.clone(); + threads.push(std::thread::spawn(move || { + barrier.wait(); + permit_tx + .send(authority.try_acquire("alice", Some(LIMIT))) + .unwrap(); + })); + } + drop(permit_tx); + barrier.wait(); + for thread in threads { + thread.join().unwrap(); + } + let permits = permit_rx.into_iter().flatten().collect::>(); + + assert_eq!(permits.len() as u64, LIMIT); + assert_eq!(authority.active("alice"), LIMIT); + drop(permits); + assert_eq!(authority.active("alice"), 0); + } +} diff --git a/src/stats/mod.rs b/src/stats/mod.rs index f237e29..f751a10 100644 --- a/src/stats/mod.rs +++ b/src/stats/mod.rs @@ -22,9 +22,11 @@ use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering}; use std::time::Instant; pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle}; +pub(crate) use self::users::UserConnectionObservation; #[allow(unused_imports)] pub use self::replay::{ReplayChecker, ReplayStats}; use self::telemetry::TelemetryPolicy; +use crate::proxy::user_connection_authority::UserConnectionAuthority; pub use self::tls_fingerprints::TlsFingerprintSnapshotRow; use crate::config::MeWriterPickMode; @@ -351,6 +353,7 @@ pub struct Stats { tls_fingerprints: tls_fingerprints::TlsFingerprintCollector, user_stats: DashMap>, quota_store: Arc, + connection_authority: Arc, user_stats_last_cleanup_epoch_secs: AtomicU64, start_time: parking_lot::RwLock>, } @@ -417,12 +420,28 @@ impl UserStats { impl Stats { pub fn new() -> Self { - Self::with_quota_store(Arc::new(QuotaStore::default())) + Self::with_process_authorities( + Arc::new(QuotaStore::default()), + Arc::new(UserConnectionAuthority::default()), + ) } + #[cfg(test)] pub(crate) fn with_quota_store(quota_store: Arc) -> Self { + Self::with_process_authorities( + quota_store, + Arc::new(UserConnectionAuthority::default()), + ) + } + + /// Creates generation telemetry around process-owned enforcement authorities. + pub(crate) fn with_process_authorities( + quota_store: Arc, + connection_authority: Arc, + ) -> Self { let stats = Self { quota_store, + connection_authority, ..Self::default() }; stats.apply_telemetry_policy(TelemetryPolicy::default()); @@ -431,10 +450,16 @@ impl Stats { stats } + /// Returns the process-scoped quota authority for test runtime construction. #[cfg(test)] pub(crate) fn quota_store(&self) -> Arc { Arc::clone(&self.quota_store) } + + /// Returns process-owned per-user connection admission. + pub(crate) fn connection_authority(&self) -> Arc { + Arc::clone(&self.connection_authority) + } } #[cfg(test)] diff --git a/src/stats/tests.rs b/src/stats/tests.rs index f90e601..a93f66c 100644 --- a/src/stats/tests.rs +++ b/src/stats/tests.rs @@ -13,6 +13,32 @@ fn test_stats_shared_counters() { assert_eq!(stats.get_connects_all(), 3); } +#[test] +fn runtime_stats_share_process_connection_admission_authority() { + let quota_store = Arc::new(QuotaStore::default()); + let authority = Arc::new(UserConnectionAuthority::default()); + let first = Stats::with_process_authorities( + Arc::clone("a_store), + Arc::clone(&authority), + ); + let second = Stats::with_process_authorities(quota_store, authority); + + let permit = first + .connection_authority() + .try_acquire("alice", Some(1)) + .unwrap(); + assert!(second + .connection_authority() + .try_acquire("alice", Some(1)) + .is_none()); + + drop(permit); + assert!(second + .connection_authority() + .try_acquire("alice", Some(1)) + .is_some()); +} + #[test] fn test_telemetry_policy_disables_core_and_user_counters() { let stats = Stats::new(); diff --git a/src/stats/users.rs b/src/stats/users.rs index 982c2c1..e092b7b 100644 --- a/src/stats/users.rs +++ b/src/stats/users.rs @@ -1,5 +1,32 @@ use super::*; +/// Mirrors one live connection into optional per-user telemetry. +#[must_use = "the observation must be retained for the connection lifetime"] +pub(crate) struct UserConnectionObservation { + stats: Arc, +} + +impl Drop for UserConnectionObservation { + fn drop(&mut self) { + decrement_current_connections(&self.stats.curr_connects); + } +} + +fn decrement_current_connections(counter: &AtomicU64) { + let mut current = counter.load(Ordering::Relaxed); + while current != 0 { + match counter.compare_exchange_weak( + current, + current - 1, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => return, + Err(actual) => current = actual, + } + } +} + impl Stats { pub fn increment_user_connects(&self, user: &str) { if !self.telemetry_user_enabled() { @@ -19,6 +46,25 @@ impl Stats { stats.curr_connects.fetch_add(1, Ordering::Relaxed); } + /// Starts an optional telemetry observation without owning admission policy. + pub(crate) fn observe_user_current_connection( + &self, + user: &str, + ) -> Option { + if !self.telemetry_user_enabled() { + return None; + } + let stats = self.get_or_create_user_stats_handle(user); + self.touch_user_stats(stats.as_ref()); + stats.curr_connects.fetch_add(1, Ordering::Relaxed); + Some(UserConnectionObservation { stats }) + } + + /// Returns the process-scoped count used by connection admission. + pub(crate) fn get_process_user_curr_connects(&self, user: &str) -> u64 { + self.connection_authority.active(user) + } + pub fn try_acquire_user_curr_connects(&self, user: &str, limit: Option) -> bool { if !self.telemetry_user_enabled() { return true; @@ -50,22 +96,7 @@ impl Stats { pub fn decrement_user_curr_connects(&self, user: &str) { if let Some(stats) = self.user_stats.get(user) { self.touch_user_stats(stats.value().as_ref()); - let counter = &stats.curr_connects; - let mut current = counter.load(Ordering::Relaxed); - loop { - if current == 0 { - break; - } - match counter.compare_exchange_weak( - current, - current - 1, - Ordering::Relaxed, - Ordering::Relaxed, - ) { - Ok(_) => break, - Err(actual) => current = actual, - } - } + decrement_current_connections(&stats.curr_connects); } } diff --git a/src/synlimit_control/command.rs b/src/synlimit_control/command.rs index 0376c8a..e20d25e 100644 --- a/src/synlimit_control/command.rs +++ b/src/synlimit_control/command.rs @@ -3,6 +3,8 @@ use tokio::process::Command; use crate::util::trusted_command::resolve_trusted_helper; +const COMMAND_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30); + pub(super) async fn run_command( binary: &str, args: &[&str], @@ -18,21 +20,26 @@ pub(super) async fn run_command( } command.stdout(std::process::Stdio::null()); command.stderr(std::process::Stdio::piped()); + command.kill_on_drop(true); let mut child = command .spawn() .map_err(|e| format!("spawn {binary} failed: {e}"))?; - if let Some(blob) = stdin - && let Some(mut writer) = child.stdin.take() - { - writer - .write_all(blob.as_bytes()) + let output = tokio::time::timeout(COMMAND_TIMEOUT, async move { + if let Some(blob) = stdin + && let Some(mut writer) = child.stdin.take() + { + writer + .write_all(blob.as_bytes()) + .await + .map_err(|e| format!("stdin write {binary} failed: {e}"))?; + } + child + .wait_with_output() .await - .map_err(|e| format!("stdin write {binary} failed: {e}"))?; - } - let output = child - .wait_with_output() - .await - .map_err(|e| format!("wait {binary} failed: {e}"))?; + .map_err(|e| format!("wait {binary} failed: {e}")) + }) + .await + .map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??; if output.status.success() { return Ok(()); } @@ -48,10 +55,11 @@ pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result usize { - match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { + fn writer_contour_rank_for_selection(contour: WriterContour) -> usize { + match contour { WriterContour::Active => 0, WriterContour::Warm => 1, WriterContour::Draining => 2, } } - pub(super) fn writer_idle_rank_for_selection( - &self, - writer: &super::super::pool::MeWriter, + fn writer_idle_rank_for_selection( + writer_id: u64, idle_since_by_writer: &HashMap, now_epoch_secs: u64, ) -> usize { - let Some(idle_since) = idle_since_by_writer.get(&writer.id).copied() else { + let Some(idle_since) = idle_since_by_writer.get(&writer_id).copied() else { return 0; }; let idle_age_secs = now_epoch_secs.saturating_sub(idle_since); @@ -102,51 +111,67 @@ impl MePool { } } - pub(super) fn writer_pick_score( + fn capture_writer_selection_key( &self, + index: usize, writer: &super::super::pool::MeWriter, idle_since_by_writer: &HashMap, now_epoch_secs: u64, - ) -> u64 { - let contour_penalty = match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { + current_generation: u64, + ) -> WriterSelectionKey { + let contour = WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)); + let contour_rank = Self::writer_contour_rank_for_selection(contour); + let contour_penalty = match contour { WriterContour::Active => 0, WriterContour::Warm => PICK_PENALTY_WARM, WriterContour::Draining => PICK_PENALTY_DRAINING, }; - let stale_penalty = if writer.generation < self.current_generation() { + let stale = (writer.generation < current_generation) as usize; + let stale_penalty = if stale != 0 { PICK_PENALTY_STALE } else { 0 }; - let degraded_penalty = if writer.degraded.load(Ordering::Relaxed) { + let degraded = writer.degraded.load(Ordering::Relaxed) as usize; + let degraded_penalty = if degraded != 0 { PICK_PENALTY_DEGRADED } else { 0 }; - let idle_penalty = - (self.writer_idle_rank_for_selection(writer, idle_since_by_writer, now_epoch_secs) - as u64) - * 100; + let idle_rank = + Self::writer_idle_rank_for_selection(writer.id, idle_since_by_writer, now_epoch_secs); + let idle_penalty = (idle_rank as u64) * 100; let queue_cap = self.writer_lifecycle.writer_cmd_channel_capacity.max(1) as u64; - let queue_remaining = writer.tx.capacity() as u64; - let queue_used = queue_cap.saturating_sub(queue_remaining.min(queue_cap)); + let queue_remaining = writer.tx.capacity(); + let queue_used = queue_cap.saturating_sub((queue_remaining as u64).min(queue_cap)); let queue_util_pct = queue_used.saturating_mul(100) / queue_cap; let queue_penalty = queue_util_pct.saturating_mul(4); let rtt_penalty = ((writer.rtt_ema_ms_x10.load(Ordering::Relaxed) as u64).saturating_add(5) / 10) .min(400); - contour_penalty + let pick_score = contour_penalty .saturating_add(stale_penalty) .saturating_add(degraded_penalty) .saturating_add(idle_penalty) .saturating_add(queue_penalty) - .saturating_add(rtt_penalty) + .saturating_add(rtt_penalty); + WriterSelectionKey { + index, + contour_rank, + stale, + degraded, + idle_rank, + queue_remaining, + addr: writer.addr, + id: writer.id, + pick_score, + } } pub(super) fn p2c_ordered_candidate_indices( &self, - candidate_indices: &[usize], + mut candidate_indices: Vec, writers_snapshot: &[super::super::pool::MeWriter], idle_since_by_writer: &HashMap, now_epoch_secs: u64, @@ -158,33 +183,26 @@ impl MePool { return Vec::new(); } - let mut sampled = Vec::::with_capacity(sample_size.min(total)); - let mut seen = HashSet::::with_capacity(total); - for offset in 0..sample_size.min(total) { - let idx = candidate_indices[(start + offset) % total]; - if seen.insert(idx) { - sampled.push(idx); - } + candidate_indices.rotate_left(start % total); + let current_generation = self.current_generation(); + let sample_size = sample_size.min(total); + let mut sampled = candidate_indices[..sample_size] + .iter() + .map(|idx| { + self.capture_writer_selection_key( + *idx, + &writers_snapshot[*idx], + idle_since_by_writer, + now_epoch_secs, + current_generation, + ) + }) + .collect::>(); + sampled.sort_by_key(|candidate| (candidate.pick_score, candidate.addr, candidate.id)); + for (target, candidate) in candidate_indices[..sample_size].iter_mut().zip(sampled) { + *target = candidate.index; } - - sampled.sort_by_key(|idx| { - let writer = &writers_snapshot[*idx]; - ( - self.writer_pick_score(writer, idle_since_by_writer, now_epoch_secs), - writer.addr, - writer.id, - ) - }); - - let mut ordered = Vec::::with_capacity(total); - ordered.extend(sampled.iter().copied()); - for offset in 0..total { - let idx = candidate_indices[(start + offset) % total]; - if seen.insert(idx) { - ordered.push(idx); - } - } - ordered + candidate_indices } pub(super) async fn ordered_candidate_indices( @@ -206,7 +224,7 @@ impl MePool { let start = self.rr.fetch_add(1, Ordering::Relaxed) as usize % candidate_indices.len(); if pick_mode == MeWriterPickMode::P2c { return self.p2c_ordered_candidate_indices( - &candidate_indices, + candidate_indices, writers_snapshot, &writer_idle_since, now_epoch_secs, @@ -220,60 +238,65 @@ impl MePool { .me_deterministic_writer_sort .load(Ordering::Relaxed) { - candidate_indices.sort_by(|lhs, rhs| { - let left = &writers_snapshot[*lhs]; - let right = &writers_snapshot[*rhs]; - let left_key = ( - self.writer_contour_rank_for_selection(left), - (left.generation < self.current_generation()) as usize, - left.degraded.load(Ordering::Relaxed) as usize, - self.writer_idle_rank_for_selection( - left, + let current_generation = self.current_generation(); + let mut captured = candidate_indices + .iter() + .map(|idx| { + self.capture_writer_selection_key( + *idx, + &writers_snapshot[*idx], &writer_idle_since, now_epoch_secs, - ), - Reverse(left.tx.capacity()), - left.addr, - left.id, - ); - let right_key = ( - self.writer_contour_rank_for_selection(right), - (right.generation < self.current_generation()) as usize, - right.degraded.load(Ordering::Relaxed) as usize, - self.writer_idle_rank_for_selection( - right, - &writer_idle_since, - now_epoch_secs, - ), - Reverse(right.tx.capacity()), - right.addr, - right.id, - ); - left_key.cmp(&right_key) - }); - } else { - candidate_indices.sort_by_key(|idx| { - let writer = &writers_snapshot[*idx]; - let degraded = writer.degraded.load(Ordering::Relaxed); - let stale = (writer.generation < self.current_generation()) as usize; + current_generation, + ) + }) + .collect::>(); + captured.sort_by_key(|candidate| { ( - self.writer_contour_rank_for_selection(writer), - stale, - degraded as usize, - self.writer_idle_rank_for_selection( - writer, - &writer_idle_since, - now_epoch_secs, - ), - Reverse(writer.tx.capacity()), + candidate.contour_rank, + candidate.stale, + candidate.degraded, + candidate.idle_rank, + Reverse(candidate.queue_remaining), + candidate.addr, + candidate.id, ) }); + for (target, candidate) in candidate_indices.iter_mut().zip(captured) { + *target = candidate.index; + } + } else { + let current_generation = self.current_generation(); + let mut captured = candidate_indices + .iter() + .map(|idx| { + self.capture_writer_selection_key( + *idx, + &writers_snapshot[*idx], + &writer_idle_since, + now_epoch_secs, + current_generation, + ) + }) + .collect::>(); + captured.sort_by_key(|candidate| { + ( + candidate.contour_rank, + candidate.stale, + candidate.degraded, + candidate.idle_rank, + Reverse(candidate.queue_remaining), + ) + }); + for (target, candidate) in candidate_indices.iter_mut().zip(captured) { + *target = candidate.index; + } } - let mut ordered = Vec::::with_capacity(candidate_indices.len()); - for offset in 0..candidate_indices.len() { - ordered.push(candidate_indices[(start + offset) % candidate_indices.len()]); + if !candidate_indices.is_empty() { + let len = candidate_indices.len(); + candidate_indices.rotate_left(start % len); } - ordered + candidate_indices } } diff --git a/src/web/http/websocket/driver.rs b/src/web/http/websocket/driver.rs index b936ea3..4245bee 100644 --- a/src/web/http/websocket/driver.rs +++ b/src/web/http/websocket/driver.rs @@ -1,3 +1,4 @@ +use std::future::Future; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -112,6 +113,40 @@ pub(super) async fn run_upgraded( type CarrierSocket = WebSocketStream; +enum DataPlaneEvent { + Incoming(I), + Down(D), +} + +#[derive(Default)] +struct FairDataSelector { + prefer_down: bool, +} + +impl FairDataSelector { + async fn select(&mut self, incoming: I, down: D) -> DataPlaneEvent + where + I: Future, + D: Future, + { + let event = if self.prefer_down { + tokio::select! { + biased; + down = down => DataPlaneEvent::Down(down), + incoming = incoming => DataPlaneEvent::Incoming(incoming), + } + } else { + tokio::select! { + biased; + incoming = incoming => DataPlaneEvent::Incoming(incoming), + down = down => DataPlaneEvent::Down(down), + } + }; + self.prefer_down = matches!(event, DataPlaneEvent::Incoming(_)); + event + } +} + async fn run_multiplex( socket: &mut CarrierSocket, runtime: &Arc, @@ -134,26 +169,31 @@ async fn run_multiplex( let write_timeout = Duration::from_secs(session.timeouts().websocket_write_secs); let maximum_message = session.limits().carrier_batch_bytes; let mut active = false; + let mut data_selector = FairDataSelector::default(); loop { let down = session.poll_down_websocket(cursor); - tokio::pin!(down); let event = tokio::select! { biased; _ = cancellation.cancelled() => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, - incoming = read_message( - socket, - runtime, - session.profile_key(), - &cancellation, - &mut read_budget, - maximum_message, - backpressure_timeout, + data = data_selector.select( + read_message( + socket, + runtime, + session.profile_key(), + &cancellation, + &mut read_budget, + maximum_message, + backpressure_timeout, + ), + down, ) => { - DriverEvent::Incoming(incoming?) + match data { + DataPlaneEvent::Incoming(incoming) => DriverEvent::Incoming(incoming?), + DataPlaneEvent::Down(down) => DriverEvent::Down(down.map_err(|_| ())?), + } } - down = &mut down => DriverEvent::Down(down.map_err(|_| ())?), }; match event { DriverEvent::Incoming((message, _budget)) => match message { @@ -363,3 +403,22 @@ enum DriverEvent { Down(crate::web::session::PollResult), Liveness, } + +#[cfg(test)] +mod fairness_tests { + use super::*; + + #[tokio::test] + async fn continuously_ready_directions_alternate() { + let mut selector = FairDataSelector::default(); + for expected_incoming in [true, false, true, false] { + let event = selector + .select(std::future::ready("incoming"), std::future::ready("down")) + .await; + assert_eq!( + matches!(event, DataPlaneEvent::Incoming(_)), + expected_incoming + ); + } + } +} diff --git a/src/web/http/websocket/driver/lane.rs b/src/web/http/websocket/driver/lane.rs index 8780d0e..45574b0 100644 --- a/src/web/http/websocket/driver/lane.rs +++ b/src/web/http/websocket/driver/lane.rs @@ -5,9 +5,9 @@ use bytes::Bytes; use tokio_tungstenite::tungstenite::protocol::Message; use tokio_util::sync::CancellationToken; -use super::CarrierSocket; +use super::{CarrierSocket, DataPlaneEvent, DriverEvent, FairDataSelector}; use super::io::{flush, process_lane, read_message, record_message, reserve_data, send}; -use crate::web::manager::{WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection}; +use crate::web::manager::{WebProcessRuntime, WebSocketConnection}; use crate::web::session::{SessionCloseReason, WebSession, WebSocketLaneReservation}; use crate::web::trace::{TraceDirection, TraceWebSocketContext}; @@ -35,26 +35,31 @@ pub(super) async fn run_lane( let write_timeout = Duration::from_secs(session.timeouts().websocket_write_secs); let maximum_message = session.limits().carrier_batch_bytes; let mut active = false; + let mut data_selector = FairDataSelector::default(); loop { let down = session.poll_down_websocket_lane(reservation.lane_identity(), cursor); - tokio::pin!(down); let event = tokio::select! { biased; _ = cancellation.cancelled() => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, - incoming = read_message( - socket, - runtime, - session.profile_key(), - &cancellation, - &mut read_budget, - maximum_message, - backpressure_timeout, + data = data_selector.select( + read_message( + socket, + runtime, + session.profile_key(), + &cancellation, + &mut read_budget, + maximum_message, + backpressure_timeout, + ), + down, ) => { - DriverEvent::Incoming(incoming?) + match data { + DataPlaneEvent::Incoming(incoming) => DriverEvent::Incoming(incoming?), + DataPlaneEvent::Down(down) => DriverEvent::Down(down.map_err(|_| ())?), + } } - down = &mut down => DriverEvent::Down(down.map_err(|_| ())?), }; match event { DriverEvent::Incoming((message, _budget)) => match message { @@ -267,9 +272,3 @@ pub(super) async fn run_lane( } } } - -enum DriverEvent { - Incoming((Message, Option)), - Down(crate::web::session::PollResult), - Liveness, -} diff --git a/src/web/session/lifecycle.rs b/src/web/session/lifecycle.rs index fb9a80a..22cedf2 100644 --- a/src/web/session/lifecycle.rs +++ b/src/web/session/lifecycle.rs @@ -1,6 +1,10 @@ +use std::sync::Arc; use std::sync::atomic::Ordering; +use std::task::Waker; use std::time::{Duration, Instant}; +use tokio::sync::Notify; + use super::WebSession; /// Stable terminal cause assigned by the first session-close winner. @@ -104,6 +108,8 @@ struct ReleasedQueues { recovery_closed_before_commit: bool, reason: SessionCloseReason, peer_gap: Duration, + stream_wakers: Vec, + lane_notifies: Vec>, } /// Deferred queue release after manager publication linearizes a supersede. @@ -276,12 +282,13 @@ impl WebSession { if reason == SessionCloseReason::CarrierSuperseded { state.negotiation_phase = SessionNegotiationPhase::Superseded; } + let mut stream_wakers = Vec::with_capacity(state.streams.len().saturating_mul(2)); for stream in state.streams.values_mut() { if let Some(waker) = stream.read_waker.take() { - waker.wake(); + stream_wakers.push(waker); } if let Some(waker) = stream.write_waker.take() { - waker.wake(); + stream_wakers.push(waker); } } state.streams.clear(); @@ -296,8 +303,9 @@ impl WebSession { let mut lane_data_items = 0usize; let mut lane_control_bytes = 0usize; let mut lane_control_items = 0usize; + let mut lane_notifies = Vec::with_capacity(state.carrier_lanes.len()); for lane in state.carrier_lanes.values_mut() { - lane.notify.notify_waiters(); + lane_notifies.push(Arc::clone(&lane.notify)); if let Some(batch) = lane.unacked.take() { batch.lease.detach(); lane_data_bytes = lane_data_bytes.saturating_add(batch.data_bytes); @@ -326,10 +334,18 @@ impl WebSession { recovery_closed_before_commit, reason, peer_gap, + stream_wakers, + lane_notifies, } } fn finish_close(&self, released: ReleasedQueues) { + for waker in released.stream_wakers { + waker.wake(); + } + for notify in released.lane_notifies { + notify.notify_waiters(); + } self.cancel.cancel(); if self.carrier().is_multiplexed() { self.down_notify.notify_waiters(); diff --git a/src/web/session/uplink_tests.rs b/src/web/session/uplink_tests.rs index 69a46b6..97d0a3b 100644 --- a/src/web/session/uplink_tests.rs +++ b/src/web/session/uplink_tests.rs @@ -1,10 +1,13 @@ use super::*; use std::net::SocketAddr; +use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; +use std::task::{Wake, Waker}; use crate::config::{ WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig, }; use crate::web::manager::WebProcessRuntime; +use crate::web::session::SessionCloseOutcome; fn session() -> Arc { session_with_automatic(false) @@ -53,6 +56,80 @@ fn session_with_automatic(automatic: bool) -> Arc { ) } +struct SessionLockProbe { + session: std::sync::Weak, + lock_was_free: Arc, +} + +impl Wake for SessionLockProbe { + fn wake(self: Arc) { + if let Some(session) = self.session.upgrade() { + self.lock_was_free + .store(session.state.try_lock().is_some(), AtomicOrdering::Release); + } + } +} + +#[test] +fn close_wakes_stream_only_after_releasing_session_lock() { + let session = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(waker), + write_waker: None, + }, + ); + } + + assert_eq!( + session.close(SessionCloseReason::ApiClose), + SessionCloseOutcome::Closed + ); + assert!(lock_was_free.load(AtomicOrdering::Acquire)); +} + +#[test] +fn supersede_completion_defers_stream_wake_until_finish() { + let session = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(waker), + write_waker: None, + }, + ); + } + + assert!(session.begin_carrier_supersede()); + let completion = session.prepare_carrier_supersede().unwrap(); + assert!(!lock_was_free.load(AtomicOrdering::Acquire)); + completion.finish(); + assert!(lock_was_free.load(AtomicOrdering::Acquire)); +} + #[test] fn uplink_retry_commits_only_one_exact_body() { let session = session();