Skip to main content

amalgam/
distributed_lock.rs

1//! Cross-node single-flight: the [`DistributedLocker`] seam.
2//!
3//! In-process stampede protection ([`KeyedLock`](crate::locking::KeyedLock))
4//! guarantees one factory per key *per node*. A distributed locker extends that
5//! to *one factory per key across the cluster*, matching FusionCache's
6//! `IFusionCacheDistributedLocker`. It is token-based: `acquire` returns an
7//! opaque token that `release` consumes, so a Redis (`SET key token NX PX`) or
8//! database implementation maps cleanly.
9
10use std::sync::Arc;
11use std::time::Duration;
12
13use async_trait::async_trait;
14use dashmap::DashMap;
15
16use crate::error::Result;
17use crate::time::{Clock, Timeout, Timestamp};
18
19/// A cross-node lock backend.
20#[async_trait]
21pub trait DistributedLocker: Send + Sync {
22    /// Attempts to acquire the lock for `key`, waiting up to `timeout`.
23    ///
24    /// Returns `Some(token)` on success (the token must be passed back to
25    /// [`release`](DistributedLocker::release)), or `None` if the wait elapsed.
26    /// `ttl` bounds how long the lock is held even if the holder dies.
27    ///
28    /// # Errors
29    /// Returns an error only on backend failure, not on a normal timeout.
30    async fn acquire(&self, key: &str, ttl: Duration, timeout: Timeout) -> Result<Option<String>>;
31
32    /// Releases a previously-acquired lock. Releasing a token that is no longer
33    /// held (e.g. it already expired) is a no-op, not an error.
34    ///
35    /// # Errors
36    /// Returns an error only on backend failure.
37    async fn release(&self, key: &str, token: &str) -> Result<()>;
38}
39
40/// An in-process reference [`DistributedLocker`].
41///
42/// Share one instance (via `Arc`) between several caches to get cross-instance
43/// single-flight within one process. A real cluster uses a Redis-backed locker.
44pub struct InMemoryDistributedLocker {
45    locks: Arc<DashMap<String, Held>>,
46    clock: Arc<dyn Clock>,
47}
48
49#[derive(Clone)]
50struct Held {
51    token: String,
52    expires_at: Timestamp,
53}
54
55impl InMemoryDistributedLocker {
56    /// Creates a locker using `clock` for lock TTL accounting.
57    #[must_use]
58    pub fn new(clock: Arc<dyn Clock>) -> Self {
59        Self {
60            locks: Arc::new(DashMap::new()),
61            clock,
62        }
63    }
64
65    /// Attempts a single non-blocking acquisition. Returns the token on success.
66    fn try_once(&self, key: &str, ttl: Duration, now: Timestamp) -> Option<String> {
67        use dashmap::mapref::entry::Entry;
68        let expires_at = now.saturating_add(ttl);
69        match self.locks.entry(key.to_owned()) {
70            Entry::Occupied(mut occupied) => {
71                if now.is_before(occupied.get().expires_at) {
72                    None // still held by someone else
73                } else {
74                    let token = new_token();
75                    occupied.insert(Held {
76                        token: token.clone(),
77                        expires_at,
78                    });
79                    Some(token)
80                }
81            }
82            Entry::Vacant(vacant) => {
83                let token = new_token();
84                vacant.insert(Held {
85                    token: token.clone(),
86                    expires_at,
87                });
88                Some(token)
89            }
90        }
91    }
92}
93
94fn new_token() -> String {
95    format!("{:016x}{:016x}", fastrand::u64(..), fastrand::u64(..))
96}
97
98#[async_trait]
99impl DistributedLocker for InMemoryDistributedLocker {
100    async fn acquire(&self, key: &str, ttl: Duration, timeout: Timeout) -> Result<Option<String>> {
101        let deadline = match timeout {
102            Timeout::After(d) => Some(std::time::Instant::now() + d),
103            Timeout::Infinite => None,
104        };
105        loop {
106            if let Some(token) = self.try_once(key, ttl, self.clock.now()) {
107                return Ok(Some(token));
108            }
109            if let Some(deadline) = deadline
110                && std::time::Instant::now() >= deadline
111            {
112                return Ok(None);
113            }
114            tokio::time::sleep(Duration::from_millis(10)).await;
115        }
116    }
117
118    async fn release(&self, key: &str, token: &str) -> Result<()> {
119        self.locks.remove_if(key, |_, held| held.token == token);
120        Ok(())
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::*;
127    use crate::time::ManualClock;
128
129    #[tokio::test]
130    async fn second_acquire_blocks_until_release() {
131        let clock = Arc::new(ManualClock::default());
132        let locker = InMemoryDistributedLocker::new(clock);
133        let token = locker
134            .acquire("k", Duration::from_secs(30), Timeout::Infinite)
135            .await
136            .unwrap()
137            .expect("first acquire succeeds");
138        // A non-waiting second attempt fails while held.
139        let second = locker
140            .acquire("k", Duration::from_secs(30), Timeout::After(Duration::ZERO))
141            .await
142            .unwrap();
143        assert!(second.is_none());
144        locker.release("k", &token).await.unwrap();
145        // Now it can be acquired again.
146        assert!(
147            locker
148                .acquire("k", Duration::from_secs(30), Timeout::After(Duration::ZERO))
149                .await
150                .unwrap()
151                .is_some()
152        );
153    }
154}