mirror of
https://github.com/telemt/telemt.git
synced 2026-10-07 18:05:57 +03:00
Races in admission + accounting + publication,+ PID fixed
This commit is contained in:
@@ -79,7 +79,11 @@ where
|
||||
|
||||
let route_snapshot = deps.route_runtime.snapshot();
|
||||
let session_id = deps.rng.u64();
|
||||
let user_session = deps.shared.register_user_session(&user, session_id);
|
||||
let Some(user_session) = deps.shared.register_user_session(&user, session_id) else {
|
||||
user_reservation.release_deferred();
|
||||
warn!(user = %user, "Disabled user rejected during final admission");
|
||||
return Err(ProxyError::UserDisabled { user });
|
||||
};
|
||||
let session_cancel = user_session.token();
|
||||
let selected_me_pool = if deps.config.general.use_middle_proxy
|
||||
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
|
||||
@@ -246,6 +250,18 @@ impl UserConnectionReservation {
|
||||
}
|
||||
self.stats.decrement_user_curr_connects(&self.user);
|
||||
}
|
||||
|
||||
/// 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(self.user.clone(), self.ip);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for UserConnectionReservation {
|
||||
|
||||
+40
-61
@@ -16,10 +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,
|
||||
refund_reserved_quota_bytes,
|
||||
};
|
||||
use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, 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.
|
||||
@@ -213,7 +210,7 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
}
|
||||
|
||||
let mut remaining_before = None;
|
||||
let mut reserved_read_bytes = 0u64;
|
||||
let mut quota_reservation = None;
|
||||
let mut read_limit = buf.remaining();
|
||||
if let Some(limit) = this.quota_limit {
|
||||
let used_before = this.user_stats.quota_used();
|
||||
@@ -231,11 +228,11 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
|
||||
let desired = read_limit as u64;
|
||||
let mut reserve_rounds = 0usize;
|
||||
while reserved_read_bytes == 0 {
|
||||
while quota_reservation.is_none() {
|
||||
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
|
||||
match this.user_stats.quota_try_reserve(desired, limit) {
|
||||
Ok(_) => {
|
||||
reserved_read_bytes = desired;
|
||||
match this.user_stats.quota_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
quota_reservation = Some(reservation);
|
||||
break;
|
||||
}
|
||||
Err(crate::stats::QuotaReserveError::LimitExceeded) => {
|
||||
@@ -248,7 +245,7 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
}
|
||||
}
|
||||
|
||||
if reserved_read_bytes == 0 {
|
||||
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();
|
||||
@@ -287,9 +284,9 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
|
||||
match read_result {
|
||||
Poll::Ready(Ok(n)) => {
|
||||
if reserved_read_bytes > n as u64 {
|
||||
let refund_bytes = reserved_read_bytes - n as u64;
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
|
||||
reservation.settle(n as u64);
|
||||
this.stats.add_quota_refund_bytes_total(refund_bytes);
|
||||
}
|
||||
if n > 0 {
|
||||
@@ -333,16 +330,16 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
Poll::Pending => {
|
||||
if reserved_read_bytes > 0 {
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes);
|
||||
this.stats.add_quota_refund_bytes_total(reserved_read_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
this.stats
|
||||
.add_quota_refund_bytes_total(reservation.reserved_bytes());
|
||||
}
|
||||
Poll::Pending
|
||||
}
|
||||
Poll::Ready(Err(err)) => {
|
||||
if reserved_read_bytes > 0 {
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes);
|
||||
this.stats.add_quota_refund_bytes_total(reserved_read_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
this.stats
|
||||
.add_quota_refund_bytes_total(reservation.reserved_bytes());
|
||||
}
|
||||
Poll::Ready(Err(err))
|
||||
}
|
||||
@@ -361,14 +358,15 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
|
||||
let mut shaper_reserved_bytes = 0u64;
|
||||
let mut shaper_reservation = None;
|
||||
let mut write_buf = buf;
|
||||
if let Some(lease) = this.traffic_lease.as_ref() {
|
||||
if !buf.is_empty() {
|
||||
loop {
|
||||
let consume = lease.try_consume(RateDirection::Down, buf.len() as u64);
|
||||
let reservation = lease.try_reserve(RateDirection::Down, buf.len() as u64);
|
||||
let consume = reservation.result();
|
||||
if consume.granted > 0 {
|
||||
shaper_reserved_bytes = consume.granted;
|
||||
shaper_reservation = Some(reservation);
|
||||
if consume.granted < buf.len() as u64 {
|
||||
write_buf = &buf[..consume.granted as usize];
|
||||
}
|
||||
@@ -398,17 +396,14 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
}
|
||||
|
||||
let mut remaining_before = None;
|
||||
let mut reserved_bytes = 0u64;
|
||||
let mut quota_reservation = None;
|
||||
if let Some(limit) = this.quota_limit {
|
||||
if !write_buf.is_empty() {
|
||||
let mut reserve_rounds = 0usize;
|
||||
while reserved_bytes == 0 {
|
||||
while quota_reservation.is_none() {
|
||||
let used_before = this.user_stats.quota_used();
|
||||
let remaining = limit.saturating_sub(used_before);
|
||||
if remaining == 0 {
|
||||
if let Some(lease) = this.traffic_lease.as_ref() {
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes);
|
||||
}
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
@@ -417,9 +412,9 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
let desired = remaining.min(write_buf.len() as u64);
|
||||
let mut saw_contention = false;
|
||||
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
|
||||
match this.user_stats.quota_try_reserve(desired, limit) {
|
||||
Ok(_) => {
|
||||
reserved_bytes = desired;
|
||||
match this.user_stats.quota_reserve(desired, limit) {
|
||||
Ok(reservation) => {
|
||||
quota_reservation = Some(reservation);
|
||||
write_buf = &write_buf[..desired as usize];
|
||||
break;
|
||||
}
|
||||
@@ -433,14 +428,13 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
}
|
||||
}
|
||||
|
||||
if reserved_bytes == 0 {
|
||||
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 let Some(lease) = this.traffic_lease.as_ref() {
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes);
|
||||
}
|
||||
let _ = this.arm_quota_wait(cx);
|
||||
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();
|
||||
@@ -451,9 +445,6 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
let used_before = this.user_stats.quota_used();
|
||||
let remaining = limit.saturating_sub(used_before);
|
||||
if remaining == 0 {
|
||||
if let Some(lease) = this.traffic_lease.as_ref() {
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes);
|
||||
}
|
||||
this.quota_exceeded.store(true, Ordering::Release);
|
||||
return Poll::Ready(Err(quota_io_error()));
|
||||
}
|
||||
@@ -463,15 +454,13 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
|
||||
match Pin::new(&mut this.inner).poll_write(cx, write_buf) {
|
||||
Poll::Ready(Ok(n)) => {
|
||||
if reserved_bytes > n as u64 {
|
||||
let refund_bytes = reserved_bytes - n as u64;
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
|
||||
reservation.settle(n as u64);
|
||||
this.stats.add_quota_refund_bytes_total(refund_bytes);
|
||||
}
|
||||
if shaper_reserved_bytes > n as u64
|
||||
&& let Some(lease) = this.traffic_lease.as_ref()
|
||||
{
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes - n as u64);
|
||||
if let Some(reservation) = shaper_reservation.take() {
|
||||
reservation.settle_written(n as u64);
|
||||
}
|
||||
if n > 0 {
|
||||
if let Some(lease) = this.traffic_lease.as_ref() {
|
||||
@@ -513,26 +502,16 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
|
||||
Poll::Ready(Ok(n))
|
||||
}
|
||||
Poll::Ready(Err(err)) => {
|
||||
if reserved_bytes > 0 {
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes);
|
||||
this.stats.add_quota_refund_bytes_total(reserved_bytes);
|
||||
}
|
||||
if shaper_reserved_bytes > 0
|
||||
&& let Some(lease) = this.traffic_lease.as_ref()
|
||||
{
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
this.stats
|
||||
.add_quota_refund_bytes_total(reservation.reserved_bytes());
|
||||
}
|
||||
Poll::Ready(Err(err))
|
||||
}
|
||||
Poll::Pending => {
|
||||
if reserved_bytes > 0 {
|
||||
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes);
|
||||
this.stats.add_quota_refund_bytes_total(reserved_bytes);
|
||||
}
|
||||
if shaper_reserved_bytes > 0
|
||||
&& let Some(lease) = this.traffic_lease.as_ref()
|
||||
{
|
||||
lease.refund(RateDirection::Down, shaper_reserved_bytes);
|
||||
if let Some(reservation) = quota_reservation.take() {
|
||||
this.stats
|
||||
.add_quota_refund_bytes_total(reservation.reserved_bytes());
|
||||
}
|
||||
Poll::Pending
|
||||
}
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
use crate::stats::UserStats;
|
||||
use std::io;
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -46,10 +45,3 @@ pub(in crate::proxy::relay) fn should_immediate_quota_check(
|
||||
) -> bool {
|
||||
remaining_before <= QUOTA_NEAR_LIMIT_BYTES || charge_bytes >= QUOTA_LARGE_CHARGE_BYTES
|
||||
}
|
||||
|
||||
pub(super) fn refund_reserved_quota_bytes(user_stats: &UserStats, reserved_bytes: u64) {
|
||||
if reserved_bytes == 0 {
|
||||
return;
|
||||
}
|
||||
user_stats.refund_quota(reserved_bytes);
|
||||
}
|
||||
|
||||
+150
-43
@@ -6,6 +6,7 @@ use std::sync::{Arc, Mutex};
|
||||
use std::time::Instant;
|
||||
|
||||
use dashmap::DashMap;
|
||||
use parking_lot::Mutex as ParkingMutex;
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
@@ -75,13 +76,18 @@ pub(crate) struct MiddleRelaySharedState {
|
||||
pub(crate) relay_idle_mark_seq: AtomicU64,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct UserAdmissionState {
|
||||
disabled_users: HashSet<String>,
|
||||
sessions_by_user: HashMap<String, HashMap<u64, CancellationToken>>,
|
||||
}
|
||||
|
||||
pub(crate) struct ProxySharedState {
|
||||
pub(crate) handshake: HandshakeSharedState,
|
||||
pub(crate) middle_relay: MiddleRelaySharedState,
|
||||
pub(crate) traffic_limiter: Arc<TrafficLimiter>,
|
||||
pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>,
|
||||
disabled_users: DashMap<String, ()>,
|
||||
active_user_sessions: DashMap<(String, u64), CancellationToken>,
|
||||
user_admission: ParkingMutex<UserAdmissionState>,
|
||||
pub(crate) conntrack_pressure_active: AtomicBool,
|
||||
pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>,
|
||||
masking_fallback_permits: Arc<Semaphore>,
|
||||
@@ -106,7 +112,18 @@ struct UserSessionGuard {
|
||||
|
||||
impl Drop for UserSessionGuard {
|
||||
fn drop(&mut self) {
|
||||
self.shared.active_user_sessions.remove(&self.key);
|
||||
let mut admission = self.shared.user_admission.lock();
|
||||
let remove_user = admission
|
||||
.sessions_by_user
|
||||
.get_mut(&self.key.0)
|
||||
.map(|sessions| {
|
||||
sessions.remove(&self.key.1);
|
||||
sessions.is_empty()
|
||||
})
|
||||
.unwrap_or(false);
|
||||
if remove_user {
|
||||
admission.sessions_by_user.remove(&self.key.0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -150,8 +167,7 @@ impl ProxySharedState {
|
||||
},
|
||||
traffic_limiter: TrafficLimiter::new(),
|
||||
direct_buffer_budget,
|
||||
disabled_users: DashMap::new(),
|
||||
active_user_sessions: DashMap::new(),
|
||||
user_admission: ParkingMutex::new(UserAdmissionState::default()),
|
||||
conntrack_pressure_active: AtomicBool::new(false),
|
||||
conntrack_close_tx: Mutex::new(None),
|
||||
masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)),
|
||||
@@ -167,68 +183,102 @@ impl ProxySharedState {
|
||||
}
|
||||
|
||||
pub(crate) fn is_user_enabled(&self, user: &str) -> bool {
|
||||
!self.disabled_users.contains_key(user)
|
||||
!self.user_admission.lock().disabled_users.contains(user)
|
||||
}
|
||||
|
||||
pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> bool {
|
||||
if enabled {
|
||||
self.disabled_users.remove(user);
|
||||
false
|
||||
} else {
|
||||
self.disabled_users.insert(user.to_string(), ()).is_none()
|
||||
pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> (bool, usize) {
|
||||
let (newly_disabled, tokens) = {
|
||||
let mut admission = self.user_admission.lock();
|
||||
if enabled {
|
||||
admission.disabled_users.remove(user);
|
||||
(false, Vec::new())
|
||||
} else {
|
||||
let newly_disabled = admission.disabled_users.insert(user.to_string());
|
||||
let tokens = admission
|
||||
.sessions_by_user
|
||||
.get(user)
|
||||
.map(|sessions| sessions.values().cloned().collect())
|
||||
.unwrap_or_default();
|
||||
(newly_disabled, tokens)
|
||||
}
|
||||
};
|
||||
for token in &tokens {
|
||||
token.cancel();
|
||||
}
|
||||
(newly_disabled, tokens.len())
|
||||
}
|
||||
|
||||
pub(crate) fn apply_user_enabled_config(
|
||||
&self,
|
||||
user_enabled: &HashMap<String, bool>,
|
||||
) -> Vec<String> {
|
||||
) -> Vec<(String, usize)> {
|
||||
let desired_disabled = user_enabled
|
||||
.iter()
|
||||
.filter_map(|(user, enabled)| (!*enabled).then_some(user.clone()))
|
||||
.collect::<HashSet<_>>();
|
||||
let current_disabled = self
|
||||
.disabled_users
|
||||
.iter()
|
||||
.map(|entry| entry.key().clone())
|
||||
.collect::<HashSet<_>>();
|
||||
|
||||
for user in current_disabled.difference(&desired_disabled) {
|
||||
self.disabled_users.remove(user);
|
||||
}
|
||||
let newly_disabled = desired_disabled
|
||||
.difference(¤t_disabled)
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
for user in desired_disabled {
|
||||
self.disabled_users.insert(user, ());
|
||||
}
|
||||
newly_disabled
|
||||
let cancellations = {
|
||||
let mut admission = self.user_admission.lock();
|
||||
let newly_disabled = desired_disabled
|
||||
.difference(&admission.disabled_users)
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
admission.disabled_users = desired_disabled;
|
||||
newly_disabled
|
||||
.into_iter()
|
||||
.map(|user| {
|
||||
let tokens = admission
|
||||
.sessions_by_user
|
||||
.get(&user)
|
||||
.map(|sessions| sessions.values().cloned().collect())
|
||||
.unwrap_or_default();
|
||||
(user, tokens)
|
||||
})
|
||||
.collect::<Vec<(String, Vec<CancellationToken>)>>()
|
||||
};
|
||||
cancellations
|
||||
.into_iter()
|
||||
.map(|(user, tokens)| {
|
||||
for token in &tokens {
|
||||
token.cancel();
|
||||
}
|
||||
(user, tokens.len())
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn register_user_session(
|
||||
self: &Arc<Self>,
|
||||
user: &str,
|
||||
session_id: u64,
|
||||
) -> UserSessionRegistration {
|
||||
) -> Option<UserSessionRegistration> {
|
||||
let token = CancellationToken::new();
|
||||
let key = (user.to_string(), session_id);
|
||||
self.active_user_sessions.insert(key.clone(), token.clone());
|
||||
UserSessionRegistration {
|
||||
let mut admission = self.user_admission.lock();
|
||||
if admission.disabled_users.contains(user) {
|
||||
return None;
|
||||
}
|
||||
admission
|
||||
.sessions_by_user
|
||||
.entry(key.0.clone())
|
||||
.or_default()
|
||||
.insert(session_id, token.clone());
|
||||
Some(UserSessionRegistration {
|
||||
token,
|
||||
_guard: UserSessionGuard {
|
||||
shared: Arc::clone(self),
|
||||
key,
|
||||
},
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize {
|
||||
let tokens = self
|
||||
.active_user_sessions
|
||||
.iter()
|
||||
.filter_map(|entry| (entry.key().0 == user).then(|| entry.value().clone()))
|
||||
.collect::<Vec<_>>();
|
||||
let tokens: Vec<CancellationToken> = self
|
||||
.user_admission
|
||||
.lock()
|
||||
.sessions_by_user
|
||||
.get(user)
|
||||
.map(|sessions| sessions.values().cloned().collect())
|
||||
.unwrap_or_default();
|
||||
for token in &tokens {
|
||||
token.cancel();
|
||||
}
|
||||
@@ -311,7 +361,7 @@ mod tests {
|
||||
|
||||
let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled);
|
||||
newly_disabled.sort();
|
||||
assert_eq!(newly_disabled, vec!["alice".to_string()]);
|
||||
assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]);
|
||||
assert!(!shared.is_user_enabled("alice"));
|
||||
assert!(shared.is_user_enabled("bob"));
|
||||
|
||||
@@ -325,9 +375,9 @@ mod tests {
|
||||
#[test]
|
||||
fn cancel_user_sessions_cancels_only_registered_matching_user() {
|
||||
let shared = ProxySharedState::new();
|
||||
let alice_1 = shared.register_user_session("alice", 1);
|
||||
let alice_2 = shared.register_user_session("alice", 2);
|
||||
let bob = shared.register_user_session("bob", 1);
|
||||
let alice_1 = shared.register_user_session("alice", 1).unwrap();
|
||||
let alice_2 = shared.register_user_session("alice", 2).unwrap();
|
||||
let bob = shared.register_user_session("bob", 1).unwrap();
|
||||
let alice_1_token = alice_1.token();
|
||||
let alice_2_token = alice_2.token();
|
||||
let bob_token = bob.token();
|
||||
@@ -339,4 +389,61 @@ mod tests {
|
||||
assert!(alice_2_token.is_cancelled());
|
||||
assert!(!bob_token.is_cancelled());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_user_cannot_register_after_the_cancellation_snapshot() {
|
||||
let shared = ProxySharedState::new();
|
||||
|
||||
assert_eq!(shared.set_user_enabled("alice", false), (true, 0));
|
||||
assert_eq!(shared.cancel_user_sessions("alice"), 0);
|
||||
|
||||
let late = shared.register_user_session("alice", 1);
|
||||
assert!(
|
||||
late.is_none(),
|
||||
"a session registered after disable returned must be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabling_user_cancels_existing_sessions_before_return() {
|
||||
let shared = ProxySharedState::new();
|
||||
let registration = shared.register_user_session("alice", 1).unwrap();
|
||||
let token = registration.token();
|
||||
|
||||
assert_eq!(shared.set_user_enabled("alice", false), (true, 1));
|
||||
assert!(token.is_cancelled());
|
||||
assert!(shared.register_user_session("alice", 2).is_none());
|
||||
|
||||
assert_eq!(shared.set_user_enabled("alice", true), (false, 0));
|
||||
assert!(shared.register_user_session("alice", 3).is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_disable_and_registration_never_leave_a_live_session() {
|
||||
const ITERATIONS: usize = 10_000;
|
||||
|
||||
let shared = ProxySharedState::new();
|
||||
let barrier = Arc::new(std::sync::Barrier::new(2));
|
||||
let register_shared = Arc::clone(&shared);
|
||||
let register_barrier = Arc::clone(&barrier);
|
||||
let register = std::thread::spawn(move || {
|
||||
let mut registrations = Vec::with_capacity(ITERATIONS);
|
||||
for session_id in 0..ITERATIONS as u64 {
|
||||
let user = format!("user-{session_id}");
|
||||
register_barrier.wait();
|
||||
registrations.push(register_shared.register_user_session(&user, session_id));
|
||||
}
|
||||
registrations
|
||||
});
|
||||
|
||||
for session_id in 0..ITERATIONS as u64 {
|
||||
let user = format!("user-{session_id}");
|
||||
barrier.wait();
|
||||
shared.set_user_enabled(&user, false);
|
||||
}
|
||||
|
||||
for registration in register.join().unwrap().into_iter().flatten() {
|
||||
assert!(registration.token().is_cancelled());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use arc_swap::ArcSwap;
|
||||
use dashmap::DashMap;
|
||||
use ipnetwork::IpNetwork;
|
||||
use parking_lot::Mutex as ParkingMutex;
|
||||
|
||||
use crate::config::RateLimitBps;
|
||||
|
||||
@@ -32,6 +33,9 @@ const REGISTRY_SHARDS: usize = 64;
|
||||
const FAIR_EPOCH_MS: u64 = 20;
|
||||
const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024;
|
||||
const CLEANUP_INTERVAL_SECS: u64 = 60;
|
||||
const PACKED_USAGE_BITS: u32 = 28;
|
||||
const PACKED_USAGE_MASK: u64 = (1u64 << PACKED_USAGE_BITS) - 1;
|
||||
const PACKED_EPOCH_MAX: u64 = (1u64 << (u64::BITS - PACKED_USAGE_BITS)) - 1;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum RateDirection {
|
||||
@@ -76,12 +80,12 @@ struct ScopeMetrics {
|
||||
struct AtomicRatePair {
|
||||
up_bps: AtomicU64,
|
||||
down_bps: AtomicU64,
|
||||
revision: ParkingMutex<u64>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct DirectionBucket {
|
||||
epoch: AtomicU64,
|
||||
used: AtomicU64,
|
||||
state: AtomicU64,
|
||||
}
|
||||
|
||||
struct UserBucket {
|
||||
@@ -93,15 +97,13 @@ struct UserBucket {
|
||||
|
||||
#[derive(Default)]
|
||||
struct CidrDirectionBucket {
|
||||
epoch: AtomicU64,
|
||||
used: AtomicU64,
|
||||
active_users: AtomicU64,
|
||||
used: DirectionBucket,
|
||||
active_users: DirectionBucket,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct CidrUserDirectionState {
|
||||
epoch: AtomicU64,
|
||||
used: AtomicU64,
|
||||
used: DirectionBucket,
|
||||
}
|
||||
|
||||
struct CidrUserShare {
|
||||
@@ -139,6 +141,7 @@ enum CidrPolicyMatch<'a> {
|
||||
|
||||
#[derive(Default)]
|
||||
struct PolicySnapshot {
|
||||
revision: u64,
|
||||
user_limits: HashMap<String, RateLimitBps>,
|
||||
cidr_rules_v4: Vec<CidrRule>,
|
||||
cidr_rules_v6: Vec<CidrRule>,
|
||||
@@ -162,9 +165,25 @@ pub struct TrafficLease {
|
||||
|
||||
pub struct TrafficLimiter {
|
||||
policy: ArcSwap<PolicySnapshot>,
|
||||
policy_update: ParkingMutex<()>,
|
||||
user_buckets: ShardedRegistry<UserBucket>,
|
||||
cidr_buckets: ShardedRegistry<CidrBucket>,
|
||||
user_scope: ScopeMetrics,
|
||||
cidr_scope: ScopeMetrics,
|
||||
last_cleanup_epoch_secs: AtomicU64,
|
||||
}
|
||||
|
||||
struct DirectionDebit<'a> {
|
||||
bucket: &'a DirectionBucket,
|
||||
epoch: u64,
|
||||
refundable: u64,
|
||||
}
|
||||
|
||||
/// Refunds uncommitted shaping budget when an I/O attempt is cancelled.
|
||||
#[must_use = "traffic reservations must be settled after the I/O attempt"]
|
||||
pub(crate) struct TrafficReservation<'a> {
|
||||
result: TrafficConsumeResult,
|
||||
user: Option<DirectionDebit<'a>>,
|
||||
cidr: Option<DirectionDebit<'a>>,
|
||||
cidr_user: Option<DirectionDebit<'a>>,
|
||||
}
|
||||
|
||||
@@ -26,203 +26,277 @@ impl ScopeMetrics {
|
||||
}
|
||||
|
||||
impl AtomicRatePair {
|
||||
pub(super) fn set(&self, limits: RateLimitBps) {
|
||||
self.up_bps.store(limits.up_bps, Ordering::Relaxed);
|
||||
self.down_bps.store(limits.down_bps, Ordering::Relaxed);
|
||||
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
|
||||
let rates = Self::default();
|
||||
rates.set(revision, limits);
|
||||
rates
|
||||
}
|
||||
|
||||
pub(super) fn set(&self, revision: u64, limits: RateLimitBps) {
|
||||
let mut current_revision = self.revision.lock();
|
||||
if revision < *current_revision {
|
||||
return;
|
||||
}
|
||||
self.up_bps.store(limits.up_bps, Ordering::Release);
|
||||
self.down_bps.store(limits.down_bps, Ordering::Release);
|
||||
*current_revision = revision;
|
||||
}
|
||||
|
||||
pub(super) fn get(&self, direction: RateDirection) -> u64 {
|
||||
match direction {
|
||||
RateDirection::Up => self.up_bps.load(Ordering::Relaxed),
|
||||
RateDirection::Down => self.down_bps.load(Ordering::Relaxed),
|
||||
RateDirection::Up => self.up_bps.load(Ordering::Acquire),
|
||||
RateDirection::Down => self.down_bps.load(Ordering::Acquire),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl DirectionBucket {
|
||||
pub(super) fn sync_epoch(&self, epoch: u64) {
|
||||
let current = self.epoch.load(Ordering::Relaxed);
|
||||
if current == epoch {
|
||||
return;
|
||||
}
|
||||
if current < epoch
|
||||
&& self
|
||||
.epoch
|
||||
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
self.used.store(0, Ordering::Relaxed);
|
||||
}
|
||||
fn unpack(state: u64) -> (u64, u64) {
|
||||
(state >> PACKED_USAGE_BITS, state & PACKED_USAGE_MASK)
|
||||
}
|
||||
|
||||
pub(super) fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 {
|
||||
if requested == 0 {
|
||||
return 0;
|
||||
fn pack(epoch: u64, used: u64) -> Option<u64> {
|
||||
if epoch > PACKED_EPOCH_MAX || used > PACKED_USAGE_MASK {
|
||||
return None;
|
||||
}
|
||||
if cap_bps == 0 {
|
||||
return requested;
|
||||
Some((epoch << PACKED_USAGE_BITS) | used)
|
||||
}
|
||||
|
||||
pub(super) fn used_at(&self, epoch: u64) -> Option<u64> {
|
||||
if epoch > PACKED_EPOCH_MAX {
|
||||
return None;
|
||||
}
|
||||
let (current_epoch, used) = Self::unpack(self.state.load(Ordering::Relaxed));
|
||||
(current_epoch == epoch).then_some(used)
|
||||
}
|
||||
|
||||
let epoch = current_epoch();
|
||||
self.sync_epoch(epoch);
|
||||
let cap_epoch = bytes_per_epoch(cap_bps);
|
||||
pub(super) fn try_reserve_at(
|
||||
&self,
|
||||
epoch: u64,
|
||||
cap: u64,
|
||||
requested: u64,
|
||||
) -> Option<DirectionDebit<'_>> {
|
||||
if requested == 0 || cap == 0 || epoch > PACKED_EPOCH_MAX {
|
||||
return None;
|
||||
}
|
||||
let cap = cap.min(PACKED_USAGE_MASK);
|
||||
|
||||
let mut observed = self.state.load(Ordering::Relaxed);
|
||||
loop {
|
||||
let used = self.used.load(Ordering::Relaxed);
|
||||
if used >= cap_epoch {
|
||||
return 0;
|
||||
let (observed_epoch, observed_used) = Self::unpack(observed);
|
||||
if observed_epoch > epoch {
|
||||
return None;
|
||||
}
|
||||
let remaining = cap_epoch.saturating_sub(used);
|
||||
let used = if observed_epoch == epoch {
|
||||
observed_used
|
||||
} else {
|
||||
0
|
||||
};
|
||||
if used >= cap {
|
||||
return None;
|
||||
}
|
||||
let remaining = cap - used;
|
||||
let grant = requested.min(remaining);
|
||||
if grant == 0 {
|
||||
return 0;
|
||||
return None;
|
||||
}
|
||||
let next = used.saturating_add(grant);
|
||||
if self
|
||||
.used
|
||||
.compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
return grant;
|
||||
let next = Self::pack(epoch, used + grant)?;
|
||||
match self.state.compare_exchange_weak(
|
||||
observed,
|
||||
next,
|
||||
Ordering::Relaxed,
|
||||
Ordering::Relaxed,
|
||||
) {
|
||||
Ok(_) => {
|
||||
return Some(DirectionDebit {
|
||||
bucket: self,
|
||||
epoch,
|
||||
refundable: grant,
|
||||
});
|
||||
}
|
||||
Err(actual) => observed = actual,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn refund(&self, bytes: u64) {
|
||||
if bytes == 0 {
|
||||
fn refund_at(&self, epoch: u64, bytes: u64) {
|
||||
if bytes == 0 || epoch > PACKED_EPOCH_MAX {
|
||||
return;
|
||||
}
|
||||
decrement_atomic_saturating(&self.used, bytes);
|
||||
|
||||
let mut observed = self.state.load(Ordering::Relaxed);
|
||||
loop {
|
||||
let (observed_epoch, used) = Self::unpack(observed);
|
||||
if observed_epoch != epoch || used == 0 {
|
||||
return;
|
||||
}
|
||||
let next = Self::pack(epoch, used.saturating_sub(bytes)).unwrap_or(observed);
|
||||
match self.state.compare_exchange_weak(
|
||||
observed,
|
||||
next,
|
||||
Ordering::Relaxed,
|
||||
Ordering::Relaxed,
|
||||
) {
|
||||
Ok(_) => return,
|
||||
Err(actual) => observed = actual,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl DirectionDebit<'_> {
|
||||
fn granted(&self) -> u64 {
|
||||
self.refundable
|
||||
}
|
||||
|
||||
pub(super) fn shrink_to(&mut self, retained: u64) {
|
||||
let retained = retained.min(self.refundable);
|
||||
self.bucket
|
||||
.refund_at(self.epoch, self.refundable - retained);
|
||||
self.refundable = retained;
|
||||
}
|
||||
|
||||
pub(super) fn settle(&mut self, committed: u64) {
|
||||
self.shrink_to(committed);
|
||||
self.refundable = 0;
|
||||
}
|
||||
|
||||
pub(super) fn commit_all(&mut self) -> u64 {
|
||||
let committed = self.refundable;
|
||||
self.refundable = 0;
|
||||
committed
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for DirectionDebit<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.bucket.refund_at(self.epoch, self.refundable);
|
||||
}
|
||||
}
|
||||
|
||||
impl UserBucket {
|
||||
pub(super) fn new(limits: RateLimitBps) -> Self {
|
||||
let rates = AtomicRatePair::default();
|
||||
rates.set(limits);
|
||||
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
|
||||
Self {
|
||||
rates,
|
||||
rates: AtomicRatePair::new(revision, limits),
|
||||
up: DirectionBucket::default(),
|
||||
down: DirectionBucket::default(),
|
||||
active_leases: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn set_rates(&self, limits: RateLimitBps) {
|
||||
self.rates.set(limits);
|
||||
pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
|
||||
self.rates.set(revision, limits);
|
||||
}
|
||||
|
||||
pub(super) fn try_consume(&self, direction: RateDirection, requested: u64) -> u64 {
|
||||
pub(super) fn try_reserve(
|
||||
&self,
|
||||
direction: RateDirection,
|
||||
requested: u64,
|
||||
) -> (u64, Option<DirectionDebit<'_>>) {
|
||||
let cap_bps = self.rates.get(direction);
|
||||
match direction {
|
||||
RateDirection::Up => self.up.try_consume(cap_bps, requested),
|
||||
RateDirection::Down => self.down.try_consume(cap_bps, requested),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn refund(&self, direction: RateDirection, bytes: u64) {
|
||||
match direction {
|
||||
RateDirection::Up => self.up.refund(bytes),
|
||||
RateDirection::Down => self.down.refund(bytes),
|
||||
if cap_bps == 0 {
|
||||
return (requested, None);
|
||||
}
|
||||
let cap = bytes_per_epoch(cap_bps);
|
||||
let debit = match direction {
|
||||
RateDirection::Up => self.up.try_reserve_at(current_epoch(), cap, requested),
|
||||
RateDirection::Down => self.down.try_reserve_at(current_epoch(), cap, requested),
|
||||
};
|
||||
let granted = debit.as_ref().map(DirectionDebit::granted).unwrap_or(0);
|
||||
(granted, debit)
|
||||
}
|
||||
}
|
||||
|
||||
impl CidrDirectionBucket {
|
||||
pub(super) fn sync_epoch(&self, epoch: u64) {
|
||||
let current = self.epoch.load(Ordering::Relaxed);
|
||||
if current == epoch {
|
||||
return;
|
||||
}
|
||||
if current < epoch
|
||||
&& self
|
||||
.epoch
|
||||
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
self.used.store(0, Ordering::Relaxed);
|
||||
self.active_users.store(0, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn try_consume(
|
||||
&self,
|
||||
user_state: &CidrUserDirectionState,
|
||||
pub(super) fn try_reserve<'a>(
|
||||
&'a self,
|
||||
user_state: &'a CidrUserDirectionState,
|
||||
cap_epoch: u64,
|
||||
requested: u64,
|
||||
) -> u64 {
|
||||
) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
|
||||
if requested == 0 || cap_epoch == 0 {
|
||||
return 0;
|
||||
return (0, None, None);
|
||||
}
|
||||
|
||||
let epoch = current_epoch();
|
||||
self.sync_epoch(epoch);
|
||||
user_state.sync_epoch_and_mark_active(epoch, &self.active_users);
|
||||
let active_users = self.active_users.load(Ordering::Relaxed).max(1);
|
||||
if !user_state.ensure_active(epoch, &self.active_users) {
|
||||
return (0, None, None);
|
||||
}
|
||||
let Some(active_users) = self.active_users.used_at(epoch) else {
|
||||
return (0, None, None);
|
||||
};
|
||||
let active_users = active_users.max(1);
|
||||
let fair_share = cap_epoch.saturating_div(active_users).max(1);
|
||||
|
||||
loop {
|
||||
let total_used = self.used.load(Ordering::Relaxed);
|
||||
if total_used >= cap_epoch {
|
||||
return 0;
|
||||
}
|
||||
let total_remaining = cap_epoch.saturating_sub(total_used);
|
||||
let user_used = user_state.used.load(Ordering::Relaxed);
|
||||
let guaranteed_remaining = fair_share.saturating_sub(user_used);
|
||||
|
||||
let grant = if guaranteed_remaining > 0 {
|
||||
requested.min(guaranteed_remaining).min(total_remaining)
|
||||
} else {
|
||||
requested.min(total_remaining).min(MAX_BORROW_CHUNK_BYTES)
|
||||
let Some(user_used) = user_state.used.used_at(epoch) else {
|
||||
return (0, None, None);
|
||||
};
|
||||
|
||||
if grant == 0 {
|
||||
return 0;
|
||||
}
|
||||
|
||||
let next_total = total_used.saturating_add(grant);
|
||||
if self
|
||||
.used
|
||||
.compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
user_state.used.fetch_add(grant, Ordering::Relaxed);
|
||||
return grant;
|
||||
let guaranteed_remaining = fair_share.saturating_sub(user_used);
|
||||
let (user_cap, desired) = if guaranteed_remaining > 0 {
|
||||
(fair_share, requested.min(guaranteed_remaining))
|
||||
} else {
|
||||
(PACKED_USAGE_MASK, requested.min(MAX_BORROW_CHUNK_BYTES))
|
||||
};
|
||||
let Some(mut user_debit) = user_state.used.try_reserve_at(epoch, user_cap, desired)
|
||||
else {
|
||||
if guaranteed_remaining > 0 {
|
||||
continue;
|
||||
}
|
||||
return (0, None, None);
|
||||
};
|
||||
let user_granted = user_debit.granted();
|
||||
let Some(aggregate_debit) = self.used.try_reserve_at(epoch, cap_epoch, user_granted)
|
||||
else {
|
||||
return (0, None, None);
|
||||
};
|
||||
let granted = aggregate_debit.granted();
|
||||
if granted < user_granted {
|
||||
user_debit.shrink_to(granted);
|
||||
}
|
||||
return (granted, Some(aggregate_debit), Some(user_debit));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn refund(&self, bytes: u64) {
|
||||
if bytes == 0 {
|
||||
return;
|
||||
}
|
||||
decrement_atomic_saturating(&self.used, bytes);
|
||||
}
|
||||
}
|
||||
|
||||
impl CidrUserDirectionState {
|
||||
pub(super) fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) {
|
||||
let current = self.epoch.load(Ordering::Relaxed);
|
||||
if current == epoch {
|
||||
return;
|
||||
pub(super) fn ensure_active(&self, epoch: u64, active_users: &DirectionBucket) -> bool {
|
||||
if epoch > PACKED_EPOCH_MAX {
|
||||
return false;
|
||||
}
|
||||
if current < epoch
|
||||
&& self
|
||||
.epoch
|
||||
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
self.used.store(0, Ordering::Relaxed);
|
||||
active_users.fetch_add(1, Ordering::Relaxed);
|
||||
let mut observed = self.used.state.load(Ordering::Relaxed);
|
||||
loop {
|
||||
let (observed_epoch, _) = DirectionBucket::unpack(observed);
|
||||
if observed_epoch == epoch {
|
||||
return true;
|
||||
}
|
||||
if observed_epoch > epoch {
|
||||
return false;
|
||||
}
|
||||
let Some(mut active_debit) = active_users.try_reserve_at(epoch, PACKED_USAGE_MASK, 1)
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
let Some(next) = DirectionBucket::pack(epoch, 0) else {
|
||||
return false;
|
||||
};
|
||||
match self.used.state.compare_exchange(
|
||||
observed,
|
||||
next,
|
||||
Ordering::Relaxed,
|
||||
Ordering::Relaxed,
|
||||
) {
|
||||
Ok(_) => {
|
||||
active_debit.commit_all();
|
||||
return true;
|
||||
}
|
||||
Err(actual) => {
|
||||
drop(active_debit);
|
||||
observed = actual;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn refund(&self, bytes: u64) {
|
||||
if bytes == 0 {
|
||||
return;
|
||||
}
|
||||
decrement_atomic_saturating(&self.used, bytes);
|
||||
}
|
||||
}
|
||||
|
||||
impl CidrUserShare {
|
||||
@@ -236,11 +310,9 @@ impl CidrUserShare {
|
||||
}
|
||||
|
||||
impl CidrBucket {
|
||||
pub(super) fn new(limits: RateLimitBps) -> Self {
|
||||
let rates = AtomicRatePair::default();
|
||||
rates.set(limits);
|
||||
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
|
||||
Self {
|
||||
rates,
|
||||
rates: AtomicRatePair::new(revision, limits),
|
||||
up: CidrDirectionBucket::default(),
|
||||
down: CidrDirectionBucket::default(),
|
||||
users: ShardedRegistry::new(REGISTRY_SHARDS),
|
||||
@@ -248,8 +320,8 @@ impl CidrBucket {
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn set_rates(&self, limits: RateLimitBps) {
|
||||
self.rates.set(limits);
|
||||
pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
|
||||
self.rates.set(revision, limits);
|
||||
}
|
||||
|
||||
pub(super) fn acquire_user_share(&self, user: &str) -> Arc<CidrUserShare> {
|
||||
@@ -268,38 +340,20 @@ impl CidrBucket {
|
||||
});
|
||||
}
|
||||
|
||||
pub(super) fn try_consume_for_user(
|
||||
&self,
|
||||
pub(super) fn try_reserve_for_user<'a>(
|
||||
&'a self,
|
||||
direction: RateDirection,
|
||||
share: &CidrUserShare,
|
||||
share: &'a CidrUserShare,
|
||||
requested: u64,
|
||||
) -> u64 {
|
||||
) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
|
||||
let cap_bps = self.rates.get(direction);
|
||||
if cap_bps == 0 {
|
||||
return requested;
|
||||
return (requested, None, None);
|
||||
}
|
||||
let cap_epoch = bytes_per_epoch(cap_bps);
|
||||
match direction {
|
||||
RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested),
|
||||
RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn refund_for_user(
|
||||
&self,
|
||||
direction: RateDirection,
|
||||
share: &CidrUserShare,
|
||||
bytes: u64,
|
||||
) {
|
||||
match direction {
|
||||
RateDirection::Up => {
|
||||
self.up.refund(bytes);
|
||||
share.up.refund(bytes);
|
||||
}
|
||||
RateDirection::Down => {
|
||||
self.down.refund(bytes);
|
||||
share.down.refund(bytes);
|
||||
}
|
||||
RateDirection::Up => self.up.try_reserve(&share.up, cap_epoch, requested),
|
||||
RateDirection::Down => self.down.try_reserve(&share.down, cap_epoch, requested),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -56,7 +56,7 @@ pub(super) fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option<String>
|
||||
pub(super) fn current_epoch() -> u64 {
|
||||
let start = limiter_epoch_start();
|
||||
let elapsed_ms = start.elapsed().as_millis() as u64;
|
||||
elapsed_ms / FAIR_EPOCH_MS
|
||||
elapsed_ms / FAIR_EPOCH_MS + 1
|
||||
}
|
||||
|
||||
pub(super) fn limiter_epoch_start() -> &'static Instant {
|
||||
|
||||
@@ -1,70 +1,93 @@
|
||||
use super::*;
|
||||
|
||||
impl TrafficLease {
|
||||
pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult {
|
||||
/// Reserves shaping budget until the associated I/O result is settled.
|
||||
pub(crate) fn try_reserve(
|
||||
&self,
|
||||
direction: RateDirection,
|
||||
requested: u64,
|
||||
) -> TrafficReservation<'_> {
|
||||
if requested == 0 {
|
||||
return TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: false,
|
||||
blocked_cidr: false,
|
||||
return TrafficReservation {
|
||||
result: TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: false,
|
||||
blocked_cidr: false,
|
||||
},
|
||||
user: None,
|
||||
cidr: None,
|
||||
cidr_user: None,
|
||||
};
|
||||
}
|
||||
|
||||
let mut granted = requested;
|
||||
let mut user_debit = None;
|
||||
if let Some(user_bucket) = self.user_bucket.as_ref() {
|
||||
let user_granted = user_bucket.try_consume(direction, granted);
|
||||
let (user_granted, debit) = user_bucket.try_reserve(direction, granted);
|
||||
user_debit = debit;
|
||||
if user_granted == 0 {
|
||||
self.limiter.observe_throttle(direction, true, false);
|
||||
return TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: true,
|
||||
blocked_cidr: false,
|
||||
return TrafficReservation {
|
||||
result: TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: true,
|
||||
blocked_cidr: false,
|
||||
},
|
||||
user: user_debit,
|
||||
cidr: None,
|
||||
cidr_user: None,
|
||||
};
|
||||
}
|
||||
granted = user_granted;
|
||||
}
|
||||
|
||||
let mut cidr_debit = None;
|
||||
let mut cidr_user_debit = None;
|
||||
if let (Some(cidr_bucket), Some(cidr_user_share)) =
|
||||
(self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref())
|
||||
{
|
||||
let cidr_granted =
|
||||
cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted);
|
||||
let (cidr_granted, aggregate_debit, share_debit) =
|
||||
cidr_bucket.try_reserve_for_user(direction, cidr_user_share, granted);
|
||||
cidr_debit = aggregate_debit;
|
||||
cidr_user_debit = share_debit;
|
||||
if cidr_granted < granted
|
||||
&& let Some(user_bucket) = self.user_bucket.as_ref()
|
||||
&& let Some(debit) = user_debit.as_mut()
|
||||
{
|
||||
user_bucket.refund(direction, granted.saturating_sub(cidr_granted));
|
||||
debit.shrink_to(cidr_granted);
|
||||
}
|
||||
if cidr_granted == 0 {
|
||||
self.limiter.observe_throttle(direction, false, true);
|
||||
return TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: false,
|
||||
blocked_cidr: true,
|
||||
return TrafficReservation {
|
||||
result: TrafficConsumeResult {
|
||||
granted: 0,
|
||||
blocked_user: false,
|
||||
blocked_cidr: true,
|
||||
},
|
||||
user: user_debit,
|
||||
cidr: cidr_debit,
|
||||
cidr_user: cidr_user_debit,
|
||||
};
|
||||
}
|
||||
granted = cidr_granted;
|
||||
}
|
||||
|
||||
TrafficConsumeResult {
|
||||
granted,
|
||||
blocked_user: false,
|
||||
blocked_cidr: false,
|
||||
TrafficReservation {
|
||||
result: TrafficConsumeResult {
|
||||
granted,
|
||||
blocked_user: false,
|
||||
blocked_cidr: false,
|
||||
},
|
||||
user: user_debit,
|
||||
cidr: cidr_debit,
|
||||
cidr_user: cidr_user_debit,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn refund(&self, direction: RateDirection, bytes: u64) {
|
||||
if bytes == 0 {
|
||||
return;
|
||||
}
|
||||
|
||||
if let Some(user_bucket) = self.user_bucket.as_ref() {
|
||||
user_bucket.refund(direction, bytes);
|
||||
}
|
||||
if let (Some(cidr_bucket), Some(cidr_user_share)) =
|
||||
(self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref())
|
||||
{
|
||||
cidr_bucket.refund_for_user(direction, cidr_user_share, bytes);
|
||||
}
|
||||
pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult {
|
||||
let reservation = self.try_reserve(direction, requested);
|
||||
let result = reservation.result();
|
||||
reservation.settle_written(result.granted);
|
||||
result
|
||||
}
|
||||
|
||||
pub fn observe_wait_ms(
|
||||
@@ -82,6 +105,27 @@ impl TrafficLease {
|
||||
}
|
||||
}
|
||||
|
||||
impl TrafficReservation<'_> {
|
||||
/// Returns the shaping decision associated with this reservation.
|
||||
pub(crate) fn result(&self) -> TrafficConsumeResult {
|
||||
self.result
|
||||
}
|
||||
|
||||
/// Commits written bytes and refunds the uncommitted remainder.
|
||||
pub(crate) fn settle_written(mut self, committed: u64) {
|
||||
let committed = committed.min(self.result.granted);
|
||||
if let Some(debit) = self.user.as_mut() {
|
||||
debit.settle(committed);
|
||||
}
|
||||
if let Some(debit) = self.cidr.as_mut() {
|
||||
debit.settle(committed);
|
||||
}
|
||||
if let Some(debit) = self.cidr_user.as_mut() {
|
||||
debit.settle(committed);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TrafficLease {
|
||||
fn drop(&mut self) {
|
||||
if let Some(bucket) = self.user_bucket.as_ref() {
|
||||
|
||||
@@ -5,6 +5,7 @@ impl TrafficLimiter {
|
||||
pub fn new() -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
policy: ArcSwap::from_pointee(PolicySnapshot::default()),
|
||||
policy_update: ParkingMutex::new(()),
|
||||
user_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
|
||||
cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
|
||||
user_scope: ScopeMetrics::default(),
|
||||
@@ -18,6 +19,11 @@ impl TrafficLimiter {
|
||||
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 filtered_users = user_limits
|
||||
.into_iter()
|
||||
.filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0)
|
||||
@@ -78,6 +84,7 @@ impl TrafficLimiter {
|
||||
.store(cidr_policy_entries as u64, Ordering::Relaxed);
|
||||
|
||||
self.policy.store(Arc::new(PolicySnapshot {
|
||||
revision,
|
||||
user_limits: filtered_users,
|
||||
cidr_rules_v4,
|
||||
cidr_rules_v6,
|
||||
@@ -86,6 +93,7 @@ impl TrafficLimiter {
|
||||
cidr_rule_keys,
|
||||
}));
|
||||
|
||||
drop(policy_update);
|
||||
self.maybe_cleanup();
|
||||
}
|
||||
|
||||
@@ -99,12 +107,12 @@ impl TrafficLimiter {
|
||||
if let Some(limit) = policy.user_limits.get(user).copied() {
|
||||
let bucket = self.user_buckets.get_or_insert_with(
|
||||
user,
|
||||
|| UserBucket::new(limit),
|
||||
|| UserBucket::new(policy.revision, limit),
|
||||
|bucket| {
|
||||
bucket.active_leases.fetch_add(1, Ordering::Relaxed);
|
||||
},
|
||||
);
|
||||
bucket.set_rates(limit);
|
||||
bucket.set_rates(policy.revision, limit);
|
||||
self.user_scope
|
||||
.active_leases
|
||||
.fetch_add(1, Ordering::Relaxed);
|
||||
@@ -121,12 +129,12 @@ impl TrafficLimiter {
|
||||
};
|
||||
let bucket = self.cidr_buckets.get_or_insert_with(
|
||||
key,
|
||||
|| CidrBucket::new(limits),
|
||||
|| CidrBucket::new(policy.revision, limits),
|
||||
|bucket| {
|
||||
bucket.active_leases.fetch_add(1, Ordering::Relaxed);
|
||||
},
|
||||
);
|
||||
bucket.set_rates(limits);
|
||||
bucket.set_rates(policy.revision, limits);
|
||||
self.cidr_scope
|
||||
.active_leases
|
||||
.fetch_add(1, Ordering::Relaxed);
|
||||
|
||||
@@ -74,3 +74,185 @@ fn auto_cidr_bucket_key_canonicalizes_network_address() {
|
||||
"auto:6:2001:db8::/64"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refund_from_an_old_epoch_does_not_reduce_the_current_epoch() {
|
||||
let bucket = DirectionBucket::default();
|
||||
let old_debit = bucket.try_reserve_at(7, 100, 80).unwrap();
|
||||
let current_debit = bucket.try_reserve_at(8, 100, 60).unwrap();
|
||||
|
||||
drop(old_debit);
|
||||
|
||||
assert_eq!(bucket.used_at(8), Some(60));
|
||||
drop(current_debit);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_rollover_cannot_publish_multiple_epoch_budgets() {
|
||||
const CONTENDERS: usize = 32;
|
||||
|
||||
let bucket = Arc::new(DirectionBucket::default());
|
||||
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
|
||||
let mut threads = Vec::with_capacity(CONTENDERS);
|
||||
for _ in 0..CONTENDERS {
|
||||
let bucket = Arc::clone(&bucket);
|
||||
let barrier = Arc::clone(&barrier);
|
||||
threads.push(std::thread::spawn(move || {
|
||||
barrier.wait();
|
||||
bucket
|
||||
.try_reserve_at(9, 100, 100)
|
||||
.map(|mut debit| debit.commit_all())
|
||||
.unwrap_or(0)
|
||||
}));
|
||||
}
|
||||
|
||||
let granted = threads
|
||||
.into_iter()
|
||||
.map(|thread| thread.join().unwrap())
|
||||
.sum::<u64>();
|
||||
assert_eq!(granted, 100);
|
||||
assert_eq!(bucket.used_at(9), Some(100));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scheduler_pressure_never_exceeds_a_packed_epoch_budget() {
|
||||
const CONTENDERS: usize = 4;
|
||||
const EPOCHS: usize = 10_000;
|
||||
|
||||
let bucket = Arc::new(DirectionBucket::default());
|
||||
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
|
||||
let mut threads = Vec::with_capacity(CONTENDERS);
|
||||
for _ in 0..CONTENDERS {
|
||||
let bucket = Arc::clone(&bucket);
|
||||
let barrier = Arc::clone(&barrier);
|
||||
threads.push(std::thread::spawn(move || {
|
||||
let mut grants = Vec::with_capacity(EPOCHS);
|
||||
for epoch in 1..=EPOCHS as u64 {
|
||||
barrier.wait();
|
||||
let granted = bucket
|
||||
.try_reserve_at(epoch, 100, 100)
|
||||
.map(|mut debit| debit.commit_all())
|
||||
.unwrap_or(0);
|
||||
grants.push(granted);
|
||||
barrier.wait();
|
||||
}
|
||||
grants
|
||||
}));
|
||||
}
|
||||
|
||||
let grants = threads
|
||||
.into_iter()
|
||||
.map(|thread| thread.join().unwrap())
|
||||
.collect::<Vec<_>>();
|
||||
for epoch_index in 0..EPOCHS {
|
||||
let granted = grants
|
||||
.iter()
|
||||
.map(|thread_grants| thread_grants[epoch_index])
|
||||
.sum::<u64>();
|
||||
assert_eq!(granted, 100);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stale_policy_revision_cannot_restore_an_old_rate() {
|
||||
let bucket = UserBucket::new(2, rate(2_000, 3_000));
|
||||
|
||||
bucket.set_rates(3, rate(4_000, 5_000));
|
||||
bucket.set_rates(2, rate(6_000, 7_000));
|
||||
|
||||
assert_eq!(bucket.rates.get(RateDirection::Up), 4_000);
|
||||
assert_eq!(bucket.rates.get(RateDirection::Down), 5_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dropped_debit_refunds_only_its_packed_epoch() {
|
||||
let bucket = DirectionBucket::default();
|
||||
let debit = bucket.try_reserve_at(11, 100, 80).unwrap();
|
||||
|
||||
drop(debit);
|
||||
|
||||
assert_eq!(bucket.used_at(11), Some(0));
|
||||
assert!(
|
||||
bucket
|
||||
.try_reserve_at(PACKED_EPOCH_MAX + 1, 100, 1)
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_first_use_counts_one_active_cidr_user() {
|
||||
const CONTENDERS: usize = 32;
|
||||
|
||||
let bucket = Arc::new(CidrDirectionBucket::default());
|
||||
let user = Arc::new(CidrUserDirectionState::default());
|
||||
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
|
||||
let mut threads = Vec::with_capacity(CONTENDERS);
|
||||
for _ in 0..CONTENDERS {
|
||||
let bucket = Arc::clone(&bucket);
|
||||
let user = Arc::clone(&user);
|
||||
let barrier = Arc::clone(&barrier);
|
||||
threads.push(std::thread::spawn(move || {
|
||||
barrier.wait();
|
||||
assert!(user.ensure_active(13, &bucket.active_users));
|
||||
}));
|
||||
}
|
||||
for thread in threads {
|
||||
thread.join().unwrap();
|
||||
}
|
||||
|
||||
assert_eq!(bucket.active_users.used_at(13), Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn configured_rate_maximum_fits_the_packed_epoch_budget() {
|
||||
assert_eq!(bytes_per_epoch(100_000_000_000), 250_000_000);
|
||||
assert!(bytes_per_epoch(100_000_000_000) <= PACKED_USAGE_MASK);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dropped_traffic_reservation_refunds_user_and_cidr_debits() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
let mut user_limits = HashMap::new();
|
||||
user_limits.insert("alice".to_string(), rate(400_000, 400_000));
|
||||
let mut cidr_limits = HashMap::new();
|
||||
cidr_limits.insert(
|
||||
CidrRateLimitKey::Network("203.0.113.0/24".parse().unwrap()),
|
||||
rate(400_000, 400_000),
|
||||
);
|
||||
limiter.apply_policy(user_limits, cidr_limits);
|
||||
let lease = limiter
|
||||
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
|
||||
.unwrap();
|
||||
|
||||
let reservation = lease.try_reserve(RateDirection::Down, 800);
|
||||
assert_eq!(reservation.result().granted, 800);
|
||||
let epoch = reservation.user.as_ref().unwrap().epoch;
|
||||
drop(reservation);
|
||||
|
||||
let user_bucket = lease.user_bucket.as_ref().unwrap();
|
||||
let cidr_bucket = lease.cidr_bucket.as_ref().unwrap();
|
||||
let cidr_user = lease.cidr_user_share.as_ref().unwrap();
|
||||
assert_eq!(user_bucket.down.used_at(epoch), Some(0));
|
||||
assert_eq!(cidr_bucket.down.used.used_at(epoch), Some(0));
|
||||
assert_eq!(cidr_user.down.used.used_at(epoch), Some(0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn partial_traffic_settlement_charges_only_committed_bytes() {
|
||||
let limiter = TrafficLimiter::new();
|
||||
let mut user_limits = HashMap::new();
|
||||
user_limits.insert("alice".to_string(), rate(400_000, 400_000));
|
||||
limiter.apply_policy(user_limits, HashMap::new());
|
||||
let lease = limiter
|
||||
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
|
||||
.unwrap();
|
||||
|
||||
let reservation = lease.try_reserve(RateDirection::Down, 800);
|
||||
let epoch = reservation.user.as_ref().unwrap().epoch;
|
||||
reservation.settle_written(300);
|
||||
|
||||
assert_eq!(
|
||||
lease.user_bucket.as_ref().unwrap().down.used_at(epoch),
|
||||
Some(300)
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user