use crate::common::{PageID, Position};
use byteorder::{ByteOrder, LittleEndian};
use umadb_dcb::{DcbError, DcbResult};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TrackingLeafNode {
pub keys: Vec<String>,
pub values: Vec<Position>,
}
impl TrackingLeafNode {
pub fn new() -> Self {
Self {
keys: Vec::new(),
values: Vec::new(),
}
}
pub fn calc_serialized_size(&self) -> usize {
let mut size = 1 + 2; for k in &self.keys {
size += 1 + k.len();
}
size += 8 * self.values.len();
size
}
pub fn serialize_into(&self, buf: &mut [u8]) -> usize {
let needed = self.calc_serialized_size();
assert!(buf.len() >= needed, "buffer too small for TrackingLeafNode");
buf[0] = 1; LittleEndian::write_u16(&mut buf[1..3], self.keys.len() as u16);
let mut off = 3;
for k in &self.keys {
let kb = k.as_bytes();
assert!(
kb.len() <= u8::MAX as usize,
"tracking key too long to serialize (len>{})",
u8::MAX
);
buf[off] = kb.len() as u8;
off += 1;
buf[off..off + kb.len()].copy_from_slice(kb);
off += kb.len();
}
for v in &self.values {
LittleEndian::write_u64(&mut buf[off..off + 8], v.0);
off += 8;
}
needed
}
pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
if slice.is_empty() {
return Err(DcbError::DeserializationError(
"tracking leaf too small".to_string(),
));
}
let ver = slice[0];
let (mut off, count): (usize, usize) = if ver == 0 {
if slice.len() < 5 {
return Err(DcbError::DeserializationError(
"tracking leaf too small (v0)".to_string(),
));
}
let cnt_u32 = LittleEndian::read_u32(&slice[1..5]);
if cnt_u32 > u16::MAX as u32 {
return Err(DcbError::DeserializationError(
"v0 tracking leaf count exceeds u16".to_string(),
));
}
(5, cnt_u32 as usize)
} else {
if slice.len() < 3 {
return Err(DcbError::DeserializationError(
"tracking leaf too small (v1)".to_string(),
));
}
(3, LittleEndian::read_u16(&slice[1..3]) as usize)
};
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
if ver == 0 {
if off + 4 > slice.len() {
return Err(DcbError::DeserializationError(
"tracking leaf truncated klen (v0)".to_string(),
));
}
let klen_u32 = LittleEndian::read_u32(&slice[off..off + 4]);
if klen_u32 > u8::MAX as u32 {
return Err(DcbError::DeserializationError(
"v0 tracking leaf key length exceeds u8".to_string(),
));
}
let klen = klen_u32 as usize;
off += 4;
if off + klen > slice.len() {
return Err(DcbError::DeserializationError(
"tracking leaf truncated key (v0)".to_string(),
));
}
let k = std::str::from_utf8(&slice[off..off + klen])
.map_err(|e| DcbError::DeserializationError(format!("invalid utf8: {e}")))?
.to_string();
off += klen;
keys.push(k);
} else {
if off + 1 > slice.len() {
return Err(DcbError::DeserializationError(
"tracking leaf truncated klen (v1)".to_string(),
));
}
let klen = slice[off] as usize;
off += 1;
if off + klen > slice.len() {
return Err(DcbError::DeserializationError(
"tracking leaf truncated key (v1)".to_string(),
));
}
let k = std::str::from_utf8(&slice[off..off + klen])
.map_err(|e| DcbError::DeserializationError(format!("invalid utf8: {e}")))?
.to_string();
off += klen;
keys.push(k);
}
}
let mut values = Vec::with_capacity(count);
for _ in 0..count {
if off + 8 > slice.len() {
return Err(DcbError::DeserializationError(
"tracking leaf truncated value".to_string(),
));
}
let v = LittleEndian::read_u64(&slice[off..off + 8]);
off += 8;
values.push(Position(v));
}
if ver == 0 {
let mut pairs: Vec<(String, Position)> = keys.into_iter().zip(values).collect();
pairs.sort_by(|a, b| a.0.cmp(&b.0));
let (new_keys, new_vals): (Vec<String>, Vec<Position>) = pairs.into_iter().unzip();
Ok(Self {
keys: new_keys,
values: new_vals,
})
} else {
Ok(Self { keys, values })
}
}
pub fn get(&self, source: &str) -> Option<Position> {
match self.keys.binary_search_by(|k| k.as_str().cmp(source)) {
Ok(i) => Some(self.values[i]),
Err(_) => None,
}
}
pub fn upsert_no_split(
&mut self,
source: &str,
pos: Position,
page_body_capacity: usize,
) -> DcbResult<()> {
match self.keys.binary_search_by(|k| k.as_str().cmp(source)) {
Ok(idx) => {
self.values[idx] = pos;
Ok(())
}
Err(ins_idx) => {
let key_len = source.len();
let additional = key_len + 8;
let current = self.calc_serialized_size();
if current + additional > page_body_capacity {
return Err(DcbError::InternalError(
"not implemented: tracking split".to_string(),
));
}
self.keys.insert(ins_idx, source.to_string());
self.values.insert(ins_idx, pos);
Ok(())
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TrackingInternalNode {
pub keys: Vec<String>,
pub child_ids: Vec<PageID>,
}
impl TrackingInternalNode {
pub fn child_index_for_key(&self, key: &str) -> usize {
match self.keys.binary_search_by(|k| k.as_str().cmp(key)) {
Ok(idx) => idx + 1,
Err(idx) => idx,
}
}
pub fn replace_child_id_at(
&mut self,
idx: usize,
old_id: PageID,
new_id: PageID,
) -> DcbResult<()> {
if idx >= self.child_ids.len() {
return Err(DcbError::DatabaseCorrupted(
"child index out of bounds".to_string(),
));
}
if self.child_ids[idx] != old_id {
return Err(DcbError::DatabaseCorrupted("Child ID mismatch".to_string()));
}
self.child_ids[idx] = new_id;
Ok(())
}
pub fn insert_promoted_at(&mut self, idx: usize, key: String, right_child: PageID) {
self.keys.insert(idx, key);
self.child_ids.insert(idx + 1, right_child);
}
pub fn split_off(&mut self) -> DcbResult<(String, Vec<String>, Vec<PageID>)> {
if self.child_ids.len() < 4 || self.keys.len() + 1 != self.child_ids.len() {
return Err(DcbError::DatabaseCorrupted(
"Cannot split tracking internal with insufficient arity".to_string(),
));
}
let mid = self.keys.len() / 2; let promoted_key = self.keys[mid].clone();
let right_keys: Vec<String> = self.keys[mid + 1..].to_vec();
let right_child_ids: Vec<PageID> = self.child_ids[mid + 1..].to_vec();
self.keys.truncate(mid);
self.child_ids.truncate(mid + 1);
Ok((promoted_key, right_keys, right_child_ids))
}
pub fn new() -> Self {
Self {
keys: Vec::new(),
child_ids: Vec::new(),
}
}
pub fn calc_serialized_size(&self) -> usize {
let mut size = 1 + 2; for k in &self.keys {
size += 1 + k.len();
}
size += 8 * (self.keys.len() + 1);
size
}
pub fn serialize_into(&self, buf: &mut [u8]) -> usize {
let needed = self.calc_serialized_size();
assert!(
buf.len() >= needed,
"buffer too small for TrackingInternalNode"
);
buf[0] = 1; LittleEndian::write_u16(&mut buf[1..3], self.keys.len() as u16);
let mut off = 3;
for k in &self.keys {
let kb = k.as_bytes();
assert!(
kb.len() <= u8::MAX as usize,
"tracking internal key too long to serialize (len>{})",
u8::MAX
);
buf[off] = kb.len() as u8;
off += 1;
buf[off..off + kb.len()].copy_from_slice(kb);
off += kb.len();
}
for id in &self.child_ids {
LittleEndian::write_u64(&mut buf[off..off + 8], id.0);
off += 8;
}
needed
}
pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
if slice.is_empty() {
return Err(DcbError::DeserializationError(
"tracking internal too small".to_string(),
));
}
let ver = slice[0];
if ver == 0 {
return Err(DcbError::DeserializationError(
"unsupported tracking internal version 0".to_string(),
));
}
if slice.len() < 3 {
return Err(DcbError::DeserializationError(
"tracking internal too small (v1)".to_string(),
));
}
let count = LittleEndian::read_u16(&slice[1..3]) as usize;
let mut off = 3;
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
if off + 1 > slice.len() {
return Err(DcbError::DeserializationError(
"tracking internal truncated klen (v1)".to_string(),
));
}
let klen = slice[off] as usize;
off += 1;
if off + klen > slice.len() {
return Err(DcbError::DeserializationError(
"tracking internal truncated key (v1)".to_string(),
));
}
let k = std::str::from_utf8(&slice[off..off + klen])
.map_err(|e| DcbError::DeserializationError(format!("invalid utf8: {e}")))?
.to_string();
off += klen;
keys.push(k);
}
let child_count = keys.len() + 1;
let need = off + 8 * child_count;
if slice.len() < need {
return Err(DcbError::DeserializationError(
"tracking internal truncated children".to_string(),
));
}
let mut child_ids = Vec::with_capacity(child_count);
for _ in 0..child_count {
let v = LittleEndian::read_u64(&slice[off..off + 8]);
off += 8;
child_ids.push(PageID(v));
}
Ok(Self { keys, child_ids })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tracking_leaf_roundtrip() {
let mut node = TrackingLeafNode::new();
node.keys = vec!["a".to_string(), "b".to_string()];
node.values = vec![Position(1), Position(2)];
let mut buf = vec![0u8; node.calc_serialized_size()];
let n = node.serialize_into(&mut buf);
assert_eq!(n, buf.len());
let dec = TrackingLeafNode::from_slice(&buf).unwrap();
assert_eq!(node, dec);
assert_eq!(dec.get("a"), Some(Position(1)));
assert_eq!(dec.get("z"), None);
}
#[test]
fn test_deserialize_v0_unsorted_sorts_and_aligns() {
let keys = vec!["b", "c", "a"];
let values = vec![Position(2), Position(3), Position(1)];
let key_bytes: Vec<Vec<u8>> = keys.iter().map(|k| k.as_bytes().to_vec()).collect();
let mut size = 1 + 4;
for kb in &key_bytes {
size += 4 + kb.len();
}
size += 8 * values.len();
let mut buf = vec![0u8; size];
buf[0] = 0; LittleEndian::write_u32(&mut buf[1..5], keys.len() as u32);
let mut off = 5;
for kb in &key_bytes {
LittleEndian::write_u32(&mut buf[off..off + 4], kb.len() as u32);
off += 4;
buf[off..off + kb.len()].copy_from_slice(kb);
off += kb.len();
}
for v in &values {
LittleEndian::write_u64(&mut buf[off..off + 8], v.0);
off += 8;
}
let dec = TrackingLeafNode::from_slice(&buf).unwrap();
assert_eq!(dec.keys, vec!["a", "b", "c"]);
assert_eq!(dec.values, vec![Position(1), Position(2), Position(3)]);
assert_eq!(dec.get("a"), Some(Position(1)));
assert_eq!(dec.get("c"), Some(Position(3)));
assert_eq!(dec.get("z"), None);
}
#[test]
fn test_upsert_maintains_sorted_and_capacity_check() {
let mut node = TrackingLeafNode::new();
let capacity = 1 + 4 + (4 + 1) + (4 + 1) + (4 + 1) + 8 * 3; node.upsert_no_split("b", Position(2), capacity).unwrap();
node.upsert_no_split("a", Position(1), capacity).unwrap();
node.upsert_no_split("c", Position(3), capacity).unwrap();
assert_eq!(
node.keys,
vec!["a".to_string(), "b".to_string(), "c".to_string()]
);
assert_eq!(node.get("b"), Some(Position(2)));
let mut node2 = TrackingLeafNode::new();
let small_capacity = 1 + 4; let err = node2
.upsert_no_split("x", Position(9), small_capacity)
.unwrap_err();
match err {
DcbError::InternalError(s) => assert!(s.contains("tracking split")),
_ => panic!("unexpected error type"),
}
}
#[test]
fn test_tracking_internal_roundtrip() {
let node = TrackingInternalNode {
keys: vec!["alpha".into(), "beta".into(), "gamma".into()],
child_ids: vec![PageID(10), PageID(20), PageID(30), PageID(40)],
};
let mut buf = vec![0u8; node.calc_serialized_size()];
let n = node.serialize_into(&mut buf);
assert_eq!(n, buf.len());
let dec = TrackingInternalNode::from_slice(&buf).unwrap();
assert_eq!(dec, node);
}
#[test]
fn test_tracking_internal_empty_keys_one_child_roundtrip() {
let node = TrackingInternalNode {
keys: vec![],
child_ids: vec![PageID(123)],
};
let mut buf = vec![0u8; node.calc_serialized_size()];
let n = node.serialize_into(&mut buf);
assert_eq!(n, buf.len());
let dec = TrackingInternalNode::from_slice(&buf).unwrap();
assert_eq!(dec, node);
}
#[test]
fn test_tracking_internal_from_slice_truncated_children_err() {
let node = TrackingInternalNode {
keys: vec!["k1".into()],
child_ids: vec![PageID(1), PageID(2)],
};
let mut buf = vec![0u8; node.calc_serialized_size()];
let _ = node.serialize_into(&mut buf);
buf.truncate(buf.len() - 4); let err = TrackingInternalNode::from_slice(&buf).unwrap_err();
match err {
DcbError::DeserializationError(_) => {}
_ => panic!("expected DeserializationError"),
}
}
}