Skip to main content

reth_eth_wire/
p2pstream.rs

1use crate::{
2    capability::SharedCapabilities,
3    disconnect::CanDisconnect,
4    errors::{P2PHandshakeError, P2PStreamError},
5    pinger::{Pinger, PingerEvent},
6    protocol::ProtocolIngressLimits,
7    DisconnectReason, HelloMessage, HelloMessageWithProtocols,
8};
9use alloy_primitives::{
10    bytes::{Buf, BufMut, Bytes, BytesMut},
11    hex,
12};
13use alloy_rlp::{Decodable, Encodable, Error as RlpError, EMPTY_LIST_CODE};
14use futures::{Sink, SinkExt, StreamExt};
15use pin_project::pin_project;
16use reth_codecs::add_arbitrary_tests;
17use reth_metrics::metrics::counter;
18use reth_primitives_traits::GotExpected;
19use std::{
20    collections::VecDeque,
21    future::Future,
22    io,
23    pin::Pin,
24    task::{ready, Context, Poll},
25    time::{Duration, Instant},
26};
27use tokio_stream::Stream;
28use tracing::{debug, trace};
29
30#[cfg(feature = "serde")]
31use serde::{Deserialize, Serialize};
32
33/// [`MAX_PAYLOAD_SIZE`] is the maximum size of an uncompressed message payload.
34/// This is defined in [EIP-706](https://eips.ethereum.org/EIPS/eip-706).
35const MAX_PAYLOAD_SIZE: usize = 16 * 1024 * 1024;
36
37/// [`MAX_RESERVED_MESSAGE_ID`] is the maximum message ID reserved for the `p2p` subprotocol. If
38/// there are any incoming messages with an ID greater than this, they are subprotocol messages.
39pub const MAX_RESERVED_MESSAGE_ID: u8 = 0x0f;
40
41/// [`MAX_P2P_MESSAGE_ID`] is the maximum message ID in use for the `p2p` subprotocol.
42const MAX_P2P_MESSAGE_ID: u8 = P2PMessageID::Pong as u8;
43
44/// Snappy framed RLP empty list payload used by fixed `p2p` ping/pong control messages.
45const SNAPPY_EMPTY_LIST_PAYLOAD: &[u8] = &[0x01, 0x00, EMPTY_LIST_CODE];
46
47/// Wire-encoded `p2p` ping control message.
48const SNAPPY_PING_MESSAGE: &[u8] = &[0x02, 0x01, 0x00, EMPTY_LIST_CODE];
49
50/// Wire-encoded `p2p` pong control message.
51const SNAPPY_PONG_MESSAGE: &[u8] = &[0x03, 0x01, 0x00, EMPTY_LIST_CODE];
52
53/// [`HANDSHAKE_TIMEOUT`] determines the amount of time to wait before determining that a `p2p`
54/// handshake has timed out.
55pub const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
56
57/// [`PING_TIMEOUT`] determines the amount of time to wait before determining that a `p2p` ping has
58/// timed out.
59const PING_TIMEOUT: Duration = Duration::from_secs(15);
60
61/// [`PING_INTERVAL`] determines the amount of time to wait between sending `p2p` ping messages
62/// when the peer is responsive.
63const PING_INTERVAL: Duration = Duration::from_secs(60);
64
65/// Maximum number of incoming pings that can arrive in a burst.
66const PING_TOKEN_BUCKET_CAPACITY: u8 = 5;
67
68/// [`MAX_P2P_CAPACITY`] is the maximum number of messages that can be buffered to be sent in the
69/// `p2p` stream.
70///
71/// Note: this default is rather low because it is expected that the [`P2PStream`] wraps an
72/// [`ECIESStream`](reth_ecies::stream::ECIESStream) which internally already buffers a few MB of
73/// encoded data.
74const MAX_P2P_CAPACITY: usize = 2;
75
76/// Maximum size of the reusable compression scratch buffer in [`P2PStream`], covering the snappy
77/// worst case of typical broadcast messages (soft-capped around 128KiB).
78///
79/// Messages with a larger compressed worst case are compressed through a one-off allocation
80/// instead, so a single oversized message neither grows the scratch buffer for the connection's
81/// lifetime nor causes shrink/regrow churn, see [`compress_frame`].
82const MAX_COMPRESS_SCRATCH_SIZE: usize = 256 * 1024;
83
84/// An un-authenticated [`P2PStream`]. This is consumed and returns a [`P2PStream`] after the
85/// `Hello` handshake is completed.
86#[pin_project]
87#[derive(Debug)]
88pub struct UnauthedP2PStream<S> {
89    #[pin]
90    inner: S,
91}
92
93impl<S> UnauthedP2PStream<S> {
94    /// Create a new `UnauthedP2PStream` from a type `S` which implements `Stream` and `Sink`.
95    pub const fn new(inner: S) -> Self {
96        Self { inner }
97    }
98
99    /// Returns a reference to the inner stream.
100    pub const fn inner(&self) -> &S {
101        &self.inner
102    }
103}
104
105impl<S> UnauthedP2PStream<S>
106where
107    S: Stream<Item = io::Result<BytesMut>> + Sink<Bytes, Error = io::Error> + Unpin,
108{
109    /// Consumes the `UnauthedP2PStream` and returns a `P2PStream` after the `Hello` handshake is
110    /// completed successfully. This also returns the `Hello` message sent by the remote peer.
111    pub async fn handshake(
112        mut self,
113        hello: HelloMessageWithProtocols,
114    ) -> Result<(P2PStream<S>, HelloMessage), P2PStreamError> {
115        trace!(?hello, "sending p2p hello to peer");
116
117        // send our hello message with the Sink
118        self.inner.send(alloy_rlp::encode(P2PMessage::Hello(hello.message())).into()).await?;
119
120        let first_message_bytes = tokio::time::timeout(HANDSHAKE_TIMEOUT, self.inner.next())
121            .await
122            .or(Err(P2PStreamError::HandshakeError(P2PHandshakeError::Timeout)))?
123            .ok_or(P2PStreamError::HandshakeError(P2PHandshakeError::NoResponse))??;
124
125        // Check that the uncompressed message length does not exceed the max payload size.
126        // Note: The first message (Hello/Disconnect) is not snappy compressed. We will check the
127        // decompressed length again for subsequent messages after the handshake.
128        if first_message_bytes.len() > MAX_PAYLOAD_SIZE {
129            return Err(P2PStreamError::MessageTooBig {
130                message_size: first_message_bytes.len(),
131                max_size: MAX_PAYLOAD_SIZE,
132            })
133        }
134
135        // The first message sent MUST be a hello OR disconnect message
136        //
137        // If the first message is a disconnect message, we should not decode using
138        // Decodable::decode, because the first message (either Disconnect or Hello) is not snappy
139        // compressed, and the Decodable implementation assumes that non-hello messages are snappy
140        // compressed.
141        let their_hello = match P2PMessage::decode(&mut &first_message_bytes[..]) {
142            Ok(P2PMessage::Hello(hello)) => Ok(hello),
143            Ok(P2PMessage::Disconnect(reason)) => {
144                if matches!(reason, DisconnectReason::TooManyPeers) {
145                    // Too many peers is a very common disconnect reason that spams the DEBUG logs
146                    trace!(%reason, "Disconnected by peer during handshake");
147                } else {
148                    debug!(%reason, "Disconnected by peer during handshake");
149                };
150                counter!("p2pstream.disconnected_errors").increment(1);
151                Err(P2PStreamError::HandshakeError(P2PHandshakeError::Disconnected(reason)))
152            }
153            Err(err) => {
154                debug!(%err, msg=%hex::encode(&first_message_bytes), "Failed to decode first message from peer");
155                Err(P2PStreamError::HandshakeError(err.into()))
156            }
157            Ok(msg) => {
158                debug!(?msg, "expected hello message but received another message");
159                Err(P2PStreamError::HandshakeError(P2PHandshakeError::NonHelloMessageInHandshake))
160            }
161        }?;
162
163        trace!(
164            hello=?their_hello,
165            "validating incoming p2p hello from peer"
166        );
167
168        if (hello.protocol_version as u8) != their_hello.protocol_version as u8 {
169            // send a disconnect message notifying the peer of the protocol version mismatch
170            self.send_disconnect(DisconnectReason::IncompatibleP2PProtocolVersion).await?;
171            return Err(P2PStreamError::MismatchedProtocolVersion(GotExpected {
172                got: their_hello.protocol_version,
173                expected: hello.protocol_version,
174            }))
175        }
176
177        // determine shared capabilities (currently returns only one capability)
178        let capability_res =
179            SharedCapabilities::try_new(hello.protocols, their_hello.capabilities.clone());
180
181        let shared_capability = match capability_res {
182            Err(err) => {
183                // we don't share any capabilities, send a disconnect message
184                self.send_disconnect(DisconnectReason::UselessPeer).await?;
185                Err(err)
186            }
187            Ok(cap) => Ok(cap),
188        }?;
189
190        let stream = P2PStream::new(self.inner, shared_capability);
191
192        Ok((stream, their_hello))
193    }
194}
195
196impl<S> UnauthedP2PStream<S>
197where
198    S: Sink<Bytes, Error = io::Error> + Unpin,
199{
200    /// Send a disconnect message during the handshake. This is sent without snappy compression.
201    pub async fn send_disconnect(
202        &mut self,
203        reason: DisconnectReason,
204    ) -> Result<(), P2PStreamError> {
205        trace!(
206            %reason,
207            "Sending disconnect message during the handshake",
208        );
209        self.inner
210            .send(Bytes::from(alloy_rlp::encode(P2PMessage::Disconnect(reason))))
211            .await
212            .map_err(P2PStreamError::Io)
213    }
214}
215
216impl<S> CanDisconnect<Bytes> for P2PStream<S>
217where
218    S: Sink<Bytes, Error = io::Error> + Unpin + Send + Sync,
219{
220    fn disconnect(
221        &mut self,
222        reason: DisconnectReason,
223    ) -> Pin<Box<dyn Future<Output = Result<(), P2PStreamError>> + Send + '_>> {
224        Box::pin(async move { self.disconnect(reason).await })
225    }
226}
227
228/// A `P2PStream` wraps over any `Stream` that yields bytes and makes it compatible with `p2p`
229/// protocol messages.
230///
231/// This stream supports multiple shared capabilities, that were negotiated during the handshake.
232///
233/// ### Message-ID based multiplexing
234///
235/// > Each capability is given as much of the message-ID space as it needs. All such capabilities
236/// > must statically specify how many message IDs they require. On connection and reception of the
237/// > Hello message, both peers have equivalent information about what capabilities they share
238/// > (including versions) and are able to form consensus over the composition of message ID space.
239///
240/// > Message IDs are assumed to be compact from ID 0x10 onwards (0x00-0x0f is reserved for the
241/// > "p2p" capability) and given to each shared (equal-version, equal-name) capability in
242/// > alphabetic order. Capability names are case-sensitive. Capabilities which are not shared are
243/// > ignored. If multiple versions are shared of the same (equal name) capability, the numerically
244/// > highest wins, others are ignored.
245///
246/// See also <https://github.com/ethereum/devp2p/blob/master/rlpx.md#message-id-based-multiplexing>
247///
248/// This stream emits _non-empty_ Bytes that start with the normalized message id, so that the first
249/// byte of each message starts from 0. If this stream only supports a single capability, for
250/// example `eth` then the first byte of each message will match
251/// [EthMessageID](reth_eth_wire_types::message::EthMessageID).
252///
253/// ### Sink behavior
254///
255/// The [`Sink`] impl batches writes: queued messages are drained into the underlying sink
256/// unflushed, and the caller is responsible for driving [`Sink::poll_flush`] to deliver them to
257/// the wire. Queued `p2p` control messages (ping/pong/disconnect) are the exception: they force a
258/// flush from [`Sink::poll_ready`]. Keepalive pings are generated in `poll_ready`, so the sink
259/// half must be polled regularly even if the caller has nothing to send.
260#[pin_project]
261#[derive(Debug)]
262pub struct P2PStream<S> {
263    #[pin]
264    inner: S,
265
266    /// The snappy encoder used for compressing outgoing messages
267    encoder: snap::raw::Encoder,
268
269    /// Reusable scratch buffer for compressing outgoing messages, see [`compress_frame`].
270    ///
271    /// Grow-only and capped at [`MAX_COMPRESS_SCRATCH_SIZE`]; kept fully initialized, so
272    /// zero-initialization is only paid when the buffer grows and each message only copies out
273    /// the exact compressed size instead of zeroing a worst-case sized buffer per message.
274    compress_scratch: Vec<u8>,
275
276    /// The snappy decoder used for decompressing incoming messages
277    decoder: snap::raw::Decoder,
278
279    /// The state machine used for keeping track of the peer's ping status.
280    pinger: Pinger,
281
282    /// Per-connection limit for incoming ping bursts.
283    ping_token_bucket: PingTokenBucket,
284
285    /// The supported capability for this stream.
286    shared_capabilities: SharedCapabilities,
287
288    /// Explicit frame limits registered by installed subprotocol handlers.
289    inbound_protocol_limits: Vec<InboundProtocolLimit>,
290
291    /// Outgoing messages buffered for sending to the underlying stream.
292    outgoing_messages: VecDeque<Bytes>,
293
294    /// Maximum number of messages that we can buffer here before the [Sink] impl returns
295    /// [`Poll::Pending`].
296    outgoing_message_buffer_capacity: usize,
297
298    /// Whether this stream is currently in the process of disconnecting by sending a disconnect
299    /// message.
300    disconnecting: bool,
301
302    /// Whether the underlying sink has accepted messages that still need to be flushed.
303    needs_flush: bool,
304
305    /// Whether a queued p2p control message needs to be flushed even if no subprotocol messages
306    /// are sent by the caller.
307    needs_control_flush: bool,
308}
309
310impl<S> P2PStream<S> {
311    /// Create a new [`P2PStream`] from the provided stream.
312    /// New [`P2PStream`]s are assumed to have completed the `p2p` handshake successfully and are
313    /// ready to send and receive subprotocol messages.
314    pub fn new(inner: S, shared_capabilities: SharedCapabilities) -> Self {
315        Self {
316            inner,
317            encoder: snap::raw::Encoder::new(),
318            compress_scratch: Vec::new(),
319            decoder: snap::raw::Decoder::new(),
320            pinger: Pinger::new(PING_INTERVAL, PING_TIMEOUT),
321            ping_token_bucket: PingTokenBucket::new(Instant::now()),
322            shared_capabilities,
323            inbound_protocol_limits: Vec::new(),
324            outgoing_messages: VecDeque::new(),
325            outgoing_message_buffer_capacity: MAX_P2P_CAPACITY,
326            disconnecting: false,
327            needs_flush: false,
328            needs_control_flush: false,
329        }
330    }
331
332    /// Returns a reference to the inner stream.
333    pub const fn inner(&self) -> &S {
334        &self.inner
335    }
336
337    /// Sets a custom outgoing message buffer capacity.
338    ///
339    /// # Panics
340    ///
341    /// If the provided capacity is `0`.
342    pub const fn set_outgoing_message_buffer_capacity(&mut self, capacity: usize) {
343        assert!(capacity != 0);
344        self.outgoing_message_buffer_capacity = capacity;
345    }
346
347    /// Returns the shared capabilities for this stream.
348    ///
349    /// This includes all the shared capabilities that were negotiated during the handshake and
350    /// their offsets based on the number of messages of each capability.
351    pub const fn shared_capabilities(&self) -> &SharedCapabilities {
352        &self.shared_capabilities
353    }
354
355    pub(crate) fn set_protocol_ingress_limits(
356        &mut self,
357        capability: &crate::capability::SharedCapability,
358        limits: ProtocolIngressLimits,
359    ) {
360        let Some(max_frame_bytes) = limits.max_frame_bytes() else { return };
361        if capability.num_messages() == 0 {
362            return
363        }
364
365        let start = capability.message_id_offset();
366        let end = u16::from(start) + u16::from(capability.num_messages());
367        let configured = InboundProtocolLimit {
368            capability: capability.capability().into_owned(),
369            start,
370            end,
371            max_frame_bytes,
372        };
373
374        if let Some(current) = self
375            .inbound_protocol_limits
376            .iter_mut()
377            .find(|current| current.capability == configured.capability)
378        {
379            *current = configured;
380        } else {
381            self.inbound_protocol_limits.push(configured);
382        }
383    }
384
385    fn inbound_protocol_limit(&self, message_id: u8) -> Option<&InboundProtocolLimit> {
386        self.inbound_protocol_limits.iter().find(|limit| limit.contains(message_id))
387    }
388
389    /// Returns `true` if the stream has outgoing capacity.
390    fn has_outgoing_capacity(&self) -> bool {
391        self.outgoing_messages.len() < self.outgoing_message_buffer_capacity
392    }
393
394    /// Queues in a _snappy_ encoded [`P2PMessage::Pong`] message.
395    fn send_pong(&mut self) {
396        self.outgoing_messages.push_back(Bytes::from_static(SNAPPY_PONG_MESSAGE));
397        self.needs_control_flush = true;
398    }
399
400    /// Queues in a _snappy_ encoded [`P2PMessage::Ping`] message.
401    pub fn send_ping(&mut self) {
402        self.outgoing_messages.push_back(Bytes::from_static(SNAPPY_PING_MESSAGE));
403        self.needs_control_flush = true;
404    }
405}
406
407/// Per-connection bucket that restores one incoming ping token per second.
408#[derive(Debug)]
409struct PingTokenBucket {
410    tokens: u8,
411    last_refill: Instant,
412}
413
414impl PingTokenBucket {
415    const fn new(now: Instant) -> Self {
416        Self { tokens: PING_TOKEN_BUCKET_CAPACITY, last_refill: now }
417    }
418
419    fn try_take(&mut self, now: Instant) -> bool {
420        let refill = now.saturating_duration_since(self.last_refill).as_secs();
421        if refill > 0 {
422            self.tokens = u64::from(self.tokens)
423                .saturating_add(refill)
424                .min(u64::from(PING_TOKEN_BUCKET_CAPACITY)) as u8;
425            self.last_refill += Duration::from_secs(refill);
426        }
427
428        if self.tokens == 0 {
429            return false
430        }
431
432        if self.tokens == PING_TOKEN_BUCKET_CAPACITY {
433            self.last_refill = now;
434        }
435        self.tokens -= 1;
436        true
437    }
438}
439
440#[derive(Debug)]
441struct InboundProtocolLimit {
442    capability: crate::Capability,
443    start: u8,
444    end: u16,
445    max_frame_bytes: usize,
446}
447
448impl InboundProtocolLimit {
449    const fn contains(&self, message_id: u8) -> bool {
450        message_id >= self.start && (message_id as u16) < self.end
451    }
452}
453
454/// Gracefully disconnects the connection by sending a disconnect message and stop reading new
455/// messages.
456pub trait DisconnectP2P {
457    /// Starts to gracefully disconnect.
458    fn start_disconnect(&mut self, reason: DisconnectReason) -> Result<(), P2PStreamError>;
459
460    /// Returns `true` if the connection is about to disconnect.
461    fn is_disconnecting(&self) -> bool;
462}
463
464impl<S> DisconnectP2P for P2PStream<S> {
465    /// Starts to gracefully disconnect the connection by sending a Disconnect message and stop
466    /// reading new messages.
467    ///
468    /// Once disconnect process has started, the [`Stream`] will terminate immediately.
469    ///
470    /// # Errors
471    ///
472    /// Returns an error only if the message fails to compress.
473    fn start_disconnect(&mut self, reason: DisconnectReason) -> Result<(), P2PStreamError> {
474        // clear any buffered messages and queue in
475        self.outgoing_messages.clear();
476        let disconnect = P2PMessage::Disconnect(reason);
477        let buf = alloy_rlp::encode(&disconnect);
478
479        // we do not add the capability offset because the disconnect message is a `p2p` reserved
480        // message
481        let compressed =
482            compress_frame(&mut self.encoder, &mut self.compress_scratch, buf[0], &buf[1..])
483                .map_err(|err| {
484                    debug!(
485                        %err,
486                        msg=%hex::encode(&buf[1..]),
487                        "error compressing disconnect"
488                    );
489                    err
490                })?;
491
492        self.outgoing_messages.push_back(compressed);
493        self.needs_control_flush = true;
494        self.disconnecting = true;
495        Ok(())
496    }
497
498    fn is_disconnecting(&self) -> bool {
499        self.disconnecting
500    }
501}
502
503impl<S> P2PStream<S>
504where
505    S: Sink<Bytes, Error = io::Error> + Unpin + Send,
506{
507    /// Disconnects the connection by sending a disconnect message.
508    ///
509    /// This future resolves once the disconnect message has been sent and the stream has been
510    /// closed.
511    pub async fn disconnect(&mut self, reason: DisconnectReason) -> Result<(), P2PStreamError> {
512        self.start_disconnect(reason)?;
513        self.close().await
514    }
515}
516
517impl<S> P2PStream<S>
518where
519    S: Sink<Bytes, Error = io::Error> + Unpin,
520{
521    /// Drains queued p2p frames into the underlying sink without flushing the underlying sink.
522    fn poll_drain_outgoing(
523        mut self: Pin<&mut Self>,
524        cx: &mut Context<'_>,
525    ) -> Poll<Result<(), P2PStreamError>> {
526        let mut this = self.as_mut().project();
527        while !this.outgoing_messages.is_empty() {
528            ready!(this.inner.as_mut().poll_ready(cx))?;
529            let message = this.outgoing_messages.pop_front().expect("checked non-empty");
530            this.inner.as_mut().start_send(message)?;
531            *this.needs_flush = true;
532        }
533
534        Poll::Ready(Ok(()))
535    }
536}
537
538// S must also be `Sink` because we need to be able to respond with ping messages to follow the
539// protocol
540impl<S> Stream for P2PStream<S>
541where
542    S: Stream<Item = io::Result<BytesMut>> + Sink<Bytes, Error = io::Error> + Unpin,
543{
544    type Item = Result<BytesMut, P2PStreamError>;
545
546    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
547        let this = self.get_mut();
548
549        if this.disconnecting {
550            // if disconnecting, stop reading messages
551            return Poll::Ready(None)
552        }
553
554        let mut ping_batch_time = None;
555
556        // we should loop here to ensure we don't return Poll::Pending if we have a message to
557        // return behind any pings we need to respond to
558        while let Poll::Ready(res) = this.inner.poll_next_unpin(cx) {
559            let bytes = match res {
560                Some(Ok(bytes)) => bytes,
561                Some(Err(err)) => return Poll::Ready(Some(Err(err.into()))),
562                None => return Poll::Ready(None),
563            };
564
565            if bytes.is_empty() {
566                // empty messages are not allowed
567                return Poll::Ready(Some(Err(P2PStreamError::EmptyProtocolMessage)))
568            }
569
570            // first decode disconnect reasons, because they can be encoded in a variety of forms
571            // over the wire, in both snappy compressed and uncompressed forms.
572            //
573            // see: [crate::disconnect::tests::test_decode_known_reasons]
574            let id = bytes[0];
575            if id == P2PMessageID::Disconnect as u8 {
576                // We can't handle the error here because disconnect reasons are encoded as both:
577                // * snappy compressed, AND
578                // * uncompressed
579                // over the network.
580                //
581                // If the decoding succeeds, we already checked the id and know this is a
582                // disconnect message, so we can return with the reason.
583                //
584                // If the decoding fails, we continue, and will attempt to decode it again if the
585                // message is snappy compressed. Failure handling in that step is the primary point
586                // where an error is returned if the disconnect reason is malformed.
587                if let Ok(reason) = DisconnectReason::decode(&mut &bytes[1..]) {
588                    return Poll::Ready(Some(Err(P2PStreamError::Disconnected(reason))))
589                }
590            }
591
592            if id == P2PMessageID::Ping as u8 || id == P2PMessageID::Pong as u8 {
593                validate_ping_pong_payload(&mut this.decoder, id, &bytes[1..])?;
594
595                if id == P2PMessageID::Ping as u8 {
596                    // Use the timestamp of the first ping for every ping in this poll, so buffered
597                    // pings form one burst. The tradeoff is that a poll that takes more than one
598                    // second can reject a later ping. Considered acceptable.
599                    let now = *ping_batch_time.get_or_insert_with(Instant::now);
600                    if !this.ping_token_bucket.try_take(now) {
601                        return Poll::Ready(Some(Err(P2PStreamError::TooManyPings)))
602                    }
603
604                    trace!("Received Ping, Sending Pong");
605                    this.send_pong();
606                    // This is required because the `Sink` may not be polled externally, and if
607                    // that happens, the pong will never be sent.
608                    cx.waker().wake_by_ref();
609                } else {
610                    // if we were waiting for a pong, this will reset the pinger state
611                    this.pinger.on_pong()?;
612                }
613                continue
614            }
615
616            if id > MAX_RESERVED_MESSAGE_ID && this.shared_capabilities.find_by_offset(id).is_none()
617            {
618                return Poll::Ready(Some(Err(P2PStreamError::UnknownSubprotocolMessageId(id))))
619            }
620
621            // first check that the compressed message length does not exceed the max
622            // payload size
623            let decompressed_len = snap::raw::decompress_len(&bytes[1..])?;
624            if decompressed_len > MAX_PAYLOAD_SIZE {
625                return Poll::Ready(Some(Err(P2PStreamError::MessageTooBig {
626                    message_size: decompressed_len,
627                    max_size: MAX_PAYLOAD_SIZE,
628                })))
629            }
630
631            let frame_len = decompressed_len + 1;
632            if let Some(limit) = this.inbound_protocol_limit(id) &&
633                frame_len > limit.max_frame_bytes
634            {
635                counter!("p2pstream.subprotocol_message_too_big").increment(1);
636                return Poll::Ready(Some(Err(P2PStreamError::SubprotocolMessageTooBig {
637                    capability: limit.capability.clone(),
638                    message_size: frame_len,
639                    max_size: limit.max_frame_bytes,
640                })))
641            }
642
643            // create a buffer to hold the decompressed message, adding a byte to the length for
644            // the message ID byte, which is the first byte in this buffer
645            let mut decompress_buf = BytesMut::zeroed(frame_len);
646
647            // each message following a successful handshake is compressed with snappy, so we need
648            // to decompress the message before we can decode it.
649            this.decoder.decompress(&bytes[1..], &mut decompress_buf[1..]).map_err(|err| {
650                debug!(
651                    %err,
652                    msg=%hex::encode(&bytes[1..]),
653                    "error decompressing p2p message"
654                );
655                err
656            })?;
657
658            match id {
659                _ if id == P2PMessageID::Hello as u8 => {
660                    // we have received a hello message outside of the handshake, so we will return
661                    // an error
662                    return Poll::Ready(Some(Err(P2PStreamError::HandshakeError(
663                        P2PHandshakeError::HelloNotInHandshake,
664                    ))))
665                }
666                _ if id == P2PMessageID::Disconnect as u8 => {
667                    // At this point, the `decompress_buf` contains the snappy decompressed
668                    // disconnect message.
669                    //
670                    // It's possible we already tried to RLP decode this, but it was snappy
671                    // compressed, so we need to RLP decode it again.
672                    let reason = DisconnectReason::decode(&mut &decompress_buf[1..]).inspect_err(|err| {
673                        debug!(
674                            %err, msg=%hex::encode(&decompress_buf[1..]), "Failed to decode disconnect message from peer"
675                        );
676                    })?;
677                    return Poll::Ready(Some(Err(P2PStreamError::Disconnected(reason))))
678                }
679                _ if id > MAX_P2P_MESSAGE_ID && id <= MAX_RESERVED_MESSAGE_ID => {
680                    // we have received an unknown reserved message
681                    return Poll::Ready(Some(Err(P2PStreamError::UnknownReservedMessageId(id))))
682                }
683                _ => {
684                    // we have received a message that is outside the `p2p` reserved message space,
685                    // so it is a subprotocol message.
686
687                    // Peers must be able to identify messages meant for different subprotocols
688                    // using a single message ID byte, and those messages must be distinct from the
689                    // lower-level `p2p` messages.
690                    //
691                    // To ensure that messages for subprotocols are distinct from messages meant
692                    // for the `p2p` capability, message IDs 0x00 - 0x0f are reserved for `p2p`
693                    // messages, so subprotocol messages must have an ID of 0x10 or higher.
694                    //
695                    // To ensure that messages for two different capabilities are distinct from
696                    // each other, all shared capabilities are first ordered lexicographically.
697                    // Message IDs are then reserved in this order, starting at 0x10, reserving a
698                    // message ID for each message the capability supports.
699                    //
700                    // For example, if the shared capabilities are `eth/67` (containing 10
701                    // messages), and "qrs/65" (containing 8 messages):
702                    //
703                    //  * The special case of `p2p`: `p2p` is reserved message IDs 0x00 - 0x0f.
704                    //  * `eth/67` is reserved message IDs 0x10 - 0x19.
705                    //  * `qrs/65` is reserved message IDs 0x1a - 0x21.
706                    //
707                    decompress_buf[0] = bytes[0] - MAX_RESERVED_MESSAGE_ID - 1;
708
709                    return Poll::Ready(Some(Ok(decompress_buf)))
710                }
711            }
712        }
713
714        Poll::Pending
715    }
716}
717
718/// Validates a compressed Ping or Pong payload without allocating from its advertised size.
719fn validate_ping_pong_payload(
720    decoder: &mut snap::raw::Decoder,
721    message_id: u8,
722    compressed_payload: &[u8],
723) -> Result<(), P2PStreamError> {
724    if snap::raw::decompress_len(compressed_payload)? != 1 {
725        return Err(P2PStreamError::InvalidPingPongPayload(message_id))
726    }
727
728    let mut payload = [0u8; 1];
729    decoder.decompress(compressed_payload, &mut payload)?;
730    if payload != [EMPTY_LIST_CODE] {
731        return Err(P2PStreamError::InvalidPingPongPayload(message_id))
732    }
733
734    Ok(())
735}
736
737impl<S> Sink<Bytes> for P2PStream<S>
738where
739    S: Sink<Bytes, Error = io::Error> + Unpin,
740{
741    type Error = P2PStreamError;
742
743    fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
744        let this = self.as_mut().get_mut();
745
746        // Poll the pinger to determine if we should send a ping; `send_ping` and
747        // `start_disconnect` set `needs_control_flush`.
748        match this.pinger.poll_ping(cx) {
749            Poll::Pending => {}
750            Poll::Ready(Ok(PingerEvent::Ping)) => {
751                this.send_ping();
752            }
753            Poll::Ready(Ok(PingerEvent::Timeout) | Err(_)) => {
754                this.start_disconnect(DisconnectReason::PingTimeout)?;
755            }
756        }
757
758        // Control messages (ping/pong/disconnect) must reach the wire even if the caller never
759        // sends a message, so they force a flush. Subprotocol messages are only drained into the
760        // underlying sink (unflushed) once the buffer is full; the caller is responsible for
761        // flushing the batch via `poll_flush`.
762        if self.needs_control_flush {
763            ready!(self.as_mut().poll_flush(cx))?;
764        } else if !self.has_outgoing_capacity() {
765            ready!(self.as_mut().poll_drain_outgoing(cx))?;
766        }
767
768        // both branches above fully drain the queue, and an empty queue always has capacity
769        debug_assert!(self.has_outgoing_capacity());
770        Poll::Ready(Ok(()))
771    }
772
773    fn start_send(self: Pin<&mut Self>, item: Bytes) -> Result<(), Self::Error> {
774        if item.len() > MAX_PAYLOAD_SIZE {
775            return Err(P2PStreamError::MessageTooBig {
776                message_size: item.len(),
777                max_size: MAX_PAYLOAD_SIZE,
778            })
779        }
780
781        if item.is_empty() {
782            // empty messages are not allowed
783            return Err(P2PStreamError::EmptyProtocolMessage)
784        }
785
786        // ensure we have free capacity
787        if !self.has_outgoing_capacity() {
788            return Err(P2PStreamError::SendBufferFull)
789        }
790
791        let this = self.project();
792
793        // all messages sent in this stream are subprotocol messages, so we need to switch the
794        // message id based on the offset
795        let compressed = compress_frame(
796            this.encoder,
797            this.compress_scratch,
798            item[0] + MAX_RESERVED_MESSAGE_ID + 1,
799            &item[1..],
800        )
801        .map_err(|err| {
802            debug!(
803                %err,
804                msg=%hex::encode(&item[1..]),
805                "error compressing p2p message"
806            );
807            err
808        })?;
809        this.outgoing_messages.push_back(compressed);
810
811        Ok(())
812    }
813
814    /// Returns `Poll::Ready(Ok(()))` when no buffered items remain.
815    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
816        ready!(self.as_mut().poll_drain_outgoing(cx))?;
817
818        let mut this = self.project();
819
820        if *this.needs_flush {
821            ready!(this.inner.as_mut().poll_flush(cx))?;
822            *this.needs_flush = false;
823        }
824        *this.needs_control_flush = false;
825
826        Poll::Ready(Ok(()))
827    }
828
829    fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
830        ready!(self.as_mut().poll_flush(cx))?;
831        ready!(self.project().inner.poll_close(cx))?;
832
833        Poll::Ready(Ok(()))
834    }
835}
836
837/// This represents only the reserved `p2p` subprotocol messages.
838#[derive(Debug, Clone, PartialEq, Eq)]
839#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
840#[cfg_attr(any(test, feature = "arbitrary"), derive(arbitrary::Arbitrary))]
841#[add_arbitrary_tests(rlp)]
842pub enum P2PMessage {
843    /// The first packet sent over the connection, and sent once by both sides.
844    Hello(HelloMessage),
845
846    /// Inform the peer that a disconnection is imminent; if received, a peer should disconnect
847    /// immediately.
848    Disconnect(DisconnectReason),
849
850    /// Requests an immediate reply of [`P2PMessage::Pong`] from the peer.
851    Ping,
852
853    /// Reply to the peer's [`P2PMessage::Ping`] packet.
854    Pong,
855}
856
857impl P2PMessage {
858    /// Gets the [`P2PMessageID`] for the given message.
859    pub const fn message_id(&self) -> P2PMessageID {
860        match self {
861            Self::Hello(_) => P2PMessageID::Hello,
862            Self::Disconnect(_) => P2PMessageID::Disconnect,
863            Self::Ping => P2PMessageID::Ping,
864            Self::Pong => P2PMessageID::Pong,
865        }
866    }
867}
868
869impl Encodable for P2PMessage {
870    /// The [`Encodable`] implementation for [`P2PMessage::Ping`] and [`P2PMessage::Pong`] encodes
871    /// the message as RLP, and prepends a snappy header to the RLP bytes for all variants except
872    /// the [`P2PMessage::Hello`] variant, because the hello message is never compressed in the
873    /// `p2p` subprotocol.
874    fn encode(&self, out: &mut dyn BufMut) {
875        (self.message_id() as u8).encode(out);
876        match self {
877            Self::Hello(msg) => msg.encode(out),
878            Self::Disconnect(msg) => msg.encode(out),
879            Self::Ping => {
880                // Ping payload is _always_ snappy encoded
881                out.put_slice(SNAPPY_EMPTY_LIST_PAYLOAD);
882            }
883            Self::Pong => {
884                // Pong payload is _always_ snappy encoded
885                out.put_slice(SNAPPY_EMPTY_LIST_PAYLOAD);
886            }
887        }
888    }
889
890    fn length(&self) -> usize {
891        let payload_len = match self {
892            Self::Hello(msg) => msg.length(),
893            Self::Disconnect(msg) => msg.length(),
894            // snappy encoded empty RLP list payload
895            Self::Ping | Self::Pong => SNAPPY_EMPTY_LIST_PAYLOAD.len(),
896        };
897        payload_len + 1 // (1 for length of p2p message id)
898    }
899}
900
901impl Decodable for P2PMessage {
902    /// The [`Decodable`] implementation for [`P2PMessage`] assumes that each of the message
903    /// variants are snappy compressed, except for the [`P2PMessage::Hello`] variant since the
904    /// hello message is never compressed in the `p2p` subprotocol.
905    ///
906    /// The [`Decodable`] implementation for [`P2PMessage::Ping`] and [`P2PMessage::Pong`] expects
907    /// a snappy encoded payload, see [`Encodable`] implementation.
908    fn decode(buf: &mut &[u8]) -> alloy_rlp::Result<Self> {
909        /// Removes the snappy prefix from the Ping/Pong buffer
910        fn advance_snappy_ping_pong_payload(buf: &mut &[u8]) -> alloy_rlp::Result<()> {
911            if buf.len() < 3 {
912                return Err(RlpError::InputTooShort)
913            }
914            if buf[..3] != [0x01, 0x00, EMPTY_LIST_CODE] {
915                return Err(RlpError::Custom("expected snappy payload"))
916            }
917            buf.advance(3);
918            Ok(())
919        }
920
921        let message_id = u8::decode(&mut &buf[..])?;
922        let id = P2PMessageID::try_from(message_id)
923            .or(Err(RlpError::Custom("unknown p2p message id")))?;
924        buf.advance(1);
925        match id {
926            P2PMessageID::Hello => Ok(Self::Hello(HelloMessage::decode(buf)?)),
927            P2PMessageID::Disconnect => Ok(Self::Disconnect(DisconnectReason::decode(buf)?)),
928            P2PMessageID::Ping => {
929                advance_snappy_ping_pong_payload(buf)?;
930                Ok(Self::Ping)
931            }
932            P2PMessageID::Pong => {
933                advance_snappy_ping_pong_payload(buf)?;
934                Ok(Self::Pong)
935            }
936        }
937    }
938}
939
940/// Message IDs for `p2p` subprotocol messages.
941#[derive(Debug, Copy, Clone, Eq, PartialEq)]
942pub enum P2PMessageID {
943    /// Message ID for the [`P2PMessage::Hello`] message.
944    Hello = 0x00,
945
946    /// Message ID for the [`P2PMessage::Disconnect`] message.
947    Disconnect = 0x01,
948
949    /// Message ID for the [`P2PMessage::Ping`] message.
950    Ping = 0x02,
951
952    /// Message ID for the [`P2PMessage::Pong`] message.
953    Pong = 0x03,
954}
955
956impl From<P2PMessage> for P2PMessageID {
957    fn from(msg: P2PMessage) -> Self {
958        match msg {
959            P2PMessage::Hello(_) => Self::Hello,
960            P2PMessage::Disconnect(_) => Self::Disconnect,
961            P2PMessage::Ping => Self::Ping,
962            P2PMessage::Pong => Self::Pong,
963        }
964    }
965}
966
967impl TryFrom<u8> for P2PMessageID {
968    type Error = P2PStreamError;
969
970    fn try_from(id: u8) -> Result<Self, Self::Error> {
971        match id {
972            0x00 => Ok(Self::Hello),
973            0x01 => Ok(Self::Disconnect),
974            0x02 => Ok(Self::Ping),
975            0x03 => Ok(Self::Pong),
976            _ => Err(P2PStreamError::UnknownReservedMessageId(id)),
977        }
978    }
979}
980
981/// Snappy-compresses an id-prefixed `p2p` message payload into a frame carrying the given wire
982/// message id.
983///
984/// Frames whose worst-case compressed size fits within [`MAX_COMPRESS_SCRATCH_SIZE`] are
985/// compressed through the reusable `scratch` buffer and copied out at their exact size; larger
986/// frames use a one-off allocation, see [`MAX_COMPRESS_SCRATCH_SIZE`].
987fn compress_frame(
988    encoder: &mut snap::raw::Encoder,
989    scratch: &mut Vec<u8>,
990    wire_id: u8,
991    payload: &[u8],
992) -> Result<Bytes, snap::Error> {
993    let needed = 1 + snap::raw::max_compress_len(payload.len());
994
995    if needed > MAX_COMPRESS_SCRATCH_SIZE {
996        let mut compressed = vec![0u8; needed];
997        let compressed_size = encoder.compress(payload, &mut compressed[1..])?;
998        compressed[0] = wire_id;
999        compressed.truncate(compressed_size + 1);
1000        return Ok(compressed.into())
1001    }
1002
1003    if scratch.len() < needed {
1004        scratch.resize(needed, 0);
1005    }
1006    let compressed_size = encoder.compress(payload, &mut scratch[1..])?;
1007    scratch[0] = wire_id;
1008    Ok(Bytes::copy_from_slice(&scratch[..compressed_size + 1]))
1009}
1010
1011#[cfg(test)]
1012mod tests {
1013    use super::*;
1014    use crate::{
1015        capability::SharedCapability, protocol::Protocol, test_utils::eth_hello, Capability,
1016        EthVersion, ProtocolVersion,
1017    };
1018    use futures::task::noop_waker_ref;
1019    use tokio::net::{TcpListener, TcpStream};
1020    use tokio_util::codec::Decoder;
1021
1022    /// A sink that records started frames and counts flushes, to observe batching behavior.
1023    #[derive(Default)]
1024    struct FlushCountingTransport {
1025        incoming: VecDeque<io::Result<BytesMut>>,
1026        sent: Vec<Bytes>,
1027        flushes: usize,
1028    }
1029
1030    #[derive(Default)]
1031    struct InboundTransport {
1032        incoming: VecDeque<BytesMut>,
1033    }
1034
1035    impl Stream for InboundTransport {
1036        type Item = io::Result<BytesMut>;
1037
1038        fn poll_next(mut self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1039            Poll::Ready(self.incoming.pop_front().map(Ok))
1040        }
1041    }
1042
1043    impl Sink<Bytes> for InboundTransport {
1044        type Error = io::Error;
1045
1046        fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1047            Poll::Ready(Ok(()))
1048        }
1049
1050        fn start_send(self: Pin<&mut Self>, _item: Bytes) -> Result<(), Self::Error> {
1051            Ok(())
1052        }
1053
1054        fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1055            Poll::Ready(Ok(()))
1056        }
1057
1058        fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1059            Poll::Ready(Ok(()))
1060        }
1061    }
1062
1063    impl Stream for FlushCountingTransport {
1064        type Item = io::Result<BytesMut>;
1065
1066        fn poll_next(mut self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1067            match self.incoming.pop_front() {
1068                Some(item) => Poll::Ready(Some(item)),
1069                None => Poll::Pending,
1070            }
1071        }
1072    }
1073
1074    impl Sink<Bytes> for FlushCountingTransport {
1075        type Error = io::Error;
1076
1077        fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1078            Poll::Ready(Ok(()))
1079        }
1080
1081        fn start_send(mut self: Pin<&mut Self>, item: Bytes) -> Result<(), Self::Error> {
1082            self.sent.push(item);
1083            Ok(())
1084        }
1085
1086        fn poll_flush(
1087            mut self: Pin<&mut Self>,
1088            _: &mut Context<'_>,
1089        ) -> Poll<Result<(), Self::Error>> {
1090            self.flushes += 1;
1091            Poll::Ready(Ok(()))
1092        }
1093
1094        fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1095            Poll::Ready(Ok(()))
1096        }
1097    }
1098
1099    fn eth_shared_capabilities() -> SharedCapabilities {
1100        SharedCapabilities::try_new(
1101            vec![EthVersion::Eth68.into()],
1102            vec![Capability::eth(EthVersion::Eth68)],
1103        )
1104        .unwrap()
1105    }
1106
1107    fn stream_with_incoming(frame: BytesMut) -> P2PStream<FlushCountingTransport> {
1108        let mut transport = FlushCountingTransport::default();
1109        transport.incoming.push_back(Ok(frame));
1110        P2PStream::new(transport, eth_shared_capabilities())
1111    }
1112
1113    fn stream_with_incoming_pings(count: usize) -> P2PStream<FlushCountingTransport> {
1114        let mut transport = FlushCountingTransport::default();
1115        transport.incoming.extend((0..count).map(|_| Ok(BytesMut::from(SNAPPY_PING_MESSAGE))));
1116        P2PStream::new(transport, eth_shared_capabilities())
1117    }
1118
1119    fn compressed_p2p_message(message_id: P2PMessageID, payload: &[u8]) -> BytesMut {
1120        let mut encoder = snap::raw::Encoder::new();
1121        let mut scratch = Vec::new();
1122        let message =
1123            compress_frame(&mut encoder, &mut scratch, message_id as u8, payload).unwrap();
1124        BytesMut::from(message.as_ref())
1125    }
1126
1127    #[tokio::test]
1128    async fn rejects_ping_pong_with_oversized_payload_before_decompression() {
1129        const SIXTEEN_MIB_SNAPPY_HEADER: [u8; 4] = [0x80, 0x80, 0x80, 0x08];
1130
1131        for message_id in [P2PMessageID::Ping, P2PMessageID::Pong] {
1132            let frame = BytesMut::from(
1133                [
1134                    message_id as u8,
1135                    SIXTEEN_MIB_SNAPPY_HEADER[0],
1136                    SIXTEEN_MIB_SNAPPY_HEADER[1],
1137                    SIXTEEN_MIB_SNAPPY_HEADER[2],
1138                    SIXTEEN_MIB_SNAPPY_HEADER[3],
1139                ]
1140                .as_slice(),
1141            );
1142            assert_eq!(snap::raw::decompress_len(&frame[1..]).unwrap(), MAX_PAYLOAD_SIZE);
1143
1144            let mut stream = stream_with_incoming(frame);
1145            let waker = noop_waker_ref();
1146            let mut cx = Context::from_waker(waker);
1147
1148            match Pin::new(&mut stream).poll_next(&mut cx) {
1149                Poll::Ready(Some(Err(P2PStreamError::InvalidPingPongPayload(id)))) => {
1150                    assert_eq!(id, message_id as u8)
1151                }
1152                result => panic!("unexpected poll result: {result:?}"),
1153            }
1154            assert!(stream.outgoing_messages.is_empty());
1155        }
1156    }
1157
1158    #[tokio::test]
1159    async fn rejects_ping_pong_with_non_list_payload() {
1160        for message_id in [P2PMessageID::Ping, P2PMessageID::Pong] {
1161            let frame = compressed_p2p_message(message_id, &[alloy_rlp::EMPTY_STRING_CODE]);
1162            let mut stream = stream_with_incoming(frame);
1163            let waker = noop_waker_ref();
1164            let mut cx = Context::from_waker(waker);
1165
1166            assert!(matches!(
1167                Pin::new(&mut stream).poll_next(&mut cx),
1168                Poll::Ready(Some(Err(P2PStreamError::InvalidPingPongPayload(id))))
1169                    if id == message_id as u8
1170            ));
1171            assert!(stream.outgoing_messages.is_empty());
1172        }
1173    }
1174
1175    #[tokio::test]
1176    async fn accepts_ping_with_empty_list_payload() {
1177        let mut stream = stream_with_incoming(BytesMut::from(SNAPPY_PING_MESSAGE));
1178        let waker = noop_waker_ref();
1179        let mut cx = Context::from_waker(waker);
1180
1181        assert!(Pin::new(&mut stream).poll_next(&mut cx).is_pending());
1182        assert_eq!(stream.outgoing_messages.len(), 1);
1183        assert_eq!(stream.outgoing_messages.front().unwrap().as_ref(), SNAPPY_PONG_MESSAGE);
1184    }
1185
1186    #[test]
1187    fn ping_token_bucket_limits_bursts_and_refills() {
1188        let now = Instant::now();
1189        let mut bucket = PingTokenBucket::new(now);
1190
1191        for _ in 0..PING_TOKEN_BUCKET_CAPACITY {
1192            assert!(bucket.try_take(now));
1193        }
1194        assert!(!bucket.try_take(now));
1195
1196        let one_refill = now + Duration::from_secs(1);
1197        assert!(bucket.try_take(one_refill));
1198        assert!(!bucket.try_take(one_refill));
1199
1200        // The extra half second checks that refill time accumulated while the bucket is full is
1201        // discarded. Otherwise, the next token could arrive less than one second after the burst.
1202        let full_refill = one_refill +
1203            Duration::from_secs(u64::from(PING_TOKEN_BUCKET_CAPACITY)) +
1204            Duration::from_millis(500);
1205        for _ in 0..PING_TOKEN_BUCKET_CAPACITY {
1206            assert!(bucket.try_take(full_refill));
1207        }
1208        assert!(!bucket.try_take(full_refill));
1209        assert!(!bucket.try_take(full_refill + Duration::from_millis(999)));
1210        assert!(bucket.try_take(full_refill + Duration::from_secs(1)));
1211    }
1212
1213    #[tokio::test]
1214    async fn accepts_ping_burst_at_token_bucket_capacity() {
1215        let mut stream = stream_with_incoming_pings(usize::from(PING_TOKEN_BUCKET_CAPACITY));
1216        let waker = noop_waker_ref();
1217        let mut cx = Context::from_waker(waker);
1218
1219        assert!(Pin::new(&mut stream).poll_next(&mut cx).is_pending());
1220        assert_eq!(stream.outgoing_messages.len(), usize::from(PING_TOKEN_BUCKET_CAPACITY));
1221    }
1222
1223    #[tokio::test]
1224    async fn rejects_ping_burst_over_token_bucket_capacity() {
1225        let mut stream = stream_with_incoming_pings(usize::from(PING_TOKEN_BUCKET_CAPACITY) + 1);
1226        let waker = noop_waker_ref();
1227        let mut cx = Context::from_waker(waker);
1228
1229        assert!(matches!(
1230            Pin::new(&mut stream).poll_next(&mut cx),
1231            Poll::Ready(Some(Err(P2PStreamError::TooManyPings)))
1232        ));
1233        assert_eq!(stream.outgoing_messages.len(), usize::from(PING_TOKEN_BUCKET_CAPACITY));
1234    }
1235
1236    #[tokio::test]
1237    async fn rejects_subprotocol_frame_before_decompression_when_declared_size_exceeds_limit() {
1238        let cap = Capability::new_static("test", 1);
1239        let shared_capabilities =
1240            SharedCapabilities::try_new(vec![Protocol::new(cap.clone(), 1)], vec![cap.clone()])
1241                .unwrap();
1242        let shared_capability = shared_capabilities.find(&cap).unwrap().clone();
1243        let wire_id = shared_capability.message_id_offset();
1244        let mut encoder = snap::raw::Encoder::new();
1245        let mut scratch = Vec::new();
1246        let accepted = compress_frame(&mut encoder, &mut scratch, wire_id, &[0; 3]).unwrap();
1247        let oversized = compress_frame(&mut encoder, &mut scratch, wire_id, &[0; 4]).unwrap();
1248        let transport = InboundTransport {
1249            incoming: [accepted, oversized]
1250                .into_iter()
1251                .map(|frame| BytesMut::from(frame.as_ref()))
1252                .collect(),
1253        };
1254        let mut stream = P2PStream::new(transport, shared_capabilities);
1255        stream.set_protocol_ingress_limits(&shared_capability, ProtocolIngressLimits::new(4));
1256
1257        let frame = stream.next().await.unwrap().unwrap();
1258        assert_eq!(frame.len(), 4);
1259
1260        assert!(matches!(
1261            stream.next().await.unwrap().unwrap_err(),
1262            P2PStreamError::SubprotocolMessageTooBig {
1263                capability,
1264                message_size: 5,
1265                max_size: 4,
1266            } if capability == cap
1267        ));
1268    }
1269
1270    #[tokio::test]
1271    async fn zero_message_protocol_does_not_replace_neighboring_frame_limit() {
1272        let zero = Capability::new_static("aaa", 1);
1273        let limited = Capability::new_static("bbb", 1);
1274        let shared_capabilities = SharedCapabilities::try_new(
1275            vec![Protocol::new(zero.clone(), 0), Protocol::new(limited.clone(), 1)],
1276            vec![zero.clone(), limited.clone()],
1277        )
1278        .unwrap();
1279        let zero_shared = shared_capabilities.find(&zero).unwrap().clone();
1280        let limited_shared = shared_capabilities.find(&limited).unwrap().clone();
1281        assert_eq!(zero_shared.message_id_offset(), limited_shared.message_id_offset());
1282
1283        let mut encoder = snap::raw::Encoder::new();
1284        let mut scratch = Vec::new();
1285        let oversized =
1286            compress_frame(&mut encoder, &mut scratch, limited_shared.message_id_offset(), &[0; 4])
1287                .unwrap();
1288        let transport =
1289            InboundTransport { incoming: VecDeque::from([BytesMut::from(oversized.as_ref())]) };
1290        let mut stream = P2PStream::new(transport, shared_capabilities);
1291        stream.set_protocol_ingress_limits(&limited_shared, ProtocolIngressLimits::new(4));
1292        stream.set_protocol_ingress_limits(&zero_shared, ProtocolIngressLimits::new(1));
1293
1294        assert!(matches!(
1295            stream.next().await.unwrap().unwrap_err(),
1296            P2PStreamError::SubprotocolMessageTooBig {
1297                capability,
1298                message_size: 5,
1299                max_size: 4,
1300            } if capability == limited
1301        ));
1302    }
1303
1304    #[tokio::test]
1305    async fn poll_ready_drains_full_subprotocol_queue_without_flushing_inner() {
1306        let mut stream =
1307            P2PStream::new(FlushCountingTransport::default(), eth_shared_capabilities());
1308        stream.set_outgoing_message_buffer_capacity(1);
1309        Pin::new(&mut stream).start_send(Bytes::from_static(&[0x00, EMPTY_LIST_CODE])).unwrap();
1310
1311        let waker = noop_waker_ref();
1312        let mut cx = Context::from_waker(waker);
1313        assert!(Pin::new(&mut stream).poll_ready(&mut cx).is_ready());
1314
1315        // the full queue was drained into the inner sink to make room, but not flushed
1316        assert_eq!(stream.inner().sent.len(), 1);
1317        assert_eq!(stream.inner().flushes, 0);
1318
1319        // the caller-driven flush pushes the batch out with a single inner flush
1320        assert!(Pin::new(&mut stream).poll_flush(&mut cx).is_ready());
1321        assert_eq!(stream.inner().flushes, 1);
1322
1323        // flushing again is a no-op on the inner sink
1324        assert!(Pin::new(&mut stream).poll_flush(&mut cx).is_ready());
1325        assert_eq!(stream.inner().flushes, 1);
1326    }
1327
1328    #[tokio::test]
1329    async fn poll_ready_flushes_queued_control_messages() {
1330        let mut stream =
1331            P2PStream::new(FlushCountingTransport::default(), eth_shared_capabilities());
1332        stream.send_ping();
1333
1334        let waker = noop_waker_ref();
1335        let mut cx = Context::from_waker(waker);
1336        assert!(Pin::new(&mut stream).poll_ready(&mut cx).is_ready());
1337
1338        // control messages must not wait for a caller-driven flush
1339        assert_eq!(stream.inner().sent.len(), 1);
1340        assert_eq!(stream.inner().flushes, 1);
1341    }
1342
1343    #[tokio::test]
1344    async fn test_can_disconnect() {
1345        reth_tracing::init_test_tracing();
1346        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1347        let local_addr = listener.local_addr().unwrap();
1348
1349        let expected_disconnect = DisconnectReason::UselessPeer;
1350
1351        let handle = tokio::spawn(async move {
1352            // roughly based off of the design of tokio::net::TcpListener
1353            let (incoming, _) = listener.accept().await.unwrap();
1354            let stream = crate::PassthroughCodec::default().framed(incoming);
1355
1356            let (server_hello, _) = eth_hello();
1357
1358            let (mut p2p_stream, _) =
1359                UnauthedP2PStream::new(stream).handshake(server_hello).await.unwrap();
1360
1361            p2p_stream.disconnect(expected_disconnect).await.unwrap();
1362        });
1363
1364        let outgoing = TcpStream::connect(local_addr).await.unwrap();
1365        let sink = crate::PassthroughCodec::default().framed(outgoing);
1366
1367        let (client_hello, _) = eth_hello();
1368
1369        let (mut p2p_stream, _) =
1370            UnauthedP2PStream::new(sink).handshake(client_hello).await.unwrap();
1371
1372        let err = p2p_stream.next().await.unwrap().unwrap_err();
1373        match err {
1374            P2PStreamError::Disconnected(reason) => assert_eq!(reason, expected_disconnect),
1375            e => panic!("unexpected err: {e}"),
1376        }
1377
1378        handle.await.unwrap();
1379    }
1380
1381    #[tokio::test]
1382    async fn test_can_disconnect_weird_disconnect_encoding() {
1383        reth_tracing::init_test_tracing();
1384        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1385        let local_addr = listener.local_addr().unwrap();
1386
1387        let expected_disconnect = DisconnectReason::SubprotocolSpecific;
1388
1389        let handle = tokio::spawn(async move {
1390            // roughly based off of the design of tokio::net::TcpListener
1391            let (incoming, _) = listener.accept().await.unwrap();
1392            let stream = crate::PassthroughCodec::default().framed(incoming);
1393
1394            let (server_hello, _) = eth_hello();
1395
1396            let (mut p2p_stream, _) =
1397                UnauthedP2PStream::new(stream).handshake(server_hello).await.unwrap();
1398
1399            // Unrolled `disconnect` method, without compression
1400            p2p_stream.outgoing_messages.clear();
1401
1402            p2p_stream.outgoing_messages.push_back(Bytes::from(alloy_rlp::encode(
1403                P2PMessage::Disconnect(DisconnectReason::SubprotocolSpecific),
1404            )));
1405            p2p_stream.disconnecting = true;
1406            p2p_stream.close().await.unwrap();
1407        });
1408
1409        let outgoing = TcpStream::connect(local_addr).await.unwrap();
1410        let sink = crate::PassthroughCodec::default().framed(outgoing);
1411
1412        let (client_hello, _) = eth_hello();
1413
1414        let (mut p2p_stream, _) =
1415            UnauthedP2PStream::new(sink).handshake(client_hello).await.unwrap();
1416
1417        let err = p2p_stream.next().await.unwrap().unwrap_err();
1418        match err {
1419            P2PStreamError::Disconnected(reason) => assert_eq!(reason, expected_disconnect),
1420            e => panic!("unexpected err: {e}"),
1421        }
1422
1423        handle.await.unwrap();
1424    }
1425
1426    #[tokio::test]
1427    async fn test_handshake_passthrough() {
1428        // create a p2p stream and server, then confirm that the two are authed
1429        // create tcpstream
1430        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1431        let local_addr = listener.local_addr().unwrap();
1432
1433        let handle = tokio::spawn(async move {
1434            // roughly based off of the design of tokio::net::TcpListener
1435            let (incoming, _) = listener.accept().await.unwrap();
1436            let stream = crate::PassthroughCodec::default().framed(incoming);
1437
1438            let (server_hello, _) = eth_hello();
1439
1440            let unauthed_stream = UnauthedP2PStream::new(stream);
1441            let (p2p_stream, _) = unauthed_stream.handshake(server_hello).await.unwrap();
1442
1443            // ensure that the two share a single capability, eth67
1444            assert_eq!(
1445                *p2p_stream.shared_capabilities.iter_caps().next().unwrap(),
1446                SharedCapability::Eth {
1447                    version: EthVersion::Eth67,
1448                    offset: MAX_RESERVED_MESSAGE_ID + 1
1449                }
1450            );
1451        });
1452
1453        let outgoing = TcpStream::connect(local_addr).await.unwrap();
1454        let sink = crate::PassthroughCodec::default().framed(outgoing);
1455
1456        let (client_hello, _) = eth_hello();
1457
1458        let unauthed_stream = UnauthedP2PStream::new(sink);
1459        let (p2p_stream, _) = unauthed_stream.handshake(client_hello).await.unwrap();
1460
1461        // ensure that the two share a single capability, eth67
1462        assert_eq!(
1463            *p2p_stream.shared_capabilities.iter_caps().next().unwrap(),
1464            SharedCapability::Eth {
1465                version: EthVersion::Eth67,
1466                offset: MAX_RESERVED_MESSAGE_ID + 1
1467            }
1468        );
1469
1470        // make sure the server receives the message and asserts before ending the test
1471        handle.await.unwrap();
1472    }
1473
1474    #[tokio::test]
1475    async fn test_handshake_disconnect() {
1476        // create a p2p stream and server, then confirm that the two are authed
1477        // create tcpstream
1478        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1479        let local_addr = listener.local_addr().unwrap();
1480
1481        let handle = tokio::spawn(async move {
1482            // roughly based off of the design of tokio::net::TcpListener
1483            let (incoming, _) = listener.accept().await.unwrap();
1484            let stream = crate::PassthroughCodec::default().framed(incoming);
1485
1486            let (server_hello, _) = eth_hello();
1487
1488            let unauthed_stream = UnauthedP2PStream::new(stream);
1489            match unauthed_stream.handshake(server_hello.clone()).await {
1490                Ok((_, hello)) => {
1491                    panic!("expected handshake to fail, instead got a successful Hello: {hello:?}")
1492                }
1493                Err(P2PStreamError::MismatchedProtocolVersion(GotExpected { got, expected })) => {
1494                    assert_ne!(expected, got);
1495                    assert_eq!(expected, server_hello.protocol_version);
1496                }
1497                Err(other_err) => {
1498                    panic!("expected mismatched protocol version error, got {other_err:?}")
1499                }
1500            }
1501        });
1502
1503        let outgoing = TcpStream::connect(local_addr).await.unwrap();
1504        let sink = crate::PassthroughCodec::default().framed(outgoing);
1505
1506        let (mut client_hello, _) = eth_hello();
1507
1508        // modify the hello to include an incompatible p2p protocol version
1509        client_hello.protocol_version = ProtocolVersion::V4;
1510
1511        let unauthed_stream = UnauthedP2PStream::new(sink);
1512        match unauthed_stream.handshake(client_hello.clone()).await {
1513            Ok((_, hello)) => {
1514                panic!("expected handshake to fail, instead got a successful Hello: {hello:?}")
1515            }
1516            Err(P2PStreamError::MismatchedProtocolVersion(GotExpected { got, expected })) => {
1517                assert_ne!(expected, got);
1518                assert_eq!(expected, client_hello.protocol_version);
1519            }
1520            Err(other_err) => {
1521                panic!("expected mismatched protocol version error, got {other_err:?}")
1522            }
1523        }
1524
1525        // make sure the server receives the message and asserts before ending the test
1526        handle.await.unwrap();
1527    }
1528
1529    #[test]
1530    fn snappy_ping_pong_consts_match_rlp_encoding() {
1531        assert_eq!(alloy_rlp::encode(P2PMessage::Ping).as_slice(), SNAPPY_PING_MESSAGE);
1532        assert_eq!(alloy_rlp::encode(P2PMessage::Pong).as_slice(), SNAPPY_PONG_MESSAGE);
1533    }
1534
1535    #[test]
1536    fn snappy_decode_encode_ping() {
1537        let snappy_ping = b"\x02\x01\0\xc0";
1538        let ping = P2PMessage::decode(&mut &snappy_ping[..]).unwrap();
1539        assert!(matches!(ping, P2PMessage::Ping));
1540        assert_eq!(alloy_rlp::encode(ping), &snappy_ping[..]);
1541    }
1542
1543    #[test]
1544    fn snappy_decode_encode_pong() {
1545        let snappy_pong = b"\x03\x01\0\xc0";
1546        let pong = P2PMessage::decode(&mut &snappy_pong[..]).unwrap();
1547        assert!(matches!(pong, P2PMessage::Pong));
1548        assert_eq!(alloy_rlp::encode(pong), &snappy_pong[..]);
1549    }
1550}