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