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#[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
38pub struct PeerChannel<Q> {
40 peer_info: Arc<PeerInfo>,
41 inner: Q,
42}
43
44impl<Q> PeerChannel<Q> {
45 pub fn remote_ip(&self) -> &SocketAddr {
47 &self.peer_info.remote_addr
48 }
49 pub fn remote_info(&self) -> &Handshake {
51 &self.peer_info.handshake_info
52 }
53 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 pub async fn receive_message(&mut self) -> Result<Msg, ChannelError> {
75 self.inner.next().await.ok_or(ChannelError::ConnectionClosed)
76 }
77
78 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 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 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
109pub type DownloadTxChannel = TxChannel<DownloaderMessage>;
111pub type DownloadRxChannel = RxChannel<UploaderMessage>;
113pub struct DownloadChannels(pub DownloadTxChannel, pub DownloadRxChannel);
115
116pub type UploadTxChannel = TxChannel<UploaderMessage>;
118pub type UploadRxChannel = RxChannel<DownloaderMessage>;
120pub struct UploadChannels(pub UploadTxChannel, pub UploadRxChannel);
122
123pub type ExtendedTxChannel = TxChannel<(ExtendedMessage, u8)>;
125pub type ExtendedRxChannel = RxChannel<ExtendedMessage>;
127pub struct ExtendedChannels(pub ExtendedTxChannel, pub ExtendedRxChannel);
129
130const HANDSHAKE_TIMEOUT: Duration = sec!(10);
133
134pub 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
163pub 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#[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#[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
244fn 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; 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; 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()); 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}