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