use std::time::Duration;
use miette::IntoDiagnostic;
use serde::{Deserialize, Serialize};
use tokio::{io::{AsyncRead,
AsyncReadExt,
AsyncWrite,
AsyncWriteExt,
BufReader,
BufWriter},
time::timeout};
use crate::{bincode_serde, compress, ok, protocol_types::*};
pub mod protocol_constants {
use super::*;
pub const MAGIC_NUMBER: u64 = 0xDEADBEEFCAFEBABE;
pub const PROTOCOL_VERSION: u64 = 1;
pub const TIMEOUT_DURATION: Duration = Duration::from_secs(1);
pub const MAX_PAYLOAD_SIZE: u64 = 10_000_000;
}
pub mod handshake {
use super::*;
pub async fn try_connect_or_timeout<W: AsyncWrite + Unpin, R: AsyncRead + Unpin>(
read_half: &mut R,
write_half: &mut W,
) -> miette::Result<()> {
let result = timeout(
protocol_constants::TIMEOUT_DURATION,
try_connect(read_half, write_half),
)
.await;
match result {
Ok(Err(handshake_err)) => {
miette::bail!("Handshake failed due to: {}", handshake_err.root_cause())
}
Err(_elapsed_err) => {
miette::bail!("Handshake timed out")
}
_ => {
ok!()
}
}
}
async fn try_connect<W: AsyncWrite + Unpin, R: AsyncRead + Unpin>(
read_half: &mut R,
write_half: &mut W,
) -> miette::Result<()> {
write_half
.write_u64(protocol_constants::MAGIC_NUMBER)
.await
.into_diagnostic()?;
write_half
.write_u64(protocol_constants::PROTOCOL_VERSION)
.await
.into_diagnostic()?;
write_half.flush().await.into_diagnostic()?;
let received_magic_number = read_half.read_u64().await.into_diagnostic()?;
if received_magic_number != protocol_constants::MAGIC_NUMBER {
miette::bail!("Invalid protocol magic number")
}
ok!()
}
pub async fn try_accept_or_timeout<W: AsyncWrite + Unpin, R: AsyncRead + Unpin>(
read_half: &mut R,
write_half: &mut W,
) -> miette::Result<()> {
let result = timeout(
protocol_constants::TIMEOUT_DURATION,
try_accept(read_half, write_half),
)
.await
.into_diagnostic();
match result {
Ok(handshake_result) => match handshake_result {
Ok(_) => ok!(),
Err(handshake_err) => {
miette::bail!(
"Handshake failed due to: {}",
handshake_err.root_cause()
)
}
},
Err(_elapsed_err) => miette::bail!("Handshake timed out"),
}
}
async fn try_accept<W: AsyncWrite + Unpin, R: AsyncRead + Unpin>(
read_half: &mut R,
write_half: &mut W,
) -> miette::Result<()> {
let received_magic_number = read_half.read_u64().await.into_diagnostic()?;
if received_magic_number != protocol_constants::MAGIC_NUMBER {
miette::bail!("Invalid protocol magic number")
}
let received_protocol_version = read_half.read_u64().await.into_diagnostic()?;
if received_protocol_version != protocol_constants::PROTOCOL_VERSION {
miette::bail!("Invalid protocol version")
}
write_half
.write_u64(protocol_constants::MAGIC_NUMBER)
.await
.into_diagnostic()?;
ok!()
}
}
#[cfg(test)]
mod tests_handshake {
use super::*;
use crate::{get_mock_socket_halves, MockSocket};
#[tokio::test]
async fn test_handshake() {
let MockSocket {
mut client_read,
mut client_write,
mut server_read,
mut server_write,
} = get_mock_socket_halves();
let client_handshake =
handshake::try_connect_or_timeout(&mut client_read, &mut client_write);
let server_handshake =
handshake::try_accept_or_timeout(&mut server_read, &mut server_write);
let (client_handshake_result, server_handshake_result) =
tokio::join!(client_handshake, server_handshake);
assert!(client_handshake_result.is_ok());
assert!(server_handshake_result.is_ok());
}
}
pub mod byte_io {
use super::*;
pub async fn try_write<W: AsyncWrite + Unpin, T: Serialize>(
buf_writer: &mut BufWriter<W>,
data: &T,
) -> miette::Result<()> {
let payload_buffer = bincode_serde::try_serialize(data)?;
let payload_buffer = compress::compress(&payload_buffer)?;
let payload_size = payload_buffer.len();
buf_writer
.write_u64(payload_size as LengthPrefixType)
.await
.into_diagnostic()?;
buf_writer
.write_all(&payload_buffer)
.await
.into_diagnostic()?;
buf_writer.flush().await.into_diagnostic()?;
Ok(())
}
pub async fn try_read<R: AsyncRead + Unpin, T: for<'d> Deserialize<'d>>(
buf_reader: &mut BufReader<R>,
) -> miette::Result<T> {
let size_of_payload = buf_reader.read_u64().await.into_diagnostic()?;
if size_of_payload > protocol_constants::MAX_PAYLOAD_SIZE {
miette::bail!("Payload size is too large")
}
let mut payload_buffer = vec![0; size_of_payload as usize];
buf_reader
.read_exact(&mut payload_buffer)
.await
.into_diagnostic()?;
let payload_buffer = compress::decompress(&payload_buffer)?;
let payload_buffer = bincode_serde::try_deserialize::<T>(&payload_buffer)?;
Ok(payload_buffer)
}
}
#[cfg(test)]
mod tests_byte_io {
use super::*;
use crate::{get_mock_socket_halves, MockSocket};
pub fn get_all_client_messages<'a>() -> Vec<&'a str> { vec!["one", "two", "three"] }
#[tokio::test]
async fn test_byte_io() {
let MockSocket {
client_read: _,
mut client_write,
mut server_read,
server_write: _,
} = get_mock_socket_halves();
for sent_payload in get_all_client_messages() {
byte_io::try_write(&mut BufWriter::new(&mut client_write), &sent_payload)
.await
.unwrap();
let received_payload: String =
byte_io::try_read(&mut BufReader::new(&mut server_read))
.await
.unwrap();
assert_eq!(received_payload, sent_payload);
}
}
}