From 8eedcbc9d95b7d5b907b8fd5f2f6a050d3ea1138 Mon Sep 17 00:00:00 2001 From: Igor Malovitsa Date: Tue, 29 Sep 2026 19:57:58 +0000 Subject: [PATCH 1/4] Fix TrieRef lookups with a node key longer than its buffer new_with_key_and_path_in copied the focus's node key into a fixed stack buffer, overflowing it (UB in release) when the key was longer, and truncated key + path when the two together didn't fit. No child key is longer than a node key, so only the first MAX_NODE_KEY_BYTES bytes of the combined key can decide a step: look up that prefix, and keep stepping through a child that begins within the node key. A combined key too long to remain at a node is invalid. --- src/trie_ref.rs | 108 +++++++++++++++++++++++++++++++++++++----------- 1 file changed, 83 insertions(+), 25 deletions(-) diff --git a/src/trie_ref.rs b/src/trie_ref.rs index 67b59d70..2d514e06 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, @@ -115,31 +115,24 @@ where { // 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]; - - let node_key_len = node_key.len(); - let path_len = path.len(); + let mut node_key: &[u8] = node_key; + let mut path: &[u8] = path; // Copy the existing node key and the first chunk of the path into the temporary buffer, then try to descend one step. - if node_key_len > 0 && path_len > 0 { + // No child key is longer than `MAX_NODE_KEY_BYTES`, so the first `MAX_NODE_KEY_BYTES` bytes of the + // combined key decide every step, however long the combined key is. + while !node_key.is_empty() && !path.is_empty() { + let node_key_len = node_key.len(); + let path_len = path.len(); + let key_from_node = node_key_len.min(MAX_NODE_KEY_BYTES); + let key_from_path = path_len.min(MAX_NODE_KEY_BYTES - key_from_node); 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(); + // SAFETY: `key_from_node + key_from_path <= MAX_NODE_KEY_BYTES`, and we copy that many bytes + // from two valid slices into `temp_key_buf`, so the resulting slice is initialized and in bounds. 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) + core::ptr::copy_nonoverlapping(node_key.as_ptr(), dst_ptr, key_from_node); + core::ptr::copy_nonoverlapping(path.as_ptr(), dst_ptr.add(key_from_node), key_from_path); + core::slice::from_raw_parts(dst_ptr, key_from_node + key_from_path) }; match node.as_tagged().node_get_child(next_node_path) { @@ -147,11 +140,25 @@ where 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..]; + node_key = &[]; + } + // The child begins within `node_key`, so step into it and keep going with the rest + Some((consumed_byte_cnt, next_node)) if consumed_byte_cnt < node_key_len => { + node = next_node; + node_key = &node_key[consumed_byte_cnt..]; + } + _ => { + // The combined key can't be walked any further from here, so it's the remaining key + // at `node`, and a key longer than a node key can't be there + if node_key_len + path_len > MAX_NODE_KEY_BYTES { + return invalid(alloc); + } + path = next_node_path; + node_key = &[]; } - // 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 { + } + if path.is_empty() { path = node_key; } @@ -1448,4 +1455,55 @@ 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); + } + + /// 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()); + } + } + } + } } From ae439ed700e172d461d4fb5862a4549d70ca9699 Mon Sep 17 00:00:00 2001 From: Igor Malovitsa Date: Tue, 29 Sep 2026 19:57:58 +0000 Subject: [PATCH 2/4] Fix TrieRef::is_shared on an empty node --- src/trie_ref.rs | 27 +++++++++++++++++++++++++-- 1 file changed, 25 insertions(+), 2 deletions(-) diff --git a/src/trie_ref.rs b/src/trie_ref.rs index 2d514e06..9bc0db48 100644 --- a/src/trie_ref.rs +++ b/src/trie_ref.rs @@ -462,7 +462,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, } } @@ -789,7 +789,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 } } @@ -1476,6 +1476,29 @@ mod tests { assert_eq!(rz.val_at(&[7u8; 60]), None); } + /// `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() { From ad6d4639c103383b4790ff9df22822651cf7987a Mon Sep 17 00:00:00 2001 From: Luke Peterson Date: Thu, 1 Oct 2026 21:06:49 -0600 Subject: [PATCH 3/4] Some optimizations to try to claw back a couple percent --- src/trie_ref.rs | 77 ++++++++++++++++++++++--------------------------- 1 file changed, 35 insertions(+), 42 deletions(-) diff --git a/src/trie_ref.rs b/src/trie_ref.rs index 9bc0db48..f781e1e8 100644 --- a/src/trie_ref.rs +++ b/src/trie_ref.rs @@ -113,52 +113,44 @@ 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`. + // 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(); let mut temp_key_buf: [MaybeUninit; MAX_NODE_KEY_BYTES] = [MaybeUninit::uninit(); MAX_NODE_KEY_BYTES]; - let mut node_key: &[u8] = node_key; - let mut path: &[u8] = path; - - // Copy the existing node key and the first chunk of the path into the temporary buffer, then try to descend one step. - // No child key is longer than `MAX_NODE_KEY_BYTES`, so the first `MAX_NODE_KEY_BYTES` bytes of the - // combined key decide every step, however long the combined key is. - while !node_key.is_empty() && !path.is_empty() { - let node_key_len = node_key.len(); - let path_len = path.len(); - let key_from_node = node_key_len.min(MAX_NODE_KEY_BYTES); - let key_from_path = path_len.min(MAX_NODE_KEY_BYTES - key_from_node); - let next_node_path = unsafe { - // SAFETY: `key_from_node + key_from_path <= MAX_NODE_KEY_BYTES`, and we copy that many bytes - // from two valid slices into `temp_key_buf`, so the resulting slice is initialized and in bounds. - let dst_ptr = temp_key_buf.as_mut_ptr().cast::(); - core::ptr::copy_nonoverlapping(node_key.as_ptr(), dst_ptr, key_from_node); - core::ptr::copy_nonoverlapping(path.as_ptr(), dst_ptr.add(key_from_node), key_from_path); - core::slice::from_raw_parts(dst_ptr, key_from_node + key_from_path) - }; - - 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..]; - node_key = &[]; - } - // The child begins within `node_key`, so step into it and keep going with the rest - Some((consumed_byte_cnt, next_node)) if consumed_byte_cnt < node_key_len => { - node = next_node; - node_key = &node_key[consumed_byte_cnt..]; - } - _ => { - // The combined key can't be walked any further from here, so it's the remaining key - // at `node`, and a key longer than a node key can't be there - if node_key_len + path_len > MAX_NODE_KEY_BYTES { - return invalid(alloc); + if node_key_len > 0 && path_len > 0 { + 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..]; } - path = next_node_path; - node_key = &[]; + _ => return invalid(alloc), } } - } - if path.is_empty() { + } else if path_len == 0 { path = node_key; } @@ -1474,6 +1466,7 @@ mod tests { 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 From 94779b16f1d2736ca26424350254ef79b608dbbb Mon Sep 17 00:00:00 2001 From: Luke Peterson Date: Thu, 1 Oct 2026 21:14:49 -0600 Subject: [PATCH 4/4] Fixing some silliness that was costing some performance --- src/trie_ref.rs | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/trie_ref.rs b/src/trie_ref.rs index f781e1e8..3b7312a0 100644 --- a/src/trie_ref.rs +++ b/src/trie_ref.rs @@ -157,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)) {