#![allow(async_fn_in_trait)]
pub mod buffer;
pub mod command;
pub mod forward;
pub mod secure;
use snafu::{ResultExt, ensure};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::buffer::{BufferGetter, CommonBuffer, FixedSizeBuffer};
use pb_mapper_core::checksum::{
AesKeyType, get_checksum, get_checksum_for_key, get_msg_header_key, process_checksum_is_ready,
valid_checksum, valid_checksum_for_key,
};
use pb_mapper_core::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor, Encryptor};
use pb_mapper_core::error::MsgDatalenExceededSnafu;
use pb_mapper_core::error::{
self, MsgDatalenValidateSnafu, MsgNetworkReadBodySnafu, MsgNetworkReadCheckSumSnafu,
MsgNetworkReadDatalenSnafu, MsgNetworkWriteBodySnafu, MsgNetworkWriteCheckSumSnafu,
MsgNetworkWriteCodecMsgSnafu, MsgNetworkWriteCodecTagSnafu, MsgNetworkWriteDatalenSnafu,
Result,
};
pub trait MessageReader {
async fn read_msg(&mut self) -> Result<&'_ [u8]>;
}
pub trait MessageWriter {
async fn write_msg(&mut self, msg: &[u8]) -> Result<()>;
async fn write_msg_mut(&mut self, msg: &mut [u8]) -> Result<()> {
self.write_msg(msg).await
}
}
const MAX_PLAINTEXT_LEN: DataLenType = 8 * 1024 * 1024;
const CODEC_TAG_LEN: DataLenType = 16;
const MAX_MSG_LEN: DataLenType = MAX_PLAINTEXT_LEN + CODEC_TAG_LEN;
pub use pb_mapper_core::DataLenType;
macro_rules! gen_read_network_with_error {
($func_name:ident, $read_method:ident, $error:expr, $return_ty:ty) => {
#[inline]
async fn $func_name<T: AsyncReadExt + Unpin>(reader: &mut T) -> Result<$return_ty> {
reader.$read_method().await.context($error)
}
};
($func_name:ident, $read_method:ident, $error:expr, $input_type:ty, $return_type:ty) => {
#[inline]
async fn $func_name<T: AsyncReadExt + Unpin>(
reader: &mut T,
input_type: $input_type,
) -> Result<$return_type> {
reader.$read_method(input_type).await.context($error)
}
};
}
macro_rules! gen_write_network_with_error {
($func_name:ident, $write_method:ident, $error:expr, $input_type:ty) => {
#[inline]
async fn $func_name<T: AsyncWriteExt + Unpin>(
writer: &mut T,
data: $input_type,
) -> Result<()> {
writer.$write_method(data).await.context($error)
}
};
}
gen_read_network_with_error!(read_checksum, read_u32, MsgNetworkReadCheckSumSnafu, u32);
gen_read_network_with_error!(read_datalen, read_u32, MsgNetworkReadDatalenSnafu, u32);
gen_write_network_with_error!(write_checksum, write_u32, MsgNetworkWriteCheckSumSnafu, u32);
gen_write_network_with_error!(write_datalen, write_u32, MsgNetworkWriteDatalenSnafu, u32);
gen_read_network_with_error!(
read_msg_body,
read_exact,
MsgNetworkReadBodySnafu,
&mut [u8],
usize
);
gen_write_network_with_error!(write_msg_body, write_all, MsgNetworkWriteBodySnafu, &[u8]);
gen_write_network_with_error!(
write_codec_msg,
write_all,
MsgNetworkWriteCodecMsgSnafu,
&[u8]
);
gen_write_network_with_error!(
write_codec_tag,
write_all,
MsgNetworkWriteCodecTagSnafu,
&[u8]
);
fn checksum_key_bytes(key: &Option<AesKeyType>) -> Option<&[u8]> {
key.as_ref().map(|key| key.as_slice())
}
#[inline]
fn checksum_matches(datalen: DataLenType, checksum: u32, key: Option<&[u8]>) -> bool {
match key {
Some(key) => valid_checksum_for_key(datalen, checksum, key),
None => valid_checksum(datalen, checksum),
}
}
#[inline]
fn checksum_for(len: DataLenType, key: Option<&[u8]>) -> Result<u32> {
match key {
Some(key) => Ok(get_checksum_for_key(len, key)),
None => {
if !process_checksum_is_ready() {
return Err(error::Error::MsgCodec {
action: "load configured credential",
detail:
"`MSG_HEADER_KEY` is required; no insecure default checksum is available"
.to_string(),
});
}
Ok(get_checksum(len))
}
}
}
#[inline]
async fn get_msg_len<T: AsyncReadExt + Unpin>(
reader: &mut T,
checksum_key: Option<&[u8]>,
) -> Result<DataLenType> {
let checksum = read_checksum(reader).await?;
let datalen = read_datalen(reader).await?;
if checksum_matches(datalen, checksum, checksum_key) {
ensure!(
datalen <= MAX_MSG_LEN,
MsgDatalenExceededSnafu {
actual: datalen,
max: MAX_MSG_LEN
}
);
Ok(datalen)
} else {
MsgDatalenValidateSnafu { datalen, checksum }.fail()?
}
}
#[inline]
async fn set_msg_len<T: AsyncWriteExt + Unpin>(
writer: &mut T,
len: DataLenType,
checksum_key: Option<&[u8]>,
) -> Result<()> {
write_checksum(writer, checksum_for(len, checksum_key)?).await?;
write_datalen(writer, len).await
}
pub struct NormalMessageReader<'a, T: AsyncReadExt + Unpin> {
reader: &'a mut T,
buffer: CommonBuffer,
checksum_key: Option<AesKeyType>,
}
impl<'a, T: AsyncReadExt + Unpin> NormalMessageReader<'a, T> {
pub fn new(reader: &'a mut T) -> Self {
Self {
reader,
buffer: CommonBuffer::new(),
checksum_key: None,
}
}
pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
self.checksum_key = Some(key);
self
}
async fn read_msg_inner(&mut self) -> Result<&'_ [u8]> {
let datalen = get_msg_len(&mut self.reader, checksum_key_bytes(&self.checksum_key)).await?;
self.buffer.fixed_resize(datalen as usize);
let n = read_msg_body(&mut self.reader, self.buffer.buffer_mut()).await?;
Ok(&self.buffer.buffer()[0..n])
}
}
impl<'a, T: AsyncReadExt + Unpin> MessageReader for NormalMessageReader<'a, T> {
async fn read_msg(&mut self) -> Result<&'_ [u8]> {
self.read_msg_inner().await
}
}
pub struct NormalMessageWriter<'a, T: AsyncWriteExt> {
writer: &'a mut T,
checksum_key: Option<AesKeyType>,
}
impl<'a, T: AsyncWriteExt + Unpin> NormalMessageWriter<'a, T> {
pub fn new(writer: &'a mut T) -> Self {
Self {
writer,
checksum_key: None,
}
}
pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
self.checksum_key = Some(key);
self
}
async fn write_msg_inner(&mut self, msg: &[u8]) -> Result<()> {
set_msg_len(
&mut self.writer,
msg.len() as u32,
checksum_key_bytes(&self.checksum_key),
)
.await?;
write_msg_body(&mut self.writer, msg).await
}
}
impl<'a, T: AsyncWriteExt + Unpin> MessageWriter for NormalMessageWriter<'a, T> {
async fn write_msg(&mut self, msg: &[u8]) -> Result<()> {
self.write_msg_inner(msg).await
}
}
pub struct CodecMessageReader<'a, T: AsyncReadExt + Unpin, D: Decryptor> {
reader: NormalMessageReader<'a, T>,
decryptor: D,
}
impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> CodecMessageReader<'a, T, D> {
pub fn new(reader: &'a mut T, decryptor: D) -> Self {
Self {
reader: NormalMessageReader::new(reader),
decryptor,
}
}
pub fn for_session_key(reader: &'a mut T, decryptor: D, key: AesKeyType) -> Self {
Self::new(reader, decryptor).with_checksum_key(key)
}
pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
self.reader.checksum_key = Some(key);
self
}
}
impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> MessageReader for CodecMessageReader<'a, T, D> {
async fn read_msg(&mut self) -> Result<&'_ [u8]> {
let n = self.reader.read_msg().await?.len();
let v = self
.decryptor
.decrypt(&mut self.reader.buffer.buffer_mut()[..n])
.map_err(|e| error::Error::MsgCodec {
action: "decrypt",
detail: format!("got {e} when we read msg"),
})?;
Ok(v)
}
}
pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> {
writer: &'a mut T,
encryptor: E,
checksum_key: Option<AesKeyType>,
}
impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> CodecMessageWriter<'a, T, E> {
pub fn new(writer: &'a mut T, encryptor: E) -> Self {
Self {
writer,
encryptor,
checksum_key: None,
}
}
pub fn for_session_key(writer: &'a mut T, encryptor: E, key: AesKeyType) -> Self {
Self::new(writer, encryptor).with_checksum_key(key)
}
pub fn with_checksum_key(mut self, key: AesKeyType) -> Self {
self.checksum_key = Some(key);
self
}
pub async fn shutdown(&mut self) -> std::io::Result<()> {
self.writer.shutdown().await
}
}
impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> MessageWriter for CodecMessageWriter<'a, T, E> {
async fn write_msg(&mut self, msg: &[u8]) -> Result<()> {
let mut buf = msg.to_vec();
let tag = self
.encryptor
.encrypt(&mut buf)
.map_err(|e| error::Error::MsgCodec {
action: "encrypt",
detail: format!("got {e} when we read msg"),
})?;
let msg_len = (buf.len() + tag.as_ref().len()) as DataLenType;
set_msg_len(self.writer, msg_len, checksum_key_bytes(&self.checksum_key)).await?;
write_codec_msg(self.writer, &buf).await?;
write_codec_tag(self.writer, tag.as_ref()).await
}
}
#[inline]
pub fn get_header_msg_reader<T: AsyncReadExt + Unpin>(
reader: &mut T,
) -> Result<CodecMessageReader<'_, T, Aes256GcmDeCodec>> {
Ok(CodecMessageReader::new(reader, get_default_decodec()?))
}
#[inline]
pub fn get_header_msg_writer<T: AsyncWriteExt + Unpin>(
writer: &mut T,
) -> Result<CodecMessageWriter<'_, T, Aes256GcmEnCodec>> {
Ok(CodecMessageWriter::new(writer, get_default_encodec()?))
}
#[inline]
pub fn get_default_encodec() -> Result<Aes256GcmEnCodec> {
let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec {
action: "load configured credential",
detail,
})?;
Aes256GcmEnCodec::try_new(&key).map_err(|e| error::Error::MsgCodec {
action: "create default encodec",
detail: format!("{e}"),
})
}
#[inline]
pub fn get_default_decodec() -> Result<Aes256GcmDeCodec> {
let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec {
action: "load configured credential",
detail,
})?;
Aes256GcmDeCodec::try_new(&key).map_err(|e| error::Error::MsgCodec {
action: "create default decodec",
detail: format!("{e}"),
})
}
#[inline]
pub fn get_encodec(key: &[u8]) -> Result<Aes256GcmEnCodec> {
Aes256GcmEnCodec::try_new(key).map_err(|e| error::Error::MsgCodec {
action: "create encodec",
detail: format!("{e}"),
})
}
#[inline]
pub fn get_decodec(key: &[u8]) -> Result<Aes256GcmDeCodec> {
Aes256GcmDeCodec::try_new(key).map_err(|e| error::Error::MsgCodec {
action: "create decodec",
detail: format!("{e}"),
})
}