//! Length-prefixed framing of postcard-encoded values. //! //! Every frame on the wire is laid out as: //! //! ```text //! 4 bytes: payload length, unsigned big-endian //! N bytes: postcard payload //! ``` use serde::Serialize; use serde::de::DeserializeOwned; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use crate::error::{NetworkError, Result}; /// Encodes a value with postcard, enforcing `max_size` on the encoded bytes. pub(crate) fn encode(value: &T, max_size: usize) -> Result> { let bytes = postcard::to_stdvec(value) .map_err(|err| NetworkError::Serialization(format!("postcard encoding failed: {err}")))?; if bytes.len() > max_size { return Err(NetworkError::MessageTooLarge); } Ok(bytes) } /// Decodes a postcard-encoded value. pub(crate) fn decode(bytes: &[u8]) -> Result { postcard::from_bytes(bytes) .map_err(|err| NetworkError::Serialization(format!("postcard decoding failed: {err}"))) } /// Writes one length-prefixed frame containing the postcard encoding of /// `value`. pub(crate) async fn write_frame(writer: &mut W, value: &T, max_size: usize) -> Result<()> where W: AsyncWrite + Unpin, T: Serialize, { let payload = encode(value, max_size)?; // `max_size` is validated to fit into u32 by the configuration. let len = payload.len() as u32; writer .write_all(&len.to_be_bytes()) .await .map_err(|err| NetworkError::Transport(format!("failed to write frame header: {err}")))?; writer .write_all(&payload) .await .map_err(|err| NetworkError::Transport(format!("failed to write frame payload: {err}")))?; Ok(()) } /// Reads one length-prefixed frame and decodes it with postcard. /// /// The length prefix is validated against `max_size` before any payload /// memory is allocated. A frame that exceeds the limit is rejected without /// reading its payload. pub(crate) async fn read_frame(reader: &mut R, max_size: usize) -> Result where R: AsyncRead + Unpin, T: DeserializeOwned, { let mut len_bytes = [0u8; 4]; reader .read_exact(&mut len_bytes) .await .map_err(|err| NetworkError::Transport(format!("failed to read frame header: {err}")))?; let len = u32::from_be_bytes(len_bytes) as usize; if len > max_size { return Err(NetworkError::Transport(format!( "incoming frame of {len} bytes exceeds limit of {max_size} bytes" ))); } let mut payload = vec![0u8; len]; reader .read_exact(&mut payload) .await .map_err(|err| NetworkError::Transport(format!("failed to read frame payload: {err}")))?; decode(&payload) } #[cfg(test)] mod tests { use serde::Deserialize; use super::*; #[derive(Debug, PartialEq, Serialize, Deserialize)] struct TestValue { text: String, number: u64, } #[tokio::test] async fn frame_round_trip() { let value = TestValue { text: "hello".into(), number: 42, }; let mut buf = Vec::new(); write_frame(&mut buf, &value, 1024).await.expect("write"); let decoded: TestValue = read_frame(&mut buf.as_slice(), 1024).await.expect("read"); assert_eq!(decoded, value); } #[tokio::test] async fn oversized_frame_is_rejected_before_allocation() { // A header declaring a huge payload with no payload bytes present. // If the length check happened after allocation, this would try to // allocate 4 GiB; instead it must fail fast on the limit check. let mut buf = Vec::new(); buf.extend_from_slice(&u32::MAX.to_be_bytes()); let result: Result = read_frame(&mut buf.as_slice(), 1024).await; assert!(matches!(result, Err(NetworkError::Transport(_)))); } #[tokio::test] async fn oversized_value_is_rejected_on_write() { let value = TestValue { text: "x".repeat(2048), number: 1, }; let mut buf = Vec::new(); let result = write_frame(&mut buf, &value, 1024).await; assert!(matches!(result, Err(NetworkError::MessageTooLarge))); assert!(buf.is_empty()); } #[tokio::test] async fn malformed_payload_is_rejected() { // Valid header, but payload bytes that do not decode as TestValue. let payload = [0xffu8; 8]; let mut buf = Vec::new(); buf.extend_from_slice(&(payload.len() as u32).to_be_bytes()); buf.extend_from_slice(&payload); let result: Result = read_frame(&mut buf.as_slice(), 1024).await; assert!(matches!(result, Err(NetworkError::Serialization(_)))); } #[tokio::test] async fn truncated_stream_is_an_error() { let value = TestValue { text: "hello".into(), number: 42, }; let mut buf = Vec::new(); write_frame(&mut buf, &value, 1024).await.expect("write"); buf.truncate(buf.len() - 1); let result: Result = read_frame(&mut buf.as_slice(), 1024).await; assert!(matches!(result, Err(NetworkError::Transport(_)))); } }