1use async_trait::async_trait;
10use bytes::Bytes;
11use klieo_core::bus::{KeyPage, KvEntry, KvStore, Lease, LeaseImpl, Revision};
12use klieo_core::error::BusError;
13use std::collections::BTreeMap;
14use std::ops::Bound;
15use std::sync::Arc;
16use std::time::{Duration, Instant};
17use tokio::sync::Mutex;
18
19#[derive(Debug, Clone)]
20struct Entry {
21 value: Bytes,
22 revision: Revision,
23 leased_until: Option<Instant>,
26}
27
28#[derive(Default)]
29struct State {
30 map: BTreeMap<(String, String), Entry>,
32}
33
34pub struct MemoryKv {
36 state: Arc<Mutex<State>>,
37}
38
39impl MemoryKv {
40 pub fn new() -> Self {
42 Self {
43 state: Arc::new(Mutex::new(State::default())),
44 }
45 }
46}
47
48impl Default for MemoryKv {
49 fn default() -> Self {
50 Self::new()
51 }
52}
53
54#[async_trait]
55impl KvStore for MemoryKv {
56 async fn get(&self, bucket: &str, key: &str) -> Result<Option<KvEntry>, BusError> {
57 let g = self.state.lock().await;
58 Ok(g.map.get(&(bucket.into(), key.into())).map(|e| KvEntry {
59 value: e.value.clone(),
60 revision: e.revision,
61 }))
62 }
63
64 async fn put(&self, bucket: &str, key: &str, value: Bytes) -> Result<Revision, BusError> {
65 let mut g = self.state.lock().await;
66 let key = (bucket.into(), key.into());
67 let next_rev = g.map.get(&key).map(|e| e.revision + 1).unwrap_or(1);
68 g.map.insert(
69 key,
70 Entry {
71 value,
72 revision: next_rev,
73 leased_until: None,
74 },
75 );
76 Ok(next_rev)
77 }
78
79 async fn cas(
80 &self,
81 bucket: &str,
82 key: &str,
83 value: Bytes,
84 expected: Option<Revision>,
85 ) -> Result<Revision, BusError> {
86 let mut g = self.state.lock().await;
87 let map_key = (bucket.into(), key.into());
88 let existing = g.map.get(&map_key);
89 match (expected, existing) {
90 (None, None) => {
91 g.map.insert(
92 map_key,
93 Entry {
94 value,
95 revision: 1,
96 leased_until: None,
97 },
98 );
99 Ok(1)
100 }
101 (None, Some(e)) => Err(BusError::CasConflict {
102 expected: 0,
103 actual: e.revision,
104 }),
105 (Some(want), Some(e)) if e.revision == want => {
106 let next = e.revision + 1;
107 g.map.insert(
108 map_key,
109 Entry {
110 value,
111 revision: next,
112 leased_until: None,
113 },
114 );
115 Ok(next)
116 }
117 (Some(want), Some(e)) => Err(BusError::CasConflict {
118 expected: want,
119 actual: e.revision,
120 }),
121 (Some(want), None) => Err(BusError::CasConflict {
122 expected: want,
123 actual: 0,
124 }),
125 }
126 }
127
128 async fn delete(&self, bucket: &str, key: &str) -> Result<(), BusError> {
129 let mut g = self.state.lock().await;
130 g.map.remove(&(bucket.into(), key.into()));
131 Ok(())
132 }
133
134 async fn keys(&self, bucket: &str) -> Result<Vec<String>, BusError> {
135 let g = self.state.lock().await;
136 Ok(g.map
137 .keys()
138 .filter(|(b, _)| b == bucket)
139 .map(|(_, k)| k.clone())
140 .collect())
141 }
142
143 async fn scan_bucket(&self, bucket: &str) -> Result<Vec<(String, Bytes)>, BusError> {
144 let g = self.state.lock().await;
145 let entries = g
146 .map
147 .iter()
148 .filter(|((b, k), _)| b == bucket && !klieo_core::bus::is_causer_key(k))
149 .map(|((_, k), e)| (k.clone(), e.value.clone()))
150 .collect();
151 Ok(entries)
152 }
153
154 async fn keys_paginated(
158 &self,
159 bucket: &str,
160 cursor: Option<String>,
161 limit: usize,
162 ) -> Result<KeyPage, BusError> {
163 let limit = limit.max(1);
164 let start = match &cursor {
165 Some(c) => Bound::Excluded((bucket.to_string(), c.clone())),
166 None => Bound::Included((bucket.to_string(), String::new())),
167 };
168 let g = self.state.lock().await;
169 let mut after_cursor = g
170 .map
171 .range::<(String, String), _>((start, Bound::Unbounded))
172 .map(|((b, k), _)| (b.as_str(), k))
173 .take_while(|(b, _)| *b == bucket)
174 .map(|(_, k)| k.clone());
175 let keys: Vec<String> = after_cursor.by_ref().take(limit).collect();
176 let next = if after_cursor.next().is_some() {
177 keys.last().cloned()
178 } else {
179 None
180 };
181 Ok(KeyPage::new(keys, next))
182 }
183
184 async fn lease(&self, bucket: &str, key: &str, ttl: Duration) -> Result<Lease, BusError> {
185 let now = Instant::now();
186 let mut g = self.state.lock().await;
187 let map_key = (bucket.to_string(), key.to_string());
188 let entry = g.map.entry(map_key.clone()).or_insert_with(|| Entry {
189 value: Bytes::new(),
190 revision: 0,
191 leased_until: None,
192 });
193 if let Some(deadline) = entry.leased_until {
194 if deadline > now {
195 return Err(BusError::Permanent("lease already held".into()));
196 }
197 }
198 entry.leased_until = Some(now + ttl);
199 Ok(Lease::new(Box::new(MemoryLease {
200 state: self.state.clone(),
201 bucket: bucket.to_string(),
202 key: key.to_string(),
203 ttl,
204 })))
205 }
206}
207
208struct MemoryLease {
210 state: Arc<Mutex<State>>,
211 bucket: String,
212 key: String,
213 ttl: Duration,
214}
215
216#[async_trait]
217impl LeaseImpl for MemoryLease {
218 async fn heartbeat(&self) -> Result<(), BusError> {
219 let mut g = self.state.lock().await;
220 let key = (self.bucket.clone(), self.key.clone());
221 match g.map.get_mut(&key) {
222 Some(e) => {
223 e.leased_until = Some(Instant::now() + self.ttl);
224 Ok(())
225 }
226 None => Err(BusError::NotFound(format!("{}/{}", self.bucket, self.key))),
227 }
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234 use bytes::Bytes;
235 use klieo_core::error::BusError;
236 use klieo_core::KvStore;
237 use std::time::Duration;
238
239 #[tokio::test]
240 async fn put_then_get_returns_value() {
241 let kv = MemoryKv::new();
242 let rev = kv.put("b", "k", Bytes::from_static(b"v1")).await.unwrap();
243 assert!(rev > 0);
244 let entry = kv.get("b", "k").await.unwrap().unwrap();
245 assert_eq!(entry.value, Bytes::from_static(b"v1"));
246 assert_eq!(entry.revision, rev);
247 }
248
249 #[tokio::test]
250 async fn kv_causer_round_trips_and_leaves_value_key_intact() {
251 use klieo_core::ids::RunId;
252 let kv = MemoryKv::new();
253 let run = RunId::new();
254 kv.put("b", "k", Bytes::from_static(b"val")).await.unwrap();
255 kv.put_causer("b", "k", run).await.unwrap();
256 assert_eq!(kv.causer_of("b", "k").await.unwrap(), Some(run));
257 assert_eq!(
259 kv.get("b", "k").await.unwrap().unwrap().value,
260 Bytes::from_static(b"val")
261 );
262 kv.put("b", "k2", Bytes::from_static(b"v2")).await.unwrap();
264 assert_eq!(kv.causer_of("b", "k2").await.unwrap(), None);
265
266 let later = RunId::new();
268 kv.put_causer("b", "k", later).await.unwrap();
269 assert_eq!(
270 kv.causer_of("b", "k").await.unwrap(),
271 Some(run),
272 "the first writer's causer survives a later stamp"
273 );
274
275 let scanned = kv.scan_bucket("b").await.unwrap();
277 assert!(
278 scanned
279 .iter()
280 .all(|(k, _)| !klieo_core::bus::is_causer_key(k)),
281 "scan_bucket must not surface causer sidecar keys; got {scanned:?}"
282 );
283 assert!(
284 scanned.iter().any(|(k, _)| k == "k"),
285 "the value key is still present in scan_bucket"
286 );
287 }
288
289 #[tokio::test]
290 async fn cas_create_if_absent_succeeds_when_key_missing() {
291 let kv = MemoryKv::new();
292 let rev = kv
293 .cas("b", "k", Bytes::from_static(b"v"), None)
294 .await
295 .unwrap();
296 assert_eq!(rev, 1);
297 }
298
299 #[tokio::test]
300 async fn cas_create_if_absent_conflicts_when_key_present() {
301 let kv = MemoryKv::new();
302 kv.put("b", "k", Bytes::from_static(b"v0")).await.unwrap();
303 let err = kv
304 .cas("b", "k", Bytes::from_static(b"v1"), None)
305 .await
306 .unwrap_err();
307 assert!(matches!(err, BusError::CasConflict { .. }));
308 }
309
310 #[tokio::test]
311 async fn cas_with_expected_revision_succeeds_on_match() {
312 let kv = MemoryKv::new();
313 let r0 = kv.put("b", "k", Bytes::from_static(b"v0")).await.unwrap();
314 let r1 = kv
315 .cas("b", "k", Bytes::from_static(b"v1"), Some(r0))
316 .await
317 .unwrap();
318 assert_eq!(r1, r0 + 1);
319 }
320
321 #[tokio::test]
322 async fn cas_with_expected_revision_conflicts_on_stale() {
323 let kv = MemoryKv::new();
324 let r0 = kv.put("b", "k", Bytes::from_static(b"v0")).await.unwrap();
325 kv.put("b", "k", Bytes::from_static(b"v1")).await.unwrap();
326 let err = kv
327 .cas("b", "k", Bytes::from_static(b"v2"), Some(r0))
328 .await
329 .unwrap_err();
330 assert!(matches!(err, BusError::CasConflict { .. }));
331 }
332
333 #[tokio::test]
334 async fn delete_removes_key() {
335 let kv = MemoryKv::new();
336 kv.put("b", "k", Bytes::from_static(b"v")).await.unwrap();
337 kv.delete("b", "k").await.unwrap();
338 assert!(kv.get("b", "k").await.unwrap().is_none());
339 }
340
341 #[tokio::test]
342 async fn lease_can_heartbeat() {
343 let kv = MemoryKv::new();
344 let lease = kv.lease("b", "k", Duration::from_millis(50)).await.unwrap();
345 lease.heartbeat().await.unwrap();
346 }
347
348 #[tokio::test]
349 async fn keys_returns_keys_in_bucket() {
350 let kv = MemoryKv::new();
351 kv.put("b1", "alpha", Bytes::from_static(b"a"))
352 .await
353 .unwrap();
354 kv.put("b1", "beta", Bytes::from_static(b"b"))
355 .await
356 .unwrap();
357 kv.put("b2", "gamma", Bytes::from_static(b"c"))
358 .await
359 .unwrap();
360 let mut keys = kv.keys("b1").await.unwrap();
361 keys.sort();
362 assert_eq!(keys, vec!["alpha".to_string(), "beta".to_string()]);
363 let keys_other = kv.keys("b2").await.unwrap();
364 assert_eq!(keys_other, vec!["gamma".to_string()]);
365 let keys_empty = kv.keys("missing").await.unwrap();
366 assert!(keys_empty.is_empty());
367 }
368
369 #[tokio::test]
370 async fn keys_paginated_walks_bucket_in_order_across_pages() {
371 let kv = MemoryKv::new();
372 for k in ["e", "a", "d", "b", "c"] {
373 kv.put("b1", k, Bytes::from_static(b"v")).await.unwrap();
374 }
375 kv.put("b2", "zzz", Bytes::from_static(b"v")).await.unwrap();
376
377 let mut seen = Vec::new();
378 let mut cursor = None;
379 loop {
380 let page = kv.keys_paginated("b1", cursor, 2).await.unwrap();
381 assert!(page.keys.len() <= 2, "page respects the limit");
382 seen.extend(page.keys.iter().cloned());
383 match page.next {
384 Some(c) => cursor = Some(c),
385 None => break,
386 }
387 }
388 assert_eq!(
389 seen,
390 vec!["a", "b", "c", "d", "e"],
391 "every key once, ascending; the b2 key never leaks in"
392 );
393 }
394
395 #[tokio::test]
396 async fn keys_paginated_empty_bucket_yields_one_terminal_page() {
397 let kv = MemoryKv::new();
398 let page = kv.keys_paginated("missing", None, 8).await.unwrap();
399 assert!(page.keys.is_empty());
400 assert!(page.next.is_none());
401 }
402
403 #[tokio::test]
404 async fn keys_paginated_limit_past_count_is_single_page() {
405 let kv = MemoryKv::new();
406 kv.put("b1", "k1", Bytes::from_static(b"v")).await.unwrap();
407 kv.put("b1", "k2", Bytes::from_static(b"v")).await.unwrap();
408 let page = kv.keys_paginated("b1", None, 16).await.unwrap();
409 assert_eq!(page.keys, vec!["k1".to_string(), "k2".to_string()]);
410 assert!(page.next.is_none());
411 }
412
413 #[tokio::test]
414 async fn scan_bucket_returns_all_pairs_in_single_read() {
415 let kv = MemoryKv::new();
416 kv.put("b1", "k1", Bytes::from_static(b"v1")).await.unwrap();
417 kv.put("b1", "k2", Bytes::from_static(b"v2")).await.unwrap();
418 kv.put("b2", "k3", Bytes::from_static(b"v3")).await.unwrap();
419
420 let mut pairs = kv.scan_bucket("b1").await.unwrap();
421 pairs.sort_by(|a, b| a.0.cmp(&b.0));
422
423 assert_eq!(pairs.len(), 2);
424 assert_eq!(pairs[0], ("k1".to_string(), Bytes::from_static(b"v1")));
425 assert_eq!(pairs[1], ("k2".to_string(), Bytes::from_static(b"v2")));
426
427 let pairs_other = kv.scan_bucket("b2").await.unwrap();
428 assert_eq!(pairs_other.len(), 1);
429 assert_eq!(pairs_other[0].0, "k3");
430
431 let pairs_empty = kv.scan_bucket("missing").await.unwrap();
432 assert!(pairs_empty.is_empty());
433 }
434}