Skip to main content

vote_commitment_tree/
kv_shard_store.rs

1//! [`KvShardStore`] — a [`ShardStore`] implementation backed by Go's Cosmos KV
2//! store via C function pointer callbacks.
3//!
4//! # Design
5//!
6//! Instead of maintaining an in-process copy of all shard data,
7//! `KvShardStore` forwards every [`ShardStore`] read and write directly to
8//! the Cosmos KV store through a set of C callbacks registered at creation
9//! time. Go registers `//export` functions that dispatch to the current
10//! block's `store.KVStore` through a stable proxy pointer.
11//!
12//! This gives `ShardTree` true lazy loading: on a cold start only the data
13//! that is actually accessed (the frontier shard + cap + checkpoints) is read.
14//! No explicit restore loop, no O(n) blob loading, no shard geometry in Go.
15//!
16//! # KV key schema (matches keys.go)
17//!
18//! | Prefix    | Key                              | Value           |
19//! |-----------|----------------------------------|-----------------|
20//! | `0x0F`    | `0x0F \|\| u64 BE shard_index`   | shard blob      |
21//! | `0x10`    | `0x10`                           | cap blob        |
22//! | `0x11`    | `0x11 \|\| u32 BE checkpoint_id` | checkpoint blob |
23//! | `0x12`    | `0x12 \|\| u32 BE checkpoint_id` | retained marker |
24//!
25//! # Buffer ownership
26//!
27//! `get` returns a C-malloc'd buffer that Rust frees with the provided
28//! `free_buf` callback after copying the value. All write callbacks receive
29//! a Rust-owned slice (pointer + length); they must copy the data if they
30//! need it to outlive the call.
31//!
32//! # Iterator protocol
33//!
34//! `iter_create(ctx, prefix, prefix_len, reverse)` returns an opaque handle
35//! (a `cgo.Handle` on the Go side). `iter_next` advances and writes
36//! C-malloc'd key + value; Rust frees each pair with `free_buf` before the
37//! next call. `iter_free` closes and drops the iterator. `iter_next` returns
38//! 0 on a valid entry, 1 when exhausted, -1 on error.
39
40use std::collections::BTreeSet;
41use std::fmt;
42use std::os::raw::c_void;
43
44use incrementalmerkletree::{Address, Level};
45use shardtree::{
46    store::{Checkpoint, ShardStore},
47    LocatedPrunableTree, LocatedTree, PrunableTree, Tree,
48};
49
50use crate::hash::{MerkleHashVote, SHARD_HEIGHT};
51use crate::serde::{read_checkpoint, read_shard_vote, write_checkpoint, write_shard_vote};
52
53// ---------------------------------------------------------------------------
54// KvError
55// ---------------------------------------------------------------------------
56
57/// Error type for [`KvShardStore`] operations.
58///
59/// Replaces `Infallible` so that KV callback failures are visible to callers
60/// rather than being silently swallowed. The three variants cover all
61/// observable failure modes:
62///
63/// - `IoError`: a KV callback returned a non-zero error code (disk full,
64///   store closed, etc.).
65/// - `Deserialization`: a blob retrieved from KV failed to decode.
66/// - `Serialization`: a shard or cap could not be encoded before writing.
67#[derive(Debug, Clone, PartialEq, Eq)]
68pub enum KvError {
69    /// A KV callback returned an error code (set, delete, or iterator failure).
70    IoError,
71    /// Shard or checkpoint data retrieved from KV could not be decoded.
72    Deserialization,
73    /// Shard or cap data could not be serialized before writing.
74    Serialization,
75}
76
77impl fmt::Display for KvError {
78    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79        match self {
80            KvError::IoError => write!(f, "KV callback returned an error"),
81            KvError::Deserialization => write!(f, "failed to deserialize KV data"),
82            KvError::Serialization => write!(f, "failed to serialize data for KV"),
83        }
84    }
85}
86
87impl std::error::Error for KvError {}
88
89// ---------------------------------------------------------------------------
90// KV key constants (must match keys.go 0x0F / 0x10 / 0x11 / 0x12)
91// ---------------------------------------------------------------------------
92
93const SHARD_PREFIX: u8 = 0x0F;
94const CAP_KEY: u8 = 0x10;
95const CHECKPOINT_PREFIX: u8 = 0x11;
96const RETAINED_CHECKPOINT_PREFIX: u8 = 0x12;
97
98fn shard_key(index: u64) -> [u8; 9] {
99    let mut k = [0u8; 9];
100    k[0] = SHARD_PREFIX;
101    k[1..].copy_from_slice(&index.to_be_bytes());
102    k
103}
104
105fn cap_key() -> [u8; 1] {
106    [CAP_KEY]
107}
108
109fn checkpoint_key(id: u32) -> [u8; 5] {
110    let mut k = [0u8; 5];
111    k[0] = CHECKPOINT_PREFIX;
112    k[1..].copy_from_slice(&id.to_be_bytes());
113    k
114}
115
116fn retained_checkpoint_key(id: u32) -> [u8; 5] {
117    let mut k = [0u8; 5];
118    k[0] = RETAINED_CHECKPOINT_PREFIX;
119    k[1..].copy_from_slice(&id.to_be_bytes());
120    k
121}
122
123// ---------------------------------------------------------------------------
124// Callback function pointer types
125// ---------------------------------------------------------------------------
126
127/// Retrieve a value from the KV store.
128///
129/// On success (key found) writes a C-malloc'd buffer to `*out_val` and its
130/// length to `*out_val_len`, then returns 0.
131/// Returns 1 if the key was not found (out pointers are unchanged).
132/// Returns -1 on error.
133pub type KvGetFn = unsafe extern "C" fn(
134    ctx: *mut c_void,
135    key: *const u8,
136    key_len: usize,
137    out_val: *mut *mut u8,
138    out_val_len: *mut usize,
139) -> i32;
140
141/// Write a key-value pair. Returns 0 on success, -1 on error.
142pub type KvSetFn = unsafe extern "C" fn(
143    ctx: *mut c_void,
144    key: *const u8,
145    key_len: usize,
146    val: *const u8,
147    val_len: usize,
148) -> i32;
149
150/// Delete a key. Returns 0 on success, -1 on error.
151pub type KvDeleteFn = unsafe extern "C" fn(ctx: *mut c_void, key: *const u8, key_len: usize) -> i32;
152
153/// Create an iterator over the given prefix.
154///
155/// `reverse` is 1 for a reverse (descending) iterator, 0 for ascending.
156/// Returns an opaque iterator handle, or null on error.
157pub type KvIterCreateFn = unsafe extern "C" fn(
158    ctx: *mut c_void,
159    prefix: *const u8,
160    prefix_len: usize,
161    reverse: u8,
162) -> *mut c_void;
163
164/// Advance the iterator and return the next key-value pair as C-malloc'd
165/// buffers. Caller frees with `free_buf`.
166///
167/// Returns 0 if a valid entry was written, 1 if exhausted, -1 on error.
168pub type KvIterNextFn = unsafe extern "C" fn(
169    iter: *mut c_void,
170    out_key: *mut *mut u8,
171    out_key_len: *mut usize,
172    out_val: *mut *mut u8,
173    out_val_len: *mut usize,
174) -> i32;
175
176/// Close and free an iterator handle.
177pub type KvIterFreeFn = unsafe extern "C" fn(iter: *mut c_void);
178
179/// Free a C-malloc'd buffer returned by a KV callback.
180pub type KvFreeBufFn = unsafe extern "C" fn(ptr: *mut u8, len: usize);
181
182// ---------------------------------------------------------------------------
183// KvCallbacks
184// ---------------------------------------------------------------------------
185
186/// Bundle of C function pointers + context passed to [`KvShardStore`].
187///
188/// # Safety
189/// All function pointers must remain valid for the lifetime of the
190/// `KvShardStore`. The `ctx` pointer must remain stable; Go achieves this
191/// via a `KvStoreProxy` whose address never changes across blocks.
192#[derive(Clone, Copy)]
193pub struct KvCallbacks {
194    pub ctx: *mut c_void,
195    pub get: KvGetFn,
196    pub set: KvSetFn,
197    pub delete: KvDeleteFn,
198    pub iter_create: KvIterCreateFn,
199    pub iter_next: KvIterNextFn,
200    pub iter_free: KvIterFreeFn,
201    pub free_buf: KvFreeBufFn,
202}
203
204// SAFETY: EndBlocker is single-threaded; all callbacks are called only on
205// the goroutine that owns the KV store.
206unsafe impl Send for KvCallbacks {}
207unsafe impl Sync for KvCallbacks {}
208
209// ---------------------------------------------------------------------------
210// Low-level helpers
211// ---------------------------------------------------------------------------
212
213impl KvCallbacks {
214    /// Fetch a value by key.
215    ///
216    /// Returns `Ok(Some(bytes))` if found, `Ok(None)` if not present, or
217    /// `Err(KvError::IoError)` if the callback signalled a hard error (rc=-1).
218    pub fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>, KvError> {
219        let mut out_ptr: *mut u8 = std::ptr::null_mut();
220        let mut out_len: usize = 0;
221        let rc = unsafe {
222            (self.get)(
223                self.ctx,
224                key.as_ptr(),
225                key.len(),
226                &mut out_ptr,
227                &mut out_len,
228            )
229        };
230        match rc {
231            0 => {
232                let val = unsafe { std::slice::from_raw_parts(out_ptr, out_len).to_vec() };
233                unsafe { (self.free_buf)(out_ptr, out_len) };
234                Ok(Some(val))
235            }
236            1 => Ok(None),              // not found
237            _ => Err(KvError::IoError), // rc=-1 or any other error code
238        }
239    }
240
241    /// Write a key-value pair. Returns `Err(KvError::IoError)` if the
242    /// callback returned a non-zero code.
243    pub fn set(&self, key: &[u8], val: &[u8]) -> Result<(), KvError> {
244        let rc = unsafe { (self.set)(self.ctx, key.as_ptr(), key.len(), val.as_ptr(), val.len()) };
245        if rc != 0 {
246            Err(KvError::IoError)
247        } else {
248            Ok(())
249        }
250    }
251
252    /// Delete a key. Returns `Err(KvError::IoError)` if the callback failed.
253    pub fn delete(&self, key: &[u8]) -> Result<(), KvError> {
254        let rc = unsafe { (self.delete)(self.ctx, key.as_ptr(), key.len()) };
255        if rc != 0 {
256            Err(KvError::IoError)
257        } else {
258            Ok(())
259        }
260    }
261
262    /// Create a forward or reverse iterator over the given prefix.
263    fn iter(&self, prefix: &[u8], reverse: bool) -> KvIter<'_> {
264        let handle =
265            unsafe { (self.iter_create)(self.ctx, prefix.as_ptr(), prefix.len(), reverse as u8) };
266        KvIter { handle, cb: self }
267    }
268}
269
270struct KvIter<'a> {
271    handle: *mut c_void,
272    cb: &'a KvCallbacks,
273}
274
275impl<'a> KvIter<'a> {
276    /// Advance and return `Some((key, value))`, or `None` when exhausted.
277    fn next(&mut self) -> Option<(Vec<u8>, Vec<u8>)> {
278        if self.handle.is_null() {
279            return None;
280        }
281        let mut key_ptr: *mut u8 = std::ptr::null_mut();
282        let mut key_len: usize = 0;
283        let mut val_ptr: *mut u8 = std::ptr::null_mut();
284        let mut val_len: usize = 0;
285        let rc = unsafe {
286            (self.cb.iter_next)(
287                self.handle,
288                &mut key_ptr,
289                &mut key_len,
290                &mut val_ptr,
291                &mut val_len,
292            )
293        };
294        if rc != 0 {
295            return None;
296        }
297        let key = unsafe { std::slice::from_raw_parts(key_ptr, key_len).to_vec() };
298        unsafe { (self.cb.free_buf)(key_ptr, key_len) };
299        let val = unsafe { std::slice::from_raw_parts(val_ptr, val_len).to_vec() };
300        unsafe { (self.cb.free_buf)(val_ptr, val_len) };
301        Some((key, val))
302    }
303}
304
305impl<'a> Drop for KvIter<'a> {
306    fn drop(&mut self) {
307        if !self.handle.is_null() {
308            unsafe { (self.cb.iter_free)(self.handle) };
309        }
310    }
311}
312
313// ---------------------------------------------------------------------------
314// KvShardStore
315// ---------------------------------------------------------------------------
316
317/// A [`ShardStore`] that stores all state in the Cosmos KV store via Go
318/// callbacks. Gives `ShardTree` true lazy loading: only the data it actually
319/// accesses is read from KV.
320pub struct KvShardStore {
321    pub(crate) cb: KvCallbacks,
322}
323
324impl KvShardStore {
325    pub fn new(cb: KvCallbacks) -> Self {
326        Self { cb }
327    }
328}
329
330// ---------------------------------------------------------------------------
331// ShardStore implementation
332// ---------------------------------------------------------------------------
333
334impl ShardStore for KvShardStore {
335    type H = MerkleHashVote;
336    type CheckpointId = u32;
337    type Error = KvError;
338
339    fn get_shard(
340        &self,
341        shard_root: Address,
342    ) -> Result<Option<LocatedPrunableTree<MerkleHashVote>>, KvError> {
343        let idx = shard_root.index();
344        let key = shard_key(idx);
345        let Some(blob) = self.cb.get(&key)? else {
346            return Ok(None);
347        };
348        match read_shard_vote(&blob) {
349            Ok(tree) => Ok(LocatedTree::from_parts(shard_root, tree).ok()),
350            Err(_) => Err(KvError::Deserialization),
351        }
352    }
353
354    fn last_shard(&self) -> Result<Option<LocatedPrunableTree<MerkleHashVote>>, KvError> {
355        let prefix = [SHARD_PREFIX];
356        let mut iter = self.cb.iter(&prefix, true /* reverse */);
357        let Some((key, val)) = iter.next() else {
358            return Ok(None);
359        };
360        if key.len() < 9 {
361            return Ok(None);
362        }
363        let idx = u64::from_be_bytes(key[1..9].try_into().unwrap());
364        let level = Level::from(SHARD_HEIGHT);
365        let addr = Address::from_parts(level, idx);
366        match read_shard_vote(&val) {
367            Ok(tree) => Ok(LocatedTree::from_parts(addr, tree).ok()),
368            Err(_) => Err(KvError::Deserialization),
369        }
370    }
371
372    fn put_shard(&mut self, subtree: LocatedPrunableTree<MerkleHashVote>) -> Result<(), KvError> {
373        let idx = subtree.root_addr().index();
374        let key = shard_key(idx);
375        let blob = write_shard_vote(subtree.root()).map_err(|_| KvError::Serialization)?;
376        self.cb.set(&key, &blob)
377    }
378
379    fn get_shard_roots(&self) -> Result<Vec<Address>, KvError> {
380        let prefix = [SHARD_PREFIX];
381        let mut iter = self.cb.iter(&prefix, false);
382        let level = Level::from(SHARD_HEIGHT);
383        let mut roots = Vec::new();
384        while let Some((key, _)) = iter.next() {
385            if key.len() < 9 {
386                continue;
387            }
388            let idx = u64::from_be_bytes(key[1..9].try_into().unwrap());
389            roots.push(Address::from_parts(level, idx));
390        }
391        Ok(roots)
392    }
393
394    fn truncate_shards(&mut self, shard_index: u64) -> Result<(), KvError> {
395        let prefix = [SHARD_PREFIX];
396        let mut iter = self.cb.iter(&prefix, false);
397        let mut to_delete = Vec::new();
398        while let Some((key, _)) = iter.next() {
399            if key.len() < 9 {
400                continue;
401            }
402            let idx = u64::from_be_bytes(key[1..9].try_into().unwrap());
403            if idx >= shard_index {
404                to_delete.push(key);
405            }
406        }
407        drop(iter);
408        for key in to_delete {
409            self.cb.delete(&key)?;
410        }
411        Ok(())
412    }
413
414    fn get_cap(&self) -> Result<PrunableTree<MerkleHashVote>, KvError> {
415        let key = cap_key();
416        let Some(blob) = self.cb.get(&key)? else {
417            return Ok(Tree::empty());
418        };
419        read_shard_vote(&blob).map_err(|_| KvError::Deserialization)
420    }
421
422    fn put_cap(&mut self, cap: PrunableTree<MerkleHashVote>) -> Result<(), KvError> {
423        let key = cap_key();
424        let blob = write_shard_vote(&cap).map_err(|_| KvError::Serialization)?;
425        self.cb.set(&key, &blob)
426    }
427
428    fn min_checkpoint_id(&self) -> Result<Option<u32>, KvError> {
429        let prefix = [CHECKPOINT_PREFIX];
430        let mut iter = self.cb.iter(&prefix, false);
431        Ok(iter.next().and_then(|(k, _)| {
432            if k.len() >= 5 {
433                Some(u32::from_be_bytes(k[1..5].try_into().unwrap()))
434            } else {
435                None
436            }
437        }))
438    }
439
440    fn max_checkpoint_id(&self) -> Result<Option<u32>, KvError> {
441        let prefix = [CHECKPOINT_PREFIX];
442        let mut iter = self.cb.iter(&prefix, true /* reverse */);
443        Ok(iter.next().and_then(|(k, _)| {
444            if k.len() >= 5 {
445                Some(u32::from_be_bytes(k[1..5].try_into().unwrap()))
446            } else {
447                None
448            }
449        }))
450    }
451
452    fn add_checkpoint(
453        &mut self,
454        checkpoint_id: u32,
455        checkpoint: Checkpoint,
456    ) -> Result<(), KvError> {
457        let key = checkpoint_key(checkpoint_id);
458        let blob = write_checkpoint(&checkpoint);
459        self.cb.set(&key, &blob)
460    }
461
462    fn checkpoint_count(&self) -> Result<usize, KvError> {
463        let prefix = [CHECKPOINT_PREFIX];
464        let mut iter = self.cb.iter(&prefix, false);
465        let mut count = 0usize;
466        while iter.next().is_some() {
467            count += 1;
468        }
469        Ok(count)
470    }
471
472    fn get_checkpoint_at_depth(
473        &self,
474        checkpoint_depth: usize,
475    ) -> Result<Option<(u32, Checkpoint)>, KvError> {
476        let prefix = [CHECKPOINT_PREFIX];
477        let mut iter = self.cb.iter(&prefix, true /* reverse */);
478        let mut seen = 0usize;
479        while let Some((key, val)) = iter.next() {
480            if seen == checkpoint_depth {
481                if key.len() < 5 {
482                    return Ok(None);
483                }
484                let id = u32::from_be_bytes(key[1..5].try_into().unwrap());
485                return Ok(read_checkpoint(&val).ok().map(|cp| (id, cp)));
486            }
487            seen += 1;
488        }
489        Ok(None)
490    }
491
492    fn get_checkpoint(&self, checkpoint_id: &u32) -> Result<Option<Checkpoint>, KvError> {
493        let key = checkpoint_key(*checkpoint_id);
494        let Some(blob) = self.cb.get(&key)? else {
495            return Ok(None);
496        };
497        Ok(read_checkpoint(&blob).ok())
498    }
499
500    fn with_checkpoints<F>(&mut self, limit: usize, mut callback: F) -> Result<(), KvError>
501    where
502        F: FnMut(&u32, &Checkpoint) -> Result<(), KvError>,
503    {
504        let prefix = [CHECKPOINT_PREFIX];
505        let mut iter = self.cb.iter(&prefix, false);
506        let mut count = 0usize;
507        while count < limit {
508            let Some((key, val)) = iter.next() else {
509                break;
510            };
511            if key.len() < 5 {
512                continue;
513            }
514            let id = u32::from_be_bytes(key[1..5].try_into().unwrap());
515            if let Ok(cp) = read_checkpoint(&val) {
516                callback(&id, &cp)?;
517            }
518            count += 1;
519        }
520        Ok(())
521    }
522
523    fn for_each_checkpoint<F>(&self, limit: usize, mut callback: F) -> Result<(), KvError>
524    where
525        F: FnMut(&u32, &Checkpoint) -> Result<(), KvError>,
526    {
527        let prefix = [CHECKPOINT_PREFIX];
528        let mut iter = self.cb.iter(&prefix, false);
529        let mut count = 0usize;
530        while count < limit {
531            let Some((key, val)) = iter.next() else {
532                break;
533            };
534            if key.len() < 5 {
535                continue;
536            }
537            let id = u32::from_be_bytes(key[1..5].try_into().unwrap());
538            if let Ok(cp) = read_checkpoint(&val) {
539                callback(&id, &cp)?;
540            }
541            count += 1;
542        }
543        Ok(())
544    }
545
546    fn update_checkpoint_with<F>(&mut self, checkpoint_id: &u32, update: F) -> Result<bool, KvError>
547    where
548        F: Fn(&mut Checkpoint) -> Result<(), KvError>,
549    {
550        let key = checkpoint_key(*checkpoint_id);
551        let Some(blob) = self.cb.get(&key)? else {
552            return Ok(false);
553        };
554        let Ok(mut cp) = read_checkpoint(&blob) else {
555            return Ok(false);
556        };
557        update(&mut cp)?;
558        let new_blob = write_checkpoint(&cp);
559        self.cb.set(&key, &new_blob)?;
560        Ok(true)
561    }
562
563    fn remove_checkpoint(&mut self, checkpoint_id: &u32) -> Result<(), KvError> {
564        let key = checkpoint_key(*checkpoint_id);
565        self.cb.delete(&key)
566    }
567
568    fn add_retained_checkpoint(&mut self, checkpoint_id: u32) -> Result<(), KvError> {
569        let key = retained_checkpoint_key(checkpoint_id);
570        self.cb.set(&key, &[])
571    }
572
573    fn remove_retained_checkpoint(&mut self, checkpoint_id: &u32) -> Result<(), KvError> {
574        let key = retained_checkpoint_key(*checkpoint_id);
575        self.cb.delete(&key)
576    }
577
578    fn retained_checkpoints(&self) -> Result<BTreeSet<u32>, KvError> {
579        let prefix = [RETAINED_CHECKPOINT_PREFIX];
580        let mut iter = self.cb.iter(&prefix, false);
581        let mut checkpoints = BTreeSet::new();
582        while let Some((key, _)) = iter.next() {
583            if key.len() < 5 {
584                continue;
585            }
586            checkpoints.insert(u32::from_be_bytes(key[1..5].try_into().unwrap()));
587        }
588        Ok(checkpoints)
589    }
590
591    fn truncate_checkpoints_retaining(&mut self, checkpoint_id: &u32) -> Result<(), KvError> {
592        // Delete all checkpoints with id < checkpoint_id; clear marks_removed
593        // on the retained checkpoint itself (matches MemoryShardStore semantics).
594        let prefix = [CHECKPOINT_PREFIX];
595        let mut iter = self.cb.iter(&prefix, false);
596        let mut to_delete = Vec::new();
597        while let Some((key, _)) = iter.next() {
598            if key.len() < 5 {
599                continue;
600            }
601            let id = u32::from_be_bytes(key[1..5].try_into().unwrap());
602            if id < *checkpoint_id {
603                to_delete.push(key);
604            } else {
605                break;
606            }
607        }
608        drop(iter);
609        for key in to_delete {
610            self.cb.delete(&key)?;
611        }
612        // Clear marks_removed on the retaining checkpoint.
613        let retain_key = checkpoint_key(*checkpoint_id);
614        if let Some(blob) = self.cb.get(&retain_key)? {
615            if let Ok(cp) = read_checkpoint(&blob) {
616                let cleared = Checkpoint::from_parts(cp.tree_state(), BTreeSet::new());
617                self.cb.set(&retain_key, &write_checkpoint(&cleared))?;
618            }
619        }
620        Ok(())
621    }
622}