tuwunel_core/utils/
mutex_map.rs1use std::{
8 fmt::Debug,
9 hash::Hash,
10 sync::{Arc, TryLockError::WouldBlock},
11};
12
13use tokio::sync::OwnedMutexGuard as Omg;
14
15use crate::{Result, err};
16
17#[derive(Debug)]
23pub struct MutexMap<Key, Val> {
24 map: Map<Key, Val>,
25}
26
27#[derive(Debug)]
33#[clippy::has_significant_drop]
34pub struct Guard<Key, Val> {
35 map: Map<Key, Val>,
36 entry: Option<Value<Val>>,
37 val: Option<Omg<Val>>,
38}
39
40type Map<Key, Val> = Arc<MapMutex<Key, Val>>;
41type MapMutex<Key, Val> = std::sync::Mutex<HashMap<Key, Val>>;
42type HashMap<Key, Val> = std::collections::HashMap<Key, Value<Val>>;
43type Value<Val> = Arc<tokio::sync::Mutex<Val>>;
44
45impl<Key, Val> MutexMap<Key, Val>
46where
47 Key: Clone + Eq + Hash + Send,
48 Val: Default + Send,
49{
50 #[must_use]
55 pub fn new() -> Self {
56 Self {
57 map: Map::new(MapMutex::new(HashMap::new())),
58 }
59 }
60
61 #[tracing::instrument(level = "trace", skip(self))]
67 pub async fn lock<K>(&self, k: &K) -> Guard<Key, Val>
68 where
69 K: Debug + Send + ?Sized + Sync + ToOwned<Owned = Key>,
70 {
71 self.entry(k).lock().await
72 }
73
74 #[tracing::instrument(level = "trace", skip(self))]
80 pub fn try_lock<K>(&self, k: &K) -> Result<Guard<Key, Val>>
81 where
82 K: Debug + Send + ?Sized + Sync + ToOwned<Owned = Key>,
83 {
84 self.entry(k).try_lock()
85 }
86
87 #[tracing::instrument(level = "trace", skip(self))]
95 pub fn try_try_lock<K>(&self, k: &K) -> Result<Guard<Key, Val>>
96 where
97 K: Debug + Send + ?Sized + Sync + ToOwned<Owned = Key>,
98 {
99 self.try_entry(k)?.try_lock()
100 }
101
102 #[must_use]
107 pub fn contains(&self, k: &Key) -> bool { self.map.lock().expect("locked").contains_key(k) }
108
109 #[must_use]
114 pub fn is_empty(&self) -> bool { self.map.lock().expect("locked").is_empty() }
115
116 #[must_use]
121 pub fn len(&self) -> usize { self.map.lock().expect("locked").len() }
122
123 #[must_use]
129 pub fn keys(&self) -> impl ExactSizeIterator<Item = Key> + Send + use<Key, Val> {
130 self.map
131 .lock()
132 .expect("locked")
133 .keys()
134 .cloned()
135 .collect::<Vec<_>>()
136 .into_iter()
137 }
138
139 fn entry<K>(&self, k: &K) -> Guard<Key, Val>
140 where
141 K: ?Sized + ToOwned<Owned = Key>,
142 {
143 let val = self
144 .map
145 .lock()
146 .expect("locked")
147 .entry(k.to_owned())
148 .or_default()
149 .clone();
150
151 self.pending(val)
152 }
153
154 fn try_entry<K>(&self, k: &K) -> Result<Guard<Key, Val>>
155 where
156 K: ?Sized + ToOwned<Owned = Key>,
157 {
158 let val = self
159 .map
160 .try_lock()
161 .map_err(|e| match e {
162 | WouldBlock => err!("would block"),
163 | _ => panic!("{e:?}"),
164 })?
165 .entry(k.to_owned())
166 .or_default()
167 .clone();
168
169 Ok(self.pending(val))
170 }
171
172 fn pending(&self, val: Value<Val>) -> Guard<Key, Val> {
173 Guard {
174 map: Arc::clone(&self.map),
175 entry: Some(val),
176 val: None,
177 }
178 }
179}
180
181impl<Key, Val> Default for MutexMap<Key, Val>
182where
183 Key: Clone + Eq + Hash + Send,
184 Val: Default + Send,
185{
186 fn default() -> Self { Self::new() }
187}
188
189impl<Key, Val> Guard<Key, Val> {
190 async fn lock(mut self) -> Self {
191 let val = self.claim();
194
195 self.val = Some(val.lock_owned().await);
196 self
197 }
198
199 fn try_lock(mut self) -> Result<Self> {
200 self.val = self
201 .claim()
202 .try_lock_owned()
203 .map_err(|_| err!("would yield"))
204 .map(Some)?;
205
206 Ok(self)
207 }
208
209 fn claim(&self) -> Value<Val> { Arc::clone(self.entry.as_ref().expect("claimed")) }
210}
211
212impl<Key, Val> Drop for Guard<Key, Val> {
213 #[tracing::instrument(name = "unlock", level = "trace", skip_all)]
214 fn drop(&mut self) {
215 self.val.take();
216
217 let mut map = self.map.lock().expect("locked");
219
220 if self
221 .entry
222 .take()
223 .is_some_and(|val| Arc::strong_count(&val) <= 2)
224 {
225 map.retain(|_, val| Arc::strong_count(val) > 1);
226 }
227 }
228}