diff --git a/.github/scripts/bench_ab.py b/.github/scripts/bench_ab.py index bf8f96a9..191548b9 100755 --- a/.github/scripts/bench_ab.py +++ b/.github/scripts/bench_ab.py @@ -165,6 +165,11 @@ def compare_bench(self, bench, rounds_so_far): cmp = self.fmt.compare_fields(avg['base'], avg['head'], 'median_ns') self.results[bench] = {f'{g}/{c}': {'base': r['base'], 'head': r['other'], 'pct': r['pct']} for (g, c), r in cmp.items()} table = re.sub(r'\x1b\[[0-9;]*m', '', self.fmt.render_divan_table(cmp)) + # a PR that adds, removes or renames cases leaves them on one side only; say so rather than drop them silently + for side, mine, theirs in (('base', avg['base'], avg['head']), ('head', avg['head'], avg['base'])): + only = sorted(f'{g}/{c}' for g, c in mine.keys() - theirs.keys()) + if only: + table += f'\nonly in {side} ({len(only)}): ' + ', '.join(only) text = (f'{bench} (base {self.short(self.base_sha)} head {self.short(self.head_sha)}' f' rounds {rounds_so_far} median ns)\n{table}\n\n') (self.out / f'cmp-{bench}.txt').write_text(text) diff --git a/Cargo.toml b/Cargo.toml index a03740e1..c90fff4b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -137,6 +137,10 @@ harness = false name = "zipper_head_owned" harness = false +[[bench]] +name = "prune" +harness = false + [[bench]] name = "act_paths" harness = false diff --git a/benches/divan_fmt.py b/benches/divan_fmt.py index 5fb79c95..00e6ff36 100644 --- a/benches/divan_fmt.py +++ b/benches/divan_fmt.py @@ -119,6 +119,8 @@ def render_divan_table(data): for group in grouped: grouped[group].sort(key=lambda item: case_sort_key(item[0])) + if not grouped: + return "(no cases)" (_, first_record) = next(iter(grouped.values()))[0] lines = [] diff --git a/benches/prune.rs b/benches/prune.rs new file mode 100644 index 00000000..595c4b17 --- /dev/null +++ b/benches/prune.rs @@ -0,0 +1,92 @@ +use divan::{Bencher, Divan, black_box}; +use pathmap::PathMap; +use pathmap::zipper::*; + +fn main() { + Divan::from_args().main(); +} + +fn fixture(path: &[u8]) -> PathMap { + let mut map = PathMap::new(); + map.set_val_at(path, 1); + map +} + +fn run_remove_val(bencher: Bencher, path: &[u8], root_len: usize, prune: bool) { + bencher.with_inputs(|| fixture(path)).bench_local_values(|mut map| { + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + black_box(wz.remove_val(prune)); + }); +} + +fn run_remove_branches(bencher: Bencher, path: &[u8], root_len: usize, prune: bool) { + bencher.with_inputs(|| fixture(path)).bench_local_values(|mut map| { + let focus = &path[..path.len() - 1]; + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&focus[root_len..]); + black_box(wz.remove_branches(prune)); + }); +} + +fn run_prune_path(bencher: Bencher, path: &[u8], root_len: usize) { + bencher.with_inputs(|| { + let mut map = PathMap::::new(); + map.create_path(path); + map + }).bench_local_values(|mut map| { + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + black_box(wz.prune_path()); + }); +} + +#[divan::bench] +fn prune_path_short_root_at_map_root(bencher: Bencher) { + run_prune_path(bencher, b"abcd", 0); +} + +#[divan::bench] +fn prune_path_short_root_inside_node(bencher: Bencher) { + run_prune_path(bencher, b"abcd", 2); +} + +#[divan::bench] +fn prune_path_long_root_inside_node(bencher: Bencher) { + let path: Vec = (0..100).collect(); + run_prune_path(bencher, &path, 95); +} + +#[divan::bench(args = [false, true])] +fn remove_val_short(bencher: Bencher, prune: bool) { + run_remove_val(bencher, b"abcd", 2, prune); +} + +#[divan::bench(args = [false, true])] +fn remove_val_long_root_above_node(bencher: Bencher, prune: bool) { + let path: Vec = (0..100).collect(); + run_remove_val(bencher, &path, 5, prune); +} + +#[divan::bench(args = [false, true])] +fn remove_val_long_root_inside_node(bencher: Bencher, prune: bool) { + let path: Vec = (0..100).collect(); + run_remove_val(bencher, &path, 95, prune); +} + +#[divan::bench(args = [false, true])] +fn remove_branches_short(bencher: Bencher, prune: bool) { + run_remove_branches(bencher, b"abcd", 2, prune); +} + +#[divan::bench(args = [false, true])] +fn remove_branches_long_root_above_node(bencher: Bencher, prune: bool) { + let path: Vec = (0..100).collect(); + run_remove_branches(bencher, &path, 5, prune); +} + +#[divan::bench(args = [false, true])] +fn remove_branches_long_root_inside_node(bencher: Bencher, prune: bool) { + let path: Vec = (0..100).collect(); + run_remove_branches(bencher, &path, 95, prune); +} diff --git a/benches/zipper_head_owned.rs b/benches/zipper_head_owned.rs index 4b4d032d..7e8534ad 100644 --- a/benches/zipper_head_owned.rs +++ b/benches/zipper_head_owned.rs @@ -2,6 +2,8 @@ use divan::{Bencher, Divan, black_box}; use pathmap::PathMap; use pathmap::zipper::*; +const REPEATS: usize = 100; + fn main() { Divan::from_args().sample_count(100).main(); } @@ -16,82 +18,84 @@ fn zipper_head_fixture() -> PathMap { map } -fn bench_read_creation(bencher: Bencher, repeats: usize, mut read_child_count: F) +fn bench_head_read_creation<'trie, H>(bencher: Bencher, head: &H) where - F: FnMut() -> usize, + H: ZipperCreation<'trie, usize>, { + let path = [7u8]; bencher.bench_local(|| { let mut observed = 0usize; - for _ in 0..repeats { - observed += read_child_count(); + for _ in 0..REPEATS { + let reader = head.read_zipper_at_borrowed_path(black_box(&path)).unwrap(); + observed += reader.child_count(); } black_box(observed); }); } -fn bench_write_creation_cleanup(bencher: Bencher, repeats: usize, mut create_and_cleanup: F) +fn bench_head_write_creation_cleanup<'trie, H, const CHECKED: bool>(bencher: Bencher, head: &H) where - F: FnMut([u8; 2]) -> usize, + H: ZipperCreation<'trie, usize>, { bencher.bench_local(|| { + let mut writers = Vec::with_capacity(REPEATS); let mut observed = 0usize; - for i in 0..repeats { - observed += create_and_cleanup([240u8, i as u8]); + for i in 0..REPEATS { + let path = black_box([240u8, i as u8]); + // The paths are disjoint. All writers stay live until cleanup below. + let writer = if CHECKED { + head.write_zipper_at_exclusive_path(path).unwrap() + } else { + unsafe { head.write_zipper_at_exclusive_path_unchecked(path) } + }; + observed += writer.path_exists() as usize; + writers.push(writer); + } + for writer in writers { + head.cleanup_write_zipper(writer); } black_box(observed); }); } -fn bench_head_read_creation<'trie, H>(bencher: Bencher, repeats: usize, head: &H) -where - H: ZipperCreation<'trie, usize>, -{ - let path = [7u8]; - - bench_read_creation(bencher, repeats, || { - let reader = head.read_zipper_at_borrowed_path(black_box(&path)).unwrap(); - reader.child_count() - }); +#[divan::bench] +fn borrowed_head_read_creation(bencher: Bencher) { + let mut map = zipper_head_fixture(); + let head = black_box(&mut map).zipper_head(); + bench_head_read_creation(bencher, &head); } -fn bench_head_write_creation_cleanup<'trie, H>(bencher: Bencher, repeats: usize, head: &H) -where - H: ZipperCreation<'trie, usize>, -{ - bench_write_creation_cleanup(bencher, repeats, |path| { - let writer = head - .write_zipper_at_exclusive_path(black_box(path)) - .unwrap(); - let observed = writer.path_exists() as usize; - head.cleanup_write_zipper(writer); - observed - }); +#[divan::bench] +fn owned_head_read_creation(bencher: Bencher) { + let map = zipper_head_fixture(); + let head = black_box(map).into_zipper_head([]); + bench_head_read_creation(bencher, &head); } -#[divan::bench(args = [1usize, 10, 100])] -fn borrowed_head_read_creation(bencher: Bencher, repeats: usize) { +#[divan::bench] +fn borrowed_head_write_creation_cleanup(bencher: Bencher) { let mut map = zipper_head_fixture(); let head = black_box(&mut map).zipper_head(); - bench_head_read_creation(bencher, repeats, &head); + bench_head_write_creation_cleanup::<_, true>(bencher, &head); } -#[divan::bench(args = [1usize, 10, 100])] -fn owned_head_read_creation(bencher: Bencher, repeats: usize) { +#[divan::bench] +fn owned_head_write_creation_cleanup(bencher: Bencher) { let map = zipper_head_fixture(); let head = black_box(map).into_zipper_head([]); - bench_head_read_creation(bencher, repeats, &head); + bench_head_write_creation_cleanup::<_, true>(bencher, &head); } -#[divan::bench(args = [1usize, 10, 100])] -fn borrowed_head_write_creation_cleanup(bencher: Bencher, repeats: usize) { +#[divan::bench] +fn borrowed_head_write_creation_cleanup_unchecked(bencher: Bencher) { let mut map = zipper_head_fixture(); let head = black_box(&mut map).zipper_head(); - bench_head_write_creation_cleanup(bencher, repeats, &head); + bench_head_write_creation_cleanup::<_, false>(bencher, &head); } -#[divan::bench(args = [1usize, 10, 100])] -fn owned_head_write_creation_cleanup(bencher: Bencher, repeats: usize) { +#[divan::bench] +fn owned_head_write_creation_cleanup_unchecked(bencher: Bencher) { let map = zipper_head_fixture(); let head = black_box(map).into_zipper_head([]); - bench_head_write_creation_cleanup(bencher, repeats, &head); + bench_head_write_creation_cleanup::<_, false>(bencher, &head); } diff --git a/src/dense_byte_node.rs b/src/dense_byte_node.rs index 07af4970..ab6bf2db 100644 --- a/src/dense_byte_node.rs +++ b/src/dense_byte_node.rs @@ -178,14 +178,14 @@ impl> ByteNode } #[inline] - pub fn remove_val(&mut self, k: u8, prune: bool) -> Option { + pub fn remove_val(&mut self, k: u8, prune_limit: usize) -> Option { if self.mask.test_bit(k) { let ix = self.mask.index_of(k) as usize; let cf = unsafe { self.values.get_unchecked_mut(ix) }; let result = cf.take_val(); - if prune && !cf.has_rec() { + if prune_limit == 0 && !cf.has_rec() { self.mask.clear_bit(k); self.values.remove(ix); } @@ -895,9 +895,9 @@ impl> TrieNode } } } - fn node_remove_val(&mut self, key: &[u8], prune: bool) -> Option { + fn node_remove_val(&mut self, key: &[u8], prune_limit: usize) -> Option { if key.len() == 1 { - self.remove_val(key[0], prune) + self.remove_val(key[0], prune_limit) } else { None } @@ -929,8 +929,11 @@ impl> TrieNode } } } - fn node_remove_dangling(&mut self, key: &[u8]) -> usize { + fn node_remove_dangling(&mut self, key: &[u8], min_keep_len: usize) -> usize { debug_assert!(key.len() > 0); + if min_keep_len >= key.len() { + return 0; + } if key.len() == 1 { let k = key[0]; if self.mask.test_bit(k) { @@ -994,7 +997,7 @@ impl> TrieNode } } } - fn node_remove_all_branches(&mut self, key: &[u8], prune: bool) -> bool { + fn node_remove_all_branches(&mut self, key: &[u8], prune_limit: usize) -> bool { if key.len() > 1 { return false; } @@ -1009,7 +1012,7 @@ impl> TrieNode true }, (true, false) => { - if prune { + if prune_limit == 0 { self.values.remove(ix); self.mask.clear_bit(k); } else { @@ -1181,7 +1184,7 @@ impl> TrieNode (Some(&ALL_BYTES[prefix..=prefix]), cf.rec().map(|cf| cf.as_tagged())) } - fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, _prune: bool) { + fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, _prune_limit: usize) { if key.len() > 0 { //We're in a non-existent path below this node return @@ -1298,7 +1301,7 @@ impl> TrieNode } } - fn take_node_at_key(&mut self, key: &[u8], prune: bool) -> Option> { + fn take_node_at_key(&mut self, key: &[u8], prune_limit: usize) -> Option> { if key.len() < 2 { debug_assert!(key.len() == 1); let k = key[0]; @@ -1308,7 +1311,7 @@ impl> TrieNode let cf = unsafe { self.values.get_unchecked_mut(ix) }; let result = cf.take_rec(); - if prune && !cf.has_val() { + if prune_limit == 0 && !cf.has_val() { self.mask.clear_bit(k); self.values.remove(ix); } diff --git a/src/empty_node.rs b/src/empty_node.rs index ce336dd1..c59b711f 100644 --- a/src/empty_node.rs +++ b/src/empty_node.rs @@ -35,13 +35,13 @@ impl TrieNode for EmptyNode { fn node_get_val(&self, _key: &[u8]) -> Option<&V> { None } - fn node_remove_val(&mut self, _key: &[u8], _prune: bool) -> Option { + fn node_remove_val(&mut self, _key: &[u8], _prune_limit: usize) -> Option { unreachable!() } fn node_create_dangling(&mut self, _key: &[u8]) -> Result<(bool, bool), TrieNodeODRc> { unreachable!() } - fn node_remove_dangling(&mut self, _key: &[u8]) -> usize { + fn node_remove_dangling(&mut self, _key: &[u8], _min_keep_len: usize) -> usize { unreachable!() } fn node_get_val_mut(&mut self, _key: &[u8]) -> Option<&mut V> { @@ -53,10 +53,10 @@ impl TrieNode for EmptyNode { fn node_set_branch(&mut self, _key: &[u8], _new_node: TrieNodeODRc) -> Result> { unreachable!() //we should head this off upstream } - fn node_remove_all_branches(&mut self, _key: &[u8], _prune: bool) -> bool { + fn node_remove_all_branches(&mut self, _key: &[u8], _prune_limit: usize) -> bool { false } - fn node_remove_unmasked_branches(&mut self, _key: &[u8], _mask: ByteMask, _prune: bool) {} + fn node_remove_unmasked_branches(&mut self, _key: &[u8], _mask: ByteMask, _prune_limit: usize) {} fn node_is_empty(&self) -> bool { true } fn new_iter_token(&self) -> IterToken { 0 @@ -115,7 +115,7 @@ impl TrieNode for EmptyNode { fn get_node_at_key(&self, _key: &[u8]) -> AbstractNodeRef<'_, V, A> { AbstractNodeRef::None } - fn take_node_at_key(&mut self, _key: &[u8], _prune: bool) -> Option> { + fn take_node_at_key(&mut self, _key: &[u8], _prune_limit: usize) -> Option> { None } fn pjoin_dyn(&self, other: TaggedNodeRef) -> AlgebraicResult> where V: Lattice { diff --git a/src/line_list_node.rs b/src/line_list_node.rs index c70b8878..85c8f24e 100644 --- a/src/line_list_node.rs +++ b/src/line_list_node.rs @@ -1803,14 +1803,14 @@ impl TrieNode for LineListNode (result.map(|payload| payload.into_val() ), created_subnode) }) } - fn node_remove_val(&mut self, key: &[u8], prune: bool) -> Option { + fn node_remove_val(&mut self, key: &[u8], prune_limit: usize) -> Option { //Removing a value is one of the ways a node can be left holding two onward children // under one key, so check the node over on the way out let result = (|| { if self.is_used_value_0() { let node_key_0 = unsafe{ self.key_unchecked::<0>() }; if node_key_0 == key { - if prune { + if prune_limit < key.len() { return Some(self.take_payload::<0>().unwrap().into_val()) } else { //If the other slot already keeps this path, then just remove the value @@ -1828,7 +1828,7 @@ impl TrieNode for LineListNode if self.is_used_value_1() { let node_key_1 = unsafe{ self.key_unchecked::<1>() }; if node_key_1 == key { - if prune { + if prune_limit < key.len() { return Some(self.take_payload::<1>().unwrap().into_val()) } else { //If the other slot already keeps this path, then remove the value @@ -1845,6 +1845,9 @@ impl TrieNode for LineListNode } None })(); + if prune_limit > 0 && prune_limit < key.len() && result.is_some() { + self.preserve_prune_limit(key, prune_limit); + } debug_assert!(validate_node(self)); result } @@ -1863,13 +1866,29 @@ impl TrieNode for LineListNode } #[inline] - fn node_remove_dangling(&mut self, key: &[u8]) -> usize { + fn node_remove_dangling(&mut self, key: &[u8], min_keep_len: usize) -> usize { debug_assert!(key.len() > 0); + debug_assert!(min_keep_len <= key.len()); + if min_keep_len >= key.len() { + return 0; + } let (key0, key1) = self.get_both_keys(); if self.is_used_child_0() { if key0 == key { let child = unsafe{ &self.val_or_child0.child }; if child.as_tagged().node_is_empty() { + if min_keep_len > 0 { + let overlap = find_prefix_overlap(key, key1); + if overlap == key.len() { + return 0; + } + if overlap < min_keep_len { + self.shorten_key_len::<0>(min_keep_len); + } else { + let _ = self.take_payload::<0>(); + } + return key.len() - overlap.max(min_keep_len); + } let pruned_bytes = if key1.len() > 0 && key[0] == key1[0] { key.len() - 1 } else { @@ -1884,6 +1903,18 @@ impl TrieNode for LineListNode if key1 == key { let child = unsafe{ &self.val_or_child1.child }; if child.as_tagged().node_is_empty() { + if min_keep_len > 0 { + let overlap = find_prefix_overlap(key, key0); + if overlap == key.len() { + return 0; + } + if overlap < min_keep_len { + self.shorten_key_len::<1>(min_keep_len); + } else { + let _ = self.take_payload::<1>(); + } + return key.len() - overlap.max(min_keep_len); + } let pruned_bytes = if key[0] == key0[0] { key.len() - 1 } else { @@ -1902,17 +1933,18 @@ impl TrieNode for LineListNode result.map(|(_, created_subnode)| created_subnode) } - fn node_remove_all_branches(&mut self, key: &[u8], prune: bool) -> bool { + fn node_remove_all_branches(&mut self, key: &[u8], prune_limit: usize) -> bool { let key_len = key.len(); let (key0, key1) = self.get_both_keys(); let key0_starts_with = starts_with(key0, key); let remove_0 = key0_starts_with && (key0.len() > key_len || self.is_child_ptr::<0>()); let remove_1 = starts_with(key1, key) && (key1.len() > key_len || self.is_child_ptr::<1>()); - self.remove_subtries(remove_0, remove_1, key0_starts_with, prune, key.len()); + self.remove_subtries(remove_0, remove_1, key0_starts_with, prune_limit < key.len(), key.len()); + if prune_limit > 0 && prune_limit < key_len && (remove_0 || remove_1) { self.preserve_prune_limit(key, prune_limit); } remove_0 || remove_1 } - fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune: bool) { + fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune_limit: usize) { let key_len = key.len(); let (key0, key1) = self.get_both_keys(); let mut remove_0 = false; @@ -1935,7 +1967,8 @@ impl TrieNode for LineListNode debug_assert!(!self.is_used_child_1() || unsafe{ self.child_in_slot::<1>().is_empty() }); } } - self.remove_subtries(remove_0, remove_1, key0_starts_with, prune, key.len()); + self.remove_subtries(remove_0, remove_1, key0_starts_with, prune_limit < key.len(), key.len()); + if prune_limit > 0 && prune_limit < key_len && (remove_0 || remove_1) { self.preserve_prune_limit(key, prune_limit); } } fn node_is_empty(&self) -> bool { @@ -2557,56 +2590,60 @@ impl TrieNode for LineListNode AbstractNodeRef::None } - fn take_node_at_key(&mut self, key: &[u8], prune: bool) -> Option> { + fn take_node_at_key(&mut self, key: &[u8], prune_limit: usize) -> Option> { debug_assert!(validate_node(self)); debug_assert!(key.len() > 0); + let result = (|| { - //Exact match with a path to a child node means take that node - let (key0, key1) = self.get_both_keys(); - if self.is_used_child_0() && key0 == key { - if prune { - return self.take_payload::<0>().map(|payload| payload.into_child()) - } else { - let child_payload = self.swap_payload::<0>(ValOrChild::Child(TrieNodeODRc::new_empty())); - return Some(child_payload.into_child()) + //Exact match with a path to a child node means take that node + let (key0, key1) = self.get_both_keys(); + if self.is_used_child_0() && key0 == key { + if prune_limit < key.len() { + return self.take_payload::<0>().map(|payload| payload.into_child()) + } else { + let child_payload = self.swap_payload::<0>(ValOrChild::Child(TrieNodeODRc::new_empty())); + return Some(child_payload.into_child()) + } } - } - if self.is_used_child_1() && key1 == key { - if prune { - return self.take_payload::<1>().map(|payload| payload.into_child()) - } else { - let child_payload = self.swap_payload::<1>(ValOrChild::Child(TrieNodeODRc::new_empty())); - return Some(child_payload.into_child()) + if self.is_used_child_1() && key1 == key { + if prune_limit < key.len() { + return self.take_payload::<1>().map(|payload| payload.into_child()) + } else { + let child_payload = self.swap_payload::<1>(ValOrChild::Child(TrieNodeODRc::new_empty())); + return Some(child_payload.into_child()) + } } - } - //Otherwise check to see if we need to make a sub-node. If we do, - // We know the new node will have only 1 slot filled - if key0.len() > key.len() && starts_with(key0, key) { - let mut new_node = Self::new_in(self.alloc.clone()); - unsafe{ new_node.set_payload_0(&key0[key.len()..], self.is_child_ptr::<0>(), ValOrChildUnion{ _unused: () }) } - new_node.val_or_child0 = if prune { - self.take_payload::<0>().unwrap().into() - } else { - self.shorten_key_len::<0>(key.len()); - self.swap_payload::<0>(ValOrChild::Child(TrieNodeODRc::new_empty())).into() - }; - debug_assert!(validate_node(&new_node)); - return Some(TrieNodeODRc::new_in(new_node, self.alloc.clone())); - } - if key1.len() > key.len() && starts_with(key1, key) { - let mut new_node = Self::new_in(self.alloc.clone()); - unsafe{ new_node.set_payload_0(&key1[key.len()..], self.is_child_ptr::<1>(), ValOrChildUnion{ _unused: () }) } - new_node.val_or_child0 = if prune { - self.take_payload::<1>().unwrap().into() - } else { - self.shorten_key_len::<1>(key.len()); - self.swap_payload::<1>(ValOrChild::Child(TrieNodeODRc::new_empty())).into() - }; - debug_assert!(validate_node(&new_node)); - return Some(TrieNodeODRc::new_in(new_node, self.alloc.clone())); - } - None + //Otherwise check to see if we need to make a sub-node. If we do, + // We know the new node will have only 1 slot filled + if key0.len() > key.len() && starts_with(key0, key) { + let mut new_node = Self::new_in(self.alloc.clone()); + unsafe{ new_node.set_payload_0(&key0[key.len()..], self.is_child_ptr::<0>(), ValOrChildUnion{ _unused: () }) } + new_node.val_or_child0 = if prune_limit < key.len() { + self.take_payload::<0>().unwrap().into() + } else { + self.shorten_key_len::<0>(key.len()); + self.swap_payload::<0>(ValOrChild::Child(TrieNodeODRc::new_empty())).into() + }; + debug_assert!(validate_node(&new_node)); + return Some(TrieNodeODRc::new_in(new_node, self.alloc.clone())); + } + if key1.len() > key.len() && starts_with(key1, key) { + let mut new_node = Self::new_in(self.alloc.clone()); + unsafe{ new_node.set_payload_0(&key1[key.len()..], self.is_child_ptr::<1>(), ValOrChildUnion{ _unused: () }) } + new_node.val_or_child0 = if prune_limit < key.len() { + self.take_payload::<1>().unwrap().into() + } else { + self.shorten_key_len::<1>(key.len()); + self.swap_payload::<1>(ValOrChild::Child(TrieNodeODRc::new_empty())).into() + }; + debug_assert!(validate_node(&new_node)); + return Some(TrieNodeODRc::new_in(new_node, self.alloc.clone())); + } + None + })(); + if result.is_some() { self.preserve_prune_limit(key, prune_limit); } + result } fn pjoin_dyn(&self, other: TaggedNodeRef) -> AlgebraicResult> where V: Lattice { @@ -2934,6 +2971,17 @@ impl TrieNode for LineListNode } impl LineListNode { + #[inline] + fn preserve_prune_limit(&mut self, key: &[u8], prune_limit: usize) { + if prune_limit > 0 && prune_limit < key.len() { + //A compressed key may span the zipper root. Keep that prefix after + //removing the payload below it. + self.node_create_dangling(&key[..prune_limit]).unwrap_or_else(|_| { + unreachable!("removing a payload must leave space for the zipper root") + }); + } + } + /// Part of the implementation of methods the remove subtries from a node fn remove_subtries(&mut self, remove_0: bool, remove_1: bool, key0_starts_with: bool, prune: bool, key_len: usize) { //NOTE: the order here is important because removing slot_0 first might shift the @@ -3413,7 +3461,7 @@ mod tests { let mut new_node = LineListNode::::new_in(global_alloc()); assert_eq!(new_node.node_set_val(&full_key, 24).map_err(|_| 0), Ok((None, false))); - let detached = new_node.take_node_at_key(&prefix, false).unwrap(); + let detached = new_node.take_node_at_key(&prefix, usize::MAX).unwrap(); assert_eq!(detached.as_tagged().node_get_val(suffix), Some(&24)); assert_eq!(new_node.key_len_0(), prefix.len()); diff --git a/src/tiny_node.rs b/src/tiny_node.rs index 6ece3c72..17929493 100644 --- a/src/tiny_node.rs +++ b/src/tiny_node.rs @@ -195,9 +195,9 @@ impl<'a, V: Clone + Send + Sync, A: Allocator> TrieNode for TinyRefNode<'a } None } - fn node_remove_val(&mut self, _key: &[u8], _prune: bool) -> Option { unreachable!() } + fn node_remove_val(&mut self, _key: &[u8], _prune_limit: usize) -> Option { unreachable!() } fn node_create_dangling(&mut self, _key: &[u8]) -> Result<(bool, bool), TrieNodeODRc> { unreachable!() } - fn node_remove_dangling(&mut self, _key: &[u8]) -> usize { unreachable!() } + fn node_remove_dangling(&mut self, _key: &[u8], _min_keep_len: usize) -> usize { unreachable!() } fn node_get_val_mut(&mut self, _key: &[u8]) -> Option<&mut V> { unreachable!() } fn node_set_val(&mut self, key: &[u8], val: V) -> Result<(Option, bool), TrieNodeODRc> { let mut replacement_node = self.into_full().unwrap(); @@ -209,8 +209,8 @@ impl<'a, V: Clone + Send + Sync, A: Allocator> TrieNode for TinyRefNode<'a replacement_node.node_set_branch(key, new_node).unwrap_or_else(|_| panic!()); Err(TrieNodeODRc::new_in(replacement_node, self.alloc.clone())) } - fn node_remove_all_branches(&mut self, _key: &[u8], _prune: bool) -> bool { unreachable!() } - fn node_remove_unmasked_branches(&mut self, _key: &[u8], _mask: ByteMask, _prune: bool) { unreachable!() } + fn node_remove_all_branches(&mut self, _key: &[u8], _prune_limit: usize) -> bool { unreachable!() } + fn node_remove_unmasked_branches(&mut self, _key: &[u8], _mask: ByteMask, _prune_limit: usize) { unreachable!() } fn node_is_empty(&self) -> bool { self.header & (1 << 7) == 0 } @@ -291,7 +291,7 @@ impl<'a, V: Clone + Send + Sync, A: Allocator> TrieNode for TinyRefNode<'a //The key must specify a path the node doesn't contains AbstractNodeRef::None } - fn take_node_at_key(&mut self, _key: &[u8], _prune: bool) -> Option> { unreachable!() } + fn take_node_at_key(&mut self, _key: &[u8], _prune_limit: usize) -> Option> { unreachable!() } fn pjoin_dyn(&self, other: TaggedNodeRef) -> AlgebraicResult> where V: Lattice { //TODO, I can streamline this quite a lot, but for now I'll just up-convert to a ListNode to test // the basic premise of the TinyRefNode @@ -344,4 +344,4 @@ mod tests { //Confirm TinyRefNode is 16 bytes assert_eq!(std::mem::size_of::>(), 16); } -} \ No newline at end of file +} diff --git a/src/trie_node.rs b/src/trie_node.rs index 01b7003d..19b7031b 100644 --- a/src/trie_node.rs +++ b/src/trie_node.rs @@ -132,10 +132,10 @@ pub(crate) trait TrieNode: TrieNodeDowncas /// /// Returns `Some(val)` with the value that was removed, otherwise returns `None` /// - /// If `prune` is `true` this method will prune dangling paths within the node, otherwise - /// it will keep the dangling path. + /// Dangling paths may be pruned down to `prune_limit` bytes of `key`. + /// `usize::MAX` disables pruning. /// WARNING: This method may leave the node empty - fn node_remove_val(&mut self, key: &[u8], prune: bool) -> Option; + fn node_remove_val(&mut self, key: &[u8], prune_limit: usize) -> Option; /// Creates a dangling path up to `key` if none exists. Does nothing if the path already exists /// @@ -152,10 +152,11 @@ pub(crate) trait TrieNode: TrieNodeDowncas /// /// Does nothing and returns 0 if `key` specifies a non-dagling or non-existent path. /// This method will not affect dangling paths other than those specified by `key`. + /// The first `min_keep_len` bytes of `key` must not be removed. /// This method may leave the node empty. /// This method should never be called with a zero-length key. If the `key` arg is longer than the /// keys contained within the node, this method should return `false` - fn node_remove_dangling(&mut self, key: &[u8]) -> usize; + fn node_remove_dangling(&mut self, key: &[u8], min_keep_len: usize) -> usize; /// Sets the downstream branch from the specified `key`. Does not affect the value at the `key` /// @@ -171,16 +172,18 @@ pub(crate) trait TrieNode: TrieNodeDowncas /// Returns `true` if one or more downstream branches were removed from the node; returns `false` if /// the node did not contain any downstream branches from the specified key /// + /// `prune_limit` is the minimum retained prefix length within this node. `usize::MAX` disables pruning. /// WARNING: This method may leave the node empty. If eager pruning of branches is desired then the /// node should subsequently be checked to see if it is empty - fn node_remove_all_branches(&mut self, key: &[u8], prune: bool) -> bool; + fn node_remove_all_branches(&mut self, key: &[u8], prune_limit: usize) -> bool; /// Uses a 256-bit mask to filter down children and values from the specified `key`. Does not affect /// the value at the `key` /// + /// `prune_limit` is the minimum retained prefix length within this node. `usize::MAX` disables pruning. /// WARNING: This method may leave the node empty. If eager pruning of branches is desired then the /// node should subsequently be checked to see if it is empty - fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune: bool); + fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune_limit: usize); /// Returns `true` if the node contains no children nor values, otherwise false fn node_is_empty(&self) -> bool; @@ -357,7 +360,8 @@ pub(crate) trait TrieNode: TrieNodeDowncas /// WARNING: This method may leave the node empty /// /// This method should never be called with `key.len() == 0` - fn take_node_at_key(&mut self, key: &[u8], prune: bool) -> Option>; + /// `prune_limit` is the minimum retained prefix length within this node. `usize::MAX` disables pruning. + fn take_node_at_key(&mut self, key: &[u8], prune_limit: usize) -> Option>; /// Allows for the implementation of the Lattice trait on different node implementations, and /// the logic to promote nodes to other node types @@ -1707,11 +1711,11 @@ mod tagged_node_ref { } } - pub fn node_remove_dangling(&mut self, key: &[u8]) -> usize { + pub fn node_remove_dangling(&mut self, key: &[u8], min_keep_len: usize) -> usize { match self { - Self::DenseByteNode(node) => node.node_remove_dangling(key), - Self::LineListNode(node) => node.node_remove_dangling(key), - Self::CellByteNode(node) => node.node_remove_dangling(key), + Self::DenseByteNode(node) => node.node_remove_dangling(key, min_keep_len), + Self::LineListNode(node) => node.node_remove_dangling(key, min_keep_len), + Self::CellByteNode(node) => node.node_remove_dangling(key, min_keep_len), } } @@ -1749,11 +1753,11 @@ mod tagged_node_ref { } } - pub fn node_remove_val(&mut self, key: &[u8], prune: bool) -> Option { + pub fn node_remove_val(&mut self, key: &[u8], prune_limit: usize) -> Option { match self { - Self::DenseByteNode(node) => node.node_remove_val(key, prune), - Self::LineListNode(node) => node.node_remove_val(key, prune), - Self::CellByteNode(node) => node.node_remove_val(key, prune), + Self::DenseByteNode(node) => node.node_remove_val(key, prune_limit), + Self::LineListNode(node) => node.node_remove_val(key, prune_limit), + Self::CellByteNode(node) => node.node_remove_val(key, prune_limit), } } @@ -1765,26 +1769,26 @@ mod tagged_node_ref { } } - pub fn node_remove_all_branches(&mut self, key: &[u8], prune: bool) -> bool { + pub fn node_remove_all_branches(&mut self, key: &[u8], prune_limit: usize) -> bool { match self { - Self::DenseByteNode(node) => node.node_remove_all_branches(key, prune), - Self::LineListNode(node) => node.node_remove_all_branches(key, prune), - Self::CellByteNode(node) => node.node_remove_all_branches(key, prune), + Self::DenseByteNode(node) => node.node_remove_all_branches(key, prune_limit), + Self::LineListNode(node) => node.node_remove_all_branches(key, prune_limit), + Self::CellByteNode(node) => node.node_remove_all_branches(key, prune_limit), } } - pub fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune: bool) { + pub fn node_remove_unmasked_branches(&mut self, key: &[u8], mask: ByteMask, prune_limit: usize) { match self { - Self::DenseByteNode(node) => node.node_remove_unmasked_branches(key, mask, prune), - Self::LineListNode(node) => node.node_remove_unmasked_branches(key, mask, prune), - Self::CellByteNode(node) => node.node_remove_unmasked_branches(key, mask, prune), + Self::DenseByteNode(node) => node.node_remove_unmasked_branches(key, mask, prune_limit), + Self::LineListNode(node) => node.node_remove_unmasked_branches(key, mask, prune_limit), + Self::CellByteNode(node) => node.node_remove_unmasked_branches(key, mask, prune_limit), } } - pub fn take_node_at_key(&mut self, key: &[u8], prune: bool) -> Option> { + pub fn take_node_at_key(&mut self, key: &[u8], prune_limit: usize) -> Option> { match self { - Self::DenseByteNode(node) => node.take_node_at_key(key, prune), - Self::LineListNode(node) => node.take_node_at_key(key, prune), - Self::CellByteNode(node) => node.take_node_at_key(key, prune), + Self::DenseByteNode(node) => node.take_node_at_key(key, prune_limit), + Self::LineListNode(node) => node.take_node_at_key(key, prune_limit), + Self::CellByteNode(node) => node.take_node_at_key(key, prune_limit), } } pub fn join_into_dyn(&mut self, other: TrieNodeODRc) -> (AlgebraicStatus, Result<(), TrieNodeODRc>) where V: Lattice { diff --git a/src/write_zipper.rs b/src/write_zipper.rs index 1da1381f..28781b44 100644 --- a/src/write_zipper.rs +++ b/src/write_zipper.rs @@ -1276,7 +1276,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC let alloc = self.alloc.clone(); let sub_branch_added = self.in_zipper_mut_static_result( |node, key| { - let new_node = if let Some(remaining) = node.take_node_at_key(key, false) { + let new_node = if let Some(remaining) = node.take_node_at_key(key, usize::MAX) { remaining } else { #[cfg(not(feature = "all_dense_nodes"))] @@ -1456,6 +1456,15 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC pub fn set_val(&mut self, val: V) -> Option { self.set_val_at(&[], val) } + #[inline(always)] + fn node_prune_limit(&self, prune: bool) -> usize { + if prune { + self.key.origin_path.len().saturating_sub(self.key.node_key_start()) + } else { + usize::MAX + } + } + /// See [ZipperWriting::remove_val] pub fn remove_val(&mut self, prune: bool) -> Option { if self.key.node_key().len() == 0 { @@ -1463,8 +1472,9 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC let root_val_ref = self.root_val.as_mut().unwrap(); return core::mem::take(unsafe{&mut **root_val_ref}) } + let prune_limit = self.node_prune_limit(prune); let mut focus_node = self.focus_stack.top_mut().unwrap(); - if let Some(result) = focus_node.node_remove_val(self.key.node_key(), prune) { + if let Some(result) = focus_node.node_remove_val(self.key.node_key(), prune_limit) { if prune { self.prune_path_internal(false); } @@ -2245,8 +2255,9 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC pub fn remove_branches(&mut self, prune: bool) -> bool { let node_key = self.key.node_key(); if node_key.len() > 0 { + let prune_limit = self.node_prune_limit(prune); let mut focus_node = self.focus_stack.top_mut().unwrap(); - if focus_node.node_remove_all_branches(node_key, prune) { + if focus_node.node_remove_all_branches(node_key, prune_limit) { if prune { self.prune_path_internal(false); } @@ -2285,26 +2296,27 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC /// See [WriteZipper::remove_unmasked_branches] pub fn remove_unmasked_branches(&mut self, mask: ByteMask, prune: bool) { let node_key = self.key.node_key(); + let prune_limit = self.node_prune_limit(prune); let mut focus_node = self.focus_stack.top_mut().unwrap(); if node_key.len() > 0 { match focus_node.node_get_child_mut(node_key) { Some((consumed_bytes, child_node)) => { if node_key.len() >= consumed_bytes && !child_node.is_empty() { - child_node.make_mut().node_remove_unmasked_branches(&node_key[consumed_bytes..], mask, prune); + child_node.make_mut().node_remove_unmasked_branches(&node_key[consumed_bytes..], mask, prune_limit.saturating_sub(consumed_bytes)); if child_node.as_tagged().node_is_empty() { - focus_node.node_remove_all_branches(&node_key[..consumed_bytes], prune); + focus_node.node_remove_all_branches(&node_key[..consumed_bytes], prune_limit); } } else { //Zipper is positioned at non-existent or dangling node. Removing anything from nothing is nothing } }, None => { - focus_node.node_remove_unmasked_branches(node_key, mask, prune); + focus_node.node_remove_unmasked_branches(node_key, mask, prune_limit); } } } else { debug_assert!(self.key.prefix_buf.len() <= self.key.origin_path.len()); //Equivalent to `self.at_root()`, but can't borrow `self` here - focus_node.node_remove_unmasked_branches(node_key, mask, prune); + focus_node.node_remove_unmasked_branches(node_key, mask, prune_limit); } if prune { self.prune_path_internal(false); @@ -2332,7 +2344,8 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC pub(crate) fn prune_path(&mut self) -> usize { let key = self.key.node_key(); if key.len() > 0 { - let node_pruned_bytes = self.focus_stack.top_mut().unwrap().node_remove_dangling(key); + let min_keep_len = self.key.origin_path.len().saturating_sub(self.key.node_key_start()); + let node_pruned_bytes = self.focus_stack.top_mut().unwrap().node_remove_dangling(key, min_keep_len); let trie_pruned_bytes = if node_pruned_bytes > 0 { self.prune_path_internal(false) } else { 0 }; @@ -2352,6 +2365,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC /// Internal method, Removes and returns the node at the zipper's focus. This method may leave behind a dangling path #[inline] fn take_focus(&mut self, prune: bool) -> Option> { + let prune_limit = self.node_prune_limit(prune); let mut focus_node = self.focus_stack.top_mut().unwrap(); let node_key = self.key.node_key(); if node_key.len() == 0 { @@ -2366,7 +2380,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC None } } else { - if let Some(new_node) = focus_node.take_node_at_key(node_key, prune) { + if let Some(new_node) = focus_node.take_node_at_key(node_key, prune_limit) { if prune { self.prune_path_internal(false); } @@ -2404,7 +2418,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC let sub_branch_added = self.in_zipper_mut_static_result( |node, key| { // A graft replaces everything below the focus - node.node_remove_all_branches(key, false); + node.node_remove_all_branches(key, usize::MAX); node.node_set_branch(key, src) }, |_, _| true); @@ -2477,7 +2491,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC // 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(|| { + let mut child = focus_node.take_node_at_key(key, usize::MAX).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")] @@ -2520,15 +2534,18 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC } else { &self.key.prefix_buf[..] }; + let root_len = self.key.origin_path.len(); let mut temp_path = path_buf; let mut ascended = false; let mut just_popped = false; let mut node_key_end = temp_path.len(); + let mut stopped_at_zipper_root = false; //This loop mirrors the behavior of `ascend_until`, popping from the node stack but leaving the path buffer alone loop { - debug_assert!(temp_path.len() >= self.key.origin_path.len()); - if temp_path.len() == 0 || temp_path.len() == self.key.origin_path.len() { + debug_assert!(temp_path.len() >= root_len); + if temp_path.len() == root_len { + stopped_at_zipper_root = root_len > 0; break } let node_key_start = self.key.node_key_start(); @@ -2536,7 +2553,7 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC //This mirrors the logic of `ascend_within_node`, but using our alternative path buffer let branch_key = self.focus_stack.top().unwrap().prior_branch_key(node_key); - let new_len = self.key.origin_path.len().max(node_key_start + branch_key.len()); + let new_len = root_len.max(node_key_start + branch_key.len()); ascended = true; temp_path = &temp_path[..new_len]; @@ -2581,27 +2598,31 @@ impl <'a, 'path, V: Clone + Send + Sync + Unpin, A: Allocator + 'a> WriteZipperC if ascended { let mut focus_node = self.focus_stack.top_mut().unwrap(); let node_key_start = self.key.node_key_start(); + if stopped_at_zipper_root { + node_key_end = temp_path.len(); + } let next_node_key = &path_buf[node_key_start..node_key_end]; //The path to the node or subnode we need to remove might not be within the focus node, // so get the actual node that we want to remove the contents from - let (mut container_node, next_node_key) = match focus_node.node_get_child_mut(next_node_key) { + let (mut container_node, next_node_key, consumed) = match focus_node.node_get_child_mut(next_node_key) { Some((consumed_bytes, new_focus)) => { if consumed_bytes < next_node_key.len() { - (new_focus.make_mut(), &next_node_key[consumed_bytes..]) + (new_focus.make_mut(), &next_node_key[consumed_bytes..], consumed_bytes) } else { - (focus_node, next_node_key) + (focus_node, next_node_key, 0) } }, - None => (focus_node, next_node_key) + None => (focus_node, next_node_key, 0) }; - let removed = container_node.node_remove_all_branches(next_node_key, true); + let prune_limit = if stopped_at_zipper_root { root_len.saturating_sub(node_key_start + consumed) } else { 0 }; + let removed = container_node.node_remove_all_branches(next_node_key, prune_limit); //If we got here, we should have either removed something, or we should be at the top of the zipper debug_assert!(removed || self.focus_stack.depth()==1); } - debug_assert!(temp_path.len() >= self.key.origin_path.len()); + debug_assert!(temp_path.len() >= root_len); let pruned_bytes = path_buf.len() - temp_path.len(); if should_ascend { @@ -2720,7 +2741,7 @@ pub(crate) fn swap_top_node<'cursor, V: Clone + Send + Sync, A: Allocator + 'cur focus_stack.backtrack(); let mut parent_node = unsafe{ focus_stack.top_mut().unwrap_unchecked() }; let parent_key = key.parent_key(); - let existing_node = parent_node.take_node_at_key(parent_key, false).unwrap(); + let existing_node = parent_node.take_node_at_key(parent_key, usize::MAX).unwrap(); let replacement_node = func(existing_node); parent_node.node_set_branch(parent_key, replacement_node).unwrap(); focus_stack.advance(|node| node.node_get_child_mut(parent_key).map(|(_, child_node)| child_node.make_mut())); @@ -3749,6 +3770,231 @@ mod tests { assert_eq!(btm2.path_exists_at(&[0, 255, 1]), false); } + /// A write zipper must not modify trie above its root via prune + #[test] + fn prune_should_not_cross_zipper_root() { + let mut map = PathMap::::new(); + map.create_path(b"ab"); + let pruned = map.write_zipper_at_path(b"ab").prune_path(); + assert!(map.path_exists_at(b"ab")); + assert_eq!(pruned, 0); + + let mut map = PathMap::::new(); + map.create_path(b"abcd"); + let mut wz = map.write_zipper_at_path(b"ab"); + wz.descend_to(b"cd"); + let pruned = wz.prune_path(); + assert!(map.path_exists_at(b"ab")); + assert!(!map.path_exists_at(b"abcd")); + assert_eq!(pruned, 2); + + // Two values split at "a". Root the zipper above, at, and below the split. + for (zipper_root, expected_pruned) in [ + (b"".as_slice(), 2), + (b"a".as_slice(), 2), + (b"ab".as_slice(), 1), + ] { + let mut map = PathMap::::new(); + map.set_val_at(b"abc", 1); + map.set_val_at(b"axd", 2); + #[cfg(not(feature = "all_dense_nodes"))] + { + let (consumed, pair_node) = map.root().unwrap().as_tagged().node_get_child(b"a").unwrap(); + assert_eq!(consumed, 1); + assert!(pair_node.as_tagged().as_list().is_some()); + } + + let mut wz = map.write_zipper_at_path(zipper_root); + wz.descend_to(&b"abc"[zipper_root.len()..]); + assert_eq!(wz.remove_val(false), Some(1)); + assert_eq!(wz.prune_path(), expected_pruned, "zipper_root={zipper_root:?}"); + assert!(map.path_exists_at(zipper_root), "zipper_root={zipper_root:?}"); + assert!(!map.path_exists_at(b"abc"), "zipper_root={zipper_root:?}"); + assert_eq!(map.get(b"axd"), Some(&2)); + } + + // Three root branches force a ByteNode. Keep the zipper root after pruning its last child. + let mut map = PathMap::::new(); + for (path, val) in [(b"a0", 1), (b"b0", 2), (b"c0", 3)] { + map.set_val_at(path, val); + } + assert!(map.root().unwrap().as_tagged().as_dense().is_some()); + let mut wz = map.write_zipper_at_path(b"a"); + wz.descend_to(b"0"); + assert_eq!(wz.remove_val(false), Some(1)); + assert_eq!(wz.prune_path(), 1); + wz.reset(); + assert_eq!(wz.prune_path(), 0); + assert!(map.path_exists_at(b"a")); + assert!(!map.path_exists_at(b"a0")); + assert_eq!(map.get(b"b0"), Some(&2)); + assert_eq!(map.get(b"c0"), Some(&3)); + } + + #[test] + fn prune_should_not_cross_zipper_root_across_nodes() { + let path: Vec = (0..100).map(|i| i as u8).collect(); + for root_len in [1, 5, 16, 32, 47, 48, 49, 64, 95] { + let mut map = PathMap::::new(); + map.create_path(&path); + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + assert_eq!(wz.prune_path(), path.len() - root_len, "root_len={root_len}"); + assert!(map.path_exists_at(&path[..root_len]), "root_len={root_len}"); + assert!(!map.path_exists_at(&path), "root_len={root_len}"); + } + } + + #[test] + fn prune_preserves_zipper_root_with_sibling_paths() { + for sibling in [&b"ax"[..], &b"abef"[..]] { + let mut map = PathMap::::new(); + map.create_path(b"abcd"); + map.set_val_at(sibling, 1); + let mut wz = map.write_zipper_at_path(b"ab"); + wz.descend_to(b"cd"); + assert_eq!(wz.prune_path(), 2); + assert!(map.path_exists_at(b"ab")); + assert!(!map.path_exists_at(b"abcd")); + assert_eq!(map.get(sibling), Some(&1)); + } + } + + #[test] + fn prune_flags_preserve_zipper_root() { + type Check = fn(&[u8], usize) -> (bool, bool, bool); + let methods: &[(&str, Check)] = &[ + ("remove_val", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + let removed = wz.remove_val(true); + (removed == Some(1), map.path_exists_at(&path[..root_len]), !map.path_exists_at(path)) + }), + ("remove_branches", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let focus = &path[..path.len() - 1]; + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&focus[root_len..]); + let removed = wz.remove_branches(true); + (removed, map.path_exists_at(&path[..root_len]), !map.path_exists_at(focus) && !map.path_exists_at(path)) + }), + ("remove_unmasked_branches", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let focus = &path[..path.len() - 1]; + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&focus[root_len..]); + wz.remove_unmasked_branches(ByteMask::EMPTY, true); + (true, map.path_exists_at(&path[..root_len]), !map.path_exists_at(focus) && !map.path_exists_at(path)) + }), + ("take_map", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let focus = &path[..path.len() - 1]; + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&focus[root_len..]); + let taken = wz.take_map(true); + (taken.is_some(), map.path_exists_at(&path[..root_len]), !map.path_exists_at(path)) + }), + ("join_into_take", |path, root_len| { + let mut source = PathMap::::new(); + source.set_val_at(path, 1); + let focus = &path[..path.len() - 1]; + let mut src_wz = source.write_zipper_at_path(&path[..root_len]); + src_wz.descend_to(&focus[root_len..]); + let mut destination = PathMap::::new(); + let mut dst_wz = destination.write_zipper_at_path(b"z"); + let status = dst_wz.join_into_take(&mut src_wz, true); + (status == AlgebraicStatus::Element, source.path_exists_at(&path[..root_len]), !source.path_exists_at(path)) + }), + ("meet_k_path_into", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let focus = &path[..path.len() - 1]; + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&focus[root_len..]); + let result = wz.meet_k_path_into(2, true); + (!result, map.path_exists_at(&path[..root_len]), !map.path_exists_at(path)) + }), + ("meet_into", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let empty = PathMap::::new(); + let rz = empty.read_zipper(); + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + let status = wz.meet_into(&rz, true); + (status == AlgebraicStatus::None, map.path_exists_at(&path[..root_len]), !map.path_exists_at(path)) + }), + ("subtract_into", |path, root_len| { + let mut map = PathMap::::new(); + map.set_val_at(path, 1); + let mut source = PathMap::::new(); + source.set_val_at(b"", 1); + let rz = source.read_zipper(); + let mut wz = map.write_zipper_at_path(&path[..root_len]); + wz.descend_to(&path[root_len..]); + let status = wz.subtract_into(&rz, true); + (status == AlgebraicStatus::None, map.path_exists_at(&path[..root_len]), !map.path_exists_at(path)) + }), + ]; + let long_path: Vec = (0..100).map(|i| i as u8).collect(); + let cases = [(b"abcd".as_slice(), 2), (long_path.as_slice(), 5), (long_path.as_slice(), 95)]; + let mut failures = Vec::new(); + for &(method, check) in methods { + for &(path, root_len) in &cases { + let (operation_succeeded, root_preserved, target_removed) = check(path, root_len); + if !operation_succeeded || !root_preserved || !target_removed { + failures.push(format!("{method}: path_len={}, root_len={root_len}: operation_succeeded={operation_succeeded}, root_preserved={root_preserved}, target_removed={target_removed}", path.len())); + } + } + } + assert!(failures.is_empty(), "{failures:?}"); + } + + #[test] + fn prune_flags_preserve_sibling_paths() { + for root in [b"".as_slice(), b"a", b"ab"] { + for remove_branches in [false, true] { + let mut map = PathMap::::new(); + map.set_val_at(b"abc", 1); + map.set_val_at(b"axd", 2); + let mut wz = map.write_zipper_at_path(root); + if remove_branches { + wz.descend_to(&b"ab"[root.len()..]); + assert!(wz.remove_branches(true)); + } else { + wz.descend_to(&b"abc"[root.len()..]); + assert_eq!(wz.remove_val(true), Some(1)); + } + assert!(map.path_exists_at(root), "root={root:?}, remove_branches={remove_branches}"); + assert!(!map.path_exists_at(b"abc")); + assert_eq!(map.get(b"axd"), Some(&2)); + } + } + + for remove_branches in [false, true] { + let mut map = PathMap::::new(); + for (path, val) in [(b"a0", 1), (b"b0", 2), (b"c0", 3)] { + map.set_val_at(path, val); + } + let mut wz = map.write_zipper_at_path(b"a"); + if remove_branches { + assert!(wz.remove_branches(true)); + } else { + wz.descend_to(b"0"); + assert_eq!(wz.remove_val(true), Some(1)); + } + assert!(map.path_exists_at(b"a")); + assert!(!map.path_exists_at(b"a0")); + assert_eq!(map.get(b"b0"), Some(&2)); + assert_eq!(map.get(b"c0"), Some(&3)); + } + } + /// A write after `prune_path` (or `meet_into(.., true)`) must reach the focus node #[test] fn write_zipper_write_after_prune_path_below_a_graft() { @@ -5654,9 +5900,9 @@ mod tests { assert_eq!(wz.child_count(), 0); assert_eq!(wz.child_mask(), ByteMask::EMPTY); - //Finally, prune again, and make sure that did what it was supposed to do - wz.prune_path(); - assert_eq!(wz.path_exists(), false); + //The zipper root remains even after its value is removed. + assert_eq!(wz.prune_path(), 0); + assert_eq!(wz.path_exists(), true); assert_eq!(wz.is_val(), false); assert_eq!(wz.child_count(), 0); assert_eq!(wz.child_mask(), ByteMask::EMPTY); diff --git a/src/zipper_head.rs b/src/zipper_head.rs index f9e50888..20357a74 100644 --- a/src/zipper_head.rs +++ b/src/zipper_head.rs @@ -385,7 +385,7 @@ pub(crate) fn prepare_exclusive_write_path<'a, 'trie: 'a, 'path: 'a, V: Clone + z.in_zipper_mut_static_result( |node, key| { let new_node = if key.len() > 0 { - if let Some(mut remaining) = node.take_node_at_key(key, false) { + if let Some(mut remaining) = node.take_node_at_key(key, usize::MAX) { make_cell_node(&mut remaining, alloc.clone()); remaining } else { @@ -438,7 +438,7 @@ fn prepare_node_at_path_end<'a, V: Clone + Send + Sync, A: Allocator>(start_node //If remaining_key is non-zero length, split and upgrade the intervening node if remaining_key.len() > 0 { let mut node_ref = node.make_mut(); - let mut new_parent = match node_ref.take_node_at_key(remaining_key, false) { + let mut new_parent = match node_ref.take_node_at_key(remaining_key, usize::MAX) { Some(downward_node) => downward_node, None => TrieNodeODRc::new_in(CellByteNode::new_in(alloc.clone()), alloc.clone()) }; diff --git a/src/zipper_tracking.rs b/src/zipper_tracking.rs index 884e345f..afefd1e8 100644 --- a/src/zipper_tracking.rs +++ b/src/zipper_tracking.rs @@ -4,8 +4,8 @@ use std::num::NonZeroU32; use std::sync::Arc; use std::sync::RwLock; -use crate::PathMap; -use crate::zipper::{ReadZipperUntracked, Zipper, ZipperAbsolutePath, ZipperForking, ZipperMoving, ZipperPath, ZipperReadOnlyValues, ZipperWriting, ZipperIteration, ZipperReadOnlyIteration, }; +use crate::write_zipper::WriteZipperOwned; +use crate::zipper::{Zipper, ZipperAbsolutePath, ZipperMoving, ZipperPath, ZipperValues, ZipperWriting, ZipperIteration, }; /// Marker to track an outstanding read zipper pub struct TrackingRead; @@ -108,14 +108,14 @@ impl Conflict { } } - fn check_for_lock_along_path<'a, A: Clone + Send + Sync + Unpin>( + fn check_for_lock_along_path( path: &[u8], - zipper: &'a mut ReadZipperUntracked, - ) -> Option<&'a A> { + zipper: &mut WriteZipperOwned, + ) -> Option { let mut current_path = path; loop { if zipper.is_val() { - return zipper.get_val(); + return zipper.val().cloned(); } else if current_path.is_empty() { return None; } else { @@ -128,18 +128,17 @@ impl Conflict { } } - fn check_for_write_conflictC>(path: &[u8], all_paths: &PathMap<()>, conflict_f: ConflictF) -> Result<(), C> { - let mut zipper = all_paths.read_zipper(); - match Conflict::check_for_lock_along_path(path, &mut zipper) { + fn check_for_write_conflictC>(path: &[u8], zipper: &mut WriteZipperOwned<()>, conflict_f: ConflictF) -> Result<(), C> { + zipper.reset(); + match Conflict::check_for_lock_along_path(path, zipper) { None => /* at this point zipper is either focued on the given path (when it exists) , or the procedure broke out early, because it was determined that the path does not exist */ { if zipper.depth() == path.len() { - let mut subtree = zipper.fork_read_zipper(); - match subtree.to_next_val() { + match zipper.to_next_val() { false => Ok(()), - true => Err(conflict_f(subtree.origin_path())), + true => Err(conflict_f(zipper.origin_path())), } } else { Ok(()) @@ -151,26 +150,22 @@ impl Conflict { fn check_for_read_conflictC>( path: &[u8], - all_paths: &PathMap, + zipper: &mut WriteZipperOwned, conflict_f: ConflictF ) -> Result<(), C> { - let mut zipper = all_paths.read_zipper(); - match Conflict::check_for_lock_along_path(path, &mut zipper) { + zipper.reset(); + match Conflict::check_for_lock_along_path(path, zipper) { None => { if zipper.depth() == path.len() { - let mut subtree = zipper.fork_read_zipper(); - match subtree.to_next_get_val() { - None => Ok(()), - Some(lock) => Err(conflict_f( - *lock, - subtree.origin_path(), - )), + match zipper.to_next_val() { + false => Ok(()), + true => Err(conflict_f(*zipper.val().unwrap(), zipper.origin_path())), } } else { Ok(()) } } - Some(lock) => Err(conflict_f(*lock, zipper.path())), + Some(lock) => Err(conflict_f(lock, zipper.path())), } } @@ -194,10 +189,15 @@ impl Conflict { #[derive(Clone, Default)] pub struct SharedTrackerPaths(Arc>); -#[derive(Clone, Default)] struct TrackerPaths { - read_paths: PathMap, - written_paths: PathMap<()>, + read_paths: WriteZipperOwned, + written_paths: WriteZipperOwned<()>, +} + +impl Default for TrackerPaths { + fn default() -> Self { + Self { read_paths: WriteZipperOwned::new(), written_paths: WriteZipperOwned::new() } + } } /// Represents the status of a specific path, returned by [SharedTrackerPaths::path_status] @@ -231,11 +231,11 @@ impl SharedTrackerPaths { pub fn path_status>(&self, path: P) -> PathStatus { let path = path.as_ref(); self.with_paths(|all_paths: &mut TrackerPaths| { - match Conflict::check_for_write_conflict(path, &all_paths.written_paths, |_| ()) { + match Conflict::check_for_write_conflict(path, &mut all_paths.written_paths, |_| ()) { Ok(()) => {}, Err(()) => return PathStatus::Unavailable } - match Conflict::check_for_read_conflict(path, &all_paths.read_paths, |_, _| ()) { + match Conflict::check_for_read_conflict(path, &mut all_paths.read_paths, |_, _| ()) { Ok(()) => {}, Err(()) => return PathStatus::AvailableForReading } @@ -245,9 +245,11 @@ impl SharedTrackerPaths { fn try_add_writer(&self, path: &[u8]) -> Result<(), Conflict> { let try_add_writer_internal = |all_paths: &mut TrackerPaths| { - Conflict::check_for_write_conflict(path, &all_paths.written_paths, Conflict::write_conflict)?; - Conflict::check_for_read_conflict(path, &all_paths.read_paths, Conflict::read_conflict)?; - let mut writer = all_paths.written_paths.write_zipper_at_path(path); + Conflict::check_for_write_conflict(path, &mut all_paths.written_paths, Conflict::write_conflict)?; + Conflict::check_for_read_conflict(path, &mut all_paths.read_paths, Conflict::read_conflict)?; + let writer = &mut all_paths.written_paths; + writer.reset(); + writer.descend_to(path); writer.set_val(()); Ok(()) }; @@ -257,8 +259,10 @@ impl SharedTrackerPaths { fn try_add_reader(&self, path: &[u8]) -> Result<(), Conflict> { let try_add_reader_internal = |all_paths: &mut TrackerPaths| { - Conflict::check_for_write_conflict(path, &all_paths.written_paths, Conflict::write_conflict)?; - let mut writer = all_paths.read_paths.write_zipper_at_path(path); + Conflict::check_for_write_conflict(path, &mut all_paths.written_paths, Conflict::write_conflict)?; + let writer = &mut all_paths.read_paths; + writer.reset(); + writer.descend_to(path); let value = writer.get_val_mut(); match value { Some(cnt) => match cnt.checked_add(1) { @@ -281,7 +285,9 @@ impl SharedTrackerPaths { /// Adds a new reader without checking to see whether it conflicts with existing writers fn add_reader_unchecked(&self, path: &[u8]) { let add_reader = |paths: &mut TrackerPaths| { - let mut writer = paths.read_paths.write_zipper_at_path(path); + let writer = &mut paths.read_paths; + writer.reset(); + writer.descend_to(path); match writer.get_val_mut() { Some(cnt) => { *cnt = unsafe { NonZero::new_unchecked(cnt.get() + 1) }; @@ -312,11 +318,15 @@ impl core::fmt::Debug for ZipperTracker { self.this_path ); let _ = writeln!(f, "\tRead Zippers:"); - for (rz, cnt) in all_paths.read_paths.iter() { + let mut read_paths = all_paths.read_paths.clone(); + read_paths.reset(); + for (rz, cnt) in read_paths.into_iter() { let _ = writeln!(f, "\t\t{rz:?} ({cnt:?})"); } let _ = writeln!(f, "\tWrite Zippers:"); - for (wz, _) in all_paths.written_paths.iter() { + let mut written_paths = all_paths.written_paths.clone(); + written_paths.reset(); + for (wz, _) in written_paths.into_iter() { let _ = writeln!(f, "\t\t{wz:?}"); } write!(f, "}}") @@ -373,7 +383,9 @@ impl ZipperTracker { fn remove_lock(all_paths: &SharedTrackerPaths, this_path: &[u8]) { let is_removed = all_paths.with_paths(|paths| { if M::tracks_reads() { - let mut write_zipper = paths.read_paths.write_zipper_at_path(this_path); + let write_zipper = &mut paths.read_paths; + write_zipper.reset(); + write_zipper.descend_to(this_path); match write_zipper.get_val_mut() { Some(cnt) => { if *cnt == NonZero::::MIN { @@ -386,10 +398,10 @@ impl ZipperTracker { None => false, } } else { - let removed = paths - .written_paths - .write_zipper_at_path(this_path) - .remove_val(true); + let write_zipper = &mut paths.written_paths; + write_zipper.reset(); + write_zipper.descend_to(this_path); + let removed = write_zipper.remove_val(true); removed.is_some() } }); @@ -404,3 +416,34 @@ impl Drop for ZipperTracker { Self::remove_lock(&self.all_paths, &self.this_path); } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn persistent_tracker_zippers_check_conflicts_and_prune_released_paths() { + let paths = SharedTrackerPaths::default(); + let first = ZipperTracker::::new(paths.clone(), b"a/one").unwrap(); + let second = ZipperTracker::::new(paths.clone(), b"b/two").unwrap(); + + assert!(ZipperTracker::::new(paths.clone(), b"a").is_err()); + assert!(ZipperTracker::::new(paths.clone(), b"a/one/child").is_err()); + assert!(ZipperTracker::::new(paths.clone(), b"b").is_err()); + let debug = format!("{first:?}"); + assert!(debug.contains("[97, 47, 111, 110, 101]")); + assert!(debug.contains("[98, 47, 116, 119, 111]")); + + drop(first); + paths.with_paths(|all| { + all.written_paths.reset(); + all.written_paths.descend_to(b"a/one"); + assert!(!all.written_paths.path_exists()); + }); + let reader = ZipperTracker::::new(paths.clone(), b"a/one").unwrap(); + assert!(ZipperTracker::::new(paths.clone(), b"a/one").is_err()); + drop(reader); + drop(second); + assert!(ZipperTracker::::new(paths, b"a").is_ok()); + } +}