o_sfu/runtime/websocket_server/
admission.rs1use 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#[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}