#![doc = include_str!("../docs/raw.md")]
use std::future::Future;
#[cfg(feature = "zstd")]
use std::io::Cursor;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use anyhow::Result;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::broadcast;
use tokio::sync::broadcast::error::RecvError;
use tokio::sync::broadcast::Sender;
use tokio::time::sleep;
use tokio::{select, spawn};
use tracing::{debug, error, info};
#[cfg(feature = "zstd")]
use zstd::{decode_all, encode_all};
pub struct RawTCPClient {
tcp_stream: TcpStream,
}
impl RawTCPClient {
pub async fn connect(host: &str, port: u16) -> Result<Self> {
Ok(Self {
tcp_stream: TcpStream::connect((host, port)).await?,
})
}
pub async fn send(&mut self, message: &[u8]) -> Result<Vec<u8>> {
write(&mut self.tcp_stream, message).await?;
read(&mut self.tcp_stream).await
}
}
#[derive(Debug, Clone)]
pub enum RawTCPResponse {
Message(Vec<u8>),
CloseConnection,
StopServer,
}
enum RequestAction {
CloseConnection,
StopServer,
None,
}
pub struct RawTCPServer<H, F>
where
H: Fn(Vec<u8>) -> F + Send + Sync + 'static,
F: Future<Output = Result<RawTCPResponse>> + Send + 'static,
{
host: String,
port: u16,
handler: H,
inactivity_timeout_ms: AtomicU64,
}
impl<H, F> RawTCPServer<H, F>
where
H: Fn(Vec<u8>) -> F + Send + Sync + 'static,
F: Future<Output = Result<RawTCPResponse>> + Send + 'static,
{
pub fn new(host: impl Into<String>, port: u16, handler: H) -> Arc<Self> {
Arc::new(Self {
host: host.into(),
port,
handler,
inactivity_timeout_ms: AtomicU64::new(0),
})
}
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 (signal_sender, mut signal_receiver) = broadcast::channel(1);
spawn(self.clone().accept_connections(signal_sender));
let inactivity_timeout_ms = self.inactivity_timeout_ms.load(Ordering::Relaxed);
loop {
let must_exit = if inactivity_timeout_ms == 0 {
signal_receiver
.recv()
.await
.expect("Unable to read message from channel")
} else {
select! {
must_exit = signal_receiver.recv() => match must_exit {
Ok(must_exit) => must_exit,
Err(RecvError::Lagged(_)) => false,
Err(_) => true,
},
_ = sleep(Duration::from_millis(inactivity_timeout_ms)) => {
info!(
port = self.port,
"Not receiving requests since more that {} seconds, stopping request handling...",
inactivity_timeout_ms / 1_000
);
true
},
}
};
if must_exit {
break;
}
}
}
async fn accept_connections(self: Arc<Self>, signal_sender: Sender<bool>) {
let listener = match TcpListener::bind((self.host.as_str(), self.port)).await {
Ok(listener) => listener,
Err(err) => {
error!(
port = self.port,
error = ?err,
"Unable to bind socket listener: {:?}",
err,
);
panic!("Unable to bind socket listener: {:?}", err);
}
};
let mut signal_receiver = signal_sender.subscribe();
loop {
select! {
result = listener.accept() => match result {
Ok((tcp_stream, _)) => {
debug!(port = self.port, "Client successfully connected");
spawn(
self.clone()
.handle_connection(tcp_stream, signal_sender.clone())
);
}
Err(err) => error!(
port = self.port,
error = ?err,
"Error while accepting connection: {:?}",
err,
),
},
must_exit = signal_receiver.recv() => match must_exit {
Ok(must_exit) => if must_exit { break },
Err(RecvError::Lagged(_)) => {},
Err(_) => break,
},
}
}
}
async fn handle_connection(
self: Arc<Self>,
mut tcp_stream: TcpStream,
signal_sender: Sender<bool>,
) {
loop {
let action = self.clone().handle_request(&mut tcp_stream).await;
signal_sender
.send(matches!(action, RequestAction::StopServer))
.expect("Unable to send message to channel");
match action {
RequestAction::CloseConnection | RequestAction::StopServer => {
break;
}
RequestAction::None => {}
}
}
}
async fn handle_request(self: Arc<Self>, tcp_stream: &mut TcpStream) -> RequestAction {
let request = match read(tcp_stream).await {
Ok(msg) => msg,
Err(err) => {
debug!(
port = self.port,
error = ?err,
"Connection closed by client",
);
return RequestAction::CloseConnection;
}
};
match (self.handler)(request).await {
Ok(action) => match action {
RawTCPResponse::Message(message) => {
if !message.is_empty() {
if let Err(error) = write(tcp_stream, &message).await {
error!(
port = self.port,
?error,
?message,
"Connection was unexpectedly closed by the client while handling \
the request. The handler is unable to return the response \
and will be forced to discard it: {:?}",
error,
);
return RequestAction::CloseConnection;
}
}
RequestAction::None
}
RawTCPResponse::CloseConnection => {
write(tcp_stream, &[]).await.expect(
"Unable to write empty message to socket before closing connection.",
);
RequestAction::CloseConnection
}
RawTCPResponse::StopServer => {
write(tcp_stream, &[]).await.expect(
"Unable to write empty message to socket before stopping TCP server.",
);
RequestAction::StopServer
}
},
Err(error) => {
error!(
port = self.port,
?error,
"Error while handling request: {:?}",
error,
);
write(tcp_stream, &[])
.await
.expect("Unable to write empty message to socket after handler error.");
RequestAction::None
}
}
}
}
#[cfg(not(feature = "zstd"))]
async fn read(tcp_stream: &mut TcpStream) -> Result<Vec<u8>> {
raw_read(tcp_stream).await
}
#[cfg(feature = "zstd")]
async fn read(tcp_stream: &mut TcpStream) -> Result<Vec<u8>> {
let compressed_response = raw_read(tcp_stream).await?;
Ok(decode_all(Cursor::new(compressed_response))?)
}
#[cfg(not(feature = "zstd"))]
async fn write(tcp_stream: &mut TcpStream, message: &[u8]) -> Result<()> {
raw_write(tcp_stream, message).await
}
#[cfg(feature = "zstd")]
async fn write(tcp_stream: &mut TcpStream, message: &[u8]) -> Result<()> {
let compressed_message = encode_all(message, 0)?;
raw_write(tcp_stream, &compressed_message).await
}
async fn raw_read(tcp_stream: &mut TcpStream) -> Result<Vec<u8>> {
tcp_stream.readable().await?;
let len = tcp_stream.read_u64_le().await?;
let mut buffer = Vec::with_capacity(len as usize);
tcp_stream.take(len).read_to_end(&mut buffer).await?;
Ok(buffer)
}
async fn raw_write(tcp_stream: &mut TcpStream, message: &[u8]) -> Result<()> {
tcp_stream.writable().await?;
tcp_stream.write_u64_le(message.len() as u64).await?;
tcp_stream.write_all(message).await?;
tcp_stream.flush().await?;
Ok(())
}
#[cfg(test)]
mod tests {
use anyhow::bail;
use serial_test::serial;
use tokio::time::Instant;
use super::*;
const HOST: &str = "127.0.0.1";
const PORT: u16 = 12345;
async fn while_server_running(fut: impl Future) {
let handle = spawn(RawTCPServer::new(HOST, PORT, handle_requests).listen());
sleep(Duration::from_millis(10)).await;
fut.await;
handle.await.unwrap();
}
async fn handle_requests(req: Vec<u8>) -> Result<RawTCPResponse> {
Ok(match &req[..] {
&[1] => RawTCPResponse::Message(vec![10]),
&[2] => RawTCPResponse::Message(vec![200]),
&[3] => bail!("Test error"),
&[4] => RawTCPResponse::CloseConnection,
_ => RawTCPResponse::StopServer,
})
}
#[tokio::test]
#[serial]
async fn test_raw_tcp_server() {
while_server_running(async {
let mut client = RawTCPClient::connect(HOST, PORT).await.unwrap();
assert_eq!(client.send(&[1]).await.unwrap(), vec![10]);
assert_eq!(client.send(&[2]).await.unwrap(), vec![200]);
assert_eq!(client.send(&[]).await.unwrap(), vec![]);
})
.await;
}
#[tokio::test]
#[serial]
async fn test_raw_tcp_server_handler_error() {
while_server_running(async {
let mut client = RawTCPClient::connect(HOST, PORT).await.unwrap();
assert_eq!(client.send(&[3]).await.unwrap(), vec![]);
client.send(&[]).await.unwrap();
})
.await;
}
#[tokio::test]
#[serial]
#[should_panic]
async fn test_raw_tcp_server_close_connection() {
while_server_running(async {
let mut client1 = RawTCPClient::connect(HOST, PORT).await.unwrap();
let mut client2 = RawTCPClient::connect(HOST, PORT).await.unwrap();
assert_eq!(client1.send(&[4]).await.unwrap(), vec![]);
assert_eq!(client2.send(&[1]).await.unwrap(), vec![10]);
client1.send(&[1]).await.unwrap();
})
.await;
}
#[tokio::test]
#[serial]
async fn test_handle_raw_tcp_requests_timeout() {
let start = Instant::now();
let handle = spawn(
RawTCPServer::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);
}
}