1use 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 WaitingForHeader,
35 WaitingForBody(BucketHead<C>),
37}
38
39#[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 WaitingForRestOfFragment(Vec<u8>),
176}
177
178#[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 #[cfg(feature = "tracing")]
214 tracing::trace!("Received DUMMY bucket from peer, ignoring.");
215 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 #[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 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 #[cfg(feature = "tracing")]
261 tracing::trace!("Received DUMMY bucket from peer, ignoring.");
262 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 drop(bucket);
286
287 let MessageState::WaitingForRestOfFragment(bytes) =
288 std::mem::replace(&mut self.state, MessageState::WaitingForBucket)
289 else {
290 unreachable!();
291 };
292
293 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 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}