mirror of
https://github.com/telemt/telemt.git
synced 2026-10-10 11:25:57 +03:00
Process-wide concurrency + Cancellation ownership fixes
This commit is contained in:
+87
-47
@@ -15,7 +15,8 @@ use crate::proxy::middle_relay::{handle_via_middle_proxy, handle_via_middle_prox
|
||||
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
|
||||
use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState};
|
||||
use crate::proxy::user_admission::UserIncarnation;
|
||||
use crate::stats::{Stats, UserQuotaHandle};
|
||||
use crate::proxy::user_connection_authority::UserConnectionPermit;
|
||||
use crate::stats::{Stats, UserConnectionObservation, UserQuotaHandle};
|
||||
use crate::stream::{BufferPool, CryptoReader, CryptoWriter};
|
||||
use crate::transport::UpstreamManager;
|
||||
use crate::transport::middle_proxy::MePool;
|
||||
@@ -238,13 +239,63 @@ where
|
||||
/// Owns one authenticated user's connection and source-IP admission slots.
|
||||
pub(crate) struct UserConnectionReservation {
|
||||
stats: Arc<Stats>,
|
||||
ip_tracker: Arc<UserIpTracker>,
|
||||
user: String,
|
||||
ip: IpAddr,
|
||||
incarnation: UserIncarnation,
|
||||
quota_handle: UserQuotaHandle,
|
||||
tracks_ip: bool,
|
||||
active: bool,
|
||||
_connection_permit: UserConnectionPermit,
|
||||
_stats_observation: Option<UserConnectionObservation>,
|
||||
ip_permit: Option<UserIpPermit>,
|
||||
released: bool,
|
||||
}
|
||||
|
||||
struct UserIpPermit {
|
||||
tracker: Arc<UserIpTracker>,
|
||||
owner: Option<UserIpOwner>,
|
||||
}
|
||||
|
||||
struct UserIpOwner {
|
||||
user: String,
|
||||
incarnation: UserIncarnation,
|
||||
ip: IpAddr,
|
||||
}
|
||||
|
||||
impl UserIpPermit {
|
||||
fn new(
|
||||
tracker: Arc<UserIpTracker>,
|
||||
user: String,
|
||||
incarnation: UserIncarnation,
|
||||
ip: IpAddr,
|
||||
) -> Self {
|
||||
Self {
|
||||
tracker,
|
||||
owner: Some(UserIpOwner {
|
||||
user,
|
||||
incarnation,
|
||||
ip,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
async fn release(mut self) {
|
||||
let Some(owner) = self.owner.as_ref() else {
|
||||
return;
|
||||
};
|
||||
self.tracker
|
||||
.remove_ip_for_incarnation(&owner.user, owner.incarnation, owner.ip)
|
||||
.await;
|
||||
self.owner = None;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for UserIpPermit {
|
||||
fn drop(&mut self) {
|
||||
let Some(owner) = self.owner.take() else {
|
||||
return;
|
||||
};
|
||||
self.tracker.enqueue_cleanup_for_incarnation(
|
||||
owner.user,
|
||||
owner.incarnation,
|
||||
owner.ip,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
impl UserConnectionReservation {
|
||||
@@ -257,6 +308,11 @@ impl UserConnectionReservation {
|
||||
tracks_ip: bool,
|
||||
) -> Self {
|
||||
let quota_handle = stats.current_user_quota_handle(&user);
|
||||
let connection_permit = stats
|
||||
.connection_authority()
|
||||
.try_acquire(&user, None)
|
||||
.expect("unlimited test connection permit must be available");
|
||||
let stats_observation = stats.observe_user_current_connection(&user);
|
||||
Self::new_for_incarnation(
|
||||
stats,
|
||||
ip_tracker,
|
||||
@@ -264,6 +320,8 @@ impl UserConnectionReservation {
|
||||
ip,
|
||||
0,
|
||||
quota_handle,
|
||||
connection_permit,
|
||||
stats_observation,
|
||||
tracks_ip,
|
||||
)
|
||||
}
|
||||
@@ -276,17 +334,20 @@ impl UserConnectionReservation {
|
||||
ip: IpAddr,
|
||||
incarnation: UserIncarnation,
|
||||
quota_handle: UserQuotaHandle,
|
||||
connection_permit: UserConnectionPermit,
|
||||
stats_observation: Option<UserConnectionObservation>,
|
||||
tracks_ip: bool,
|
||||
) -> Self {
|
||||
let ip_permit = tracks_ip.then(|| {
|
||||
UserIpPermit::new(ip_tracker, user, incarnation, ip)
|
||||
});
|
||||
Self {
|
||||
stats,
|
||||
ip_tracker,
|
||||
user,
|
||||
ip,
|
||||
incarnation,
|
||||
quota_handle,
|
||||
tracks_ip,
|
||||
active: true,
|
||||
_connection_permit: connection_permit,
|
||||
_stats_observation: stats_observation,
|
||||
ip_permit,
|
||||
released: false,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -297,50 +358,24 @@ impl UserConnectionReservation {
|
||||
|
||||
/// Releases both admission counters through the asynchronous cleanup path.
|
||||
pub(crate) async fn release(mut self) {
|
||||
if !self.active {
|
||||
return;
|
||||
if let Some(ip_permit) = self.ip_permit.take() {
|
||||
ip_permit.release().await;
|
||||
}
|
||||
self.active = false;
|
||||
if self.tracks_ip {
|
||||
self.ip_tracker
|
||||
.remove_ip_for_incarnation(&self.user, self.incarnation, self.ip)
|
||||
.await;
|
||||
}
|
||||
self.stats.decrement_user_curr_connects(&self.user);
|
||||
self.released = true;
|
||||
}
|
||||
|
||||
/// Defers IP cleanup when admission fails after the asynchronous reservation step.
|
||||
pub(crate) fn release_deferred(mut self) {
|
||||
if !self.active {
|
||||
return;
|
||||
}
|
||||
self.active = false;
|
||||
self.stats.decrement_user_curr_connects(&self.user);
|
||||
if self.tracks_ip {
|
||||
self.ip_tracker.enqueue_cleanup_for_incarnation(
|
||||
self.user.clone(),
|
||||
self.incarnation,
|
||||
self.ip,
|
||||
);
|
||||
}
|
||||
self.released = true;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for UserConnectionReservation {
|
||||
fn drop(&mut self) {
|
||||
if !self.active {
|
||||
if self.released {
|
||||
return;
|
||||
}
|
||||
self.active = false;
|
||||
self.stats.increment_session_drop_fallback_total();
|
||||
self.stats.decrement_user_curr_connects(&self.user);
|
||||
if self.tracks_ip {
|
||||
self.ip_tracker.enqueue_cleanup_for_incarnation(
|
||||
self.user.clone(),
|
||||
self.incarnation,
|
||||
self.ip,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -400,17 +435,20 @@ async fn acquire_user_connection_reservation_for_incarnation(
|
||||
.or((config.access.user_max_tcp_conns_global_each > 0)
|
||||
.then_some(config.access.user_max_tcp_conns_global_each))
|
||||
.map(|value| value as u64);
|
||||
if !stats.try_acquire_user_curr_connects(user, limit) {
|
||||
let Some(connection_permit) = stats
|
||||
.connection_authority()
|
||||
.try_acquire(user, limit)
|
||||
else {
|
||||
return Err(ProxyError::ConnectionLimitExceeded {
|
||||
user: user.to_string(),
|
||||
});
|
||||
}
|
||||
};
|
||||
let stats_observation = stats.observe_user_current_connection(user);
|
||||
|
||||
if let Err(reason) = ip_tracker
|
||||
.check_and_add_for_incarnation(user, incarnation, peer_addr.ip())
|
||||
.await
|
||||
{
|
||||
stats.decrement_user_curr_connects(user);
|
||||
warn!(
|
||||
user = %user,
|
||||
ip = %peer_addr.ip(),
|
||||
@@ -429,6 +467,8 @@ async fn acquire_user_connection_reservation_for_incarnation(
|
||||
peer_addr.ip(),
|
||||
incarnation,
|
||||
quota_handle,
|
||||
connection_permit,
|
||||
stats_observation,
|
||||
true,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -155,18 +155,20 @@ impl RunningClientHandler {
|
||||
.or((config.access.user_max_tcp_conns_global_each > 0)
|
||||
.then_some(config.access.user_max_tcp_conns_global_each))
|
||||
.map(|v| v as u64);
|
||||
if !stats.try_acquire_user_curr_connects(user, limit) {
|
||||
let Some(_connection_permit) = stats
|
||||
.connection_authority()
|
||||
.try_acquire(user, limit)
|
||||
else {
|
||||
return Err(ProxyError::ConnectionLimitExceeded {
|
||||
user: user.to_string(),
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
match ip_tracker.check_and_add(user, peer_addr.ip()).await {
|
||||
Ok(()) => {
|
||||
ip_tracker.remove_ip(user, peer_addr.ip()).await;
|
||||
}
|
||||
Err(reason) => {
|
||||
stats.decrement_user_curr_connects(user);
|
||||
warn!(
|
||||
user = %user,
|
||||
ip = %peer_addr.ip(),
|
||||
@@ -178,8 +180,6 @@ impl RunningClientHandler {
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
stats.decrement_user_curr_connects(user);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,12 +2,16 @@ use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
use parking_lot::{Mutex as ParkingMutex, MutexGuard as ParkingMutexGuard};
|
||||
use tokio::sync::watch;
|
||||
|
||||
use crate::stats::Stats;
|
||||
use crate::stream::BufferPool;
|
||||
|
||||
use super::shared_state::ProxySharedState;
|
||||
// Process controller and system-memory sampling remain outside data-plane accounting.
|
||||
mod controller;
|
||||
pub(crate) use controller::{
|
||||
resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
|
||||
};
|
||||
#[cfg(test)]
|
||||
use controller::connection_fill_pct;
|
||||
|
||||
/// Accounting granularity for process-wide Direct copy-buffer reservations.
|
||||
pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024;
|
||||
@@ -70,6 +74,8 @@ pub(crate) struct DirectBufferBudget {
|
||||
hard_limit_bytes: u64,
|
||||
target_bytes: AtomicU64,
|
||||
reserved_bytes: AtomicU64,
|
||||
active_controller_generation: AtomicU64,
|
||||
controller_update: ParkingMutex<()>,
|
||||
pressure_generation: AtomicU64,
|
||||
pressure_tx: watch::Sender<u64>,
|
||||
memory_total_bytes: AtomicU64,
|
||||
@@ -94,6 +100,8 @@ impl DirectBufferBudget {
|
||||
hard_limit_bytes,
|
||||
target_bytes: AtomicU64::new(hard_limit_bytes),
|
||||
reserved_bytes: AtomicU64::new(0),
|
||||
active_controller_generation: AtomicU64::new(0),
|
||||
controller_update: ParkingMutex::new(()),
|
||||
pressure_generation: AtomicU64::new(0),
|
||||
pressure_tx,
|
||||
memory_total_bytes: AtomicU64::new(0),
|
||||
@@ -120,6 +128,22 @@ impl DirectBufferBudget {
|
||||
self.pressure_tx.subscribe()
|
||||
}
|
||||
|
||||
/// Transfers adaptive-target writes to the active runtime generation.
|
||||
pub(crate) fn activate_controller(&self, generation: u64) {
|
||||
let _controller_update = self.controller_update.lock();
|
||||
self.active_controller_generation
|
||||
.fetch_max(generation, Ordering::AcqRel);
|
||||
}
|
||||
|
||||
fn begin_controller_update(
|
||||
&self,
|
||||
generation: u64,
|
||||
) -> Option<ParkingMutexGuard<'_, ()>> {
|
||||
let controller_update = self.controller_update.lock();
|
||||
(self.active_controller_generation.load(Ordering::Acquire) == generation)
|
||||
.then_some(controller_update)
|
||||
}
|
||||
|
||||
/// Reserves bytes against either the adaptive target or the absolute ceiling.
|
||||
pub(crate) fn try_reserve(
|
||||
self: &Arc<Self>,
|
||||
@@ -322,222 +346,6 @@ impl Drop for DirectBufferLease {
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolves the startup hard ceiling from config, cgroup, and host memory.
|
||||
pub(crate) async fn resolve_direct_buffer_hard_limit(configured: usize) -> usize {
|
||||
if configured != 0 {
|
||||
return align_down(configured);
|
||||
}
|
||||
let sample = read_system_memory_sample().await;
|
||||
if sample.total_bytes == 0 {
|
||||
return AUTO_HARD_FALLBACK_BYTES;
|
||||
}
|
||||
let derived = (sample.total_bytes / 4)
|
||||
.clamp(AUTO_HARD_MIN_BYTES as u64, AUTO_HARD_MAX_BYTES as u64)
|
||||
.min(sample.total_bytes);
|
||||
align_down(derived as usize).max(DIRECT_BUFFER_UNIT_BYTES)
|
||||
}
|
||||
|
||||
/// Runs the control-plane loop for Direct budget and shared pool pressure.
|
||||
pub(crate) async fn run_direct_buffer_budget_controller(
|
||||
budget: Arc<DirectBufferBudget>,
|
||||
buffer_pool: Arc<BufferPool>,
|
||||
stats: Arc<Stats>,
|
||||
shared: Arc<ProxySharedState>,
|
||||
max_connections: u32,
|
||||
) {
|
||||
let mut interval = tokio::time::interval(CONTROL_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
|
||||
let mut healthy_streak = 0u8;
|
||||
let mut previous_denied = 0u64;
|
||||
let mut previous_fallback = 0u64;
|
||||
let mut previous_rejected = 0u64;
|
||||
let pool_trim_low = buffer_pool
|
||||
.max_buffers()
|
||||
.min(BUFFER_POOL_TRIM_LOW_WATERMARK);
|
||||
let pool_trim_high = buffer_pool
|
||||
.max_buffers()
|
||||
.min(BUFFER_POOL_TRIM_HIGH_WATERMARK);
|
||||
let mut pool_trim_armed = true;
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
let sample = read_system_memory_sample().await;
|
||||
budget.update_system_sample(sample);
|
||||
|
||||
let snapshot = budget.snapshot();
|
||||
let denied_delta = snapshot
|
||||
.promotion_denied_total
|
||||
.saturating_sub(previous_denied);
|
||||
previous_denied = snapshot.promotion_denied_total;
|
||||
let fallback_delta = snapshot
|
||||
.minimum_fallback_total
|
||||
.saturating_sub(previous_fallback);
|
||||
previous_fallback = snapshot.minimum_fallback_total;
|
||||
let rejected_delta = snapshot
|
||||
.admission_rejected_total
|
||||
.saturating_sub(previous_rejected);
|
||||
previous_rejected = snapshot.admission_rejected_total;
|
||||
|
||||
let connection_pct = connection_fill_pct(stats.as_ref(), max_connections);
|
||||
let memory_available_pct = percentage(sample.available_bytes, sample.total_bytes);
|
||||
let target_utilization_pct = percentage(snapshot.reserved_bytes, snapshot.target_bytes);
|
||||
let pressure = shared.conntrack_pressure_active()
|
||||
|| connection_pct.is_some_and(|value| value >= 85)
|
||||
|| memory_available_pct.is_some_and(|value| value <= 15)
|
||||
|| target_utilization_pct.is_some_and(|value| value >= 90)
|
||||
|| denied_delta > 0
|
||||
|| fallback_delta > 0
|
||||
|| rejected_delta > 0;
|
||||
|
||||
if !pressure {
|
||||
pool_trim_armed = true;
|
||||
} else if pool_trim_armed && buffer_pool.pooled() > pool_trim_high {
|
||||
buffer_pool.trim_to(pool_trim_low);
|
||||
pool_trim_armed = false;
|
||||
}
|
||||
|
||||
let pool_snapshot = buffer_pool.stats();
|
||||
stats.set_buffer_pool_gauges(
|
||||
pool_snapshot.pooled,
|
||||
pool_snapshot.allocated,
|
||||
pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled),
|
||||
);
|
||||
stats.set_buffer_pool_replaced_nonstandard_total(pool_snapshot.replaced_nonstandard);
|
||||
|
||||
let headroom_target = if sample.total_bytes == 0 {
|
||||
snapshot.hard_limit_bytes
|
||||
} else {
|
||||
snapshot
|
||||
.reserved_bytes
|
||||
.saturating_add(sample.available_bytes / 4)
|
||||
.min(snapshot.hard_limit_bytes)
|
||||
};
|
||||
|
||||
if pressure {
|
||||
healthy_streak = 0;
|
||||
let reduced = snapshot.target_bytes.saturating_mul(3) / 4;
|
||||
budget.set_target_bytes(reduced.min(headroom_target));
|
||||
continue;
|
||||
}
|
||||
|
||||
let healthy = memory_available_pct.is_none_or(|value| value >= 30)
|
||||
&& connection_pct.is_none_or(|value| value <= 70);
|
||||
if !healthy {
|
||||
healthy_streak = 0;
|
||||
if headroom_target < snapshot.target_bytes {
|
||||
budget.set_target_bytes(headroom_target);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
healthy_streak = healthy_streak.saturating_add(1);
|
||||
if healthy_streak >= HEALTHY_RECOVERY_SAMPLES {
|
||||
healthy_streak = 0;
|
||||
let increment = (snapshot.target_bytes / 16).max(4 * 1024 * 1024);
|
||||
budget.set_target_bytes(
|
||||
snapshot
|
||||
.target_bytes
|
||||
.saturating_add(increment)
|
||||
.min(headroom_target),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn connection_fill_pct(stats: &Stats, max_connections: u32) -> Option<u8> {
|
||||
if max_connections == 0 {
|
||||
return None;
|
||||
}
|
||||
Some(
|
||||
((stats.get_current_connections_total().saturating_mul(100)) / u64::from(max_connections))
|
||||
.min(100) as u8,
|
||||
)
|
||||
}
|
||||
|
||||
fn percentage(value: u64, total: u64) -> Option<u8> {
|
||||
if total == 0 {
|
||||
return None;
|
||||
}
|
||||
Some(((value.saturating_mul(100)) / total).min(100) as u8)
|
||||
}
|
||||
|
||||
async fn read_system_memory_sample() -> SystemMemorySample {
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
let meminfo = tokio::fs::read_to_string("/proc/meminfo")
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let status = tokio::fs::read_to_string("/proc/self/status")
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let host_total = parse_kib_field(&meminfo, "MemTotal:");
|
||||
let host_available = parse_kib_field(&meminfo, "MemAvailable:");
|
||||
let process_rss = parse_kib_field(&status, "VmRSS:");
|
||||
|
||||
let cgroup_v2_max = read_cgroup_limit("/sys/fs/cgroup/memory.max").await;
|
||||
let cgroup_v2_current = read_u64_file("/sys/fs/cgroup/memory.current").await;
|
||||
let cgroup_v1_max = read_cgroup_limit("/sys/fs/cgroup/memory/memory.limit_in_bytes").await;
|
||||
let cgroup_v1_current = read_u64_file("/sys/fs/cgroup/memory/memory.usage_in_bytes").await;
|
||||
let cgroup_max = cgroup_v2_max.or(cgroup_v1_max);
|
||||
let cgroup_current = cgroup_v2_current.or(cgroup_v1_current);
|
||||
|
||||
let total = match (host_total, cgroup_max) {
|
||||
(0, Some(limit)) => limit,
|
||||
(host, Some(limit)) => host.min(limit),
|
||||
(host, None) => host,
|
||||
};
|
||||
let cgroup_available = cgroup_max
|
||||
.zip(cgroup_current)
|
||||
.map(|(limit, current)| limit.saturating_sub(current));
|
||||
let available = match (host_available, cgroup_available) {
|
||||
(0, Some(value)) => value,
|
||||
(host, Some(value)) => host.min(value),
|
||||
(host, None) => host,
|
||||
};
|
||||
return SystemMemorySample {
|
||||
total_bytes: total,
|
||||
available_bytes: available,
|
||||
process_rss_bytes: process_rss,
|
||||
};
|
||||
}
|
||||
#[cfg(not(target_os = "linux"))]
|
||||
{
|
||||
SystemMemorySample::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
async fn read_cgroup_limit(path: &str) -> Option<u64> {
|
||||
let raw = tokio::fs::read_to_string(path).await.ok()?;
|
||||
let raw = raw.trim();
|
||||
if raw == "max" {
|
||||
return None;
|
||||
}
|
||||
let value = raw.parse::<u64>().ok()?;
|
||||
(value < (1u64 << 60)).then_some(value)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
async fn read_u64_file(path: &str) -> Option<u64> {
|
||||
tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn parse_kib_field(raw: &str, key: &str) -> u64 {
|
||||
raw.lines()
|
||||
.find_map(|line| {
|
||||
let value = line.strip_prefix(key)?.split_whitespace().next()?;
|
||||
value.parse::<u64>().ok()
|
||||
})
|
||||
.unwrap_or(0)
|
||||
.saturating_mul(1024)
|
||||
}
|
||||
|
||||
fn align_up(bytes: usize) -> usize {
|
||||
bytes
|
||||
.div_ceil(DIRECT_BUFFER_UNIT_BYTES)
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
use super::*;
|
||||
use crate::proxy::shared_state::ProxySharedState;
|
||||
use crate::stats::Stats;
|
||||
use crate::stream::BufferPool;
|
||||
|
||||
/// Resolves the startup hard ceiling from config, cgroup, and host memory.
|
||||
pub(crate) async fn resolve_direct_buffer_hard_limit(configured: usize) -> usize {
|
||||
if configured != 0 {
|
||||
return align_down(configured);
|
||||
}
|
||||
let sample = read_system_memory_sample().await;
|
||||
if sample.total_bytes == 0 {
|
||||
return AUTO_HARD_FALLBACK_BYTES;
|
||||
}
|
||||
let derived = (sample.total_bytes / 4)
|
||||
.clamp(AUTO_HARD_MIN_BYTES as u64, AUTO_HARD_MAX_BYTES as u64)
|
||||
.min(sample.total_bytes);
|
||||
align_down(derived as usize).max(DIRECT_BUFFER_UNIT_BYTES)
|
||||
}
|
||||
|
||||
/// Runs the control-plane loop for Direct budget and shared pool pressure.
|
||||
pub(crate) async fn run_direct_buffer_budget_controller(
|
||||
source_generation: u64,
|
||||
budget: Arc<DirectBufferBudget>,
|
||||
buffer_pool: Arc<BufferPool>,
|
||||
stats: Arc<Stats>,
|
||||
shared: Arc<ProxySharedState>,
|
||||
connection_slots: Arc<Semaphore>,
|
||||
max_connections: u32,
|
||||
) {
|
||||
let mut interval = tokio::time::interval(CONTROL_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
|
||||
let mut healthy_streak = 0u8;
|
||||
let mut previous_denied = 0u64;
|
||||
let mut previous_fallback = 0u64;
|
||||
let mut previous_rejected = 0u64;
|
||||
let pool_trim_low = buffer_pool
|
||||
.max_buffers()
|
||||
.min(BUFFER_POOL_TRIM_LOW_WATERMARK);
|
||||
let pool_trim_high = buffer_pool
|
||||
.max_buffers()
|
||||
.min(BUFFER_POOL_TRIM_HIGH_WATERMARK);
|
||||
let mut pool_trim_armed = true;
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
if budget.active_controller_generation.load(Ordering::Acquire) != source_generation {
|
||||
continue;
|
||||
}
|
||||
let sample = read_system_memory_sample().await;
|
||||
let Some(_controller_update) = budget.begin_controller_update(source_generation) else {
|
||||
continue;
|
||||
};
|
||||
budget.update_system_sample(sample);
|
||||
|
||||
let snapshot = budget.snapshot();
|
||||
let denied_delta = snapshot
|
||||
.promotion_denied_total
|
||||
.saturating_sub(previous_denied);
|
||||
previous_denied = snapshot.promotion_denied_total;
|
||||
let fallback_delta = snapshot
|
||||
.minimum_fallback_total
|
||||
.saturating_sub(previous_fallback);
|
||||
previous_fallback = snapshot.minimum_fallback_total;
|
||||
let rejected_delta = snapshot
|
||||
.admission_rejected_total
|
||||
.saturating_sub(previous_rejected);
|
||||
previous_rejected = snapshot.admission_rejected_total;
|
||||
|
||||
let connection_pct = connection_fill_pct(connection_slots.as_ref(), max_connections);
|
||||
let memory_available_pct = percentage(sample.available_bytes, sample.total_bytes);
|
||||
let target_utilization_pct = percentage(snapshot.reserved_bytes, snapshot.target_bytes);
|
||||
let pressure = shared.conntrack_pressure_active()
|
||||
|| connection_pct.is_some_and(|value| value >= 85)
|
||||
|| memory_available_pct.is_some_and(|value| value <= 15)
|
||||
|| target_utilization_pct.is_some_and(|value| value >= 90)
|
||||
|| denied_delta > 0
|
||||
|| fallback_delta > 0
|
||||
|| rejected_delta > 0;
|
||||
|
||||
if !pressure {
|
||||
pool_trim_armed = true;
|
||||
} else if pool_trim_armed && buffer_pool.pooled() > pool_trim_high {
|
||||
buffer_pool.trim_to(pool_trim_low);
|
||||
pool_trim_armed = false;
|
||||
}
|
||||
|
||||
let pool_snapshot = buffer_pool.stats();
|
||||
stats.set_buffer_pool_gauges(
|
||||
pool_snapshot.pooled,
|
||||
pool_snapshot.allocated,
|
||||
pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled),
|
||||
);
|
||||
stats.set_buffer_pool_replaced_nonstandard_total(pool_snapshot.replaced_nonstandard);
|
||||
|
||||
let headroom_target = if sample.total_bytes == 0 {
|
||||
snapshot.hard_limit_bytes
|
||||
} else {
|
||||
snapshot
|
||||
.reserved_bytes
|
||||
.saturating_add(sample.available_bytes / 4)
|
||||
.min(snapshot.hard_limit_bytes)
|
||||
};
|
||||
|
||||
if pressure {
|
||||
healthy_streak = 0;
|
||||
let reduced = snapshot.target_bytes.saturating_mul(3) / 4;
|
||||
budget.set_target_bytes(reduced.min(headroom_target));
|
||||
continue;
|
||||
}
|
||||
|
||||
let healthy = memory_available_pct.is_none_or(|value| value >= 30)
|
||||
&& connection_pct.is_none_or(|value| value <= 70);
|
||||
if !healthy {
|
||||
healthy_streak = 0;
|
||||
if headroom_target < snapshot.target_bytes {
|
||||
budget.set_target_bytes(headroom_target);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
healthy_streak = healthy_streak.saturating_add(1);
|
||||
if healthy_streak >= HEALTHY_RECOVERY_SAMPLES {
|
||||
healthy_streak = 0;
|
||||
let increment = (snapshot.target_bytes / 16).max(4 * 1024 * 1024);
|
||||
budget.set_target_bytes(
|
||||
snapshot
|
||||
.target_bytes
|
||||
.saturating_add(increment)
|
||||
.min(headroom_target),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn connection_fill_pct(
|
||||
connection_slots: &Semaphore,
|
||||
max_connections: u32,
|
||||
) -> Option<u8> {
|
||||
if max_connections == 0 {
|
||||
return None;
|
||||
}
|
||||
let max_connections = max_connections as usize;
|
||||
let active = max_connections.saturating_sub(
|
||||
connection_slots
|
||||
.available_permits()
|
||||
.min(max_connections),
|
||||
);
|
||||
Some((active.saturating_mul(100) / max_connections).min(100) as u8)
|
||||
}
|
||||
|
||||
fn percentage(value: u64, total: u64) -> Option<u8> {
|
||||
if total == 0 {
|
||||
return None;
|
||||
}
|
||||
Some(((value.saturating_mul(100)) / total).min(100) as u8)
|
||||
}
|
||||
|
||||
async fn read_system_memory_sample() -> SystemMemorySample {
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
let meminfo = tokio::fs::read_to_string("/proc/meminfo")
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let status = tokio::fs::read_to_string("/proc/self/status")
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let host_total = parse_kib_field(&meminfo, "MemTotal:");
|
||||
let host_available = parse_kib_field(&meminfo, "MemAvailable:");
|
||||
let process_rss = parse_kib_field(&status, "VmRSS:");
|
||||
|
||||
let cgroup_v2_max = read_cgroup_limit("/sys/fs/cgroup/memory.max").await;
|
||||
let cgroup_v2_current = read_u64_file("/sys/fs/cgroup/memory.current").await;
|
||||
let cgroup_v1_max = read_cgroup_limit("/sys/fs/cgroup/memory/memory.limit_in_bytes").await;
|
||||
let cgroup_v1_current = read_u64_file("/sys/fs/cgroup/memory/memory.usage_in_bytes").await;
|
||||
let cgroup_max = cgroup_v2_max.or(cgroup_v1_max);
|
||||
let cgroup_current = cgroup_v2_current.or(cgroup_v1_current);
|
||||
|
||||
let total = match (host_total, cgroup_max) {
|
||||
(0, Some(limit)) => limit,
|
||||
(host, Some(limit)) => host.min(limit),
|
||||
(host, None) => host,
|
||||
};
|
||||
let cgroup_available = cgroup_max
|
||||
.zip(cgroup_current)
|
||||
.map(|(limit, current)| limit.saturating_sub(current));
|
||||
let available = match (host_available, cgroup_available) {
|
||||
(0, Some(value)) => value,
|
||||
(host, Some(value)) => host.min(value),
|
||||
(host, None) => host,
|
||||
};
|
||||
return SystemMemorySample {
|
||||
total_bytes: total,
|
||||
available_bytes: available,
|
||||
process_rss_bytes: process_rss,
|
||||
};
|
||||
}
|
||||
#[cfg(not(target_os = "linux"))]
|
||||
{
|
||||
SystemMemorySample::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
async fn read_cgroup_limit(path: &str) -> Option<u64> {
|
||||
let raw = tokio::fs::read_to_string(path).await.ok()?;
|
||||
let raw = raw.trim();
|
||||
if raw == "max" {
|
||||
return None;
|
||||
}
|
||||
let value = raw.parse::<u64>().ok()?;
|
||||
(value < (1u64 << 60)).then_some(value)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
async fn read_u64_file(path: &str) -> Option<u64> {
|
||||
tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn parse_kib_field(raw: &str, key: &str) -> u64 {
|
||||
raw.lines()
|
||||
.find_map(|line| {
|
||||
let value = line.strip_prefix(key)?.split_whitespace().next()?;
|
||||
value.parse::<u64>().ok()
|
||||
})
|
||||
.unwrap_or(0)
|
||||
.saturating_mul(1024)
|
||||
}
|
||||
@@ -15,6 +15,7 @@ use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch};
|
||||
use tokio::time::timeout;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
use tracing::{debug, info, trace, warn};
|
||||
|
||||
use crate::config::{ConntrackPressureProfile, ProxyConfig};
|
||||
@@ -155,7 +156,7 @@ const ME_D2C_FLUSH_BATCH_MAX_FRAMES_MIN: usize = 1;
|
||||
const ME_D2C_FLUSH_BATCH_MAX_BYTES_MIN: usize = 4096;
|
||||
const ME_D2C_FRAME_BUF_SHRINK_HYSTERESIS_FACTOR: usize = 2;
|
||||
const ME_D2C_SINGLE_WRITE_COALESCE_MAX_BYTES: usize = 128 * 1024;
|
||||
const QUOTA_RESERVE_SPIN_RETRIES: usize = 32;
|
||||
const QUOTA_RESERVE_ATTEMPTS_PER_ROUND: usize = 4;
|
||||
const QUOTA_RESERVE_BACKOFF_MIN_MS: u64 = 1;
|
||||
const QUOTA_RESERVE_BACKOFF_MAX_MS: u64 = 16;
|
||||
const QUOTA_RESERVE_MAX_BACKOFF_ROUNDS: usize = 16;
|
||||
|
||||
@@ -22,7 +22,7 @@ pub(super) async fn reserve_user_quota_with_yield(
|
||||
let mut backoff_ms = QUOTA_RESERVE_BACKOFF_MIN_MS;
|
||||
let mut backoff_rounds = 0usize;
|
||||
loop {
|
||||
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
|
||||
for _ in 0..QUOTA_RESERVE_ATTEMPTS_PER_ROUND {
|
||||
match quota_handle.try_reserve(bytes, limit) {
|
||||
Ok(reservation) => return Ok(reservation.commit()),
|
||||
Err(QuotaReserveError::LimitExceeded) => {
|
||||
@@ -30,7 +30,6 @@ pub(super) async fn reserve_user_quota_with_yield(
|
||||
}
|
||||
Err(QuotaReserveError::Contended) => {
|
||||
stats.increment_quota_contention_total();
|
||||
std::hint::spin_loop();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,11 +2,15 @@ use super::*;
|
||||
|
||||
// Bounded C2ME sender and downstream writer tasks.
|
||||
mod tasks;
|
||||
// Child-task ownership aborts relay tasks when the parent future is cancelled.
|
||||
mod children;
|
||||
// Conntrack close classification.
|
||||
mod close_reason;
|
||||
|
||||
use children::RelayChildTasks;
|
||||
use close_reason::classify_conntrack_close_reason;
|
||||
use tasks::{run_c2me_sender, run_me_writer};
|
||||
|
||||
struct RelayConnLease {
|
||||
connection: Option<ConnLease>,
|
||||
conn_id: u64,
|
||||
@@ -185,7 +189,7 @@ where
|
||||
let c2me_byte_semaphore = Arc::new(Semaphore::new(c2me_byte_budget));
|
||||
let (c2me_tx, c2me_rx) = mpsc::channel::<C2MeCommand>(c2me_channel_capacity);
|
||||
let me_pool_c2me = me_pool.clone();
|
||||
let mut c2me_sender = tokio::spawn(run_c2me_sender(
|
||||
let c2me_sender = AbortOnDropHandle::new(tokio::spawn(run_c2me_sender(
|
||||
c2me_rx,
|
||||
me_pool_c2me,
|
||||
conn_id,
|
||||
@@ -193,7 +197,7 @@ where
|
||||
peer,
|
||||
translated_local_addr,
|
||||
effective_tag_array,
|
||||
));
|
||||
)));
|
||||
|
||||
let (stop_tx, stop_rx) = oneshot::channel::<()>();
|
||||
let flow_cancel = CancellationToken::new();
|
||||
@@ -208,7 +212,7 @@ where
|
||||
let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone();
|
||||
let bytes_me2c_clone = bytes_me2c.clone();
|
||||
let d2c_flush_policy = MeD2cFlushPolicy::from_config(&config);
|
||||
let mut me_writer = tokio::spawn(run_me_writer(
|
||||
let me_writer = AbortOnDropHandle::new(tokio::spawn(run_me_writer(
|
||||
crypto_writer,
|
||||
me_rx_task,
|
||||
stats_clone,
|
||||
@@ -226,7 +230,13 @@ where
|
||||
session_started_at,
|
||||
conn_id,
|
||||
stop_rx,
|
||||
));
|
||||
)));
|
||||
let mut child_tasks = RelayChildTasks {
|
||||
c2me_sender,
|
||||
me_writer,
|
||||
flow_cancel: flow_cancel.clone(),
|
||||
stop_tx: Some(stop_tx),
|
||||
};
|
||||
|
||||
let mut main_result: Result<()> = Ok(());
|
||||
let mut client_closed = false;
|
||||
@@ -454,28 +464,30 @@ where
|
||||
}
|
||||
|
||||
drop(c2me_tx);
|
||||
let c2me_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut c2me_sender).await {
|
||||
let c2me_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut child_tasks.c2me_sender).await {
|
||||
Ok(joined) => {
|
||||
joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME sender join error: {e}"))))
|
||||
}
|
||||
Err(_) => {
|
||||
stats.increment_me_child_join_timeout_total();
|
||||
stats.increment_me_child_abort_total();
|
||||
c2me_sender.abort();
|
||||
child_tasks.c2me_sender.abort();
|
||||
Err(ProxyError::Proxy("ME sender join timeout".into()))
|
||||
}
|
||||
};
|
||||
|
||||
flow_cancel.cancel();
|
||||
let _ = stop_tx.send(());
|
||||
let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut me_writer).await {
|
||||
if let Some(stop_tx) = child_tasks.stop_tx.take() {
|
||||
let _ = stop_tx.send(());
|
||||
}
|
||||
let mut writer_result = match timeout(ME_CHILD_JOIN_TIMEOUT, &mut child_tasks.me_writer).await {
|
||||
Ok(joined) => {
|
||||
joined.unwrap_or_else(|e| Err(ProxyError::Proxy(format!("ME writer join error: {e}"))))
|
||||
}
|
||||
Err(_) => {
|
||||
stats.increment_me_child_join_timeout_total();
|
||||
stats.increment_me_child_abort_total();
|
||||
me_writer.abort();
|
||||
child_tasks.me_writer.abort();
|
||||
Err(ProxyError::Proxy("ME writer join timeout".into()))
|
||||
}
|
||||
};
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
use super::*;
|
||||
|
||||
pub(super) struct RelayChildTasks {
|
||||
pub(super) c2me_sender: AbortOnDropHandle<Result<()>>,
|
||||
pub(super) me_writer: AbortOnDropHandle<Result<()>>,
|
||||
pub(super) flow_cancel: CancellationToken,
|
||||
pub(super) stop_tx: Option<oneshot::Sender<()>>,
|
||||
}
|
||||
|
||||
impl Drop for RelayChildTasks {
|
||||
fn drop(&mut self) {
|
||||
self.flow_cancel.cancel();
|
||||
if let Some(stop_tx) = self.stop_tx.take() {
|
||||
let _ = stop_tx.send(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
struct DropSignal(Arc<AtomicUsize>);
|
||||
|
||||
impl Drop for DropSignal {
|
||||
fn drop(&mut self) {
|
||||
self.0.fetch_add(1, Ordering::AcqRel);
|
||||
}
|
||||
}
|
||||
|
||||
async fn pending_child(signal: DropSignal) -> Result<()> {
|
||||
let _signal = signal;
|
||||
std::future::pending::<()>().await;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_child_scope_drop_aborts_both_children() {
|
||||
let dropped = Arc::new(AtomicUsize::new(0));
|
||||
let flow_cancel = CancellationToken::new();
|
||||
let (stop_tx, stop_rx) = oneshot::channel();
|
||||
let child_tasks = RelayChildTasks {
|
||||
c2me_sender: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(
|
||||
Arc::clone(&dropped),
|
||||
)))),
|
||||
me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(
|
||||
Arc::clone(&dropped),
|
||||
)))),
|
||||
flow_cancel: flow_cancel.clone(),
|
||||
stop_tx: Some(stop_tx),
|
||||
};
|
||||
|
||||
drop(child_tasks);
|
||||
assert!(flow_cancel.is_cancelled());
|
||||
assert!(stop_rx.await.is_ok());
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while dropped.load(Ordering::Acquire) != 2 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("both relay child futures must be dropped after scope cancellation");
|
||||
}
|
||||
}
|
||||
@@ -74,6 +74,8 @@ pub mod session_eviction;
|
||||
pub mod shared_state;
|
||||
pub mod traffic_limiter;
|
||||
pub(crate) mod user_admission;
|
||||
// Process-wide per-user connection admission remains independent from telemetry.
|
||||
pub(crate) mod user_connection_authority;
|
||||
|
||||
pub use client::ClientHandler;
|
||||
#[allow(unused_imports)]
|
||||
|
||||
+49
-69
@@ -16,7 +16,7 @@ mod quota;
|
||||
pub(super) use self::combined::CombinedStream;
|
||||
pub(super) use self::counters::SharedCounters;
|
||||
pub(super) use self::quota::is_quota_io_error;
|
||||
use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error};
|
||||
use self::quota::{QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL, quota_io_error};
|
||||
pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check};
|
||||
|
||||
/// Transparent I/O wrapper that tracks per-user statistics and activity.
|
||||
@@ -218,48 +218,33 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
let mut quota_reservation = None;
|
||||
let mut read_limit = buf.remaining();
|
||||
if let Some(limit) = this.quota_limit {
|
||||
let used_before = this.quota_handle.used();
|
||||
let remaining = limit.saturating_sub(used_before);
|
||||
if remaining == 0 {
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
remaining_before = Some(remaining);
|
||||
read_limit = read_limit.min(remaining as usize);
|
||||
if read_limit == 0 {
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
|
||||
let desired = read_limit as u64;
|
||||
let mut reserve_rounds = 0usize;
|
||||
while quota_reservation.is_none() {
|
||||
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
|
||||
match this.quota_handle.try_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
quota_reservation = Some(reservation);
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::LimitExceeded) => {
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::Contended) => {
|
||||
this.stats.increment_quota_contention_total();
|
||||
}
|
||||
for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL {
|
||||
let used_before = this.quota_handle.used();
|
||||
let remaining = limit.saturating_sub(used_before);
|
||||
if remaining == 0 {
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
let desired = remaining.min(read_limit as u64);
|
||||
match this.quota_handle.try_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
remaining_before = Some(remaining);
|
||||
read_limit = desired as usize;
|
||||
quota_reservation = Some(reservation);
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::LimitExceeded)
|
||||
| Err(crate::stats::QuotaReserveError::Contended) => {
|
||||
this.stats.increment_quota_contention_total();
|
||||
}
|
||||
}
|
||||
|
||||
if quota_reservation.is_none() {
|
||||
reserve_rounds = reserve_rounds.saturating_add(1);
|
||||
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
|
||||
this.stats.increment_quota_contention_timeout_total();
|
||||
if this.arm_quota_wait(cx).is_pending() {
|
||||
return Poll::Pending;
|
||||
}
|
||||
reserve_rounds = 0;
|
||||
}
|
||||
}
|
||||
if quota_reservation.is_none() {
|
||||
this.stats.increment_quota_contention_timeout_total();
|
||||
if this.arm_quota_wait(cx).is_ready() {
|
||||
cx.waker().wake_by_ref();
|
||||
}
|
||||
return Poll::Pending;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -404,8 +389,7 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
let mut quota_reservation = None;
|
||||
if let Some(limit) = this.quota_limit {
|
||||
if !write_buf.is_empty() {
|
||||
let mut reserve_rounds = 0usize;
|
||||
while quota_reservation.is_none() {
|
||||
for _ in 0..QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL {
|
||||
let used_before = this.quota_handle.used();
|
||||
let remaining = limit.saturating_sub(used_before);
|
||||
if remaining == 0 {
|
||||
@@ -415,36 +399,32 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
remaining_before = Some(remaining);
|
||||
|
||||
let desired = remaining.min(write_buf.len() as u64);
|
||||
let mut saw_contention = false;
|
||||
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
|
||||
match this.quota_handle.try_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
quota_reservation = Some(reservation);
|
||||
write_buf = &write_buf[..desired as usize];
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::LimitExceeded) => {
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::Contended) => {
|
||||
this.stats.increment_quota_contention_total();
|
||||
saw_contention = true;
|
||||
}
|
||||
match this.quota_handle.try_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
quota_reservation = Some(reservation);
|
||||
write_buf = &write_buf[..desired as usize];
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::LimitExceeded)
|
||||
| Err(crate::stats::QuotaReserveError::Contended) => {
|
||||
this.stats.increment_quota_contention_total();
|
||||
}
|
||||
}
|
||||
|
||||
if quota_reservation.is_none() {
|
||||
reserve_rounds = reserve_rounds.saturating_add(1);
|
||||
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
|
||||
this.stats.increment_quota_contention_timeout_total();
|
||||
Self::arm_wait(&mut this.quota_wait, false, false);
|
||||
let _ =
|
||||
Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up);
|
||||
return Poll::Pending;
|
||||
} else if saw_contention {
|
||||
std::hint::spin_loop();
|
||||
}
|
||||
}
|
||||
if quota_reservation.is_none() {
|
||||
this.stats.increment_quota_contention_timeout_total();
|
||||
Self::arm_wait(&mut this.quota_wait, false, false);
|
||||
if Self::poll_wait(
|
||||
&mut this.quota_wait,
|
||||
cx,
|
||||
None,
|
||||
RateDirection::Up,
|
||||
)
|
||||
.is_ready()
|
||||
{
|
||||
cx.waker().wake_by_ref();
|
||||
}
|
||||
return Poll::Pending;
|
||||
}
|
||||
} else {
|
||||
let used_before = this.quota_handle.used();
|
||||
|
||||
@@ -27,8 +27,7 @@ const QUOTA_NEAR_LIMIT_BYTES: u64 = 64 * 1024;
|
||||
const QUOTA_LARGE_CHARGE_BYTES: u64 = 16 * 1024;
|
||||
const QUOTA_ADAPTIVE_INTERVAL_MIN_BYTES: u64 = 4 * 1024;
|
||||
const QUOTA_ADAPTIVE_INTERVAL_MAX_BYTES: u64 = 64 * 1024;
|
||||
pub(super) const QUOTA_RESERVE_SPIN_RETRIES: usize = 64;
|
||||
pub(super) const QUOTA_RESERVE_MAX_ROUNDS: usize = 8;
|
||||
pub(super) const QUOTA_RESERVE_MAX_ATTEMPTS_PER_POLL: usize = 4;
|
||||
|
||||
#[inline]
|
||||
pub(in crate::proxy::relay) fn quota_adaptive_interval_bytes(remaining_before: u64) -> u64 {
|
||||
|
||||
@@ -123,6 +123,19 @@ impl ProxySharedState {
|
||||
pub(crate) fn new_with_direct_buffer_budget_and_user_admission(
|
||||
direct_buffer_budget: Arc<DirectBufferBudget>,
|
||||
user_admission: Arc<UserAdmissionAuthority>,
|
||||
) -> Arc<Self> {
|
||||
Self::new_with_process_authorities(
|
||||
direct_buffer_budget,
|
||||
TrafficLimiter::new(),
|
||||
user_admission,
|
||||
)
|
||||
}
|
||||
|
||||
/// Creates generation-local caches around process-owned data-plane authorities.
|
||||
pub(crate) fn new_with_process_authorities(
|
||||
direct_buffer_budget: Arc<DirectBufferBudget>,
|
||||
traffic_limiter: Arc<TrafficLimiter>,
|
||||
user_admission: Arc<UserAdmissionAuthority>,
|
||||
) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
handshake: HandshakeSharedState {
|
||||
@@ -163,7 +176,7 @@ impl ProxySharedState {
|
||||
relay_idle_registry: RelayIdleCandidateRegistry::default(),
|
||||
relay_idle_mark_seq: AtomicU64::new(0),
|
||||
},
|
||||
traffic_limiter: TrafficLimiter::new(),
|
||||
traffic_limiter,
|
||||
direct_buffer_budget,
|
||||
user_admission,
|
||||
conntrack_pressure_active: AtomicBool::new(false),
|
||||
|
||||
@@ -276,14 +276,13 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() {
|
||||
|
||||
ip_tracker.set_user_limit(&user, 1).await;
|
||||
ip_tracker.check_and_add(&user, ip).await.unwrap();
|
||||
stats.increment_user_curr_connects(&user);
|
||||
|
||||
assert_eq!(ip_tracker.get_active_ip_count(&user).await, 1);
|
||||
assert_eq!(stats.get_user_curr_connects(&user), 1);
|
||||
|
||||
let reservation =
|
||||
UserConnectionReservation::new(stats.clone(), ip_tracker.clone(), user.clone(), ip, true);
|
||||
|
||||
assert_eq!(stats.get_user_curr_connects(&user), 1);
|
||||
|
||||
// Drop the reservation synchronously without any tokio::spawn/await yielding!
|
||||
drop(reservation);
|
||||
|
||||
@@ -304,6 +303,117 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() {
|
||||
assert_eq!(ip_tracker.get_active_ip_count(&user).await, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_ip_admission_releases_process_connection_permit() {
|
||||
let ip_tracker = Arc::new(UserIpTracker::new());
|
||||
let stats = Arc::new(Stats::new());
|
||||
let user = "cancelled-admission-user";
|
||||
let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap();
|
||||
let mut config = ProxyConfig::default();
|
||||
config
|
||||
.access
|
||||
.user_max_tcp_conns
|
||||
.insert(user.to_string(), 1);
|
||||
|
||||
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
|
||||
let (release_tx, release_rx) = tokio::sync::oneshot::channel();
|
||||
let held_tracker = Arc::clone(&ip_tracker);
|
||||
let held_user = user.to_string();
|
||||
let holder = tokio::spawn(async move {
|
||||
held_tracker
|
||||
.hold_user_shard_for_tests(&held_user, entered_tx, release_rx)
|
||||
.await;
|
||||
});
|
||||
entered_rx.await.unwrap();
|
||||
|
||||
let acquire_stats = Arc::clone(&stats);
|
||||
let acquire_tracker = Arc::clone(&ip_tracker);
|
||||
let acquire_config = config.clone();
|
||||
let acquire = tokio::spawn(async move {
|
||||
acquire_user_connection_reservation(
|
||||
user,
|
||||
&acquire_config,
|
||||
acquire_stats,
|
||||
peer_addr,
|
||||
acquire_tracker,
|
||||
)
|
||||
.await
|
||||
});
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while stats.get_process_user_curr_connects(user) != 1 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("connection permit must be acquired before IP admission completes");
|
||||
|
||||
acquire.abort();
|
||||
let acquire_result = acquire.await;
|
||||
assert!(
|
||||
acquire_result
|
||||
.as_ref()
|
||||
.is_err_and(tokio::task::JoinError::is_cancelled)
|
||||
);
|
||||
assert_eq!(stats.get_process_user_curr_connects(user), 0);
|
||||
assert_eq!(stats.get_user_curr_connects(user), 0);
|
||||
assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 0);
|
||||
|
||||
let _ = release_tx.send(());
|
||||
holder.await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_async_release_preserves_ip_cleanup_ownership() {
|
||||
let ip_tracker = Arc::new(UserIpTracker::new());
|
||||
let stats = Arc::new(Stats::new());
|
||||
let user = "cancelled-release-user";
|
||||
let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap();
|
||||
let mut config = ProxyConfig::default();
|
||||
config
|
||||
.access
|
||||
.user_max_tcp_conns
|
||||
.insert(user.to_string(), 1);
|
||||
|
||||
let reservation = acquire_user_connection_reservation(
|
||||
user,
|
||||
&config,
|
||||
Arc::clone(&stats),
|
||||
peer_addr,
|
||||
Arc::clone(&ip_tracker),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(stats.get_process_user_curr_connects(user), 1);
|
||||
assert_eq!(ip_tracker.get_active_ip_count(user).await, 1);
|
||||
|
||||
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
|
||||
let (release_tx, release_rx) = tokio::sync::oneshot::channel();
|
||||
let held_tracker = Arc::clone(&ip_tracker);
|
||||
let held_user = user.to_string();
|
||||
let holder = tokio::spawn(async move {
|
||||
held_tracker
|
||||
.hold_user_shard_for_tests(&held_user, entered_tx, release_rx)
|
||||
.await;
|
||||
});
|
||||
entered_rx.await.unwrap();
|
||||
|
||||
let release = tokio::spawn(reservation.release());
|
||||
tokio::task::yield_now().await;
|
||||
release.abort();
|
||||
assert!(release.await.unwrap_err().is_cancelled());
|
||||
|
||||
assert_eq!(stats.get_process_user_curr_connects(user), 0);
|
||||
assert_eq!(stats.get_user_curr_connects(user), 0);
|
||||
assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 1);
|
||||
|
||||
let _ = release_tx.send(());
|
||||
holder.await.unwrap();
|
||||
ip_tracker.drain_cleanup_queue().await;
|
||||
assert_eq!(ip_tracker.get_active_ip_count(user).await, 0);
|
||||
assert_eq!(ip_tracker.cleanup_queue_len_for_tests(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_task_abort_releases_user_gate_and_ip_reservation() {
|
||||
let tg_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
@@ -2813,7 +2923,11 @@ async fn tcp_limit_rejection_does_not_reserve_ip_or_trigger_rollback() {
|
||||
.insert("user".to_string(), 1);
|
||||
|
||||
let stats = Stats::new();
|
||||
stats.increment_user_curr_connects("user");
|
||||
let _existing_connection = stats
|
||||
.connection_authority()
|
||||
.try_acquire("user", Some(1))
|
||||
.expect("existing connection must occupy the process admission slot");
|
||||
let _existing_observation = stats.observe_user_current_connection("user");
|
||||
|
||||
let ip_tracker = UserIpTracker::new();
|
||||
let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap();
|
||||
@@ -2853,7 +2967,11 @@ async fn zero_tcp_limit_uses_global_fallback_and_rejects_without_side_effects()
|
||||
config.access.user_max_tcp_conns_global_each = 1;
|
||||
|
||||
let stats = Stats::new();
|
||||
stats.increment_user_curr_connects("user");
|
||||
let _existing_connection = stats
|
||||
.connection_authority()
|
||||
.try_acquire("user", Some(1))
|
||||
.expect("existing connection must occupy the process admission slot");
|
||||
let _existing_observation = stats.observe_user_current_connection("user");
|
||||
let ip_tracker = UserIpTracker::new();
|
||||
let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap();
|
||||
|
||||
@@ -2914,7 +3032,11 @@ async fn global_tcp_fallback_applies_when_per_user_limit_is_missing() {
|
||||
config.access.user_max_tcp_conns_global_each = 1;
|
||||
|
||||
let stats = Stats::new();
|
||||
stats.increment_user_curr_connects("user");
|
||||
let _existing_connection = stats
|
||||
.connection_authority()
|
||||
.try_acquire("user", Some(1))
|
||||
.expect("existing connection must occupy the process admission slot");
|
||||
let _existing_observation = stats.observe_user_current_connection("user");
|
||||
let ip_tracker = UserIpTracker::new();
|
||||
let peer_addr: SocketAddr = "198.51.100.213:50003".parse().unwrap();
|
||||
|
||||
@@ -4024,7 +4146,11 @@ async fn concurrent_limit_rejections_from_mixed_ips_leave_no_ip_footprint() {
|
||||
|
||||
let config = Arc::new(config);
|
||||
let stats = Arc::new(Stats::new());
|
||||
stats.increment_user_curr_connects("user");
|
||||
let _existing_connection = stats
|
||||
.connection_authority()
|
||||
.try_acquire("user", Some(1))
|
||||
.expect("existing connection must occupy the process admission slot");
|
||||
let _existing_observation = stats.observe_user_current_connection("user");
|
||||
let ip_tracker = Arc::new(UserIpTracker::new());
|
||||
|
||||
let mut tasks = tokio::task::JoinSet::new();
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use super::*;
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
#[test]
|
||||
fn lease_drop_releases_the_complete_reservation() {
|
||||
@@ -36,3 +37,62 @@ fn growth_and_shrink_keep_accounting_balanced() {
|
||||
drop(lease);
|
||||
assert_eq!(budget.snapshot().reserved_bytes, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_generations_share_one_absolute_reservation_envelope() {
|
||||
let first_generation = DirectBufferBudget::new(16 * 1024);
|
||||
let second_generation = Arc::clone(&first_generation);
|
||||
let first = first_generation
|
||||
.try_reserve(12 * 1024, true)
|
||||
.expect("first generation reservation must fit");
|
||||
|
||||
assert!(second_generation.try_reserve(8 * 1024, true).is_none());
|
||||
assert_eq!(second_generation.snapshot().reserved_bytes, 12 * 1024);
|
||||
|
||||
drop(first);
|
||||
assert!(second_generation.try_reserve(8 * 1024, true).is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stale_runtime_cannot_reclaim_direct_controller_ownership() {
|
||||
let budget = DirectBufferBudget::new(16 * 1024);
|
||||
budget.activate_controller(2);
|
||||
budget.activate_controller(1);
|
||||
|
||||
assert_eq!(
|
||||
budget.active_controller_generation.load(Ordering::Acquire),
|
||||
2
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn controller_handoff_waits_for_inflight_update_and_fences_old_generation() {
|
||||
let budget = DirectBufferBudget::new(16 * 1024);
|
||||
budget.activate_controller(1);
|
||||
let update = budget.begin_controller_update(1).unwrap();
|
||||
let (activated_tx, activated_rx) = std::sync::mpsc::channel();
|
||||
let next_budget = Arc::clone(&budget);
|
||||
let activation = std::thread::spawn(move || {
|
||||
next_budget.activate_controller(2);
|
||||
activated_tx.send(()).unwrap();
|
||||
});
|
||||
|
||||
assert!(activated_rx
|
||||
.recv_timeout(Duration::from_millis(50))
|
||||
.is_err());
|
||||
drop(update);
|
||||
activated_rx.recv_timeout(Duration::from_secs(1)).unwrap();
|
||||
activation.join().unwrap();
|
||||
|
||||
assert!(budget.begin_controller_update(1).is_none());
|
||||
assert!(budget.begin_controller_update(2).is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_pressure_uses_process_wide_slot_ownership() {
|
||||
let slots = Arc::new(Semaphore::new(10));
|
||||
let _old_generation = Arc::clone(&slots).try_acquire_many_owned(3).unwrap();
|
||||
let _new_generation = Arc::clone(&slots).try_acquire_many_owned(2).unwrap();
|
||||
|
||||
assert_eq!(connection_fill_pct(slots.as_ref(), 10), Some(50));
|
||||
}
|
||||
|
||||
@@ -142,6 +142,7 @@ enum CidrPolicyMatch<'a> {
|
||||
#[derive(Default)]
|
||||
struct PolicySnapshot {
|
||||
revision: u64,
|
||||
source_generation: u64,
|
||||
user_limits: HashMap<String, RateLimitBps>,
|
||||
cidr_rules_v4: Vec<CidrRule>,
|
||||
cidr_rules_v6: Vec<CidrRule>,
|
||||
@@ -175,7 +176,6 @@ pub struct TrafficLease {
|
||||
pub struct TrafficLimiter {
|
||||
policy: ArcSwap<PolicySnapshot>,
|
||||
policy_update: ParkingMutex<()>,
|
||||
published_revision: AtomicU64,
|
||||
user_buckets: ShardedRegistry<UserBucket>,
|
||||
cidr_buckets: ShardedRegistry<CidrBucket>,
|
||||
user_scope: ScopeMetrics,
|
||||
|
||||
@@ -2,20 +2,16 @@ use super::*;
|
||||
|
||||
impl TrafficLease {
|
||||
fn current_binding(&self) -> Arc<TrafficLeaseBinding> {
|
||||
let published_revision = self.limiter.published_revision.load(Ordering::Acquire);
|
||||
let policy = self.limiter.policy.load();
|
||||
let current = self.binding.load_full();
|
||||
if current.revision == published_revision {
|
||||
if current.revision == policy.revision {
|
||||
return current;
|
||||
}
|
||||
drop(policy);
|
||||
|
||||
let refresh = self.refresh.lock();
|
||||
let published_revision = self.limiter.published_revision.load(Ordering::Acquire);
|
||||
let current = self.binding.load_full();
|
||||
if current.revision == published_revision {
|
||||
return current;
|
||||
}
|
||||
let policy_update = self.limiter.policy_update.lock();
|
||||
let _refresh = self.refresh.lock();
|
||||
let policy = self.limiter.policy.load_full();
|
||||
let current = self.binding.load_full();
|
||||
if current.revision == policy.revision {
|
||||
return current;
|
||||
}
|
||||
@@ -23,9 +19,6 @@ impl TrafficLease {
|
||||
.limiter
|
||||
.build_binding(&self.user, self.client_ip, &policy);
|
||||
self.binding.store(Arc::clone(&next));
|
||||
drop(policy_update);
|
||||
drop(refresh);
|
||||
self.limiter.maybe_cleanup();
|
||||
next
|
||||
}
|
||||
|
||||
|
||||
@@ -6,7 +6,6 @@ impl TrafficLimiter {
|
||||
Arc::new(Self {
|
||||
policy: ArcSwap::from_pointee(PolicySnapshot::default()),
|
||||
policy_update: ParkingMutex::new(()),
|
||||
published_revision: AtomicU64::new(0),
|
||||
user_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
|
||||
cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
|
||||
user_scope: ScopeMetrics::default(),
|
||||
@@ -15,16 +14,31 @@ impl TrafficLimiter {
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub fn apply_policy(
|
||||
&self,
|
||||
user_limits: HashMap<String, RateLimitBps>,
|
||||
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
|
||||
) {
|
||||
let policy_update = self.policy_update.lock();
|
||||
// Revision wrap could otherwise let an old lease restore stale rates.
|
||||
let Some(revision) = self.policy.load().revision.checked_add(1) else {
|
||||
return;
|
||||
};
|
||||
let _ = self.apply_policy_inner(None, user_limits, cidr_limits);
|
||||
}
|
||||
|
||||
/// Publishes policy only when the source runtime is not older than the active source.
|
||||
pub(crate) fn apply_policy_from_source(
|
||||
&self,
|
||||
source_generation: u64,
|
||||
user_limits: HashMap<String, RateLimitBps>,
|
||||
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
|
||||
) -> bool {
|
||||
self.apply_policy_inner(Some(source_generation), user_limits, cidr_limits)
|
||||
}
|
||||
|
||||
fn apply_policy_inner(
|
||||
&self,
|
||||
source_generation: Option<u64>,
|
||||
user_limits: HashMap<String, RateLimitBps>,
|
||||
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
|
||||
) -> bool {
|
||||
let filtered_users = user_limits
|
||||
.into_iter()
|
||||
.filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0)
|
||||
@@ -77,6 +91,17 @@ impl TrafficLimiter {
|
||||
let cidr_policy_entries =
|
||||
cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len();
|
||||
|
||||
let policy_update = self.policy_update.lock();
|
||||
let current = self.policy.load_full();
|
||||
if source_generation.is_some_and(|source| source < current.source_generation) {
|
||||
return false;
|
||||
}
|
||||
// Revision wrap could otherwise let an old lease restore stale rates.
|
||||
let Some(revision) = current.revision.checked_add(1) else {
|
||||
return false;
|
||||
};
|
||||
let source_generation = source_generation.unwrap_or(current.source_generation);
|
||||
|
||||
self.user_scope
|
||||
.policy_entries
|
||||
.store(filtered_users.len() as u64, Ordering::Relaxed);
|
||||
@@ -86,6 +111,7 @@ impl TrafficLimiter {
|
||||
|
||||
self.policy.store(Arc::new(PolicySnapshot {
|
||||
revision,
|
||||
source_generation,
|
||||
user_limits: filtered_users,
|
||||
cidr_rules_v4,
|
||||
cidr_rules_v6,
|
||||
@@ -93,10 +119,10 @@ impl TrafficLimiter {
|
||||
cidr_auto_rules_v6,
|
||||
cidr_rule_keys,
|
||||
}));
|
||||
self.published_revision.store(revision, Ordering::Release);
|
||||
|
||||
drop(policy_update);
|
||||
self.maybe_cleanup();
|
||||
true
|
||||
}
|
||||
|
||||
pub fn acquire_lease(
|
||||
@@ -104,11 +130,8 @@ impl TrafficLimiter {
|
||||
user: &str,
|
||||
client_ip: IpAddr,
|
||||
) -> Option<Arc<TrafficLease>> {
|
||||
let policy_update = self.policy_update.lock();
|
||||
let policy = self.policy.load_full();
|
||||
let binding = self.build_binding(user, client_ip, &policy);
|
||||
drop(policy_update);
|
||||
self.maybe_cleanup();
|
||||
Some(Arc::new(TrafficLease {
|
||||
limiter: Arc::clone(self),
|
||||
user: user.to_string(),
|
||||
|
||||
@@ -5,6 +5,79 @@ fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps {
|
||||
RateLimitBps { up_bps, down_bps }
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stale_runtime_cannot_overwrite_newer_rate_policy() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
let mut newer = HashMap::new();
|
||||
newer.insert("alice".to_string(), rate(2_000, 3_000));
|
||||
assert!(limiter.apply_policy_from_source(2, newer, HashMap::new()));
|
||||
|
||||
let mut stale = HashMap::new();
|
||||
stale.insert("alice".to_string(), rate(1_000, 1_000));
|
||||
assert!(!limiter.apply_policy_from_source(1, stale, HashMap::new()));
|
||||
|
||||
let policy = limiter.policy.load_full();
|
||||
assert_eq!(policy.source_generation, 2);
|
||||
assert_eq!(policy.user_limits["alice"].up_bps, 2_000);
|
||||
assert_eq!(policy.user_limits["alice"].down_bps, 3_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn active_runtime_can_publish_same_generation_rate_update() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
assert!(limiter.apply_policy_from_source(4, HashMap::new(), HashMap::new()));
|
||||
let mut updated = HashMap::new();
|
||||
updated.insert("alice".to_string(), rate(4_000, 5_000));
|
||||
|
||||
assert!(limiter.apply_policy_from_source(4, updated, HashMap::new()));
|
||||
|
||||
let policy = limiter.policy.load_full();
|
||||
assert_eq!(policy.source_generation, 4);
|
||||
assert_eq!(policy.user_limits["alice"].up_bps, 4_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lease_acquisition_and_refresh_do_not_wait_for_policy_publication_lock() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
let mut initial = HashMap::new();
|
||||
initial.insert("alice".to_string(), rate(1_000, 1_000));
|
||||
limiter.apply_policy(initial, HashMap::new());
|
||||
let lease = limiter
|
||||
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
|
||||
.unwrap();
|
||||
|
||||
let mut updated = HashMap::new();
|
||||
updated.insert("alice".to_string(), rate(2_000, 2_000));
|
||||
limiter.apply_policy(updated, HashMap::new());
|
||||
|
||||
let publication = limiter.policy_update.lock();
|
||||
let (completed_tx, completed_rx) = std::sync::mpsc::channel();
|
||||
let acquire_limiter = Arc::clone(&limiter);
|
||||
let acquire_tx = completed_tx.clone();
|
||||
let acquire = std::thread::spawn(move || {
|
||||
let _lease = acquire_limiter
|
||||
.acquire_lease("bob", "203.0.113.8".parse().unwrap())
|
||||
.unwrap();
|
||||
acquire_tx.send(()).unwrap();
|
||||
});
|
||||
let refresh = std::thread::spawn(move || {
|
||||
let _ = lease.try_consume(RateDirection::Up, 1);
|
||||
completed_tx.send(()).unwrap();
|
||||
});
|
||||
|
||||
let first_completed = completed_rx
|
||||
.recv_timeout(std::time::Duration::from_secs(1))
|
||||
.is_ok();
|
||||
let second_completed = completed_rx
|
||||
.recv_timeout(std::time::Duration::from_secs(1))
|
||||
.is_ok();
|
||||
drop(publication);
|
||||
acquire.join().unwrap();
|
||||
refresh.join().unwrap();
|
||||
|
||||
assert!(first_completed && second_completed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_cidr_rule_wins_over_auto_template() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use dashmap::DashMap;
|
||||
use dashmap::mapref::entry::Entry;
|
||||
|
||||
/// Process-wide per-user connection admission shared by runtime generations.
|
||||
#[derive(Default)]
|
||||
pub(crate) struct UserConnectionAuthority {
|
||||
active: DashMap<String, u64>,
|
||||
}
|
||||
|
||||
/// Owns one exact connection slot until the authenticated connection exits.
|
||||
#[must_use = "connection permits must be retained for the connection lifetime"]
|
||||
pub(crate) struct UserConnectionPermit {
|
||||
authority: Arc<UserConnectionAuthority>,
|
||||
user: String,
|
||||
}
|
||||
|
||||
impl UserConnectionAuthority {
|
||||
/// Acquires a slot without consulting optional telemetry state.
|
||||
pub(crate) fn try_acquire(
|
||||
self: &Arc<Self>,
|
||||
user: &str,
|
||||
limit: Option<u64>,
|
||||
) -> Option<UserConnectionPermit> {
|
||||
match self.active.entry(user.to_string()) {
|
||||
Entry::Occupied(mut entry) => {
|
||||
if limit.is_some_and(|max| *entry.get() >= max) {
|
||||
return None;
|
||||
}
|
||||
let next = entry.get().checked_add(1)?;
|
||||
*entry.get_mut() = next;
|
||||
}
|
||||
Entry::Vacant(entry) => {
|
||||
if limit == Some(0) {
|
||||
return None;
|
||||
}
|
||||
entry.insert(1);
|
||||
}
|
||||
}
|
||||
Some(UserConnectionPermit {
|
||||
authority: Arc::clone(self),
|
||||
user: user.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Returns the authoritative active connection count for one username.
|
||||
pub(crate) fn active(&self, user: &str) -> u64 {
|
||||
self.active.get(user).map(|entry| *entry).unwrap_or(0)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn tracked_users(&self) -> usize {
|
||||
self.active.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for UserConnectionPermit {
|
||||
fn drop(&mut self) {
|
||||
let Entry::Occupied(mut entry) = self.authority.active.entry(self.user.clone()) else {
|
||||
debug_assert!(false, "connection permit owner entry disappeared");
|
||||
return;
|
||||
};
|
||||
debug_assert!(*entry.get() > 0, "connection permit counter underflow");
|
||||
if *entry.get() <= 1 {
|
||||
entry.remove();
|
||||
} else {
|
||||
*entry.get_mut() -= 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::Barrier;
|
||||
|
||||
#[test]
|
||||
fn permit_drop_releases_and_removes_zero_entry() {
|
||||
let authority = Arc::new(UserConnectionAuthority::default());
|
||||
let permit = authority.try_acquire("alice", Some(1)).unwrap();
|
||||
assert_eq!(authority.active("alice"), 1);
|
||||
assert!(authority.try_acquire("alice", Some(1)).is_none());
|
||||
|
||||
drop(permit);
|
||||
|
||||
assert_eq!(authority.active("alice"), 0);
|
||||
assert_eq!(authority.tracked_users(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_acquire_never_exceeds_limit() {
|
||||
const CONTENDERS: usize = 64;
|
||||
const LIMIT: u64 = 7;
|
||||
|
||||
let authority = Arc::new(UserConnectionAuthority::default());
|
||||
let barrier = Arc::new(Barrier::new(CONTENDERS + 1));
|
||||
let (permit_tx, permit_rx) = std::sync::mpsc::channel();
|
||||
let mut threads = Vec::with_capacity(CONTENDERS);
|
||||
for _ in 0..CONTENDERS {
|
||||
let authority = Arc::clone(&authority);
|
||||
let barrier = Arc::clone(&barrier);
|
||||
let permit_tx = permit_tx.clone();
|
||||
threads.push(std::thread::spawn(move || {
|
||||
barrier.wait();
|
||||
permit_tx
|
||||
.send(authority.try_acquire("alice", Some(LIMIT)))
|
||||
.unwrap();
|
||||
}));
|
||||
}
|
||||
drop(permit_tx);
|
||||
barrier.wait();
|
||||
for thread in threads {
|
||||
thread.join().unwrap();
|
||||
}
|
||||
let permits = permit_rx.into_iter().flatten().collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(permits.len() as u64, LIMIT);
|
||||
assert_eq!(authority.active("alice"), LIMIT);
|
||||
drop(permits);
|
||||
assert_eq!(authority.active("alice"), 0);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user