pb-mapper-protocol 0.5.1

Message framing and authenticated sessions for pb-mapper
Documentation
//! Define message protocols and tools for reading and writing
//! messages
//!
//! The reader and writer traits are `async fn` in a public trait, which cannot
//! state its auto-trait bounds. That is deliberate: these are only ever awaited
//! on the connection task that owns the stream, never sent across one.
#![allow(async_fn_in_trait)]

pub mod buffer;
pub mod command;
pub mod forward;
mod frame_read;
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,
    MsgNetworkWriteBodySnafu, MsgNetworkWriteCheckSumSnafu, MsgNetworkWriteCodecMsgSnafu,
    MsgNetworkWriteCodecTagSnafu, MsgNetworkWriteDatalenSnafu, Result,
};

/// This message protocol contains header and body, and the header
/// includes checksum, datalen,respectively, u32, u32, where datalen
/// represents the length of the body, checksum is used to check the
/// datalen field. This is just the most basic pedestal protocol, in
/// order to solve the sticky packet problem with TCP streams. We
/// can build more advanced communication on top of this protocol, for
/// example, we can use json or other forms of data representation
///
/// ```text
/// ┌─────────────┐
/// │ u32 checksum│
/// │ u32 datalen │
/// └─────────────┘
/// ┌─────────────┐
/// │ Body        │
/// │(Actual Data)│
/// └─────────────┘
/// ```
pub trait MessageReader {
    async fn read_msg(&mut self) -> Result<&'_ [u8]>;
}

pub trait MessageWriter {
    async fn write_msg(&mut self, msg: &[u8]) -> Result<()>;

    /// TODO: Implement this method to fix the encryption zero-copy data corruption bug
    ///
    /// This method should be used for encryption scenarios where:
    /// - Zero-copy performance is needed
    /// - The caller can provide mutable data
    /// - No unsafe transmutation is required
    ///
    /// Default implementation falls back to the immutable version for compatibility.
    /// Encryption implementations should override this for true zero-copy operation.
    async fn write_msg_mut(&mut self, msg: &mut [u8]) -> Result<()> {
        // Default implementation: delegate to immutable version
        // This maintains backward compatibility but doesn't solve the corruption issue
        self.write_msg(msg).await
    }
}

/// Maximum plaintext payload size read from local services.
const MAX_PLAINTEXT_LEN: DataLenType = 8 * 1024 * 1024;
/// AES-GCM tag length (bytes). Keep in sync with ring's tag length.
const CODEC_TAG_LEN: DataLenType = 16;
/// Maximum value of `datalen` to prevent Out of Memory.
/// For encrypted frames, the tag is appended to the payload.
const MAX_MSG_LEN: DataLenType = MAX_PLAINTEXT_LEN + CODEC_TAG_LEN;

// Defined in `pb-mapper-core` so that the checksum and error types can name it
// without depending on this module. Re-exported rather than redeclared: a second
// `pub type` would be a distinct name for the same width, and the two would read
// as unrelated at the crate boundary.
pub use pb_mapper_core::DataLenType;

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_write_network_with_error!(write_checksum, write_u32, MsgNetworkWriteCheckSumSnafu, u32);

gen_write_network_with_error!(write_datalen, write_u32, MsgNetworkWriteDatalenSnafu, u32);

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 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,
    frame: frame_read::FrameRead<8>,
    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(),
            frame: frame_read::FrameRead::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]> {
        self.frame
            .header(self.reader)
            .await
            .context(MsgNetworkReadCheckSumSnafu)?;
        let header = self.frame.header;
        let checksum = u32::from_be_bytes([header[0], header[1], header[2], header[3]]);
        let datalen = u32::from_be_bytes([header[4], header[5], header[6], header[7]]);
        ensure!(
            checksum_matches(datalen, checksum, checksum_key_bytes(&self.checksum_key)),
            MsgDatalenValidateSnafu { datalen, checksum }
        );
        ensure!(
            datalen <= MAX_MSG_LEN,
            MsgDatalenExceededSnafu {
                actual: datalen,
                max: MAX_MSG_LEN
            }
        );
        self.buffer.fixed_resize(datalen as usize);
        self.frame
            .body(self.reader, self.buffer.buffer_mut())
            .await
            .context(MsgNetworkReadBodySnafu)?;
        self.frame.finish();
        Ok(self.buffer.buffer())
    }
}

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,
        }
    }

    /// Bind the length checksum to `key` instead of the process credential.
    /// Isolated relays keep a remote `MSG_HEADER_KEY` while speaking with a
    /// different local administrator key.
    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)
    }
}

/// NOTE: We copy input data before encryption to avoid mutating shared buffers.
/// This trades some performance for correctness until a zero-copy mutable API is added.
pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> {
    writer: &'a mut T,
    encryptor: E,
    /// `None` uses the process `MSG_HEADER_KEY` hash. Isolated relays must set
    /// this to the session key so continuation frames stay decryptable.
    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}"),
    })
}