1mod 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
28const 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 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 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}