1#![allow(async_fn_in_trait)]
8
9pub mod buffer;
10pub mod command;
11pub mod forward;
12mod frame_read;
13pub mod secure;
14use snafu::{ResultExt, ensure};
15use tokio::io::{AsyncReadExt, AsyncWriteExt};
16
17use crate::buffer::{BufferGetter, CommonBuffer, FixedSizeBuffer};
18use pb_mapper_core::checksum::{
19 AesKeyType, get_checksum, get_checksum_for_key, get_msg_header_key, process_checksum_is_ready,
20 valid_checksum, valid_checksum_for_key,
21};
22use pb_mapper_core::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor, Encryptor};
23use pb_mapper_core::error::MsgDatalenExceededSnafu;
24use pb_mapper_core::error::{
25 self, MsgDatalenValidateSnafu, MsgNetworkReadBodySnafu, MsgNetworkReadCheckSumSnafu,
26 MsgNetworkWriteBodySnafu, MsgNetworkWriteCheckSumSnafu, MsgNetworkWriteCodecMsgSnafu,
27 MsgNetworkWriteCodecTagSnafu, MsgNetworkWriteDatalenSnafu, Result,
28};
29
30pub trait MessageReader {
49 async fn read_msg(&mut self) -> Result<&'_ [u8]>;
50}
51
52pub trait MessageWriter {
53 async fn write_msg(&mut self, msg: &[u8]) -> Result<()>;
54
55 async fn write_msg_mut(&mut self, msg: &mut [u8]) -> Result<()> {
65 self.write_msg(msg).await
68 }
69}
70
71const MAX_PLAINTEXT_LEN: DataLenType = 8 * 1024 * 1024;
73const CODEC_TAG_LEN: DataLenType = 16;
75const MAX_MSG_LEN: DataLenType = MAX_PLAINTEXT_LEN + CODEC_TAG_LEN;
78
79pub use pb_mapper_core::DataLenType;
84
85macro_rules! gen_write_network_with_error {
86 ($func_name:ident, $write_method:ident, $error:expr, $input_type:ty) => {
87 #[inline]
88 async fn $func_name<T: AsyncWriteExt + Unpin>(
89 writer: &mut T,
90 data: $input_type,
91 ) -> Result<()> {
92 writer.$write_method(data).await.context($error)
93 }
94 };
95}
96
97gen_write_network_with_error!(write_checksum, write_u32, MsgNetworkWriteCheckSumSnafu, u32);
98
99gen_write_network_with_error!(write_datalen, write_u32, MsgNetworkWriteDatalenSnafu, u32);
100
101gen_write_network_with_error!(write_msg_body, write_all, MsgNetworkWriteBodySnafu, &[u8]);
102
103gen_write_network_with_error!(
104 write_codec_msg,
105 write_all,
106 MsgNetworkWriteCodecMsgSnafu,
107 &[u8]
108);
109
110gen_write_network_with_error!(
111 write_codec_tag,
112 write_all,
113 MsgNetworkWriteCodecTagSnafu,
114 &[u8]
115);
116
117fn checksum_key_bytes(key: &Option<AesKeyType>) -> Option<&[u8]> {
118 key.as_ref().map(|key| key.as_slice())
119}
120
121#[inline]
122fn checksum_matches(datalen: DataLenType, checksum: u32, key: Option<&[u8]>) -> bool {
123 match key {
124 Some(key) => valid_checksum_for_key(datalen, checksum, key),
125 None => valid_checksum(datalen, checksum),
126 }
127}
128
129#[inline]
130fn checksum_for(len: DataLenType, key: Option<&[u8]>) -> Result<u32> {
131 match key {
132 Some(key) => Ok(get_checksum_for_key(len, key)),
133 None => {
134 if !process_checksum_is_ready() {
135 return Err(error::Error::MsgCodec {
136 action: "load configured credential",
137 detail:
138 "`MSG_HEADER_KEY` is required; no insecure default checksum is available"
139 .to_string(),
140 });
141 }
142 Ok(get_checksum(len))
143 }
144 }
145}
146
147#[inline]
148async fn set_msg_len<T: AsyncWriteExt + Unpin>(
149 writer: &mut T,
150 len: DataLenType,
151 checksum_key: Option<&[u8]>,
152) -> Result<()> {
153 write_checksum(writer, checksum_for(len, checksum_key)?).await?;
154 write_datalen(writer, len).await
155}
156
157pub struct NormalMessageReader<'a, T: AsyncReadExt + Unpin> {
158 reader: &'a mut T,
159 buffer: CommonBuffer,
160 frame: frame_read::FrameRead<8>,
161 checksum_key: Option<AesKeyType>,
162}
163
164impl<'a, T: AsyncReadExt + Unpin> NormalMessageReader<'a, T> {
165 pub fn new(reader: &'a mut T) -> Self {
166 Self {
167 reader,
168 buffer: CommonBuffer::new(),
169 frame: frame_read::FrameRead::new(),
170 checksum_key: None,
171 }
172 }
173
174 pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
175 self.checksum_key = Some(key);
176 self
177 }
178
179 async fn read_msg_inner(&mut self) -> Result<&'_ [u8]> {
180 self.frame
181 .header(self.reader)
182 .await
183 .context(MsgNetworkReadCheckSumSnafu)?;
184 let header = self.frame.header;
185 let checksum = u32::from_be_bytes([header[0], header[1], header[2], header[3]]);
186 let datalen = u32::from_be_bytes([header[4], header[5], header[6], header[7]]);
187 ensure!(
188 checksum_matches(datalen, checksum, checksum_key_bytes(&self.checksum_key)),
189 MsgDatalenValidateSnafu { datalen, checksum }
190 );
191 ensure!(
192 datalen <= MAX_MSG_LEN,
193 MsgDatalenExceededSnafu {
194 actual: datalen,
195 max: MAX_MSG_LEN
196 }
197 );
198 self.buffer.fixed_resize(datalen as usize);
199 self.frame
200 .body(self.reader, self.buffer.buffer_mut())
201 .await
202 .context(MsgNetworkReadBodySnafu)?;
203 self.frame.finish();
204 Ok(self.buffer.buffer())
205 }
206}
207
208impl<'a, T: AsyncReadExt + Unpin> MessageReader for NormalMessageReader<'a, T> {
209 async fn read_msg(&mut self) -> Result<&'_ [u8]> {
210 self.read_msg_inner().await
211 }
212}
213
214pub struct NormalMessageWriter<'a, T: AsyncWriteExt> {
215 writer: &'a mut T,
216 checksum_key: Option<AesKeyType>,
217}
218
219impl<'a, T: AsyncWriteExt + Unpin> NormalMessageWriter<'a, T> {
220 pub fn new(writer: &'a mut T) -> Self {
221 Self {
222 writer,
223 checksum_key: None,
224 }
225 }
226
227 pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
228 self.checksum_key = Some(key);
229 self
230 }
231
232 async fn write_msg_inner(&mut self, msg: &[u8]) -> Result<()> {
233 set_msg_len(
234 &mut self.writer,
235 msg.len() as u32,
236 checksum_key_bytes(&self.checksum_key),
237 )
238 .await?;
239
240 write_msg_body(&mut self.writer, msg).await
241 }
242}
243
244impl<'a, T: AsyncWriteExt + Unpin> MessageWriter for NormalMessageWriter<'a, T> {
245 async fn write_msg(&mut self, msg: &[u8]) -> Result<()> {
246 self.write_msg_inner(msg).await
247 }
248}
249
250pub struct CodecMessageReader<'a, T: AsyncReadExt + Unpin, D: Decryptor> {
251 reader: NormalMessageReader<'a, T>,
252 decryptor: D,
253}
254
255impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> CodecMessageReader<'a, T, D> {
256 pub fn new(reader: &'a mut T, decryptor: D) -> Self {
257 Self {
258 reader: NormalMessageReader::new(reader),
259 decryptor,
260 }
261 }
262
263 pub fn for_session_key(reader: &'a mut T, decryptor: D, key: AesKeyType) -> Self {
267 Self::new(reader, decryptor).with_checksum_key(key)
268 }
269
270 pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
271 self.reader.checksum_key = Some(key);
272 self
273 }
274}
275
276impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> MessageReader for CodecMessageReader<'a, T, D> {
277 async fn read_msg(&mut self) -> Result<&'_ [u8]> {
278 let n = self.reader.read_msg().await?.len();
279 let v = self
280 .decryptor
281 .decrypt(&mut self.reader.buffer.buffer_mut()[..n])
282 .map_err(|e| error::Error::MsgCodec {
283 action: "decrypt",
284 detail: format!("got {e} when we read msg"),
285 })?;
286 Ok(v)
287 }
288}
289
290pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> {
293 writer: &'a mut T,
294 encryptor: E,
295 checksum_key: Option<AesKeyType>,
298}
299
300impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> CodecMessageWriter<'a, T, E> {
301 pub fn new(writer: &'a mut T, encryptor: E) -> Self {
302 Self {
303 writer,
304 encryptor,
305 checksum_key: None,
306 }
307 }
308
309 pub fn for_session_key(writer: &'a mut T, encryptor: E, key: AesKeyType) -> Self {
310 Self::new(writer, encryptor).with_checksum_key(key)
311 }
312
313 pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
314 self.checksum_key = Some(key);
315 self
316 }
317
318 pub async fn shutdown(&mut self) -> std::io::Result<()> {
319 self.writer.shutdown().await
320 }
321}
322
323impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> MessageWriter for CodecMessageWriter<'a, T, E> {
324 async fn write_msg(&mut self, msg: &[u8]) -> Result<()> {
325 let mut buf = msg.to_vec();
326 let tag = self
327 .encryptor
328 .encrypt(&mut buf)
329 .map_err(|e| error::Error::MsgCodec {
330 action: "encrypt",
331 detail: format!("got {e} when we read msg"),
332 })?;
333 let msg_len = (buf.len() + tag.as_ref().len()) as DataLenType;
334
335 set_msg_len(self.writer, msg_len, checksum_key_bytes(&self.checksum_key)).await?;
336 write_codec_msg(self.writer, &buf).await?;
337 write_codec_tag(self.writer, tag.as_ref()).await
338 }
339}
340
341#[inline]
342pub fn get_header_msg_reader<T: AsyncReadExt + Unpin>(
343 reader: &mut T,
344) -> Result<CodecMessageReader<'_, T, Aes256GcmDeCodec>> {
345 Ok(CodecMessageReader::new(reader, get_default_decodec()?))
346}
347
348#[inline]
349pub fn get_header_msg_writer<T: AsyncWriteExt + Unpin>(
350 writer: &mut T,
351) -> Result<CodecMessageWriter<'_, T, Aes256GcmEnCodec>> {
352 Ok(CodecMessageWriter::new(writer, get_default_encodec()?))
353}
354
355#[inline]
356pub fn get_default_encodec() -> Result<Aes256GcmEnCodec> {
357 let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec {
358 action: "load configured credential",
359 detail,
360 })?;
361 Aes256GcmEnCodec::try_new(&key).map_err(|e| error::Error::MsgCodec {
362 action: "create default encodec",
363 detail: format!("{e}"),
364 })
365}
366
367#[inline]
368pub fn get_default_decodec() -> Result<Aes256GcmDeCodec> {
369 let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec {
370 action: "load configured credential",
371 detail,
372 })?;
373 Aes256GcmDeCodec::try_new(&key).map_err(|e| error::Error::MsgCodec {
374 action: "create default decodec",
375 detail: format!("{e}"),
376 })
377}
378
379#[inline]
380pub fn get_encodec(key: &[u8]) -> Result<Aes256GcmEnCodec> {
381 Aes256GcmEnCodec::try_new(key).map_err(|e| error::Error::MsgCodec {
382 action: "create encodec",
383 detail: format!("{e}"),
384 })
385}
386
387#[inline]
388pub fn get_decodec(key: &[u8]) -> Result<Aes256GcmDeCodec> {
389 Aes256GcmDeCodec::try_new(key).map_err(|e| error::Error::MsgCodec {
390 action: "create decodec",
391 detail: format!("{e}"),
392 })
393}