Process-wide concurrency + Cancellation ownership fixes

This commit is contained in:
Alexey
2026-09-23 22:30:51 +03:00
parent baa9bfbb01
commit f1107c21d9
47 changed files with 2177 additions and 762 deletions
+87 -47
View File
@@ -15,7 +15,8 @@ use crate::proxy::middle_relay::{handle_via_middle_proxy, handle_via_middle_prox
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::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,
))
}
+5 -5
View File
@@ -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(())
}
}
+28 -220
View File
@@ -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)
}
+2 -1
View File
@@ -15,6 +15,7 @@ use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch};
use tokio::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;
+1 -2
View File
@@ -22,7 +22,7 @@ pub(super) async fn reserve_user_quota_with_yield(
let mut backoff_ms = QUOTA_RESERVE_BACKOFF_MIN_MS;
let mut backoff_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();
}
}
}
+21 -9
View File
@@ -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");
}
}
+2
View File
@@ -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
View File
@@ -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();
+1 -2
View File
@@ -27,8 +27,7 @@ const QUOTA_NEAR_LIMIT_BYTES: u64 = 64 * 1024;
const QUOTA_LARGE_CHARGE_BYTES: u64 = 16 * 1024;
const QUOTA_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 {
+14 -1
View File
@@ -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),
+133 -7
View File
@@ -276,14 +276,13 @@ async fn user_connection_reservation_drop_enqueues_cleanup_synchronously() {
ip_tracker.set_user_limit(&user, 1).await;
ip_tracker.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));
}
+1 -1
View File
@@ -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,
+5 -12
View File
@@ -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
}
+33 -10
View File
@@ -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(),
+73
View File
@@ -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();
+123
View File
@@ -0,0 +1,123 @@
use std::sync::Arc;
use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
/// Process-wide per-user connection admission shared by runtime generations.
#[derive(Default)]
pub(crate) struct UserConnectionAuthority {
active: DashMap<String, u64>,
}
/// Owns one exact connection slot until the authenticated connection exits.
#[must_use = "connection permits must be retained for the connection lifetime"]
pub(crate) struct UserConnectionPermit {
authority: Arc<UserConnectionAuthority>,
user: String,
}
impl UserConnectionAuthority {
/// Acquires a slot without consulting optional telemetry state.
pub(crate) fn try_acquire(
self: &Arc<Self>,
user: &str,
limit: Option<u64>,
) -> Option<UserConnectionPermit> {
match self.active.entry(user.to_string()) {
Entry::Occupied(mut entry) => {
if limit.is_some_and(|max| *entry.get() >= max) {
return None;
}
let next = entry.get().checked_add(1)?;
*entry.get_mut() = next;
}
Entry::Vacant(entry) => {
if limit == Some(0) {
return None;
}
entry.insert(1);
}
}
Some(UserConnectionPermit {
authority: Arc::clone(self),
user: user.to_string(),
})
}
/// Returns the authoritative active connection count for one username.
pub(crate) fn active(&self, user: &str) -> u64 {
self.active.get(user).map(|entry| *entry).unwrap_or(0)
}
#[cfg(test)]
fn tracked_users(&self) -> usize {
self.active.len()
}
}
impl Drop for UserConnectionPermit {
fn drop(&mut self) {
let Entry::Occupied(mut entry) = self.authority.active.entry(self.user.clone()) else {
debug_assert!(false, "connection permit owner entry disappeared");
return;
};
debug_assert!(*entry.get() > 0, "connection permit counter underflow");
if *entry.get() <= 1 {
entry.remove();
} else {
*entry.get_mut() -= 1;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Barrier;
#[test]
fn permit_drop_releases_and_removes_zero_entry() {
let authority = Arc::new(UserConnectionAuthority::default());
let permit = authority.try_acquire("alice", Some(1)).unwrap();
assert_eq!(authority.active("alice"), 1);
assert!(authority.try_acquire("alice", Some(1)).is_none());
drop(permit);
assert_eq!(authority.active("alice"), 0);
assert_eq!(authority.tracked_users(), 0);
}
#[test]
fn concurrent_acquire_never_exceeds_limit() {
const CONTENDERS: usize = 64;
const LIMIT: u64 = 7;
let authority = Arc::new(UserConnectionAuthority::default());
let barrier = Arc::new(Barrier::new(CONTENDERS + 1));
let (permit_tx, permit_rx) = std::sync::mpsc::channel();
let mut threads = Vec::with_capacity(CONTENDERS);
for _ in 0..CONTENDERS {
let authority = Arc::clone(&authority);
let barrier = Arc::clone(&barrier);
let permit_tx = permit_tx.clone();
threads.push(std::thread::spawn(move || {
barrier.wait();
permit_tx
.send(authority.try_acquire("alice", Some(LIMIT)))
.unwrap();
}));
}
drop(permit_tx);
barrier.wait();
for thread in threads {
thread.join().unwrap();
}
let permits = permit_rx.into_iter().flatten().collect::<Vec<_>>();
assert_eq!(permits.len() as u64, LIMIT);
assert_eq!(authority.active("alice"), LIMIT);
drop(permits);
assert_eq!(authority.active("alice"), 0);
}
}