Skip to main content

pb_mapper_protocol/
lib.rs

1//! Define message protocols and tools for reading and writing
2//! messages
3//!
4//! The reader and writer traits are `async fn` in a public trait, which cannot
5//! state its auto-trait bounds. That is deliberate: these are only ever awaited
6//! on the connection task that owns the stream, never sent across one.
7#![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
30/// This message protocol contains header and body, and the header
31/// includes checksum, datalen,respectively, u32, u32, where datalen
32/// represents the length of the body, checksum is used to check the
33/// datalen field. This is just the most basic pedestal protocol, in
34/// order to solve the sticky packet problem with TCP streams. We
35/// can build more advanced communication on top of this protocol, for
36/// example, we can use json or other forms of data representation
37///
38/// ```text
39/// ┌─────────────┐
40/// │ u32 checksum│
41/// │ u32 datalen │
42/// └─────────────┘
43/// ┌─────────────┐
44/// │ Body        │
45/// │(Actual Data)│
46/// └─────────────┘
47/// ```
48pub 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    /// TODO: Implement this method to fix the encryption zero-copy data corruption bug
56    ///
57    /// This method should be used for encryption scenarios where:
58    /// - Zero-copy performance is needed
59    /// - The caller can provide mutable data
60    /// - No unsafe transmutation is required
61    ///
62    /// Default implementation falls back to the immutable version for compatibility.
63    /// Encryption implementations should override this for true zero-copy operation.
64    async fn write_msg_mut(&mut self, msg: &mut [u8]) -> Result<()> {
65        // Default implementation: delegate to immutable version
66        // This maintains backward compatibility but doesn't solve the corruption issue
67        self.write_msg(msg).await
68    }
69}
70
71/// Maximum plaintext payload size read from local services.
72const MAX_PLAINTEXT_LEN: DataLenType = 8 * 1024 * 1024;
73/// AES-GCM tag length (bytes). Keep in sync with ring's tag length.
74const CODEC_TAG_LEN: DataLenType = 16;
75/// Maximum value of `datalen` to prevent Out of Memory.
76/// For encrypted frames, the tag is appended to the payload.
77const MAX_MSG_LEN: DataLenType = MAX_PLAINTEXT_LEN + CODEC_TAG_LEN;
78
79// Defined in `pb-mapper-core` so that the checksum and error types can name it
80// without depending on this module. Re-exported rather than redeclared: a second
81// `pub type` would be a distinct name for the same width, and the two would read
82// as unrelated at the crate boundary.
83pub 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    /// Bind the length checksum to `key` instead of the process credential.
292    /// Isolated relays keep a remote `MSG_HEADER_KEY` while speaking with a
293    /// different local administrator key.
294    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
318/// NOTE: We copy input data before encryption to avoid mutating shared buffers.
319/// This trades some performance for correctness until a zero-copy mutable API is added.
320pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> {
321    writer: &'a mut T,
322    encryptor: E,
323    /// `None` uses the process `MSG_HEADER_KEY` hash. Isolated relays must set
324    /// this to the session key so continuation frames stay decryptable.
325    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}