use std::io::{self, Read, Write};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
const FRAME_HEADER_SIZE: usize = 4;
pub fn encode_message(payload: &[u8]) -> Result<Vec<u8>, io::Error> {
let len = payload.len();
if len > MAX_MESSAGE_SIZE as usize {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Message too large: {} bytes (max: {})", len, MAX_MESSAGE_SIZE),
));
}
let mut frame = Vec::with_capacity(FRAME_HEADER_SIZE + len);
frame.extend_from_slice(&(len as u32).to_be_bytes());
frame.extend_from_slice(payload);
Ok(frame)
}
pub fn decode_message<R: Read>(reader: &mut R) -> Result<Vec<u8>, io::Error> {
let mut len_bytes = [0u8; FRAME_HEADER_SIZE];
reader.read_exact(&mut len_bytes)?;
let len = u32::from_be_bytes(len_bytes) as usize;
if len > MAX_MESSAGE_SIZE as usize {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("Message too large: {} bytes", len),
));
}
let mut payload = vec![0u8; len];
reader.read_exact(&mut payload)?;
Ok(payload)
}
pub async fn send_message<W>(writer: &mut W, payload: &[u8]) -> Result<(), io::Error>
where
W: AsyncWriteExt + Unpin,
{
let frame = encode_message(payload)?;
eprintln!(">>> SENDING PACKET <<<");
eprintln!(" Length: {} bytes", payload.len());
eprintln!(" Frame size (with header): {} bytes", frame.len());
eprintln!(" Payload hex: {}", hex::encode(&payload[..payload.len().min(256)]));
if payload.len() > 256 {
eprintln!(" (truncated, showing first 256 bytes)");
}
if let Ok(s) = std::str::from_utf8(payload) {
eprintln!(" Payload text: {}", s);
}
eprintln!();
writer.write_all(&frame).await?;
writer.flush().await?;
Ok(())
}
pub async fn recv_message<R>(reader: &mut R) -> Result<Vec<u8>, io::Error>
where
R: AsyncReadExt + Unpin,
{
let mut len_bytes = [0u8; FRAME_HEADER_SIZE];
reader.read_exact(&mut len_bytes).await?;
let len = u32::from_be_bytes(len_bytes) as usize;
if len > MAX_MESSAGE_SIZE as usize {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("Message too large: {} bytes", len),
));
}
let mut payload = vec![0u8; len];
reader.read_exact(&mut payload).await?;
eprintln!("<<< RECEIVED PACKET >>>");
eprintln!(" Length: {} bytes", len);
eprintln!(" Payload hex: {}", hex::encode(&payload[..payload.len().min(256)]));
if payload.len() > 256 {
eprintln!(" (truncated, showing first 256 bytes)");
}
if let Ok(s) = std::str::from_utf8(&payload) {
eprintln!(" Payload text: {}", s);
}
eprintln!();
Ok(payload)
}
pub fn write_frame<W: Write>(writer: &mut W, payload: &[u8]) -> Result<(), io::Error> {
let frame = encode_message(payload)?;
writer.write_all(&frame)?;
writer.flush()?;
Ok(())
}
pub struct FramedMessage {
pub payload: Vec<u8>,
}
impl FramedMessage {
pub fn new(payload: Vec<u8>) -> Self {
Self { payload }
}
pub fn encode(&self) -> Result<Vec<u8>, io::Error> {
encode_message(&self.payload)
}
pub fn as_str(&self) -> Result<&str, std::str::Utf8Error> {
std::str::from_utf8(&self.payload)
}
pub fn len(&self) -> usize {
self.payload.len()
}
pub fn is_empty(&self) -> bool {
self.payload.is_empty()
}
}
pub struct MessageFraming;
impl MessageFraming {
pub fn encode_json<T: serde::Serialize>(message: &T) -> Result<Vec<u8>, io::Error> {
let json = serde_json::to_vec(message).map_err(|e| {
io::Error::new(io::ErrorKind::InvalidData, format!("JSON serialization failed: {}", e))
})?;
encode_message(&json)
}
pub fn decode_json<T: serde::de::DeserializeOwned, R: Read>(reader: &mut R) -> Result<T, io::Error> {
let payload = decode_message(reader)?;
serde_json::from_slice(&payload).map_err(|e| {
io::Error::new(io::ErrorKind::InvalidData, format!("JSON deserialization failed: {}", e))
})
}
pub async fn send_json<T: serde::Serialize, W>(writer: &mut W, message: &T) -> Result<(), io::Error>
where
W: AsyncWriteExt + Unpin,
{
let json = serde_json::to_vec(message).map_err(|e| {
io::Error::new(io::ErrorKind::InvalidData, format!("JSON serialization failed: {}", e))
})?;
send_message(writer, &json).await
}
pub async fn recv_json<T: serde::de::DeserializeOwned, R>(reader: &mut R) -> Result<T, io::Error>
where
R: AsyncReadExt + Unpin,
{
let payload = recv_message(reader).await?;
serde_json::from_slice(&payload).map_err(|e| {
io::Error::new(io::ErrorKind::InvalidData, format!("JSON deserialization failed: {}", e))
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn test_encode_decode() {
let message = b"Hello, World!";
let encoded = encode_message(message).unwrap();
let len = u32::from_be_bytes([encoded[0], encoded[1], encoded[2], encoded[3]]);
assert_eq!(len, message.len() as u32);
let mut cursor = Cursor::new(encoded);
let decoded = decode_message(&mut cursor).unwrap();
assert_eq!(decoded, message);
}
#[test]
fn test_max_size_error() {
let oversized = vec![0u8; (MAX_MESSAGE_SIZE + 1) as usize];
assert!(encode_message(&oversized).is_err());
}
#[test]
fn test_empty_message() {
let message = b"";
let encoded = encode_message(message).unwrap();
let mut cursor = Cursor::new(encoded);
let decoded = decode_message(&mut cursor).unwrap();
assert_eq!(decoded, message);
}
}