Skip to main content

reth_trie_common/
trie_node_v2.rs

1//! Version 2 types related to representing nodes in an MPT.
2
3use crate::BranchNodeMasks;
4use alloc::vec::Vec;
5use alloy_primitives::hex;
6use alloy_rlp::{bytes, Decodable, Encodable, EMPTY_STRING_CODE};
7use alloy_trie::{
8    nodes::{BranchNodeRef, ExtensionNode, ExtensionNodeRef, LeafNode, RlpNode, TrieNode},
9    Nibbles, TrieMask,
10};
11use core::fmt;
12
13/// Carries all information needed by a sparse trie to reveal a particular node.
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub struct ProofTrieNodeV2 {
16    /// Path of the node.
17    pub path: Nibbles,
18    /// The node itself.
19    pub node: TrieNodeV2,
20    /// Tree and hash masks for the node, if known.
21    /// Both masks are always set together (from database branch nodes).
22    pub masks: Option<BranchNodeMasks>,
23}
24
25impl ProofTrieNodeV2 {
26    /// Creates an empty `ProofTrieNodeV2` with an empty root node. Useful as a placeholder when
27    /// taking a node out of a slice via [`core::mem::replace`].
28    pub fn empty() -> Self {
29        Self { path: Nibbles::default(), node: TrieNodeV2::EmptyRoot, masks: None }
30    }
31
32    /// Converts an iterator of `(path, TrieNode, masks)` tuples into `Vec<ProofTrieNodeV2>`,
33    /// merging extension nodes into their child branch nodes.
34    ///
35    /// The input **must** be sorted in depth-first order (children before parents) for extension
36    /// merging to work correctly.
37    pub fn from_sorted_trie_nodes(
38        iter: impl IntoIterator<Item = (Nibbles, TrieNode, Option<BranchNodeMasks>)>,
39    ) -> Vec<Self> {
40        let iter = iter.into_iter();
41        let mut result = Vec::with_capacity(iter.size_hint().0);
42
43        for (path, node, masks) in iter {
44            match node {
45                TrieNode::EmptyRoot => {
46                    result.push(Self { path, node: TrieNodeV2::EmptyRoot, masks });
47                }
48                TrieNode::Leaf(leaf) => {
49                    result.push(Self { path, node: TrieNodeV2::Leaf(leaf), masks });
50                }
51                TrieNode::Branch(branch) => {
52                    result.push(Self {
53                        path,
54                        node: TrieNodeV2::Branch(BranchNodeV2 {
55                            key: Nibbles::new(),
56                            branch_rlp_node: None,
57                            stack: branch.stack,
58                            state_mask: branch.state_mask,
59                        }),
60                        masks,
61                    });
62                }
63                TrieNode::Extension(ext) => {
64                    // In depth-first order, the child branch comes BEFORE the parent
65                    // extension. The child branch should be the last item we added to
66                    // result, at path extension.path + extension.key.
67                    let expected_branch_path = path.join(&ext.key);
68
69                    // Check if the last item in result is the child branch
70                    if let Some(last) = result.last_mut() &&
71                        last.path == expected_branch_path &&
72                        let TrieNodeV2::Branch(branch_v2) = &mut last.node
73                    {
74                        debug_assert!(
75                            branch_v2.key.is_empty(),
76                            "Branch at {:?} already has extension key {:?}",
77                            last.path,
78                            branch_v2.key
79                        );
80                        branch_v2.key = ext.key;
81                        branch_v2.branch_rlp_node = Some(ext.child);
82                        last.path = path;
83                    }
84
85                    // If we reach here, the extension's child is not a branch in the
86                    // result. This happens when the child branch is hashed (not revealed
87                    // in the proof). In V2 format, extension nodes are always combined
88                    // with their child branch, so we skip extension nodes whose child
89                    // isn't revealed.
90                }
91            }
92        }
93
94        result
95    }
96}
97
98/// Enum representing an MPT trie node.
99///
100/// This is a V2 representiation, differing from [`TrieNode`] in that branch and extension nodes are
101/// compressed into a single node.
102#[derive(PartialEq, Eq, Clone, Debug)]
103#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
104pub enum TrieNodeV2 {
105    /// Variant representing empty root node.
106    EmptyRoot,
107    /// Variant representing a [`BranchNodeV2`].
108    Branch(BranchNodeV2),
109    /// Variant representing a [`LeafNode`].
110    Leaf(LeafNode),
111    /// Variant representing an [`ExtensionNode`].
112    ///
113    /// This will only be used for extension nodes for which child is not inlined. This variant
114    /// will never be produced by proof workers that will always reveal a full path to a requested
115    /// leaf.
116    Extension(ExtensionNode),
117}
118
119impl Encodable for TrieNodeV2 {
120    fn length(&self) -> usize {
121        match self {
122            Self::EmptyRoot => 1,
123            Self::Leaf(leaf) => leaf.as_ref().length(),
124            Self::Branch(branch) => branch.length(),
125            Self::Extension(ext) => ext.length(),
126        }
127    }
128
129    fn encode(&self, out: &mut dyn bytes::BufMut) {
130        match self {
131            Self::EmptyRoot => {
132                out.put_u8(EMPTY_STRING_CODE);
133            }
134            Self::Leaf(leaf) => {
135                leaf.as_ref().encode(out);
136            }
137            Self::Branch(branch) => branch.encode(out),
138            Self::Extension(ext) => {
139                ext.encode(out);
140            }
141        }
142    }
143}
144
145impl Decodable for TrieNodeV2 {
146    fn decode(buf: &mut &[u8]) -> Result<Self, alloy_rlp::Error> {
147        match TrieNode::decode(buf)? {
148            TrieNode::EmptyRoot => Ok(Self::EmptyRoot),
149            TrieNode::Leaf(leaf) => Ok(Self::Leaf(leaf)),
150            TrieNode::Branch(branch) => Ok(Self::Branch(BranchNodeV2::new(
151                Default::default(),
152                branch.stack,
153                branch.state_mask,
154                None,
155            ))),
156            TrieNode::Extension(ext) => {
157                if ext.child.is_hash() {
158                    Ok(Self::Extension(ext))
159                } else {
160                    let Self::Branch(mut branch) = Self::decode(&mut ext.child.as_ref())? else {
161                        return Err(alloy_rlp::Error::Custom(
162                            "extension node child is not a branch",
163                        ));
164                    };
165
166                    branch.key = ext.key;
167
168                    Ok(Self::Branch(branch))
169                }
170            }
171        }
172    }
173}
174
175/// A branch node in an Ethereum Merkle Patricia Trie.
176///
177/// Branch node is a 17-element array consisting of 16 slots that correspond to each hexadecimal
178/// character and an additional slot for a value. We do exclude the node value since all paths have
179/// a fixed size.
180///
181/// This node also encompasses the possible parent extension node of a branch via the `key` field.
182#[derive(PartialEq, Eq, Clone, Default)]
183#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
184pub struct BranchNodeV2 {
185    /// The key for the branch's parent extension. if key is empty then the branch does not have a
186    /// parent extension.
187    pub key: Nibbles,
188    /// The collection of RLP encoded children.
189    pub stack: Vec<RlpNode>,
190    /// The bitmask indicating the presence of children at the respective nibble positions
191    pub state_mask: TrieMask,
192    /// [`RlpNode`] encoding of the branch node. Always provided when `key` is not empty (i.e this
193    /// is an extension node).
194    pub branch_rlp_node: Option<RlpNode>,
195}
196
197impl fmt::Debug for BranchNodeV2 {
198    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
199        f.debug_struct("BranchNode")
200            .field("key", &self.key)
201            .field("stack", &self.stack.iter().map(hex::encode).collect::<Vec<_>>())
202            .field("state_mask", &self.state_mask)
203            .field("branch_rlp_node", &self.branch_rlp_node)
204            .finish()
205    }
206}
207
208impl BranchNodeV2 {
209    /// Creates a new branch node with the given short key, stack, and state mask.
210    pub const fn new(
211        key: Nibbles,
212        stack: Vec<RlpNode>,
213        state_mask: TrieMask,
214        branch_rlp_node: Option<RlpNode>,
215    ) -> Self {
216        Self { key, stack, state_mask, branch_rlp_node }
217    }
218}
219
220impl Encodable for BranchNodeV2 {
221    fn encode(&self, out: &mut dyn bytes::BufMut) {
222        if self.key.is_empty() {
223            BranchNodeRef::new(&self.stack, self.state_mask).encode(out);
224            return;
225        }
226
227        let branch_rlp_node = self
228            .branch_rlp_node
229            .as_ref()
230            .expect("branch_rlp_node must always be present for extension nodes");
231
232        ExtensionNodeRef::new(&self.key, branch_rlp_node.as_slice()).encode(out);
233    }
234
235    fn length(&self) -> usize {
236        if self.key.is_empty() {
237            return BranchNodeRef::new(&self.stack, self.state_mask).length()
238        }
239
240        let branch_rlp_node = self
241            .branch_rlp_node
242            .as_ref()
243            .expect("branch_rlp_node must always be present for extension nodes");
244
245        ExtensionNodeRef::new(&self.key, branch_rlp_node.as_slice()).length()
246    }
247}
248
249#[cfg(test)]
250mod tests {
251    use super::*;
252    use alloy_primitives::B256;
253    use proptest::prelude::*;
254
255    fn assert_roundtrip_and_length(node: TrieNodeV2) {
256        let encoded = alloy_rlp::encode(&node);
257        assert_eq!(node.length(), encoded.len());
258        let mut buf = encoded.as_slice();
259        assert_eq!(TrieNodeV2::decode(&mut buf).unwrap(), node);
260        assert!(buf.is_empty());
261    }
262
263    #[test]
264    fn trie_node_variants_rlp_roundtrip_and_length() {
265        assert_roundtrip_and_length(TrieNodeV2::EmptyRoot);
266        assert_roundtrip_and_length(TrieNodeV2::Branch(BranchNodeV2::default()));
267        assert_roundtrip_and_length(TrieNodeV2::Extension(ExtensionNode::new(
268            Nibbles::from_nibbles([1]),
269            RlpNode::word_rlp(&B256::repeat_byte(0xaa)),
270        )));
271    }
272
273    proptest! {
274        #[test]
275        fn leaf_rlp_roundtrip_and_length(
276            key in proptest::collection::vec(0u8..16, 0..64),
277            value in proptest::collection::vec(any::<u8>(), 0..128),
278        ) {
279            let node = TrieNodeV2::Leaf(LeafNode::new(Nibbles::from_nibbles(key), value));
280            assert_roundtrip_and_length(node);
281        }
282    }
283}