Skip to main content

librqbit_peer_protocol/
lib.rs

1// BitTorrent peer protocol implementation: parsing, serialization etc.
2//
3// Can be used outside of librqbit.
4
5mod double_buf;
6pub mod extended;
7
8use std::hint::unreachable_unchecked;
9
10use buffers::{ByteBuf, ByteBufOwned};
11use byteorder::{BE, ByteOrder};
12use bytes::Bytes;
13use clone_to_owned::CloneToOwned;
14use extended::PeerExtendedMessageIds;
15use librqbit_core::{constants::CHUNK_SIZE, hash_id::Id20, lengths::ChunkInfo};
16use serde_derive::{Deserialize, Serialize};
17
18pub use crate::double_buf::DoubleBufHelper;
19
20use self::extended::ExtendedMessage;
21
22const INTEGER_LEN: usize = 4;
23const MSGID_LEN: usize = 1;
24const PREAMBLE_LEN: usize = INTEGER_LEN + MSGID_LEN;
25const PIECE_MESSAGE_PREAMBLE_LEN: usize = PREAMBLE_LEN + INTEGER_LEN * 2;
26pub const PIECE_MESSAGE_DEFAULT_LEN: usize = PIECE_MESSAGE_PREAMBLE_LEN + CHUNK_SIZE as usize;
27
28// extended message ut_metadata request is the largest known message.
29const MAX_MSG_LEN_LEN_JUST_IN_CASE_EXTRA: usize = 64;
30pub const MAX_MSG_LEN: usize = PREAMBLE_LEN
31    + 1
32    + b"d8:msg_typei1e5:piecei42e10:total_sizei16384ee".len()
33    + CHUNK_SIZE as usize
34    + MAX_MSG_LEN_LEN_JUST_IN_CASE_EXTRA;
35
36const PSTR_BT1: &str = "BitTorrent protocol";
37
38type MsgId = u8;
39
40const MSGID_CHOKE: MsgId = 0;
41const MSGID_UNCHOKE: MsgId = 1;
42const MSGID_INTERESTED: MsgId = 2;
43const MSGID_NOT_INTERESTED: MsgId = 3;
44const MSGID_HAVE: MsgId = 4;
45const MSGID_BITFIELD: MsgId = 5;
46const MSGID_REQUEST: MsgId = 6;
47const MSGID_PIECE: MsgId = 7;
48const MSGID_CANCEL: MsgId = 8;
49const MSGID_EXTENDED: MsgId = 20;
50
51pub const EXTENDED_UT_METADATA_KEY: &[u8] = b"ut_metadata";
52pub const MY_EXTENDED_UT_METADATA: u8 = 3;
53
54pub const EXTENDED_UT_PEX_KEY: &[u8] = b"ut_pex";
55pub const MY_EXTENDED_UT_PEX: u8 = 1;
56
57#[derive(Clone, Copy)]
58pub struct MsgIdDebug(MsgId);
59impl MsgIdDebug {
60    const fn name(&self) -> Option<&'static str> {
61        let n = match self.0 {
62            MSGID_CHOKE => "choke",
63            MSGID_UNCHOKE => "unchoke",
64            MSGID_INTERESTED => "interested",
65            MSGID_NOT_INTERESTED => "not_interested",
66            MSGID_HAVE => "have",
67            MSGID_BITFIELD => "bitfield",
68            MSGID_REQUEST => "request",
69            MSGID_PIECE => "piece",
70            MSGID_CANCEL => "cancel",
71            MSGID_EXTENDED => "extended",
72            _ => return None,
73        };
74        Some(n)
75    }
76}
77impl core::fmt::Debug for MsgIdDebug {
78    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
79        match self.name() {
80            Some(name) => f.write_str(name),
81            None => write!(f, "<unknown msg_id {}>", self.0),
82        }
83    }
84}
85
86#[derive(thiserror::Error, Debug)]
87pub enum MessageDeserializeError {
88    #[error("not enough data (msgid={1:?}): expected at least {0} more bytes")]
89    NotEnoughData(usize, Option<MsgIdDebug>),
90    #[error("need a contiguous input to deserialize")]
91    NeedContiguous,
92    #[error("unsupported message id {0}")]
93    UnsupportedMessageId(u8),
94    #[error(transparent)]
95    Bencode(#[from] bencode::DeserializeError),
96    #[error("incorrect message length msg_id={msg_id:?}, expected={expected}, received={received}")]
97    IncorrectMsgLen {
98        received: u32,
99        expected: u32,
100        msg_id: MsgIdDebug,
101    },
102    #[error("ut_metadata:data received {received_len} >= total_size is {total_size}")]
103    UtMetadataBufLargerThanTotalSize { total_size: u32, received_len: u32 },
104    #[error("ut_metadata:data length must be <= {CHUNK_SIZE} but received {0} bytes")]
105    UtMetadataTooLarge(u32),
106    #[error("ut_metadata: trailing bytes when decoding")]
107    UtMetadataTrailingBytes,
108    #[error("ut_metadata: missing total_size")]
109    UtMetadataMissingTotalSize,
110    #[error("ut_metadata: unrecognized message type: {0}")]
111    UtMetadataTypeUnknown(u32),
112    #[error("ut_metadata: received piece {received_piece} > total pieces {total_pieces}")]
113    UtMetadataPieceOutOfBounds {
114        total_pieces: u32,
115        received_piece: u32,
116    },
117    #[error("ut_metadata: expected size {expected_size} != received size {received_size}")]
118    UtMetadataSizeMismatch {
119        expected_size: u32,
120        received_size: u32,
121    },
122    #[error("pstr doesn't match {PSTR_BT1:?}")]
123    HandshakePstrWrongContent,
124    #[error("pstr should be 19 bytes long but got {0}")]
125    HandshakePstrWrongLength(u8),
126}
127
128pub fn serialize_piece_preamble(chunk: &ChunkInfo, mut buf: &mut [u8]) -> usize {
129    let len_prefix = MSGID_LEN as u32 + INTEGER_LEN as u32 * 2 + chunk.size;
130    BE::write_u32(&mut buf[0..4], len_prefix);
131    buf[4] = MSGID_PIECE;
132
133    buf = &mut buf[5..];
134    BE::write_u32(&mut buf[0..4], chunk.piece_index.get());
135    BE::write_u32(&mut buf[4..8], chunk.offset);
136
137    PIECE_MESSAGE_PREAMBLE_LEN
138}
139
140pub struct Piece<B> {
141    pub index: u32,
142    pub begin: u32,
143    block_0: B,
144    block_1: B,
145}
146
147impl<B: AsRef<[u8]>> std::fmt::Debug for Piece<B> {
148    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
149        f.debug_struct("Piece")
150            .field("index", &self.index)
151            .field("begin", &self.begin)
152            .field("len", &self.len())
153            .field("len_0", &self.block_0.as_ref().len())
154            .field("len_1", &self.block_1.as_ref().len())
155            .finish_non_exhaustive()
156    }
157}
158
159impl CloneToOwned for Piece<ByteBuf<'_>> {
160    type Target = Piece<ByteBufOwned>;
161
162    fn clone_to_owned(&self, within_buffer: Option<&Bytes>) -> Self::Target {
163        Piece {
164            index: self.index,
165            begin: self.begin,
166            block_0: self.block_0.clone_to_owned(within_buffer),
167            block_1: self.block_1.clone_to_owned(within_buffer),
168        }
169    }
170}
171
172impl<B: AsRef<[u8]>> Piece<B> {
173    #[allow(clippy::len_without_is_empty)]
174    pub fn len(&self) -> usize {
175        self.block_0.as_ref().len() + self.block_1.as_ref().len()
176    }
177
178    pub fn serialize_unchecked_len(&self, mut buf: &mut [u8]) -> usize {
179        buf[0..4].copy_from_slice(&self.index.to_be_bytes());
180        buf[4..8].copy_from_slice(&self.begin.to_be_bytes());
181        buf = &mut buf[8..];
182
183        let b0 = self.block_0.as_ref();
184        let b1 = self.block_1.as_ref();
185
186        buf[..b0.len()].copy_from_slice(b0);
187        buf = &mut buf[b0.len()..];
188        buf[..b1.len()].copy_from_slice(b1);
189        8 + b0.len() + b1.len()
190    }
191}
192
193impl Piece<ByteBufOwned> {
194    pub fn as_borrowed(&self) -> Piece<ByteBuf<'_>> {
195        Piece {
196            index: self.index,
197            begin: self.begin,
198            block_0: self.block_0.as_ref().into(),
199            block_1: self.block_1.as_ref().into(),
200        }
201    }
202}
203
204impl<'a> Piece<ByteBuf<'a>> {
205    pub fn data(&self) -> (&'a [u8], &'a [u8]) {
206        (self.block_0.0, self.block_1.0)
207    }
208
209    pub fn from_data(index: u32, begin: u32, block: &'a [u8]) -> Self {
210        Piece {
211            index,
212            begin,
213            block_0: ByteBuf(block),
214            block_1: ByteBuf(&[]),
215        }
216    }
217}
218
219#[derive(Debug)]
220pub enum Message<'a> {
221    Request(Request),
222    Cancel(Request),
223    Bitfield(ByteBuf<'a>),
224    KeepAlive,
225    Have(u32),
226    Choke,
227    Unchoke,
228    Interested,
229    NotInterested,
230    Piece(Piece<ByteBuf<'a>>),
231    Extended(ExtendedMessage<ByteBuf<'a>>),
232}
233
234#[derive(thiserror::Error, Debug)]
235pub enum SerializeError {
236    #[error("not enough space in buffer")]
237    NoSpaceInBuffer,
238    #[error(transparent)]
239    Bencode(#[from] bencode::SerializeError),
240    #[error("need peer's handshake to serialize ut_metadata, or peer does't support ut_metadata")]
241    NeedUtMetadata,
242    #[error("need peer's handshake to serialize ut_pex, or peer does't support ut_pex")]
243    NeedPex,
244}
245
246impl From<std::io::Error> for SerializeError {
247    fn from(_: std::io::Error) -> Self {
248        Self::NoSpaceInBuffer
249    }
250}
251
252impl Message<'_> {
253    pub fn serialize(
254        &self,
255        out: &mut [u8],
256        peer_extended_messages: &dyn Fn() -> PeerExtendedMessageIds,
257    ) -> Result<usize, SerializeError> {
258        macro_rules! check_len {
259            ($l:expr) => {
260                if out.len() < $l {
261                    return Err(SerializeError::NoSpaceInBuffer);
262                }
263            };
264        }
265
266        macro_rules! write_preamble {
267            ($msg_len:expr, $msg_id:expr) => {
268                out[0..4].copy_from_slice(&(($msg_len + 1u32).to_be_bytes()));
269                out[4] = $msg_id;
270            };
271        }
272
273        match self {
274            Message::Request(request) | Message::Cancel(request) => {
275                const TOTAL_LEN: usize = PREAMBLE_LEN + INTEGER_LEN * 3;
276                check_len!(TOTAL_LEN);
277                let msg_id = match self {
278                    Message::Request(..) => MSGID_REQUEST,
279                    Message::Cancel(..) => MSGID_CANCEL,
280                    _ => unsafe { unreachable_unchecked() },
281                };
282                write_preamble!((INTEGER_LEN * 3) as u32, msg_id);
283                request.serialize_unchecked_len(&mut out[PREAMBLE_LEN..]);
284                Ok(TOTAL_LEN)
285            }
286            Message::Bitfield(b) => {
287                let block_len = b.as_ref().len();
288                let total_len: usize = PREAMBLE_LEN + block_len;
289                check_len!(total_len);
290                write_preamble!(block_len as u32, MSGID_BITFIELD);
291                out[PREAMBLE_LEN..PREAMBLE_LEN + block_len].copy_from_slice(b.as_ref());
292                Ok(total_len)
293            }
294            Message::Choke | Message::Unchoke | Message::Interested | Message::NotInterested => {
295                check_len!(PREAMBLE_LEN);
296                let msg_id = match self {
297                    Message::Choke => MSGID_CHOKE,
298                    Message::Unchoke => MSGID_UNCHOKE,
299                    Message::Interested => MSGID_INTERESTED,
300                    Message::NotInterested => MSGID_NOT_INTERESTED,
301                    _ => unsafe { unreachable_unchecked() },
302                };
303                write_preamble!(0, msg_id);
304                Ok(PREAMBLE_LEN)
305            }
306            Message::Piece(p) => {
307                let block_len = p.len();
308                let payload_len = INTEGER_LEN * 2 + block_len;
309                let total_len = PREAMBLE_LEN + payload_len;
310                check_len!(total_len);
311                write_preamble!(payload_len as u32, MSGID_PIECE);
312                p.serialize_unchecked_len(&mut out[PREAMBLE_LEN..]);
313                Ok(total_len)
314            }
315            Message::KeepAlive => {
316                check_len!(4);
317                out[0..4].copy_from_slice(&0u32.to_be_bytes());
318                Ok(4)
319            }
320            Message::Have(v) => {
321                check_len!(PREAMBLE_LEN + INTEGER_LEN);
322                write_preamble!(INTEGER_LEN as u32, MSGID_HAVE);
323                out[5..9].copy_from_slice(&v.to_be_bytes());
324                Ok(9)
325            }
326            Message::Extended(e) => {
327                check_len!(PREAMBLE_LEN + 2);
328                let msg_len = e.serialize(&mut out[PREAMBLE_LEN..], peer_extended_messages)?;
329                write_preamble!(msg_len as u32, MSGID_EXTENDED);
330                Ok(PREAMBLE_LEN + msg_len)
331            }
332        }
333    }
334}
335
336impl Message<'_> {
337    pub fn deserialize<'a>(
338        buf: &'a [u8],
339        buf2: &'a [u8],
340    ) -> Result<(Message<'a>, usize), MessageDeserializeError> {
341        let mut buf = DoubleBufHelper::new(buf, buf2);
342        let len_prefix = buf
343            .read_u32_be()
344            .map_err(|rem| MessageDeserializeError::NotEnoughData(rem, None))?;
345        let total_len = len_prefix as usize + 4;
346        if len_prefix == 0 {
347            return Ok((Message::KeepAlive, total_len));
348        }
349
350        let msg_id = buf.read_u8().ok_or(MessageDeserializeError::NotEnoughData(
351            len_prefix as usize,
352            None,
353        ))?;
354
355        let msg_len = len_prefix as usize - 1;
356        if buf.len() < msg_len {
357            return Err(MessageDeserializeError::NotEnoughData(
358                msg_len - buf.len(),
359                Some(MsgIdDebug(msg_id)),
360            ));
361        }
362
363        macro_rules! check_msg_len {
364            ($expected:expr) => {{
365                if msg_len != $expected {
366                    return Err(MessageDeserializeError::IncorrectMsgLen {
367                        received: len_prefix - 1,
368                        expected: $expected,
369                        msg_id: MsgIdDebug(msg_id),
370                    });
371                }
372            }};
373            (min $expected:expr) => {{
374                if msg_len < $expected {
375                    return Err(MessageDeserializeError::IncorrectMsgLen {
376                        received: len_prefix - 1,
377                        expected: $expected,
378                        msg_id: MsgIdDebug(msg_id),
379                    });
380                }
381            }};
382        }
383
384        match msg_id {
385            MSGID_CHOKE => {
386                check_msg_len!(0);
387                Ok((Message::Choke, total_len))
388            }
389            MSGID_UNCHOKE => {
390                check_msg_len!(0);
391                Ok((Message::Unchoke, total_len))
392            }
393            MSGID_INTERESTED => {
394                check_msg_len!(0);
395                Ok((Message::Interested, total_len))
396            }
397            MSGID_NOT_INTERESTED => {
398                check_msg_len!(0);
399                Ok((Message::NotInterested, total_len))
400            }
401            MSGID_HAVE => {
402                check_msg_len!(4);
403                let have = buf.read_u32_be().unwrap();
404                Ok((Message::Have(have), total_len))
405            }
406            MSGID_BITFIELD => {
407                check_msg_len!(min 1);
408                // In practice, as bitfield is always (almost) the first message, it should be contiguous.
409                let data = buf
410                    .get_contiguous(msg_len)
411                    .ok_or(MessageDeserializeError::NeedContiguous)?;
412                Ok((Message::Bitfield(ByteBuf::from(data)), total_len))
413            }
414            MSGID_REQUEST | MSGID_CANCEL => {
415                check_msg_len!(12);
416                const I32: usize = 4;
417                const I32_3: usize = I32 * 3;
418                let req = buf.consume::<I32_3>().unwrap();
419                let request = Request {
420                    index: BE::read_u32(&req[0..I32]),
421                    begin: BE::read_u32(&req[I32..I32 * 2]),
422                    length: BE::read_u32(&req[I32 * 2..I32 * 3]),
423                };
424                let req = if msg_id == MSGID_REQUEST {
425                    Message::Request(request)
426                } else {
427                    Message::Cancel(request)
428                };
429                Ok((req, total_len))
430            }
431            MSGID_PIECE => {
432                const MIN_PAYLOAD: usize = 1;
433                const MIN_LENGTH: usize = INTEGER_LEN * 2 + MIN_PAYLOAD;
434                if msg_len < MIN_LENGTH {
435                    return Err(MessageDeserializeError::IncorrectMsgLen {
436                        expected: MIN_LENGTH as u32,
437                        received: msg_len as u32,
438                        msg_id: MsgIdDebug(msg_id),
439                    });
440                }
441
442                let index = buf.read_u32_be().unwrap();
443                let begin = buf.read_u32_be().unwrap();
444
445                let block_len = msg_len - INTEGER_LEN * 2;
446                let (block_0, block_1) = buf.consume_variable(block_len).unwrap();
447
448                Ok((
449                    Message::Piece(Piece {
450                        index,
451                        begin,
452                        block_0: block_0.into(),
453                        block_1: block_1.into(),
454                    }),
455                    total_len,
456                ))
457            }
458            MSGID_EXTENDED => Ok((
459                Message::Extended(ExtendedMessage::deserialize(buf.with_max_len(msg_len))?),
460                PREAMBLE_LEN + msg_len,
461            )),
462            msg_id => Err(MessageDeserializeError::UnsupportedMessageId(msg_id)),
463        }
464    }
465}
466
467#[derive(Debug, PartialEq, Eq)]
468pub struct Handshake {
469    pub reserved: u64,
470    pub info_hash: Id20,
471    pub peer_id: Id20,
472}
473
474impl Handshake {
475    pub fn new(info_hash: Id20, peer_id: Id20) -> Handshake {
476        debug_assert_eq!(PSTR_BT1.len(), 19);
477
478        let mut reserved: u64 = 0;
479        // supports extended messaging
480        reserved |= 1 << 20;
481
482        Handshake {
483            reserved,
484            info_hash,
485            peer_id,
486        }
487    }
488
489    pub fn deserialize(b: &[u8]) -> Result<(Handshake, usize), MessageDeserializeError> {
490        const LEN: usize = 1 + PSTR_BT1.len() + 8 + 20 + 20;
491        if b.len() < LEN {
492            return Err(MessageDeserializeError::NotEnoughData(LEN - b.len(), None));
493        }
494        if b[0] as usize != PSTR_BT1.len() {
495            return Err(MessageDeserializeError::HandshakePstrWrongLength(b[0]));
496        }
497        if &b[1..20] != PSTR_BT1.as_bytes() {
498            return Err(MessageDeserializeError::HandshakePstrWrongContent);
499        }
500
501        let h = Handshake {
502            reserved: BE::read_u64(&b[20..28]),
503            info_hash: Id20::new(b[28..48].try_into().unwrap()),
504            peer_id: Id20::new(b[48..68].try_into().unwrap()),
505        };
506        Ok((h, LEN))
507    }
508
509    pub fn supports_extended(&self) -> bool {
510        self.reserved.to_be_bytes()[5] & 0x10 > 0
511    }
512
513    #[must_use]
514    pub fn serialize_unchecked_len(&self, buf: &mut [u8]) -> usize {
515        debug_assert_eq!(PSTR_BT1.len(), 19);
516        buf[0] = 19;
517        buf[1..20].copy_from_slice(PSTR_BT1.as_bytes());
518        buf[20..28].copy_from_slice(&self.reserved.to_be_bytes());
519        buf[28..48].copy_from_slice(&self.info_hash.0);
520        buf[48..68].copy_from_slice(&self.peer_id.0);
521        68
522    }
523}
524
525#[derive(Serialize, Deserialize, Debug, Clone, Copy)]
526pub struct Request {
527    pub index: u32,
528    pub begin: u32,
529    pub length: u32,
530}
531
532impl Request {
533    pub fn new(index: u32, begin: u32, length: u32) -> Self {
534        Self {
535            index,
536            begin,
537            length,
538        }
539    }
540
541    pub fn serialize_unchecked_len(&self, buf: &mut [u8]) -> usize {
542        buf[0..4].copy_from_slice(&self.index.to_be_bytes());
543        buf[4..8].copy_from_slice(&self.begin.to_be_bytes());
544        buf[8..12].copy_from_slice(&self.length.to_be_bytes());
545        12
546    }
547}
548
549#[cfg(test)]
550mod tests {
551    use anyhow::Context;
552
553    use crate::extended::handshake::ExtendedHandshake;
554
555    const EXTENDED: &[u8] = include_bytes!("../../librqbit/resources/test/extended-handshake.bin");
556
557    use super::*;
558    #[test]
559    fn test_handshake_serialize() {
560        let info_hash = Id20::new([
561            1u8, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20,
562        ]);
563        let peer_id = Id20::new([
564            1u8, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20,
565        ]);
566        let mut buf = [0u8; 100];
567        let se = Handshake::new(info_hash, peer_id);
568        let len = se.serialize_unchecked_len(&mut buf);
569        assert_eq!(len, 20 + 20 + 8 + 19 + 1);
570        assert_eq!(buf[0], 19);
571        assert_eq!(&buf[1..20], PSTR_BT1.as_bytes());
572        assert_eq!(&buf[28..48], &info_hash.0);
573        assert_eq!(&buf[48..68], &peer_id.0);
574
575        let (de, dlen) = Handshake::deserialize(&buf).unwrap();
576        assert_eq!(dlen, len);
577        assert_eq!(se, de);
578    }
579
580    #[test]
581    fn test_extended_serialize() {
582        let msg = Message::Extended(ExtendedMessage::Handshake(ExtendedHandshake::new()));
583        let mut out = [0u8; 100];
584        msg.serialize(&mut out, &Default::default).unwrap();
585        dbg!(out);
586    }
587
588    #[test]
589    fn test_deserialize_serialize_extended_non_contiguous() {
590        for split_point in 0..EXTENDED.len() {
591            let (first, second) = EXTENDED.split_at(split_point);
592            let res = Message::deserialize(first, second);
593            if split_point > PREAMBLE_LEN + 1 && split_point < EXTENDED.len() {
594                assert!(
595                    matches!(res, Err(MessageDeserializeError::NeedContiguous)),
596                    "expected NeedContiguous: {split_point}"
597                )
598            } else {
599                let (msg, len) = res
600                    .inspect_err(|e| panic!("split_point={split_point:?}; error: {e:#}"))
601                    .unwrap();
602                assert!(matches!(msg, Message::Extended(..)));
603                assert_eq!(len, EXTENDED.len());
604            }
605        }
606    }
607
608    #[test]
609    fn test_deserialize_piece() {
610        const LEN: usize = 100;
611        const EXTRA: usize = 100;
612        let mut buf = [0u8; LEN + EXTRA];
613
614        #[allow(clippy::needless_range_loop)]
615        for id in 0..buf.len() {
616            buf[id] = id as u8;
617        }
618
619        let block_len = LEN - PREAMBLE_LEN - INTEGER_LEN * 2;
620        let len_prefix: u32 = (block_len + INTEGER_LEN * 2 + MSGID_LEN) as u32;
621        let index: u32 = 42;
622        let begin: u32 = 43;
623
624        buf[0..4].copy_from_slice(&len_prefix.to_be_bytes());
625        buf[4] = MSGID_PIECE;
626        buf[5..9].copy_from_slice(&index.to_be_bytes());
627        buf[9..13].copy_from_slice(&begin.to_be_bytes());
628
629        for split_point in 0..buf.len() {
630            dbg!(split_point);
631            let (first, second) = buf.split_at(split_point);
632            let (msg, len) = Message::deserialize(first, second).unwrap();
633
634            let piece = match &msg {
635                Message::Piece(piece) => piece,
636                other => panic!("expected piece got {other:?}"),
637            };
638
639            assert_eq!(piece.len(), block_len);
640            assert_eq!(piece.index, index);
641            assert_eq!(piece.begin, begin);
642            assert_eq!(len, LEN);
643
644            let mut tmp = [0u8; 100];
645            let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
646            assert_eq!(slen, len);
647            assert_eq!(buf[..len], tmp[..len]);
648
649            let (first, second) = piece.data();
650
651            assert_eq!(first.len() + second.len(), block_len);
652            assert_eq!(first, &buf[13..13 + first.len()]);
653            assert_eq!(
654                second,
655                &buf[13 + first.len()..13 + first.len() + second.len()]
656            );
657        }
658    }
659
660    #[test]
661    fn test_deserialize_request() {
662        let mut buf = [0u8; 100];
663
664        let len_prefix: u32 = (MSGID_LEN + INTEGER_LEN * 3) as u32;
665        let index: u32 = 42;
666        let begin: u32 = 43;
667        let length: u32 = 44;
668
669        buf[0..4].copy_from_slice(&len_prefix.to_be_bytes());
670        buf[4] = MSGID_REQUEST;
671        buf[5..9].copy_from_slice(&index.to_be_bytes());
672        buf[9..13].copy_from_slice(&begin.to_be_bytes());
673        buf[13..17].copy_from_slice(&length.to_be_bytes());
674
675        for split_point in 0..buf.len() {
676            dbg!(split_point);
677            let (first, second) = buf.split_at(split_point);
678            let (msg, len) = Message::deserialize(first, second).unwrap();
679
680            let request = match msg {
681                Message::Request(req) => req,
682                other => panic!("expected request got {other:?}"),
683            };
684
685            assert_eq!(request.index, index);
686            assert_eq!(request.begin, begin);
687            assert_eq!(request.length, length);
688            assert_eq!(len, 17);
689
690            let mut tmp = [0u8; 100];
691            let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
692            assert_eq!(slen, len);
693            assert_eq!(buf[..len], tmp[..len]);
694        }
695    }
696
697    #[test]
698    fn test_keepalive() {
699        let buf = [0u8; 100];
700
701        for split_point in 0..buf.len() {
702            let (first, second) = buf.split_at(split_point);
703            let (msg, len) = Message::deserialize(first, second).unwrap();
704            assert!(matches!(msg, Message::KeepAlive));
705            assert_eq!(len, 4);
706            let mut tmp = [0u8; 100];
707            let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
708            assert_eq!(slen, len);
709            assert_eq!(buf[..len], tmp[..len]);
710        }
711    }
712
713    #[test]
714    fn test_have() {
715        let mut buf = [0u8; 100];
716        buf[0..4].copy_from_slice(&5u32.to_be_bytes());
717        buf[4] = MSGID_HAVE;
718        buf[5..9].copy_from_slice(&42u32.to_be_bytes());
719
720        for split_point in 0..buf.len() {
721            let (first, second) = buf.split_at(split_point);
722            let (msg, len) = Message::deserialize(first, second).unwrap();
723            assert!(matches!(msg, Message::Have(42)));
724            assert_eq!(len, 9);
725            let mut tmp = [0u8; 100];
726            let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
727            assert_eq!(slen, len);
728            assert_eq!(buf[..len], tmp[..len]);
729        }
730    }
731
732    #[test]
733    fn test_bitfield() {
734        let mut buf = [0u8; 100];
735        buf[0..4].copy_from_slice(&43u32.to_be_bytes());
736        buf[4] = MSGID_BITFIELD;
737        for byte in buf[5..47].iter_mut() {
738            *byte = 0b10101010;
739        }
740
741        for split_point in 0..buf.len() {
742            let (first, second) = buf.split_at(split_point);
743            let res = Message::deserialize(first, second);
744            if (6..47).contains(&split_point) {
745                assert!(
746                    matches!(res, Err(MessageDeserializeError::NeedContiguous)),
747                    "expected NeedContiguous: split_point={split_point}"
748                );
749                continue;
750            }
751            let (msg, len) = res.context(split_point).unwrap();
752            let bf = match &msg {
753                Message::Bitfield(bf) => bf,
754                other => panic!("expected bitfield, got {other:?}"),
755            };
756            assert_eq!(len, 47);
757            assert_eq!(bf.as_ref().len(), 42);
758            for byte in bf.as_ref() {
759                assert_eq!(*byte, 0b10101010);
760            }
761            let mut tmp = [0u8; 100];
762            let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
763            assert_eq!(slen, len);
764            assert_eq!(buf[..len], tmp[..len]);
765        }
766    }
767
768    #[test]
769    fn test_no_data_messages() {
770        let mut buf = [0u8; 100];
771
772        for msgid in [
773            MSGID_CHOKE,
774            MSGID_UNCHOKE,
775            MSGID_INTERESTED,
776            MSGID_NOT_INTERESTED,
777        ] {
778            buf[0..4].copy_from_slice(&1u32.to_be_bytes());
779            buf[4] = msgid;
780            for split_point in 0..buf.len() {
781                let (first, second) = buf.split_at(split_point);
782                let (msg, len) = Message::deserialize(first, second).unwrap();
783                match (msgid, &msg) {
784                    (MSGID_CHOKE, Message::Choke)
785                    | (MSGID_UNCHOKE, Message::Unchoke)
786                    | (MSGID_INTERESTED, Message::Interested)
787                    | (MSGID_NOT_INTERESTED, Message::NotInterested) => {}
788                    (msgid, msg) => panic!("msgid={msgid}, msg={msg:?}"),
789                }
790                assert_eq!(len, 5);
791                let mut tmp = [0u8; 100];
792                let slen = msg.serialize(&mut tmp, &|| Default::default()).unwrap();
793                assert_eq!(slen, len);
794                assert_eq!(buf[..len], tmp[..len]);
795            }
796        }
797    }
798}