From 68dae2f0cae9c5b8a4480991a72669ac6c2ba3a9 Mon Sep 17 00:00:00 2001 From: Valery Date: Mon, 20 Jul 2026 15:09:33 +0100 Subject: [PATCH] perf(mpt): specialize common witness RLP forms --- crates/mpt/src/tests.rs | 18 ++++++++++ crates/mpt/src/trie.rs | 79 +++++++++++++++++++++++------------------ 2 files changed, 63 insertions(+), 34 deletions(-) diff --git a/crates/mpt/src/tests.rs b/crates/mpt/src/tests.rs index 444b1c929..d63baffbc 100644 --- a/crates/mpt/src/tests.rs +++ b/crates/mpt/src/tests.rs @@ -232,6 +232,24 @@ fn test_serde_keccak_trie() -> Result<(), Error> { Ok(()) } +#[cfg(feature = "host")] +#[test] +fn test_serde_digest_root() -> Result<(), Error> { + let bump = bumpalo::Bump::new(); + let digest = bump.alloc_slice_copy(&[0xabu8; 32]); + let mut trie = Mpt::new(&bump); + let root_id = trie.add_node(NodeData::Digest(digest), None); + trie.set_root_id(root_id); + + let encoded = trie.encode_trie(); + let mut encoded_slice = encoded.as_slice(); + let recovered_trie = Mpt::decode_trie(&bump, &mut encoded_slice, trie.num_nodes())?; + + assert!(encoded_slice.is_empty()); + assert_eq!(recovered_trie.hash(), B256::from_slice(digest)); + Ok(()) +} + /// Test that deleting a key that would cause branch collapse fails if sibling is unresolved Digest. /// This tests the soundness fix: we must not assume an unresolved Digest is a Branch node. #[test] diff --git a/crates/mpt/src/trie.rs b/crates/mpt/src/trie.rs index f00fcd921..86105e23d 100644 --- a/crates/mpt/src/trie.rs +++ b/crates/mpt/src/trie.rs @@ -103,6 +103,24 @@ unsafe fn advance_unchecked<'a>(buf: &mut &'a [u8], cnt: usize) -> &'a [u8] { bytes } +/// Decodes and advances past one RLP item, returning its full encoding and payload. +/// Empty strings dominate sparse branch children, so handle their single-byte encoding without +/// entering the generic RLP header decoder. +#[inline(always)] +fn decode_rlp_item<'a>(buf: &mut &'a [u8]) -> Result<(&'a [u8], &'a [u8]), Error> { + let item_start = *buf; + if item_start.first() == Some(&alloy_rlp::EMPTY_STRING_CODE) { + *buf = &item_start[1..]; + return Ok((&item_start[..1], &item_start[1..1])); + } + + let alloy_rlp::Header { payload_length, .. } = alloy_rlp::Header::decode(buf)?; + // SAFETY: the header was decoded, so the item contains its declared payload. + let payload = unsafe { advance_unchecked(buf, payload_length) }; + let item_length = item_start.len() - buf.len(); + Ok((&item_start[..item_length], payload)) +} + /// Whether `slice` is the RLP encoding of an empty node reference: a single `EMPTY_STRING_CODE` /// byte. Written as an explicit pattern match because comparing against [`NULL_NODE_REF_SLICE`] /// with `==` compiles to a `memcmp` call, whose overhead dwarfs this one-byte check — and trie @@ -253,6 +271,21 @@ impl<'a> Mpt<'a> { bytes: &mut &'a [u8], expected_node_ref: NodeRef<'a>, ) -> Result { + // Unresolved nodes are encoded as a 32-byte RLP string followed by three alignment bytes. + // They are common at the witness boundary and have a fixed representation, so avoid the + // generic header decoder and the general node-kind dispatch. + if bytes.first() == Some(&(alloy_rlp::EMPTY_STRING_CODE + 32)) { + // SAFETY: as elsewhere in this decoder, the witness encoding is expected to contain + // the declared RLP payload and its alignment padding. + let encoded = unsafe { advance_unchecked(bytes, 33) }; + unsafe { advance_unchecked(bytes, 3) }; + let digest = &encoded[1..]; + if !bytes_eq(digest, expected_node_ref.as_slice()) { + return Err(Error::NodeRefMismatch); + } + return Ok(self.add_node(NodeData::Digest(digest), Some(NodeRef::Digest(digest)))); + } + let rlp_node_header_start = *bytes; let alloy_rlp::Header { list, payload_length } = alloy_rlp::Header::decode(bytes)?; @@ -298,56 +331,41 @@ impl<'a> Mpt<'a> { return Ok(node_id); } - // first payload item - let item0_header_start = payload; - let alloy_rlp::Header { payload_length: item0_payload_length, .. } = - alloy_rlp::Header::decode(&mut payload)?; - // SAFETY: we already decoded the header, so we know the payload length. - let item0_payload_start = unsafe { advance_unchecked(&mut payload, item0_payload_length) }; - let item0_length = item0_header_start.len() - payload.len(); - - // second payload item - let item1_header_start = payload; - let alloy_rlp::Header { payload_length: item1_payload_length, .. } = - alloy_rlp::Header::decode(&mut payload)?; - // SAFETY: we already decoded the header, so we know the payload length. - let item1_payload_start = unsafe { advance_unchecked(&mut payload, item1_payload_length) }; - let item1_length = item1_header_start.len() - payload.len(); + let (item0, item0_payload) = decode_rlp_item(&mut payload)?; + let (item1, item1_payload) = decode_rlp_item(&mut payload)?; if payload.is_empty() { // either an extension or leaf - let path = &item0_payload_start[..item0_payload_length]; + let path = item0_payload; let prefix = path[0]; if (prefix & (2 << 4)) == 0 { // extension node - let ext_node_expected_ref = - NodeRef::from_rlp_slice(&item1_header_start[..item1_length]); + let ext_node_expected_ref = NodeRef::from_rlp_slice(item1); let ext_node_id = self.decode_trie_internal(bytes, ext_node_expected_ref)?; let node_data = NodeData::Extension(path, ext_node_id); return Ok(self.add_node(node_data, Some(node_ref))); } else { // leaf node - let value = &item1_payload_start[..item1_payload_length]; - let node_data = NodeData::Leaf(path, value); + let node_data = NodeData::Leaf(path, item1_payload); return Ok(self.add_node(node_data, Some(node_ref))); } } // branch - let child0_expected_node_ref = NodeRef::from_rlp_slice(&item0_header_start[..item0_length]); let child0 = { - if is_null_ref(child0_expected_node_ref.as_slice()) { + if is_null_ref(item0) { None } else { + let child0_expected_node_ref = NodeRef::from_rlp_slice(item0); Some(branch_child_id(self.decode_trie_internal(bytes, child0_expected_node_ref)?)) } }; - let child1_expected_node_ref = NodeRef::from_rlp_slice(&item1_header_start[..item1_length]); let child1 = { - if is_null_ref(child1_expected_node_ref.as_slice()) { + if is_null_ref(item1) { None } else { + let child1_expected_node_ref = NodeRef::from_rlp_slice(item1); Some(branch_child_id(self.decode_trie_internal(bytes, child1_expected_node_ref)?)) } }; @@ -371,20 +389,13 @@ impl<'a> Mpt<'a> { // Initialize remaining elements for child in &mut childs[2..] { - let item_header_start = payload; - let alloy_rlp::Header { payload_length: item_payload_length, .. } = - alloy_rlp::Header::decode(&mut payload)?; - // SAFETY: we already decoded the header, so we know the payload length. - unsafe { advance_unchecked(&mut payload, item_payload_length) }; - let item_length = item_header_start.len() - payload.len(); - - let child_expected_node_ref = - NodeRef::from_rlp_slice(&item_header_start[..item_length]); + let (item, _) = decode_rlp_item(&mut payload)?; *child = MaybeUninit::new({ - if is_null_ref(child_expected_node_ref.as_slice()) { + if is_null_ref(item) { None } else { + let child_expected_node_ref = NodeRef::from_rlp_slice(item); Some(branch_child_id( self.decode_trie_internal(bytes, child_expected_node_ref)?, ))