1use crate::BlockAccessLists;
9use alloc::vec::Vec;
10use alloy_primitives::{Bytes, B256, KECCAK256_EMPTY, U256};
11use alloy_rlp::{BufMut, Decodable, Encodable, Header, RlpDecodable, RlpEncodable};
12use alloy_trie::{TrieAccount, EMPTY_ROOT_HASH};
13use reth_codecs_derive::add_arbitrary_tests;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Hash)]
17#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
18#[repr(u8)]
19pub enum SnapVersion {
20 #[default]
22 V2 = 2,
23}
24
25impl SnapVersion {
26 pub const fn message_count(self) -> u8 {
29 match self {
30 Self::V2 => 10,
31 }
32 }
33
34 pub const fn supports_message_id(self, id: u8) -> bool {
39 match self {
40 Self::V2 => {
42 id <= SnapMessageId::ByteCodes as u8 ||
43 id == SnapMessageId::GetBlockAccessLists as u8 ||
44 id == SnapMessageId::BlockAccessLists as u8
45 }
46 }
47 }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
52pub enum SnapMessageId {
53 GetAccountRange = 0x00,
55 AccountRange = 0x01,
58 GetStorageRanges = 0x02,
60 StorageRanges = 0x03,
62 GetByteCodes = 0x04,
64 ByteCodes = 0x05,
66 GetBlockAccessLists = 0x08,
68 BlockAccessLists = 0x09,
70}
71
72impl SnapMessageId {
73 pub const fn response(self) -> Option<Self> {
76 match self {
77 Self::GetAccountRange => Some(Self::AccountRange),
78 Self::GetStorageRanges => Some(Self::StorageRanges),
79 Self::GetByteCodes => Some(Self::ByteCodes),
80 Self::GetBlockAccessLists => Some(Self::BlockAccessLists),
81 Self::AccountRange | Self::StorageRanges | Self::ByteCodes | Self::BlockAccessLists => {
82 None
83 }
84 }
85 }
86}
87
88#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
91#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
92#[add_arbitrary_tests(rlp)]
93pub struct GetAccountRangeMessage {
94 pub request_id: u64,
96 pub root_hash: B256,
98 pub starting_hash: B256,
100 pub limit_hash: B256,
102 pub response_bytes: u64,
104}
105
106#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
108#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
109#[add_arbitrary_tests(rlp)]
110pub struct AccountData {
111 pub hash: B256,
113 pub body: SlimAccountBody,
115}
116
117impl AccountData {
118 pub fn from_trie_account(hash: B256, account: &TrieAccount) -> Self {
120 Self { hash, body: account.into() }
121 }
122
123 #[allow(clippy::clone_on_copy)]
127 pub fn trie_account(&self) -> TrieAccount {
128 self.body.0.clone()
129 }
130
131 #[allow(clippy::missing_const_for_fn)]
133 pub fn into_trie_entry(self) -> (B256, TrieAccount) {
134 (self.hash, self.body.0)
135 }
136}
137
138#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
141#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
142#[add_arbitrary_tests(rlp)]
143pub struct AccountRangeMessage {
144 pub request_id: u64,
146 pub accounts: Vec<AccountData>,
148 pub proof: Vec<Bytes>,
150}
151
152#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
155#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
156#[add_arbitrary_tests(rlp)]
157pub struct GetStorageRangesMessage {
158 pub request_id: u64,
160 pub root_hash: B256,
162 pub account_hashes: Vec<B256>,
164 pub starting_hash: RangeBound,
167 pub limit_hash: RangeBound,
170 pub response_bytes: u64,
172}
173
174#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
178#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
179pub struct RangeBound(Option<B256>);
180
181impl RangeBound {
182 pub const fn unwrap_or(self, default: B256) -> B256 {
184 match self.0 {
185 Some(hash) => hash,
186 None => default,
187 }
188 }
189}
190
191impl From<B256> for RangeBound {
192 fn from(hash: B256) -> Self {
193 Self(Some(hash))
194 }
195}
196
197impl Encodable for RangeBound {
198 fn encode(&self, out: &mut dyn BufMut) {
199 match self.0 {
200 Some(hash) => hash.encode(out),
201 None => Bytes::new().encode(out),
202 }
203 }
204
205 fn length(&self) -> usize {
206 match self.0 {
207 Some(hash) => hash.length(),
208 None => Bytes::new().length(),
209 }
210 }
211}
212
213impl Decodable for RangeBound {
214 fn decode(buf: &mut &[u8]) -> alloy_rlp::Result<Self> {
215 let bytes = Bytes::decode(buf)?;
216 match bytes.len() {
217 0 => Ok(Self(None)),
218 32 => Ok(Self(Some(B256::from_slice(&bytes)))),
219 _ => Err(alloy_rlp::Error::UnexpectedLength),
220 }
221 }
222}
223
224#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
226#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
227#[add_arbitrary_tests(rlp)]
228pub struct StorageData {
229 pub hash: B256,
231 pub data: Bytes,
233}
234
235impl StorageData {
236 pub fn from_value(hash: B256, value: U256) -> Self {
238 Self { hash, data: alloy_rlp::encode(value).into() }
239 }
240
241 pub fn value(&self) -> alloy_rlp::Result<U256> {
243 alloy_rlp::decode_exact(&self.data)
244 }
245}
246
247#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
252#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
253#[add_arbitrary_tests(rlp)]
254pub struct StorageRangesMessage {
255 pub request_id: u64,
257 pub slots: Vec<Vec<StorageData>>,
259 pub proof: Vec<Bytes>,
261}
262
263#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
266#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
267#[add_arbitrary_tests(rlp)]
268pub struct GetByteCodesMessage {
269 pub request_id: u64,
271 pub hashes: Vec<B256>,
273 pub response_bytes: u64,
275}
276
277#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
280#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
281#[add_arbitrary_tests(rlp)]
282pub struct ByteCodesMessage {
283 pub request_id: u64,
285 pub codes: Vec<Bytes>,
287}
288
289#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
291#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
292#[add_arbitrary_tests(rlp)]
293pub struct GetBlockAccessListsMessage {
294 pub request_id: u64,
296 pub block_hashes: Vec<B256>,
298 pub response_bytes: u64,
300}
301
302#[derive(Debug, Clone, PartialEq, Eq, RlpEncodable, RlpDecodable)]
304#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
305#[add_arbitrary_tests(rlp)]
306pub struct BlockAccessListsMessage {
307 pub request_id: u64,
309 pub block_access_lists: BlockAccessLists,
311}
312
313#[derive(Debug, Clone, PartialEq, Eq)]
315pub enum SnapProtocolMessage {
316 GetAccountRange(GetAccountRangeMessage),
318 AccountRange(AccountRangeMessage),
320 GetStorageRanges(GetStorageRangesMessage),
322 StorageRanges(StorageRangesMessage),
324 GetByteCodes(GetByteCodesMessage),
326 ByteCodes(ByteCodesMessage),
328 GetBlockAccessLists(GetBlockAccessListsMessage),
330 BlockAccessLists(BlockAccessListsMessage),
332}
333
334#[derive(thiserror::Error, Debug)]
336pub enum SnapProtocolError {
337 #[error("empty snap message")]
339 Empty,
340 #[error("message id {0:#x} is invalid for snap/{1:?}")]
343 UnsupportedMessageId(u8, SnapVersion),
344 #[error("RLP error: {0}")]
346 Rlp(#[from] alloy_rlp::Error),
347}
348
349impl SnapProtocolMessage {
350 pub const fn message_id(&self) -> SnapMessageId {
354 match self {
355 Self::GetAccountRange(_) => SnapMessageId::GetAccountRange,
356 Self::AccountRange(_) => SnapMessageId::AccountRange,
357 Self::GetStorageRanges(_) => SnapMessageId::GetStorageRanges,
358 Self::StorageRanges(_) => SnapMessageId::StorageRanges,
359 Self::GetByteCodes(_) => SnapMessageId::GetByteCodes,
360 Self::ByteCodes(_) => SnapMessageId::ByteCodes,
361 Self::GetBlockAccessLists(_) => SnapMessageId::GetBlockAccessLists,
362 Self::BlockAccessLists(_) => SnapMessageId::BlockAccessLists,
363 }
364 }
365
366 pub const fn request_id(&self) -> u64 {
368 match self {
369 Self::GetAccountRange(m) => m.request_id,
370 Self::AccountRange(m) => m.request_id,
371 Self::GetStorageRanges(m) => m.request_id,
372 Self::StorageRanges(m) => m.request_id,
373 Self::GetByteCodes(m) => m.request_id,
374 Self::ByteCodes(m) => m.request_id,
375 Self::GetBlockAccessLists(m) => m.request_id,
376 Self::BlockAccessLists(m) => m.request_id,
377 }
378 }
379
380 pub const fn is_response(&self) -> bool {
382 matches!(
383 self,
384 Self::AccountRange(_) |
385 Self::StorageRanges(_) |
386 Self::ByteCodes(_) |
387 Self::BlockAccessLists(_)
388 )
389 }
390
391 pub const fn set_request_id(&mut self, request_id: u64) {
394 match self {
395 Self::GetAccountRange(m) => m.request_id = request_id,
396 Self::AccountRange(m) => m.request_id = request_id,
397 Self::GetStorageRanges(m) => m.request_id = request_id,
398 Self::StorageRanges(m) => m.request_id = request_id,
399 Self::GetByteCodes(m) => m.request_id = request_id,
400 Self::ByteCodes(m) => m.request_id = request_id,
401 Self::GetBlockAccessLists(m) => m.request_id = request_id,
402 Self::BlockAccessLists(m) => m.request_id = request_id,
403 }
404 }
405
406 pub fn encode(&self) -> Bytes {
408 let mut buf = Vec::new();
409 buf.push(self.message_id() as u8);
411
412 match self {
414 Self::GetAccountRange(msg) => msg.encode(&mut buf),
415 Self::AccountRange(msg) => msg.encode(&mut buf),
416 Self::GetStorageRanges(msg) => msg.encode(&mut buf),
417 Self::StorageRanges(msg) => msg.encode(&mut buf),
418 Self::GetByteCodes(msg) => msg.encode(&mut buf),
419 Self::ByteCodes(msg) => msg.encode(&mut buf),
420 Self::GetBlockAccessLists(msg) => msg.encode(&mut buf),
421 Self::BlockAccessLists(msg) => msg.encode(&mut buf),
422 }
423
424 Bytes::from(buf)
425 }
426
427 pub fn decode(message_id: u8, buf: &mut &[u8]) -> Result<Self, alloy_rlp::Error> {
429 macro_rules! decode_snap_message_variant {
431 ($message_id:expr, $buf:expr, $id:expr, $variant:ident, $msg_type:ty) => {
432 if $message_id == $id as u8 {
433 return Ok(Self::$variant(<$msg_type>::decode($buf)?));
434 }
435 };
436 }
437
438 decode_snap_message_variant!(
440 message_id,
441 buf,
442 SnapMessageId::GetAccountRange,
443 GetAccountRange,
444 GetAccountRangeMessage
445 );
446 decode_snap_message_variant!(
447 message_id,
448 buf,
449 SnapMessageId::AccountRange,
450 AccountRange,
451 AccountRangeMessage
452 );
453 decode_snap_message_variant!(
454 message_id,
455 buf,
456 SnapMessageId::GetStorageRanges,
457 GetStorageRanges,
458 GetStorageRangesMessage
459 );
460 decode_snap_message_variant!(
461 message_id,
462 buf,
463 SnapMessageId::StorageRanges,
464 StorageRanges,
465 StorageRangesMessage
466 );
467 decode_snap_message_variant!(
468 message_id,
469 buf,
470 SnapMessageId::GetByteCodes,
471 GetByteCodes,
472 GetByteCodesMessage
473 );
474 decode_snap_message_variant!(
475 message_id,
476 buf,
477 SnapMessageId::ByteCodes,
478 ByteCodes,
479 ByteCodesMessage
480 );
481 decode_snap_message_variant!(
482 message_id,
483 buf,
484 SnapMessageId::GetBlockAccessLists,
485 GetBlockAccessLists,
486 GetBlockAccessListsMessage
487 );
488 decode_snap_message_variant!(
489 message_id,
490 buf,
491 SnapMessageId::BlockAccessLists,
492 BlockAccessLists,
493 BlockAccessListsMessage
494 );
495
496 Err(alloy_rlp::Error::Custom("Unknown message ID"))
497 }
498
499 pub fn decode_versioned(version: SnapVersion, bytes: &[u8]) -> Result<Self, SnapProtocolError> {
504 let (&id, mut body) = bytes.split_first().ok_or(SnapProtocolError::Empty)?;
505 if !version.supports_message_id(id) {
506 return Err(SnapProtocolError::UnsupportedMessageId(id, version));
507 }
508 let msg = Self::decode(id, &mut body)?;
509 if !body.is_empty() {
510 return Err(SnapProtocolError::Rlp(alloy_rlp::Error::UnexpectedLength));
511 }
512 Ok(msg)
513 }
514}
515
516#[derive(Debug, Clone, PartialEq, Eq)]
518pub struct SlimAccountBody(TrieAccount);
519
520impl SlimAccountBody {
521 fn as_rlp(&self) -> SlimAccountBodyRef<'_> {
522 SlimAccountBodyRef {
523 nonce: self.0.nonce,
524 balance: self.0.balance,
525 storage_root: SlimAccountBodyRef::shorten(&self.0.storage_root, EMPTY_ROOT_HASH),
526 code_hash: SlimAccountBodyRef::shorten(&self.0.code_hash, KECCAK256_EMPTY),
527 }
528 }
529
530 fn restore(value: &[u8], empty: B256) -> alloy_rlp::Result<B256> {
532 match value {
533 [] => Ok(empty),
534 _ => B256::try_from(value).map_err(|_| alloy_rlp::Error::UnexpectedLength),
535 }
536 }
537}
538
539impl From<&TrieAccount> for SlimAccountBody {
540 #[allow(clippy::clone_on_copy)]
541 fn from(account: &TrieAccount) -> Self {
542 Self(account.clone())
543 }
544}
545
546impl Encodable for SlimAccountBody {
547 fn encode(&self, out: &mut dyn BufMut) {
548 self.as_rlp().encode(out);
549 }
550
551 fn length(&self) -> usize {
552 self.as_rlp().length()
553 }
554}
555
556impl Decodable for SlimAccountBody {
557 fn decode(buf: &mut &[u8]) -> alloy_rlp::Result<Self> {
558 let mut payload = Header::decode_bytes(buf, true)?;
559 let nonce = u64::decode(&mut payload)?;
560 let balance = U256::decode(&mut payload)?;
561 let storage_root =
562 Self::restore(Header::decode_bytes(&mut payload, false)?, EMPTY_ROOT_HASH)?;
563 let code_hash = Self::restore(Header::decode_bytes(&mut payload, false)?, KECCAK256_EMPTY)?;
564 if !payload.is_empty() {
565 return Err(alloy_rlp::Error::UnexpectedLength)
566 }
567 Ok(Self(TrieAccount::new(nonce, balance, storage_root, code_hash)))
568 }
569}
570
571#[cfg(any(test, feature = "arbitrary"))]
572impl<'a> arbitrary::Arbitrary<'a> for SlimAccountBody {
573 fn arbitrary(u: &mut arbitrary::Unstructured<'a>) -> arbitrary::Result<Self> {
574 let storage_root = if u.arbitrary()? { u.arbitrary()? } else { EMPTY_ROOT_HASH };
575 let code_hash = if u.arbitrary()? { u.arbitrary()? } else { KECCAK256_EMPTY };
576 Ok(Self(TrieAccount::new(u.arbitrary()?, u.arbitrary()?, storage_root, code_hash)))
577 }
578}
579
580#[derive(RlpEncodable)]
582struct SlimAccountBodyRef<'a> {
583 nonce: u64,
585 balance: U256,
587 storage_root: &'a [u8],
589 code_hash: &'a [u8],
591}
592
593impl<'a> SlimAccountBodyRef<'a> {
594 fn shorten(value: &'a B256, empty: B256) -> &'a [u8] {
596 if *value == empty {
597 &[]
598 } else {
599 value.as_slice()
600 }
601 }
602}
603
604#[cfg(test)]
605mod tests {
606 use super::*;
607 use test_case::test_case;
608
609 fn b256_from_u64(value: u64) -> B256 {
611 B256::left_padding_from(&value.to_be_bytes())
612 }
613
614 fn test_roundtrip(original: SnapProtocolMessage) {
616 let encoded = original.encode();
617
618 assert_eq!(encoded[0], original.message_id() as u8);
620
621 let mut buf = &encoded[1..];
622 let decoded = SnapProtocolMessage::decode(encoded[0], &mut buf).unwrap();
623
624 assert_eq!(decoded, original);
626 }
627
628 #[derive(alloy_rlp::RlpEncodable)]
632 struct GethStorageRequest {
633 request_id: u64,
634 root_hash: B256,
635 account_hashes: Vec<B256>,
636 origin: Bytes,
637 limit: Bytes,
638 response_bytes: u64,
639 }
640
641 #[test]
642 fn test_all_message_roundtrips() {
643 assert_eq!(SnapVersion::V2.message_count(), 10);
644
645 test_roundtrip(SnapProtocolMessage::GetAccountRange(GetAccountRangeMessage {
646 request_id: 42,
647 root_hash: b256_from_u64(123),
648 starting_hash: b256_from_u64(456),
649 limit_hash: b256_from_u64(789),
650 response_bytes: 1024,
651 }));
652
653 test_roundtrip(SnapProtocolMessage::AccountRange(AccountRangeMessage {
654 request_id: 42,
655 accounts: vec![AccountData::from_trie_account(
656 b256_from_u64(123),
657 &trie_account(EMPTY_ROOT_HASH, KECCAK256_EMPTY),
658 )],
659 proof: vec![Bytes::from(vec![4, 5, 6])],
660 }));
661
662 test_roundtrip(SnapProtocolMessage::GetStorageRanges(GetStorageRangesMessage {
663 request_id: 42,
664 root_hash: b256_from_u64(123),
665 account_hashes: vec![b256_from_u64(456)],
666 starting_hash: b256_from_u64(789).into(),
667 limit_hash: b256_from_u64(101112).into(),
668 response_bytes: 2048,
669 }));
670
671 test_roundtrip(SnapProtocolMessage::GetStorageRanges(GetStorageRangesMessage {
673 request_id: 43,
674 root_hash: b256_from_u64(123),
675 account_hashes: vec![b256_from_u64(456), b256_from_u64(789)],
676 starting_hash: RangeBound::default(),
677 limit_hash: RangeBound::default(),
678 response_bytes: 2048,
679 }));
680
681 test_roundtrip(SnapProtocolMessage::StorageRanges(StorageRangesMessage {
682 request_id: 42,
683 slots: vec![vec![StorageData {
684 hash: b256_from_u64(123),
685 data: Bytes::from(vec![1, 2, 3]),
686 }]],
687 proof: vec![Bytes::from(vec![4, 5, 6])],
688 }));
689
690 test_roundtrip(SnapProtocolMessage::GetByteCodes(GetByteCodesMessage {
691 request_id: 42,
692 hashes: vec![b256_from_u64(123)],
693 response_bytes: 1024,
694 }));
695
696 test_roundtrip(SnapProtocolMessage::ByteCodes(ByteCodesMessage {
697 request_id: 42,
698 codes: vec![Bytes::from(vec![1, 2, 3])],
699 }));
700
701 test_roundtrip(SnapProtocolMessage::GetBlockAccessLists(GetBlockAccessListsMessage {
702 request_id: 42,
703 block_hashes: vec![b256_from_u64(123), b256_from_u64(456)],
704 response_bytes: 4096,
705 }));
706
707 test_roundtrip(SnapProtocolMessage::BlockAccessLists(BlockAccessListsMessage {
708 request_id: 42,
709 block_access_lists: BlockAccessLists(vec![
710 Some(Bytes::from_static(&[alloy_rlp::EMPTY_LIST_CODE])),
711 Some(Bytes::from_static(&[0xc1, alloy_rlp::EMPTY_LIST_CODE])),
712 ]),
713 }));
714 }
715
716 #[test]
717 fn test_unknown_message_id() {
718 let data = Bytes::from(vec![1, 2, 3, 4]);
720 let mut buf = data.as_ref();
721
722 let result = SnapProtocolMessage::decode(255, &mut buf);
724
725 assert!(result.is_err());
726 if let Err(e) = result {
727 assert_eq!(e.to_string(), "Unknown message ID");
728 }
729 }
730
731 #[test]
732 fn test_snap_v2_message_validity() {
733 let v2 = SnapVersion::V2;
734 for id in 0x00..=0x05 {
736 assert!(v2.supports_message_id(id), "snap/2 should accept {id:#x}");
737 }
738 assert!(!v2.supports_message_id(0x06));
740 assert!(!v2.supports_message_id(0x07));
741 assert!(v2.supports_message_id(SnapMessageId::GetBlockAccessLists as u8));
743 assert!(v2.supports_message_id(SnapMessageId::BlockAccessLists as u8));
744 assert!(!v2.supports_message_id(0x0a));
745 assert!(!v2.supports_message_id(0xff));
746 }
747
748 #[test_case(
749 SnapProtocolMessage::GetAccountRange(GetAccountRangeMessage {
750 request_id: 1, root_hash: B256::ZERO, starting_hash: B256::ZERO,
751 limit_hash: B256::ZERO, response_bytes: 0,
752 }), 1, false ; "get_account_range is a request"
753 )]
754 #[test_case(
755 SnapProtocolMessage::AccountRange(AccountRangeMessage {
756 request_id: 2, accounts: vec![], proof: vec![],
757 }), 2, true ; "account_range is a response"
758 )]
759 #[test_case(
760 SnapProtocolMessage::GetStorageRanges(GetStorageRangesMessage {
761 request_id: 3, root_hash: B256::ZERO, account_hashes: vec![],
762 starting_hash: B256::ZERO.into(), limit_hash: B256::ZERO.into(), response_bytes: 0,
763 }), 3, false ; "get_storage_ranges is a request"
764 )]
765 #[test_case(
766 SnapProtocolMessage::StorageRanges(StorageRangesMessage {
767 request_id: 4, slots: vec![], proof: vec![],
768 }), 4, true ; "storage_ranges is a response"
769 )]
770 #[test_case(
771 SnapProtocolMessage::GetByteCodes(GetByteCodesMessage {
772 request_id: 5, hashes: vec![], response_bytes: 0,
773 }), 5, false ; "get_byte_codes is a request"
774 )]
775 #[test_case(
776 SnapProtocolMessage::ByteCodes(ByteCodesMessage { request_id: 6, codes: vec![] }),
777 6, true ; "byte_codes is a response"
778 )]
779 #[test_case(
780 SnapProtocolMessage::GetBlockAccessLists(GetBlockAccessListsMessage {
781 request_id: 7, block_hashes: vec![], response_bytes: 0,
782 }), 7, false ; "get_block_access_lists is a request"
783 )]
784 #[test_case(
785 SnapProtocolMessage::BlockAccessLists(BlockAccessListsMessage {
786 request_id: 8, block_access_lists: BlockAccessLists(vec![]),
787 }), 8, true ; "block_access_lists is a response"
788 )]
789 fn request_id_and_is_response(msg: SnapProtocolMessage, expected_id: u64, is_response: bool) {
790 assert_eq!(msg.request_id(), expected_id);
791 assert_eq!(msg.is_response(), is_response);
792 }
793
794 #[test_case(
795 SnapProtocolMessage::GetAccountRange(GetAccountRangeMessage {
796 request_id: 1, root_hash: B256::ZERO, starting_hash: B256::ZERO,
797 limit_hash: B256::ZERO, response_bytes: 0,
798 }) ; "get_account_range"
799 )]
800 #[test_case(
801 SnapProtocolMessage::AccountRange(AccountRangeMessage {
802 request_id: 1, accounts: vec![], proof: vec![],
803 }) ; "account_range"
804 )]
805 #[test_case(
806 SnapProtocolMessage::GetStorageRanges(GetStorageRangesMessage {
807 request_id: 1, root_hash: B256::ZERO, account_hashes: vec![],
808 starting_hash: B256::ZERO.into(), limit_hash: B256::ZERO.into(), response_bytes: 0,
809 }) ; "get_storage_ranges"
810 )]
811 #[test_case(
812 SnapProtocolMessage::StorageRanges(StorageRangesMessage {
813 request_id: 1, slots: vec![], proof: vec![],
814 }) ; "storage_ranges"
815 )]
816 #[test_case(
817 SnapProtocolMessage::GetByteCodes(GetByteCodesMessage {
818 request_id: 1, hashes: vec![], response_bytes: 0,
819 }) ; "get_byte_codes"
820 )]
821 #[test_case(
822 SnapProtocolMessage::ByteCodes(ByteCodesMessage { request_id: 1, codes: vec![] }) ;
823 "byte_codes"
824 )]
825 #[test_case(
826 SnapProtocolMessage::GetBlockAccessLists(GetBlockAccessListsMessage {
827 request_id: 1, block_hashes: vec![], response_bytes: 0,
828 }) ; "get_block_access_lists"
829 )]
830 #[test_case(
831 SnapProtocolMessage::BlockAccessLists(BlockAccessListsMessage {
832 request_id: 1, block_access_lists: BlockAccessLists(vec![]),
833 }) ; "block_access_lists"
834 )]
835 fn per_variant_request_id_and_round_trip(mut msg: SnapProtocolMessage) {
836 msg.set_request_id(42);
838 assert_eq!(msg.request_id(), 42);
839
840 let decoded =
842 SnapProtocolMessage::decode_versioned(SnapVersion::V2, &msg.encode()).unwrap();
843 assert_eq!(decoded, msg);
844 }
845
846 #[test]
847 fn decode_versioned_rejects_empty() {
848 assert!(matches!(
850 SnapProtocolMessage::decode_versioned(SnapVersion::V2, &[]),
851 Err(SnapProtocolError::Empty)
852 ));
853 }
854
855 #[test]
856 fn decode_versioned_rejects_trie_node_ids_in_v2() {
857 for id in [0x06u8, 0x07] {
860 assert!(matches!(
861 SnapProtocolMessage::decode_versioned(SnapVersion::V2, &[id]),
862 Err(SnapProtocolError::UnsupportedMessageId(got, SnapVersion::V2)) if got == id
863 ));
864 }
865 }
866
867 #[test]
868 fn decode_versioned_reports_malformed_body() {
869 assert!(matches!(
872 SnapProtocolMessage::decode_versioned(SnapVersion::V2, &[0x08, 0xff]),
873 Err(SnapProtocolError::Rlp(_))
874 ));
875 }
876
877 #[test]
878 fn decode_versioned_rejects_trailing_bytes() {
879 let original = SnapProtocolMessage::GetBlockAccessLists(GetBlockAccessListsMessage {
882 request_id: 7,
883 block_hashes: vec![b256_from_u64(1)],
884 response_bytes: 1024,
885 });
886 let mut framed = original.encode().to_vec();
887 framed.push(0xff);
888 assert!(matches!(
889 SnapProtocolMessage::decode_versioned(SnapVersion::V2, &framed),
890 Err(SnapProtocolError::Rlp(alloy_rlp::Error::UnexpectedLength))
891 ));
892 }
893
894 #[test]
895 fn get_storage_ranges_decodes_geths_empty_origin_and_limit() {
896 let body = alloy_rlp::encode(GethStorageRequest {
897 request_id: 21,
898 root_hash: B256::ZERO,
899 account_hashes: vec![B256::repeat_byte(1), B256::repeat_byte(2)],
900 origin: Bytes::new(),
901 limit: Bytes::new(),
902 response_bytes: 1024,
903 });
904 let mut framed = vec![SnapMessageId::GetStorageRanges as u8];
905 framed.extend_from_slice(&body);
906
907 let decoded = SnapProtocolMessage::decode_versioned(SnapVersion::V2, &framed).unwrap();
908 let SnapProtocolMessage::GetStorageRanges(msg) = decoded else {
909 panic!("expected a GetStorageRanges message");
910 };
911 assert_eq!(msg.starting_hash.unwrap_or(B256::ZERO), B256::ZERO);
912 assert_eq!(msg.limit_hash.unwrap_or(B256::repeat_byte(0xff)), B256::repeat_byte(0xff));
913 }
914
915 fn trie_account(storage_root: B256, code_hash: B256) -> TrieAccount {
916 TrieAccount::new(7, U256::from(42), storage_root, code_hash)
917 }
918
919 #[test]
920 fn account_range_matches_nested_account_wire_encoding() {
921 let wire = alloy_primitives::hex!(
923 "ea01e7e6a00101010101010101010101010101010101010101010101010101010101010101c4072a8080c0"
924 );
925 let message = AccountRangeMessage {
926 request_id: 1,
927 accounts: vec![AccountData::from_trie_account(
928 B256::repeat_byte(1),
929 &trie_account(EMPTY_ROOT_HASH, KECCAK256_EMPTY),
930 )],
931 proof: vec![],
932 };
933
934 assert_eq!(alloy_rlp::encode(&message), wire);
935 assert_eq!(alloy_rlp::decode_exact::<AccountRangeMessage>(&wire).unwrap(), message);
936 }
937
938 #[test]
939 fn account_data_rejects_byte_string_wrapped_body() {
940 let wire = alloy_primitives::hex!(
941 "e7a0010101010101010101010101010101010101010101010101010101010101010185c4072a8080"
942 );
943 assert!(alloy_rlp::decode_exact::<AccountData>(&wire).is_err());
944 }
945
946 #[test]
947 fn account_data_consumes_only_its_own_list() {
948 let account = AccountData::from_trie_account(
949 B256::repeat_byte(1),
950 &trie_account(B256::repeat_byte(2), B256::repeat_byte(3)),
951 );
952 let mut wire = alloy_rlp::encode(&account);
953 assert_eq!(wire.len(), account.length());
954 wire.push(0x80);
955 let mut input = wire.as_slice();
956 assert_eq!(AccountData::decode(&mut input).unwrap(), account);
957 assert_eq!(input, &[0x80]);
958 }
959
960 #[test]
961 fn account_data_rejects_missing_truncated_and_extra_body_items() {
962 for wire in [
963 "e1a00101010101010101010101010101010101010101010101010101010101010101",
964 "e4a00101010101010101010101010101010101010101010101010101010101010101c4072a",
965 "e7a00101010101010101010101010101010101010101010101010101010101010101c4072a808080",
966 ] {
967 let bytes = alloy_primitives::hex::decode(wire).unwrap();
968 assert!(alloy_rlp::decode_exact::<AccountData>(&bytes).is_err());
969 }
970 }
971
972 #[test]
973 fn slim_body_elides_empty_storage_and_code() {
974 let account = trie_account(EMPTY_ROOT_HASH, KECCAK256_EMPTY);
975 let hash = B256::repeat_byte(1);
976 let encoded = AccountData::from_trie_account(hash, &account);
977
978 let body = alloy_rlp::encode(&encoded.body);
979 assert_eq!(body, alloy_primitives::hex!("c4072a8080"));
980 assert_eq!(alloy_rlp::decode_exact::<SlimAccountBody>(&body).unwrap(), encoded.body);
981 assert_eq!(encoded.trie_account(), account);
982 assert_eq!(encoded.into_trie_entry(), (hash, account));
983 }
984
985 #[test]
986 #[allow(clippy::clone_on_copy)]
987 fn slim_body_keeps_non_default_storage_and_code() {
988 let account = trie_account(B256::repeat_byte(2), B256::repeat_byte(3));
989 let encoded = AccountData::from_trie_account(B256::repeat_byte(1), &account);
990
991 let body = alloy_rlp::encode(&encoded.body);
992 assert_eq!(body, alloy_rlp::encode(account.clone()));
993 assert_eq!(alloy_rlp::decode_exact::<SlimAccountBody>(&body).unwrap(), encoded.body);
994 assert_eq!(encoded.trie_account(), account);
995 }
996
997 #[test_case(1)]
998 #[test_case(16)]
999 #[test_case(31)]
1000 #[test_case(33)]
1001 fn slim_body_rejects_invalid_hash_lengths(len: usize) {
1002 let invalid = vec![0xaa; len];
1003 for (storage_root, code_hash) in
1004 [(invalid.as_slice(), &[][..]), (&[][..], invalid.as_slice())]
1005 {
1006 let body = alloy_rlp::encode(SlimAccountBodyRef {
1007 nonce: 7,
1008 balance: U256::from(42),
1009 storage_root,
1010 code_hash,
1011 });
1012 assert_eq!(
1013 alloy_rlp::decode_exact::<SlimAccountBody>(&body),
1014 Err(alloy_rlp::Error::UnexpectedLength)
1015 );
1016 }
1017 }
1018
1019 #[test_case("c0")]
1020 #[test_case("c3072a80")]
1021 #[test_case("c5072a808080")]
1022 #[test_case("c4072a808000")]
1023 fn slim_body_rejects_missing_and_extra_fields(wire: &str) {
1024 let bytes = alloy_primitives::hex::decode(wire).unwrap();
1025 assert!(alloy_rlp::decode_exact::<SlimAccountBody>(&bytes).is_err());
1026 }
1027
1028 #[test]
1029 fn storage_data_carries_the_trie_leaf_encoding() {
1030 let value = U256::from(1234);
1031 let slot = StorageData::from_value(B256::repeat_byte(4), value);
1032
1033 assert_eq!(slot.data.as_ref(), alloy_rlp::encode(value));
1036 assert_eq!(slot.value().unwrap(), value);
1037 }
1038
1039 #[test]
1040 fn storage_data_rejects_trailing_bytes() {
1041 let mut slot = StorageData::from_value(B256::repeat_byte(4), U256::from(1));
1042 slot.data = [slot.data.as_ref(), &[0x00]].concat().into();
1043
1044 assert!(slot.value().is_err());
1045 }
1046}