diff --git a/src/trie_ref.rs b/src/trie_ref.rs index 67b59d70..3b7312a0 100644 --- a/src/trie_ref.rs +++ b/src/trie_ref.rs @@ -101,7 +101,7 @@ fn trie_ref_from_key_and_path_in<'a, 'paths, V, A, R, RootValF, BuildF, InvalidF mut node: &'a TrieNodeODRc, root_val_f: RootValF, node_key: &'paths [u8], - mut path: &'paths [u8], + path: &'paths [u8], alloc: A, build: BuildF, invalid: InvalidF, @@ -113,43 +113,42 @@ where BuildF: FnOnce(&'a TrieNodeODRc, &[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; 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; 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::(); - 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::().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::(), 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::(); + 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::(); + 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; @@ -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)) { @@ -455,7 +453,7 @@ impl 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, } } @@ -782,7 +780,7 @@ impl 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 } } @@ -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::::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::::new(); + map.set_val_at(&[1u8], 1); + map.remove_val_at(&[1u8], false); + let empty = PathMap::::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::::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::::new(); + let mut keys: Vec> = 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::>(); + 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()); + } + } + } + } }