Skip to main content

chunkedge_protocol/
encode.rs

1use std::io::Write;
2
3#[cfg(feature = "encryption")]
4use aes::cipher::KeyIvInit;
5use anyhow::ensure;
6use bytes::{BufMut, BytesMut};
7use chunkedge_binary::{Encode, VarInt};
8use tracing::warn;
9
10use crate::{CompressionThreshold, MAX_PACKET_SIZE, Packet};
11
12/// The AES block cipher with a 128 bit key, using the CFB-8 mode of
13/// operation.
14#[cfg(feature = "encryption")]
15type Cipher = cfb8::Encryptor<aes::Aes128>;
16
17#[derive(Default)]
18pub struct PacketEncoder {
19    buf: BytesMut,
20    #[cfg(feature = "compression")]
21    compress_buf: Vec<u8>,
22    #[cfg(feature = "compression")]
23    threshold: CompressionThreshold,
24    #[cfg(feature = "encryption")]
25    cipher: Option<Cipher>,
26}
27
28impl PacketEncoder {
29    pub fn new() -> Self {
30        Self::default()
31    }
32
33    #[inline]
34    pub fn append_bytes(&mut self, bytes: &[u8]) {
35        self.buf.extend_from_slice(bytes)
36    }
37
38    pub fn prepend_packet<P>(&mut self, pkt: &P) -> anyhow::Result<()>
39    where
40        P: Packet + Encode,
41    {
42        let start_len = self.buf.len();
43        self.append_packet(pkt)?;
44
45        let end_len = self.buf.len();
46        let total_packet_len = end_len - start_len;
47
48        // 1) Move everything back by the length of the packet.
49        // 2) Move the packet to the new space at the front.
50        // 3) Truncate the old packet away.
51        self.buf.put_bytes(0, total_packet_len);
52        self.buf.copy_within(..end_len, total_packet_len);
53        self.buf.copy_within(total_packet_len + start_len.., 0);
54        self.buf.truncate(end_len);
55
56        Ok(())
57    }
58
59    #[allow(clippy::needless_borrows_for_generic_args)]
60    pub fn append_packet<P>(&mut self, pkt: &P) -> anyhow::Result<()>
61    where
62        P: Packet + Encode,
63    {
64        let start_len = self.buf.len();
65
66        pkt.encode_with_id((&mut self.buf).writer())?;
67
68        let data_len = self.buf.len() - start_len;
69
70        #[cfg(feature = "compression")]
71        if self.threshold.0 >= 0 {
72            use std::io::Read;
73
74            use flate2::Compression;
75            use flate2::bufread::ZlibEncoder;
76
77            if data_len >= self.threshold.0 as usize {
78                let mut z = ZlibEncoder::new(&self.buf[start_len..], Compression::new(4));
79
80                self.compress_buf.clear();
81
82                let data_len_size = VarInt(data_len as i32).written_size();
83
84                let packet_len = data_len_size + z.read_to_end(&mut self.compress_buf)?;
85
86                ensure!(
87                    packet_len <= MAX_PACKET_SIZE as usize,
88                    "packet exceeds maximum length"
89                );
90
91                drop(z);
92
93                self.buf.truncate(start_len);
94
95                let mut writer = (&mut self.buf).writer();
96
97                VarInt(packet_len as i32).encode(&mut writer)?;
98                VarInt(data_len as i32).encode(&mut writer)?;
99                self.buf.extend_from_slice(&self.compress_buf);
100            } else {
101                let data_len_size = 1;
102                let packet_len = data_len_size + data_len;
103
104                ensure!(
105                    packet_len <= MAX_PACKET_SIZE as usize,
106                    "packet exceeds maximum length"
107                );
108
109                let packet_len_size = VarInt(packet_len as i32).written_size();
110
111                let data_prefix_len = packet_len_size + data_len_size;
112
113                self.buf.put_bytes(0, data_prefix_len);
114                self.buf
115                    .copy_within(start_len..start_len + data_len, start_len + data_prefix_len);
116
117                let mut front = &mut self.buf[start_len..];
118
119                VarInt(packet_len as i32).encode(&mut front)?;
120                // Zero for no compression on this packet.
121                VarInt(0).encode(front)?;
122            }
123
124            return Ok(());
125        }
126
127        let packet_len = data_len;
128
129        ensure!(
130            packet_len <= MAX_PACKET_SIZE as usize,
131            "packet exceeds maximum length"
132        );
133
134        let packet_len_size = VarInt(packet_len as i32).written_size();
135
136        self.buf.put_bytes(0, packet_len_size);
137        self.buf
138            .copy_within(start_len..start_len + data_len, start_len + packet_len_size);
139
140        let front = &mut self.buf[start_len..];
141        VarInt(packet_len as i32).encode(front)?;
142
143        Ok(())
144    }
145
146    /// Takes all the packets written so far and encrypts them if encryption is
147    /// enabled.
148    pub fn take(&mut self) -> BytesMut {
149        #[cfg(feature = "encryption")]
150        if let Some(cipher) = &mut self.cipher {
151            cipher.encrypt(&mut self.buf);
152        }
153
154        self.buf.split()
155    }
156
157    pub fn clear(&mut self) {
158        self.buf.clear();
159    }
160
161    #[cfg(feature = "compression")]
162    pub fn set_compression(&mut self, threshold: CompressionThreshold) {
163        self.threshold = threshold;
164    }
165
166    /// Initializes the cipher with the given key. All future packets **and any
167    /// that have not been [taken] yet** are encrypted.
168    ///
169    /// [taken]: Self::take
170    ///
171    /// # Panics
172    ///
173    /// Panics if encryption is already enabled.
174    #[cfg(feature = "encryption")]
175    pub fn enable_encryption(&mut self, key: &[u8; 16]) {
176        assert!(self.cipher.is_none(), "encryption is already enabled");
177        self.cipher = Some(Cipher::new_from_slices(key, key).expect("invalid key"));
178    }
179}
180
181/// Types that can have packets written to them.
182pub trait WritePacket {
183    /// Writes a packet to this object. Encoding errors are typically logged and
184    /// discarded.
185    fn write_packet<P>(&mut self, packet: &P)
186    where
187        P: Packet + Encode,
188    {
189        if let Err(e) = self.write_packet_fallible(packet) {
190            warn!("failed to write packet '{}': {e:#}", P::NAME);
191        }
192    }
193
194    /// Writes a packet to this object. The result of encoding the packet is
195    /// returned.
196    fn write_packet_fallible<P>(&mut self, packet: &P) -> anyhow::Result<()>
197    where
198        P: Packet + Encode;
199
200    /// Copies raw packet data directly into this object. Don't use this unless
201    /// you know what you're doing.
202    fn write_packet_bytes(&mut self, bytes: &[u8]);
203}
204
205impl<W: WritePacket> WritePacket for &mut W {
206    fn write_packet_fallible<P>(&mut self, packet: &P) -> anyhow::Result<()>
207    where
208        P: Packet + Encode,
209    {
210        (*self).write_packet_fallible(packet)
211    }
212
213    fn write_packet_bytes(&mut self, bytes: &[u8]) {
214        (*self).write_packet_bytes(bytes)
215    }
216}
217
218impl<T: WritePacket> WritePacket for bevy_ecs::world::Mut<'_, T> {
219    fn write_packet_fallible<P>(&mut self, packet: &P) -> anyhow::Result<()>
220    where
221        P: Packet + Encode,
222    {
223        self.as_mut().write_packet_fallible(packet)
224    }
225
226    fn write_packet_bytes(&mut self, bytes: &[u8]) {
227        self.as_mut().write_packet_bytes(bytes)
228    }
229}
230
231/// An implementor of [`WritePacket`] backed by a `Vec` mutable reference.
232///
233/// Packets are written by appending to the contained vec. If an error occurs
234/// while writing, the written bytes are truncated away.
235#[derive(Debug)]
236pub struct PacketWriter<'a> {
237    pub buf: &'a mut Vec<u8>,
238    pub threshold: CompressionThreshold,
239}
240
241impl<'a> PacketWriter<'a> {
242    pub fn new(buf: &'a mut Vec<u8>, threshold: CompressionThreshold) -> Self {
243        Self { buf, threshold }
244    }
245}
246
247impl WritePacket for PacketWriter<'_> {
248    #[cfg_attr(not(feature = "compression"), track_caller)]
249    fn write_packet_fallible<P>(&mut self, pkt: &P) -> anyhow::Result<()>
250    where
251        P: Packet + Encode,
252    {
253        let start = self.buf.len();
254
255        let res;
256
257        if self.threshold.0 >= 0 {
258            #[cfg(feature = "compression")]
259            {
260                res = encode_packet_compressed(self.buf, pkt, self.threshold.0 as u32);
261            }
262
263            #[cfg(not(feature = "compression"))]
264            {
265                panic!("\"compression\" feature must be enabled to write compressed packets");
266            }
267        } else {
268            res = encode_packet(self.buf, pkt)
269        };
270
271        if res.is_err() {
272            self.buf.truncate(start);
273        }
274
275        res
276    }
277
278    fn write_packet_bytes(&mut self, bytes: &[u8]) {
279        if let Err(e) = self.buf.write_all(bytes) {
280            warn!("failed to write packet bytes: {e:#}");
281        }
282    }
283}
284
285impl WritePacket for PacketEncoder {
286    fn write_packet_fallible<P>(&mut self, packet: &P) -> anyhow::Result<()>
287    where
288        P: Packet + Encode,
289    {
290        self.append_packet(packet)
291    }
292
293    fn write_packet_bytes(&mut self, bytes: &[u8]) {
294        self.append_bytes(bytes)
295    }
296}
297
298fn encode_packet<P>(buf: &mut Vec<u8>, pkt: &P) -> anyhow::Result<()>
299where
300    P: Packet + Encode,
301{
302    let start_len = buf.len();
303
304    pkt.encode_with_id(&mut *buf)?;
305
306    let packet_len = buf.len() - start_len;
307
308    ensure!(
309        packet_len <= MAX_PACKET_SIZE as usize,
310        "packet exceeds maximum length"
311    );
312
313    let packet_len_size = VarInt(packet_len as i32).written_size();
314
315    buf.put_bytes(0, packet_len_size);
316    buf.copy_within(
317        start_len..start_len + packet_len,
318        start_len + packet_len_size,
319    );
320
321    let front = &mut buf[start_len..];
322    VarInt(packet_len as i32).encode(front)?;
323
324    Ok(())
325}
326
327#[cfg(feature = "compression")]
328#[allow(clippy::needless_borrows_for_generic_args)]
329fn encode_packet_compressed<P>(buf: &mut Vec<u8>, pkt: &P, threshold: u32) -> anyhow::Result<()>
330where
331    P: Packet + Encode,
332{
333    use std::io::Read;
334
335    use flate2::Compression;
336    use flate2::bufread::ZlibEncoder;
337
338    let start_len = buf.len();
339
340    pkt.encode_with_id(&mut *buf)?;
341
342    let data_len = buf.len() - start_len;
343
344    if data_len >= threshold as usize {
345        let mut z = ZlibEncoder::new(&buf[start_len..], Compression::new(4));
346
347        let mut scratch = vec![];
348
349        let packet_len = VarInt(data_len as i32).written_size() + z.read_to_end(&mut scratch)?;
350
351        ensure!(
352            packet_len <= MAX_PACKET_SIZE as usize,
353            "packet exceeds maximum length"
354        );
355
356        drop(z);
357
358        buf.truncate(start_len);
359
360        VarInt(packet_len as i32).encode(&mut *buf)?;
361        VarInt(data_len as i32).encode(&mut *buf)?;
362        buf.extend_from_slice(&scratch);
363    } else {
364        let data_len_size = 1;
365        let packet_len = data_len_size + data_len;
366
367        ensure!(
368            packet_len <= MAX_PACKET_SIZE as usize,
369            "packet exceeds maximum length"
370        );
371
372        let packet_len_size = VarInt(packet_len as i32).written_size();
373
374        let data_prefix_len = packet_len_size + data_len_size;
375
376        buf.put_bytes(0, data_prefix_len);
377        buf.copy_within(start_len..start_len + data_len, start_len + data_prefix_len);
378
379        let mut front = &mut buf[start_len..];
380
381        VarInt(packet_len as i32).encode(&mut front)?;
382        // Zero for no compression on this packet.
383        VarInt(0).encode(front)?;
384    }
385
386    Ok(())
387}