use std::io::Cursor;
use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
use log::debug;
use serde_cbor::de::from_slice;
use serde_cbor::ser::to_vec;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::error::Error;
use crate::network::message::*;
pub use super::socket::*;
pub const PACKET_SIZE: usize = 1280;
pub async fn send_message<T>(message: T, stream: &mut GenericStream) -> Result<(), Error>
where
T: Into<Message>,
{
let message: Message = message.into();
debug!("Sending message: {message:#?}",);
let payload = to_vec(&message).map_err(|err| Error::MessageDeserialization(err.to_string()))?;
send_bytes(&payload, stream).await
}
pub async fn send_bytes(payload: &[u8], stream: &mut GenericStream) -> Result<(), Error> {
let message_size = payload.len() as u64;
let mut header = Vec::new();
WriteBytesExt::write_u64::<BigEndian>(&mut header, message_size).unwrap();
stream
.write_all(&header)
.await
.map_err(|err| Error::IoError("sending request size header".to_string(), err))?;
for chunk in payload.chunks(PACKET_SIZE) {
stream
.write_all(chunk)
.await
.map_err(|err| Error::IoError("sending payload chunk".to_string(), err))?;
}
Ok(())
}
pub async fn receive_bytes(stream: &mut GenericStream) -> Result<Vec<u8>, Error> {
let mut header = vec![0; 8];
stream
.read_exact(&mut header)
.await
.map_err(|err| Error::IoError("reading request size header".to_string(), err))?;
let mut header = Cursor::new(header);
let message_size = ReadBytesExt::read_u64::<BigEndian>(&mut header)? as usize;
let mut payload_bytes = Vec::with_capacity(message_size);
while payload_bytes.len() < message_size {
let remaining_bytes = message_size - payload_bytes.len();
let mut chunk_buffer: Vec<u8> = if remaining_bytes < PACKET_SIZE {
vec![0; remaining_bytes]
} else {
vec![0; PACKET_SIZE]
};
let received_bytes = stream
.read(&mut chunk_buffer)
.await
.map_err(|err| Error::IoError("reading next chunk".to_string(), err))?;
if received_bytes == 0 {
return Err(Error::Connection(
"Connection went away while receiving payload.".into(),
));
}
payload_bytes.extend_from_slice(&chunk_buffer[0..received_bytes]);
}
Ok(payload_bytes)
}
pub async fn receive_message(stream: &mut GenericStream) -> Result<Message, Error> {
let payload_bytes = receive_bytes(stream).await?;
if payload_bytes.is_empty() {
return Err(Error::EmptyPayload);
}
let message: Message =
from_slice(&payload_bytes).map_err(|err| Error::MessageDeserialization(err.to_string()))?;
debug!("Received message: {message:#?}");
Ok(message)
}
#[cfg(test)]
mod test {
use std::time::Duration;
use async_trait::async_trait;
use pretty_assertions::assert_eq;
use tokio::net::{TcpListener, TcpStream};
use tokio::task;
use super::*;
use crate::network::socket::Stream as PueueStream;
#[async_trait]
impl Listener for TcpListener {
async fn accept<'a>(&'a self) -> Result<GenericStream, Error> {
let (stream, _) = self.accept().await?;
Ok(Box::new(stream))
}
}
impl PueueStream for TcpStream {}
#[tokio::test]
async fn test_single_huge_payload() -> Result<(), Error> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let payload = "a".repeat(100_000);
let message = create_success_message(payload);
let original_bytes = to_vec(&message).expect("Failed to serialize message.");
let listener: GenericListener = Box::new(listener);
task::spawn(async move {
let mut stream = listener.accept().await.unwrap();
let message_bytes = receive_bytes(&mut stream).await.unwrap();
let message: Message = from_slice(&message_bytes).unwrap();
send_message(message, &mut stream).await.unwrap();
});
let mut client: GenericStream = Box::new(TcpStream::connect(&addr).await?);
send_message(message, &mut client).await?;
let response_bytes = receive_bytes(&mut client).await?;
let _message: Message = from_slice(&response_bytes)
.map_err(|err| Error::MessageDeserialization(err.to_string()))?;
assert_eq!(response_bytes, original_bytes);
Ok(())
}
#[tokio::test]
async fn test_successive_messages() -> Result<(), Error> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let listener: GenericListener = Box::new(listener);
task::spawn(async move {
let mut stream = listener.accept().await.unwrap();
send_message(create_success_message("message_a"), &mut stream)
.await
.unwrap();
send_message(create_success_message("message_b"), &mut stream)
.await
.unwrap();
});
let mut client: GenericStream = Box::new(TcpStream::connect(&addr).await?);
tokio::time::sleep(Duration::from_millis(500)).await;
let message_a = receive_message(&mut client).await.expect("First message");
let message_b = receive_message(&mut client).await.expect("Second message");
assert_eq!(Message::Success("message_a".to_string()), message_a);
assert_eq!(Message::Success("message_b".to_string()), message_b);
Ok(())
}
}