1use 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#[derive(Debug, Clone, PartialEq, Eq)]
15pub struct ProofTrieNodeV2 {
16 pub path: Nibbles,
18 pub node: TrieNodeV2,
20 pub masks: Option<BranchNodeMasks>,
23}
24
25impl ProofTrieNodeV2 {
26 pub fn empty() -> Self {
29 Self { path: Nibbles::default(), node: TrieNodeV2::EmptyRoot, masks: None }
30 }
31
32 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 let expected_branch_path = path.join(&ext.key);
68
69 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 }
91 }
92 }
93
94 result
95 }
96}
97
98#[derive(PartialEq, Eq, Clone, Debug)]
103#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
104pub enum TrieNodeV2 {
105 EmptyRoot,
107 Branch(BranchNodeV2),
109 Leaf(LeafNode),
111 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#[derive(PartialEq, Eq, Clone, Default)]
183#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
184pub struct BranchNodeV2 {
185 pub key: Nibbles,
188 pub stack: Vec<RlpNode>,
190 pub state_mask: TrieMask,
192 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 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}