use super::CompactArc;
use std::marker::PhantomData;
use std::mem;
use std::ops::Bound;
use std::ptr;
use std::sync::atomic::{AtomicU32, Ordering};
const MAX_KEYS: usize = 128;
const _: () = assert!(
MAX_KEYS <= 255,
"MAX_KEYS must be <= 255 (NodePath uses u8 indices)"
);
const MIN_KEYS: usize = (MAX_KEYS - 1) / 2;
#[repr(C)]
struct NodeHeader {
len: u16,
is_leaf: u8,
_pad1: u8,
drop_count: AtomicU32,
}
pub struct NodePtr<V: Clone> {
ptr: CompactArc<[u8]>,
_marker: PhantomData<V>,
}
impl<V: Clone> Clone for NodePtr<V> {
fn clone(&self) -> Self {
let new_ptr = self.ptr.clone();
let header = unsafe { &*(self.ptr.data_ptr_mut() as *const NodeHeader) };
header.drop_count.fetch_add(1, Ordering::Relaxed);
Self {
ptr: new_ptr,
_marker: PhantomData,
}
}
}
impl<V: Clone> Drop for NodePtr<V> {
fn drop(&mut self) {
let header = unsafe { &*(self.ptr.data_ptr_mut() as *const NodeHeader) };
let old_count = header.drop_count.fetch_sub(1, Ordering::AcqRel);
if old_count != 1 {
return;
}
let len = self.len();
if self.is_leaf() {
unsafe {
let v_ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
for i in 0..len {
ptr::drop_in_place(v_ptr.add(i));
}
}
} else {
unsafe {
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
for i in 0..=len {
ptr::drop_in_place(c_ptr.add(i));
}
}
}
}
}
impl<V: Clone> NodePtr<V> {
const fn keys_offset() -> usize {
mem::size_of::<NodeHeader>()
}
const fn values_offset() -> usize {
Self::keys_offset() + ((MAX_KEYS + 1) * 8)
}
const fn children_offset() -> usize {
Self::keys_offset() + ((MAX_KEYS + 1) * 8)
}
fn new_leaf() -> Self {
let size = Self::values_offset() + ((MAX_KEYS + 1) * mem::size_of::<V>());
let vec = vec![0u8; size];
let mut ptr = NodePtr {
ptr: CompactArc::from_vec(vec),
_marker: PhantomData,
};
let header = ptr.header_mut();
header.len = 0;
header.is_leaf = 1;
header.drop_count = AtomicU32::new(1);
ptr
}
fn new_internal() -> Self {
let size = Self::children_offset() + ((MAX_KEYS + 2) * mem::size_of::<NodePtr<V>>());
let vec = vec![0u8; size];
let mut ptr = NodePtr {
ptr: CompactArc::from_vec(vec),
_marker: PhantomData,
};
let header = ptr.header_mut();
header.len = 0;
header.is_leaf = 0;
header.drop_count = AtomicU32::new(1);
ptr
}
fn make_mut(&mut self) -> &mut Self {
let header = unsafe { &*(self.ptr.data_ptr_mut() as *const NodeHeader) };
if header.drop_count.load(Ordering::Acquire) != 1 {
*self = self.deep_clone();
}
self
}
fn deep_clone(&self) -> Self {
let len = self.len();
if self.is_leaf() {
let mut new_node = NodePtr::new_leaf();
unsafe {
let k_src = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *const i64;
let k_dst = new_node.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy_nonoverlapping(k_src, k_dst, len);
}
unsafe {
let v_src = self.ptr.data_ptr_mut().add(Self::values_offset()) as *const V;
let v_dst = new_node.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
for i in 0..len {
ptr::write(v_dst.add(i), (*v_src.add(i)).clone());
new_node.set_len(i + 1);
}
}
new_node
} else {
let mut new_node = NodePtr::new_internal();
new_node.set_len(len);
unsafe {
let k_src = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *const i64;
let k_dst = new_node.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy_nonoverlapping(k_src, k_dst, len);
}
unsafe {
let c_src =
self.ptr.data_ptr_mut().add(Self::children_offset()) as *const NodePtr<V>;
let c_dst =
new_node.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
for i in 0..=len {
ptr::write(c_dst.add(i), (*c_src.add(i)).clone());
}
}
new_node
}
}
fn header(&self) -> &NodeHeader {
unsafe { &*(self.ptr.data_ptr_mut() as *const NodeHeader) }
}
fn header_mut(&mut self) -> &mut NodeHeader {
unsafe { &mut *(self.ptr.data_ptr_mut() as *mut NodeHeader) }
}
fn is_leaf(&self) -> bool {
self.header().is_leaf == 1
}
fn len(&self) -> usize {
self.header().len as usize
}
fn set_len(&mut self, len: usize) {
self.header_mut().len = len as u16;
}
fn keys(&self) -> &[i64] {
let len = self.len();
unsafe {
let ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *const i64;
std::slice::from_raw_parts(ptr, len)
}
}
fn values(&self) -> &[V] {
assert!(self.is_leaf());
let len = self.len();
unsafe {
let ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *const V;
std::slice::from_raw_parts(ptr, len)
}
}
fn values_mut_slice(&mut self) -> &mut [V] {
assert!(self.is_leaf());
let len = self.len();
unsafe {
let ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
std::slice::from_raw_parts_mut(ptr, len)
}
}
fn children(&self) -> &[NodePtr<V>] {
assert!(!self.is_leaf());
let len = self.len() + 1; unsafe {
let ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *const NodePtr<V>;
std::slice::from_raw_parts(ptr, len)
}
}
fn children_mut_slice(&mut self) -> &mut [NodePtr<V>] {
assert!(!self.is_leaf());
let len = self.len() + 1;
unsafe {
let ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
std::slice::from_raw_parts_mut(ptr, len)
}
}
fn child(&self, index: usize) -> &NodePtr<V> {
&self.children()[index]
}
fn child_mut(&mut self, index: usize) -> &mut NodePtr<V> {
&mut self.children_mut_slice()[index]
}
fn search(&self, key: i64) -> Result<usize, usize> {
self.keys().binary_search(&key)
}
fn push_leaf(&mut self, key: i64, value: V) {
assert!(self.is_leaf());
let len = self.len();
assert!(len <= MAX_KEYS);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(len), key);
let v_ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
ptr::write(v_ptr.add(len), value);
}
self.set_len(len + 1);
}
fn remove_leaf(&mut self, index: usize) -> V {
assert!(self.is_leaf());
let len = self.len();
assert!(index < len);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy(k_ptr.add(index + 1), k_ptr.add(index), len - index - 1);
let v_ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
let val = ptr::read(v_ptr.add(index));
ptr::copy(v_ptr.add(index + 1), v_ptr.add(index), len - index - 1);
self.set_len(len - 1);
val
}
}
fn insert_leaf(&mut self, index: usize, key: i64, value: V) {
assert!(self.is_leaf());
let len = self.len();
assert!(len <= MAX_KEYS);
assert!(
index <= len,
"insert_leaf: index {} > len {} for key {}",
index,
len,
key
);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
let p_key = k_ptr.add(index);
ptr::copy(p_key, p_key.add(1), len - index);
ptr::write(p_key, key);
let v_ptr = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
let p_val = v_ptr.add(index);
ptr::copy(p_val, p_val.add(1), len - index);
ptr::write(p_val, value);
}
self.set_len(len + 1);
}
fn push_internal(&mut self, key: i64, child: NodePtr<V>) {
assert!(!self.is_leaf());
let len = self.len();
assert!(len <= MAX_KEYS);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(len), key);
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::write(c_ptr.add(len + 1), child);
}
self.set_len(len + 1);
}
fn split_internal(&mut self) -> (i64, NodePtr<V>) {
let len = self.len();
let mid = len / 2;
let med_key = self.keys()[mid];
let mut right = NodePtr::new_internal();
let right_keys_count = len - mid - 1;
let right_children_count = right_keys_count + 1;
unsafe {
let k_src = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
let k_dst = right.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy_nonoverlapping(k_src.add(mid + 1), k_dst, right_keys_count);
let c_src = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
let c_dst = right.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::copy_nonoverlapping(c_src.add(mid + 1), c_dst, right_children_count);
}
right.set_len(right_keys_count);
self.set_len(mid);
(med_key, right)
}
fn split_internal_rightmost(&mut self) -> (i64, NodePtr<V>) {
let len = self.len();
debug_assert!(
len == MAX_KEYS + 1,
"split_internal_rightmost expects overflow node"
);
let med_key = self.keys()[len - 1];
let mut right = NodePtr::new_internal();
unsafe {
let c_src = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
let c_dst = right.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::write(c_dst, ptr::read(c_src.add(len)));
}
right.set_len(0); self.set_len(len - 1);
(med_key, right)
}
fn borrow_from_left(&mut self, index: usize) {
assert!(!self.is_leaf());
let is_left_leaf = self.child(index - 1).is_leaf();
if is_left_leaf {
let (key, val) = {
let left = self.child_mut(index - 1).make_mut();
let left_len = left.len();
let key = left.keys()[left_len - 1];
let val = unsafe {
let val_ptr = left.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
ptr::read(val_ptr.add(left_len - 1))
};
left.set_len(left_len - 1);
(key, val)
};
let current = self.child_mut(index).make_mut();
current.insert_leaf(0, key, val);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(index - 1), key);
}
} else {
let (child, key) = {
let left = self.child_mut(index - 1).make_mut();
let left_len = left.len();
let key = left.keys()[left_len - 1];
let child = unsafe {
let c_ptr =
left.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::read(c_ptr.add(left_len))
};
left.set_len(left_len - 1);
(child, key)
};
let separator_idx = index - 1;
let separator = self.keys()[separator_idx];
let current = self.child_mut(index).make_mut();
current.insert_internal_at_start(separator, child);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(separator_idx), key);
}
}
}
fn insert_internal(&mut self, index: usize, key: i64, child: NodePtr<V>) {
assert!(!self.is_leaf());
let len = self.len();
assert!(len <= MAX_KEYS);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy(k_ptr.add(index), k_ptr.add(index + 1), len - index);
ptr::write(k_ptr.add(index), key);
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::copy(c_ptr.add(index + 1), c_ptr.add(index + 2), len - index);
ptr::write(c_ptr.add(index + 1), child);
}
self.set_len(len + 1);
}
fn insert_internal_at_start(&mut self, key: i64, child: NodePtr<V>) {
let len = self.len();
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy(k_ptr, k_ptr.add(1), len);
ptr::write(k_ptr, key);
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::copy(c_ptr, c_ptr.add(1), len + 1);
ptr::write(c_ptr, child);
}
self.set_len(len + 1);
}
fn borrow_from_right(&mut self, index: usize) {
assert!(!self.is_leaf());
let is_right_leaf = self.child(index + 1).is_leaf();
if is_right_leaf {
let (key, val) = {
let right = self.child_mut(index + 1).make_mut();
let key = right.keys()[0];
let val = right.remove_leaf(0);
(key, val)
};
let current = self.child_mut(index).make_mut();
current.push_leaf(key, val);
let right = self.child(index + 1);
let new_sep = right.keys()[0];
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(index), new_sep);
}
} else {
let (child, key) = {
let right = self.child_mut(index + 1).make_mut();
let key = right.keys()[0];
let child = unsafe {
let c_ptr =
right.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::read(c_ptr)
};
right.remove_internal_at_start();
(child, key)
};
let separator = self.keys()[index];
let current = self.child_mut(index).make_mut();
current.push_internal(separator, child);
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::write(k_ptr.add(index), key);
}
}
}
fn remove_internal_at_start(&mut self) {
let len = self.len();
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy(k_ptr.add(1), k_ptr, len - 1);
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::copy(c_ptr.add(1), c_ptr, len);
}
self.set_len(len - 1);
}
fn merge_with_left(&mut self, index: usize) {
let separator = self.keys()[index - 1];
let right_raw = self.child_mut(index) as *mut NodePtr<V>;
let right = unsafe { (*right_raw).make_mut() };
let is_leaf = right.is_leaf();
let r_len = right.len();
if is_leaf {
let mut keys_vals: Vec<(i64, V)> = Vec::with_capacity(r_len);
unsafe {
let k_ptr = right.ptr.data_ptr_mut().add(Self::keys_offset()) as *const i64;
let v_ptr = right.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
for i in 0..r_len {
let key = *k_ptr.add(i);
let val = ptr::read(v_ptr.add(i)); keys_vals.push((key, val));
}
}
right.set_len(0);
let left = self.child_mut(index - 1).make_mut();
for (k, v) in keys_vals {
left.push_leaf(k, v);
}
} else {
let keys: Vec<i64> = right.keys().to_vec();
let children: Vec<NodePtr<V>> = unsafe {
let c_ptr =
right.ptr.data_ptr_mut().add(Self::children_offset()) as *const NodePtr<V>;
(0..=r_len).map(|i| ptr::read(c_ptr.add(i))).collect()
};
unsafe {
let header = &*(right.ptr.data_ptr_mut() as *const NodeHeader);
header.drop_count.store(2, Ordering::Release);
}
right.set_len(0);
let left = self.child_mut(index - 1).make_mut();
let mut children_iter = children.into_iter();
left.push_internal(separator, children_iter.next().unwrap());
for (key, child) in keys.into_iter().zip(children_iter) {
left.push_internal(key, child);
}
}
self.remove_key_and_child(index - 1, index);
}
fn merge_with_right(&mut self, index: usize) {
self.merge_with_left(index + 1)
}
fn remove_key_and_child(&mut self, key_idx: usize, child_idx: usize) {
let len = self.len();
unsafe {
let k_ptr = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy(
k_ptr.add(key_idx + 1),
k_ptr.add(key_idx),
len - key_idx - 1,
);
let c_ptr = self.ptr.data_ptr_mut().add(Self::children_offset()) as *mut NodePtr<V>;
ptr::drop_in_place(c_ptr.add(child_idx));
ptr::copy(
c_ptr.add(child_idx + 1),
c_ptr.add(child_idx),
len - child_idx,
);
}
self.set_len(len - 1);
}
fn split_leaf(&mut self) -> (i64, NodePtr<V>) {
let mid = self.len() / 2;
let right_count = self.len() - mid;
let mut right = NodePtr::new_leaf();
unsafe {
let k_src = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
let k_dst = right.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
ptr::copy_nonoverlapping(k_src.add(mid), k_dst, right_count);
let v_src = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
let v_dst = right.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
ptr::copy_nonoverlapping(v_src.add(mid), v_dst, right_count);
}
right.set_len(right_count);
self.set_len(mid);
let median = right.keys()[0];
(median, right)
}
fn split_leaf_rightmost(&mut self) -> (i64, NodePtr<V>) {
let len = self.len();
debug_assert!(
len == MAX_KEYS + 1,
"split_leaf_rightmost expects overflow node"
);
let mut right = NodePtr::new_leaf();
unsafe {
let k_src = self.ptr.data_ptr_mut().add(Self::keys_offset()) as *const i64;
let k_dst = right.ptr.data_ptr_mut().add(Self::keys_offset()) as *mut i64;
let v_src = self.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
let v_dst = right.ptr.data_ptr_mut().add(Self::values_offset()) as *mut V;
ptr::write(k_dst, *k_src.add(len - 1));
ptr::write(v_dst, ptr::read(v_src.add(len - 1)));
}
right.set_len(1);
self.set_len(len - 1);
let median = right.keys()[0];
(median, right)
}
}
const MAX_TREE_DEPTH: usize = 16;
#[derive(Clone, Copy)]
pub struct NodePath {
indices: [u8; MAX_TREE_DEPTH],
len: u8,
}
impl Default for NodePath {
fn default() -> Self {
Self {
indices: [0; MAX_TREE_DEPTH],
len: 0,
}
}
}
impl NodePath {
pub fn new() -> Self {
Self::default()
}
#[cold]
#[inline(never)]
fn depth_overflow() -> ! {
panic!("B-Tree depth exceeded maximum")
}
pub fn push(&mut self, idx: usize) {
if (self.len as usize) < MAX_TREE_DEPTH {
self.indices[self.len as usize] = idx as u8;
self.len += 1;
} else {
Self::depth_overflow();
}
}
fn get(&self, depth: usize) -> usize {
self.indices[depth] as usize
}
fn iter(&self) -> impl Iterator<Item = usize> + '_ {
(0..self.len as usize).map(|i| self.indices[i] as usize)
}
}
pub struct CowBTree<V: Clone> {
root: Option<NodePtr<V>>,
max_key: i64,
len: usize,
}
impl<V: Clone> Default for CowBTree<V> {
fn default() -> Self {
Self::new()
}
}
impl<V: Clone> Clone for CowBTree<V> {
#[inline]
fn clone(&self) -> Self {
Self {
root: self.root.clone(),
max_key: self.max_key,
len: self.len,
}
}
}
impl<V: Clone> CowBTree<V> {
#[inline]
pub fn new() -> Self {
assert!(
mem::align_of::<V>() <= 8,
"CowBTree value type alignment must be <= 8"
);
Self {
root: None,
max_key: 0,
len: 0,
}
}
#[inline]
pub fn len(&self) -> usize {
self.len
}
#[inline]
pub fn is_empty(&self) -> bool {
self.root.is_none()
}
#[inline]
pub fn get(&self, key: i64) -> Option<&V> {
let mut node = self.root.as_ref()?;
loop {
match node.search(key) {
Ok(i) => {
if node.is_leaf() {
return Some(&node.values()[i]);
}
node = node.child(i + 1);
}
Err(i) => {
if node.is_leaf() {
return None;
}
node = node.child(i);
}
}
}
}
#[inline]
pub fn get_mut(&mut self, key: i64) -> Option<&mut V> {
let (path, leaf_idx) = self.search_path(key);
let leaf_idx = leaf_idx.ok()?;
self.get_mut_with_path(key, &path, leaf_idx)
}
fn get_mut_with_path(&mut self, _key: i64, path: &NodePath, leaf_idx: usize) -> Option<&mut V> {
let root = self.root.as_mut()?;
let mut node = root.make_mut();
for idx in path.iter() {
let child = node.child_mut(idx);
node = child.make_mut();
}
if leaf_idx < node.len() {
Some(&mut node.values_mut_slice()[leaf_idx])
} else {
None
}
}
#[inline]
pub fn contains_key(&self, key: i64) -> bool {
self.get(key).is_some()
}
pub fn insert(&mut self, key: i64, value: V) -> Option<V> {
if self.root.is_none() {
let mut node = NodePtr::new_leaf();
node.push_leaf(key, value);
self.root = Some(node);
self.max_key = key;
self.len = 1;
return None;
}
if self.is_key_greater_than_max(key) {
let root = self.root.as_mut().unwrap();
let result = Self::insert_rightmost(root, key, value);
self.max_key = key;
return match result {
InsertResult::Done(old) => {
if old.is_none() {
self.len += 1;
}
old
}
InsertResult::Split(median, right) => {
let old_root = self.root.take().unwrap();
let mut new_root = NodePtr::new_internal();
unsafe {
let c_ptr = new_root
.ptr
.data_ptr_mut()
.add(NodePtr::<V>::children_offset())
as *mut NodePtr<V>;
ptr::write(c_ptr, old_root);
}
new_root.push_internal(median, right);
self.root = Some(new_root);
self.len += 1;
None
}
};
}
let root = self.root.as_mut().unwrap();
let result = Self::insert_recursive(root, key, value);
if key > self.max_key {
self.max_key = key;
}
match result {
InsertResult::Done(old) => {
if old.is_none() {
self.len += 1;
}
old
}
InsertResult::Split(median, right) => {
let old_root = self.root.take().unwrap();
let mut new_root = NodePtr::new_internal();
unsafe {
let c_ptr = new_root
.ptr
.data_ptr_mut()
.add(NodePtr::<V>::children_offset())
as *mut NodePtr<V>;
ptr::write(c_ptr, old_root);
}
new_root.push_internal(median, right);
self.root = Some(new_root);
self.len += 1;
None
}
}
}
#[inline]
fn is_key_greater_than_max(&self, key: i64) -> bool {
self.root.is_some() && key > self.max_key
}
fn insert_rightmost(node: &mut NodePtr<V>, key: i64, value: V) -> InsertResult<V> {
let (res, _) = Self::insert_rightmost_return_ptr(node, key, value);
res
}
fn insert_rightmost_return_ptr(
node: &mut NodePtr<V>,
key: i64,
value: V,
) -> (InsertResult<V>, *mut V) {
let node = node.make_mut();
if node.is_leaf() {
node.push_leaf(key, value);
unsafe {
let len = node.len();
let v_ptr = node.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
let ptr = v_ptr.add(len - 1);
if node.len() > MAX_KEYS {
let (median, right) = node.split_leaf_rightmost();
let v_ptr_new =
right.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
(InsertResult::Split(median, right), v_ptr_new)
} else {
(InsertResult::Done(None), ptr)
}
}
} else {
let last_idx = node.len();
let child = node.child_mut(last_idx);
let (result, ptr) = Self::insert_rightmost_return_ptr(child, key, value);
match result {
InsertResult::Done(old) => (InsertResult::Done(old), ptr),
InsertResult::Split(median, right) => {
node.push_internal(median, right);
if node.len() > MAX_KEYS {
let (m, r) = node.split_internal_rightmost();
(InsertResult::Split(m, r), ptr)
} else {
(InsertResult::Done(None), ptr)
}
}
}
}
}
fn insert_recursive(node: &mut NodePtr<V>, key: i64, value: V) -> InsertResult<V> {
let node = node.make_mut();
if node.is_leaf() {
match node.search(key) {
Ok(i) => unsafe {
let v_ptr =
node.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
let old = ptr::read(v_ptr.add(i));
ptr::write(v_ptr.add(i), value);
InsertResult::Done(Some(old))
},
Err(i) => {
node.insert_leaf(i, key, value);
if node.len() > MAX_KEYS {
let (median, right) = if i == node.len() - 1 {
node.split_leaf_rightmost()
} else {
node.split_leaf()
};
InsertResult::Split(median, right)
} else {
InsertResult::Done(None)
}
}
}
} else {
let i = match node.search(key) {
Ok(i) => i + 1,
Err(i) => i,
};
let result = Self::insert_recursive(node.child_mut(i), key, value);
match result {
InsertResult::Done(old) => InsertResult::Done(old),
InsertResult::Split(median, right) => {
node.insert_internal(i, median, right);
if node.len() > MAX_KEYS {
let (m, r) = if i == node.len() - 1 {
node.split_internal_rightmost()
} else {
node.split_internal()
};
InsertResult::Split(m, r)
} else {
InsertResult::Done(None)
}
}
}
}
}
pub fn remove(&mut self, key: i64) -> Option<V> {
let root = self.root.as_mut()?;
let result = Self::remove_recursive(root, key);
if result.is_some() {
self.len -= 1;
if self.len == 0 {
self.root = None;
self.max_key = 0;
} else if self.max_key == key {
self.refresh_max_key();
}
if let Some(ref mut root) = self.root {
let root = root.make_mut();
if !root.is_leaf() && root.len() == 0 {
let child_node = root.child_mut(0).clone();
self.root = Some(child_node);
} else if root.is_leaf() && root.len() == 0 {
self.root = None;
self.max_key = 0;
}
}
}
result
}
fn search_path(&self, key: i64) -> (NodePath, Result<usize, usize>) {
let mut path = NodePath::new();
if self.root.is_none() {
return (path, Err(0));
}
let mut node = self.root.as_ref().unwrap();
loop {
match node.search(key) {
Ok(i) => {
if node.is_leaf() {
return (path, Ok(i));
}
path.push(i + 1);
node = node.child(i + 1);
}
Err(i) => {
if node.is_leaf() {
return (path, Err(i));
}
path.push(i);
node = node.child(i);
}
}
}
}
fn insert_with_path(
node_ptr: &mut NodePtr<V>,
key: i64,
value: V,
path: &NodePath,
depth: usize,
leaf_idx: usize,
) -> (InsertResult<V>, *mut V) {
let node = node_ptr.make_mut();
if node.is_leaf() {
let i = leaf_idx;
node.insert_leaf(i, key, value);
if node.len() > MAX_KEYS {
let (median, right_node) = node.split_leaf();
let mid = MAX_KEYS.div_ceil(2);
let ptr = if i < mid {
unsafe {
let v_ptr =
node.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
v_ptr.add(i)
}
} else {
let right_idx = i - mid;
unsafe {
let v_ptr = right_node
.ptr
.data_ptr_mut()
.add(NodePtr::<V>::values_offset())
as *mut V;
v_ptr.add(right_idx)
}
};
(InsertResult::Split(median, right_node), ptr)
} else {
unsafe {
let v_ptr =
node.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
let ptr = v_ptr.add(i);
(InsertResult::Done(None), ptr)
}
}
} else {
let i = path.get(depth);
let (result, ptr) =
Self::insert_with_path(node.child_mut(i), key, value, path, depth + 1, leaf_idx);
match result {
InsertResult::Done(old) => (InsertResult::Done(old), ptr),
InsertResult::Split(median, right) => {
node.insert_internal(i, median, right);
if node.len() > MAX_KEYS {
let (m, r) = if i == node.len() - 1 {
node.split_internal_rightmost()
} else {
node.split_internal()
};
(InsertResult::Split(m, r), ptr)
} else {
(InsertResult::Done(None), ptr)
}
}
}
}
}
fn refresh_max_key(&mut self) {
if let Some(root) = &self.root {
let mut node = root;
loop {
if node.is_leaf() {
self.max_key = node.keys().last().copied().unwrap_or(0);
break;
}
node = node.children().last().unwrap();
}
} else {
self.max_key = 0;
}
}
fn remove_recursive(node: &mut NodePtr<V>, key: i64) -> Option<V> {
let node = node.make_mut();
if node.is_leaf() {
match node.search(key) {
Ok(i) => Some(node.remove_leaf(i)),
Err(_) => None,
}
} else {
let i = match node.search(key) {
Ok(i) => i + 1,
Err(i) => i,
};
if node.child(i).len() <= MIN_KEYS {
Self::ensure_child_can_lose_key(node, i);
}
let new_i = match node.search(key) {
Ok(i) => i + 1,
Err(i) => i,
};
let i = new_i.min(node.len()); Self::remove_recursive(node.child_mut(i), key)
}
}
fn ensure_child_can_lose_key(node: &mut NodePtr<V>, i: usize) {
let can_borrow_left = i > 0 && node.child(i - 1).len() > MIN_KEYS;
let can_borrow_right = i < node.len() && node.child(i + 1).len() > MIN_KEYS;
if can_borrow_left {
node.borrow_from_left(i);
} else if can_borrow_right {
node.borrow_from_right(i);
} else if i > 0 {
node.merge_with_left(i);
} else if i < node.len() {
node.merge_with_right(i);
}
}
pub fn iter_chunks(&self) -> impl Iterator<Item = (&[i64], &[V])> {
CowBTreeChunkIter::new(self.root.as_ref())
}
pub fn iter(&self) -> impl Iterator<Item = (&i64, &V)> {
self.iter_chunks()
.flat_map(|(keys, values)| keys.iter().zip(values.iter()))
}
pub fn keys(&self) -> impl Iterator<Item = i64> + '_ {
self.iter().map(|(k, _)| *k)
}
pub fn values(&self) -> impl Iterator<Item = &V> {
self.iter().map(|(_, v)| v)
}
pub fn range_chunks<R>(&self, range: R) -> impl Iterator<Item = (&[i64], &[V])>
where
R: std::ops::RangeBounds<i64>,
{
CowBTreeRangeChunkIter::new(self.root.as_ref(), range)
}
pub fn range<R>(&self, range: R) -> impl Iterator<Item = (&i64, &V)>
where
R: std::ops::RangeBounds<i64>,
{
self.range_chunks(range)
.flat_map(|(keys, values)| keys.iter().zip(values.iter()))
}
pub fn iter_rev_chunks(&self) -> impl Iterator<Item = (&[i64], &[V])> {
CowBTreeRevChunkIter::new(self.root.as_ref())
}
pub fn iter_rev(&self) -> impl Iterator<Item = (&i64, &V)> {
self.iter_rev_chunks()
.flat_map(|(keys, values)| keys.iter().zip(values.iter()).rev())
}
pub fn range_rev_chunks<R>(&self, range: R) -> impl Iterator<Item = (&[i64], &[V])>
where
R: std::ops::RangeBounds<i64>,
{
CowBTreeRevRangeChunkIter::new(self.root.as_ref(), range)
}
pub fn range_rev<R>(&self, range: R) -> impl Iterator<Item = (&i64, &V)>
where
R: std::ops::RangeBounds<i64>,
{
self.range_rev_chunks(range)
.flat_map(|(keys, values)| keys.iter().zip(values.iter()).rev())
}
pub fn clear(&mut self) {
self.root = None;
self.max_key = 0;
self.len = 0;
}
#[inline]
pub fn max_key(&self) -> Option<i64> {
if self.root.is_some() {
Some(self.max_key)
} else {
None
}
}
fn insert_rightmost_entry(&mut self, key: i64, value: V) -> *mut V {
let root = self.root.as_mut().unwrap();
let (result, ptr) = Self::insert_rightmost_return_ptr(root, key, value);
self.max_key = key;
match result {
InsertResult::Done(old) => {
if old.is_none() {
self.len += 1;
}
}
InsertResult::Split(median, right) => {
let old_root = self.root.take().unwrap();
let mut new_root = NodePtr::new_internal();
unsafe {
let c_ptr = new_root
.ptr
.data_ptr_mut()
.add(NodePtr::<V>::children_offset())
as *mut NodePtr<V>;
ptr::write(c_ptr, old_root);
}
new_root.push_internal(median, right);
self.root = Some(new_root);
self.len += 1;
}
}
ptr
}
pub fn entry(&mut self, key: i64) -> Entry<'_, V> {
if self.is_key_greater_than_max(key) {
return Entry::Vacant(VacantEntry {
tree: self,
key,
path: NodePath::new(), leaf_idx: 0, is_rightmost: true,
});
}
let (path, leaf_result) = self.search_path(key);
match leaf_result {
Ok(idx) => Entry::Occupied(OccupiedEntry {
tree: self,
key,
path,
leaf_idx: idx,
}),
Err(idx) => Entry::Vacant(VacantEntry {
tree: self,
key,
path,
leaf_idx: idx,
is_rightmost: false,
}),
}
}
fn insert_using_path(
&mut self,
key: i64,
value: V,
path: &NodePath,
leaf_idx: usize,
) -> *mut V {
if self.root.is_none() {
self.insert(key, value);
return self.get_mut(key).unwrap() as *mut V;
}
let root = self.root.as_mut().unwrap();
let (result, ptr) = Self::insert_with_path(root, key, value, path, 0, leaf_idx);
if key > self.max_key {
self.max_key = key;
}
match result {
InsertResult::Done(old) => {
if old.is_none() {
self.len += 1;
}
}
InsertResult::Split(median, right) => {
let old_root = self.root.take().unwrap();
let mut new_root = NodePtr::new_internal();
unsafe {
let c_ptr = new_root
.ptr
.data_ptr_mut()
.add(NodePtr::<V>::children_offset())
as *mut NodePtr<V>;
ptr::write(c_ptr, old_root);
}
new_root.push_internal(median, right);
self.root = Some(new_root);
self.len += 1;
}
}
ptr
}
}
pub enum Entry<'a, V: Clone> {
Occupied(OccupiedEntry<'a, V>),
Vacant(VacantEntry<'a, V>),
}
impl<'a, V: Clone> Entry<'a, V> {
pub fn or_insert(self, default: V) -> &'a mut V {
match self {
Entry::Occupied(entry) => entry.into_mut(),
Entry::Vacant(entry) => entry.insert(default),
}
}
pub fn and_modify<F>(self, f: F) -> Self
where
F: FnOnce(&mut V),
{
match self {
Entry::Occupied(mut entry) => {
f(entry.get_mut());
Entry::Occupied(entry)
}
Entry::Vacant(entry) => Entry::Vacant(entry),
}
}
pub fn key(&self) -> i64 {
match self {
Entry::Occupied(entry) => entry.key(),
Entry::Vacant(entry) => entry.key(),
}
}
}
pub struct OccupiedEntry<'a, V: Clone> {
tree: &'a mut CowBTree<V>,
key: i64,
path: NodePath,
leaf_idx: usize,
}
impl<'a, V: Clone> OccupiedEntry<'a, V> {
#[inline]
pub fn key(&self) -> i64 {
self.key
}
#[inline]
pub fn get(&self) -> &V {
let mut node = self.tree.root.as_ref().unwrap();
for idx in self.path.iter() {
node = node.child(idx);
}
&node.values()[self.leaf_idx]
}
pub fn get_mut(&mut self) -> &mut V {
self.tree
.get_mut_with_path(self.key, &self.path, self.leaf_idx)
.unwrap()
}
pub fn into_mut(self) -> &'a mut V {
self.tree
.get_mut_with_path(self.key, &self.path, self.leaf_idx)
.unwrap()
}
pub fn insert(&mut self, value: V) -> V {
let node = self.tree.root.as_mut().unwrap();
let mut node = node.make_mut();
for idx in self.path.iter() {
let child = node.child_mut(idx);
node = child.make_mut();
}
unsafe {
let v_ptr = node.ptr.data_ptr_mut().add(NodePtr::<V>::values_offset()) as *mut V;
let ptr = v_ptr.add(self.leaf_idx);
let old = ptr::read(ptr);
ptr::write(ptr, value);
old
}
}
}
pub struct VacantEntry<'a, V: Clone> {
tree: &'a mut CowBTree<V>,
key: i64,
path: NodePath,
leaf_idx: usize,
is_rightmost: bool,
}
impl<'a, V: Clone> VacantEntry<'a, V> {
#[inline]
pub fn key(&self) -> i64 {
self.key
}
#[inline]
pub fn insert(self, value: V) -> &'a mut V {
if self.is_rightmost {
let ptr = self.tree.insert_rightmost_entry(self.key, value);
unsafe { &mut *ptr }
} else {
let ptr = self
.tree
.insert_using_path(self.key, value, &self.path, self.leaf_idx);
unsafe { &mut *ptr }
}
}
}
enum InsertResult<V: Clone> {
Done(Option<V>),
Split(i64, NodePtr<V>),
}
struct CowBTreeChunkIter<'a, V: Clone> {
stack: Vec<(&'a NodePtr<V>, usize)>,
current_leaf: Option<&'a NodePtr<V>>,
}
impl<'a, V: Clone> CowBTreeChunkIter<'a, V> {
fn new(root: Option<&'a NodePtr<V>>) -> Self {
let mut iter = Self {
stack: Vec::new(),
current_leaf: None,
};
if let Some(root) = root {
iter.descend_to_leftmost(root);
}
iter
}
fn descend_to_leftmost(&mut self, mut node: &'a NodePtr<V>) {
while !node.is_leaf() {
self.stack.push((node, 1));
node = node.child(0);
}
self.current_leaf = Some(node);
}
}
impl<'a, V: Clone> Iterator for CowBTreeChunkIter<'a, V> {
type Item = (&'a [i64], &'a [V]);
fn next(&mut self) -> Option<Self::Item> {
if let Some(leaf) = self.current_leaf.take() {
return Some((leaf.keys(), leaf.values()));
}
loop {
let (node, idx) = self.stack.last_mut()?;
if *idx < node.len() + 1 {
let child_idx = *idx;
*idx += 1;
let child = node.child(child_idx);
self.descend_to_leftmost(child);
if let Some(leaf) = self.current_leaf.take() {
return Some((leaf.keys(), leaf.values()));
}
} else {
self.stack.pop();
}
}
}
}
struct CowBTreeRangeChunkIter<'a, V: Clone, R> {
stack: Vec<(&'a NodePtr<V>, usize)>,
range: R,
current_leaf: Option<&'a NodePtr<V>>,
current_idx: usize,
finished: bool,
}
impl<'a, V: Clone, R: std::ops::RangeBounds<i64>> CowBTreeRangeChunkIter<'a, V, R> {
fn new(root: Option<&'a NodePtr<V>>, range: R) -> Self {
let mut iter = Self {
stack: Vec::new(),
range,
current_leaf: None,
current_idx: 0,
finished: false,
};
if let Some(root) = root {
iter.seek_to_start(root);
} else {
iter.finished = true;
}
iter
}
fn seek_to_start(&mut self, mut node: &'a NodePtr<V>) {
let start_key = match self.range.start_bound() {
Bound::Included(&k) => Some(k),
Bound::Excluded(&k) => Some(k),
Bound::Unbounded => None,
};
loop {
if node.is_leaf() {
let keys = node.keys();
let mut idx = if let Some(k) = start_key {
match keys.binary_search(&k) {
Ok(i) => i,
Err(i) => i,
}
} else {
0
};
if let Bound::Excluded(&k) = self.range.start_bound() {
if idx < keys.len() && keys[idx] == k {
idx += 1;
}
}
self.current_leaf = Some(node);
self.current_idx = idx;
break;
} else {
let idx = if let Some(k) = start_key {
match node.search(k) {
Ok(i) => i + 1,
Err(i) => i,
}
} else {
0
};
self.stack.push((node, idx + 1));
node = node.child(idx);
}
}
}
}
impl<'a, V: Clone, R: std::ops::RangeBounds<i64>> Iterator for CowBTreeRangeChunkIter<'a, V, R> {
type Item = (&'a [i64], &'a [V]);
fn next(&mut self) -> Option<Self::Item> {
if self.finished {
return None;
}
loop {
if let Some(leaf) = self.current_leaf {
let keys = leaf.keys();
let values = leaf.values();
if self.current_idx < keys.len() {
let start = self.current_idx;
let end = match self.range.end_bound() {
Bound::Unbounded => keys.len(),
Bound::Included(&k) => {
if keys.last().unwrap() <= &k {
keys.len()
} else {
let pos = keys[start..].partition_point(|&x| x <= k);
self.finished = true;
start + pos
}
}
Bound::Excluded(&k) => {
if keys.last().unwrap() < &k {
keys.len()
} else {
let pos = keys[start..].partition_point(|&x| x < k);
self.finished = true;
start + pos
}
}
};
if start >= end {
self.finished = true;
self.current_leaf = None;
return None;
}
self.current_idx = end;
let result = (&keys[start..end], &values[start..end]);
if end == keys.len() && !self.finished {
self.current_leaf = None;
} else {
self.finished = true;
self.current_leaf = None;
}
return Some(result);
} else {
self.current_leaf = None;
}
}
if self.finished {
return None;
}
if let Some((node, idx)) = self.stack.last_mut() {
if *idx < node.len() + 1 {
let child_idx = *idx;
*idx += 1;
let mut child = node.child(child_idx);
loop {
if child.is_leaf() {
self.current_leaf = Some(child);
self.current_idx = 0;
break;
} else {
self.stack.push((child, 1));
child = child.child(0);
}
}
} else {
self.stack.pop();
if self.stack.is_empty() {
self.finished = true;
return None;
}
}
} else {
self.finished = true;
return None;
}
}
}
}
impl<V: Clone + std::fmt::Debug> std::fmt::Debug for CowBTree<V> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_map().entries(self.iter()).finish()
}
}
struct CowBTreeRevChunkIter<'a, V: Clone> {
stack: Vec<(&'a NodePtr<V>, usize)>,
current_leaf: Option<&'a NodePtr<V>>,
}
impl<'a, V: Clone> CowBTreeRevChunkIter<'a, V> {
fn new(root: Option<&'a NodePtr<V>>) -> Self {
let mut iter = Self {
stack: Vec::new(),
current_leaf: None,
};
if let Some(root) = root {
iter.descend_to_rightmost(root);
}
iter
}
fn descend_to_rightmost(&mut self, mut node: &'a NodePtr<V>) {
while !node.is_leaf() {
let last_child = node.len(); self.stack.push((node, last_child));
node = node.child(last_child);
}
self.current_leaf = Some(node);
}
}
impl<'a, V: Clone> Iterator for CowBTreeRevChunkIter<'a, V> {
type Item = (&'a [i64], &'a [V]);
fn next(&mut self) -> Option<Self::Item> {
if let Some(leaf) = self.current_leaf.take() {
return Some((leaf.keys(), leaf.values()));
}
loop {
let (node, next_plus_one) = self.stack.last_mut()?;
if *next_plus_one > 0 {
let child_idx = *next_plus_one - 1;
*next_plus_one = child_idx;
let child = node.child(child_idx);
self.descend_to_rightmost(child);
if let Some(leaf) = self.current_leaf.take() {
return Some((leaf.keys(), leaf.values()));
}
} else {
self.stack.pop();
}
}
}
}
struct CowBTreeRevRangeChunkIter<'a, V: Clone, R> {
stack: Vec<(&'a NodePtr<V>, usize)>,
range: R,
current_leaf: Option<&'a NodePtr<V>>,
current_end_idx: usize,
finished: bool,
}
impl<'a, V: Clone, R: std::ops::RangeBounds<i64>> CowBTreeRevRangeChunkIter<'a, V, R> {
fn new(root: Option<&'a NodePtr<V>>, range: R) -> Self {
let mut iter = Self {
stack: Vec::new(),
range,
current_leaf: None,
current_end_idx: 0,
finished: false,
};
if let Some(root) = root {
iter.seek_to_end(root);
} else {
iter.finished = true;
}
iter
}
fn descend_to_rightmost(&mut self, mut node: &'a NodePtr<V>) {
while !node.is_leaf() {
let last_child = node.len();
self.stack.push((node, last_child));
node = node.child(last_child);
}
self.current_leaf = Some(node);
self.current_end_idx = node.len();
}
fn seek_to_end(&mut self, mut node: &'a NodePtr<V>) {
let end_key = match self.range.end_bound() {
Bound::Included(&k) | Bound::Excluded(&k) => Some(k),
Bound::Unbounded => None,
};
if end_key.is_none() {
self.descend_to_rightmost(node);
return;
}
let k = end_key.unwrap();
loop {
if node.is_leaf() {
let keys = node.keys();
let idx = match keys.binary_search(&k) {
Ok(i) => match self.range.end_bound() {
Bound::Included(_) => i + 1,
_ => i,
},
Err(i) => i,
};
if idx > 0 {
self.current_leaf = Some(node);
self.current_end_idx = idx;
}
break;
} else {
let child_idx = match node.search(k) {
Ok(i) => i + 1,
Err(i) => i,
};
self.stack.push((node, child_idx));
node = node.child(child_idx);
}
}
}
}
impl<'a, V: Clone, R: std::ops::RangeBounds<i64>> Iterator for CowBTreeRevRangeChunkIter<'a, V, R> {
type Item = (&'a [i64], &'a [V]);
fn next(&mut self) -> Option<Self::Item> {
if self.finished {
return None;
}
loop {
if let Some(leaf) = self.current_leaf {
let keys = leaf.keys();
let values = leaf.values();
let end = self.current_end_idx;
let start = if end == 0 {
end } else {
match self.range.start_bound() {
Bound::Unbounded => 0,
Bound::Included(&k) => {
if keys[0] >= k {
0 } else {
self.finished = true;
keys[..end].partition_point(|&x| x < k)
}
}
Bound::Excluded(&k) => {
if keys[0] > k {
0
} else {
self.finished = true;
keys[..end].partition_point(|&x| x <= k)
}
}
}
};
self.current_leaf = None;
if start < end {
return Some((&keys[start..end], &values[start..end]));
}
if self.finished {
return None;
}
}
let (node, next_plus_one) = self.stack.last_mut()?;
if *next_plus_one > 0 {
let child_idx = *next_plus_one - 1;
*next_plus_one = child_idx;
let child = node.child(child_idx);
self.descend_to_rightmost(child);
} else {
self.stack.pop();
}
}
}
}
include!("cow_btree/tests.rs");