librqbit_peer_protocol/extended/
mod.rs1use std::io::Cursor;
2
3use bencode::BencodeValue;
4use bencode::bencode_serialize_to_writer;
5use buffers::ByteBuf;
6use buffers::ByteBufT;
7use byteorder::WriteBytesExt;
8use serde_derive::Deserialize;
9use serde_derive::Serialize;
10use ut_pex::UtPex;
11
12use crate::DoubleBufHelper;
13use crate::MSGID_EXTENDED;
14use crate::MY_EXTENDED_UT_PEX;
15use crate::SerializeError;
16
17use self::{handshake::ExtendedHandshake, ut_metadata::UtMetadata};
18
19use super::MessageDeserializeError;
20
21pub mod handshake;
22pub mod ut_metadata;
23pub mod ut_pex;
24
25use super::MY_EXTENDED_UT_METADATA;
26
27#[derive(Debug, Default, Serialize, Deserialize, Clone, Copy, PartialEq, Eq)]
28pub struct PeerExtendedMessageIds {
29 pub ut_metadata: Option<u8>,
30 pub ut_pex: Option<u8>,
31}
32
33impl PeerExtendedMessageIds {
34 pub fn my() -> Self {
35 Self {
36 ut_metadata: Some(MY_EXTENDED_UT_METADATA),
37 ut_pex: Some(MY_EXTENDED_UT_PEX),
38 }
39 }
40}
41
42#[derive(Debug, Eq, PartialEq)]
43pub enum ExtendedMessage<ByteBuf: ByteBufT> {
44 Handshake(ExtendedHandshake<ByteBuf>),
45 UtMetadata(UtMetadata<ByteBuf>),
46 UtPex(UtPex<ByteBuf>),
47 Dyn(u8, BencodeValue<ByteBuf>),
48}
49
50impl<'a> ExtendedMessage<ByteBuf<'a>> {
51 pub fn serialize(
52 &self,
53 out: &mut [u8],
54 peer_extended_msg_ids: &dyn Fn() -> PeerExtendedMessageIds,
55 ) -> Result<usize, SerializeError> {
56 let mut out = Cursor::new(out);
57 match self {
58 ExtendedMessage::Dyn(msg_id, v) => {
59 out.write_u8(*msg_id)?;
60 bencode_serialize_to_writer(v, &mut out)?;
61 }
62 ExtendedMessage::Handshake(h) => {
63 out.write_u8(0)?;
64 bencode_serialize_to_writer(h, &mut out)?;
65 }
66 ExtendedMessage::UtMetadata(u) => {
67 let emsg_id = peer_extended_msg_ids()
68 .ut_metadata
69 .ok_or(SerializeError::NeedUtMetadata)?;
70 out.write_u8(emsg_id)?;
71 u.serialize(&mut out)?;
72 }
73 ExtendedMessage::UtPex(m) => {
74 let emsg_id = peer_extended_msg_ids()
75 .ut_pex
76 .ok_or(SerializeError::NeedPex)?;
77 out.write_u8(emsg_id)?;
78 bencode_serialize_to_writer(m, &mut out)?;
79 }
80 }
81 Ok(out.position() as usize)
82 }
83
84 pub fn deserialize(mut buf: DoubleBufHelper<'a>) -> Result<Self, MessageDeserializeError> {
85 let msg_id = crate::MsgIdDebug(MSGID_EXTENDED);
86 let emsg_id = buf
87 .read_u8()
88 .ok_or(MessageDeserializeError::NotEnoughData(1, Some(msg_id)))?;
89
90 fn from_bytes_contig<'a, T>(buf: &DoubleBufHelper<'a>) -> Result<T, MessageDeserializeError>
91 where
92 T: serde::de::Deserialize<'a>,
93 {
94 let buf = buf
95 .get_contiguous(buf.len())
96 .ok_or(MessageDeserializeError::NeedContiguous)?;
97 bencode::from_bytes(buf).map_err(|e| {
98 tracing::trace!("error deserializing extended: {e:#}");
99 MessageDeserializeError::Bencode(e.into_kind())
100 })
101 }
102
103 match emsg_id {
104 0 => Ok(ExtendedMessage::Handshake(from_bytes_contig(&buf)?)),
105 MY_EXTENDED_UT_METADATA => {
106 Ok(ExtendedMessage::UtMetadata(UtMetadata::deserialize(buf)?))
107 }
108 MY_EXTENDED_UT_PEX => Ok(ExtendedMessage::UtPex(from_bytes_contig(&buf)?)),
109 _ => Ok(ExtendedMessage::Dyn(emsg_id, from_bytes_contig(&buf)?)),
110 }
111 }
112}
113
114#[cfg(test)]
115mod tests {
116 use buffers::ByteBuf;
117
118 use crate::{
119 DoubleBufHelper, MessageDeserializeError,
120 extended::{
121 ExtendedMessage, PeerExtendedMessageIds,
122 ut_metadata::{UtMetadata, UtMetadataData},
123 },
124 };
125
126 #[track_caller]
127 fn ut_metadata_trailing_bytes_is_error(msg: ExtendedMessage<ByteBuf>) {
128 let mut buf = [0u8; 100];
129 let sz = msg
130 .serialize(&mut buf, &|| PeerExtendedMessageIds::my())
131 .unwrap();
132
133 let deserialized =
134 ExtendedMessage::deserialize(DoubleBufHelper::new(&buf[..sz], &[])).unwrap();
135 assert_eq!(msg, deserialized);
136
137 let res = ExtendedMessage::deserialize(DoubleBufHelper::new(&buf[..sz + 1], &[]));
138 assert!(
139 matches!(
140 res,
141 Err(MessageDeserializeError::UtMetadataTrailingBytes
142 | MessageDeserializeError::UtMetadataSizeMismatch {
143 expected_size: 5,
144 received_size: 6
145 })
146 ),
147 "expected trailing bytes error, got {res:?}"
148 )
149 }
150
151 #[test]
152 fn test_ut_metadata_trailing_bytes_is_error() {
153 ut_metadata_trailing_bytes_is_error(ExtendedMessage::UtMetadata(UtMetadata::Request(42)));
154 ut_metadata_trailing_bytes_is_error(ExtendedMessage::UtMetadata(UtMetadata::Reject(43)));
155 ut_metadata_trailing_bytes_is_error(ExtendedMessage::UtMetadata(UtMetadata::Data(
156 UtMetadataData::from_bytes(0, 5, b"\x42\x42\x42\x42\x42"[..].into()),
157 )));
158 }
159
160 #[test]
161 fn test_ut_metadata_non_contiguous() {
162 let mut buf = [0u8; 100];
163 let msg = ExtendedMessage::UtMetadata(UtMetadata::Data(UtMetadataData::from_bytes(
164 0,
165 5,
166 b"\x42\x42\x42\x42\x42"[..].into(),
167 )));
168 let sz = msg
169 .serialize(&mut buf, &|| PeerExtendedMessageIds::my())
170 .unwrap();
171 let bencode_sz = buf[..sz].iter().position(|byte| *byte == 0x42).unwrap();
172
173 for split_point in 0..sz {
174 let (d0, d1) = buf[..sz].split_at(split_point);
175 let buf = DoubleBufHelper::new(d0, d1);
176 let res = ExtendedMessage::deserialize(buf);
177 if (2..bencode_sz).contains(&split_point) {
178 assert!(
179 matches!(res, Err(MessageDeserializeError::NeedContiguous)),
180 "expected NeedContiguous, got {res:?}, split_point={split_point}, bencode_sz={bencode_sz}"
181 );
182 continue;
183 }
184 let de = res.unwrap();
185 match de {
186 ExtendedMessage::UtMetadata(UtMetadata::Data(d)) => {
187 assert_eq!(d.piece(), 0);
188 assert_eq!(d.len(), 5);
189 let mut debuf = [0u8; 5];
190 d.copy_to_slice(&mut debuf);
191 assert_eq!(debuf, b"\x42\x42\x42\x42\x42"[..]);
192 }
193 _ => panic!("bad msg"),
194 }
195 }
196 }
197}