Files
frid/crates/federation-net/src/wire.rs
T
2026-07-10 12:45:47 +03:00

153 lines
5.1 KiB
Rust

//! 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<T: Serialize>(value: &T, max_size: usize) -> Result<Vec<u8>> {
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<T: DeserializeOwned>(bytes: &[u8]) -> Result<T> {
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<W, T>(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<R, T>(reader: &mut R, max_size: usize) -> Result<T>
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<TestValue> = 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<TestValue> = 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<TestValue> = read_frame(&mut buf.as_slice(), 1024).await;
assert!(matches!(result, Err(NetworkError::Transport(_))));
}
}