Skip to main content

cuprate_levin/
codec.rs

1// Rust Levin Library
2// Written in 2023 by
3//   Cuprate Contributors
4//
5// Permission is hereby granted, free of charge, to any person obtaining a copy
6// of this software and associated documentation files (the "Software"), to deal
7// in the Software without restriction, including without limitation the rights
8// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9// copies of the Software, and to permit persons to whom the Software is
10// furnished to do so, subject to the following conditions:
11//
12// The above copyright notice and this permission notice shall be included in all
13// copies or substantial portions of the Software.
14//
15
16//! A tokio-codec for levin buckets
17
18use std::{fmt::Debug, marker::PhantomData};
19
20use bytes::{Buf, BufMut, BytesMut};
21use tokio_util::codec::{Decoder, Encoder};
22
23use cuprate_helper::cast::u64_to_usize;
24
25use crate::{
26    header::{Flags, HEADER_SIZE},
27    message::{make_dummy_message, LevinMessage},
28    Bucket, BucketBuilder, BucketError, BucketHead, LevinBody, LevinCommand, MessageType, Protocol,
29};
30
31#[derive(Debug, Clone)]
32pub enum LevinBucketState<C> {
33    /// Waiting for the peer to send a header.
34    WaitingForHeader,
35    /// Waiting for a peer to send a body.
36    WaitingForBody(BucketHead<C>),
37}
38
39/// The levin tokio-codec for decoding and encoding raw levin buckets
40///
41#[derive(Debug, Clone)]
42pub struct LevinBucketCodec<C> {
43    state: LevinBucketState<C>,
44    protocol: Protocol,
45    handshake_message_seen: bool,
46}
47
48impl<C> Default for LevinBucketCodec<C> {
49    fn default() -> Self {
50        Self {
51            state: LevinBucketState::WaitingForHeader,
52            protocol: Protocol::default(),
53            handshake_message_seen: false,
54        }
55    }
56}
57
58impl<C> LevinBucketCodec<C> {
59    pub const fn new(protocol: Protocol) -> Self {
60        Self {
61            state: LevinBucketState::WaitingForHeader,
62            protocol,
63            handshake_message_seen: false,
64        }
65    }
66}
67
68impl<C: LevinCommand + Debug> Decoder for LevinBucketCodec<C> {
69    type Item = Bucket<C>;
70    type Error = BucketError;
71    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
72        loop {
73            match &self.state {
74                LevinBucketState::WaitingForHeader => {
75                    if src.len() < HEADER_SIZE {
76                        return Ok(None);
77                    }
78
79                    let head = BucketHead::<C>::from_bytes(src);
80
81                    if head.signature != self.protocol.signature {
82                        #[cfg(feature = "tracing")]
83                        tracing::debug!("Peer sent a levin header with an invalid signature.");
84
85                        return Err(BucketError::InvalidHeaderSignature);
86                    }
87
88                    #[cfg(feature = "tracing")]
89                    tracing::trace!(
90                        "Received new bucket header, command: {:?}, waiting for body, body len: {}",
91                        head.command,
92                        head.size
93                    );
94
95                    if head.size > self.protocol.max_packet_size
96                        || head.size > head.command.bucket_size_limit()
97                    {
98                        #[cfg(feature = "tracing")]
99                        tracing::debug!("Peer sent message which is too large.");
100
101                        return Err(BucketError::BucketExceededMaxSize);
102                    }
103
104                    if !self.handshake_message_seen {
105                        if head.size > self.protocol.max_packet_size_before_handshake {
106                            #[cfg(feature = "tracing")]
107                            tracing::debug!("Peer sent message which is too large.");
108
109                            return Err(BucketError::BucketExceededMaxSize);
110                        }
111
112                        if head.command.is_handshake() {
113                            #[cfg(feature = "tracing")]
114                            tracing::debug!(
115                                "Peer handshake message seen, increasing bucket size limit."
116                            );
117
118                            self.handshake_message_seen = true;
119                        }
120                    }
121
122                    drop(std::mem::replace(
123                        &mut self.state,
124                        LevinBucketState::WaitingForBody(head),
125                    ));
126                }
127                LevinBucketState::WaitingForBody(head) => {
128                    let body_len = u64_to_usize(head.size);
129                    if src.len() < body_len {
130                        src.reserve(body_len - src.len());
131                        return Ok(None);
132                    }
133
134                    let LevinBucketState::WaitingForBody(header) =
135                        std::mem::replace(&mut self.state, LevinBucketState::WaitingForHeader)
136                    else {
137                        unreachable!()
138                    };
139
140                    #[cfg(feature = "tracing")]
141                    tracing::trace!("Received full bucket for command: {:?}", header.command);
142
143                    return Ok(Some(Bucket {
144                        header,
145                        body: src.copy_to_bytes(body_len),
146                    }));
147                }
148            }
149        }
150    }
151}
152
153impl<C: LevinCommand> Encoder<Bucket<C>> for LevinBucketCodec<C> {
154    type Error = BucketError;
155    fn encode(&mut self, item: Bucket<C>, dst: &mut BytesMut) -> Result<(), Self::Error> {
156        if let Some(additional) = (HEADER_SIZE + item.body.len()).checked_sub(dst.capacity()) {
157            dst.reserve(additional);
158        }
159
160        item.header.write_bytes_into(dst);
161        dst.put_slice(&item.body);
162        Ok(())
163    }
164}
165
166#[derive(Default, Debug, Clone)]
167enum MessageState {
168    #[default]
169    WaitingForBucket,
170    /// Waiting for the rest of a fragmented message.
171    ///
172    /// We keep the fragmented message as a Vec<u8> instead of [`Bytes`](bytes::Bytes) as [`Bytes`](bytes::Bytes) could point to a
173    /// large allocation even if the [`Bytes`](bytes::Bytes) itself is small, so is not safe to keep around for long.
174    /// To prevent this attack vector completely we just use Vec<u8> for fragmented messages.
175    WaitingForRestOfFragment(Vec<u8>),
176}
177
178/// A tokio-codec for levin messages or in other words the decoded body
179/// of a levin bucket.
180#[derive(Debug, Clone)]
181pub struct LevinMessageCodec<T: LevinBody> {
182    message_ty: PhantomData<T>,
183    bucket_codec: LevinBucketCodec<T::Command>,
184    state: MessageState,
185}
186
187impl<T: LevinBody> Default for LevinMessageCodec<T> {
188    fn default() -> Self {
189        Self {
190            message_ty: Default::default(),
191            bucket_codec: Default::default(),
192            state: Default::default(),
193        }
194    }
195}
196
197impl<T: LevinBody> Decoder for LevinMessageCodec<T> {
198    type Item = T;
199    type Error = BucketError;
200    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
201        loop {
202            match &mut self.state {
203                MessageState::WaitingForBucket => {
204                    let Some(mut bucket) = self.bucket_codec.decode(src)? else {
205                        return Ok(None);
206                    };
207
208                    let flags = &bucket.header.flags;
209
210                    if flags.contains(Flags::DUMMY) {
211                        // Dummy message
212
213                        #[cfg(feature = "tracing")]
214                        tracing::trace!("Received DUMMY bucket from peer, ignoring.");
215                        // We may have another bucket in `src`.
216                        continue;
217                    }
218
219                    if flags.contains(Flags::END_FRAGMENT) {
220                        return Err(BucketError::InvalidHeaderFlags(
221                            "Flag end fragment received before a start fragment",
222                        ));
223                    }
224
225                    if flags.contains(Flags::START_FRAGMENT) {
226                        // monerod does not require a start flag before starting a fragmented message,
227                        // but will always produce one, so it is ok for us to require one.
228
229                        #[cfg(feature = "tracing")]
230                        tracing::debug!("Bucket is a fragment, waiting for rest of message.");
231
232                        self.state = MessageState::WaitingForRestOfFragment(bucket.body.to_vec());
233
234                        continue;
235                    }
236
237                    // Normal, non fragmented bucket
238
239                    let message_type = MessageType::from_flags_and_have_to_return(
240                        bucket.header.flags,
241                        bucket.header.have_to_return_data,
242                    )?;
243
244                    return Ok(Some(T::decode_message(
245                        &mut bucket.body,
246                        message_type,
247                        bucket.header.command,
248                    )?));
249                }
250                MessageState::WaitingForRestOfFragment(bytes) => {
251                    let Some(bucket) = self.bucket_codec.decode(src)? else {
252                        return Ok(None);
253                    };
254
255                    let flags = &bucket.header.flags;
256
257                    if flags.contains(Flags::DUMMY) {
258                        // Dummy message
259
260                        #[cfg(feature = "tracing")]
261                        tracing::trace!("Received DUMMY bucket from peer, ignoring.");
262                        // We may have another bucket in `src`.
263                        continue;
264                    }
265
266                    let max_size = u64_to_usize(if self.bucket_codec.handshake_message_seen {
267                        self.bucket_codec.protocol.max_packet_size
268                    } else {
269                        self.bucket_codec.protocol.max_packet_size_before_handshake
270                    });
271
272                    if bytes.len().saturating_add(bucket.body.len()) > max_size {
273                        return Err(BucketError::InvalidFragmentedMessage(
274                            "Fragmented message exceeded maximum size",
275                        ));
276                    }
277
278                    #[cfg(feature = "tracing")]
279                    tracing::trace!("Received another bucket fragment.");
280
281                    bytes.extend_from_slice(bucket.body.as_ref());
282
283                    if flags.contains(Flags::END_FRAGMENT) {
284                        // make sure we only look at the internal bucket and don't use this.
285                        drop(bucket);
286
287                        let MessageState::WaitingForRestOfFragment(bytes) =
288                            std::mem::replace(&mut self.state, MessageState::WaitingForBucket)
289                        else {
290                            unreachable!();
291                        };
292
293                        // Check there are enough bytes in the fragment to build a header.
294                        if bytes.len() < HEADER_SIZE {
295                            return Err(BucketError::InvalidFragmentedMessage(
296                                "Fragmented message is not large enough to build a bucket.",
297                            ));
298                        }
299
300                        let mut header_bytes = BytesMut::from(&bytes[0..HEADER_SIZE]);
301
302                        let header = BucketHead::<T::Command>::from_bytes(&mut header_bytes);
303
304                        if header.size > header.command.bucket_size_limit() {
305                            return Err(BucketError::BucketExceededMaxSize);
306                        }
307
308                        // Check the fragmented message contains enough bytes to build the message.
309                        if bytes.len().saturating_sub(HEADER_SIZE) < u64_to_usize(header.size) {
310                            return Err(BucketError::InvalidFragmentedMessage(
311                                "Fragmented message does not have enough bytes to fill bucket body",
312                            ));
313                        }
314
315                        #[cfg(feature = "tracing")]
316                        tracing::debug!(
317                            "Received final fragment, combined message command: {:?}.",
318                            header.command
319                        );
320
321                        let message_type = MessageType::from_flags_and_have_to_return(
322                            header.flags,
323                            header.have_to_return_data,
324                        )?;
325
326                        if header.command.is_handshake() {
327                            #[cfg(feature = "tracing")]
328                            tracing::debug!(
329                                "Peer handshake message seen, increasing bucket size limit."
330                            );
331
332                            self.bucket_codec.handshake_message_seen = true;
333                        }
334
335                        return Ok(Some(T::decode_message(
336                            &mut &bytes[HEADER_SIZE..],
337                            message_type,
338                            header.command,
339                        )?));
340                    }
341                }
342            }
343        }
344    }
345}
346
347impl<T: LevinBody> Encoder<LevinMessage<T>> for LevinMessageCodec<T> {
348    type Error = BucketError;
349    fn encode(&mut self, item: LevinMessage<T>, dst: &mut BytesMut) -> Result<(), Self::Error> {
350        match item {
351            LevinMessage::Body(body) => {
352                let mut bucket_builder = BucketBuilder::new(&self.bucket_codec.protocol);
353                body.encode(&mut bucket_builder)?;
354                let bucket = bucket_builder.finish();
355                self.bucket_codec.encode(bucket, dst)
356            }
357            LevinMessage::Bucket(bucket) => self.bucket_codec.encode(bucket, dst),
358            LevinMessage::Dummy(size) => {
359                let bucket = make_dummy_message(&self.bucket_codec.protocol, size);
360                self.bucket_codec.encode(bucket, dst)
361            }
362        }
363    }
364}