use crate::common::{PageID, Position};
use crate::slice_reader::SliceReader;
use std::io::{Cursor, Write};
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]) -> DcbResult<usize> {
let mut cursor = Cursor::new(buf);
cursor.write_all(&[1])?;
let klen = self.keys.len() as u16;
cursor.write_all(&klen.to_le_bytes())?;
for k in &self.keys {
let kb = k.as_bytes();
let kb_len = u8::try_from(kb.len()).map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"tracking key too long to serialize (len > 255)",
)
})?;
cursor.write_all(&[kb_len])?;
cursor.write_all(kb)?;
}
for v in &self.values {
cursor.write_all(&v.0.to_le_bytes())?;
}
Ok(cursor.position() as usize)
}
pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
let mut reader = SliceReader::new(slice);
let ver = reader.read_u8()?;
let count = if ver == 0 {
let cnt_u32 = reader.read_u32()?;
u16::try_from(cnt_u32).map_err(|_| {
DcbError::DeserializationError("v0 tracking leaf count exceeds u16".to_string())
})? as usize
} else {
reader.read_u16()? as usize
};
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
let klen = if ver == 0 {
let klen_u32 = reader.read_u32()?;
u8::try_from(klen_u32).map_err(|_| {
DcbError::DeserializationError(
"v0 tracking leaf key length exceeds u8".to_string(),
)
})? as usize
} else {
reader.read_u8()? as usize
};
let k = reader.read_string(klen)?;
keys.push(k);
}
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(reader.read_position()?);
}
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]) -> DcbResult<usize> {
let mut cursor = Cursor::new(buf);
cursor.write_all(&[1])?;
let klen = self.keys.len() as u16;
cursor.write_all(&klen.to_le_bytes())?;
for k in &self.keys {
let kb = k.as_bytes();
let kb_len = u8::try_from(kb.len()).map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"tracking internal key too long to serialize (len > 255)",
)
})?;
cursor.write_all(&[kb_len])?;
cursor.write_all(kb)?;
}
for id in &self.child_ids {
cursor.write_all(&id.0.to_le_bytes())?;
}
Ok(cursor.position() as usize)
}
pub fn from_slice(slice: &[u8]) -> DcbResult<Self> {
let mut reader = SliceReader::new(slice);
let ver = reader.read_u8()?;
if ver == 0 {
return Err(DcbError::DeserializationError(
"unsupported tracking internal version 0".to_string(),
));
}
let count = reader.read_u16()? as usize;
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
let klen = reader.read_u8()? as usize;
let k = reader.read_string(klen)?;
keys.push(k);
}
let child_count = keys.len() + 1;
let mut child_ids = Vec::with_capacity(child_count);
for _ in 0..child_count {
child_ids.push(reader.read_page_id()?);
}
Ok(Self { keys, child_ids })
}
}
#[cfg(test)]
mod tests {
use super::*;
use byteorder::{ByteOrder, LittleEndian};
#[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).unwrap();
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).unwrap();
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).unwrap();
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).unwrap();
buf.truncate(buf.len() - 4); let err = TrackingInternalNode::from_slice(&buf).unwrap_err();
match err {
DcbError::DeserializationError(_) => {}
_ => panic!("expected DeserializationError"),
}
}
}