Skip to main content

librqbit_peer_protocol/extended/
mod.rs

1use 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}