Skip to main content

o_sfu/runtime/websocket_server/
admission.rs

1use std::{
2    collections::HashMap,
3    sync::{Arc, Mutex, MutexGuard, PoisonError},
4};
5
6use tokio::sync::{OwnedSemaphorePermit, Semaphore};
7
8#[derive(Debug, Clone)]
9pub(crate) struct PreAuthWebSocketAdmission {
10    global: Arc<Semaphore>,
11    per_origin_capacity: usize,
12    origins: Arc<Mutex<HashMap<Arc<str>, OriginAdmission>>>,
13}
14
15#[derive(Debug, Clone)]
16struct OriginAdmission {
17    semaphore: Arc<Semaphore>,
18}
19
20/// holds global and origin pre-auth capacity until authentication releases it
21/// or the upgraded socket is dropped
22///
23/// dropping the permit removes idle origin buckets after the last origin permit
24/// returns
25#[derive(Debug)]
26pub(super) struct PreAuthWebSocketPermit {
27    _global_permit: OwnedSemaphorePermit,
28    origin_permit: Option<OwnedSemaphorePermit>,
29    origin: Arc<str>,
30    origins: Arc<Mutex<HashMap<Arc<str>, OriginAdmission>>>,
31    per_origin_capacity: usize,
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub(super) enum PreAuthWebSocketAdmissionRejection {
36    Global,
37    Origin,
38}
39
40impl PreAuthWebSocketAdmission {
41    #[must_use]
42    pub(crate) fn new(global_capacity: usize, per_origin_capacity: usize) -> Self {
43        debug_assert!(global_capacity > 0);
44        debug_assert!(per_origin_capacity > 0);
45        Self {
46            global: Arc::new(Semaphore::new(global_capacity)),
47            per_origin_capacity,
48            origins: Arc::new(Mutex::new(HashMap::new())),
49        }
50    }
51
52    pub(super) fn try_acquire(
53        &self,
54        origin: Arc<str>,
55    ) -> Result<PreAuthWebSocketPermit, PreAuthWebSocketAdmissionRejection> {
56        let global_permit = Arc::clone(&self.global)
57            .try_acquire_owned()
58            .map_err(|_error| PreAuthWebSocketAdmissionRejection::Global)?;
59        let mut origins = lock_origins(&self.origins);
60        let origin_admission =
61            origins
62                .entry(Arc::clone(&origin))
63                .or_insert_with(|| OriginAdmission {
64                    semaphore: Arc::new(Semaphore::new(self.per_origin_capacity)),
65                });
66        let origin_permit = Arc::clone(&origin_admission.semaphore)
67            .try_acquire_owned()
68            .map_err(|_error| PreAuthWebSocketAdmissionRejection::Origin)?;
69        drop(origins);
70        Ok(PreAuthWebSocketPermit {
71            _global_permit: global_permit,
72            origin_permit: Some(origin_permit),
73            origin,
74            origins: Arc::clone(&self.origins),
75            per_origin_capacity: self.per_origin_capacity,
76        })
77    }
78}
79
80impl Drop for PreAuthWebSocketPermit {
81    fn drop(&mut self) {
82        drop(self.origin_permit.take());
83        let mut origins = lock_origins(&self.origins);
84        let should_remove = origins.get(&self.origin).is_some_and(|admission| {
85            admission.semaphore.available_permits() == self.per_origin_capacity
86        });
87        if should_remove {
88            origins.remove(&self.origin);
89        }
90    }
91}
92
93fn lock_origins(
94    origins: &Mutex<HashMap<Arc<str>, OriginAdmission>>,
95) -> MutexGuard<'_, HashMap<Arc<str>, OriginAdmission>> {
96    origins.lock().unwrap_or_else(PoisonError::into_inner)
97}