o_sfu/runtime/websocket_server/
admission.rs1use 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#[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
95fn 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
106static 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 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
143pub(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;