Skip to main content

o_sfu/runtime/websocket_server/
admission.rs

1use std::{
2    collections::HashMap,
3    mem,
4    net::{IpAddr, Ipv6Addr},
5    sync::{Arc, LazyLock, Mutex, MutexGuard, PoisonError},
6    time::{Duration, Instant},
7};
8
9use tokio::sync::{OwnedSemaphorePermit, Semaphore};
10
11#[derive(Debug, Clone)]
12pub(crate) struct PreAuthWebSocketAdmission {
13    global: Arc<Semaphore>,
14    per_origin_capacity: usize,
15    origins: Arc<Mutex<HashMap<Option<IpAddr>, Arc<Semaphore>>>>,
16}
17
18/// holds global and origin pre-auth capacity until authentication releases it
19/// or the upgraded socket is dropped
20///
21/// dropping the permit removes idle origin buckets after the last origin permit
22/// returns
23#[derive(Debug)]
24pub(super) struct PreAuthWebSocketPermit {
25    _global_permit: OwnedSemaphorePermit,
26    origin_permit: Option<OwnedSemaphorePermit>,
27    origin: Option<IpAddr>,
28    origins: Arc<Mutex<HashMap<Option<IpAddr>, Arc<Semaphore>>>>,
29    per_origin_capacity: usize,
30}
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub(super) enum PreAuthWebSocketAdmissionRejection {
34    Global,
35    Origin,
36}
37
38impl PreAuthWebSocketAdmission {
39    #[must_use]
40    pub(crate) fn new(global_capacity: usize, per_origin_capacity: usize) -> Self {
41        debug_assert!(global_capacity > 0);
42        debug_assert!(per_origin_capacity > 0);
43        Self {
44            global: Arc::new(Semaphore::new(global_capacity)),
45            per_origin_capacity,
46            origins: Arc::new(Mutex::new(HashMap::new())),
47        }
48    }
49
50    pub(super) fn try_acquire(
51        &self,
52        origin: Option<IpAddr>,
53    ) -> Result<PreAuthWebSocketPermit, PreAuthWebSocketAdmissionRejection> {
54        let origin = origin.map(origin_bucket);
55        let global_permit = Arc::clone(&self.global)
56            .try_acquire_owned()
57            .map_err(|_error| PreAuthWebSocketAdmissionRejection::Global)?;
58        let mut origins = lock_origins(&self.origins);
59        let origin_semaphore = origins
60            .entry(origin)
61            .or_insert_with(|| Arc::new(Semaphore::new(self.per_origin_capacity)));
62        let origin_permit = Arc::clone(origin_semaphore)
63            .try_acquire_owned()
64            .map_err(|_error| PreAuthWebSocketAdmissionRejection::Origin)?;
65        drop(origins);
66        Ok(PreAuthWebSocketPermit {
67            _global_permit: global_permit,
68            origin_permit: Some(origin_permit),
69            origin,
70            origins: Arc::clone(&self.origins),
71            per_origin_capacity: self.per_origin_capacity,
72        })
73    }
74}
75
76impl Drop for PreAuthWebSocketPermit {
77    fn drop(&mut self) {
78        drop(self.origin_permit.take());
79        let mut origins = lock_origins(&self.origins);
80        let should_remove = origins
81            .get(&self.origin)
82            .is_some_and(|semaphore| semaphore.available_permits() == self.per_origin_capacity);
83        if should_remove {
84            origins.remove(&self.origin);
85        }
86    }
87}
88
89fn lock_origins(
90    origins: &Mutex<HashMap<Option<IpAddr>, Arc<Semaphore>>>,
91) -> MutexGuard<'_, HashMap<Option<IpAddr>, Arc<Semaphore>>> {
92    origins.lock().unwrap_or_else(PoisonError::into_inner)
93}
94
95/// One IPv6 /64 shares a bucket so rotating interface addresses cannot evade
96/// admission. IPv4-mapped IPv6 shares the IPv4 bucket (RFC 4291 section 2.5.5.2).
97fn origin_bucket(address: IpAddr) -> IpAddr {
98    match address.to_canonical() {
99        address @ IpAddr::V4(_) => address,
100        IpAddr::V6(address) => {
101            IpAddr::V6(Ipv6Addr::from_bits(address.to_bits() & (u128::MAX << 64)))
102        }
103    }
104}
105
106// All origins share one budget. An attacker cannot allocate logging state by
107// rotating addresses or make a flood expensive merely by getting rejected.
108static REJECTION_LOG_BUDGET: LazyLock<Mutex<RejectionLogBudget>> =
109    LazyLock::new(|| Mutex::new(RejectionLogBudget::new(Instant::now())));
110const REJECTION_LOG_INTERVAL: Duration = Duration::from_secs(1);
111const REJECTION_LOG_BURST: u8 = 5;
112
113struct RejectionLogBudget {
114    window_started: Instant,
115    remaining: u8,
116    suppressed: u64,
117}
118
119impl RejectionLogBudget {
120    fn new(now: Instant) -> Self {
121        Self {
122            window_started: now,
123            remaining: REJECTION_LOG_BURST,
124            suppressed: 0,
125        }
126    }
127
128    /// Reports suppressed rejections on the next admitted log after refill.
129    fn admit(&mut self, now: Instant) -> Option<u64> {
130        if now.saturating_duration_since(self.window_started) >= REJECTION_LOG_INTERVAL {
131            self.window_started = now;
132            self.remaining = REJECTION_LOG_BURST;
133        }
134        if self.remaining == 0 {
135            self.suppressed = self.suppressed.saturating_add(1);
136            return None;
137        }
138        self.remaining -= 1;
139        Some(mem::take(&mut self.suppressed))
140    }
141}
142
143/// Reserves one rejection log and returns the count suppressed since the last log.
144///
145/// The budget covers both rejected upgrades and authentication failures.
146pub(super) fn admit_rejection_log() -> Option<u64> {
147    REJECTION_LOG_BUDGET
148        .lock()
149        .unwrap_or_else(PoisonError::into_inner)
150        .admit(Instant::now())
151}
152
153#[cfg(test)]
154#[path = "TESTS/admission.rs"]
155mod tests;