Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 111 additions & 38 deletions src/trie_ref.rs
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ fn trie_ref_from_key_and_path_in<'a, 'paths, V, A, R, RootValF, BuildF, InvalidF
mut node: &'a TrieNodeODRc<V, A>,
root_val_f: RootValF,
node_key: &'paths [u8],
mut path: &'paths [u8],
path: &'paths [u8],
alloc: A,
build: BuildF,
invalid: InvalidF,
Expand All @@ -113,43 +113,42 @@ where
BuildF: FnOnce(&'a TrieNodeODRc<V, A>, &[u8], Option<&'a V>, A) -> R,
InvalidF: FnOnce(A) -> R,
{
// A temporary buffer on the stack, if we need to assemble a combined key from both the `node_key` and `path`.
let mut temp_key_buf: [MaybeUninit<u8>; MAX_NODE_KEY_BYTES] = [MaybeUninit::uninit(); MAX_NODE_KEY_BYTES];

// A focus key longer than any key in a node cannot exist here.
if node_key.len() > MAX_NODE_KEY_BYTES {
return invalid(alloc);
}
let mut path = path;
let node_key_len = node_key.len();
let path_len = path.len();

// Copy the existing node key and the first chunk of the path into the temporary buffer, then try to descend one step.
let mut temp_key_buf: [MaybeUninit<u8>; MAX_NODE_KEY_BYTES] = [MaybeUninit::uninit(); MAX_NODE_KEY_BYTES];
if node_key_len > 0 && path_len > 0 {
let next_node_path = unsafe {
// SAFETY: `temp_key_buf` has capacity for `MAX_NODE_KEY_BYTES` bytes. We copy exactly
// `node_key_len` bytes from `node_key`, which is a valid slice, then append at most the
// remaining buffer capacity from the valid slice `path`. Both destination ranges are
// within the stack buffer and do not overlap the sources.
let src_ptr = node_key.as_ptr();
let dst_ptr = temp_key_buf.as_mut_ptr().cast::<u8>();
core::ptr::copy_nonoverlapping(src_ptr, dst_ptr, node_key_len);

let remaining_len = (MAX_NODE_KEY_BYTES - node_key_len).min(path_len);
let src_ptr = path.as_ptr();
let dst_ptr = temp_key_buf.as_mut_ptr().cast::<u8>().add(node_key_len);
core::ptr::copy_nonoverlapping(src_ptr, dst_ptr, remaining_len);

let total_buf_len = node_key_len + remaining_len;
// SAFETY: The first `total_buf_len` bytes of `temp_key_buf` were initialized by the
// copies above, and `total_buf_len <= MAX_NODE_KEY_BYTES`, so this slice is valid for
// reads for the duration of this function.
core::slice::from_raw_parts(temp_key_buf.as_mut_ptr().cast::<u8>(), total_buf_len)
};

match node.as_tagged().node_get_child(next_node_path) {
// Only step into the child if path remains, or we'd answer with the focus value.
Some((consumed_byte_cnt, next_node)) if consumed_byte_cnt >= node_key_len && consumed_byte_cnt < node_key_len + path_len => {
node = next_node;
path = &path[consumed_byte_cnt-node_key_len..];
let available = MAX_NODE_KEY_BYTES - node_key_len;
if path_len <= available {
// The candidate fits in the node-key buffer. Let the normal walker handle child boundaries.
path = unsafe {
// SAFETY: node_key_len + path_len <= MAX_NODE_KEY_BYTES.
let dst = temp_key_buf.as_mut_ptr().cast::<u8>();
core::ptr::copy_nonoverlapping(node_key.as_ptr(), dst, node_key_len);
core::ptr::copy_nonoverlapping(path.as_ptr(), dst.add(node_key_len), path_len);
core::slice::from_raw_parts(dst, node_key_len + path_len)
};
} else {
let next_node_path = unsafe {
// SAFETY: the copies exactly fill the stack buffer.
let dst = temp_key_buf.as_mut_ptr().cast::<u8>();
core::ptr::copy_nonoverlapping(node_key.as_ptr(), dst, node_key_len);
core::ptr::copy_nonoverlapping(path.as_ptr(), dst.add(node_key_len), available);
core::slice::from_raw_parts(dst, MAX_NODE_KEY_BYTES)
};
match node.as_tagged().node_get_child(next_node_path) {
// The combined key is longer than the buffer, so a child found within
// this buffer necessarily leaves part of `path` to traverse.
Some((consumed, next)) if consumed >= node_key_len => {
node = next;
path = &path[consumed - node_key_len..];
}
_ => return invalid(alloc),
}
// If the child begins within `node_key`, let the general walker handle the combined key and path.
_ => path = next_node_path,
}
} else if path_len == 0 {
path = node_key;
Expand All @@ -158,9 +157,8 @@ where
let (node, key, val) = if path.is_empty() {
(node, &[] as &[u8], root_val_f())
} else {
node_along_path(node, path, None, true)
node_along_path(node, path, None, false)
};
let (node, key, val) = node_along_path(node, key, val, false);
if key.len() > MAX_NODE_KEY_BYTES ||
(!key.is_empty() && !node.as_tagged().node_contains_partial_key(key))
{
Expand Down Expand Up @@ -455,7 +453,7 @@ impl<V: Clone + Send + Sync + Unpin, A: Allocator> ZipperConcrete for TrieRefBor
}
fn is_shared(&self) -> bool {
match self.focus_node {
Some(node) => self.node_key().is_empty() && node.refcount() > 1,
Some(node) => self.node_key().is_empty() && !node.is_empty() && node.refcount() > 1,
None => false,
}
}
Expand Down Expand Up @@ -782,7 +780,7 @@ impl<V: Clone + Send + Sync + Unpin, A: Allocator> ZipperConcrete for TrieRefOwn
}
fn is_shared(&self) -> bool {
match &self.focus_node {
Some(node) => self.node_key().is_empty() && node.refcount() > 1,
Some(node) => self.node_key().is_empty() && !node.is_empty() && node.refcount() > 1,
None => false
}
}
Expand Down Expand Up @@ -1448,4 +1446,79 @@ mod tests {
assert!(!trie_ref.is_shared());
assert_eq!(trie_ref.shared_node_id(), None);
}

/// `val_at` and `trie_ref_at_path` where the focus key plus the path are longer than a node key
#[test]
fn trie_ref_long_node_key_and_path() {
let mut map = PathMap::<u64>::new();
map.set_val_at(&[0u8; 70], 5);
map.set_val_at(&[1u8], 6);
for (focus, rest) in [(10usize, 60usize), (30, 40), (47, 23), (69, 1)] {
let mut rz = map.read_zipper();
rz.descend_to(&vec![0u8; focus]);
assert_eq!(rz.val_at(&vec![0u8; rest]), Some(&5), "{focus}+{rest}");
assert_eq!(rz.trie_ref_at_path(&vec![0u8; rest]).val(), Some(&5), "{focus}+{rest}");
assert_eq!(rz.val_at(&vec![0u8; rest + 1]), None, "{focus}+{rest}");
}
//A focus far below anything in the trie
let mut rz = map.read_zipper();
rz.descend_to(&[7u8; 60]);
assert_eq!(rz.val_at(&[1u8]), None);
assert_eq!(rz.val_at(&[7u8; 60]), None);
assert!(!rz.trie_ref_at_path(&[1u8]).path_exists());
}

/// `is_shared` where the focus is the empty sentinel node
#[test]
fn trie_ref_is_shared_on_empty_node() {
let mut map = PathMap::<u64>::new();
map.set_val_at(&[1u8], 1);
map.remove_val_at(&[1u8], false);
let empty = PathMap::<u64>::new();
for m in [&map, &empty] {
for path in [&[][..], &[1u8][..]] {
let t = m.trie_ref_at_path(path);
let _ = (t.is_shared(), t.shared_node_id());
}
}
let mut src = PathMap::<u64>::new();
src.set_val_at(&[2u8, 3], 1);
{
let mut wz = src.write_zipper_at_path(&[2u8]);
wz.remove_branches(false);
}
let t = src.trie_ref_at_path(&[2u8]);
let _ = (t.is_shared(), t.shared_node_id());
}

/// Every split of a long path into a focus and a lookup path agrees with a lookup from the root
#[test]
fn trie_ref_long_key_every_split() {
let mut map = PathMap::<u64>::new();
let mut keys: Vec<Vec<u8>> = vec![];
for (n, len) in [(0u8, 100usize), (1, 130), (2, 49), (3, 48), (4, 47)] {
let mut k = vec![n; len];
keys.push(k.clone());
//A branch part way along, and one just past a node key's length
for at in [30usize, 48, 49, 96] {
if at < len { k[at] = 9; keys.push(k[..(at + 5).min(len)].to_vec()); k[at] = n; }
}
}
for (i, k) in keys.iter().enumerate() { map.set_val_at(k, i as u64); }
for k in keys.iter() {
for focus in 0..=k.len() {
let mut rz = map.read_zipper();
rz.descend_to(&k[..focus]);
for rest in 0..=(k.len() - focus + 1).min(120) {
let mut path = k[focus..].iter().copied().take(rest).collect::<Vec<u8>>();
while path.len() < rest { path.push(7); }
let mut full = k[..focus].to_vec();
full.extend(&path);
let want = map.val_at(&full);
assert_eq!(rz.val_at(&path), want, "focus {focus} rest {rest} key {}", k.len());
assert_eq!(rz.trie_ref_at_path(&path).val(), want, "focus {focus} rest {rest} key {}", k.len());
}
}
}
}
}