use std::cell::RefCell;
use std::fmt::Debug;
use std::future::Future;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use anyhow::{bail, Result};
#[cfg(feature = "postcard")]
use postcard::{from_bytes, to_allocvec};
#[cfg(feature = "postcard")]
use serde::{Deserialize, Serialize};
use crate::raw::{RawTCPClient, RawTCPResponse, RawTCPServer};
pub trait SerializeMessage: Sized + Send + Sync + 'static {
fn serialize(&self) -> Result<Vec<u8>>;
fn deserialize(message: &[u8]) -> Result<Self>;
}
#[cfg(feature = "postcard")]
impl<T> SerializeMessage for T
where
T: Serialize + for<'a> Deserialize<'a> + Send + Sync + 'static,
{
fn serialize(&self) -> Result<Vec<u8>> {
Ok(to_allocvec(self)?)
}
fn deserialize(message: &[u8]) -> Result<Self> {
Ok(from_bytes(message)?)
}
}
pub struct TCPClient<Q, A>
where
Q: SerializeMessage,
A: SerializeMessage,
{
raw_client: RawTCPClient,
phantom_q: PhantomData<Q>,
phantom_a: PhantomData<A>,
}
impl<Q, A> TCPClient<Q, A>
where
Q: SerializeMessage,
A: SerializeMessage,
{
pub async fn connect(host: &str, port: u16) -> Result<Self> {
Ok(Self {
raw_client: RawTCPClient::connect(host, port).await?,
phantom_q: PhantomData,
phantom_a: PhantomData,
})
}
pub async fn send(&mut self, message: Q) -> Result<Option<A>> {
let raw_message = message.serialize()?;
let raw_response = self.raw_client.send(&raw_message).await?;
if raw_response.is_empty() {
return Ok(None);
}
Ok(Some(A::deserialize(&raw_response)?))
}
}
#[derive(Debug)]
pub enum TCPResponse<A>
where
A: SerializeMessage,
{
Message(A),
CloseConnection,
StopServer,
}
pub struct TCPServer<Q, A, H, F>
where
Q: SerializeMessage,
A: SerializeMessage,
H: Fn(Q) -> F + Send + Sync + 'static,
F: Future<Output = Result<TCPResponse<A>>> + Send + 'static,
{
host: String,
port: u16,
handler: H,
bad_request_response: Mutex<RefCell<Option<fn() -> TCPResponse<A>>>>,
inactivity_timeout_ms: AtomicU64,
phantom_q: PhantomData<Q>,
}
impl<Q, A, H, F> TCPServer<Q, A, H, F>
where
A: SerializeMessage,
Q: SerializeMessage,
H: Fn(Q) -> F + Send + Sync + 'static,
F: Future<Output = Result<TCPResponse<A>>> + Send + 'static,
{
pub fn new(host: impl Into<String>, port: u16, handler: H) -> Arc<Self> {
Arc::new(Self {
host: host.into(),
port,
handler,
bad_request_response: Mutex::new(RefCell::new(None)),
inactivity_timeout_ms: AtomicU64::new(0),
phantom_q: PhantomData,
})
}
pub fn with_bad_request_handler(
self: Arc<Self>,
bad_request_response: fn() -> TCPResponse<A>,
) -> Arc<Self> {
*self.bad_request_response.lock().unwrap().borrow_mut() = Some(bad_request_response);
self
}
pub fn with_inactivity_timeout(self: Arc<Self>, timeout_ms: u64) -> Arc<Self> {
self.inactivity_timeout_ms
.store(timeout_ms, Ordering::Relaxed);
self
}
pub async fn listen(self: Arc<Self>) {
let cloned_self = self.clone();
RawTCPServer::new(self.host.clone(), self.port, move |req| {
cloned_self.clone().handle_raw_message(req)
})
.with_inactivity_timeout(self.inactivity_timeout_ms.load(Ordering::Relaxed))
.listen()
.await;
}
async fn handle_raw_message(self: Arc<Self>, raw_request: Vec<u8>) -> Result<RawTCPResponse> {
let action = match Q::deserialize(&raw_request) {
Ok(request) => (self.handler)(request).await?,
Err(err) => {
if let Some(bad_request_handler) =
self.bad_request_response.lock().unwrap().borrow().clone()
{
bad_request_handler()
} else {
bail!(
"Bad request, unable to deserialize (this might be caused \
by incompatible client version or mismatch in compression configuration): {}",
err
);
}
}
};
Ok(match action {
TCPResponse::Message(response) => RawTCPResponse::Message(response.serialize()?),
TCPResponse::CloseConnection => RawTCPResponse::CloseConnection,
TCPResponse::StopServer => RawTCPResponse::StopServer,
})
}
}
#[cfg(test)]
mod tests {
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use serial_test::serial;
use tokio::spawn;
use tokio::time::sleep;
use super::*;
const HOST: &str = "127.0.0.1";
const PORT: u16 = 12345;
#[derive(Debug, Serialize, Deserialize)]
enum Request {
Hello,
Double { num: u64 },
Sum { a: u64, b: u64 },
Close,
CauseError,
Stop,
}
#[derive(Debug, Serialize, Deserialize)]
enum OtherRequest {
One,
Two,
Three,
Four,
Five,
Six,
ImDifferent,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
enum Response {
World,
Result(u64),
}
async fn while_server_running(fut: impl Future) {
let handle = spawn(
TCPServer::new(HOST, PORT, handle_requests)
.with_inactivity_timeout(15)
.listen(),
);
sleep(Duration::from_millis(10)).await;
fut.await;
handle.await.unwrap();
}
async fn while_server_running_custom_bad_request(fut: impl Future) {
let handle = spawn(
TCPServer::new(HOST, PORT, handle_requests)
.with_inactivity_timeout(15)
.with_bad_request_handler(|| TCPResponse::Message(Response::World))
.listen(),
);
sleep(Duration::from_millis(10)).await;
fut.await;
handle.await.unwrap();
}
async fn handle_requests(req: Request) -> Result<TCPResponse<Response>> {
Ok(match req {
Request::Hello => TCPResponse::Message(Response::World),
Request::Double { num } => TCPResponse::Message(Response::Result(2 * num)),
Request::Sum { a, b } => TCPResponse::Message(Response::Result(a + b)),
Request::Close => TCPResponse::CloseConnection,
Request::CauseError => bail!("An error occurred".to_string()),
Request::Stop => TCPResponse::StopServer,
})
}
#[tokio::test]
#[serial]
async fn test_tcp_server() {
while_server_running(async {
let mut client = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
assert_eq!(
client.send(Request::Hello).await.unwrap().unwrap(),
Response::World
);
assert_eq!(
client
.send(Request::Double { num: 3 })
.await
.unwrap()
.unwrap(),
Response::Result(6)
);
assert_eq!(
client
.send(Request::Sum { a: 3, b: 5 })
.await
.unwrap()
.unwrap(),
Response::Result(8)
);
assert_eq!(client.send(Request::Stop).await.unwrap(), None);
})
.await;
}
#[tokio::test]
#[serial]
async fn test_tcp_server_handler_error() {
while_server_running(async {
let mut client = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
assert_eq!(client.send(Request::CauseError).await.unwrap(), None);
assert_eq!(client.send(Request::Stop).await.unwrap(), None);
})
.await;
}
#[tokio::test]
#[serial]
async fn test_tcp_server_bad_request_error() {
let start = Instant::now();
while_server_running(async {
let mut client = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
assert_eq!(client.send(OtherRequest::ImDifferent).await.unwrap(), None);
})
.await;
let elapsed_ms = start.elapsed().as_millis();
assert!(elapsed_ms >= 15, "Elapsed time: {} ms", elapsed_ms);
}
#[tokio::test]
#[serial]
async fn test_tcp_server_bad_request_error_custom() {
let start = Instant::now();
while_server_running_custom_bad_request(async {
let mut client = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
assert_eq!(
client.send(OtherRequest::ImDifferent).await.unwrap(),
Some(Response::World)
);
})
.await;
let elapsed_ms = start.elapsed().as_millis();
assert!(elapsed_ms >= 15, "Elapsed time: {} ms", elapsed_ms);
}
#[tokio::test]
#[serial]
#[should_panic]
async fn test_tcp_server_close_connection() {
while_server_running(async {
let mut client1 = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
let mut client2 = TCPClient::<_, Response>::connect(HOST, PORT).await.unwrap();
assert_eq!(client1.send(Request::Close).await.unwrap(), None);
assert_eq!(
client2.send(Request::Hello).await.unwrap().unwrap(),
Response::World
);
client1.send(Request::Hello).await.unwrap();
})
.await;
}
#[tokio::test]
#[serial]
async fn test_handle_tcp_requests_timeout() {
let start = Instant::now();
let handle = spawn(
TCPServer::new(HOST, PORT, handle_requests)
.with_inactivity_timeout(15)
.listen(),
);
handle.await.unwrap();
let elapsed_ms = start.elapsed().as_millis();
assert!(
14 <= elapsed_ms && elapsed_ms <= 16,
"Elapsed time: {} ms",
elapsed_ms
);
}
}