1use std::{
2 hash::Hash,
3 sync::{Arc, Mutex, MutexGuard, PoisonError, RwLock, RwLockReadGuard, RwLockWriteGuard},
4};
5
6use dashmap::{DashMap, DashSet};
7
8type SharedLockResult<'a, T, R> = Result<R, PoisonError<MutexGuard<'a, T>>>;
9type SharedRwReadResult<'a, T, R> = Result<R, PoisonError<RwLockReadGuard<'a, T>>>;
10type SharedRwWriteResult<'a, T, R> = Result<R, PoisonError<RwLockWriteGuard<'a, T>>>;
11
12#[derive(Debug)]
14pub struct SharedLock<T>(Arc<Mutex<T>>);
15
16impl<T> Clone for SharedLock<T> {
17 fn clone(&self) -> Self {
18 SharedLock(Arc::clone(&self.0))
19 }
20}
21
22impl<T> SharedLock<T> {
23 pub fn new(v: T) -> Self {
25 SharedLock(Arc::new(Mutex::new(v)))
26 }
27
28 pub fn with<F, R>(&self, f: F) -> SharedLockResult<'_, T, R>
33 where
34 F: FnOnce(&mut T) -> R,
35 {
36 let mut lock = self.0.lock()?;
37 Ok(f(&mut *lock))
38 }
39
40 pub fn get(&self) -> SharedLockResult<'_, T, T>
42 where
43 T: Clone,
44 {
45 self.with(|v| v.clone())
46 }
47
48 pub fn set(&self, value: T) -> SharedLockResult<'_, T, ()> {
50 self.with(|v| *v = value)?;
51 Ok(())
52 }
53}
54
55#[derive(Debug)]
57pub struct SharedRw<T>(Arc<RwLock<T>>);
58
59impl<T> Clone for SharedRw<T> {
60 fn clone(&self) -> Self {
61 SharedRw(Arc::clone(&self.0))
62 }
63}
64
65impl<T> SharedRw<T> {
66 pub fn new(v: T) -> Self {
68 SharedRw(Arc::new(RwLock::new(v)))
69 }
70
71 pub fn read<F, R>(&self, f: F) -> SharedRwReadResult<'_, T, R>
76 where
77 F: FnOnce(&T) -> R,
78 {
79 let guard = self.0.read()?;
80 Ok(f(&*guard))
81 }
82
83 pub fn write<F, R>(&self, f: F) -> SharedRwWriteResult<'_, T, R>
88 where
89 F: FnOnce(&mut T) -> R,
90 {
91 let mut guard = self.0.write()?;
92 Ok(f(&mut *guard))
93 }
94
95 pub fn get(&self) -> SharedRwReadResult<'_, T, T>
97 where
98 T: Clone,
99 {
100 self.read(|v| v.clone())
101 }
102
103 pub fn set(&self, value: T) -> SharedRwWriteResult<'_, T, ()> {
105 self.write(|v| *v = value)?;
106 Ok(())
107 }
108}
109
110#[derive(Debug)]
112pub struct SharedMap<K: Eq + Clone + Hash, V>(Arc<DashMap<K, V>>);
113
114impl<K: Eq + Clone + Hash, V> Clone for SharedMap<K, V> {
115 fn clone(&self) -> Self {
116 SharedMap(Arc::clone(&self.0))
117 }
118}
119
120impl<K: Eq + Hash + Clone, V> SharedMap<K, V> {
121 pub fn new() -> Self {
123 SharedMap(Arc::new(DashMap::new()))
124 }
125
126 pub fn get_cloned(&self, key: &K) -> Option<V>
133 where
134 V: Clone,
135 {
136 self.0.get(key).map(|refc| refc.value().clone())
137 }
138
139 pub fn with<F, R>(&self, key: &K, f: F) -> Option<R>
144 where
145 F: FnOnce(&V) -> R,
146 {
147 let guard = self.0.get(key)?;
148 let result = f(guard.value());
149 drop(guard);
150 Some(result)
151 }
152
153 pub fn with_mut<F, R>(&self, key: &K, f: F) -> Option<R>
158 where
159 F: FnOnce(&mut V) -> R,
160 {
161 let mut guard = self.0.get_mut(key)?;
162 let result = f(guard.value_mut());
163 Some(result)
164 }
165
166 pub fn with_mut_or_insert_with<F, D, R>(&self, key: K, default: D, f: F) -> R
168 where
169 F: FnOnce(&mut V) -> R,
170 D: FnOnce() -> V,
171 {
172 let mut entry = self.0.entry(key).or_insert_with(default);
173 f(entry.value_mut())
174 }
175
176 pub fn with_mut_or_default<F, R>(&self, key: K, f: F) -> R
178 where
179 V: Default,
180 F: FnOnce(&mut V) -> R,
181 {
182 self.with_mut_or_insert_with(key, V::default, f)
183 }
184
185 pub fn for_each<F, Ret>(&self, mut f: F)
190 where
191 F: FnMut(K, &V) -> Ret,
192 {
193 for entry in self.0.iter() {
194 f(entry.key().clone(), entry.value());
195 }
196 }
197
198 pub fn for_each_mut<F, Ret>(&self, mut f: F)
203 where
204 F: FnMut(K, &mut V) -> Ret,
205 {
206 for mut entry in self.0.iter_mut() {
207 f(entry.key().clone(), entry.value_mut());
208 }
209 }
210
211 pub fn try_for_each_mut<F, E>(&self, mut f: F) -> Result<(), E>
216 where
217 F: FnMut(K, &mut V) -> Result<(), E>,
218 {
219 for mut entry in self.0.iter_mut() {
220 f(entry.key().clone(), entry.value_mut())?;
221 }
222 Ok(())
223 }
224
225 pub fn try_for_each<F, E>(&self, mut f: F) -> Result<(), E>
230 where
231 F: FnMut(K, &V) -> Result<(), E>,
232 {
233 for entry in self.0.iter() {
234 f(entry.key().clone(), entry.value())?;
235 }
236 Ok(())
237 }
238
239 pub fn insert(&self, key: K, value: V) -> Option<V> {
241 self.0.insert(key, value)
242 }
243
244 pub fn remove(&self, key: &K) -> Option<(K, V)> {
246 self.0.remove(key)
247 }
248
249 pub fn contains_key(&self, key: &K) -> bool {
251 self.0.contains_key(key)
252 }
253
254 pub fn retain<F>(&self, f: F)
259 where
260 F: FnMut(&K, &mut V) -> bool,
261 {
262 self.0.retain(f);
263 }
264
265 pub fn keys(&self) -> Vec<K>
267 where
268 K: Clone,
269 {
270 self.0.iter().map(|e| e.key().clone()).collect()
271 }
272
273 pub fn len(&self) -> usize {
275 self.0.len()
276 }
277
278 pub fn is_empty(&self) -> bool {
280 self.0.is_empty()
281 }
282
283 pub fn clear(&self) {
285 self.0.clear()
286 }
287}
288
289impl<K: Eq + Hash + Clone, V> Default for SharedMap<K, V> {
290 fn default() -> Self {
291 Self::new()
292 }
293}
294
295#[derive(Debug)]
297pub struct SharedSet<K: Eq + Clone + Hash>(Arc<DashSet<K>>);
298
299impl<K: Eq + Clone + Hash> Clone for SharedSet<K> {
300 fn clone(&self) -> Self {
301 SharedSet(Arc::clone(&self.0))
302 }
303}
304
305impl<K: Eq + Hash + Clone> SharedSet<K> {
306 pub fn new() -> Self {
308 SharedSet(Arc::new(DashSet::new()))
309 }
310
311 pub fn get_cloned(&self, key: &K) -> Option<K> {
318 self.0.get(key).map(|item| (*item).clone())
319 }
320
321 pub fn with<F, R>(&self, key: &K, f: F) -> Option<R>
326 where
327 F: FnOnce(&K) -> R,
328 {
329 let guard = self.0.get(key)?;
330 let result = f(guard.key());
331 drop(guard);
332 Some(result)
333 }
334
335 pub fn insert(&self, key: K) -> bool {
337 self.0.insert(key)
338 }
339
340 pub fn remove(&self, key: &K) -> Option<K> {
342 self.0.remove(key)
343 }
344
345 pub fn remove_if<F>(&self, key: &K, f: F) -> Option<K>
347 where
348 F: FnOnce(&K) -> bool,
349 {
350 self.0.remove_if(key, f)
351 }
352
353 pub fn contains(&self, key: &K) -> bool {
355 self.0.contains(key)
356 }
357
358 pub fn for_each<F, Ret>(&self, mut f: F)
363 where
364 F: FnMut(&K) -> Ret,
365 {
366 for entry in self.0.iter() {
367 f(entry.key());
368 }
369 }
370
371 pub fn try_for_each<F, E>(&self, mut f: F) -> Result<(), E>
376 where
377 F: FnMut(&K) -> Result<(), E>,
378 {
379 for entry in self.0.iter() {
380 f(entry.key())?;
381 }
382 Ok(())
383 }
384
385 pub fn retain<F>(&self, f: F)
390 where
391 F: FnMut(&K) -> bool,
392 {
393 self.0.retain(f);
394 }
395
396 pub fn items(&self) -> Vec<K> {
398 self.0.iter().map(|entry| entry.key().clone()).collect()
399 }
400
401 pub fn len(&self) -> usize {
403 self.0.len()
404 }
405
406 pub fn is_empty(&self) -> bool {
408 self.0.is_empty()
409 }
410
411 pub fn clear(&self) {
413 self.0.clear()
414 }
415}
416
417impl<K: Eq + Hash + Clone> Default for SharedSet<K> {
418 fn default() -> Self {
419 Self::new()
420 }
421}
422
423#[cfg(test)]
424mod tests {
425 use super::*;
426
427 #[test]
428 fn shared_basic_usage() {
429 let v = SharedLock::new(10);
430
431 let _ = v.with(|x| *x += 5);
432 assert_eq!(v.get().unwrap(), 15);
433
434 let _ = v.set(42);
435 assert_eq!(v.get().unwrap(), 42);
436 }
437
438 #[test]
439 fn shared_rw_usage() {
440 let v = SharedRw::new(100);
441
442 let a = v.read(|x| *x).unwrap();
443 let b = v.read(|x| *x).unwrap();
444 assert_eq!(a, b);
445
446 let _ = v.write(|x| *x += 1);
447 assert_eq!(v.get().unwrap(), 101);
448 }
449
450 #[test]
451 fn shared_map_usage() {
452 let map = SharedMap::new();
453
454 map.insert("a", 1);
455 map.insert("b", 2);
456
457 let val = map.with(&"a", |v| *v).unwrap();
458 assert_eq!(val, 1);
459
460 map.with_mut(&"a", |v| *v += 10);
461 assert_eq!(map.with(&"a", |v| *v).unwrap(), 11);
462
463 map.with_mut_or_default("c", |v| *v += 3);
464 assert_eq!(map.get_cloned(&"c"), Some(3));
465
466 let mut sum = 0;
467 map.for_each(|_, v| sum += v);
468 assert_eq!(sum, 16);
469
470 map.remove(&"a");
471 assert!(!map.contains_key(&"a"));
472 }
473
474 #[test]
475 fn shared_set_usage() {
476 let set = SharedSet::new();
477
478 assert!(set.insert("a"));
479 assert!(!set.insert("a"));
480 assert!(set.contains(&"a"));
481 assert_eq!(set.get_cloned(&"a"), Some("a"));
482 assert_eq!(set.with(&"a", |item| item.len()), Some(1));
483
484 assert_eq!(set.remove_if(&"a", |item| item.starts_with("z")), None);
485 assert!(set.contains(&"a"));
486 assert_eq!(set.remove_if(&"a", |item| item.starts_with("a")), Some("a"));
487 assert!(!set.contains(&"a"));
488
489 set.insert("a");
490 set.insert("b");
491 let mut items = set.items();
492 items.sort();
493 assert_eq!(items, vec!["a", "b"]);
494
495 let mut iterated = Vec::new();
496 set.for_each(|item| iterated.push(item.to_string()));
497 iterated.sort();
498 assert_eq!(iterated, vec!["a".to_string(), "b".to_string()]);
499
500 let result: Result<(), ()> =
501 set.try_for_each(|item| if *item == "b" { Err(()) } else { Ok(()) });
502 assert!(result.is_err());
503
504 set.retain(|item| *item != "a");
505 assert_eq!(set.items(), vec!["b"]);
506
507 set.clear();
508 assert!(set.is_empty());
509 }
510}