nurtex_protocol/connection/
writer.rs1use std::io::Read;
2use std::sync::atomic::{AtomicI32, Ordering};
3use std::{fmt::Debug, sync::Arc};
4
5use bytes::{Bytes, BytesMut};
6use flate2::Compression;
7use flate2::bufread::ZlibEncoder;
8use nurtex_codec::types::variable::VarI32;
9use nurtex_encrypt::AesEncryptor;
10use tokio::io::{AsyncWrite, AsyncWriteExt};
11use tokio::net::tcp::OwnedWriteHalf;
12use tokio::sync::Mutex;
13
14use crate::ProtocolPacket;
15use crate::connection::ServersidePacket;
16
17pub struct ConnectionWriter {
19 pub write_stream: Option<OwnedWriteHalf>,
21
22 pub encryptor: Arc<Mutex<Option<AesEncryptor>>>,
24
25 pub compression_threshold: Arc<AtomicI32>,
27}
28
29impl ConnectionWriter {
30 pub async fn write_packet(&mut self, packet: ServersidePacket) -> std::io::Result<()> {
32 let Some(write_half) = &mut self.write_stream else {
33 return Err(std::io::Error::new(std::io::ErrorKind::NotConnected, "write stream not initialized"));
34 };
35
36 let serialized = match packet {
37 ServersidePacket::Handshake(p) => serialize_packet(&p)?,
38 ServersidePacket::Status(p) => serialize_packet(&p)?,
39 ServersidePacket::Login(p) => serialize_packet(&p)?,
40 ServersidePacket::Configuration(p) => serialize_packet(&p)?,
41 ServersidePacket::Play(p) => serialize_packet(&p)?,
42 };
43
44 let compression_threshold = self.compression_threshold.load(Ordering::SeqCst);
45 let mut encryptor_guard = self.encryptor.lock().await;
46
47 write_raw_packet(&serialized, write_half, compression_threshold, &mut *encryptor_guard).await
48 }
49
50 pub async fn shutdown(&mut self) -> std::io::Result<()> {
52 if let Some(write_half) = &mut self.write_stream {
53 write_half.shutdown().await?;
54 }
55
56 Ok(())
57 }
58}
59
60pub async fn write_packet<P, W>(packet: &P, stream: &mut W, compression_threshold: i32, cipher: &mut Option<AesEncryptor>) -> std::io::Result<()>
62where
63 P: ProtocolPacket + Debug,
64 W: AsyncWrite + Unpin + Send,
65{
66 let raw_packet = serialize_packet(packet)?;
67 write_raw_packet(&raw_packet, stream, compression_threshold, cipher).await
68}
69
70pub fn serialize_packet<P>(packet: &P) -> std::io::Result<Bytes>
72where
73 P: ProtocolPacket + Debug,
74{
75 let mut buf = Vec::new();
76
77 (packet.id() as i32)
78 .write_var(&mut buf)
79 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
80
81 packet.write(&mut buf).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
82
83 if buf.len() > 8388608 {
84 return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "packet too large"));
85 }
86
87 Ok(Bytes::from(buf))
88}
89
90pub async fn write_raw_packet<W>(raw_packet: &[u8], stream: &mut W, compression_threshold: i32, cipher: &mut Option<AesEncryptor>) -> std::io::Result<()>
92where
93 W: AsyncWrite + Unpin + Send,
94{
95 let network_packet = encode_to_network_packet(raw_packet, compression_threshold, cipher)?;
96 stream.write_all(&network_packet).await
97}
98
99pub fn encode_to_network_packet(raw_packet: &[u8], compression_threshold: i32, cipher: &mut Option<AesEncryptor>) -> std::io::Result<BytesMut> {
101 let mut buf = BytesMut::new();
102
103 if compression_threshold >= 0 {
104 compression_encoder(raw_packet, compression_threshold, &mut buf)?;
105 } else {
106 buf.extend_from_slice(raw_packet);
107 }
108
109 let mut frame = BytesMut::new();
110 let mut length = Vec::new();
111
112 (buf.len() as i32)
113 .write_var(&mut length)
114 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
115
116 frame.extend_from_slice(&length);
117 frame.extend_from_slice(&buf);
118
119 if let Some(cipher) = cipher {
120 nurtex_encrypt::encrypt_data(cipher, &mut frame);
121 }
122
123 Ok(frame)
124}
125
126pub fn compression_encoder(data: &[u8], compression_threshold: i32, buf: &mut BytesMut) -> std::io::Result<()> {
128 let n = data.len();
129
130 if n < compression_threshold as usize {
131 let mut temp = Vec::new();
132 0.write_var(&mut temp).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
133 buf.extend_from_slice(&temp);
134 buf.extend_from_slice(data);
135 } else {
136 let mut deflater = ZlibEncoder::new(data, Compression::default());
137 let mut compressed_data = Vec::new();
138 deflater.read_to_end(&mut compressed_data)?;
139
140 let mut length = Vec::new();
141
142 (n as i32).write_var(&mut length).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
143
144 buf.extend_from_slice(&length);
145 buf.extend_from_slice(&compressed_data);
146 }
147
148 Ok(())
149}