diff --git a/src/bitmap.rs b/src/bitmap.rs index e7fabfa..cff0890 100644 --- a/src/bitmap.rs +++ b/src/bitmap.rs @@ -42,12 +42,23 @@ impl Bitmap { self.0 |= 1 << id.0; } + /// Clears the bit at `id`. + pub(crate) fn clear(&mut self, id: StrideId) { + self.0 &= !(1 << id.0); + } + /// Returns the number of entries in the bitmap. #[inline(always)] pub(crate) fn pop_count(&self) -> u32 { self.0.count_ones() } + /// Returns true if the bitmap is empty. + #[inline(always)] + pub(crate) fn is_empty(&self) -> bool { + self.0 == 0 + } + /// Returns a vec with the sorted positions of the populated bits. pub fn bit_positions(&self) -> Vec { let mut bitmap = self.0; diff --git a/src/lib.rs b/src/lib.rs index 684758e..5d88491 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -34,7 +34,6 @@ use alloc::collections::btree_map::BTreeMap; use alloc::vec; use alloc::vec::Vec; use bitmap::*; -use core::cmp::min; use value_index::ValueIndex; /// The maximum number of bits we can consume from the prefix at a time. @@ -312,8 +311,12 @@ where /// assert!(!trie.contains_key((u32::from_be_bytes([192, 168, 0, 0]), 16))); /// ``` pub fn contains_key(&self, prefix: P) -> bool { - let (parent_node, prefix_id, _) = self.find_parent_node(prefix); - self.entries[parent_node].contains_key(&prefix_id) + self.find_parent_node(prefix).map_or_else( + || false, + |(parent_node, prefix_id)| { + self.entries[parent_node].contains_key(&prefix_id) + }, + ) } /// Returns the number of entries in the trie. @@ -365,7 +368,7 @@ where /// assert_eq!(trie.get((u32::from_be_bytes([10, 0, 0, 0]), 16)), None); /// ``` pub fn get(&self, prefix: P) -> Option<&V> { - let (parent_node, prefix_id, _) = self.find_parent_node(prefix); + let (parent_node, prefix_id) = self.find_parent_node(prefix)?; self.entries[parent_node] .get(&prefix_id) .and_then(|(_, vi)| vi.get().map(|i| &self.values[i])) @@ -389,7 +392,7 @@ where /// assert_eq!(trie.get((u32::from_be_bytes([10, 0, 0, 0]), 8)), Some(&10)); /// ``` pub fn get_mut(&mut self, prefix: P) -> Option<&mut V> { - let (parent_node, prefix_id, _) = self.find_parent_node(prefix); + let (parent_node, prefix_id) = self.find_parent_node(prefix)?; self.entries[parent_node] .get(&prefix_id) .and_then(|(_, vi)| vi.get()) @@ -511,48 +514,40 @@ where /// assert!(!trie.contains_key((u32::from_be_bytes([10, 1, 0, 0]), 16))); /// ``` pub fn remove(&mut self, prefix: P) -> Option { - let (parent_node, prefix_id, default_value_index) = - self.find_parent_node(prefix); - - self.entries[parent_node].remove(&prefix_id).map(|(_, v)| { - // Update the leaf ranges - self.calculate_leaf_ranges(parent_node, default_value_index); - - // Update the value indices in all the leaves and entries - for higher_v in self - .leaves - .iter_mut() - .chain( - self.entries - .iter_mut() - .flat_map(|s| s.values_mut().map(|(_, vi)| vi)), - ) - .filter(|higher_v| **higher_v > v) - { - higher_v.decrement(); - } + let value_index = self.remove_entry(0, prefix, 0, ValueIndex::NONE)?; + + // Update the value indices in all the leaves and entries + for higher_v in self + .leaves + .iter_mut() + .chain( + self.entries + .iter_mut() + .flat_map(|s| s.values_mut().map(|(_, vi)| vi)), + ) + .filter(|higher_v| **higher_v > value_index) + { + higher_v.decrement(); + } - // SAFETY: The value is guaranteed to exist because it was just removed from the - // entry map. - self.values.remove(v.get().unwrap()) - }) + // SAFETY: The value is guaranteed to exist because it was just removed from the + // entry map. + Some(self.values.remove(value_index.get().unwrap())) } - /// Find the final parent node and the `PrefixId` of the given key. - fn find_parent_node(&self, prefix: P) -> (usize, PrefixId, ValueIndex) { + /// Find the final parent node and the `PrefixId` of the given key if it exists. + fn find_parent_node(&self, prefix: P) -> Option<(usize, PrefixId)> { let address = prefix.address(); let prefix_length = prefix.prefix_length(); let mut offset = 0; let mut parent_node_index = 0; let mut parent_node = &self.nodes[parent_node_index]; - let mut default_value_index = ValueIndex::NONE; while prefix_length >= offset + STRIDE { let local_id = StrideId::from_address(address, offset, STRIDE); - default_value_index = self.get_default(parent_node_index, local_id); if !parent_node.node_bitmap.contains(local_id) { - break; + return None; } parent_node_index = parent_node.get_child_index(local_id); @@ -561,11 +556,85 @@ where offset += STRIDE; } - let remaining_length = min(prefix_length - offset, STRIDE - 1); let prefix_id = - PrefixId::from_address(address, offset, remaining_length); + PrefixId::from_address(address, offset, prefix_length - offset); + + Some((parent_node_index, prefix_id)) + } + + /// Recursively searches for an entry to remove, cleaning up empty internal nodes on the way back. + fn remove_entry( + &mut self, + parent_node_index: usize, + prefix: P, + offset: u8, + default_value_index: ValueIndex, + ) -> Option { + let address = prefix.address(); + let prefix_length = prefix.prefix_length(); + + if prefix_length >= offset + STRIDE { + let local_id = StrideId::from_address(address, offset, STRIDE); + + if !self.nodes[parent_node_index].node_bitmap.contains(local_id) { + return None; + } - (parent_node_index, prefix_id, default_value_index) + let child_default = self.get_default(parent_node_index, local_id); + let child_index = + self.nodes[parent_node_index].get_child_index(local_id); + + let value_index = self.remove_entry( + child_index, + prefix, + offset + STRIDE, + child_default, + )?; + + if self.nodes[child_index].node_bitmap.is_empty() + && self.entries[child_index].is_empty() + { + self.remove_node(child_index, parent_node_index, local_id); + } + + Some(value_index) + } else { + let prefix_id = + PrefixId::from_address(address, offset, prefix_length - offset); + + self.entries[parent_node_index].remove(&prefix_id).map(|(_, v)| { + // Update the leaf ranges + self.calculate_leaf_ranges( + parent_node_index, + default_value_index, + ); + + v + }) + } + } + + /// Remove a node from the trie, updating leaf ranges and node bases. + fn remove_node( + &mut self, + node_index: usize, + parent_index: usize, + local_id: StrideId, + ) { + let leaf_base = self.nodes[node_index].leaf_base as usize; + + self.nodes[parent_index].node_bitmap.clear(local_id); + self.nodes.remove(node_index); + self.leaves.remove(leaf_base); + self.entries.remove(node_index); + + for node in &mut self.nodes[node_index..] { + node.leaf_base -= 1; + } + + for node in &mut self.nodes[parent_index + 1..] { + node.node_base -= 1; + } } /// Find the next base node and leaf node index for a given parent node index. diff --git a/tests/api/remove.rs b/tests/api/remove.rs index e3569fa..c870a2c 100644 --- a/tests/api/remove.rs +++ b/tests/api/remove.rs @@ -27,3 +27,22 @@ fn remove_and_reinsert_gives_new_value() { trie.insert((u32_strides!(1), 6), 99); assert_eq!(trie.lookup(u32_strides!(1, 63)), Some(&99)); } + +#[test] +fn removed_child_with_same_stride_doesnt_find_missing_key() { + let prefixes = [ + ((u32::from_be_bytes([15, 0, 0, 0]), 25u8), 1), + ((u32::from_be_bytes([8, 0, 0, 0]), 5u8), 2), + ]; + + let mut trie = Poptrie::new(); + for (p, v) in &prefixes { + assert!(trie.insert(*p, *v).is_none()); + } + + for (p, v) in &prefixes { + assert!(trie.contains_key(*p)); + assert_eq!(trie.remove(*p), Some(*v)); + assert!(!trie.contains_key(*p)); + } +} diff --git a/tests/common/reference_model.rs b/tests/common/reference_model.rs index 99de494..eb12b5b 100644 --- a/tests/common/reference_model.rs +++ b/tests/common/reference_model.rs @@ -34,9 +34,9 @@ where self.map.insert((masked_prefix, prefix_length), value); } - pub fn remove(&mut self, prefix: Ipv4Addr, prefix_length: u8) { + pub fn remove(&mut self, prefix: Ipv4Addr, prefix_length: u8) -> Option { let masked_prefix = mask_prefix(prefix, prefix_length); - self.map.remove(&(masked_prefix, prefix_length)); + self.map.remove(&(masked_prefix, prefix_length)) } } diff --git a/tests/proptests/lookup.rs b/tests/proptests/lookup.rs index 3292b8d..69aaffe 100644 --- a/tests/proptests/lookup.rs +++ b/tests/proptests/lookup.rs @@ -64,8 +64,10 @@ proptest! { prop_assert!(poptrie.contains_key((prefix.addr.to_bits(), prefix.prefix_len))); // Delete - poptrie.remove((prefix.addr.to_bits(), prefix.prefix_len)); - reference.remove(prefix.addr, prefix.prefix_len); + prop_assert_eq!( + poptrie.remove((prefix.addr.to_bits(), prefix.prefix_len)), + reference.remove(prefix.addr, prefix.prefix_len) + ); // Assert that the poptrie no longer contains the prefix after deletion prop_assert!(!poptrie.contains_key((prefix.addr.to_bits(), prefix.prefix_len)));