Skip to main content

klieo_bus_memory/
kv.rs

1//! In-process `KvStore` implementation.
2//!
3//! State: a `BTreeMap<(bucket, key), Entry>` guarded by a single
4//! `tokio::sync::Mutex`. Per-entry monotonically-increasing revision.
5//! Leases share the same map; heartbeating extends the lease's
6//! `expires_at`. The map is key-ordered so `keys_paginated` ranges from a
7//! cursor in `O(log n + page)` without rescanning the bucket.
8
9use 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    /// `Some(deadline)` — entry is leased until `deadline`. `None` — no
24    /// active lease.
25    leased_until: Option<Instant>,
26}
27
28#[derive(Default)]
29struct State {
30    /// `(bucket, key) -> Entry`.
31    map: BTreeMap<(String, String), Entry>,
32}
33
34/// In-process `KvStore` impl.
35pub struct MemoryKv {
36    state: Arc<Mutex<State>>,
37}
38
39impl MemoryKv {
40    /// Build an empty store.
41    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    /// Ranges the key-ordered map from the cursor, reading only `limit`(+1) keys
155    /// — no full-bucket rescan, unlike the trait default. `limit` is floored to
156    /// 1 so a zero never stalls a page-walk.
157    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
208/// Lease handle. Heartbeat extends `leased_until` by `ttl` from now.
209struct 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        // the value key is unaffected by the sidecar causer key
258        assert_eq!(
259            kv.get("b", "k").await.unwrap().unwrap().value,
260            Bytes::from_static(b"val")
261        );
262        // an unstamped key has no causer
263        kv.put("b", "k2", Bytes::from_static(b"v2")).await.unwrap();
264        assert_eq!(kv.causer_of("b", "k2").await.unwrap(), None);
265
266        // first-writer-wins: a later put_causer must NOT overwrite the causer.
267        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        // scan_bucket hides the reserved causer sidecar keys.
276        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}