1use std::sync::Arc;
8use std::time::{Duration, Instant};
9
10use async_trait::async_trait;
11use tokio::sync::Mutex;
12
13#[derive(Debug)]
15pub enum LockError {
16 AlreadyHeld,
18 Expired,
20 Io(String),
22}
23
24impl std::fmt::Display for LockError {
25 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26 match self {
27 LockError::AlreadyHeld => write!(f, "lock already held"),
28 LockError::Expired => write!(f, "lock lease expired"),
29 LockError::Io(s) => write!(f, "lock io error: {s}"),
30 }
31 }
32}
33
34impl std::error::Error for LockError {}
35
36type LockMap = Arc<Mutex<std::collections::HashMap<String, Instant>>>;
38
39pub struct LockGuard {
41 map: LockMap,
42 key: String,
43}
44
45impl LockGuard {
46 pub fn key(&self) -> &str {
47 &self.key
48 }
49
50 pub async fn extend(&mut self, _dur: Duration) -> Result<(), LockError> {
52 Err(LockError::Expired)
53 }
54}
55
56impl Drop for LockGuard {
57 fn drop(&mut self) {
58 let map = Arc::clone(&self.map);
59 let key = self.key.clone();
60 tokio::spawn(async move {
62 let mut m = map.lock().await;
63 m.remove(&key);
64 });
65 }
66}
67
68#[async_trait]
70pub trait DistributedLock: Send + Sync {
71 async fn acquire(&self, key: &str, lease: Duration, timeout: Duration) -> Result<LockGuard, LockError>;
73
74 async fn release(&self, key: &str) -> Result<(), LockError>;
76
77 async fn extend(&self, key: &str, lease: Duration) -> Result<(), LockError>;
79}
80
81pub struct InProcLock {
84 held: LockMap,
85}
86
87impl InProcLock {
88 pub fn new() -> Self {
89 Self {
90 held: Arc::new(Mutex::new(std::collections::HashMap::new())),
91 }
92 }
93
94 fn is_expired(map: &std::collections::HashMap<String, Instant>, key: &str) -> bool {
95 if let Some(expiry) = map.get(key) {
96 *expiry <= Instant::now()
97 } else {
98 false
99 }
100 }
101}
102
103impl Default for InProcLock {
104 fn default() -> Self {
105 Self::new()
106 }
107}
108
109#[async_trait]
110impl DistributedLock for InProcLock {
111 async fn acquire(&self, key: &str, lease: Duration, timeout: Duration) -> Result<LockGuard, LockError> {
112 let deadline = Instant::now() + timeout;
113 loop {
114 {
115 let mut map = self.held.lock().await;
116 map.retain(|_, v| *v > Instant::now());
118 if !map.contains_key(key) {
119 map.insert(key.to_string(), Instant::now() + lease);
120 return Ok(LockGuard {
121 map: Arc::clone(&self.held),
122 key: key.to_string(),
123 });
124 }
125 }
126 if Instant::now() >= deadline {
127 return Err(LockError::AlreadyHeld);
128 }
129 tokio::time::sleep(Duration::from_millis(10)).await;
130 }
131 }
132
133 async fn release(&self, key: &str) -> Result<(), LockError> {
134 let mut map = self.held.lock().await;
135 map.remove(key);
136 Ok(())
137 }
138
139 async fn extend(&self, key: &str, lease: Duration) -> Result<(), LockError> {
140 let mut map = self.held.lock().await;
141 if let Some(entry) = map.get_mut(key) {
142 *entry = Instant::now() + lease;
143 Ok(())
144 } else {
145 Err(LockError::AlreadyHeld)
146 }
147 }
148}
149
150#[cfg(test)]
151mod tests {
152 use super::*;
153
154 #[tokio::test]
155 async fn lock_acquire_release() {
156 let lock = InProcLock::new();
157 let guard = lock.acquire("key", Duration::from_secs(10), Duration::from_secs(1)).await.unwrap();
158 assert!(lock.release("key").await.is_ok());
159 drop(guard);
160 }
161
162 #[tokio::test]
163 async fn lock_rejects_second_acquire() {
164 let lock = InProcLock::new();
165 let _guard1 = lock.acquire("key", Duration::from_secs(10), Duration::from_secs(1)).await.unwrap();
166 let result = lock.acquire("key", Duration::from_secs(10), Duration::from_millis(50)).await;
168 assert!(result.is_err());
169 drop(_guard1);
170 }
171
172 #[tokio::test]
173 async fn lock_auto_releases_on_drop() {
174 let lock = InProcLock::new();
175 let guard = lock.acquire("k", Duration::from_secs(10), Duration::from_secs(1)).await.unwrap();
176 drop(guard);
177 let result = lock.acquire("k", Duration::from_secs(10), Duration::from_millis(50)).await;
179 assert!(result.is_ok(), "lock should be free after guard drop");
180 }
181
182 #[tokio::test]
183 async fn lock_expires_after_lease() {
184 let lock = InProcLock::new();
185 let _guard = lock.acquire("key", Duration::from_millis(20), Duration::from_millis(5)).await.unwrap();
186 drop(_guard);
187 tokio::time::sleep(Duration::from_millis(30)).await;
188 let result = lock.acquire("key", Duration::from_millis(20), Duration::from_millis(5)).await;
190 assert!(result.is_ok());
191 }
192}