Process-wide concurrency + Cancellation ownership fixes

This commit is contained in:
Alexey
2026-09-23 22:30:51 +03:00
parent baa9bfbb01
commit f1107c21d9
47 changed files with 2177 additions and 762 deletions
+3 -3
View File
@@ -314,9 +314,9 @@ async fn recompute_connections_payload(
let mut active_users = 0usize; let mut active_users = 0usize;
for entry in shared.stats.iter_user_stats() { for entry in shared.stats.iter_user_stats() {
let user_stats = entry.value(); let user_stats = entry.value();
let current_connections = user_stats let current_connections = shared
.curr_connects .stats
.load(std::sync::atomic::Ordering::Relaxed); .get_process_user_curr_connects(entry.key());
let total_octets = user_stats let total_octets = user_stats
.octets_from_client .octets_from_client
.load(std::sync::atomic::Ordering::Relaxed) .load(std::sync::atomic::Ordering::Relaxed)
+1 -1
View File
@@ -71,7 +71,7 @@ pub(in crate::api) async fn users_from_config(
.filter(|limit| *limit > 0) .filter(|limit| *limit > 0)
.or((cfg.access.user_max_unique_ips_global_each > 0) .or((cfg.access.user_max_unique_ips_global_each > 0)
.then_some(cfg.access.user_max_unique_ips_global_each)), .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: active_ip_list.len(),
active_unique_ips_list: active_ip_list, active_unique_ips_list: active_ip_list,
recent_unique_ips: recent_ip_list.len(), recent_unique_ips: recent_ip_list.len(),
+18 -11
View File
@@ -1,5 +1,6 @@
use std::collections::BTreeSet; use std::collections::BTreeSet;
use std::net::IpAddr; use std::net::IpAddr;
use std::time::Duration;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tokio::process::Command; 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<String>) -> Result<(), String> { async fn run_command(binary: &str, args: &[&str], stdin: Option<String>) -> Result<(), String> {
const COMMAND_TIMEOUT: Duration = Duration::from_secs(30);
#[cfg(unix)] #[cfg(unix)]
let Some(command_path) = resolve_trusted_helper(binary) else { let Some(command_path) = resolve_trusted_helper(binary) else {
return Err(format!("{binary} is not available")); return Err(format!("{binary} is not available"));
@@ -377,21 +379,26 @@ async fn run_command(binary: &str, args: &[&str], stdin: Option<String>) -> Resu
} }
command.stdout(std::process::Stdio::null()); command.stdout(std::process::Stdio::null());
command.stderr(std::process::Stdio::piped()); command.stderr(std::process::Stdio::piped());
command.kill_on_drop(true);
let mut child = command let mut child = command
.spawn() .spawn()
.map_err(|error| format!("spawn {binary} failed: {error}"))?; .map_err(|error| format!("spawn {binary} failed: {error}"))?;
if let Some(blob) = stdin let output = tokio::time::timeout(COMMAND_TIMEOUT, async move {
&& let Some(mut writer) = child.stdin.take() if let Some(blob) = stdin
{ && let Some(mut writer) = child.stdin.take()
writer {
.write_all(blob.as_bytes()) writer
.write_all(blob.as_bytes())
.await
.map_err(|error| format!("stdin write {binary} failed: {error}"))?;
}
child
.wait_with_output()
.await .await
.map_err(|error| format!("stdin write {binary} failed: {error}"))?; .map_err(|error| format!("wait {binary} failed: {error}"))
} })
let output = child .await
.wait_with_output() .map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??;
.await
.map_err(|error| format!("wait {binary} failed: {error}"))?;
if output.status.success() { if output.status.success() {
return Ok(()); return Ok(());
} }
+35 -12
View File
@@ -38,11 +38,16 @@ struct UserIpShard {
#[derive(Debug, Default)] #[derive(Debug, Default)]
struct CleanupShard { struct CleanupShard {
queue: Mutex<HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>>, queue: Mutex<CleanupQueue>,
} }
type CleanupQueue =
HashMap<String, HashMap<UserIncarnation, HashMap<IpAddr, usize>>>;
type CleanupBatch = HashMap<(String, UserIncarnation, IpAddr), usize>;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
struct UserIpLimitPolicy { struct UserIpLimitPolicy {
source_generation: u64,
max_ips: Arc<HashMap<String, usize>>, max_ips: Arc<HashMap<String, usize>>,
default_max_ips: usize, default_max_ips: usize,
mode: UserMaxUniqueIpsMode, mode: UserMaxUniqueIpsMode,
@@ -52,6 +57,7 @@ struct UserIpLimitPolicy {
impl Default for UserIpLimitPolicy { impl Default for UserIpLimitPolicy {
fn default() -> Self { fn default() -> Self {
Self { Self {
source_generation: 0,
max_ips: Arc::new(HashMap::new()), max_ips: Arc::new(HashMap::new()),
default_max_ips: 0, default_max_ips: 0,
mode: UserMaxUniqueIpsMode::ActiveWindow, mode: UserMaxUniqueIpsMode::ActiveWindow,
@@ -70,6 +76,7 @@ pub struct UserIpTracker {
recent_cap_rejects: Arc<AtomicU64>, recent_cap_rejects: Arc<AtomicU64>,
cleanup_deferred_releases: Arc<AtomicU64>, cleanup_deferred_releases: Arc<AtomicU64>,
limit_policy: Arc<ArcSwap<UserIpLimitPolicy>>, limit_policy: Arc<ArcSwap<UserIpLimitPolicy>>,
policy_update: Arc<Mutex<()>>,
last_compact_epoch_secs: Arc<AtomicU64>, last_compact_epoch_secs: Arc<AtomicU64>,
cleanup_queue_len: Arc<AtomicU64>, cleanup_queue_len: Arc<AtomicU64>,
cleanup_shards: Arc<Box<[CleanupShard]>>, cleanup_shards: Arc<Box<[CleanupShard]>>,
@@ -121,6 +128,7 @@ impl UserIpTracker {
recent_cap_rejects: Arc::new(AtomicU64::new(0)), recent_cap_rejects: Arc::new(AtomicU64::new(0)),
cleanup_deferred_releases: Arc::new(AtomicU64::new(0)), cleanup_deferred_releases: Arc::new(AtomicU64::new(0)),
limit_policy: Arc::new(ArcSwap::from_pointee(UserIpLimitPolicy::default())), limit_policy: Arc::new(ArcSwap::from_pointee(UserIpLimitPolicy::default())),
policy_update: Arc::new(Mutex::new(())),
last_compact_epoch_secs: Arc::new(AtomicU64::new(0)), last_compact_epoch_secs: Arc::new(AtomicU64::new(0)),
cleanup_queue_len: Arc::new(AtomicU64::new(0)), cleanup_queue_len: Arc::new(AtomicU64::new(0)),
cleanup_shards: Arc::new(cleanup_shards), cleanup_shards: Arc::new(cleanup_shards),
@@ -196,19 +204,21 @@ impl UserIpTracker {
} }
pub(super) fn pop_one_cleanup( pub(super) fn pop_one_cleanup(
queue: &mut HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>, queue: &mut CleanupQueue,
) -> Option<(String, UserIncarnation, IpAddr, usize)> { ) -> Option<(String, UserIncarnation, IpAddr, usize)> {
let owner = queue.keys().next().cloned()?; let user = queue.keys().next().cloned()?;
let ip = queue.get(&owner)?.keys().next().copied()?; let incarnation = queue.get(&user)?.keys().next().copied()?;
let count = queue.get_mut(&owner)?.remove(&ip)?; let ip = queue.get(&user)?.get(&incarnation)?.keys().next().copied()?;
let remove_user = queue let incarnations = queue.get_mut(&user)?;
.get(&owner) let ips = incarnations.get_mut(&incarnation)?;
.map(|user_queue| user_queue.is_empty()) let count = ips.remove(&ip)?;
.unwrap_or(false); if ips.is_empty() {
if remove_user { incarnations.remove(&incarnation);
queue.remove(&owner);
} }
Some((owner.0, owner.1, ip, count)) if incarnations.is_empty() {
queue.remove(&user);
}
Some((user, incarnation, ip, count))
} }
#[cfg(test)] #[cfg(test)]
@@ -224,6 +234,19 @@ impl UserIpTracker {
#[cfg(not(test))] #[cfg(not(test))]
pub(super) fn observe_cleanup_poison_for_tests(&self) {} 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 { pub(super) fn now_epoch_secs() -> u64 {
std::time::SystemTime::now() std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
+65 -30
View File
@@ -2,47 +2,84 @@ use super::*;
impl UserIpTracker { impl UserIpTracker {
pub async fn set_limit_policy(&self, mode: UserMaxUniqueIpsMode, window_secs: u64) { pub async fn set_limit_policy(&self, mode: UserMaxUniqueIpsMode, window_secs: u64) {
self.limit_policy.rcu(|current| { let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| {
Arc::new(UserIpLimitPolicy { self.policy_update.clear_poison();
mode, poisoned.into_inner()
window_secs: window_secs.max(1),
..(**current).clone()
})
}); });
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) { pub async fn set_user_limit(&self, username: &str, max_ips: usize) {
let username = username.to_string(); let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| {
self.limit_policy.rcu(|current| { self.policy_update.clear_poison();
let mut limits = current.max_ips.as_ref().clone(); poisoned.into_inner()
limits.insert(username.clone(), max_ips);
Arc::new(UserIpLimitPolicy {
max_ips: Arc::new(limits),
..(**current).clone()
})
}); });
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) { pub async fn remove_user_limit(&self, username: &str) {
self.limit_policy.rcu(|current| { let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| {
let mut limits = current.max_ips.as_ref().clone(); self.policy_update.clear_poison();
limits.remove(username); poisoned.into_inner()
Arc::new(UserIpLimitPolicy {
max_ips: Arc::new(limits),
..(**current).clone()
})
}); });
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<String, usize>) { pub async fn load_limits(&self, default_limit: usize, limits: &HashMap<String, usize>) {
let limits = Arc::new(limits.clone()); let _policy_update = self.policy_update.lock().unwrap_or_else(|poisoned| {
self.limit_policy.rcu(|current| { self.policy_update.clear_poison();
Arc::new(UserIpLimitPolicy { poisoned.into_inner()
max_ips: Arc::clone(&limits),
default_max_ips: default_limit,
..(**current).clone()
})
}); });
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<String, usize>,
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( pub(super) fn prune_recent(
@@ -70,7 +107,6 @@ impl UserIpTracker {
ip: IpAddr, ip: IpAddr,
) -> Result<(), String> { ) -> Result<(), String> {
self.drain_cleanup_for_user(username).await; self.drain_cleanup_for_user(username).await;
self.maybe_compact_empty_users().await;
let policy = self.limit_policy.load(); let policy = self.limit_policy.load();
let limit = Self::user_limit(&policy, username); let limit = Self::user_limit(&policy, username);
let mode = policy.mode; let mode = policy.mode;
@@ -218,7 +254,6 @@ impl UserIpTracker {
incarnation: UserIncarnation, incarnation: UserIncarnation,
ip: IpAddr, ip: IpAddr,
) { ) {
self.maybe_compact_empty_users().await;
let shard_idx = Self::shard_idx(username); let shard_idx = Self::shard_idx(username);
let mut shard = self.shards[shard_idx].write().await; let mut shard = self.shards[shard_idx].write().await;
if shard.incarnations.get(username).copied() != Some(incarnation) { if shard.incarnations.get(username).copied() != Some(incarnation) {
+147 -46
View File
@@ -1,5 +1,68 @@
use super::*; 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<Item = (&(String, UserIncarnation, IpAddr), &usize)> {
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 { impl UserIpTracker {
/// Queues a deferred active IP cleanup for a later async drain. /// Queues a deferred active IP cleanup for a later async drain.
pub fn enqueue_cleanup(&self, user: String, ip: IpAddr) { pub fn enqueue_cleanup(&self, user: String, ip: IpAddr) {
@@ -18,8 +81,13 @@ impl UserIpTracker {
let cleanup_shard = &self.cleanup_shards[shard_idx]; let cleanup_shard = &self.cleanup_shards[shard_idx];
match cleanup_shard.queue.lock() { match cleanup_shard.queue.lock() {
Ok(mut queue) => { Ok(mut queue) => {
let user_queue = queue.entry((user, incarnation)).or_default(); let count = queue
let count = user_queue.entry(ip).or_insert(0); .entry(user)
.or_default()
.entry(incarnation)
.or_default()
.entry(ip)
.or_insert(0);
if *count == 0 { if *count == 0 {
self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed);
} }
@@ -29,8 +97,13 @@ impl UserIpTracker {
} }
Err(poisoned) => { Err(poisoned) => {
let mut queue = poisoned.into_inner(); let mut queue = poisoned.into_inner();
let user_queue = queue.entry((user.clone(), incarnation)).or_default(); let count = queue
let count = user_queue.entry(ip).or_insert(0); .entry(user.clone())
.or_default()
.entry(incarnation)
.or_default()
.entry(ip)
.or_insert(0);
if *count == 0 { if *count == 0 {
self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed);
} }
@@ -52,6 +125,26 @@ impl UserIpTracker {
self.cleanup_queue_len.load(Ordering::Relaxed) as usize 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::<usize>()
};
match cleanup_shard.queue.lock() {
Ok(queue) => count(&queue),
Err(poisoned) => count(&poisoned.into_inner()),
}
})
.sum()
}
#[cfg(test)] #[cfg(test)]
pub(crate) fn cleanup_queue_mutex_for_tests( pub(crate) fn cleanup_queue_mutex_for_tests(
&self, &self,
@@ -73,12 +166,13 @@ impl UserIpTracker {
return; return;
} }
let shard_idx = Self::shard_idx(user); 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 cleanup_shard = &self.cleanup_shards[shard_idx];
let to_remove = match cleanup_shard.queue.lock() { 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) => { Err(poisoned) => {
let mut queue = poisoned.into_inner(); 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(); cleanup_shard.queue.clear_poison();
drained drained
} }
@@ -86,25 +180,24 @@ impl UserIpTracker {
if to_remove.is_empty() { if to_remove.is_empty() {
return; return;
} }
let removed_queue_entries = to_remove
.iter()
.map(|(_, ips)| ips.len())
.sum::<usize>();
self.cleanup_queue_len
.fetch_sub(removed_queue_entries as u64, Ordering::Relaxed);
let mut shard = self.shards[shard_idx].write().await; let mut shard = self.shards[shard_idx].write().await;
let mut removed_active_entries = 0usize; let mut removed_active_entries = 0usize;
for (incarnation, ips) in to_remove { for ((queued_user, incarnation, ip), pending_count) in to_remove.entries() {
if shard.incarnations.get(user).copied() != Some(incarnation) { if shard.incarnations.get(queued_user).copied() != Some(*incarnation) {
continue; continue;
} }
for (ip, pending_count) in ips { removed_active_entries = removed_active_entries.saturating_add(
removed_active_entries = removed_active_entries.saturating_add( Self::apply_active_cleanup(
Self::apply_active_cleanup(&mut shard.active_ips, user, ip, pending_count), &mut shard.active_ips,
); queued_user,
} *ip,
*pending_count,
),
);
} }
Self::decrement_counter(&self.active_entry_count, removed_active_entries); 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) { pub(super) async fn drain_cleanup_shard(&self, shard_idx: usize) {
@@ -119,18 +212,20 @@ impl UserIpTracker {
if queue.is_empty() { if queue.is_empty() {
return; return;
} }
let mut drained = let mut drained = HashMap::with_capacity(CLEANUP_DRAIN_BATCH_LIMIT);
HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT));
for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT {
let Some((user, incarnation, ip, count)) = let Some((user, incarnation, ip, count)) =
Self::pop_one_cleanup(&mut queue) Self::pop_one_cleanup(&mut queue)
else { else {
break; break;
}; };
self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed);
drained.insert((user, incarnation, ip), count); drained.insert((user, incarnation, ip), count);
} }
drained DetachedCleanupBatch {
tracker: self,
shard_idx,
entries: drained,
}
} }
Err(poisoned) => { Err(poisoned) => {
let mut queue = poisoned.into_inner(); let mut queue = poisoned.into_inner();
@@ -138,55 +233,61 @@ impl UserIpTracker {
cleanup_shard.queue.clear_poison(); cleanup_shard.queue.clear_poison();
return; return;
} }
let mut drained = let mut drained = HashMap::with_capacity(CLEANUP_DRAIN_BATCH_LIMIT);
HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT));
for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT {
let Some((user, incarnation, ip, count)) = let Some((user, incarnation, ip, count)) =
Self::pop_one_cleanup(&mut queue) Self::pop_one_cleanup(&mut queue)
else { else {
break; break;
}; };
self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed);
drained.insert((user, incarnation, ip), count); drained.insert((user, incarnation, ip), count);
} }
cleanup_shard.queue.clear_poison(); cleanup_shard.queue.clear_poison();
drained DetachedCleanupBatch {
tracker: self,
shard_idx,
entries: drained,
}
} }
} }
}; };
drop(_drain_guard);
if to_remove.is_empty() { if to_remove.is_empty() {
return; return;
} }
let mut shard = self.shards[shard_idx].write().await; let mut shard = self.shards[shard_idx].write().await;
let mut removed_active_entries = 0usize; let mut removed_active_entries = 0usize;
for ((user, incarnation, ip), pending_count) in to_remove { for ((user, incarnation, ip), pending_count) in to_remove.entries() {
if shard.incarnations.get(&user).copied() != Some(incarnation) { if shard.incarnations.get(user).copied() != Some(*incarnation) {
continue; continue;
} }
removed_active_entries = removed_active_entries.saturating_add( 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); Self::decrement_counter(&self.active_entry_count, removed_active_entries);
drop(shard);
to_remove.commit();
} }
} }
fn drain_user_cleanup( fn detach_user_cleanup<'a>(
queue: &mut HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>, tracker: &'a UserIpTracker,
shard_idx: usize,
queue: &mut CleanupQueue,
user: &str, user: &str,
) -> Vec<(UserIncarnation, HashMap<IpAddr, usize>)> { ) -> DetachedCleanupBatch<'a> {
let owners = queue let mut entries = CleanupBatch::new();
.keys() if let Some(incarnations) = queue.remove(user) {
.filter(|(queued_user, _)| queued_user == user) for (incarnation, ips) in incarnations {
.cloned() for (ip, count) in ips {
.collect::<Vec<_>>(); entries.insert((user.to_string(), incarnation, ip), count);
owners }
.into_iter() }
.filter_map(|owner| { }
let incarnation = owner.1; DetachedCleanupBatch {
queue.remove(&owner).map(|ips| (incarnation, ips)) tracker,
}) shard_idx,
.collect() entries,
}
} }
+20 -5
View File
@@ -68,10 +68,13 @@ impl UserIpTracker {
} }
} }
pub async fn run_periodic_maintenance(self: Arc<Self>) { pub async fn run_periodic_maintenance(self: Arc<Self>, source_generation: u64) {
let mut interval = tokio::time::interval(Duration::from_secs(1)); let mut interval = tokio::time::interval(Duration::from_secs(1));
loop { loop {
interval.tick().await; interval.tick().await;
if self.limit_policy.load().source_generation != source_generation {
continue;
}
self.drain_cleanup_queue().await; self.drain_cleanup_queue().await;
self.maybe_compact_empty_users().await; self.maybe_compact_empty_users().await;
} }
@@ -263,6 +266,10 @@ impl UserIpTracker {
} }
pub async fn clear_all(&self) { 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() { for shard_lock in self.shards.iter() {
let mut shard = shard_lock.write().await; let mut shard = shard_lock.write().await;
shard.active_ips.clear(); shard.active_ips.clear();
@@ -271,16 +278,24 @@ impl UserIpTracker {
} }
self.active_entry_count.store(0, Ordering::Relaxed); self.active_entry_count.store(0, Ordering::Relaxed);
self.recent_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() { for cleanup_shard in self.cleanup_shards.iter() {
match cleanup_shard.queue.lock() { let queue = match cleanup_shard.queue.lock() {
Ok(mut queue) => queue.clear(), Ok(queue) => queue,
Err(poisoned) => { Err(poisoned) => {
poisoned.into_inner().clear(); let queue = poisoned.into_inner();
cleanup_shard.queue.clear_poison(); 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); 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 { pub async fn is_ip_active(&self, username: &str, ip: IpAddr) -> bool {
+79 -4
View File
@@ -3,6 +3,8 @@ use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
mod cleanup_invariants;
fn test_ipv4(oct1: u8, oct2: u8, oct3: u8, oct4: u8) -> IpAddr { fn test_ipv4(oct1: u8, oct2: u8, oct3: u8, oct4: u8) -> IpAddr {
IpAddr::V4(Ipv4Addr::new(oct1, oct2, oct3, oct4)) 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)); 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)] #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_policy_replacement_never_exposes_partial_limit_map() { async fn concurrent_policy_replacement_never_exposes_partial_limit_map() {
const USER_COUNT: usize = 4_096; 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.last_compact_epoch_secs.store(0, Ordering::Relaxed);
tracker tracker.maybe_compact_empty_users().await;
.check_and_add("trigger-user", test_ipv4(10, 3, 0, 2))
.await
.unwrap();
let shard_idx = UserIpTracker::shard_idx(&stale_user); let shard_idx = UserIpTracker::shard_idx(&stale_user);
let shard = tracker.shards[shard_idx].read().await; let shard = tracker.shards[shard_idx].read().await;
+107
View File
@@ -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);
}
+25 -6
View File
@@ -22,6 +22,11 @@ use crate::tls_front::TlsFrontCache;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; 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 SESSION_STOP_TIMEOUT: Duration = Duration::from_secs(5);
const BACKGROUND_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); const SESSION_ADMISSION_CLOSED: usize = 1 << (usize::BITS - 1);
@@ -145,12 +150,17 @@ impl RuntimeTaskScope {
self.cancel.clone() self.cancel.clone()
} }
/// Cancels the scope and waits within the bounded background-task budget. /// Synchronously closes task admission and signals every tracked task.
pub(crate) async fn stop(&self) { pub(crate) fn begin_stop(&self) {
self.admission.close(); self.admission.close();
self.admission.wait_for_registrations().await;
self.cancel.cancel(); self.cancel.cancel();
self.tracker.close(); 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; 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. /// Waits for registered sessions and cancels them when the deadline expires.
pub(crate) async fn drain_sessions(&self, timeout: Duration) -> bool { pub(crate) async fn drain_sessions(&self, timeout: Duration) -> bool {
self.stop_accepting_sessions(); self.stop_accepting_sessions();
let mut cancellation_guard = SessionDrainCancellationGuard::new(self);
self.session_admission.wait_for_registrations().await; self.session_admission.wait_for_registrations().await;
self.sessions.close(); self.sessions.close();
if tokio::time::timeout(timeout, self.sessions.wait()) if tokio::time::timeout(timeout, self.sessions.wait())
.await .await
.is_ok() .is_ok()
{ {
cancellation_guard.disarm();
return true; return true;
} }
cancellation_guard.disarm();
self.stop_sessions().await; self.stop_sessions().await;
false false
} }
/// Cancels all sessions and waits within the bounded session-stop budget. fn begin_stop_sessions(&self) {
pub(crate) async fn stop_sessions(&self) {
self.stop_accepting_sessions(); self.stop_accepting_sessions();
self.session_admission.wait_for_registrations().await;
self.session_cancel.cancel(); self.session_cancel.cancel();
self.sessions.close(); 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; let _ = tokio::time::timeout(SESSION_STOP_TIMEOUT, self.sessions.wait()).await;
} }
@@ -335,6 +352,8 @@ impl RuntimeGeneration {
impl Drop for RuntimeGeneration { impl Drop for RuntimeGeneration {
fn drop(&mut self) { fn drop(&mut self) {
self.background_tasks.begin_stop();
self.begin_stop_sessions();
if let Some(pool) = self.me_pool.as_ref() { if let Some(pool) = self.me_pool.as_ref() {
pool.begin_shutdown(); pool.begin_shutdown();
} }
+157
View File
@@ -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<Notify>);
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");
}
}
+30 -13
View File
@@ -3,7 +3,7 @@ use std::net::{IpAddr, SocketAddr};
use std::sync::Arc; use std::sync::Arc;
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
use tokio::sync::{RwLock, watch}; use tokio::sync::{RwLock, Semaphore, watch};
use tracing::{error, info}; use tracing::{error, info};
use crate::api; 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::direct_buffer_budget::{DirectBufferBudget, resolve_direct_buffer_hard_limit};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::traffic_limiter::TrafficLimiter;
use crate::proxy::user_admission::UserAdmissionAuthority; use crate::proxy::user_admission::UserAdmissionAuthority;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
use crate::startup::{COMPONENT_API_BOOTSTRAP, COMPONENT_NETWORK_PROBE}; use crate::startup::{COMPONENT_API_BOOTSTRAP, COMPONENT_NETWORK_PROBE};
use crate::stats::telemetry::TelemetryPolicy; use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{QuotaStore, Stats}; use crate::stats::{QuotaStore, Stats};
@@ -47,10 +49,16 @@ pub(super) async fn run_telemt_core(
} = bootstrap::bootstrap(privilege_drop_requested).await?; } = bootstrap::bootstrap(privilege_drop_requested).await?;
let quota_store = Arc::new(QuotaStore::default()); 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 tls_full_cert_budget = Arc::new(TlsFullCertBudget::new());
let process_control_plane = control_plane::ProcessControlPlane::new(); let process_control_plane = control_plane::ProcessControlPlane::new();
let runtime_task_scope = generation::RuntimeTaskScope::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)); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
let quota_state_path = config.general.quota_state_path.clone(); let quota_state_path = config.general.quota_state_path.clone();
let quota_state = let quota_state =
@@ -72,14 +80,11 @@ pub(super) async fn run_telemt_core(
.with_dns_overrides(&config.network.dns_overrides)?, .with_dns_overrides(&config.network.dns_overrides)?,
); );
let ip_tracker = Arc::new(UserIpTracker::new()); let ip_tracker = Arc::new(UserIpTracker::new());
ip_tracker let _ = ip_tracker
.load_limits( .apply_policy_from_source(
1,
config.access.user_max_unique_ips_global_each, config.access.user_max_unique_ips_global_each,
&config.access.user_max_unique_ips, &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_mode,
config.access.user_max_unique_ips_window_secs, config.access.user_max_unique_ips_window_secs,
) )
@@ -102,14 +107,22 @@ pub(super) async fn run_telemt_core(
let direct_buffer_hard_limit = let direct_buffer_hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await; resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(direct_buffer_hard_limit); let direct_buffer_budget = DirectBufferBudget::new(direct_buffer_hard_limit);
direct_buffer_budget.activate_controller(1);
info!( info!(
hard_limit_bytes = direct_buffer_hard_limit, hard_limit_bytes = direct_buffer_hard_limit,
configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes, configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes,
"Direct relay buffer budget initialized" "Direct relay buffer budget initialized"
); );
let user_admission = UserAdmissionAuthority::new_with_quota_store(quota_store.clone()); 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(), direct_buffer_budget.clone(),
traffic_limiter,
user_admission, user_admission,
); );
let _ = shared_state.activate_user_config_source( let _ = shared_state.activate_user_config_source(
@@ -118,10 +131,12 @@ pub(super) async fn run_telemt_core(
&config.access.users, &config.access.users,
&config.access.user_enabled, &config.access.user_enabled,
); );
shared_state.traffic_limiter.apply_policy( let max_connections_limit = if config.server.max_connections == 0 {
config.access.user_rate_limits.clone(), Semaphore::MAX_PERMITS
config.access.cidr_rate_limits.clone(), } 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_trace = WebTraceStore::new(config.web.debug.clone(), &config.web.limits);
let web_runtime_control = WebRuntimeControl::new(); let web_runtime_control = WebRuntimeControl::new();
@@ -303,6 +318,7 @@ pub(super) async fn run_telemt_core(
ip_tracker.clone(), ip_tracker.clone(),
shared_state.clone(), shared_state.clone(),
direct_buffer_budget, direct_buffer_budget,
max_connections,
route_runtime.clone(), route_runtime.clone(),
api_me_pool.clone(), api_me_pool.clone(),
runtime_task_scope.clone(), runtime_task_scope.clone(),
@@ -333,6 +349,7 @@ pub(super) async fn run_telemt_core(
runtime.max_connections, runtime.max_connections,
runtime_task_scope, runtime_task_scope,
); );
runtime_task_scope_guard.disarm();
let active_runtime = Arc::new(ArcSwap::from(runtime_generation)); let active_runtime = Arc::new(ArcSwap::from(runtime_generation));
let bound = listeners::bind_listeners( let bound = listeners::bind_listeners(
&runtime.config, &runtime.config,
+34 -7
View File
@@ -175,9 +175,14 @@ impl ReloadSupervisor {
resolved.effective, resolved.effective,
&self.config_path, &self.config_path,
self.quota_store.clone(), self.quota_store.clone(),
old_runtime.stats.connection_authority(),
self.runtime_log_filter.clone(), self.runtime_log_filter.clone(),
self.tls_full_cert_budget.clone(), self.tls_full_cert_budget.clone(),
old_runtime.proxy_shared.user_admission(), 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 .await
{ {
@@ -300,15 +305,37 @@ impl ReloadSupervisor {
} else { } else {
None 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 replaced = {
let listener_manager = self.listener_manager.lock().await; 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(); old_runtime.stop_accepting_sessions();
listener_manager.activate_runtime_generation(new_runtime.clone()) listener_manager.activate_runtime_generation(new_runtime.clone())
}; };
+32 -31
View File
@@ -12,11 +12,13 @@ use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker; use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, run_probe}; use crate::network::probe::{decide_network_capabilities, run_probe};
use crate::proxy::direct_buffer_budget::{ 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::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::traffic_limiter::TrafficLimiter;
use crate::proxy::user_admission::UserAdmissionAuthority; use crate::proxy::user_admission::UserAdmissionAuthority;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
use crate::startup::StartupTracker; use crate::startup::StartupTracker;
use crate::stats::beobachten::BeobachtenStore; use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy; use crate::stats::telemetry::TelemetryPolicy;
@@ -27,7 +29,9 @@ use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
use super::admission; use super::admission;
use super::generation::{RuntimeGeneration, RuntimeTaskScope}; use super::generation::{
RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard,
};
use super::listeners::listener_rebind_supported; use super::listeners::listener_rebind_supported;
use super::runtime_tasks::RuntimeLogFilter; use super::runtime_tasks::RuntimeLogFilter;
use super::{me_startup, runtime_tasks, tls_bootstrap}; use super::{me_startup, runtime_tasks, tls_bootstrap};
@@ -49,9 +53,14 @@ pub(crate) async fn prepare_runtime(
config: ProxyConfig, config: ProxyConfig,
config_path: &Path, config_path: &Path,
quota_store: Arc<QuotaStore>, quota_store: Arc<QuotaStore>,
connection_authority: Arc<UserConnectionAuthority>,
runtime_log_filter: RuntimeLogFilter, runtime_log_filter: RuntimeLogFilter,
tls_full_cert_budget: Arc<TlsFullCertBudget>, tls_full_cert_budget: Arc<TlsFullCertBudget>,
user_admission: Arc<UserAdmissionAuthority>, user_admission: Arc<UserAdmissionAuthority>,
ip_tracker: Arc<UserIpTracker>,
traffic_limiter: Arc<TrafficLimiter>,
direct_buffer_budget: Arc<DirectBufferBudget>,
max_connections: Arc<Semaphore>,
) -> Result<PreparedRuntime, String> { ) -> Result<PreparedRuntime, String> {
let user_admission_epoch = user_admission.epoch(); let user_admission_epoch = user_admission.epoch();
config config
@@ -63,7 +72,11 @@ pub(crate) async fn prepare_runtime(
.as_secs(); .as_secs();
let startup_tracker = Arc::new(StartupTracker::new(started_at_epoch_secs)); let startup_tracker = Arc::new(StartupTracker::new(started_at_epoch_secs));
let task_scope = RuntimeTaskScope::new(); 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)); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
let upstream_manager = Arc::new( let upstream_manager = Arc::new(
@@ -80,31 +93,11 @@ pub(crate) async fn prepare_runtime(
.with_dns_overrides(&config.network.dns_overrides) .with_dns_overrides(&config.network.dns_overrides)
.map_err(|error| format!("DNS override preparation failed: {}", error))?, .map_err(|error| format!("DNS override preparation failed: {}", error))?,
); );
let ip_tracker = Arc::new(UserIpTracker::new()); let proxy_shared = ProxySharedState::new_with_process_authorities(
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(
direct_buffer_budget.clone(), direct_buffer_budget.clone(),
traffic_limiter,
user_admission, user_admission,
); );
proxy_shared.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(),
);
let probe = run_probe( let probe = run_probe(
&config.network, &config.network,
@@ -184,12 +177,6 @@ pub(crate) async fn prepare_runtime(
Duration::from_secs(config.access.replay_window_secs), Duration::from_secs(config.access.replay_window_secs),
)); ));
let buffer_pool = Arc::new(BufferPool::with_config(64 * 1024, 4096)); 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 (config_watcher_activation, config_watcher_activation_rx) = watch::channel(false);
let watches = runtime_tasks::spawn_runtime_tasks( let watches = runtime_tasks::spawn_runtime_tasks(
generation_id, generation_id,
@@ -288,10 +275,12 @@ pub(crate) async fn prepare_runtime(
conntrack_scope.cancellation_token(), conntrack_scope.cancellation_token(),
)); ));
task_scope.spawn(run_direct_buffer_budget_controller( task_scope.spawn(run_direct_buffer_budget_controller(
generation_id,
direct_buffer_budget, direct_buffer_budget,
buffer_pool.clone(), buffer_pool.clone(),
stats.clone(), stats.clone(),
proxy_shared.clone(), proxy_shared.clone(),
max_connections.clone(),
config.server.max_connections, config.server.max_connections,
)); ));
let generation = RuntimeGeneration::new( let generation = RuntimeGeneration::new(
@@ -313,6 +302,7 @@ pub(crate) async fn prepare_runtime(
max_connections, max_connections,
task_scope, task_scope,
); );
task_scope_guard.disarm();
drop(admission_tx); drop(admission_tx);
Ok(PreparedRuntime { Ok(PreparedRuntime {
@@ -414,6 +404,17 @@ pub(crate) fn resolve_reload_config(
effective.server.metrics_listen = old.server.metrics_listen.clone(); effective.server.metrics_listen = old.server.metrics_listen.clone();
effective.server.metrics_port = old.server.metrics_port; 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 { if old.general.quota_state_path != desired.general.quota_state_path {
fields.push("general.quota_state_path".to_string()); fields.push("general.quota_state_path".to_string());
effective.general.quota_state_path = old.general.quota_state_path.clone(); effective.general.quota_state_path = old.general.quota_state_path.clone();
+33
View File
@@ -94,6 +94,39 @@ fn global_mss_profiles_are_deferred_with_the_listener_socket_group() {
assert!(!resolved.runtime_changed); 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] #[test]
fn mixed_reload_retains_process_state_and_applies_runtime_state() { fn mixed_reload_retains_process_state_and_applies_runtime_state() {
let old = ProxyConfig::default(); let old = ProxyConfig::default();
+3 -7
View File
@@ -56,6 +56,7 @@ pub(super) async fn prepare_runtime(
ip_tracker: Arc<UserIpTracker>, ip_tracker: Arc<UserIpTracker>,
shared_state: Arc<ProxySharedState>, shared_state: Arc<ProxySharedState>,
direct_buffer_budget: Arc<DirectBufferBudget>, direct_buffer_budget: Arc<DirectBufferBudget>,
max_connections: Arc<Semaphore>,
route_runtime: Arc<RouteRuntimeController>, route_runtime: Arc<RouteRuntimeController>,
api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>, api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>,
runtime_task_scope: RuntimeTaskScope, runtime_task_scope: RuntimeTaskScope,
@@ -69,13 +70,6 @@ pub(super) async fn prepare_runtime(
let beobachten = Arc::new(BeobachtenStore::new()); let beobachten = Arc::new(BeobachtenStore::new());
let rng = Arc::new(SecureRandom::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 me2dc_fallback = config.general.me2dc_fallback;
let me_init_retry_attempts = config.general.me_init_retry_attempts; let me_init_retry_attempts = config.general.me_init_retry_attempts;
if use_middle_proxy && !decision.ipv4_me && !decision.ipv6_me { if use_middle_proxy && !decision.ipv4_me && !decision.ipv6_me {
@@ -346,10 +340,12 @@ pub(super) async fn prepare_runtime(
conntrack_scope.cancellation_token(), conntrack_scope.cancellation_token(),
)); ));
runtime_task_scope.spawn(run_direct_buffer_budget_controller( runtime_task_scope.spawn(run_direct_buffer_budget_controller(
1,
direct_buffer_budget, direct_buffer_budget,
buffer_pool.clone(), buffer_pool.clone(),
stats, stats,
shared_state, shared_state,
max_connections.clone(),
config.server.max_connections, config.server.max_connections,
)); ));
+16 -37
View File
@@ -138,7 +138,9 @@ pub(crate) async fn spawn_runtime_tasks(
let ip_tracker_maintenance = ip_tracker.clone(); let ip_tracker_maintenance = ip_tracker.clone();
task_scope.spawn(async move { 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<IpAddr> = probe.detected_ipv4.map(IpAddr::V4); let detected_ip_v4: Option<IpAddr> = 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 ip_tracker_policy = ip_tracker.clone();
let mut config_rx_ip_limits = config_rx.clone(); let mut config_rx_ip_limits = config_rx.clone();
task_scope.spawn(async move { task_scope.spawn(async move {
let mut prev_limits = config_rx_ip_limits let mut previous = config_rx_ip_limits.borrow().access.clone();
.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;
loop { loop {
if config_rx_ip_limits.changed().await.is_err() { 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(); let cfg = config_rx_ip_limits.borrow_and_update().clone();
if prev_limits != cfg.access.user_max_unique_ips if previous.user_max_unique_ips != cfg.access.user_max_unique_ips
|| prev_global_each != cfg.access.user_max_unique_ips_global_each || 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 let _ = ip_tracker_policy
.load_limits( .apply_policy_from_source(
generation_id,
cfg.access.user_max_unique_ips_global_each, cfg.access.user_max_unique_ips_global_each,
&cfg.access.user_max_unique_ips, &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_mode,
cfg.access.user_max_unique_ips_window_secs, cfg.access.user_max_unique_ips_window_secs,
) )
.await; .await;
prev_mode = cfg.access.user_max_unique_ips_mode; previous = cfg.access.clone();
prev_window = cfg.access.user_max_unique_ips_window_secs;
} }
} }
}); });
let limiter = shared_state.traffic_limiter.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(); let mut config_rx_rate_limits = config_rx.clone();
task_scope.spawn(async move { task_scope.spawn(async move {
let mut prev_user_limits = config_rx_rate_limits 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 if prev_user_limits != cfg.access.user_rate_limits
|| prev_cidr_limits != cfg.access.cidr_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.user_rate_limits.clone(),
cfg.access.cidr_rate_limits.clone(), cfg.access.cidr_rate_limits.clone(),
); );
+1 -1
View File
@@ -153,7 +153,7 @@ pub(super) async fn render(
out, out,
"telemt_user_connections_current{{user=\"{}\"}} {}", "telemt_user_connections_current{{user=\"{}\"}} {}",
user, user,
s.curr_connects.load(std::sync::atomic::Ordering::Relaxed) stats.get_process_user_curr_connects(user)
); );
let _ = writeln!( let _ = writeln!(
out, out,
+4
View File
@@ -69,6 +69,10 @@ async fn test_render_metrics_format() {
stats.increment_me_endpoint_quarantine_draining_suppressed_total(); stats.increment_me_endpoint_quarantine_draining_suppressed_total();
stats.increment_user_connects("alice"); stats.increment_user_connects("alice");
stats.increment_user_curr_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_from("alice", 1024);
stats.add_user_octets_to("alice", 2048); stats.add_user_octets_to("alice", 2048);
stats.increment_user_msgs_from("alice"); stats.increment_user_msgs_from("alice");
+87 -47
View File
@@ -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::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState}; use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState};
use crate::proxy::user_admission::UserIncarnation; 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::stream::{BufferPool, CryptoReader, CryptoWriter};
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
@@ -238,13 +239,63 @@ where
/// Owns one authenticated user's connection and source-IP admission slots. /// Owns one authenticated user's connection and source-IP admission slots.
pub(crate) struct UserConnectionReservation { pub(crate) struct UserConnectionReservation {
stats: Arc<Stats>, stats: Arc<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
incarnation: UserIncarnation,
quota_handle: UserQuotaHandle, quota_handle: UserQuotaHandle,
tracks_ip: bool, _connection_permit: UserConnectionPermit,
active: bool, _stats_observation: Option<UserConnectionObservation>,
ip_permit: Option<UserIpPermit>,
released: bool,
}
struct UserIpPermit {
tracker: Arc<UserIpTracker>,
owner: Option<UserIpOwner>,
}
struct UserIpOwner {
user: String,
incarnation: UserIncarnation,
ip: IpAddr,
}
impl UserIpPermit {
fn new(
tracker: Arc<UserIpTracker>,
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 { impl UserConnectionReservation {
@@ -257,6 +308,11 @@ impl UserConnectionReservation {
tracks_ip: bool, tracks_ip: bool,
) -> Self { ) -> Self {
let quota_handle = stats.current_user_quota_handle(&user); 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( Self::new_for_incarnation(
stats, stats,
ip_tracker, ip_tracker,
@@ -264,6 +320,8 @@ impl UserConnectionReservation {
ip, ip,
0, 0,
quota_handle, quota_handle,
connection_permit,
stats_observation,
tracks_ip, tracks_ip,
) )
} }
@@ -276,17 +334,20 @@ impl UserConnectionReservation {
ip: IpAddr, ip: IpAddr,
incarnation: UserIncarnation, incarnation: UserIncarnation,
quota_handle: UserQuotaHandle, quota_handle: UserQuotaHandle,
connection_permit: UserConnectionPermit,
stats_observation: Option<UserConnectionObservation>,
tracks_ip: bool, tracks_ip: bool,
) -> Self { ) -> Self {
let ip_permit = tracks_ip.then(|| {
UserIpPermit::new(ip_tracker, user, incarnation, ip)
});
Self { Self {
stats, stats,
ip_tracker,
user,
ip,
incarnation,
quota_handle, quota_handle,
tracks_ip, _connection_permit: connection_permit,
active: true, _stats_observation: stats_observation,
ip_permit,
released: false,
} }
} }
@@ -297,50 +358,24 @@ impl UserConnectionReservation {
/// Releases both admission counters through the asynchronous cleanup path. /// Releases both admission counters through the asynchronous cleanup path.
pub(crate) async fn release(mut self) { pub(crate) async fn release(mut self) {
if !self.active { if let Some(ip_permit) = self.ip_permit.take() {
return; ip_permit.release().await;
} }
self.active = false; self.released = true;
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);
} }
/// Defers IP cleanup when admission fails after the asynchronous reservation step. /// Defers IP cleanup when admission fails after the asynchronous reservation step.
pub(crate) fn release_deferred(mut self) { pub(crate) fn release_deferred(mut self) {
if !self.active { self.released = true;
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,
);
}
} }
} }
impl Drop for UserConnectionReservation { impl Drop for UserConnectionReservation {
fn drop(&mut self) { fn drop(&mut self) {
if !self.active { if self.released {
return; return;
} }
self.active = false;
self.stats.increment_session_drop_fallback_total(); 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) .or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each)) .then_some(config.access.user_max_tcp_conns_global_each))
.map(|value| value as u64); .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 { return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(), user: user.to_string(),
}); });
} };
let stats_observation = stats.observe_user_current_connection(user);
if let Err(reason) = ip_tracker if let Err(reason) = ip_tracker
.check_and_add_for_incarnation(user, incarnation, peer_addr.ip()) .check_and_add_for_incarnation(user, incarnation, peer_addr.ip())
.await .await
{ {
stats.decrement_user_curr_connects(user);
warn!( warn!(
user = %user, user = %user,
ip = %peer_addr.ip(), ip = %peer_addr.ip(),
@@ -429,6 +467,8 @@ async fn acquire_user_connection_reservation_for_incarnation(
peer_addr.ip(), peer_addr.ip(),
incarnation, incarnation,
quota_handle, quota_handle,
connection_permit,
stats_observation,
true, true,
)) ))
} }
+5 -5
View File
@@ -155,18 +155,20 @@ impl RunningClientHandler {
.or((config.access.user_max_tcp_conns_global_each > 0) .or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each)) .then_some(config.access.user_max_tcp_conns_global_each))
.map(|v| v as u64); .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 { return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(), user: user.to_string(),
}); });
} };
match ip_tracker.check_and_add(user, peer_addr.ip()).await { match ip_tracker.check_and_add(user, peer_addr.ip()).await {
Ok(()) => { Ok(()) => {
ip_tracker.remove_ip(user, peer_addr.ip()).await; ip_tracker.remove_ip(user, peer_addr.ip()).await;
} }
Err(reason) => { Err(reason) => {
stats.decrement_user_curr_connects(user);
warn!( warn!(
user = %user, user = %user,
ip = %peer_addr.ip(), ip = %peer_addr.ip(),
@@ -178,8 +180,6 @@ impl RunningClientHandler {
}); });
} }
} }
stats.decrement_user_curr_connects(user);
Ok(()) Ok(())
} }
} }
+28 -220
View File
@@ -2,12 +2,16 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
use parking_lot::{Mutex as ParkingMutex, MutexGuard as ParkingMutexGuard};
use tokio::sync::watch; use tokio::sync::watch;
use crate::stats::Stats; // Process controller and system-memory sampling remain outside data-plane accounting.
use crate::stream::BufferPool; mod controller;
pub(crate) use controller::{
use super::shared_state::ProxySharedState; 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. /// Accounting granularity for process-wide Direct copy-buffer reservations.
pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024; pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024;
@@ -70,6 +74,8 @@ pub(crate) struct DirectBufferBudget {
hard_limit_bytes: u64, hard_limit_bytes: u64,
target_bytes: AtomicU64, target_bytes: AtomicU64,
reserved_bytes: AtomicU64, reserved_bytes: AtomicU64,
active_controller_generation: AtomicU64,
controller_update: ParkingMutex<()>,
pressure_generation: AtomicU64, pressure_generation: AtomicU64,
pressure_tx: watch::Sender<u64>, pressure_tx: watch::Sender<u64>,
memory_total_bytes: AtomicU64, memory_total_bytes: AtomicU64,
@@ -94,6 +100,8 @@ impl DirectBufferBudget {
hard_limit_bytes, hard_limit_bytes,
target_bytes: AtomicU64::new(hard_limit_bytes), target_bytes: AtomicU64::new(hard_limit_bytes),
reserved_bytes: AtomicU64::new(0), reserved_bytes: AtomicU64::new(0),
active_controller_generation: AtomicU64::new(0),
controller_update: ParkingMutex::new(()),
pressure_generation: AtomicU64::new(0), pressure_generation: AtomicU64::new(0),
pressure_tx, pressure_tx,
memory_total_bytes: AtomicU64::new(0), memory_total_bytes: AtomicU64::new(0),
@@ -120,6 +128,22 @@ impl DirectBufferBudget {
self.pressure_tx.subscribe() 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<ParkingMutexGuard<'_, ()>> {
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. /// Reserves bytes against either the adaptive target or the absolute ceiling.
pub(crate) fn try_reserve( pub(crate) fn try_reserve(
self: &Arc<Self>, self: &Arc<Self>,
@@ -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<DirectBufferBudget>,
buffer_pool: Arc<BufferPool>,
stats: Arc<Stats>,
shared: Arc<ProxySharedState>,
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<u8> {
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<u8> {
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<u64> {
let raw = tokio::fs::read_to_string(path).await.ok()?;
let raw = raw.trim();
if raw == "max" {
return None;
}
let value = raw.parse::<u64>().ok()?;
(value < (1u64 << 60)).then_some(value)
}
#[cfg(target_os = "linux")]
async fn read_u64_file(path: &str) -> Option<u64> {
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::<u64>().ok()
})
.unwrap_or(0)
.saturating_mul(1024)
}
fn align_up(bytes: usize) -> usize { fn align_up(bytes: usize) -> usize {
bytes bytes
.div_ceil(DIRECT_BUFFER_UNIT_BYTES) .div_ceil(DIRECT_BUFFER_UNIT_BYTES)
@@ -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<DirectBufferBudget>,
buffer_pool: Arc<BufferPool>,
stats: Arc<Stats>,
shared: Arc<ProxySharedState>,
connection_slots: Arc<Semaphore>,
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<u8> {
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<u8> {
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<u64> {
let raw = tokio::fs::read_to_string(path).await.ok()?;
let raw = raw.trim();
if raw == "max" {
return None;
}
let value = raw.parse::<u64>().ok()?;
(value < (1u64 << 60)).then_some(value)
}
#[cfg(target_os = "linux")]
async fn read_u64_file(path: &str) -> Option<u64> {
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::<u64>().ok()
})
.unwrap_or(0)
.saturating_mul(1024)
}
+2 -1
View File
@@ -15,6 +15,7 @@ use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch}; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch};
use tokio::time::timeout; use tokio::time::timeout;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use tokio_util::task::AbortOnDropHandle;
use tracing::{debug, info, trace, warn}; use tracing::{debug, info, trace, warn};
use crate::config::{ConntrackPressureProfile, ProxyConfig}; 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_FLUSH_BATCH_MAX_BYTES_MIN: usize = 4096;
const ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR: usize = 2; const ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR: usize = 2;
const ME_D2C_SINGLE_WRITE_COALESCE_MAX_BYTES: usize = 128 * 1024; 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_MIN_MS: u64 = 1;
const QUOTA_RESERVE_BACKOFF_MAX_MS: u64 = 16; const QUOTA_RESERVE_BACKOFF_MAX_MS: u64 = 16;
const QUOTA_RESERVE_MAX_BACKOFF_ROUNDS: usize = 16; const QUOTA_RESERVE_MAX_BACKOFF_ROUNDS: usize = 16;
+1 -2
View File
@@ -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_ms = QUOTA_RESERVE_BACKOFF_MIN_MS;
let mut backoff_rounds = 0usize; let mut backoff_rounds = 0usize;
loop { loop {
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { for _ in 0..QUOTA_RESERVE_ATTEMPTS_PER_ROUND {
match quota_handle.try_reserve(bytes, limit) { match quota_handle.try_reserve(bytes, limit) {
Ok(reservation) => return Ok(reservation.commit()), Ok(reservation) => return Ok(reservation.commit()),
Err(QuotaReserveError::LimitExceeded) => { Err(QuotaReserveError::LimitExceeded) => {
@@ -30,7 +30,6 @@ pub(super) async fn reserve_user_quota_with_yield(
} }
Err(QuotaReserveError::Contended) => { Err(QuotaReserveError::Contended) => {
stats.increment_quota_contention_total(); stats.increment_quota_contention_total();
std::hint::spin_loop();
} }
} }
} }
+21 -9
View File
@@ -2,11 +2,15 @@ use super::*;
// Bounded C2ME sender and downstream writer tasks. // Bounded C2ME sender and downstream writer tasks.
mod tasks; mod tasks;
// Child-task ownership aborts relay tasks when the parent future is cancelled.
mod children;
// Conntrack close classification. // Conntrack close classification.
mod close_reason; mod close_reason;
use children::RelayChildTasks;
use close_reason::classify_conntrack_close_reason; use close_reason::classify_conntrack_close_reason;
use tasks::{run_c2me_sender, run_me_writer}; use tasks::{run_c2me_sender, run_me_writer};
struct RelayConnLease { struct RelayConnLease {
connection: Option<ConnLease>, connection: Option<ConnLease>,
conn_id: u64, conn_id: u64,
@@ -185,7 +189,7 @@ where
let c2me_byte_semaphore = Arc::new(Semaphore::new(c2me_byte_budget)); let c2me_byte_semaphore = Arc::new(Semaphore::new(c2me_byte_budget));
let (c2me_tx, c2me_rx) = mpsc::channel::<C2MeCommand>(c2me_channel_capacity); let (c2me_tx, c2me_rx) = mpsc::channel::<C2MeCommand>(c2me_channel_capacity);
let me_pool_c2me = me_pool.clone(); 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, c2me_rx,
me_pool_c2me, me_pool_c2me,
conn_id, conn_id,
@@ -193,7 +197,7 @@ where
peer, peer,
translated_local_addr, translated_local_addr,
effective_tag_array, effective_tag_array,
)); )));
let (stop_tx, stop_rx) = oneshot::channel::<()>(); let (stop_tx, stop_rx) = oneshot::channel::<()>();
let flow_cancel = CancellationToken::new(); let flow_cancel = CancellationToken::new();
@@ -208,7 +212,7 @@ where
let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone(); let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone();
let bytes_me2c_clone = bytes_me2c.clone(); let bytes_me2c_clone = bytes_me2c.clone();
let d2c_flush_policy = MeD2cFlushPolicy::from_config(&config); 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, crypto_writer,
me_rx_task, me_rx_task,
stats_clone, stats_clone,
@@ -226,7 +230,13 @@ where
session_started_at, session_started_at,
conn_id, conn_id,
stop_rx, 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 main_result: Result<()> = Ok(());
let mut client_closed = false; let mut client_closed = false;
@@ -454,28 +464,30 @@ where
} }
drop(c2me_tx); 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) => { Ok(joined) => {
joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME sender join error: {e}")))) joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME sender join error: {e}"))))
} }
Err(_) => { Err(_) => {
stats.increment_me_child_join_timeout_total(); stats.increment_me_child_join_timeout_total();
stats.increment_me_child_abort_total(); stats.increment_me_child_abort_total();
c2me_sender.abort(); child_tasks.c2me_sender.abort();
Err(ProxyError::Proxy("ME sender join timeout".into())) Err(ProxyError::Proxy("ME sender join timeout".into()))
} }
}; };
flow_cancel.cancel(); flow_cancel.cancel();
let _ = stop_tx.send(()); if let Some(stop_tx) = child_tasks.stop_tx.take() {
let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut me_writer).await { let _ = stop_tx.send(());
}
let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut child_tasks.me_writer).await {
Ok(joined) => { Ok(joined) => {
joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME writer join error: {e}")))) joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME writer join error: {e}"))))
} }
Err(_) => { Err(_) => {
stats.increment_me_child_join_timeout_total(); stats.increment_me_child_join_timeout_total();
stats.increment_me_child_abort_total(); stats.increment_me_child_abort_total();
me_writer.abort(); child_tasks.me_writer.abort();
Err(ProxyError::Proxy("ME writer join timeout".into())) Err(ProxyError::Proxy("ME writer join timeout".into()))
} }
}; };
@@ -0,0 +1,65 @@
use super::*;
pub(super) struct RelayChildTasks {
pub(super) c2me_sender: AbortOnDropHandle<Result<()>>,
pub(super) me_writer: AbortOnDropHandle<Result<()>>,
pub(super) flow_cancel: CancellationToken,
pub(super) stop_tx: Option<oneshot::Sender<()>>,
}
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<AtomicUsize>);
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");
}
}
+2
View File
@@ -74,6 +74,8 @@ pub mod session_eviction;
pub mod shared_state; pub mod shared_state;
pub mod traffic_limiter; pub mod traffic_limiter;
pub(crate) mod user_admission; 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; pub use client::ClientHandler;
#[allow(unused_imports)] #[allow(unused_imports)]
+49 -69
View File
@@ -16,7 +16,7 @@ mod quota;
pub(super) use self::combined::CombinedStream; pub(super) use self::combined::CombinedStream;
pub(super) use self::counters::SharedCounters; pub(super) use self::counters::SharedCounters;
pub(super) use self::quota::is_quota_io_error; 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}; pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check};
/// Transparent I/O wrapper that tracks per-user statistics and activity. /// Transparent I/O wrapper that tracks per-user statistics and activity.
@@ -218,48 +218,33 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
let mut quota_reservation = None; let mut quota_reservation = None;
let mut read_limit = buf.remaining(); let mut read_limit = buf.remaining();
if let Some(limit) = this.quota_limit { if let Some(limit) = this.quota_limit {
let used_before = this.quota_handle.used(); for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL {
let remaining = limit.saturating_sub(used_before); let used_before = this.quota_handle.used();
if remaining == 0 { let remaining = limit.saturating_sub(used_before);
this.quota_exceeded.store(true, Ordering::Release); if remaining == 0 {
return Poll::Ready(Err(quota_io_error())); 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); let desired = remaining.min(read_limit as u64);
if read_limit == 0 { match this.quota_handle.try_reserve(desired, limit) {
this.quota_exceeded.store(true, Ordering::Release); Ok(reservation) => {
return Poll::Ready(Err(quota_io_error())); remaining_before = Some(remaining);
} read_limit = desired as usize;
quota_reservation = Some(reservation);
let desired = read_limit as u64; break;
let mut reserve_rounds = 0usize; }
while quota_reservation.is_none() { Err(crate::stats::QuotaReserveError::LimitExceeded)
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { | Err(crate::stats::QuotaReserveError::Contended) => {
match this.quota_handle.try_reserve(desired, limit) { this.stats.increment_quota_contention_total();
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();
}
} }
} }
}
if quota_reservation.is_none() { if quota_reservation.is_none() {
reserve_rounds = reserve_rounds.saturating_add(1); this.stats.increment_quota_contention_timeout_total();
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { if this.arm_quota_wait(cx).is_ready() {
this.stats.increment_quota_contention_timeout_total(); cx.waker().wake_by_ref();
if this.arm_quota_wait(cx).is_pending() {
return Poll::Pending;
}
reserve_rounds = 0;
}
} }
return Poll::Pending;
} }
} }
@@ -404,8 +389,7 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
let mut quota_reservation = None; let mut quota_reservation = None;
if let Some(limit) = this.quota_limit { if let Some(limit) = this.quota_limit {
if !write_buf.is_empty() { if !write_buf.is_empty() {
let mut reserve_rounds = 0usize; for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL {
while quota_reservation.is_none() {
let used_before = this.quota_handle.used(); let used_before = this.quota_handle.used();
let remaining = limit.saturating_sub(used_before); let remaining = limit.saturating_sub(used_before);
if remaining == 0 { if remaining == 0 {
@@ -415,36 +399,32 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
remaining_before = Some(remaining); remaining_before = Some(remaining);
let desired = remaining.min(write_buf.len() as u64); let desired = remaining.min(write_buf.len() as u64);
let mut saw_contention = false; match this.quota_handle.try_reserve(desired, limit) {
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { Ok(reservation) => {
match this.quota_handle.try_reserve(desired, limit) { quota_reservation = Some(reservation);
Ok(reservation) => { write_buf = &write_buf[..desired as usize];
quota_reservation = Some(reservation); break;
write_buf = &write_buf[..desired as usize]; }
break; Err(crate::stats::QuotaReserveError::LimitExceeded)
} | Err(crate::stats::QuotaReserveError::Contended) => {
Err(crate::stats::QuotaReserveError::LimitExceeded) => { this.stats.increment_quota_contention_total();
break;
}
Err(crate::stats::QuotaReserveError::Contended) => {
this.stats.increment_quota_contention_total();
saw_contention = true;
}
} }
} }
}
if quota_reservation.is_none() { if quota_reservation.is_none() {
reserve_rounds = reserve_rounds.saturating_add(1); this.stats.increment_quota_contention_timeout_total();
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { Self::arm_wait(&mut this.quota_wait, false, false);
this.stats.increment_quota_contention_timeout_total(); if Self::poll_wait(
Self::arm_wait(&mut this.quota_wait, false, false); &mut this.quota_wait,
let _ = cx,
Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up); None,
return Poll::Pending; RateDirection::Up,
} else if saw_contention { )
std::hint::spin_loop(); .is_ready()
} {
cx.waker().wake_by_ref();
} }
return Poll::Pending;
} }
} else { } else {
let used_before = this.quota_handle.used(); let used_before = this.quota_handle.used();
+1 -2
View File
@@ -27,8 +27,7 @@ const QUOTA_NEAR_LIMIT_BYTES: u64 = 64 * 1024;
const QUOTA_LARGE_CHARGE_BYTES: u64 = 16 * 1024; const QUOTA_LARGE_CHARGE_BYTES: u64 = 16 * 1024;
const QUOTA_ADAPTIVE_INTERVAL_MIN_BYTES: u64 = 4 * 1024; const QUOTA_ADAPTIVE_INTERVAL_MIN_BYTES: u64 = 4 * 1024;
const QUOTA_ADAPTIVE_INTERVAL_MAX_BYTES: u64 = 64 * 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_ATTEMPTS_PER_POLL: usize = 4;
pub(super) const QUOTA_RESERVE_MAX_ROUNDS: usize = 8;
#[inline] #[inline]
pub(in crate::proxy::relay) fn quota_adaptive_interval_bytes(remaining_before: u64) -> u64 { pub(in crate::proxy::relay) fn quota_adaptive_interval_bytes(remaining_before: u64) -> u64 {
+14 -1
View File
@@ -123,6 +123,19 @@ impl ProxySharedState {
pub(crate) fn new_with_direct_buffer_budget_and_user_admission( pub(crate) fn new_with_direct_buffer_budget_and_user_admission(
direct_buffer_budget: Arc<DirectBufferBudget>, direct_buffer_budget: Arc<DirectBufferBudget>,
user_admission: Arc<UserAdmissionAuthority>, user_admission: Arc<UserAdmissionAuthority>,
) -> Arc<Self> {
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<DirectBufferBudget>,
traffic_limiter: Arc<TrafficLimiter>,
user_admission: Arc<UserAdmissionAuthority>,
) -> Arc<Self> { ) -> Arc<Self> {
Arc::new(Self { Arc::new(Self {
handshake: HandshakeSharedState { handshake: HandshakeSharedState {
@@ -163,7 +176,7 @@ impl ProxySharedState {
relay_idle_registry: RelayIdleCandidateRegistry::default(), relay_idle_registry: RelayIdleCandidateRegistry::default(),
relay_idle_mark_seq: AtomicU64::new(0), relay_idle_mark_seq: AtomicU64::new(0),
}, },
traffic_limiter: TrafficLimiter::new(), traffic_limiter,
direct_buffer_budget, direct_buffer_budget,
user_admission, user_admission,
conntrack_pressure_active: AtomicBool::new(false), conntrack_pressure_active: AtomicBool::new(false),
+133 -7
View File
@@ -276,14 +276,13 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() {
ip_tracker.set_user_limit(&user, 1).await; ip_tracker.set_user_limit(&user, 1).await;
ip_tracker.check_and_add(&user, ip).await.unwrap(); 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!(ip_tracker.get_active_ip_count(&user).await, 1);
assert_eq!(stats.get_user_curr_connects(&user), 1);
let reservation = let reservation =
UserConnectionReservation::new(stats.clone(), ip_tracker.clone(), user.clone(), ip, true); 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 the reservation synchronously without any tokio::spawn/await yielding!
drop(reservation); 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); 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] #[tokio::test]
async fn relay_task_abort_releases_user_gate_and_ip_reservation() { async fn relay_task_abort_releases_user_gate_and_ip_reservation() {
let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); 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); .insert("user".to_string(), 1);
let stats = Stats::new(); 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 ip_tracker = UserIpTracker::new();
let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap(); 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; config.access.user_max_tcp_conns_global_each = 1;
let stats = Stats::new(); 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 ip_tracker = UserIpTracker::new();
let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap(); 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; config.access.user_max_tcp_conns_global_each = 1;
let stats = Stats::new(); 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 ip_tracker = UserIpTracker::new();
let peer_addr: SocketAddr = "198.51.100.213:50003".parse().unwrap(); 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 config = Arc::new(config);
let stats = Arc::new(Stats::new()); 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 ip_tracker = Arc::new(UserIpTracker::new());
let mut tasks = tokio::task::JoinSet::new(); let mut tasks = tokio::task::JoinSet::new();
@@ -1,4 +1,5 @@
use super::*; use super::*;
use tokio::sync::Semaphore;
#[test] #[test]
fn lease_drop_releases_the_complete_reservation() { fn lease_drop_releases_the_complete_reservation() {
@@ -36,3 +37,62 @@ fn growth_and_shrink_keep_accounting_balanced() {
drop(lease); drop(lease);
assert_eq!(budget.snapshot().reserved_bytes, 0); 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));
}
+1 -1
View File
@@ -142,6 +142,7 @@ enum CidrPolicyMatch<'a> {
#[derive(Default)] #[derive(Default)]
struct PolicySnapshot { struct PolicySnapshot {
revision: u64, revision: u64,
source_generation: u64,
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_rules_v4: Vec<CidrRule>, cidr_rules_v4: Vec<CidrRule>,
cidr_rules_v6: Vec<CidrRule>, cidr_rules_v6: Vec<CidrRule>,
@@ -175,7 +176,6 @@ pub struct TrafficLease {
pub struct TrafficLimiter { pub struct TrafficLimiter {
policy: ArcSwap<PolicySnapshot>, policy: ArcSwap<PolicySnapshot>,
policy_update: ParkingMutex<()>, policy_update: ParkingMutex<()>,
published_revision: AtomicU64,
user_buckets: ShardedRegistry<UserBucket>, user_buckets: ShardedRegistry<UserBucket>,
cidr_buckets: ShardedRegistry<CidrBucket>, cidr_buckets: ShardedRegistry<CidrBucket>,
user_scope: ScopeMetrics, user_scope: ScopeMetrics,
+5 -12
View File
@@ -2,20 +2,16 @@ use super::*;
impl TrafficLease { impl TrafficLease {
fn current_binding(&self) -> Arc<TrafficLeaseBinding> { fn current_binding(&self) -> Arc<TrafficLeaseBinding> {
let published_revision = self.limiter.published_revision.load(Ordering::Acquire); let policy = self.limiter.policy.load();
let current = self.binding.load_full(); let current = self.binding.load_full();
if current.revision == published_revision { if current.revision == policy.revision {
return current; return current;
} }
drop(policy);
let refresh = self.refresh.lock(); 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 policy = self.limiter.policy.load_full(); let policy = self.limiter.policy.load_full();
let current = self.binding.load_full();
if current.revision == policy.revision { if current.revision == policy.revision {
return current; return current;
} }
@@ -23,9 +19,6 @@ impl TrafficLease {
.limiter .limiter
.build_binding(&self.user, self.client_ip, &policy); .build_binding(&self.user, self.client_ip, &policy);
self.binding.store(Arc::clone(&next)); self.binding.store(Arc::clone(&next));
drop(policy_update);
drop(refresh);
self.limiter.maybe_cleanup();
next next
} }
+33 -10
View File
@@ -6,7 +6,6 @@ impl TrafficLimiter {
Arc::new(Self { Arc::new(Self {
policy: ArcSwap::from_pointee(PolicySnapshot::default()), policy: ArcSwap::from_pointee(PolicySnapshot::default()),
policy_update: ParkingMutex::new(()), policy_update: ParkingMutex::new(()),
published_revision: AtomicU64::new(0),
user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), user_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
user_scope: ScopeMetrics::default(), user_scope: ScopeMetrics::default(),
@@ -15,16 +14,31 @@ impl TrafficLimiter {
}) })
} }
#[cfg(test)]
pub fn apply_policy( pub fn apply_policy(
&self, &self,
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>, cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
) { ) {
let policy_update = self.policy_update.lock(); let _ = self.apply_policy_inner(None, user_limits, cidr_limits);
// Revision wrap could otherwise let an old lease restore stale rates. }
let Some(revision) = self.policy.load().revision.checked_add(1) else {
return; /// 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<String, RateLimitBps>,
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
) -> bool {
self.apply_policy_inner(Some(source_generation), user_limits, cidr_limits)
}
fn apply_policy_inner(
&self,
source_generation: Option<u64>,
user_limits: HashMap<String, RateLimitBps>,
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
) -> bool {
let filtered_users = user_limits let filtered_users = user_limits
.into_iter() .into_iter()
.filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0)
@@ -77,6 +91,17 @@ impl TrafficLimiter {
let cidr_policy_entries = let cidr_policy_entries =
cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len(); 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 self.user_scope
.policy_entries .policy_entries
.store(filtered_users.len() as u64, Ordering::Relaxed); .store(filtered_users.len() as u64, Ordering::Relaxed);
@@ -86,6 +111,7 @@ impl TrafficLimiter {
self.policy.store(Arc::new(PolicySnapshot { self.policy.store(Arc::new(PolicySnapshot {
revision, revision,
source_generation,
user_limits: filtered_users, user_limits: filtered_users,
cidr_rules_v4, cidr_rules_v4,
cidr_rules_v6, cidr_rules_v6,
@@ -93,10 +119,10 @@ impl TrafficLimiter {
cidr_auto_rules_v6, cidr_auto_rules_v6,
cidr_rule_keys, cidr_rule_keys,
})); }));
self.published_revision.store(revision, Ordering::Release);
drop(policy_update); drop(policy_update);
self.maybe_cleanup(); self.maybe_cleanup();
true
} }
pub fn acquire_lease( pub fn acquire_lease(
@@ -104,11 +130,8 @@ impl TrafficLimiter {
user: &str, user: &str,
client_ip: IpAddr, client_ip: IpAddr,
) -> Option<Arc<TrafficLease>> { ) -> Option<Arc<TrafficLease>> {
let policy_update = self.policy_update.lock();
let policy = self.policy.load_full(); let policy = self.policy.load_full();
let binding = self.build_binding(user, client_ip, &policy); let binding = self.build_binding(user, client_ip, &policy);
drop(policy_update);
self.maybe_cleanup();
Some(Arc::new(TrafficLease { Some(Arc::new(TrafficLease {
limiter: Arc::clone(self), limiter: Arc::clone(self),
user: user.to_string(), user: user.to_string(),
+73
View File
@@ -5,6 +5,79 @@ fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps {
RateLimitBps { up_bps, down_bps } 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] #[test]
fn explicit_cidr_rule_wins_over_auto_template() { fn explicit_cidr_rule_wins_over_auto_template() {
let limiter = TrafficLimiter::new(); let limiter = TrafficLimiter::new();
+123
View File
@@ -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<String, u64>,
}
/// 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<UserConnectionAuthority>,
user: String,
}
impl UserConnectionAuthority {
/// Acquires a slot without consulting optional telemetry state.
pub(crate) fn try_acquire(
self: &Arc<Self>,
user: &str,
limit: Option<u64>,
) -> Option<UserConnectionPermit> {
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::<Vec<_>>();
assert_eq!(permits.len() as u64, LIMIT);
assert_eq!(authority.active("alice"), LIMIT);
drop(permits);
assert_eq!(authority.active("alice"), 0);
}
}
+26 -1
View File
@@ -22,9 +22,11 @@ use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering};
use std::time::Instant; use std::time::Instant;
pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle}; pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle};
pub(crate) use self::users::UserConnectionObservation;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use self::replay::{ReplayChecker, ReplayStats}; pub use self::replay::{ReplayChecker, ReplayStats};
use self::telemetry::TelemetryPolicy; use self::telemetry::TelemetryPolicy;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
pub use self::tls_fingerprints::TlsFingerprintSnapshotRow; pub use self::tls_fingerprints::TlsFingerprintSnapshotRow;
use crate::config::MeWriterPickMode; use crate::config::MeWriterPickMode;
@@ -351,6 +353,7 @@ pub struct Stats {
tls_fingerprints: tls_fingerprints::TlsFingerprintCollector, tls_fingerprints: tls_fingerprints::TlsFingerprintCollector,
user_stats: DashMap<String, Arc<UserStats>>, user_stats: DashMap<String, Arc<UserStats>>,
quota_store: Arc<QuotaStore>, quota_store: Arc<QuotaStore>,
connection_authority: Arc<UserConnectionAuthority>,
user_stats_last_cleanup_epoch_secs: AtomicU64, user_stats_last_cleanup_epoch_secs: AtomicU64,
start_time: parking_lot::RwLock<Option<Instant>>, start_time: parking_lot::RwLock<Option<Instant>>,
} }
@@ -417,12 +420,28 @@ impl UserStats {
impl Stats { impl Stats {
pub fn new() -> Self { 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<QuotaStore>) -> Self { pub(crate) fn with_quota_store(quota_store: Arc<QuotaStore>) -> 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<QuotaStore>,
connection_authority: Arc<UserConnectionAuthority>,
) -> Self {
let stats = Self { let stats = Self {
quota_store, quota_store,
connection_authority,
..Self::default() ..Self::default()
}; };
stats.apply_telemetry_policy(TelemetryPolicy::default()); stats.apply_telemetry_policy(TelemetryPolicy::default());
@@ -431,10 +450,16 @@ impl Stats {
stats stats
} }
/// Returns the process-scoped quota authority for test runtime construction.
#[cfg(test)] #[cfg(test)]
pub(crate) fn quota_store(&self) -> Arc<QuotaStore> { pub(crate) fn quota_store(&self) -> Arc<QuotaStore> {
Arc::clone(&self.quota_store) Arc::clone(&self.quota_store)
} }
/// Returns process-owned per-user connection admission.
pub(crate) fn connection_authority(&self) -> Arc<UserConnectionAuthority> {
Arc::clone(&self.connection_authority)
}
} }
#[cfg(test)] #[cfg(test)]
+26
View File
@@ -13,6 +13,32 @@ fn test_stats_shared_counters() {
assert_eq!(stats.get_connects_all(), 3); 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(&quota_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] #[test]
fn test_telemetry_policy_disables_core_and_user_counters() { fn test_telemetry_policy_disables_core_and_user_counters() {
let stats = Stats::new(); let stats = Stats::new();
+47 -16
View File
@@ -1,5 +1,32 @@
use super::*; 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<UserStats>,
}
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 { impl Stats {
pub fn increment_user_connects(&self, user: &str) { pub fn increment_user_connects(&self, user: &str) {
if !self.telemetry_user_enabled() { if !self.telemetry_user_enabled() {
@@ -19,6 +46,25 @@ impl Stats {
stats.curr_connects.fetch_add(1, Ordering::Relaxed); 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<UserConnectionObservation> {
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<u64>) -> bool { pub fn try_acquire_user_curr_connects(&self, user: &str, limit: Option<u64>) -> bool {
if !self.telemetry_user_enabled() { if !self.telemetry_user_enabled() {
return true; return true;
@@ -50,22 +96,7 @@ impl Stats {
pub fn decrement_user_curr_connects(&self, user: &str) { pub fn decrement_user_curr_connects(&self, user: &str) {
if let Some(stats) = self.user_stats.get(user) { if let Some(stats) = self.user_stats.get(user) {
self.touch_user_stats(stats.value().as_ref()); self.touch_user_stats(stats.value().as_ref());
let counter = &stats.curr_connects; decrement_current_connections(&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,
}
}
} }
} }
+22 -14
View File
@@ -3,6 +3,8 @@ use tokio::process::Command;
use crate::util::trusted_command::resolve_trusted_helper; 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( pub(super) async fn run_command(
binary: &str, binary: &str,
args: &[&str], args: &[&str],
@@ -18,21 +20,26 @@ pub(super) async fn run_command(
} }
command.stdout(std::process::Stdio::null()); command.stdout(std::process::Stdio::null());
command.stderr(std::process::Stdio::piped()); command.stderr(std::process::Stdio::piped());
command.kill_on_drop(true);
let mut child = command let mut child = command
.spawn() .spawn()
.map_err(|e| format!("spawn {binary} failed: {e}"))?; .map_err(|e| format!("spawn {binary} failed: {e}"))?;
if let Some(blob) = stdin let output = tokio::time::timeout(COMMAND_TIMEOUT, async move {
&& let Some(mut writer) = child.stdin.take() if let Some(blob) = stdin
{ && let Some(mut writer) = child.stdin.take()
writer {
.write_all(blob.as_bytes()) writer
.write_all(blob.as_bytes())
.await
.map_err(|e| format!("stdin write {binary} failed: {e}"))?;
}
child
.wait_with_output()
.await .await
.map_err(|e| format!("stdin write {binary} failed: {e}"))?; .map_err(|e| format!("wait {binary} failed: {e}"))
} })
let output = child .await
.wait_with_output() .map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??;
.await
.map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() { if output.status.success() {
return Ok(()); return Ok(());
} }
@@ -48,10 +55,11 @@ pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result<St
let Some(command_path) = resolve_trusted_helper(binary) else { let Some(command_path) = resolve_trusted_helper(binary) else {
return Err(format!("{binary} is not available")); return Err(format!("{binary} is not available"));
}; };
let output = Command::new(command_path) let mut command = Command::new(command_path);
.args(args) command.args(args).kill_on_drop(true);
.output() let output = tokio::time::timeout(COMMAND_TIMEOUT, command.output())
.await .await
.map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))?
.map_err(|e| format!("wait {binary} failed: {e}"))?; .map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() { if output.status.success() {
return Ok(String::from_utf8_lossy(&output.stdout).to_string()); return Ok(String::from_utf8_lossy(&output.stdout).to_string());
+121 -98
View File
@@ -1,5 +1,6 @@
use std::cmp::Reverse; use std::cmp::Reverse;
use std::collections::{HashMap, HashSet}; use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use super::super::MePool; use super::super::MePool;
@@ -10,6 +11,18 @@ use super::{
}; };
use crate::config::MeWriterPickMode; use crate::config::MeWriterPickMode;
struct WriterSelectionKey {
index: usize,
contour_rank: usize,
stale: usize,
degraded: usize,
idle_rank: usize,
queue_remaining: usize,
addr: SocketAddr,
id: u64,
pick_score: u64,
}
impl MePool { impl MePool {
pub(super) async fn candidate_indices_for_dc( pub(super) async fn candidate_indices_for_dc(
&self, &self,
@@ -72,24 +85,20 @@ impl MePool {
} }
} }
pub(super) fn writer_contour_rank_for_selection( fn writer_contour_rank_for_selection(contour: WriterContour) -> usize {
&self, match contour {
writer: &super::super::pool::MeWriter,
) -> usize {
match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) {
WriterContour::Active => 0, WriterContour::Active => 0,
WriterContour::Warm => 1, WriterContour::Warm => 1,
WriterContour::Draining => 2, WriterContour::Draining => 2,
} }
} }
pub(super) fn writer_idle_rank_for_selection( fn writer_idle_rank_for_selection(
&self, writer_id: u64,
writer: &super::super::pool::MeWriter,
idle_since_by_writer: &HashMap<u64, u64>, idle_since_by_writer: &HashMap<u64, u64>,
now_epoch_secs: u64, now_epoch_secs: u64,
) -> usize { ) -> 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; return 0;
}; };
let idle_age_secs = now_epoch_secs.saturating_sub(idle_since); 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, &self,
index: usize,
writer: &super::super::pool::MeWriter, writer: &super::super::pool::MeWriter,
idle_since_by_writer: &HashMap<u64, u64>, idle_since_by_writer: &HashMap<u64, u64>,
now_epoch_secs: u64, now_epoch_secs: u64,
) -> u64 { current_generation: u64,
let contour_penalty = match WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) { ) -> 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::Active => 0,
WriterContour::Warm => PICK_PENALTY_WARM, WriterContour::Warm => PICK_PENALTY_WARM,
WriterContour::Draining => PICK_PENALTY_DRAINING, 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 PICK_PENALTY_STALE
} else { } else {
0 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 PICK_PENALTY_DEGRADED
} else { } else {
0 0
}; };
let idle_penalty = let idle_rank =
(self.writer_idle_rank_for_selection(writer, idle_since_by_writer, now_epoch_secs) Self::writer_idle_rank_for_selection(writer.id, idle_since_by_writer, now_epoch_secs);
as u64) let idle_penalty = (idle_rank as u64) * 100;
* 100;
let queue_cap = self.writer_lifecycle.writer_cmd_channel_capacity.max(1) as u64; let queue_cap = self.writer_lifecycle.writer_cmd_channel_capacity.max(1) as u64;
let queue_remaining = writer.tx.capacity() as u64; let queue_remaining = writer.tx.capacity();
let queue_used = queue_cap.saturating_sub(queue_remaining.min(queue_cap)); 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_util_pct = queue_used.saturating_mul(100) / queue_cap;
let queue_penalty = queue_util_pct.saturating_mul(4); let queue_penalty = queue_util_pct.saturating_mul(4);
let rtt_penalty = let rtt_penalty =
((writer.rtt_ema_ms_x10.load(Ordering::Relaxed) as u64).saturating_add(5) / 10) ((writer.rtt_ema_ms_x10.load(Ordering::Relaxed) as u64).saturating_add(5) / 10)
.min(400); .min(400);
contour_penalty let pick_score = contour_penalty
.saturating_add(stale_penalty) .saturating_add(stale_penalty)
.saturating_add(degraded_penalty) .saturating_add(degraded_penalty)
.saturating_add(idle_penalty) .saturating_add(idle_penalty)
.saturating_add(queue_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( pub(super) fn p2c_ordered_candidate_indices(
&self, &self,
candidate_indices: &[usize], mut candidate_indices: Vec<usize>,
writers_snapshot: &[super::super::pool::MeWriter], writers_snapshot: &[super::super::pool::MeWriter],
idle_since_by_writer: &HashMap<u64, u64>, idle_since_by_writer: &HashMap<u64, u64>,
now_epoch_secs: u64, now_epoch_secs: u64,
@@ -158,33 +183,26 @@ impl MePool {
return Vec::new(); return Vec::new();
} }
let mut sampled = Vec::<usize>::with_capacity(sample_size.min(total)); candidate_indices.rotate_left(start % total);
let mut seen = HashSet::<usize>::with_capacity(total); let current_generation = self.current_generation();
for offset in 0..sample_size.min(total) { let sample_size = sample_size.min(total);
let idx = candidate_indices[(start + offset) % total]; let mut sampled = candidate_indices[..sample_size]
if seen.insert(idx) { .iter()
sampled.push(idx); .map(|idx| {
} self.capture_writer_selection_key(
*idx,
&writers_snapshot[*idx],
idle_since_by_writer,
now_epoch_secs,
current_generation,
)
})
.collect::<Vec<_>>();
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;
} }
candidate_indices
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::<usize>::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
} }
pub(super) async fn 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(); let start = self.rr.fetch_add(1, Ordering::Relaxed) as usize % candidate_indices.len();
if pick_mode == MeWriterPickMode::P2c { if pick_mode == MeWriterPickMode::P2c {
return self.p2c_ordered_candidate_indices( return self.p2c_ordered_candidate_indices(
&candidate_indices, candidate_indices,
writers_snapshot, writers_snapshot,
&writer_idle_since, &writer_idle_since,
now_epoch_secs, now_epoch_secs,
@@ -220,60 +238,65 @@ impl MePool {
.me_deterministic_writer_sort .me_deterministic_writer_sort
.load(Ordering::Relaxed) .load(Ordering::Relaxed)
{ {
candidate_indices.sort_by(|lhs, rhs| { let current_generation = self.current_generation();
let left = &writers_snapshot[*lhs]; let mut captured = candidate_indices
let right = &writers_snapshot[*rhs]; .iter()
let left_key = ( .map(|idx| {
self.writer_contour_rank_for_selection(left), self.capture_writer_selection_key(
(left.generation < self.current_generation()) as usize, *idx,
left.degraded.load(Ordering::Relaxed) as usize, &writers_snapshot[*idx],
self.writer_idle_rank_for_selection(
left,
&writer_idle_since, &writer_idle_since,
now_epoch_secs, now_epoch_secs,
), current_generation,
Reverse(left.tx.capacity()), )
left.addr, })
left.id, .collect::<Vec<_>>();
); captured.sort_by_key(|candidate| {
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;
( (
self.writer_contour_rank_for_selection(writer), candidate.contour_rank,
stale, candidate.stale,
degraded as usize, candidate.degraded,
self.writer_idle_rank_for_selection( candidate.idle_rank,
writer, Reverse(candidate.queue_remaining),
&writer_idle_since, candidate.addr,
now_epoch_secs, candidate.id,
),
Reverse(writer.tx.capacity()),
) )
}); });
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::<Vec<_>>();
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::<usize>::with_capacity(candidate_indices.len()); if !candidate_indices.is_empty() {
for offset in 0..candidate_indices.len() { let len = candidate_indices.len();
ordered.push(candidate_indices[(start + offset) % candidate_indices.len()]); candidate_indices.rotate_left(start % len);
} }
ordered candidate_indices
} }
} }
+70 -11
View File
@@ -1,3 +1,4 @@
use std::future::Future;
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
@@ -112,6 +113,40 @@ pub(super) async fn run_upgraded(
type CarrierSocket = WebSocketStream<ConnectionIo>; type CarrierSocket = WebSocketStream<ConnectionIo>;
enum DataPlaneEvent<I, D> {
Incoming(I),
Down(D),
}
#[derive(Default)]
struct FairDataSelector {
prefer_down: bool,
}
impl FairDataSelector {
async fn select<I, D>(&mut self, incoming: I, down: D) -> DataPlaneEvent<I::Output, D::Output>
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( async fn run_multiplex(
socket: &mut CarrierSocket, socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>, runtime: &Arc<WebProcessRuntime>,
@@ -134,26 +169,31 @@ async fn run_multiplex(
let write_timeout = Duration::from_secs(session.timeouts().websocket_write_secs); let write_timeout = Duration::from_secs(session.timeouts().websocket_write_secs);
let maximum_message = session.limits().carrier_batch_bytes; let maximum_message = session.limits().carrier_batch_bytes;
let mut active = false; let mut active = false;
let mut data_selector = FairDataSelector::default();
loop { loop {
let down = session.poll_down_websocket(cursor); let down = session.poll_down_websocket(cursor);
tokio::pin!(down);
let event = tokio::select! { let event = tokio::select! {
biased; biased;
_ = cancellation.cancelled() => return Err(()), _ = cancellation.cancelled() => return Err(()),
_ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()),
_ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness,
incoming = read_message( data = data_selector.select(
socket, read_message(
runtime, socket,
session.profile_key(), runtime,
&cancellation, session.profile_key(),
&mut read_budget, &cancellation,
maximum_message, &mut read_budget,
backpressure_timeout, 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 { match event {
DriverEvent::Incoming((message, _budget)) => match message { DriverEvent::Incoming((message, _budget)) => match message {
@@ -363,3 +403,22 @@ enum DriverEvent {
Down(crate::web::session::PollResult), Down(crate::web::session::PollResult),
Liveness, 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
);
}
}
}
+18 -19
View File
@@ -5,9 +5,9 @@ use bytes::Bytes;
use tokio_tungstenite::tungstenite::protocol::Message; use tokio_tungstenite::tungstenite::protocol::Message;
use tokio_util::sync::CancellationToken; 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 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::session::{SessionCloseReason, WebSession, WebSocketLaneReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext}; 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 write_timeout = Duration::from_secs(session.timeouts().websocket_write_secs);
let maximum_message = session.limits().carrier_batch_bytes; let maximum_message = session.limits().carrier_batch_bytes;
let mut active = false; let mut active = false;
let mut data_selector = FairDataSelector::default();
loop { loop {
let down = session.poll_down_websocket_lane(reservation.lane_identity(), cursor); let down = session.poll_down_websocket_lane(reservation.lane_identity(), cursor);
tokio::pin!(down);
let event = tokio::select! { let event = tokio::select! {
biased; biased;
_ = cancellation.cancelled() => return Err(()), _ = cancellation.cancelled() => return Err(()),
_ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()),
_ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness,
incoming = read_message( data = data_selector.select(
socket, read_message(
runtime, socket,
session.profile_key(), runtime,
&cancellation, session.profile_key(),
&mut read_budget, &cancellation,
maximum_message, &mut read_budget,
backpressure_timeout, 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 { match event {
DriverEvent::Incoming((message, _budget)) => match message { DriverEvent::Incoming((message, _budget)) => match message {
@@ -267,9 +272,3 @@ pub(super) async fn run_lane(
} }
} }
} }
enum DriverEvent {
Incoming((Message, Option<WebSocketBudgetLease>)),
Down(crate::web::session::PollResult),
Liveness,
}
+19 -3
View File
@@ -1,6 +1,10 @@
use std::sync::Arc;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::task::Waker;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tokio::sync::Notify;
use super::WebSession; use super::WebSession;
/// Stable terminal cause assigned by the first session-close winner. /// Stable terminal cause assigned by the first session-close winner.
@@ -104,6 +108,8 @@ struct ReleasedQueues {
recovery_closed_before_commit: bool, recovery_closed_before_commit: bool,
reason: SessionCloseReason, reason: SessionCloseReason,
peer_gap: Duration, peer_gap: Duration,
stream_wakers: Vec<Waker>,
lane_notifies: Vec<Arc<Notify>>,
} }
/// Deferred queue release after manager publication linearizes a supersede. /// Deferred queue release after manager publication linearizes a supersede.
@@ -276,12 +282,13 @@ impl WebSession {
if reason == SessionCloseReason::CarrierSuperseded { if reason == SessionCloseReason::CarrierSuperseded {
state.negotiation_phase = SessionNegotiationPhase::Superseded; 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() { for stream in state.streams.values_mut() {
if let Some(waker) = stream.read_waker.take() { if let Some(waker) = stream.read_waker.take() {
waker.wake(); stream_wakers.push(waker);
} }
if let Some(waker) = stream.write_waker.take() { if let Some(waker) = stream.write_waker.take() {
waker.wake(); stream_wakers.push(waker);
} }
} }
state.streams.clear(); state.streams.clear();
@@ -296,8 +303,9 @@ impl WebSession {
let mut lane_data_items = 0usize; let mut lane_data_items = 0usize;
let mut lane_control_bytes = 0usize; let mut lane_control_bytes = 0usize;
let mut lane_control_items = 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() { 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() { if let Some(batch) = lane.unacked.take() {
batch.lease.detach(); batch.lease.detach();
lane_data_bytes = lane_data_bytes.saturating_add(batch.data_bytes); lane_data_bytes = lane_data_bytes.saturating_add(batch.data_bytes);
@@ -326,10 +334,18 @@ impl WebSession {
recovery_closed_before_commit, recovery_closed_before_commit,
reason, reason,
peer_gap, peer_gap,
stream_wakers,
lane_notifies,
} }
} }
fn finish_close(&self, released: ReleasedQueues) { 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(); self.cancel.cancel();
if self.carrier().is_multiplexed() { if self.carrier().is_multiplexed() {
self.down_notify.notify_waiters(); self.down_notify.notify_waiters();
+77
View File
@@ -1,10 +1,13 @@
use super::*; use super::*;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
use std::task::{Wake, Waker};
use crate::config::{ use crate::config::{
WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig, WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
}; };
use crate::web::manager::WebProcessRuntime; use crate::web::manager::WebProcessRuntime;
use crate::web::session::SessionCloseOutcome;
fn session() -> Arc<WebSession> { fn session() -> Arc<WebSession> {
session_with_automatic(false) session_with_automatic(false)
@@ -53,6 +56,80 @@ fn session_with_automatic(automatic: bool) -> Arc<WebSession> {
) )
} }
struct SessionLockProbe {
session: std::sync::Weak<WebSession>,
lock_was_free: Arc<AtomicBool>,
}
impl Wake for SessionLockProbe {
fn wake(self: Arc<Self>) {
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] #[test]
fn uplink_retry_commits_only_one_exact_body() { fn uplink_retry_commits_only_one_exact_body() {
let session = session(); let session = session();