diff --git a/benches/binary_keys.rs b/benches/binary_keys.rs index 9a29a512..664d2528 100644 --- a/benches/binary_keys.rs +++ b/benches/binary_keys.rs @@ -45,6 +45,94 @@ fn binary_insert(bencher: Bencher, n: u64) { divan::black_box_drop(out) } +// Every branch in these fixtures has at most two children. Short paths use +// all eight three-byte binary keys; long paths branch at four spaced bytes. +fn short_key(mask: u8) -> [u8; 3] { + [ + b'0' + ((mask >> 2) & 1), + b'0' + ((mask >> 1) & 1), + b'0' + (mask & 1), + ] +} + +fn seed_val(map: &mut PathMap, key: &[u8], val: u64) { + map.write_zipper_at_path(key).set_val(val); +} + +fn short_map(target_len: usize, create: bool) -> PathMap { + let target = short_key(7); + let mut map = PathMap::new(); + for mask in 0..8 { + let key = short_key(mask); + if !create || !key.starts_with(&target[..target_len]) { + seed_val(&mut map, &key, mask as u64); + } + } + if !create && target_len < target.len() { + seed_val(&mut map, &target[..target_len], 0); + } + assert_eq!(map.path_exists_at(&target[..target_len]), !create); + map +} + +fn long_key(len: usize, mask: u8) -> Vec { + let mut key = vec![b'-'; len]; + for (bit, index) in [0, len / 4, len / 2, 3 * len / 4].into_iter().enumerate() { + key[index] = b'0' + ((mask >> (3 - bit)) & 1); + } + key +} + +fn long_map(len: usize, create: bool) -> PathMap { + let mut map = PathMap::new(); + for mask in 0..16 { + if !create || mask != 15 { + seed_val(&mut map, &long_key(len, mask), mask as u64); + } + } + assert_eq!(map.path_exists_at(long_key(len, 15)), !create); + map +} + +#[divan::bench(sample_size = 64, args = [0usize, 1, 2, 3])] +fn binary_set_val_at_short_replace(bencher: Bencher, key_len: usize) { + let key = short_key(7); + let mut map = short_map(key_len, false); + bencher.bench_local(|| { + black_box(&mut map).set_val_at(black_box(&key[..key_len]), black_box(1)); + }); +} + +// The empty path is the root, so creating a new path starts at length one. +#[divan::bench(sample_size = 16, args = [1usize, 2, 3])] +fn binary_set_val_at_short_create(bencher: Bencher, key_len: usize) { + let key = short_key(7); + let out = bencher.with_inputs(|| short_map(key_len, true)).bench_local_values(|mut map| { + black_box(&mut map).set_val_at(black_box(&key[..key_len]), black_box(1)); + map + }); + divan::black_box_drop(out); +} + +#[divan::bench(args = [160usize, 256])] +fn binary_set_val_at_long_replace(bencher: Bencher, key_len: usize) { + let key = long_key(key_len, 15); + let mut map = long_map(key_len, false); + bencher.bench_local(|| { + black_box(&mut map).set_val_at(black_box(&key), black_box(1)); + }); +} + +#[divan::bench(sample_size = 16, args = [160usize, 256])] +fn binary_set_val_at_long_create(bencher: Bencher, key_len: usize) { + let key = long_key(key_len, 15); + let out = bencher.with_inputs(|| long_map(key_len, true)).bench_local_values(|mut map| { + black_box(&mut map).set_val_at(black_box(&key), black_box(1)); + map + }); + divan::black_box_drop(out); +} + #[divan::bench(args = [250, 500, 1000, 2000, 4000, 8000])] fn binary_get(bencher: Bencher, n: u64) { diff --git a/src/experimental.rs b/src/experimental.rs index f39c2caa..ab69ec52 100644 --- a/src/experimental.rs +++ b/src/experimental.rs @@ -159,6 +159,7 @@ impl ZipperWriting for NullZipper { fn get_val_or_set_mut(&mut self, default: V) -> &mut V { Box::leak(Box::new(default)) } fn get_val_or_set_mut_with(&mut self, func: F) -> &mut V where F: FnOnce() -> V { Box::leak(Box::new(func())) } fn set_val(&mut self, _val: V) -> Option { None } + fn set_val_at>(&mut self, path: K, val: V) -> Option { None } fn remove_val(&mut self, _prune: bool) -> Option { None } fn zipper_head<'z>(&'z mut self) -> Self::ZipperHead<'z> { todo!() } fn graft>(&mut self, _read_zipper: &Z) {} diff --git a/src/trie_map.rs b/src/trie_map.rs index a9cc1da0..02dc69f4 100644 --- a/src/trie_map.rs +++ b/src/trie_map.rs @@ -337,21 +337,19 @@ impl PathMap { self.path_exists_at(k) } - /// Inserts `v` into the map at `path`. Panics if `path` has a zero length + /// Inserts `v` into the map at `path`. /// /// Returns `Some(replaced_val)` if an existing value was replaced, otherwise returns `None` if /// the value was added to the map without replacing anything. pub fn set_val_at>(&mut self, path: K, v: V) -> Option { let path = path.as_ref(); - - //NOTE: Here is the old impl traversing without the zipper. Kept here for benchmarking purposes - // However, the zipper version is basically identical performance, within the margin of error - // traverse_to_leaf_static_result(&mut self.root, k, - // |node, remaining_key| node.node_set_val(remaining_key, v), - // |_new_leaf_node, _remaining_key| None) - - let mut zipper = self.write_zipper_at_path(path); - zipper.set_val(v) + if path.is_empty() { + return core::mem::replace(self.root_val_mut(), Some(v)); + } + let (old_val, _) = with_node_at_path_mut(self.get_or_init_root_mut(), path, + |node, remaining_key| node.node_set_val(remaining_key, v), + |_, _| (None, true)); + old_val } /// Alias for [Self::set_val_at], so `PathMap` "feels" like other Rust collections diff --git a/src/trie_node.rs b/src/trie_node.rs index 27366659..09fcabe8 100644 --- a/src/trie_node.rs +++ b/src/trie_node.rs @@ -2659,6 +2659,26 @@ pub(crate) fn node_along_path_mut<'a, 'k, V: Clone + Send + Sync, A: Allocator>( (key, node) } +/// Applies a node operation at a path, replacing the node if it needs to be upgraded. +#[inline] +pub(crate) fn with_node_at_path_mut(root: &mut TrieNodeODRc, path: &[u8], node_f: NodeF, retry_f: RetryF) -> R +where + V: Clone + Send + Sync, + A: Allocator, + NodeF: FnOnce(&mut TaggedNodeRefMut<'_, V, A>, &[u8]) -> Result>, + RetryF: FnOnce(&mut TaggedNodeRefMut<'_, V, A>, &[u8]) -> R, +{ + debug_assert!(!path.is_empty()); + let (remaining_key, node) = node_along_path_mut(root, path, true); + match node_f(&mut node.make_mut(), remaining_key) { + Ok(result) => result, + Err(replacement_node) => { + *node = replacement_node; + retry_f(&mut node.make_mut(), remaining_key) + } + } +} + /// Ensures the node is a CellByteNode /// /// Returns `true` if the node was upgraded and `false` if it already was a CellByteNode diff --git a/src/write_zipper.rs b/src/write_zipper.rs index b38eed05..1c1582ec 100644 --- a/src/write_zipper.rs +++ b/src/write_zipper.rs @@ -50,7 +50,15 @@ pub trait ZipperWriting: Wri /// /// Returns `Some(replaced_val)` if an existing value was replaced, otherwise returns `None` if /// the value was added without replacing anything. - fn set_val(&mut self, val: V) -> Option; + fn set_val(&mut self, val: V) -> Option { + self.set_val_at([], val) + } + + /// Sets the value at a path relative to the zipper's focus + /// + /// Returns `Some(replaced_val)` if an existing value was replaced, otherwise returns `None` if + /// the value was added without replacing anything. + fn set_val_at>(&mut self, path: K, val: V) -> Option; /// Deprecated alias for [ZipperWriting::set_val] #[deprecated] //GOAT-old-names @@ -351,6 +359,7 @@ impl ZipperWriting for &mut Z whe fn get_val_or_set_mut(&mut self, default: V) -> &mut V { (**self).get_val_or_set_mut(default) } fn get_val_or_set_mut_with(&mut self, func: F) -> &mut V where F: FnOnce() -> V { (**self).get_val_or_set_mut_with(func) } fn set_val(&mut self, val: V) -> Option { (**self).set_val(val) } + fn set_val_at>(&mut self, path: K, val: V) -> Option { (**self).set_val_at(path, val) } fn remove_val(&mut self, prune: bool) -> Option { (**self).remove_val(prune) } fn zipper_head<'z>(&'z mut self) -> Self::ZipperHead<'z> { (**self).zipper_head() } fn graft>(&mut self, read_zipper: &RZ) { (**self).graft(read_zipper) } @@ -520,6 +529,7 @@ impl<'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> ZipperWriting fn get_val_or_set_mut(&mut self, default: V) -> &mut V { self.z.get_val_or_set_mut(default) } fn get_val_or_set_mut_with(&mut self, func: F) -> &mut V where F: FnOnce() -> V { self.z.get_val_or_set_mut_with(func) } fn set_val(&mut self, val: V) -> Option { self.z.set_val(val) } + fn set_val_at>(&mut self, path: K, val: V) -> Option { self.z.set_val_at(path, val) } fn remove_val(&mut self, prune: bool) -> Option { self.z.remove_val(prune) } fn zipper_head<'z>(&'z mut self) -> Self::ZipperHead<'z> { self.z.zipper_head() } fn graft>(&mut self, read_zipper: &Z) { self.z.graft(read_zipper) } @@ -690,6 +700,7 @@ impl<'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> ZipperWriting fn get_val_or_set_mut(&mut self, default: V) -> &mut V { self.z.get_val_or_set_mut(default) } fn get_val_or_set_mut_with(&mut self, func: F) -> &mut V where F: FnOnce() -> V { self.z.get_val_or_set_mut_with(func) } fn set_val(&mut self, val: V) -> Option { self.z.set_val(val) } + fn set_val_at>(&mut self, path: K, val: V) -> Option { self.z.set_val_at(path, val) } fn remove_val(&mut self, prune: bool) -> Option { self.z.remove_val(prune) } fn zipper_head<'z>(&'z mut self) -> Self::ZipperHead<'z> { self.z.zipper_head() } fn graft>(&mut self, read_zipper: &Z) { self.z.graft(read_zipper) } @@ -830,6 +841,7 @@ impl ZipperWriting for Write fn get_val_or_set_mut(&mut self, default: V) -> &mut V { self.z.get_val_or_set_mut(default) } fn get_val_or_set_mut_with(&mut self, func: F) -> &mut V where F: FnOnce() -> V { self.z.get_val_or_set_mut_with(func) } fn set_val(&mut self, val: V) -> Option { self.z.set_val(val) } + fn set_val_at>(&mut self, path: K, val: V) -> Option { self.z.set_val_at(path, val) } fn remove_val(&mut self, prune: bool) -> Option { self.z.remove_val(prune) } fn zipper_head<'z>(&'z mut self) -> Self::ZipperHead<'z> { self.z.zipper_head() } fn graft>(&mut self, read_zipper: &Z) { self.z.graft(read_zipper) } @@ -1439,21 +1451,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC } /// See [ZipperWriting::set_val] pub fn set_val(&mut self, val: V) -> Option { - if self.key.node_key().len() == 0 { - debug_assert!(self.at_root()); - let root_val_ref = self.root_val.as_mut().unwrap(); - let mut temp_val = Some(val); - core::mem::swap(unsafe{&mut **root_val_ref}, &mut temp_val); - return temp_val - } - let (old_val, created_subnode) = self.in_zipper_mut_static_result( - |node, remaining_key| node.node_set_val(remaining_key, val), - |_new_leaf_node, _remaining_key| (None, true)); - if created_subnode { - self.mend_root(); - self.descend_to_internal(); - } - old_val + self.set_val_at(&[], val) } /// See [ZipperWriting::remove_val] pub fn remove_val(&mut self, prune: bool) -> Option { @@ -1716,7 +1714,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC self.set_node_at_child_path(&[child_byte], node) } if let Some(val) = src_root_val { - let _ = self.set_val_at_child_path(&[child_byte], val); + let _ = self.set_val_at(&[child_byte], val); } } } @@ -1736,12 +1734,26 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC } } - /// Sets a child value one byte below the focus + /// Sets a value at a path relative to the focus #[inline] - fn set_val_at_child_path(&mut self, path: &[u8], val: V) -> Option { - let (old_val, created_subnode) = self.with_node_at_path(path, - |node, remaining_key| node.node_set_val(remaining_key, val), - |_new_leaf_node, _remaining_key| (None, true)); + fn set_val_at>(&mut self, path: K, val: V) -> Option { + let path = path.as_ref(); + + //Special case for the root val + if path.is_empty() && self.key.node_key().is_empty() { + debug_assert!(self.at_root()); + let root_val_ref = self.root_val.as_mut().unwrap(); + return core::mem::replace(unsafe { &mut **root_val_ref }, Some(val)); + } + let (old_val, created_subnode) = if path.is_empty() { + self.in_zipper_mut_static_result( + |node, remaining_key| node.node_set_val(remaining_key, val), + |_new_leaf_node, _remaining_key| (None, true)) + } else { + self.with_node_at_path(path, + |node, remaining_key| node.node_set_val(remaining_key, val), + |_new_leaf_node, _remaining_key| (None, true)) + }; if created_subnode { self.mend_root(); self.descend_to_internal(); @@ -2440,30 +2452,49 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC { let key = self.key.node_key(); let mut focus_node = self.focus_stack.top_mut().unwrap(); - if let Some((key_bytes, child_node)) = focus_node.node_get_child_mut(key) { + if !key.is_empty() && let Some((key_bytes, child_node)) = focus_node.node_get_child_mut(key) { debug_assert_eq!(key_bytes, key.len()); - let (key, node) = node_along_path_mut(child_node, path, true); - let mut node_ref = node.make_mut(); - match node_f(&mut node_ref, key) { + with_node_at_path_mut(child_node, path, node_f, retry_f) + } else if key.is_empty() { + // At the zipper root there is no focus key to combine with `path`. + // Walk existing children first, as write_zipper_at_path does. + drop(focus_node); + with_node_at_path_mut(self.focus_stack.root_mut().unwrap(), path, node_f, retry_f) + } else if key.len() + path.len() <= MAX_NODE_KEY_BYTES { + let mut key_buf = [0u8; MAX_NODE_KEY_BYTES]; + key_buf[..key.len()].copy_from_slice(key); + key_buf[key.len()..key.len()+path.len()].copy_from_slice(path); + let full_key = &key_buf[..key.len()+path.len()]; + drop(focus_node); + self.in_zipper_mut_static_result( + |focus_node, _| node_f(focus_node, full_key), + |focus_node, _| retry_f(focus_node, full_key), + ) + } else { + // The focus is not represented by its own node yet. Take any subtree + // below it, apply the operation there, then back-fill the focus path + // using the same node insertion machinery as set_val. + let mut child = focus_node.take_node_at_key(key, false).filter(|node| !node.is_empty()).unwrap_or_else(|| { + #[cfg(not(feature = "all_dense_nodes"))] + { TrieNodeODRc::new_in(crate::line_list_node::LineListNode::new_in(self.alloc.clone()), self.alloc.clone()) } + #[cfg(feature = "all_dense_nodes")] + { TrieNodeODRc::new_in(crate::dense_byte_node::DenseByteNode::new_in(self.alloc.clone()), self.alloc.clone()) } + }); + let result = match node_f(&mut child.make_mut(), path) { Ok(result) => result, Err(replacement_node) => { - *node = replacement_node; - retry_f(&mut node.make_mut(), key) + child = replacement_node; + retry_f(&mut child.make_mut(), path) }, - } - } else { + }; + drop(focus_node); self.in_zipper_mut_static_result( - |focus_node, partial_key| { - let mut key_buf = [0u8; MAX_NODE_KEY_BYTES]; - key_buf[0..partial_key.len()].copy_from_slice(partial_key); - //GOAT, currently this will panic if the path is too long to fit in the buffer, which means this internal API - // isn't suitable for general-purpose path-based ops yet, but we're using it to deal with single-byte ops - key_buf[partial_key.len()..partial_key.len()+path.len()].copy_from_slice(path); - let full_key = &key_buf[0..partial_key.len()+path.len()]; - node_f(focus_node, full_key) - }, - retry_f - ) + |node, key| node.node_set_branch(key, child), + |_, _| true, + ); + self.mend_root(); + self.descend_to_internal(); + result } } @@ -6132,6 +6163,72 @@ mod tests { assert_eq!(keys(&m), ["cax", "cbx", "cdx", "d"]); } + /// `graft_child_maps` and `graft_masked_branches` below a root path too long for one node key + #[test] + fn graft_child_maps_long_root() { + for root_len in [47usize, 48, 60, 200] { + let root = vec![0u8; root_len]; + let mut map = PathMap::::new(); + { + let mut wz = map.write_zipper_at_path(&root); + wz.graft_child_maps(ByteMask::from_iter([1u8, 3]), [PathMap::single([2u8], 5), PathMap::single([], 6)], false); + } + let mut want = root.clone(); + want.extend([1u8, 2]); + assert_eq!(map.val_at(&want), Some(&5), "root {root_len}"); + want.truncate(root_len); + want.push(3); + assert_eq!(map.val_at(&want).is_some(), cfg!(feature = "graft_root_vals"), "root {root_len}"); + + let mut src = PathMap::::new(); + src.set_val_at([4u8, 4], 9); + let mut map = PathMap::::new(); + { + let mut wz = map.write_zipper_at_path(&root); + wz.graft_masked_branches(&src.read_zipper(), ByteMask::from_iter([4u8]), false); + } + let mut want = root.clone(); + want.extend([4u8, 4]); + assert_eq!(map.val_at(&want), Some(&9), "root {root_len}"); + } + } + + #[test] + fn set_val_at_below_long_missing_focus() { + for focus_len in [48usize, 60, 200] { + let focus = vec![0u8; focus_len]; + let child_path = vec![1u8; 96]; + let mut map = PathMap::::new(); + { + let mut zipper = map.write_zipper_at_path(&focus); + assert_eq!(zipper.set_val_at(&child_path, 7), None); + } + let mut full_path = focus; + full_path.extend_from_slice(&child_path); + assert_eq!(map.val_at(&full_path), Some(&7)); + assert_eq!(map.set_val_at(&full_path, 8), Some(7)); + assert_eq!(map.val_at(&full_path), Some(&8)); + } + } + + #[test] + fn set_val_at_empty_path_sets_focus() { + let mut map = PathMap::::new(); + { + let mut zipper = map.write_zipper_at_path(b"focus"); + assert_eq!(zipper.set_val_at(&[], 7), None); + assert_eq!(zipper.set_val_at(&[], 8), Some(7)); + } + assert_eq!(map.val_at(b"focus"), Some(&8)); + { + let mut zipper = map.write_zipper(); + assert_eq!(zipper.set_val_at(&[], 9), None); + assert_eq!(zipper.set_val_at(&[], 10), Some(9)); + } + assert_eq!(map.val_at([]), Some(&10)); + assert_eq!(map.val_at(b"focus"), Some(&8)); + } + #[test] fn write_zipper_graft_masked_branches_test4() { // Upper bound 0: remove_unset=true with an empty mask.