chunkedge_protocol/
encode.rs1use 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#[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 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 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 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 #[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
181pub trait WritePacket {
183 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 fn write_packet_fallible<P>(&mut self, packet: &P) -> anyhow::Result<()>
197 where
198 P: Packet + Encode;
199
200 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#[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 VarInt(0).encode(front)?;
384 }
385
386 Ok(())
387}