1use std::collections::HashMap;
5use std::collections::hash_map::{DefaultHasher, RandomState};
6use std::error::Error;
7use std::fmt;
8use std::hash::{BuildHasher, Hash, Hasher};
9use std::sync::Arc;
10use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
11use std::time::Duration;
12
13use parking_lot::{ArcMutexGuard, ArcRwLockReadGuard, ArcRwLockWriteGuard, Mutex, RwLock};
14use tokio::task::JoinHandle;
15use tokio::time::Instant;
16use tracing::info;
17
18use mysten_common::sync::execution_permit::release_execution_permit;
19#[cfg(test)]
20use mysten_common::sync::execution_permit::set_execution_permit;
21use mysten_metrics::spawn_monitored_task;
22
23type OwnedMutexGuard<T> = ArcMutexGuard<parking_lot::RawMutex, T>;
24type OwnedRwLockReadGuard<T> = ArcRwLockReadGuard<parking_lot::RawRwLock, T>;
25type OwnedRwLockWriteGuard<T> = ArcRwLockWriteGuard<parking_lot::RawRwLock, T>;
26
27pub trait Lock: Send + Sync + Default {
28 type Guard;
29 type ReadGuard;
30 fn lock_owned(self: Arc<Self>) -> Self::Guard;
31 fn try_lock_owned(self: Arc<Self>) -> Option<Self::Guard>;
32 fn read_lock_owned(self: Arc<Self>) -> Self::ReadGuard;
33}
34
35impl Lock for Mutex<()> {
36 type Guard = OwnedMutexGuard<()>;
37 type ReadGuard = Self::Guard;
38
39 fn lock_owned(self: Arc<Self>) -> Self::Guard {
40 self.lock_arc()
41 }
42
43 fn try_lock_owned(self: Arc<Self>) -> Option<Self::Guard> {
44 self.try_lock_arc()
45 }
46
47 fn read_lock_owned(self: Arc<Self>) -> Self::ReadGuard {
48 self.lock_arc()
49 }
50}
51
52impl Lock for RwLock<()> {
53 type Guard = OwnedRwLockWriteGuard<()>;
54 type ReadGuard = OwnedRwLockReadGuard<()>;
55
56 fn lock_owned(self: Arc<Self>) -> Self::Guard {
57 self.write_arc()
58 }
59
60 fn try_lock_owned(self: Arc<Self>) -> Option<Self::Guard> {
61 self.try_write_arc()
62 }
63
64 fn read_lock_owned(self: Arc<Self>) -> Self::ReadGuard {
65 self.read_arc()
66 }
67}
68
69type InnerLockTable<K, L> = HashMap<K, Arc<L>>;
70pub struct LockTable<K: Hash, L: Lock> {
72 random_state: RandomState,
73 lock_table: Arc<Vec<RwLock<InnerLockTable<K, L>>>>,
74 _k: std::marker::PhantomData<K>,
75 _cleaner: JoinHandle<()>,
76 stop: Arc<AtomicBool>,
77 size: Arc<AtomicUsize>,
78}
79
80pub type MutexTable<K> = LockTable<K, Mutex<()>>;
81pub type RwLockTable<K> = LockTable<K, RwLock<()>>;
82
83#[derive(Debug)]
84pub enum TryAcquireLockError {
85 LockTableLocked,
86 LockEntryLocked,
87}
88
89impl fmt::Display for TryAcquireLockError {
90 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
91 write!(fmt, "operation would block")
92 }
93}
94
95impl Error for TryAcquireLockError {}
96pub type MutexGuard = OwnedMutexGuard<()>;
97pub type RwLockGuard = OwnedRwLockReadGuard<()>;
98
99impl<K: Hash + Eq + Send + Sync + 'static, L: Lock + 'static> LockTable<K, L> {
100 pub fn new_with_cleanup(
101 num_shards: usize,
102 cleanup_period: Duration,
103 cleanup_initial_delay: Duration,
104 cleanup_entries_threshold: usize,
105 ) -> Self {
106 let num_shards = if cfg!(msim) { 4 } else { num_shards };
107
108 let lock_table: Arc<Vec<RwLock<InnerLockTable<K, L>>>> = Arc::new(
109 (0..num_shards)
110 .map(|_| RwLock::new(HashMap::new()))
111 .collect(),
112 );
113 let cloned = lock_table.clone();
114 let stop = Arc::new(AtomicBool::new(false));
115 let stop_cloned = stop.clone();
116 let size: Arc<AtomicUsize> = Arc::new(AtomicUsize::new(0));
117 let size_cloned = size.clone();
118 Self {
119 random_state: RandomState::new(),
120 lock_table,
121 _k: std::marker::PhantomData {},
122 _cleaner: spawn_monitored_task!(async move {
123 tokio::time::sleep(cleanup_initial_delay).await;
124 let mut previous_cleanup_instant = Instant::now();
125 while !stop_cloned.load(Ordering::SeqCst) {
126 if size_cloned.load(Ordering::SeqCst) >= cleanup_entries_threshold
127 || previous_cleanup_instant.elapsed() >= cleanup_period
128 {
129 let num_removed = Self::cleanup(cloned.clone());
130 size_cloned.fetch_sub(num_removed, Ordering::SeqCst);
131 previous_cleanup_instant = Instant::now();
132 }
133 tokio::time::sleep(Duration::from_secs(1)).await;
134 }
135 info!("Stopping mutex table cleanup!");
136 }),
137 stop,
138 size,
139 }
140 }
141
142 pub fn new(num_shards: usize) -> Self {
143 Self::new_with_cleanup(
144 num_shards,
145 Duration::from_secs(10),
146 Duration::from_secs(10),
147 10_000,
148 )
149 }
150
151 pub fn size(&self) -> usize {
152 self.size.load(Ordering::SeqCst)
153 }
154
155 pub fn cleanup(lock_table: Arc<Vec<RwLock<InnerLockTable<K, L>>>>) -> usize {
156 let mut num_removed: usize = 0;
157 for shard in lock_table.iter() {
158 let map = shard.try_write();
159 if map.is_none() {
160 continue;
161 }
162 map.unwrap().retain(|_k, v| {
163 if Arc::strong_count(v) == 1 {
167 num_removed += 1;
168 false
169 } else {
170 true
171 }
172 });
173 }
174 num_removed
175 }
176
177 fn get_lock_idx(&self, key: &K) -> usize {
178 let mut hasher = if !cfg!(test) {
179 self.random_state.build_hasher()
180 } else {
181 DefaultHasher::new()
183 };
184
185 key.hash(&mut hasher);
186 let hash: usize = hasher.finish().try_into().unwrap();
188 hash % self.lock_table.len()
189 }
190
191 pub fn acquire_locks<I>(&self, object_iter: I) -> Vec<L::Guard>
192 where
193 I: Iterator<Item = K>,
194 K: Ord,
195 {
196 let mut objects: Vec<K> = object_iter.into_iter().collect();
197 objects.sort_unstable();
198 objects.dedup();
199
200 let mut guards = Vec::with_capacity(objects.len());
201 for object in objects.into_iter() {
202 guards.push(self.acquire_lock(object));
203 }
204 guards
205 }
206
207 pub fn acquire_read_locks(&self, mut objects: Vec<K>) -> Vec<L::ReadGuard>
208 where
209 K: Ord,
210 {
211 objects.sort_unstable();
212 objects.dedup();
213 let mut guards = Vec::with_capacity(objects.len());
214 for object in objects.into_iter() {
215 guards.push(self.get_lock(object).read_lock_owned());
216 }
217 guards
218 }
219
220 pub fn get_lock(&self, k: K) -> Arc<L> {
221 let lock_idx = self.get_lock_idx(&k);
222 let element = {
223 let map = self.lock_table[lock_idx].read();
224 map.get(&k).cloned()
225 };
226 if let Some(element) = element {
227 element
228 } else {
229 {
232 let mut map = self.lock_table[lock_idx].write();
233 map.entry(k)
234 .or_insert_with(|| {
235 self.size.fetch_add(1, Ordering::SeqCst);
236 Arc::new(L::default())
237 })
238 .clone()
239 }
240 }
241 }
242
243 pub fn acquire_lock(&self, k: K) -> L::Guard {
253 let lock = self.get_lock(k);
254 if let Some(guard) = lock.clone().try_lock_owned() {
255 return guard;
256 }
257
258 release_execution_permit();
259
260 #[cfg(msim)]
261 loop {
262 msim::task::yield_blocking();
263 if let Some(guard) = lock.clone().try_lock_owned() {
264 return guard;
265 }
266 }
267
268 #[cfg(not(msim))]
269 lock.lock_owned()
270 }
271
272 pub fn try_acquire_lock(&self, k: K) -> Result<L::Guard, TryAcquireLockError> {
273 let lock_idx = self.get_lock_idx(&k);
274 let element = {
275 let map = self.lock_table[lock_idx]
276 .try_read()
277 .ok_or(TryAcquireLockError::LockTableLocked)?;
278 map.get(&k).cloned()
279 };
280 if let Some(element) = element {
281 let lock = element.try_lock_owned();
282 lock.ok_or(TryAcquireLockError::LockEntryLocked)
283 } else {
284 let element = {
286 let mut map = self.lock_table[lock_idx]
287 .try_write()
288 .ok_or(TryAcquireLockError::LockTableLocked)?;
289 map.entry(k)
290 .or_insert_with(|| {
291 self.size.fetch_add(1, Ordering::SeqCst);
292 Arc::new(L::default())
293 })
294 .clone()
295 };
296 let lock = element.try_lock_owned();
297 lock.ok_or(TryAcquireLockError::LockEntryLocked)
298 }
299 }
300}
301
302impl<K: Hash, L: Lock> Drop for LockTable<K, L> {
303 fn drop(&mut self) {
304 self.stop.store(true, Ordering::SeqCst);
305 }
306}
307
308#[cfg(test)]
309struct DropFlag(Arc<AtomicBool>);
310
311#[cfg(test)]
312impl Drop for DropFlag {
313 fn drop(&mut self) {
314 self.0.store(true, Ordering::SeqCst);
315 }
316}
317
318#[tokio::test]
319async fn test_acquire_lock_keeps_execution_permit_when_uncontended() {
320 let mutex_table = MutexTable::<String>::new(1);
321 let released = Arc::new(AtomicBool::new(false));
322 let _permit = set_execution_permit(Box::new(DropFlag(released.clone())));
323
324 let _guard = mutex_table.acquire_lock("key".to_string());
325
326 assert!(
327 !released.load(Ordering::SeqCst),
328 "permit must be kept when the lock is immediately available"
329 );
330}
331
332#[cfg(not(msim))]
333#[tokio::test]
334async fn test_acquire_lock_releases_execution_permit_when_contended() {
335 let mutex_table = Arc::new(MutexTable::<String>::new(1));
336 let holder = mutex_table.acquire_lock("key".to_string());
337 let released = Arc::new(AtomicBool::new(false));
338
339 let waiter = std::thread::spawn({
340 let mutex_table = mutex_table.clone();
341 let released = released.clone();
342 move || {
343 let _permit = set_execution_permit(Box::new(DropFlag(released)));
344 drop(mutex_table.acquire_lock("key".to_string()));
345 }
346 });
347
348 while !released.load(Ordering::SeqCst) {
349 std::thread::sleep(Duration::from_millis(1));
350 }
351 drop(holder);
352 waiter.join().unwrap();
353}
354
355#[cfg(all(msim, test))]
356#[sui_macros::sim_test]
357async fn test_contended_acquire_lock_releases_execution_permit_and_yields() {
358 let mutex_table = Arc::new(MutexTable::<String>::new(1));
359 let holder_acquired = Arc::new(AtomicBool::new(false));
360 let release_holder = Arc::new(AtomicBool::new(false));
361
362 let holder = tokio::task::spawn_blocking({
363 let mutex_table = mutex_table.clone();
364 let holder_acquired = holder_acquired.clone();
365 let release_holder = release_holder.clone();
366 move || {
367 let guard = mutex_table.acquire_lock("key".to_string());
368 holder_acquired.store(true, Ordering::SeqCst);
369 while !release_holder.load(Ordering::SeqCst) {
370 msim::task::yield_blocking();
371 }
372 drop(guard);
373 }
374 });
375
376 while !holder_acquired.load(Ordering::SeqCst) {
377 tokio::time::sleep(Duration::from_millis(1)).await;
378 }
379
380 let released = Arc::new(AtomicBool::new(false));
381 let waiter = tokio::task::spawn_blocking({
382 let released = released.clone();
383 move || {
384 let _permit = set_execution_permit(Box::new(DropFlag(released)));
385 drop(mutex_table.acquire_lock("key".to_string()));
386 }
387 });
388
389 while !released.load(Ordering::SeqCst) {
390 tokio::time::sleep(Duration::from_millis(1)).await;
391 }
392 release_holder.store(true, Ordering::SeqCst);
393
394 holder.await.unwrap();
395 waiter.await.unwrap();
396}
397
398#[cfg(test)]
399#[sui_macros::sui_test]
400async fn test_mutex_table_concurrent_in_same_bucket() {
403 let mutex_table = Arc::new(MutexTable::<String>::new(1));
404 let holder_acquired = Arc::new(AtomicBool::new(false));
405 let release_holder = Arc::new(AtomicBool::new(false));
406
407 let holder = tokio::task::spawn_blocking({
408 let mutex_table = mutex_table.clone();
409 let holder_acquired = holder_acquired.clone();
410 let release_holder = release_holder.clone();
411 move || {
412 let guard = mutex_table.acquire_lock("john".to_string());
413 holder_acquired.store(true, Ordering::SeqCst);
414 while !release_holder.load(Ordering::SeqCst) {
415 #[cfg(msim)]
416 msim::task::yield_blocking();
417 #[cfg(not(msim))]
418 std::thread::yield_now();
419 }
420 drop(guard);
421 }
422 });
423
424 while !holder_acquired.load(Ordering::SeqCst) {
425 tokio::time::sleep(Duration::from_millis(1)).await;
426 }
427
428 let waiter_contended = Arc::new(AtomicBool::new(false));
429
430 let waiter = tokio::task::spawn_blocking({
431 let mutex_table = mutex_table.clone();
432 let waiter_contended = waiter_contended.clone();
433 move || {
434 let _permit = set_execution_permit(Box::new(DropFlag(waiter_contended)));
435 drop(mutex_table.acquire_lock("john".to_string()));
436 }
437 });
438
439 while !waiter_contended.load(Ordering::SeqCst) {
440 tokio::time::sleep(Duration::from_millis(1)).await;
441 }
442
443 assert!(
444 mutex_table.try_acquire_lock("jane".to_string()).is_ok(),
445 "a waiter on one key must not hold the shared shard lock"
446 );
447
448 release_holder.store(true, Ordering::SeqCst);
449 holder.await.unwrap();
450 waiter.await.unwrap();
451}
452
453#[tokio::test]
454async fn test_mutex_table() {
455 let mutex_table =
457 MutexTable::<String>::new_with_cleanup(1, Duration::from_secs(10), Duration::MAX, 1000);
458 let john1 = mutex_table.try_acquire_lock("john".to_string());
459 assert!(john1.is_ok());
460 let john2 = mutex_table.try_acquire_lock("john".to_string());
461 assert!(john2.is_err());
462 drop(john1);
463 let john2 = mutex_table.try_acquire_lock("john".to_string());
464 assert!(john2.is_ok());
465 let jane = mutex_table.try_acquire_lock("jane".to_string());
466 assert!(jane.is_ok());
467 MutexTable::cleanup(mutex_table.lock_table.clone());
468 let map = mutex_table.lock_table.first().as_ref().unwrap().try_read();
469 assert!(map.is_some());
470 assert_eq!(map.unwrap().len(), 2);
471 drop(john2);
472 MutexTable::cleanup(mutex_table.lock_table.clone());
473 let map = mutex_table.lock_table.first().as_ref().unwrap().try_read();
474 assert!(map.is_some());
475 assert_eq!(map.unwrap().len(), 1);
476 drop(jane);
477 MutexTable::cleanup(mutex_table.lock_table.clone());
478 let map = mutex_table.lock_table.first().as_ref().unwrap().try_read();
479 assert!(map.is_some());
480 assert!(map.unwrap().is_empty());
481}
482
483#[tokio::test]
484async fn test_acquire_locks() {
485 let mutex_table =
486 RwLockTable::<String>::new_with_cleanup(1, Duration::from_secs(10), Duration::MAX, 1000);
487 let object_1 = "object 1".to_string();
488 let object_2 = "object 2".to_string();
489 let object_3 = "object 3".to_string();
490
491 let objects = vec![
493 object_1.clone(),
494 object_2.clone(),
495 object_2,
496 object_1.clone(),
497 object_3,
498 object_1,
499 ];
500
501 let locks = mutex_table.acquire_locks(objects.clone().into_iter());
502 assert_eq!(locks.len(), 3);
503
504 for object in objects.clone() {
505 assert!(mutex_table.try_acquire_lock(object).is_err());
506 }
507
508 drop(locks);
509 let locks = mutex_table.acquire_locks(objects.into_iter());
510 assert_eq!(locks.len(), 3);
511}
512
513#[tokio::test]
514async fn test_read_locks() {
515 let mutex_table =
516 RwLockTable::<String>::new_with_cleanup(1, Duration::from_secs(10), Duration::MAX, 1000);
517 let lock = "lock".to_string();
518 let locks1 = mutex_table.acquire_read_locks(vec![lock.clone()]);
519 assert!(mutex_table.try_acquire_lock(lock.clone()).is_err());
520 let locks2 = mutex_table.acquire_read_locks(vec![lock.clone()]);
521 drop(locks1);
522 drop(locks2);
523 assert!(mutex_table.try_acquire_lock(lock.clone()).is_ok());
524}
525
526#[tokio::test(flavor = "current_thread", start_paused = true)]
527async fn test_mutex_table_bg_cleanup() {
528 let mutex_table = MutexTable::<String>::new_with_cleanup(
529 1,
530 Duration::from_secs(5),
531 Duration::from_secs(1),
532 1000,
533 );
534 let lock1 = mutex_table.try_acquire_lock("lock1".to_string());
535 let lock2 = mutex_table.try_acquire_lock("lock2".to_string());
536 let lock3 = mutex_table.try_acquire_lock("lock3".to_string());
537 let lock4 = mutex_table.try_acquire_lock("lock4".to_string());
538 let lock5 = mutex_table.try_acquire_lock("lock5".to_string());
539 assert!(lock1.is_ok());
540 assert!(lock2.is_ok());
541 assert!(lock3.is_ok());
542 assert!(lock4.is_ok());
543 assert!(lock5.is_ok());
544 MutexTable::cleanup(mutex_table.lock_table.clone());
546 let lock11 = mutex_table.try_acquire_lock("lock1".to_string());
548 let lock22 = mutex_table.try_acquire_lock("lock2".to_string());
549 let lock33 = mutex_table.try_acquire_lock("lock3".to_string());
550 let lock44 = mutex_table.try_acquire_lock("lock4".to_string());
551 let lock55 = mutex_table.try_acquire_lock("lock5".to_string());
552 assert!(lock11.is_err());
553 assert!(lock22.is_err());
554 assert!(lock33.is_err());
555 assert!(lock44.is_err());
556 assert!(lock55.is_err());
557 drop(lock1);
559 drop(lock2);
560 drop(lock3);
561 drop(lock4);
562 drop(lock5);
563 tokio::time::sleep(Duration::from_secs(10)).await;
565 for entry in mutex_table.lock_table.iter() {
566 let locked = entry.read();
567 assert!(locked.is_empty());
568 }
569}
570
571#[tokio::test(flavor = "current_thread", start_paused = true)]
572async fn test_mutex_table_bg_cleanup_with_size_threshold() {
573 let mutex_table =
575 MutexTable::<String>::new_with_cleanup(1, Duration::MAX, Duration::from_secs(1), 5);
576 let lock1 = mutex_table.try_acquire_lock("lock1".to_string());
577 let lock2 = mutex_table.try_acquire_lock("lock2".to_string());
578 let lock3 = mutex_table.try_acquire_lock("lock3".to_string());
579 let lock4 = mutex_table.try_acquire_lock("lock4".to_string());
580 let lock5 = mutex_table.try_acquire_lock("lock5".to_string());
581 assert!(lock1.is_ok());
582 assert!(lock2.is_ok());
583 assert!(lock3.is_ok());
584 assert!(lock4.is_ok());
585 assert!(lock5.is_ok());
586 MutexTable::cleanup(mutex_table.lock_table.clone());
588 let lock11 = mutex_table.try_acquire_lock("lock1".to_string());
590 let lock22 = mutex_table.try_acquire_lock("lock2".to_string());
591 let lock33 = mutex_table.try_acquire_lock("lock3".to_string());
592 let lock44 = mutex_table.try_acquire_lock("lock4".to_string());
593 let lock55 = mutex_table.try_acquire_lock("lock5".to_string());
594 assert!(lock11.is_err());
595 assert!(lock22.is_err());
596 assert!(lock33.is_err());
597 assert!(lock44.is_err());
598 assert!(lock55.is_err());
599 assert_eq!(mutex_table.size(), 5);
600 drop(lock1);
602 drop(lock2);
603 drop(lock3);
604 drop(lock4);
605 drop(lock5);
606 tokio::task::yield_now().await;
607 tokio::time::advance(Duration::from_secs(5)).await;
609 tokio::task::yield_now().await;
610 assert_eq!(mutex_table.size(), 0);
611 for entry in mutex_table.lock_table.iter() {
612 let locked = entry.read();
613 assert!(locked.is_empty());
614 }
615}