use std::io;
use thiserror::Error;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum FrameError {
#[error(transparent)]
Io(#[from] io::Error),
#[error("payload length {length} exceeds the maximum of {max} bytes")]
PayloadTooLarge { length: u64, max: u64 },
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct RawMessage {
pub message_type: u8,
pub payload: Vec<u8>,
}
impl RawMessage {
pub fn new(message_type: u8, payload: Vec<u8>) -> Self {
Self {
message_type,
payload,
}
}
}
pub async fn read_frame<R>(reader: &mut R, max_payload_len: u32) -> Result<RawMessage, FrameError>
where
R: AsyncRead + Unpin,
{
let message_type = reader.read_u8().await?;
let length = reader.read_u32_le().await?;
if length > max_payload_len {
return Err(FrameError::PayloadTooLarge {
length: length.into(),
max: max_payload_len.into(),
});
}
let mut payload = vec![0; length as usize];
reader.read_exact(&mut payload).await?;
Ok(RawMessage {
message_type,
payload,
})
}
pub async fn write_frame<W>(writer: &mut W, message: &RawMessage) -> Result<(), FrameError>
where
W: AsyncWrite + Unpin,
{
let length: u32 =
message
.payload
.len()
.try_into()
.map_err(|_| FrameError::PayloadTooLarge {
length: message.payload.len() as u64,
max: u32::MAX.into(),
})?;
writer.write_u8(message.message_type).await?;
writer.write_u32_le(length).await?;
writer.write_all(&message.payload).await?;
writer.flush().await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn frame_round_trips() {
let (mut client, mut server) = tokio::io::duplex(1024);
let message = RawMessage::new(0x08, vec![4, 1, 2, 3, 4]);
write_frame(&mut client, &message).await.unwrap();
let read = read_frame(&mut server, 1024).await.unwrap();
assert_eq!(read, message);
}
#[tokio::test]
async fn frame_layout_matches_spec() {
let (mut client, mut server) = tokio::io::duplex(1024);
let message = RawMessage::new(0x01, vec![3, 3]);
write_frame(&mut client, &message).await.unwrap();
let mut bytes = [0u8; 7];
server.read_exact(&mut bytes).await.unwrap();
assert_eq!(bytes, [0x01, 2, 0, 0, 0, 3, 3]);
}
#[tokio::test]
async fn empty_payload_round_trips() {
let (mut client, mut server) = tokio::io::duplex(1024);
let message = RawMessage::new(0x05, vec![]);
write_frame(&mut client, &message).await.unwrap();
let read = read_frame(&mut server, 1024).await.unwrap();
assert_eq!(read, message);
}
#[tokio::test]
async fn oversized_payload_is_rejected_before_reading() {
let (mut client, mut server) = tokio::io::duplex(1024);
let message = RawMessage::new(0x0b, vec![0; 512]);
write_frame(&mut client, &message).await.unwrap();
let error = read_frame(&mut server, 256).await.unwrap_err();
assert!(matches!(
error,
FrameError::PayloadTooLarge {
length: 512,
max: 256
}
));
}
#[tokio::test]
async fn payload_at_exact_limit_is_accepted() {
let (mut client, mut server) = tokio::io::duplex(1024);
let message = RawMessage::new(0x0b, vec![7; 256]);
write_frame(&mut client, &message).await.unwrap();
let read = read_frame(&mut server, 256).await.unwrap();
assert_eq!(read, message);
}
#[tokio::test]
async fn truncated_frame_fails_with_io_error() {
let (mut client, mut server) = tokio::io::duplex(1024);
client.write_all(&[0x08, 10, 0, 0, 0, 1, 2]).await.unwrap();
drop(client);
let error = read_frame(&mut server, 1024).await.unwrap_err();
assert!(matches!(error, FrameError::Io(_)));
}
}