use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u32);
pub type PageKey = Vec<u32>;
pub fn page_key(input_ids: &[u32], page_size: usize) -> PageKey {
input_ids[..page_size.min(input_ids.len())].to_vec()
}
pub fn align_down(a: usize, b: usize) -> usize {
(a / b) * b
}
pub fn align_ceil(a: usize, b: usize) -> usize {
a.div_ceil(b) * b
}
pub fn match_len(a: &[u32], b: &[u32]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
#[derive(Debug, Clone)]
pub struct Node {
pub key: Vec<u32>,
pub value: Vec<u32>,
pub parent: Option<NodeId>,
pub children: HashMap<PageKey, NodeId>,
pub ref_count: u32,
pub timestamp: i64,
pub mamba_value: Option<u32>,
pub mamba_ref_count: u32,
pub swa_tombstone: bool,
pub swa_ref_count: u32,
pub swa_uuid: Option<u64>,
}
impl Node {
fn new(timestamp: i64) -> Self {
Node {
key: Vec::new(),
value: Vec::new(),
parent: None,
children: HashMap::new(),
ref_count: 0,
timestamp,
mamba_value: None,
mamba_ref_count: 0,
swa_tombstone: false,
swa_ref_count: 0,
swa_uuid: None,
}
}
pub fn length(&self) -> usize {
self.key.len()
}
pub fn is_leaf(&self) -> bool {
self.children.is_empty()
}
pub fn is_root(&self) -> bool {
self.parent.is_none()
}
}
#[derive(Debug)]
pub struct RadixTree {
nodes: Vec<Node>,
page_size: usize,
uuid: u64,
}
pub const ROOT: NodeId = NodeId(0);
impl RadixTree {
pub fn new(page_size: usize) -> Self {
assert!(page_size > 0, "page_size must be positive");
let mut root = Node::new(0);
root.ref_count = 1;
RadixTree {
nodes: vec![root],
page_size,
uuid: 0,
}
}
pub fn page_size(&self) -> usize {
self.page_size
}
pub fn node(&self, id: NodeId) -> &Node {
&self.nodes[id.0 as usize]
}
pub fn node_mut(&mut self, id: NodeId) -> &mut Node {
&mut self.nodes[id.0 as usize]
}
pub fn parent(&self, id: NodeId) -> Option<NodeId> {
self.node(id).parent
}
pub fn alloc(&mut self, timestamp: i64) -> NodeId {
self.nodes.push(Node::new(timestamp));
NodeId((self.nodes.len() - 1) as u32)
}
pub fn next_uuid(&mut self) -> u64 {
self.uuid += 1;
self.uuid
}
pub fn set_key_value(&mut self, id: NodeId, key: Vec<u32>, value: Vec<u32>) {
assert_eq!(
key.len(),
value.len(),
"a node must hold one page index per token"
);
let node = self.node_mut(id);
node.key = key;
node.value = value;
}
pub fn set_parent(&mut self, id: NodeId, parent: NodeId) {
let key = page_key(&self.node(id).key, self.page_size);
self.node_mut(id).parent = Some(parent);
self.node_mut(parent).children.insert(key, id);
}
pub fn child(&self, id: NodeId, input_ids: &[u32]) -> Option<NodeId> {
self.node(id)
.children
.get(&page_key(input_ids, self.page_size))
.copied()
}
pub fn unlink(&mut self, id: NodeId) -> NodeId {
let parent = self.node(id).parent.expect("the root is never unlinked");
let key = page_key(&self.node(id).key, self.page_size);
self.node_mut(parent).children.remove(&key);
self.node_mut(id).parent = None;
parent
}
pub fn split_at(&mut self, id: NodeId, pos: usize) -> NodeId {
let length = self.node(id).length();
assert!(
pos > 0 && pos < length,
"split_at({pos}) is not interior to a node of length {length}"
);
let parent = self.node(id).parent.expect("the root is never split");
let timestamp = self.node(id).timestamp;
let prefix = self.alloc(timestamp);
let (head_key, tail_key) = {
let key = &self.node(id).key;
(key[..pos].to_vec(), key[pos..].to_vec())
};
let (head_value, tail_value) = {
let value = &self.node(id).value;
(value[..pos].to_vec(), value[pos..].to_vec())
};
self.set_key_value(prefix, head_key, head_value);
self.set_parent(prefix, parent);
{
let original = self.node(id).clone();
let head = self.node_mut(prefix);
head.ref_count = original.ref_count;
head.swa_ref_count = original.swa_ref_count;
head.swa_tombstone = original.swa_tombstone;
head.swa_uuid = original.swa_uuid;
}
self.node_mut(id).swa_uuid = None;
self.set_key_value(id, tail_key, tail_value);
self.set_parent(id, prefix);
prefix
}
pub fn path_value(&self, id: NodeId) -> Vec<u32> {
let mut chunks: Vec<&[u32]> = Vec::new();
let mut cur = id;
while !self.node(cur).is_root() {
chunks.push(&self.node(cur).value);
cur = self.node(cur).parent.expect("non-root has a parent");
}
chunks.reverse();
chunks.concat()
}
pub fn path_len(&self, id: NodeId) -> usize {
let mut total = 0;
let mut cur = id;
while !self.node(cur).is_root() {
total += self.node(cur).length();
cur = self.node(cur).parent.expect("non-root has a parent");
}
total
}
pub fn walk(&self) -> Vec<NodeId> {
let mut out = Vec::new();
let mut stack = vec![ROOT];
while let Some(id) = stack.pop() {
for child in self.node(id).children.values() {
stack.push(*child);
out.push(*child);
}
}
out
}
pub fn leaves(&self) -> Vec<NodeId> {
self.walk()
.into_iter()
.filter(|id| self.node(*id).is_leaf())
.collect()
}
pub fn check_structure(&self) {
for id in self.walk() {
let node = self.node(id);
let parent = node.parent.expect("a walked node is not the root");
assert_eq!(
self.node(parent)
.children
.get(&page_key(&node.key, self.page_size)),
Some(&id),
"a node must be registered under its own first page key"
);
assert_eq!(node.key.len(), node.value.len(), "one page index per token");
assert_eq!(
node.length() % self.page_size,
0,
"a node must hold whole pages"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tree_with_chain(page_size: usize, key: &[u32], value: &[u32]) -> (RadixTree, NodeId) {
let mut tree = RadixTree::new(page_size);
let id = tree.alloc(1);
tree.set_key_value(id, key.to_vec(), value.to_vec());
tree.set_parent(id, ROOT);
(tree, id)
}
#[test]
fn a_page_key_covers_a_whole_page_not_a_token() {
assert_eq!(page_key(&[1, 7, 7, 2, 3, 3], 4), vec![1, 7, 7, 2]);
assert_eq!(page_key(&[1, 7, 7, 2], 4), vec![1, 7, 7, 2]);
assert_ne!(page_key(&[1, 7, 7, 2], 4), page_key(&[1, 7, 7, 9], 4));
assert_eq!(page_key(&[1, 7], 4), vec![1, 7]);
assert_ne!(page_key(&[1, 7], 4), page_key(&[1, 7, 7, 2], 4));
}
#[test]
fn reuse_is_bounded_to_whole_pages() {
let key: Vec<u32> = (100..108).collect();
assert_eq!(match_len(&key, &key), 8);
assert_eq!(align_down(match_len(&key, &key), 4), 8);
let mut q: Vec<u32> = (100..106).collect();
q.extend([0, 0]);
assert_eq!(match_len(&key, &q), 6);
assert_eq!(align_down(6, 4), 4, "six matched tokens reuse one page");
assert_eq!(match_len(&key, &[100, 101, 102]), 3);
assert_eq!(match_len(&key, &[0, 101, 102]), 0);
}
#[test]
fn split_keeps_the_original_id_as_the_suffix() {
let (mut tree, id) = tree_with_chain(1, &[1, 2, 3, 4], &[10, 11, 12, 13]);
let prefix = tree.split_at(id, 2);
assert_ne!(prefix, id);
assert_eq!(tree.node(prefix).key, vec![1, 2]);
assert_eq!(tree.node(prefix).value, vec![10, 11]);
assert_eq!(tree.node(id).key, vec![3, 4]);
assert_eq!(tree.node(id).value, vec![12, 13]);
assert_eq!(tree.node(id).parent, Some(prefix));
assert_eq!(tree.node(prefix).parent, Some(ROOT));
assert_eq!(tree.path_value(id), vec![10, 11, 12, 13]);
tree.check_structure();
}
#[test]
fn split_copies_locks_and_window_state_to_both_halves() {
let (mut tree, id) = tree_with_chain(1, &[1, 2, 3, 4], &[10, 11, 12, 13]);
{
let node = tree.node_mut(id);
node.ref_count = 2;
node.swa_ref_count = 1;
node.swa_tombstone = true;
node.swa_uuid = Some(77);
node.mamba_value = Some(99);
node.mamba_ref_count = 1;
node.timestamp = 42;
}
let prefix = tree.split_at(id, 2);
assert_eq!(tree.node(prefix).ref_count, 2);
assert_eq!(tree.node(prefix).swa_ref_count, 1);
assert!(tree.node(prefix).swa_tombstone);
assert_eq!(tree.node(prefix).timestamp, 42);
assert_eq!(tree.node(id).ref_count, 2);
assert_eq!(tree.node(prefix).swa_uuid, Some(77));
assert_eq!(tree.node(id).swa_uuid, None);
assert_eq!(tree.node(prefix).mamba_value, None);
assert_eq!(tree.node(prefix).mamba_ref_count, 0);
assert_eq!(tree.node(id).mamba_value, Some(99));
assert_eq!(tree.node(id).mamba_ref_count, 1);
}
#[test]
#[should_panic(expected = "not interior")]
fn splitting_at_an_edge_is_a_bug_not_a_no_op() {
let (mut tree, id) = tree_with_chain(1, &[1, 2], &[10, 11]);
tree.split_at(id, 0);
}
#[test]
fn unlinking_makes_a_node_unreachable_and_returns_its_parent() {
let (mut tree, id) = tree_with_chain(1, &[1, 2], &[10, 11]);
assert_eq!(tree.leaves(), vec![id]);
assert_eq!(tree.unlink(id), ROOT);
assert!(tree.walk().is_empty());
assert!(tree.node(ROOT).is_leaf());
}
}