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;
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
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_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    /// Bind the length checksum to `key` instead of the process credential.
264    /// Isolated relays keep a remote `MSG_HEADER_KEY` while speaking with a
265    /// different local administrator key.
266    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
290/// NOTE: We copy input data before encryption to avoid mutating shared buffers.
291/// This trades some performance for correctness until a zero-copy mutable API is added.
292pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> {
293    writer: &'a mut T,
294    encryptor: E,
295    /// `None` uses the process `MSG_HEADER_KEY` hash. Isolated relays must set
296    /// this to the session key so continuation frames stay decryptable.
297    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}