use std::{
cmp::{Ord, Ordering},
default::Default,
fmt,
fmt::{Debug, Formatter},
ptr::NonNull,
};
#[derive(Clone, PartialEq)]
pub enum BTree<K: Ord + Clone, V, const Q: usize> {
Internal(NonNull<InternalNode<K, V, Q>>),
Leaf(NonNull<LeafNode<K, V, Q>>),
}
#[derive(Clone, PartialEq)]
pub struct InternalNode<K: Ord + Clone, V, const Q: usize> {
keys: Vec<K>,
children: Vec<BTree<K, V, Q>>,
}
#[derive(Clone, PartialEq)]
pub struct LeafNode<K: Ord + Clone, V, const Q: usize> {
keys: Vec<K>,
values: Vec<V>,
next: Option<NonNull<LeafNode<K, V, Q>>>,
prev: Option<NonNull<LeafNode<K, V, Q>>>,
}
impl<K, V, const Q: usize> Debug for BTree<K, V, Q>
where
K: Debug + Ord + Clone,
V: Debug,
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
match self {
BTree::Internal(ptr) => {
let node = unsafe { ptr.as_ref() };
f.debug_struct("InternalNode")
.field("entries", node)
.finish()
}
BTree::Leaf(ptr) => {
let node = unsafe { ptr.as_ref() };
let prev_key = unsafe { node.prev.map(|ptr| &ptr.as_ref().keys[0]) };
let next_key= unsafe { node.next.map(|ptr| &ptr.as_ref().keys[0]) };
f.debug_struct("LeafNode")
.field("entries", node)
.field("next_0th_key", &next_key)
.field("prev_0th_key", &prev_key)
.finish()
}
}
}
}
impl<K, V, const Q: usize> Debug for InternalNode<K, V, Q>
where
K: Debug + Ord + Clone,
V: Debug,
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let mut m = &mut f.debug_map();
for i in 0..self.children.len() {
if i == 0 {
m = m.entry(&"null", &self.children[i]);
} else {
m = m.entry(&self.keys[i - 1], &self.children[i])
}
}
m.finish()
}
}
impl<K, V, const Q: usize> Debug for LeafNode<K, V, Q>
where
K: Debug + Ord + Clone,
V: Debug,
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let mut m = &mut f.debug_map();
for i in 0..self.keys.len() {
m = m.entry(&self.keys[i], &self.values[i]);
}
m.finish()
}
}
impl<K, V, const Q: usize> Default for BTree<K, V, Q>
where
K: Ord + Clone,
{
fn default() -> Self {
assert!(Q > 2, "branching factor Q must be greater than 2");
let node = unsafe { NonNull::new_unchecked(Box::into_raw(Box::new(LeafNode::default()))) };
BTree::Leaf(node)
}
}
impl<K, V, const Q: usize> Default for LeafNode<K, V, Q>
where
K: Ord + Clone,
{
fn default() -> Self {
assert!(Q > 2, "branching factor Q must be greater than 2");
LeafNode {
keys: Vec::with_capacity(Q),
values: Vec::with_capacity(Q),
next: None,
prev: None,
}
}
}
impl<K, V, const Q: usize> InternalNode<K, V, Q>
where
K: Ord + Clone + Debug,
{
fn new(child: BTree<K, V, Q>) -> Self {
assert!(Q > 2, "branching factor Q must be greater than 2");
let mut children = Vec::with_capacity(Q);
children.push(child);
InternalNode {
keys: Vec::with_capacity(Q),
children,
}
}
fn insert(&mut self, key: K, child: BTree<K, V, Q>) {
let idx = self.keys.binary_search(&key).unwrap_or_else(|idx| idx);
self.keys.insert(idx, key);
self.children.insert(idx + 1, child);
}
}
impl<K, V, const Q: usize> LeafNode<K, V, Q>
where
K: Ord + Clone,
{
fn insert(&mut self, key: K, value: V) {
let idx = self.keys.binary_search(&key).unwrap_or_else(|idx| idx);
self.keys.insert(idx, key);
self.values.insert(idx, value);
}
}
impl<K, V, const Q: usize> BTree<K, V, Q>
where
K: Ord + Clone + Debug,
V: Debug
{
fn insert_inner(&mut self, key: K, value: V) -> Option<(K, BTree<K, V, Q>)> {
match self {
BTree::Leaf(ref mut _leaf) => {
let leaf = unsafe { _leaf.as_mut() };
leaf.insert(key, value);
if leaf.keys.len() == Q {
let mid = leaf.keys.len() / 2;
let mut right: LeafNode<K, V, Q> = Default::default();
right.keys = leaf.keys.split_off(mid);
right.values = leaf.values.split_off(mid);
let right = Box::new(right);
let split_key = right.keys[0].clone();
let mut right = unsafe { NonNull::new_unchecked(Box::into_raw(right)) };
unsafe { right.as_mut().prev = Some(*_leaf) };
leaf.next = Some(right);
Some((split_key, BTree::Leaf(right)))
} else {
None
}
}
BTree::Internal(ref mut node) => {
let node = unsafe { node.as_mut() };
let idx = node.keys.binary_search(&key).unwrap_or_else(|idx| idx);
if let Some((split_key, child)) = node.children[idx].insert_inner(key, value) {
node.insert(split_key, child);
if node.keys.len() == Q {
let mid = node.keys.len() / 2;
let right_keys = node.keys.split_off(mid + 1);
let right_children = node.children.split_off(mid + 1);
let split_key = node.keys.pop().unwrap();
let split_child = node.children.pop().unwrap();
let mut right = InternalNode::new(split_child);
right.keys.extend(right_keys);
right.children.extend(right_children);
let right = Box::new(right);
let right = unsafe { NonNull::new_unchecked(Box::into_raw(right)) };
Some((split_key, BTree::Internal(right)))
} else {
None
}
} else {
None
}
}
}
}
pub fn insert(&mut self, key: K, value: V) {
if let Some((right_key, right)) = self.insert_inner(key, value) {
let new_root = InternalNode::new(right);
let new_root = Box::new(new_root);
let new_root = unsafe { NonNull::new_unchecked(Box::into_raw(new_root)) };
let left = std::mem::replace(self, BTree::Internal(new_root));
match self {
BTree::Internal(ref mut node) => {
let node = unsafe { node.as_mut() };
node.keys.push(right_key);
node.children.insert(0, left);
}
BTree::Leaf(_) => panic!("expected root to be Internal node!"),
}
}
}
pub fn get(&self, key: &K) -> Option<&V> {
match self {
BTree::Internal(ref node) => {
let idx = unsafe { node.as_ref() }
.keys
.binary_search(key)
.unwrap_or_else(|idx| idx);
unsafe { node.as_ref() }.children[idx].get(key)
}
BTree::Leaf(ref node) => {
let idx = unsafe { node.as_ref() }.keys.binary_search(key).ok();
match idx {
Some(idx) => Some(&(unsafe { node.as_ref() }.values[idx])),
None => None,
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::fmt::Debug;
fn is_sorted<T: Ord>(items: &Vec<T>) -> bool {
let (is_sorted, _) = items.iter().fold((true, None), |(is_sorted, prev), curr| {
if let Some(prev) = prev {
(is_sorted && prev <= curr, Some(curr))
} else {
(is_sorted, Some(curr))
}
});
is_sorted
}
fn assert_at_node<K: Debug>(
cond: bool,
left_parent_key: Option<&K>,
level: usize,
msg: String,
) {
assert!(
cond,
"In node with parent key {:#?} at level {}: {}",
left_parent_key, level, msg
);
}
fn assert_is_b_tree_inner<K: Ord + Clone + Debug, V, const Q: usize>(
node: &BTree<K, V, Q>,
parent_keys: (Option<&K>, Option<&K>),
level: usize,
) {
let (left_parent_key, right_parent_key) = parent_keys;
match node {
BTree::Internal(ptr) => {
let node = unsafe { ptr.as_ref() };
assert_at_node(
is_sorted(&node.keys),
left_parent_key,
level,
"keys are not sorted".to_string(),
);
node.keys
.iter()
.for_each(|key| match (left_parent_key, right_parent_key) {
(Some(left), Some(right)) => assert_at_node(
key >= left && key < right,
left_parent_key,
level,
format!(
"key {:#?} < left parent key {:?} or >= right parent key {:?}",
key, left, right
),
),
(Some(left), None) => assert_at_node(
key >= left,
left_parent_key,
level,
format!("key {:?} < left parent key {:?}", key, left),
),
(None, Some(right)) => assert_at_node(
key < right,
left_parent_key,
level,
format!("key {:?} >= right parent key {:?}", key, right),
),
(None, None) => assert_at_node(
level == 0,
left_parent_key,
level,
"(None, None) case of parent keys for non-root node encountered!"
.to_string(),
),
});
assert!(node.children.len() > 1);
node.children.iter().enumerate().for_each(|(i, child)| {
if i == 0 {
let right = &node.keys[i];
assert_is_b_tree_inner(child, (None, Some(right)), level + 1);
} else if i == node.keys.len() {
let left = &node.keys[i - 1];
assert_is_b_tree_inner(child, (Some(left), None), level + 1);
} else {
let left = &node.keys[i - 1];
let right = &node.keys[i];
assert_is_b_tree_inner(child, (Some(left), Some(right)), level + 1);
}
});
}
BTree::Leaf(ptr) => {
let node = unsafe { ptr.as_ref() };
assert_at_node(
is_sorted(&node.keys),
left_parent_key,
level,
"keys are not sorted".to_string(),
);
if left_parent_key.is_none() && right_parent_key.is_none() {
assert_at_node(
node.keys.len() < Q,
left_parent_key,
level,
format!("root node that's leaf has {} > Q keys", node.keys.len()),
);
return;
}
assert_at_node(
node.keys.len() >= Q / 2,
left_parent_key,
level,
format!("leaf node has {} < Q / 2 keys", node.keys.len()),
);
let (left_leaf_last_key, is_lte) = node.prev.map_or((None, true), |ptr| {
let left_leaf = unsafe { ptr.as_ref() };
if let Some(left_leaf_last_key) = left_leaf.keys.last() {
(
Some(left_leaf_last_key),
left_leaf_last_key <= node.keys.first().unwrap(),
)
} else {
(None, true)
}
});
assert_at_node(
is_lte,
left_parent_key,
level,
format!(
"last key {:?} of left leaf node > {:?}, the first key of current leaf node",
node.keys.first().unwrap(),
left_leaf_last_key
),
);
let (right_leaf_first_key, is_gte) = node.next.map_or((None, true), |ptr| {
let right_leaf = unsafe { ptr.as_ref() };
if let Some(right_leaf_first_key) = right_leaf.keys.first() {
(
Some(right_leaf_first_key),
right_leaf_first_key >= node.keys.last().unwrap(),
)
} else {
(None, true)
}
});
assert_at_node(
is_gte,
left_parent_key,
level,
format!(
"first key {:?} of right leaf node < {:?}, the last key of current leaf node",
node.keys.last().unwrap(),
right_leaf_first_key
),
)
}
}
}
fn assert_is_b_tree<K: Ord + Clone + Debug, V: Debug, const Q: usize>(root: &BTree<K, V, Q>) {
println!("{:#?}", root);
assert_is_b_tree_inner(root, (None, None), 0)
}
#[test]
fn test_get_insert_basic() {
let mut tree: BTree<i32, String, 4> = Default::default();
let vals = vec![
(-5, "hi"),
(2, "howdy"),
(-1, "yo"),
(6, "salutations"),
(-99, "greetings"),
(30, "wilkommen"),
(99, "ohayou"),
(-400, "nihao"),
];
for &(k, v) in vals.iter() {
tree.insert(k, v.to_string());
assert_is_b_tree(&tree);
}
assert_eq!(tree.get(&-99), Some(&"greetings".to_string()));
assert_eq!(tree.get(&99), Some(&"ohayou".to_string()));
assert!(tree.get(&0).is_none());
assert!(tree.get(&7).is_none());
}
}