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