1use 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#[derive(Debug, Clone, PartialEq, Eq)]
68pub enum KvError {
69 IoError,
71 Deserialization,
73 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
89const 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
123pub 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
141pub 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
150pub type KvDeleteFn = unsafe extern "C" fn(ctx: *mut c_void, key: *const u8, key_len: usize) -> i32;
152
153pub 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
164pub 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
176pub type KvIterFreeFn = unsafe extern "C" fn(iter: *mut c_void);
178
179pub type KvFreeBufFn = unsafe extern "C" fn(ptr: *mut u8, len: usize);
181
182#[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
204unsafe impl Send for KvCallbacks {}
207unsafe impl Sync for KvCallbacks {}
208
209impl KvCallbacks {
214 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), _ => Err(KvError::IoError), }
239 }
240
241 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 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 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 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
313pub struct KvShardStore {
321 pub(crate) cb: KvCallbacks,
322}
323
324impl KvShardStore {
325 pub fn new(cb: KvCallbacks) -> Self {
326 Self { cb }
327 }
328}
329
330impl 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 );
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 );
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 );
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 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 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}