chio_kernel/admission_operation/
sequencer.rs1use std::collections::HashMap;
2use std::sync::{Arc, Mutex, MutexGuard, OnceLock, Weak};
3
4use super::*;
5
6type MutationSequencers = HashMap<String, Weak<Mutex<()>>>;
7
8#[derive(Clone)]
9pub struct AdmissionMutationSequencer {
10 inner: Arc<Mutex<()>>,
11}
12
13pub struct AdmissionMutationGuard<'a> {
14 _guard: MutexGuard<'a, ()>,
15}
16
17impl AdmissionMutationSequencer {
18 pub fn for_fence(fence: &StoreMutationFence) -> Result<Self, AdmissionOperationError> {
19 validate_store_fence(fence)?;
20 static SEQUENCERS: OnceLock<Mutex<MutationSequencers>> = OnceLock::new();
21
22 let key = format!(
23 "{}\0{}\0{}",
24 fence.store_uuid, fence.lease_id, fence.owner_epoch
25 );
26 let registry = SEQUENCERS.get_or_init(|| Mutex::new(HashMap::new()));
27 let mut sequencers = registry
28 .lock()
29 .map_err(|_| AdmissionOperationError::MutationSequencerPoisoned)?;
30 sequencers.retain(|_, sequencer| sequencer.strong_count() > 0);
31 let inner = if let Some(sequencer) = sequencers.get(&key).and_then(Weak::upgrade) {
32 sequencer
33 } else {
34 let sequencer = Arc::new(Mutex::new(()));
35 sequencers.insert(key, Arc::downgrade(&sequencer));
36 sequencer
37 };
38 Ok(Self { inner })
39 }
40
41 pub fn lock(&self) -> Result<AdmissionMutationGuard<'_>, AdmissionOperationError> {
42 self.inner
43 .lock()
44 .map(|guard| AdmissionMutationGuard { _guard: guard })
45 .map_err(|_| AdmissionOperationError::MutationSequencerPoisoned)
46 }
47}
48
49#[cfg(test)]
50mod tests {
51 use std::sync::mpsc;
52 use std::time::Duration;
53
54 use super::*;
55
56 #[test]
57 fn identical_store_fences_share_one_mutation_sequence() {
58 let fence = StoreMutationFence {
59 store_uuid: "store-1".to_owned(),
60 lease_id: "lease-1".to_owned(),
61 owner_epoch: 1,
62 };
63 let first = AdmissionMutationSequencer::for_fence(&fence).expect("first sequencer");
64 let second = AdmissionMutationSequencer::for_fence(&fence).expect("second sequencer");
65 let held = first.lock().expect("first lock");
66 let (sender, receiver) = mpsc::channel();
67 let waiter = std::thread::spawn(move || {
68 sender.send("waiting").expect("send waiting");
69 let _guard = second.lock().expect("second lock");
70 sender.send("acquired").expect("send acquired");
71 });
72
73 assert_eq!(receiver.recv().expect("receive waiting"), "waiting");
74 assert!(receiver.recv_timeout(Duration::from_millis(25)).is_err());
75 drop(held);
76 assert_eq!(
77 receiver
78 .recv_timeout(Duration::from_secs(1))
79 .expect("receive acquired"),
80 "acquired"
81 );
82 waiter.join().expect("join waiter");
83 }
84}