Skip to main content

mtorrent_core/pwp/
channels.rs

1use super::MAX_BLOCK_SIZE;
2use super::handshake::*;
3use super::message::*;
4use crate::pe;
5use bytes::BufMut;
6use futures_channel::mpsc;
7use futures_util::future::LocalBoxFuture;
8use futures_util::{FutureExt, SinkExt, StreamExt};
9use local_async_utils::prelude::*;
10use mtorrent_utils::split_stream::SplitStream;
11use std::future::Future;
12use std::io;
13use std::mem::MaybeUninit;
14use std::net::SocketAddr;
15use std::sync::Arc;
16use std::time::Duration;
17use thiserror::Error;
18use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
19use tokio::time::{sleep, timeout};
20use tokio::{select, task, try_join};
21
22/// Error type for receiving messages from,
23/// or sending them to, a [`PeerChannel`].
24#[derive(Debug, Error, Clone, Copy, PartialEq, Eq)]
25pub enum ChannelError {
26    #[error("timeout")]
27    Timeout,
28    #[error("connection closed")]
29    ConnectionClosed,
30}
31
32struct PeerInfo {
33    handshake_info: Handshake,
34    remote_addr: SocketAddr,
35    encrypted: bool,
36}
37
38/// Channel for communicating with a single peer.
39pub struct PeerChannel<Q> {
40    peer_info: Arc<PeerInfo>,
41    inner: Q,
42}
43
44impl<Q> PeerChannel<Q> {
45    /// Address of the peer.
46    pub fn remote_ip(&self) -> &SocketAddr {
47        &self.peer_info.remote_addr
48    }
49    /// Handshake received from the peer upon connecting.
50    pub fn remote_info(&self) -> &Handshake {
51        &self.peer_info.handshake_info
52    }
53    /// Whether the connection to the peer is encrypted.
54    pub fn is_encrypted(&self) -> bool {
55        self.peer_info.encrypted
56    }
57}
58
59impl<Q: Clone> Clone for PeerChannel<Q> {
60    fn clone(&self) -> Self {
61        Self {
62            peer_info: self.peer_info.clone(),
63            inner: self.inner.clone(),
64        }
65    }
66}
67
68type RxChannel<Msg> = PeerChannel<mpsc::Receiver<Msg>>;
69type TxChannel<Msg> = PeerChannel<mpsc::Sender<Option<Msg>>>;
70
71impl<Msg> RxChannel<Msg> {
72    /// Wait for the next message received from the peer.
73    /// Returns [`ChannelError::ConnectionClosed`] if actor has exited.
74    pub async fn receive_message(&mut self) -> Result<Msg, ChannelError> {
75        self.inner.next().await.ok_or(ChannelError::ConnectionClosed)
76    }
77
78    /// Wait up to `deadline` for the next message received from the peer.
79    /// Returns [`ChannelError::ConnectionClosed`] if actor has exited,
80    /// or [`ChannelError::Timeout`] if the deadline has elapsed.
81    pub async fn receive_message_timed(&mut self, deadline: Duration) -> Result<Msg, ChannelError> {
82        timeout(deadline, self.receive_message()).await.or(Err(ChannelError::Timeout))?
83    }
84}
85
86impl<Msg> TxChannel<Msg> {
87    /// Send a single message to the peer. The async call will only return once
88    /// the message has been successfully written to the socket.
89    /// Returns [`ChannelError::ConnectionClosed`] if actor has exited.
90    pub async fn send_message(&mut self, msg: Msg) -> Result<(), ChannelError> {
91        self.inner.send(Some(msg)).await?;
92        self.inner.send(None).await?;
93        Ok(())
94    }
95
96    /// Send a single message to the peer. The async call will return either once
97    /// the message has been successfully written to the socket, or when `deadline` has elapsed.
98    /// Returns [`ChannelError::ConnectionClosed`] if actor has exited,
99    /// or [`ChannelError::Timeout`] if the deadline has elapsed.
100    pub async fn send_message_timed(
101        &mut self,
102        msg: Msg,
103        deadline: Duration,
104    ) -> Result<(), ChannelError> {
105        timeout(deadline, self.send_message(msg)).await.or(Err(ChannelError::Timeout))?
106    }
107}
108
109/// Channel for sending messages related to us downloading data from the peer.
110pub type DownloadTxChannel = TxChannel<DownloaderMessage>;
111/// Channel for receiving messages related to the peer uploading data to us.
112pub type DownloadRxChannel = RxChannel<UploaderMessage>;
113/// Channels for downloading data from a single peer.
114pub struct DownloadChannels(pub DownloadTxChannel, pub DownloadRxChannel);
115
116/// Channel for sending messages related to us uploading data to the peer.
117pub type UploadTxChannel = TxChannel<UploaderMessage>;
118/// Channel for receiving messages related to the peer downloading data from us.
119pub type UploadRxChannel = RxChannel<DownloaderMessage>;
120/// Channels for uploading data to a single peer.
121pub struct UploadChannels(pub UploadTxChannel, pub UploadRxChannel);
122
123/// Channel for sending extended protocol messages to the peer.
124pub type ExtendedTxChannel = TxChannel<(ExtendedMessage, u8)>;
125/// Channel for receiving extended protocol messages from the peer.
126pub type ExtendedRxChannel = RxChannel<ExtendedMessage>;
127/// Channels for exchanging extended protocol messages with a single peer.
128pub struct ExtendedChannels(pub ExtendedTxChannel, pub ExtendedRxChannel);
129
130// ------
131
132const HANDSHAKE_TIMEOUT: Duration = sec!(10);
133
134/// Perform handshake on an inbound connection from a peer, and set up [`PeerChannel`]s.
135/// Must be called inside [`tokio::runtime::LocalRuntime`](https://docs.rs/tokio/latest/tokio/runtime/struct.LocalRuntime.html).
136pub async fn channels_for_inbound_connection<S>(
137    local_peer_id: &[u8; 20],
138    info_hash: &[u8; 20],
139    extension_protocol_enabled: bool,
140    remote_addr: SocketAddr,
141    socket: S,
142    mut crypto: Option<pe::Crypto>,
143) -> io::Result<(DownloadChannels, UploadChannels, Option<ExtendedChannels>)>
144where
145    S: AsyncRead + AsyncWrite + SplitStream + 'static,
146{
147    let local_handshake = Handshake {
148        peer_id: *local_peer_id,
149        info_hash: *info_hash,
150        reserved: reserved_bits(extension_protocol_enabled),
151    };
152    let (socket, remote_handshake) = timeout(
153        HANDSHAKE_TIMEOUT,
154        do_handshake_incoming(&remote_addr, socket, &local_handshake, crypto.as_mut()),
155    )
156    .await??;
157    let (download, upload, extensions, runner) =
158        setup_channels(socket, remote_addr, remote_handshake, extension_protocol_enabled, crypto);
159    task::spawn_local(runner);
160    Ok((download, upload, extensions))
161}
162
163/// Perform handshake on an outbound connection to a peer, and set up [`PeerChannel`]s.
164/// Must be called inside [`tokio::runtime::LocalRuntime`](https://docs.rs/tokio/latest/tokio/runtime/struct.LocalRuntime.html).
165pub async fn channels_for_outbound_connection<S>(
166    local_peer_id: &[u8; 20],
167    info_hash: &[u8; 20],
168    extension_protocol_enabled: bool,
169    remote_addr: SocketAddr,
170    socket: S,
171    remote_peer_id: Option<&[u8; 20]>,
172    mut crypto: Option<pe::Crypto>,
173) -> io::Result<(DownloadChannels, UploadChannels, Option<ExtendedChannels>)>
174where
175    S: AsyncRead + AsyncWrite + SplitStream + 'static,
176{
177    let local_handshake = Handshake {
178        peer_id: *local_peer_id,
179        info_hash: *info_hash,
180        reserved: reserved_bits(extension_protocol_enabled),
181    };
182    let (socket, remote_handshake) = timeout(
183        HANDSHAKE_TIMEOUT,
184        do_handshake_outgoing(
185            &remote_addr,
186            socket,
187            &local_handshake,
188            remote_peer_id,
189            crypto.as_mut(),
190        ),
191    )
192    .await??;
193    let (download, upload, extensions, runner) =
194        setup_channels(socket, remote_addr, remote_handshake, extension_protocol_enabled, crypto);
195    task::spawn_local(runner);
196    Ok((download, upload, extensions))
197}
198
199/// Set up [`PeerChannel`]s for a fake stream without performing a handshake.
200/// Must be called inside [`tokio::runtime::LocalRuntime`](https://docs.rs/tokio/latest/tokio/runtime/struct.LocalRuntime.html).
201#[cfg(feature = "mocks")]
202pub fn channels_from_mock<S>(
203    peer_addr: SocketAddr,
204    remote_handshake: Handshake,
205    extension_protocol_enabled: bool,
206    mock_socket: S,
207) -> (DownloadChannels, UploadChannels, Option<ExtendedChannels>)
208where
209    S: AsyncRead + AsyncWrite + Unpin + 'static,
210{
211    let (download, upload, extensions, runner) = setup_channels(
212        StreamHolder(mock_socket),
213        peer_addr,
214        remote_handshake,
215        extension_protocol_enabled,
216        None,
217    );
218    tokio::task::spawn_local(async move {
219        let _ = runner.await;
220    });
221    (download, upload, extensions)
222}
223
224// ------
225
226#[cfg(any(feature = "mocks", test))]
227struct StreamHolder<S>(S)
228where
229    S: AsyncRead + AsyncWrite + Unpin;
230
231#[cfg(any(feature = "mocks", test))]
232impl<S> SplitStream for StreamHolder<S>
233where
234    S: AsyncRead + AsyncWrite + Unpin + 'static,
235{
236    type Ingress<'i> = local_split::ReadHalf<&'i mut S>;
237    type Egress<'e> = local_split::WriteHalf<&'e mut S>;
238
239    fn split(&mut self) -> (Self::Ingress<'_>, Self::Egress<'_>) {
240        local_split::split(&mut self.0)
241    }
242}
243
244// ------
245
246fn setup_channels<S>(
247    stream: S,
248    remote_addr: SocketAddr,
249    remote_handshake: Handshake,
250    extended_protocol_enabled: bool,
251    crypto: Option<pe::Crypto>,
252) -> (
253    DownloadChannels,
254    UploadChannels,
255    Option<ExtendedChannels>,
256    impl Future<Output = io::Result<()>>,
257)
258where
259    S: SplitStream + 'static,
260{
261    const MAX_INCOMING_QUEUE: usize = 20;
262
263    let extended_protocol_supported = is_extension_protocol_enabled(&remote_handshake.reserved);
264
265    let info = Arc::new(PeerInfo {
266        handshake_info: remote_handshake,
267        remote_addr,
268        encrypted: crypto.is_some(),
269    });
270
271    let (local_uploader_msg_in, local_uploader_msg_out) =
272        mpsc::channel::<Option<UploaderMessage>>(0);
273    let (local_downloader_msg_in, local_downloader_msg_out) =
274        mpsc::channel::<Option<DownloaderMessage>>(0);
275
276    let (remote_uploader_msg_in, remote_uploader_msg_out) =
277        mpsc::channel::<UploaderMessage>(MAX_INCOMING_QUEUE);
278    let (remote_downloader_msg_in, remote_downloader_msg_out) =
279        mpsc::channel::<DownloaderMessage>(MAX_INCOMING_QUEUE);
280
281    let (local_extended_msg_out, remote_extended_msg_in, extended_channels) =
282        if extended_protocol_supported && extended_protocol_enabled {
283            let (local_extended_msg_in, local_extended_msg_out) =
284                mpsc::channel::<Option<(ExtendedMessage, u8)>>(0);
285            let (remote_extended_msg_in, remote_extended_msg_out) =
286                mpsc::channel::<ExtendedMessage>(MAX_INCOMING_QUEUE);
287
288            let extended_rx = ExtendedRxChannel {
289                inner: remote_extended_msg_out,
290                peer_info: info.clone(),
291            };
292            let extended_tx = ExtendedTxChannel {
293                inner: local_extended_msg_in,
294                peer_info: info.clone(),
295            };
296            (
297                Some(local_extended_msg_out),
298                Some(remote_extended_msg_in),
299                Some(ExtendedChannels(extended_tx, extended_rx)),
300            )
301        } else {
302            (None, None, None)
303        };
304
305    let receiver = IngressProcessor {
306        remote_ip: remote_addr,
307        ul_msg_sink: remote_uploader_msg_in,
308        dl_msg_sink: remote_downloader_msg_in,
309        ext_msg_sink: remote_extended_msg_in,
310    };
311    let sender = EgressProcessor {
312        remote_ip: remote_addr,
313        dl_msg_source: local_downloader_msg_out,
314        ul_msg_source: local_uploader_msg_out,
315        ext_msg_source: local_extended_msg_out,
316    };
317
318    let download_rx = DownloadRxChannel {
319        inner: remote_uploader_msg_out,
320        peer_info: info.clone(),
321    };
322    let download_tx = DownloadTxChannel {
323        inner: local_downloader_msg_in,
324        peer_info: info.clone(),
325    };
326
327    let upload_rx = UploadRxChannel {
328        inner: remote_downloader_msg_out,
329        peer_info: info.clone(),
330    };
331    let upload_tx = UploadTxChannel {
332        inner: local_uploader_msg_in,
333        peer_info: info.clone(),
334    };
335
336    (
337        DownloadChannels(download_tx, download_rx),
338        UploadChannels(upload_tx, upload_rx),
339        extended_channels,
340        make_io_task(stream, receiver, sender, crypto),
341    )
342}
343
344fn make_io_task<'s>(
345    mut stream: impl SplitStream + 's,
346    receiver: IngressProcessor,
347    sender: EgressProcessor,
348    crypto: Option<pe::Crypto>,
349) -> LocalBoxFuture<'s, io::Result<()>> {
350    async fn run_io(
351        ingress: IngressProcessor,
352        egress: EgressProcessor,
353        source: impl AsyncReadExt + Unpin,
354        sink: impl AsyncWriteExt + Unpin,
355    ) -> io::Result<()> {
356        let remote_addr = egress.remote_ip;
357        try_join!(biased; ingress.read_messages(source), egress.write_messages(sink)).inspect_err(
358            |e| {
359                if !matches!(e.kind(), io::ErrorKind::BrokenPipe | io::ErrorKind::UnexpectedEof) {
360                    log::warn!("Peer runner for {remote_addr} exited: {e}");
361                }
362            },
363        )?;
364        Ok(())
365    }
366
367    if let Some(pe::Crypto {
368        encryptor,
369        decryptor,
370    }) = crypto
371    {
372        async move {
373            let (source, sink) = stream.split();
374            let source = pe::DecryptingBufReader::new(source, decryptor);
375            let sink = pe::EncryptingWriter::new(sink, encryptor);
376            run_io(receiver, sender, source, sink).await
377        }
378        .boxed_local()
379    } else {
380        async move {
381            let (source, sink) = stream.split();
382            run_io(receiver, sender, source, sink).await
383        }
384        .boxed_local()
385    }
386}
387
388struct IngressProcessor {
389    remote_ip: SocketAddr,
390    ul_msg_sink: mpsc::Sender<UploaderMessage>,
391    dl_msg_sink: mpsc::Sender<DownloaderMessage>,
392    ext_msg_sink: Option<mpsc::Sender<ExtendedMessage>>,
393}
394
395impl IngressProcessor {
396    const RECV_TIMEOUT: Duration = sec!(120);
397
398    async fn read_messages<S: AsyncReadExt + Unpin>(mut self, mut source: S) -> io::Result<()> {
399        async fn read_one(
400            buffer: &mut [MaybeUninit<u8>],
401            mut source: impl AsyncReadExt + Unpin,
402        ) -> io::Result<PeerMessage> {
403            let msg_len = source.read_u32().await? as usize;
404            if msg_len > buffer.len() {
405                return Err(io::Error::new(
406                    io::ErrorKind::OutOfMemory,
407                    format!("msg len ({msg_len}) exceeds buffer size"),
408                ));
409            }
410            let mut readbuf = ReadBuf::uninit(buffer).limit(msg_len);
411            while 0 != source.read_buf(&mut readbuf).await? {}
412            let received = PeerMessage::decode_body(msg_len, &mut readbuf.into_inner().filled())?;
413            Ok(received)
414        }
415
416        const MAX_MSG_LEN: usize = MAX_BLOCK_SIZE + 512; // metadata block + bencoded header
417        let mut buffer = [MaybeUninit::<u8>::uninit(); MAX_MSG_LEN];
418
419        loop {
420            macro_rules! forward_and_continue {
421                ($msg:expr, $sink:expr) => {{
422                    log::trace!("{} => {}", self.remote_ip, $msg);
423                    $sink
424                        .send($msg)
425                        .await
426                        .map_err(|e| io::Error::new(io::ErrorKind::Other, Box::new(e)))?;
427                    continue;
428                }};
429            }
430
431            let received =
432                timeout(Self::RECV_TIMEOUT, read_one(&mut buffer, &mut source)).await??;
433
434            let received = match UploaderMessage::try_from(received) {
435                Ok(msg) => forward_and_continue!(msg, self.ul_msg_sink),
436                Err(received) => received,
437            };
438            let received = match DownloaderMessage::try_from(received) {
439                Ok(msg) => forward_and_continue!(msg, self.dl_msg_sink),
440                Err(received) => received,
441            };
442            let received = if let Some(ext_msg_sink) = &mut self.ext_msg_sink {
443                match ExtendedMessage::try_from(received) {
444                    Ok(msg) => forward_and_continue!(msg, ext_msg_sink),
445                    Err(received) => received,
446                }
447            } else {
448                received
449            };
450            if matches!(received, PeerMessage::KeepAlive) {
451                log::trace!("{} => {:?}", self.remote_ip, received);
452            } else {
453                log::error!("{} => unknown message: {:?}", self.remote_ip, received)
454            }
455        }
456    }
457}
458
459struct EgressProcessor {
460    remote_ip: SocketAddr,
461    dl_msg_source: mpsc::Receiver<Option<DownloaderMessage>>,
462    ul_msg_source: mpsc::Receiver<Option<UploaderMessage>>,
463    ext_msg_source: Option<mpsc::Receiver<Option<(ExtendedMessage, u8)>>>,
464}
465
466impl EgressProcessor {
467    const PING_INTERVAL: Duration = sec!(30);
468
469    async fn write_messages<S: AsyncWriteExt + Unpin>(mut self, mut sink: S) -> io::Result<()> {
470        fn channel_closed_err() -> io::Error {
471            io::Error::new(io::ErrorKind::BrokenPipe, "Channel closed")
472        }
473
474        async fn write_one(
475            buffer: &mut [MaybeUninit<u8>],
476            mut sink: impl AsyncWriteExt + Unpin,
477            msg: impl Into<PeerMessage>,
478        ) -> io::Result<()> {
479            let mut readbuf = ReadBuf::uninit(buffer);
480            msg.into().encode(&mut readbuf)?;
481            sink.write_all(readbuf.filled()).await?;
482            sink.flush().await?;
483            Ok(())
484        }
485
486        const MAX_MSG_LEN: usize = 32 * 1024 + 64; // 32 KiB outbound block + some header
487        let mut buffer = [MaybeUninit::<u8>::uninit(); MAX_MSG_LEN];
488
489        macro_rules! process_one {
490            ($($ext_msg_source:expr)?) => {
491                select! {
492                    biased;
493                    dl_msg = self.dl_msg_source.next() => {
494                        if let Some(msg) = dl_msg.ok_or_else(channel_closed_err)? {
495                            log::trace!("{} <= {}", self.remote_ip, msg);
496                            write_one(&mut buffer, &mut sink, msg).await?;
497                        }
498                    }
499                    $(ext_msg = $ext_msg_source.next() => {
500                        if let Some(msg) = ext_msg.ok_or_else(channel_closed_err)? {
501                            log::trace!("{} <= {}", self.remote_ip, msg.0);
502                            write_one(&mut buffer, &mut sink, msg).await?;
503                        }
504                    })?
505                    ul_msg = self.ul_msg_source.next() => {
506                        if let Some(msg) = ul_msg.ok_or_else(channel_closed_err)? {
507                            log::trace!("{} <= {}", self.remote_ip, msg);
508                            write_one(&mut buffer, &mut sink, msg).await?;
509                        }
510                    }
511                    _ = sleep(Self::PING_INTERVAL) => {
512                        let ping_msg = PeerMessage::KeepAlive;
513                        log::trace!("{} <= {:?}", self.remote_ip, &ping_msg);
514                        write_one(&mut buffer, &mut sink, ping_msg).await?;
515                    }
516                }
517            };
518        }
519
520        if let Some(ext_src) = &mut self.ext_msg_source {
521            loop {
522                process_one!(ext_src);
523            }
524        } else {
525            loop {
526                process_one!();
527            }
528        }
529    }
530}
531
532impl From<mpsc::SendError> for ChannelError {
533    fn from(_: mpsc::SendError) -> Self {
534        ChannelError::ConnectionClosed
535    }
536}
537
538impl From<ChannelError> for io::Error {
539    fn from(ce: ChannelError) -> Self {
540        match ce {
541            ChannelError::Timeout => {
542                io::Error::new(io::ErrorKind::TimedOut, "Peer channel timeout")
543            }
544            ChannelError::ConnectionClosed => io::Error::from(io::ErrorKind::BrokenPipe),
545        }
546    }
547}
548
549#[cfg(test)]
550mod tests {
551    use super::*;
552    use futures_util::join;
553    use std::collections::HashMap;
554    use std::net::{Ipv4Addr, SocketAddrV4};
555    use std::pin::Pin;
556    use std::task::{Context, Poll};
557    use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
558    use tokio::{io, task, time};
559    use tokio_test::io::Builder as MockBuilder;
560    use tokio_test::task::spawn;
561    use tokio_test::{assert_pending, assert_ready};
562
563    fn buffer_with(msgs: &[PeerMessage]) -> Vec<u8> {
564        let mut buf = Vec::new();
565        for msg in msgs {
566            msg.encode(&mut buf).unwrap();
567        }
568        buf
569    }
570
571    macro_rules! msgs {
572        ($($arg:expr),+ $(,)? ) => {
573            buffer_with(&[$($arg),+]).as_ref()
574        };
575    }
576
577    struct FakeSink(mpsc::UnboundedSender<u8>);
578    impl AsyncWrite for FakeSink {
579        fn poll_write(
580            self: Pin<&mut Self>,
581            _cx: &mut Context<'_>,
582            buf: &[u8],
583        ) -> Poll<Result<usize, io::Error>> {
584            for byte in buf {
585                self.0.unbounded_send(*byte).unwrap();
586            }
587            Poll::Ready(Ok(buf.len()))
588        }
589
590        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
591            Poll::Ready(Ok(()))
592        }
593
594        fn poll_shutdown(
595            self: Pin<&mut Self>,
596            _cx: &mut Context<'_>,
597        ) -> Poll<Result<(), io::Error>> {
598            self.0.close_channel();
599            Poll::Ready(Ok(()))
600        }
601    }
602    impl AsyncRead for FakeSink {
603        fn poll_read(
604            self: Pin<&mut Self>,
605            _cx: &mut Context<'_>,
606            _buf: &mut ReadBuf<'_>,
607        ) -> Poll<std::io::Result<()>> {
608            Poll::Pending
609        }
610    }
611
612    const HANDSHAKE_WITH_BEP10_SUPPORT: Handshake = Handshake {
613        reserved: ReservedBits {
614            data: *b"\x00\x00\x00\x00\x00\x10\x00\x00",
615            ..ReservedBits::ZERO
616        },
617        peer_id: [0u8; 20],
618        info_hash: [0u8; 20],
619    };
620
621    macro_rules! setup_channels {
622        ($stream:expr, $($args:expr),+ $(,)?) => {
623            setup_channels(StreamHolder($stream), $($args),+, None)
624        };
625    }
626
627    #[tokio::test]
628    async fn test_read_downloader_message() {
629        let socket = MockBuilder::new().read(msgs![PeerMessage::Interested]).build();
630        let (mut download, mut upload, extended, runner) = setup_channels!(
631            socket,
632            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
633            Default::default(),
634            false,
635        );
636        assert!(extended.is_none());
637
638        let upload_fut = async move {
639            let result = upload.1.receive_message().await;
640            assert!(matches!(result, Ok(DownloaderMessage::Interested)));
641        };
642
643        let run_fut = async move {
644            let _ = runner.await;
645        };
646
647        let download_fut = async move {
648            let result = download.1.receive_message().await;
649            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
650        };
651
652        join!(upload_fut, run_fut, download_fut);
653    }
654
655    #[tokio::test]
656    async fn test_read_uploader_message() {
657        let socket = MockBuilder::new().read(msgs![PeerMessage::Unchoke]).build();
658        let (mut download, mut upload, extended, runner) = setup_channels!(
659            socket,
660            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
661            Default::default(),
662            false,
663        );
664        assert!(extended.is_none());
665
666        let download_fut = async move {
667            let result = download.1.receive_message().await;
668            assert!(matches!(result, Ok(UploaderMessage::Unchoke)));
669        };
670
671        let run_fut = async move {
672            let _ = runner.await;
673        };
674
675        let upload_fut = async move {
676            let result = upload.1.receive_message().await;
677            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
678        };
679
680        join!(download_fut, run_fut, upload_fut);
681    }
682
683    #[tokio::test]
684    async fn test_read_extended_message() {
685        let socket = MockBuilder::new()
686            .read(msgs![
687                PeerMessage::Extended {
688                    id: 0,
689                    data: Vec::from(
690                        b"d1:md11:ut_metadatai1e6:ut_pexi2ee1:pi6881e1:v13:\xc2\xb5Torrent 1.2e",
691                    ),
692                },
693                PeerMessage::Extended {
694                    id: Extension::Metadata.local_id(),
695                    data: Vec::from("d8:msg_typei2e5:piecei3ee"),
696                },
697            ])
698            .build();
699
700        let (mut download, mut upload, extended, runner) = setup_channels!(
701            socket,
702            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
703            HANDSHAKE_WITH_BEP10_SUPPORT,
704            true,
705        );
706        assert!(extended.is_some());
707
708        let extended_fut = async move {
709            let ExtendedChannels(_tx, mut rx) = extended.unwrap();
710            let result = rx.receive_message().await;
711            let received = result.unwrap();
712            let expected_data = ExtendedHandshake {
713                extensions: HashMap::from([(Extension::Metadata, 1), (Extension::PeerExchange, 2)]),
714                listen_port: Some(6881),
715                client_type: Some("µTorrent 1.2".to_owned()),
716                ..Default::default()
717            };
718            assert!(matches!(received, ExtendedMessage::Handshake(data) if *data == expected_data));
719
720            let result = rx.receive_message().await;
721            let received = result.unwrap();
722            assert!(matches!(received, ExtendedMessage::MetadataReject { piece: 3 }));
723        };
724
725        let run_fut = async move {
726            let _ = runner.await;
727        };
728
729        let upload_fut = async move {
730            let result = upload.1.receive_message().await;
731            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
732        };
733
734        let download_fut = async move {
735            let result = download.1.receive_message().await;
736            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
737        };
738
739        join!(extended_fut, run_fut, upload_fut, download_fut);
740    }
741
742    #[tokio::test]
743    async fn test_read_uploader_and_downloader_and_extended_messages() {
744        let socket = MockBuilder::new()
745            .read(msgs![
746                PeerMessage::KeepAlive,
747                PeerMessage::Interested,
748                PeerMessage::Unchoke,
749                PeerMessage::KeepAlive,
750                PeerMessage::Extended {
751                    id: 0,
752                    data: Vec::from(b"d1:md11:ut_metadatai1eee"),
753                },
754            ])
755            .build();
756        let (mut download, mut upload, extended, runner) = setup_channels!(
757            socket,
758            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
759            HANDSHAKE_WITH_BEP10_SUPPORT,
760            true,
761        );
762
763        let upload_fut = async move {
764            let result = upload.1.receive_message().await;
765            assert!(matches!(result, Ok(DownloaderMessage::Interested)));
766        };
767
768        let download_fut = async move {
769            let result = download.1.receive_message().await;
770            assert!(matches!(result, Ok(UploaderMessage::Unchoke)));
771        };
772
773        let extended_fut = async move {
774            let result = extended.unwrap().1.receive_message().await;
775            let received = result.unwrap();
776            let expected_data = ExtendedHandshake {
777                extensions: HashMap::from([(Extension::Metadata, 1)]),
778                ..Default::default()
779            };
780            assert!(matches!(received, ExtendedMessage::Handshake(data) if *data == expected_data));
781        };
782
783        let run_fut = async move {
784            let result = runner.await;
785            let error = result.unwrap_err();
786            assert_eq!(io::ErrorKind::UnexpectedEof, error.kind());
787        };
788
789        join!(upload_fut, download_fut, extended_fut, run_fut);
790    }
791
792    #[tokio::test]
793    async fn test_read_error() {
794        let socket = MockBuilder::new()
795            .read_error(io::Error::from(io::ErrorKind::OutOfMemory))
796            .build();
797        let (mut download, mut upload, _, runner) = setup_channels!(
798            socket,
799            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
800            Default::default(),
801            false,
802        );
803
804        let download_fut = async move {
805            let result = download.1.receive_message().await;
806            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
807        };
808
809        let upload_fut = async move {
810            let result = upload.1.receive_message().await;
811            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
812        };
813
814        let run_fut = async move {
815            let result = runner.await;
816            let error = result.unwrap_err();
817            assert_eq!(io::ErrorKind::OutOfMemory, error.kind(), "{error}");
818        };
819
820        join!(download_fut, upload_fut, run_fut);
821    }
822
823    #[tokio::test(flavor = "local")]
824    async fn test_write_downloader_message() {
825        let socket = MockBuilder::new().write(msgs![PeerMessage::Interested]).wait(sec!(0)).build();
826        let (mut download, _upload, _, runner) = setup_channels!(
827            socket,
828            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
829            Default::default(),
830            false,
831        );
832
833        task::spawn_local(async move {
834            let _ = runner.await;
835        });
836
837        let result = download.0.send_message(DownloaderMessage::Interested).await;
838        assert!(result.is_ok(), "{result:?}");
839    }
840
841    #[tokio::test(flavor = "local")]
842    async fn test_write_uploader_message() {
843        let socket = MockBuilder::new().write(msgs![PeerMessage::Unchoke]).wait(sec!(0)).build();
844        let (_download, mut upload, _, runner) = setup_channels!(
845            socket,
846            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
847            Default::default(),
848            false,
849        );
850
851        task::spawn_local(async move {
852            let _ = runner.await;
853        });
854
855        let result = upload.0.send_message(UploaderMessage::Unchoke).await;
856        assert!(result.is_ok(), "{result:?}");
857    }
858
859    #[tokio::test(flavor = "local")]
860    async fn test_write_extended_messages() {
861        let socket = MockBuilder::new()
862            .write(msgs![
863                PeerMessage::Extended {
864                    id: 1,
865                    data: Vec::from("d8:msg_typei2e5:piecei3ee"),
866                },
867                PeerMessage::Extended {
868                    id: 0,
869                    data: Vec::from(
870                        b"d1:md11:ut_metadatai1e6:ut_pexi2ee1:pi6881e1:v13:\xc2\xb5Torrent 1.2e"
871                    ),
872                }
873            ])
874            .wait(sec!(0))
875            .build();
876        let (_download, _upload, extended, runner) = setup_channels!(
877            socket,
878            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
879            HANDSHAKE_WITH_BEP10_SUPPORT,
880            true,
881        );
882        assert!(extended.is_some());
883        let ExtendedChannels(mut tx, _rx) = extended.unwrap();
884
885        task::spawn_local(async move {
886            let _ = runner.await;
887        });
888
889        let result = tx.send_message((ExtendedMessage::MetadataReject { piece: 3 }, 1)).await;
890        assert!(result.is_ok());
891
892        let hs_data = ExtendedHandshake {
893            extensions: HashMap::from([(Extension::Metadata, 1), (Extension::PeerExchange, 2)]),
894            listen_port: Some(6881),
895            client_type: Some("µTorrent 1.2".to_owned()),
896            ..Default::default()
897        };
898        let result = tx.send_message((ExtendedMessage::Handshake(Box::new(hs_data)), 42)).await;
899        assert!(result.is_ok());
900    }
901
902    #[tokio::test]
903    async fn test_write_error() {
904        let socket = MockBuilder::new()
905            .write_error(io::Error::from(io::ErrorKind::OutOfMemory))
906            .build();
907        let (mut download, upload, _, runner) = setup_channels!(
908            socket,
909            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
910            Default::default(),
911            false,
912        );
913
914        let send_msg_fut = async move {
915            let result = download.0.send_message(DownloaderMessage::Interested).await;
916            assert!(matches!(result, Err(ChannelError::ConnectionClosed)));
917        };
918
919        let run_fut = async move {
920            let result = runner.await;
921            let error = result.unwrap_err();
922            assert_eq!(io::ErrorKind::OutOfMemory, error.kind(), "{error}");
923        };
924
925        join!(send_msg_fut, run_fut);
926        drop(upload);
927    }
928
929    #[tokio::test]
930    async fn test_writing_downloader_message_takes_priority_over_uploader_message() {
931        for _ in 0..50 {
932            let socket = MockBuilder::new()
933                .write(msgs![PeerMessage::Interested])
934                .write(msgs![PeerMessage::Piece {
935                    index: 0,
936                    begin: 0,
937                    block: vec![0u8; 1024]
938                }])
939                .wait(sec!(0))
940                .build();
941            let (mut download, mut upload, _, runner) = setup_channels!(
942                socket,
943                SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
944                Default::default(),
945                false,
946            );
947
948            let mut send_uploader_msg_fut = spawn(upload.0.send_message(UploaderMessage::Block(
949                BlockInfo {
950                    piece_index: 0,
951                    in_piece_offset: 0,
952                    block_length: 16384,
953                },
954                vec![0u8; 1024],
955            )));
956            let mut send_downloader_msg_fut =
957                spawn(download.0.send_message(DownloaderMessage::Interested));
958            let mut runner_fut = spawn(runner);
959
960            assert_pending!(send_uploader_msg_fut.poll());
961            assert_pending!(send_downloader_msg_fut.poll());
962
963            while matches!(send_uploader_msg_fut.poll(), Poll::Pending)
964                && matches!(send_downloader_msg_fut.poll(), Poll::Pending)
965            {
966                assert_pending!(runner_fut.poll());
967            }
968        }
969    }
970
971    #[tokio::test]
972    async fn test_writing_extended_message_takes_priority_over_uploader_message() {
973        for _ in 0..50 {
974            let socket = MockBuilder::new()
975                .write(msgs![PeerMessage::Extended {
976                    id: 0,
977                    data: Vec::from(
978                        b"d1:md11:ut_metadatai1e6:ut_pexi2ee1:pi6881e1:v13:\xc2\xb5Torrent 1.2e"
979                    ),
980                }])
981                .write(msgs![PeerMessage::Bitfield {
982                    bitfield: Bitfield::repeat(true, 42),
983                }])
984                .wait(sec!(0))
985                .build();
986            let (_download, mut upload, extended, runner) = setup_channels!(
987                socket,
988                SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
989                HANDSHAKE_WITH_BEP10_SUPPORT,
990                true,
991            );
992            let mut extended = extended.unwrap();
993
994            let mut send_uploader_msg_fut =
995                spawn(upload.0.send_message(UploaderMessage::Bitfield(Bitfield::repeat(true, 42))));
996            let mut send_extended_msg_fut = spawn(extended.0.send_message((
997                ExtendedMessage::Handshake(Box::new(ExtendedHandshake {
998                    extensions: HashMap::from([
999                        (Extension::Metadata, 1),
1000                        (Extension::PeerExchange, 2),
1001                    ]),
1002                    listen_port: Some(6881),
1003                    client_type: Some("µTorrent 1.2".to_owned()),
1004                    ..Default::default()
1005                })),
1006                42,
1007            )));
1008            let mut runner_fut = spawn(runner);
1009
1010            assert_pending!(send_uploader_msg_fut.poll());
1011            assert_pending!(send_extended_msg_fut.poll());
1012
1013            while matches!(send_uploader_msg_fut.poll(), Poll::Pending)
1014                && matches!(send_extended_msg_fut.poll(), Poll::Pending)
1015            {
1016                assert_pending!(runner_fut.poll());
1017            }
1018        }
1019    }
1020
1021    #[tokio::test(start_paused = true)]
1022    async fn test_downloader_channel_send_backpressure() {
1023        let socket = MockBuilder::new()
1024            .wait(sec!(1))
1025            .write(msgs![PeerMessage::Interested])
1026            .wait(sec!(1))
1027            .write(msgs![PeerMessage::NotInterested])
1028            .wait(sec!(1))
1029            .build();
1030
1031        let (mut download, _upload, _, runner) = setup_channels!(
1032            socket,
1033            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1034            Default::default(),
1035            false,
1036        );
1037
1038        let mut runner_fut = spawn(runner);
1039        {
1040            let mut send_fut = spawn(download.0.send_message(DownloaderMessage::Interested));
1041            assert_pending!(send_fut.poll());
1042
1043            assert_pending!(runner_fut.poll());
1044            assert_pending!(send_fut.poll());
1045
1046            time::sleep(sec!(1)).await;
1047            assert_pending!(runner_fut.poll());
1048            assert!(assert_ready!(send_fut.poll()).is_ok());
1049        }
1050        {
1051            let mut send_fut = spawn(download.0.send_message(DownloaderMessage::NotInterested));
1052            assert_pending!(send_fut.poll());
1053
1054            assert_pending!(runner_fut.poll());
1055            assert_pending!(send_fut.poll());
1056
1057            time::sleep(sec!(1)).await;
1058            assert_pending!(runner_fut.poll());
1059            assert!(assert_ready!(send_fut.poll()).is_ok());
1060        }
1061    }
1062
1063    #[tokio::test(start_paused = true)]
1064    async fn test_uploader_channel_send_backpressure() {
1065        let socket = MockBuilder::new()
1066            .wait(sec!(1))
1067            .write(msgs![PeerMessage::Choke])
1068            .wait(sec!(1))
1069            .write(msgs![PeerMessage::Unchoke])
1070            .wait(sec!(1))
1071            .build();
1072
1073        let (_download, mut upload, _, runner) = setup_channels!(
1074            socket,
1075            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1076            Default::default(),
1077            false,
1078        );
1079
1080        let mut runner_fut = spawn(runner);
1081        {
1082            let mut send_fut = spawn(upload.0.send_message(UploaderMessage::Choke));
1083            assert_pending!(send_fut.poll());
1084
1085            assert_pending!(runner_fut.poll());
1086            assert_pending!(send_fut.poll());
1087
1088            time::sleep(sec!(1)).await;
1089            assert_pending!(runner_fut.poll());
1090            assert!(assert_ready!(send_fut.poll()).is_ok());
1091        }
1092        {
1093            let mut send_fut = spawn(upload.0.send_message(UploaderMessage::Unchoke));
1094            assert_pending!(send_fut.poll());
1095
1096            assert_pending!(runner_fut.poll());
1097            assert_pending!(send_fut.poll());
1098
1099            time::sleep(sec!(1)).await;
1100            assert_pending!(runner_fut.poll());
1101            assert!(assert_ready!(send_fut.poll()).is_ok());
1102        }
1103    }
1104
1105    #[tokio::test(start_paused = true)]
1106    async fn test_extended_channel_send_backpressure() {
1107        let socket = MockBuilder::new()
1108            .wait(sec!(1))
1109            .write(msgs![PeerMessage::Extended {
1110                id: 1,
1111                data: Vec::from("d8:msg_typei2e5:piecei3ee"),
1112            }])
1113            .wait(sec!(1))
1114            .write(msgs![PeerMessage::Extended {
1115                id: 1,
1116                data: Vec::from("d8:msg_typei0e5:piecei3ee"),
1117            }])
1118            .wait(sec!(1))
1119            .build();
1120        let (_download, _upload, extended, runner) = setup_channels!(
1121            socket,
1122            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1123            HANDSHAKE_WITH_BEP10_SUPPORT,
1124            true,
1125        );
1126        let mut extended = extended.unwrap();
1127
1128        let mut runner_fut = spawn(runner);
1129        {
1130            let mut send_fut =
1131                spawn(extended.0.send_message((ExtendedMessage::MetadataReject { piece: 3 }, 1)));
1132            assert_pending!(send_fut.poll());
1133
1134            assert_pending!(runner_fut.poll());
1135            assert_pending!(send_fut.poll());
1136
1137            time::sleep(sec!(1)).await;
1138            assert_pending!(runner_fut.poll());
1139            assert!(assert_ready!(send_fut.poll()).is_ok());
1140        }
1141        {
1142            let mut send_fut =
1143                spawn(extended.0.send_message((ExtendedMessage::MetadataRequest { piece: 3 }, 1)));
1144            assert_pending!(send_fut.poll());
1145
1146            assert_pending!(runner_fut.poll());
1147            assert_pending!(send_fut.poll());
1148
1149            time::sleep(sec!(1)).await;
1150            assert_pending!(runner_fut.poll());
1151            assert!(assert_ready!(send_fut.poll()).is_ok());
1152        }
1153    }
1154
1155    #[tokio::test(start_paused = true)]
1156    async fn test_clone_channel_and_send_msgs_concurrently() {
1157        let socket = MockBuilder::new()
1158            .write(msgs![PeerMessage::Have { piece_index: 0 }])
1159            .write(msgs![PeerMessage::Unchoke])
1160            .wait(sec!(1))
1161            .build();
1162
1163        let (_download, UploadChannels(mut tx, _), _, runner) = setup_channels!(
1164            socket,
1165            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1166            Default::default(),
1167            false,
1168        );
1169        let mut runner_fut = spawn(runner);
1170
1171        let mut tx_clone = tx.clone();
1172
1173        let mut send_have_fut = spawn(tx.send_message(UploaderMessage::Have { piece_index: 0 }));
1174        assert_pending!(send_have_fut.poll());
1175
1176        let mut send_unchoke_fut = spawn(tx_clone.send_message(UploaderMessage::Unchoke));
1177        assert_pending!(send_unchoke_fut.poll());
1178
1179        assert_pending!(runner_fut.poll()); // this used to panic
1180        assert_pending!(send_have_fut.poll());
1181        assert_pending!(send_unchoke_fut.poll());
1182
1183        assert_pending!(runner_fut.poll());
1184        assert_ready!(send_have_fut.poll()).expect("send_message() returned Error");
1185        assert_ready!(send_unchoke_fut.poll()).expect("send_message() returned Error");
1186    }
1187
1188    #[tokio::test(start_paused = true)]
1189    async fn test_send_keepalive_every_30s() {
1190        task::LocalSet::new()
1191            .run_until(async {
1192                let mut buf = Vec::<u8>::new();
1193                let (writer, mut reader) = mpsc::unbounded::<u8>();
1194
1195                let (_download, _upload, _, runner) = setup_channels!(
1196                    FakeSink(writer),
1197                    SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1198                    Default::default(),
1199                    false,
1200                );
1201
1202                task::spawn_local(async move {
1203                    let _ = runner.await;
1204                });
1205
1206                time::sleep(sec!(30)).await;
1207                assert!(reader.try_recv().is_err());
1208
1209                task::yield_now().await;
1210                while let Ok(byte) = reader.try_recv() {
1211                    buf.push(byte);
1212                }
1213                assert_eq!(4, buf.len());
1214                assert_eq!(&[0u8; 4], &buf[..4]);
1215
1216                time::sleep(sec!(30)).await;
1217                assert!(reader.try_recv().is_err());
1218
1219                task::yield_now().await;
1220                while let Ok(byte) = reader.try_recv() {
1221                    buf.push(byte);
1222                }
1223                assert_eq!(8, buf.len());
1224                assert_eq!(&[0u8; 4], &buf[4..8]);
1225            })
1226            .await;
1227    }
1228
1229    #[tokio::test(start_paused = true, flavor = "local")]
1230    async fn test_receiver_times_out_after_2_min() {
1231        let (sock1, _sock2) = io::duplex(0);
1232        let (_download, _upload, _, runner) = setup_channels!(
1233            sock1,
1234            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1235            Default::default(),
1236            false,
1237        );
1238
1239        let (mut result_sender, mut result_receiver) = mpsc::channel::<io::Result<()>>(1);
1240
1241        task::spawn_local(async move {
1242            let result = runner.await;
1243            result_sender.try_send(result).unwrap();
1244        });
1245
1246        time::sleep(sec!(120)).await;
1247        assert!(result_receiver.try_recv().is_err());
1248
1249        task::yield_now().await;
1250        let error = result_receiver.try_recv().expect("Runner not finished");
1251        assert_eq!(io::ErrorKind::TimedOut, error.unwrap_err().kind());
1252    }
1253
1254    #[tokio::test(start_paused = true)]
1255    async fn test_channel_send_timeout() {
1256        task::LocalSet::new()
1257            .run_until(async {
1258                const TIMEOUT: Duration = sec!(10);
1259
1260                let (sock1, _sock2) = io::duplex(0);
1261                let (mut download, mut _upload, _, _runner) = setup_channels!(
1262                    sock1,
1263                    SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1264                    Default::default(),
1265                    false,
1266                );
1267
1268                let (mut result_sender, mut result_receiver) =
1269                    mpsc::channel::<Result<(), ChannelError>>(1);
1270
1271                task::spawn_local(async move {
1272                    let result = download
1273                        .0
1274                        .send_message_timed(DownloaderMessage::NotInterested, TIMEOUT)
1275                        .await;
1276                    result_sender.try_send(result).unwrap();
1277                });
1278
1279                time::sleep(TIMEOUT).await;
1280                assert!(result_receiver.try_recv().is_err());
1281
1282                task::yield_now().await;
1283                let result = result_receiver.try_recv().expect("send not finished");
1284                assert!(matches!(result, Err(ChannelError::Timeout)));
1285            })
1286            .await;
1287    }
1288
1289    #[tokio::test(start_paused = true)]
1290    async fn test_channel_receive_timeout() {
1291        task::LocalSet::new()
1292            .run_until(async {
1293                const TIMEOUT: Duration = sec!(10);
1294
1295                let (sock1, _sock2) = io::duplex(0);
1296                let (mut _download, mut upload, _, _runner) = setup_channels!(
1297                    sock1,
1298                    SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1299                    Default::default(),
1300                    false,
1301                );
1302
1303                let (mut result_sender, mut result_receiver) =
1304                    mpsc::channel::<Result<DownloaderMessage, ChannelError>>(1);
1305
1306                task::spawn_local(async move {
1307                    let result = upload.1.receive_message_timed(TIMEOUT).await;
1308                    result_sender.try_send(result).unwrap();
1309                });
1310
1311                time::sleep(TIMEOUT).await;
1312                assert!(result_receiver.try_recv().is_err());
1313
1314                task::yield_now().await;
1315                let result = result_receiver.try_recv().expect("receive not finished");
1316                assert!(matches!(result, Err(ChannelError::Timeout)));
1317            })
1318            .await;
1319    }
1320
1321    #[tokio::test(start_paused = true, flavor = "local")]
1322    async fn test_channel_receive_zero_timeout() {
1323        let socket = MockBuilder::new()
1324            .read(msgs![
1325                PeerMessage::Have { piece_index: 42 },
1326                PeerMessage::Have { piece_index: 43 },
1327            ])
1328            .wait(sec!(0))
1329            .build();
1330
1331        let (mut download, _upload, _, runner) = setup_channels!(
1332            socket,
1333            SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 6666)),
1334            Default::default(),
1335            false,
1336        );
1337
1338        task::spawn_local(async move {
1339            let _ = runner.await;
1340        });
1341
1342        task::yield_now().await;
1343
1344        let res = download.1.receive_message_timed(sec!(0)).await;
1345        let msg = res.unwrap();
1346        assert!(matches!(msg, UploaderMessage::Have { piece_index: 42 }));
1347
1348        let res = download.1.receive_message_timed(sec!(0)).await;
1349        let msg = res.unwrap();
1350        assert!(matches!(msg, UploaderMessage::Have { piece_index: 43 }));
1351
1352        let res = download.1.receive_message_timed(sec!(0)).await;
1353        let err = res.unwrap_err();
1354        assert!(matches!(err, ChannelError::Timeout));
1355    }
1356}