use std::time::{Duration, Instant};
use dusa_collection_utils::core::errors::{ErrorArrayItem, Errors};
use dusa_collection_utils::core::logger::LogLevel;
use dusa_collection_utils::log;
use tokio::io::{self, AsyncRead, AsyncWrite, ReadHalf, WriteHalf};
use tokio::sync::{mpsc, oneshot};
use tokio::task::JoinHandle;
use tokio::time::MissedTickBehavior;
use crate::protocol::{
flags::MsgType,
handshake::{NoiseIdentity, rekey_initiator_rw, rekey_responder_rw},
heartbeat::{Heartbeat, HeartbeatState},
message::{ConnectionCtx, ProtocolMessage, SessionAck, SessionRequest, read_message_raw_buffered},
proto::Proto,
};
#[derive(Debug)]
pub enum DriverMessage<APP> {
Data(ProtocolMessage<APP>),
Open(ProtocolMessage<SessionRequest>),
OpenAck(ProtocolMessage<SessionAck>),
}
#[derive(Debug, Clone, Copy)]
pub struct DriverConfig {
pub heartbeat_interval: Duration,
pub outbound_buffer: usize,
pub inbound_buffer: usize,
}
impl Default for DriverConfig {
fn default() -> Self {
Self {
heartbeat_interval: Duration::from_secs(15),
outbound_buffer: 64,
inbound_buffer: 64,
}
}
}
pub enum ConnectionRole {
Initiator { remote_static_pubkey: [u8; 32] },
Responder { identity: NoiseIdentity },
}
enum Control {
Rekey(oneshot::Sender<Result<(), ErrorArrayItem>>),
Shutdown,
}
pub struct ConnectionHandle<APP> {
outbound_tx: mpsc::Sender<ProtocolMessage<APP>>,
inbound_rx: mpsc::Receiver<DriverMessage<APP>>,
control_tx: mpsc::Sender<Control>,
task: JoinHandle<io::Result<()>>,
}
impl<APP> ConnectionHandle<APP>
where
APP: serde::de::DeserializeOwned
+ serde::Serialize
+ std::fmt::Debug
+ Clone
+ Unpin
+ Send
+ 'static,
{
pub async fn send(&self, msg: ProtocolMessage<APP>) -> Result<(), ErrorArrayItem> {
self.outbound_tx
.send(msg)
.await
.map_err(|_| ErrorArrayItem::new(Errors::ConnectionError, "driver task has stopped"))
}
pub async fn recv(&mut self) -> Option<DriverMessage<APP>> {
self.inbound_rx.recv().await
}
pub async fn rekey(&self) -> Result<(), ErrorArrayItem> {
let (ack, done) = oneshot::channel();
self.control_tx
.send(Control::Rekey(ack))
.await
.map_err(|_| ErrorArrayItem::new(Errors::ConnectionError, "driver task has stopped"))?;
done.await
.map_err(|_| ErrorArrayItem::new(Errors::ConnectionError, "driver task has stopped"))?
}
pub fn is_finished(&self) -> bool {
self.task.is_finished()
}
pub async fn shutdown(mut self) -> io::Result<()> {
let _ = self.control_tx.send(Control::Shutdown).await;
self.inbound_rx.close();
match self.task.await {
Ok(result) => result,
Err(join_err) => Err(io::Error::other(join_err.to_string())),
}
}
}
pub struct ConnectionDriver;
impl ConnectionDriver {
pub fn spawn<S, APP>(
stream: S,
ctx: ConnectionCtx,
role: ConnectionRole,
proto: Proto,
config: DriverConfig,
) -> ConnectionHandle<APP>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
APP: serde::de::DeserializeOwned
+ serde::Serialize
+ std::fmt::Debug
+ Clone
+ Unpin
+ Send
+ 'static,
{
let (read_half, write_half) = tokio::io::split(stream);
let (outbound_tx, outbound_rx) = mpsc::channel(config.outbound_buffer);
let (inbound_tx, inbound_rx) = mpsc::channel(config.inbound_buffer);
let (control_tx, control_rx) = mpsc::channel(4);
let task = tokio::spawn(driver_loop::<S, APP>(
read_half,
write_half,
ctx,
role,
proto,
config,
DriverChannels { outbound_rx, inbound_tx, control_rx },
));
ConnectionHandle {
outbound_tx,
inbound_rx,
control_tx,
task,
}
}
}
struct DriverChannels<APP> {
outbound_rx: mpsc::Receiver<ProtocolMessage<APP>>,
inbound_tx: mpsc::Sender<DriverMessage<APP>>,
control_rx: mpsc::Receiver<Control>,
}
async fn driver_loop<S, APP>(
mut read_half: ReadHalf<S>,
mut write_half: WriteHalf<S>,
mut ctx: ConnectionCtx,
role: ConnectionRole,
proto: Proto,
config: DriverConfig,
channels: DriverChannels<APP>,
) -> io::Result<()>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
APP: serde::de::DeserializeOwned
+ serde::Serialize
+ std::fmt::Debug
+ Clone
+ Unpin
+ Send
+ 'static,
{
let DriverChannels { mut outbound_rx, inbound_tx, mut control_rx } = channels;
let mut heartbeat = Heartbeat::new(config.heartbeat_interval, Instant::now());
let mut ticker = tokio::time::interval(config.heartbeat_interval);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
let mut frame_buf: Vec<u8> = Vec::new();
loop {
tokio::select! {
read_result = read_message_raw_buffered(&mut read_half, Some(&mut ctx), &mut frame_buf) => {
let (header, payload) = match read_result {
Ok(v) => v,
Err(err) => {
log!(LogLevel::Error, "driver read error: {err}");
return Err(err);
}
};
match header.msg_type() {
MsgType::Heartbeat => {
heartbeat.mark_recv(Instant::now());
}
MsgType::Data => {
let msg: ProtocolMessage<APP> = ProtocolMessage::finish(header, &payload)?;
if inbound_tx.send(DriverMessage::Data(msg)).await.is_err() {
return Ok(());
}
}
MsgType::Open => {
let msg: ProtocolMessage<SessionRequest> =
ProtocolMessage::finish(header, &payload)?;
if inbound_tx.send(DriverMessage::Open(msg)).await.is_err() {
return Ok(());
}
}
MsgType::OpenAck => {
let msg: ProtocolMessage<SessionAck> =
ProtocolMessage::finish(header, &payload)?;
if inbound_tx.send(DriverMessage::OpenAck(msg)).await.is_err() {
return Ok(());
}
}
MsgType::Rekey => match &role {
ConnectionRole::Responder { identity } => {
match rekey_responder_rw(&mut read_half, &mut write_half, identity).await {
Ok(new_ctx) => {
ctx = new_ctx;
heartbeat.mark_recv(Instant::now());
log!(LogLevel::Info, "rekeyed (peer-initiated)");
}
Err(err) => {
log!(LogLevel::Error, "peer-initiated rekey failed: {err}");
return Err(err);
}
}
}
ConnectionRole::Initiator { .. } => {
log!(
LogLevel::Warn,
"ignoring unexpected Rekey signal from peer -- this side is the Noise initiator"
);
}
},
MsgType::Close => {
log!(LogLevel::Info, "peer closed the connection");
return Ok(());
}
other => {
log!(LogLevel::Warn, "driver ignoring unexpected message type: {other:?}");
}
}
}
outbound = outbound_rx.recv() => {
let Some(msg) = outbound else {
return Ok(());
};
if let Err(err) = msg.write_to(&mut write_half, proto, Some(&mut ctx)).await {
log!(LogLevel::Error, "driver write error: {err}");
return Err(err);
}
}
ctrl = control_rx.recv() => {
match ctrl {
Some(Control::Shutdown) | None => return Ok(()),
Some(Control::Rekey(ack)) => {
let result = match &role {
ConnectionRole::Initiator { remote_static_pubkey } => {
rekey_initiator_rw(&mut read_half, &mut write_half, &mut ctx, remote_static_pubkey)
.await
.map(|new_ctx| {
ctx = new_ctx;
heartbeat.mark_sent(Instant::now());
log!(LogLevel::Info, "rekeyed (self-initiated)");
})
.map_err(|err| ErrorArrayItem::new(Errors::Network, err.to_string()))
}
ConnectionRole::Responder { .. } => Err(ErrorArrayItem::new(
Errors::Unauthorized,
"only the connection's original Noise initiator can self-initiate a rekey",
)),
};
let _ = ack.send(result);
}
}
}
_ = ticker.tick() => {
if heartbeat.should_send(Instant::now()) {
let hb: ProtocolMessage<()> = ProtocolMessage::new(ctx.params, MsgType::Heartbeat, ())?;
if let Err(err) = hb.write_to(&mut write_half, proto, Some(&mut ctx)).await {
log!(LogLevel::Error, "driver heartbeat write error: {err}");
return Err(err);
}
heartbeat.mark_sent(Instant::now());
}
if heartbeat.state(Instant::now()) == HeartbeatState::Timeout {
log!(LogLevel::Error, "peer heartbeat timeout -- tearing down connection");
return Err(io::Error::new(io::ErrorKind::TimedOut, "peer heartbeat timeout"));
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::network::send_receive::{establish_connection_initiator, establish_connection_responder};
use crate::protocol::flags::ConnectionParams;
#[tokio::test]
async fn full_duplex_push_without_reply() {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client_stream, mut server_stream) = tokio::io::duplex(8192);
let (client_ctx, server_ctx) = tokio::join!(
establish_connection_initiator(&mut client_stream, &remote_pub, ConnectionParams::ENCRYPTED),
establish_connection_responder(&mut server_stream, &identity),
);
let client_ctx = client_ctx.unwrap();
let server_ctx = server_ctx.unwrap();
let config = DriverConfig::default();
let mut client: ConnectionHandle<Vec<u8>> = ConnectionDriver::spawn(
client_stream,
client_ctx,
ConnectionRole::Initiator { remote_static_pubkey: remote_pub },
Proto::TCP,
config,
);
let mut server: ConnectionHandle<Vec<u8>> = ConnectionDriver::spawn(
server_stream,
server_ctx,
ConnectionRole::Responder { identity },
Proto::TCP,
config,
);
for seq in 0..3u8 {
client
.send(ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Data, vec![seq]).unwrap())
.await
.unwrap();
}
for seq in 0..3u8 {
match server.recv().await.unwrap() {
DriverMessage::Data(msg) => assert_eq!(msg.payload, vec![seq]),
other => panic!("unexpected message: {other:?}"),
}
}
server
.send(ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Data, b"pong".to_vec()).unwrap())
.await
.unwrap();
match client.recv().await.unwrap() {
DriverMessage::Data(msg) => assert_eq!(msg.payload, b"pong".to_vec()),
other => panic!("unexpected message: {other:?}"),
}
let _ = client.shutdown().await;
let _ = server.shutdown().await;
}
#[tokio::test(start_paused = true)]
async fn heartbeat_timeout_tears_down_driver() {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client_stream, mut server_stream) = tokio::io::duplex(8192);
let (client_ctx, server_ctx) = tokio::join!(
establish_connection_initiator(&mut client_stream, &remote_pub, ConnectionParams::ENCRYPTED),
establish_connection_responder(&mut server_stream, &identity),
);
let client_ctx = client_ctx.unwrap();
let _server_ctx = server_ctx.unwrap();
let _server_stream = server_stream;
let config = DriverConfig {
heartbeat_interval: Duration::from_millis(50),
..Default::default()
};
let mut client: ConnectionHandle<Vec<u8>> = ConnectionDriver::spawn(
client_stream,
client_ctx,
ConnectionRole::Initiator { remote_static_pubkey: remote_pub },
Proto::TCP,
config,
);
tokio::time::advance(config.heartbeat_interval * 6).await;
assert!(client.recv().await.is_none());
assert!(client.shutdown().await.is_err());
}
#[tokio::test]
async fn rekey_while_sending_preserves_message_order() {
let identity = NoiseIdentity::generate().unwrap();
let remote_pub = identity.public_key();
let (mut client_stream, mut server_stream) = tokio::io::duplex(8192);
let (client_ctx, server_ctx) = tokio::join!(
establish_connection_initiator(&mut client_stream, &remote_pub, ConnectionParams::ENCRYPTED),
establish_connection_responder(&mut server_stream, &identity),
);
let client_ctx = client_ctx.unwrap();
let server_ctx = server_ctx.unwrap();
let config = DriverConfig::default();
let client: ConnectionHandle<Vec<u8>> = ConnectionDriver::spawn(
client_stream,
client_ctx,
ConnectionRole::Initiator { remote_static_pubkey: remote_pub },
Proto::TCP,
config,
);
let mut server: ConnectionHandle<Vec<u8>> = ConnectionDriver::spawn(
server_stream,
server_ctx,
ConnectionRole::Responder { identity },
Proto::TCP,
config,
);
for seq in 0..5u8 {
client
.send(ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Data, vec![seq]).unwrap())
.await
.unwrap();
}
client.rekey().await.unwrap();
for seq in 5..10u8 {
client
.send(ProtocolMessage::new(ConnectionParams::ENCRYPTED, MsgType::Data, vec![seq]).unwrap())
.await
.unwrap();
}
for seq in 0..10u8 {
match server.recv().await.unwrap() {
DriverMessage::Data(msg) => assert_eq!(msg.payload, vec![seq]),
other => panic!("unexpected message: {other:?}"),
}
}
assert!(server.rekey().await.is_err());
let _ = client.shutdown().await;
let _ = server.shutdown().await;
}
}