use crate::crypto::{SessionCipher, StreamEncrypter}; use crate::error::Error; use crate::factory::MessageFactory; use crate::frame::{FrameCodec, FrameHeader, MessageFrame, HEADER_LEN}; use crate::io::ByteStreamWriter; use crate::message::Message; use std::sync::Arc; use std::time::Duration; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; use tokio::net::TcpStream; use tokio::sync::mpsc; use tokio::task::JoinHandle; use tokio::time::timeout; #[derive(Debug, thiserror::Error)] pub enum MessagingError { #[error("io: {0}")] Io(#[from] std::io::Error), #[error("codec: {0}")] Codec(#[from] Error), #[error("payload of {actual} byte(s) exceeds the {limit} byte limit")] PayloadTooLarge { actual: usize, limit: usize }, #[error("peer stayed idle for longer than the read timeout")] IdleTimeout, #[error("the write half is gone")] WriterGone, } #[derive(Debug, Clone)] pub struct MessagingConfig { pub read_timeout: Duration, pub max_payload_len: usize, pub outbound_queue_len: usize, } impl Default for MessagingConfig { fn default() -> Self { Self { read_timeout: Duration::from_secs(90), max_payload_len: 1024 * 1024, outbound_queue_len: 64, } } } #[derive(Debug, Clone, PartialEq, Eq)] pub struct OutboundMessage { pub message_type: u16, pub message_version: u16, pub payload: Vec, } impl OutboundMessage { pub fn of(message: &dyn Message) -> Result { let mut writer = ByteStreamWriter::with_capacity(256); message.encode_payload(&mut writer)?; Ok(Self { message_type: message.message_type(), message_version: message.message_version(), payload: writer.into_inner(), }) } pub fn raw(message_type: u16, message_version: u16, payload: Vec) -> Self { Self { message_type, message_version, payload, } } } pub struct Incoming { pub header: FrameHeader, pub payload: Vec, pub message: Option>, } impl Incoming { pub fn message_type(&self) -> u16 { self.header.message_type } pub fn downcast(&self) -> Option<&T> { self.message .as_ref() .and_then(|message| message.as_any().downcast_ref::()) } } #[derive(Clone)] pub struct MessagingSender { sender: mpsc::Sender, } impl MessagingSender { pub async fn send(&self, message: OutboundMessage) -> Result<(), MessagingError> { self.sender .send(message) .await .map_err(|_| MessagingError::WriterGone) } pub async fn send_message(&self, message: &dyn Message) -> Result<(), MessagingError> { self.send(OutboundMessage::of(message)?).await } pub fn is_closed(&self) -> bool { self.sender.is_closed() } } pub struct Messaging { reader: OwnedReadHalf, inbound: Box, factory: Arc, config: MessagingConfig, sender: Option, writer_task: Option>, } impl Messaging { pub fn attach( stream: TcpStream, cipher: SessionCipher, factory: Arc, config: MessagingConfig, ) -> std::io::Result { stream.set_nodelay(true)?; let (reader, writer) = stream.into_split(); let (inbound, outbound) = cipher.split(); let (sender, receiver) = mpsc::channel(config.outbound_queue_len.max(1)); let writer_task = tokio::spawn(write_loop(writer, receiver, outbound)); Ok(Self { reader, inbound, factory, config, sender: Some(MessagingSender { sender }), writer_task: Some(writer_task), }) } pub fn sender(&self) -> MessagingSender { self.sender .clone() .expect("messaging sender is available until shutdown") } pub async fn next_message(&mut self) -> Result, MessagingError> { let mut header_bytes = [0u8; HEADER_LEN]; match timeout( self.config.read_timeout, self.reader.read_exact(&mut header_bytes), ) .await { Err(_) => return Err(MessagingError::IdleTimeout), Ok(Err(error)) if error.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None), Ok(Err(error)) => return Err(error.into()), Ok(Ok(_)) => {} } let header = FrameHeader::parse(&header_bytes); if header.payload_len > self.config.max_payload_len { return Err(MessagingError::PayloadTooLarge { actual: header.payload_len, limit: self.config.max_payload_len, }); } let mut cipher_text = vec![0u8; header.payload_len]; timeout( self.config.read_timeout, self.reader.read_exact(&mut cipher_text), ) .await .map_err(|_| MessagingError::IdleTimeout)??; let payload = self.inbound.decrypt(&cipher_text)?; let message = match self .factory .create(header.message_type, header.message_version, &payload) { Ok(message) => Some(message), Err(error) => { tracing::warn!( message_type = header.message_type, bytes = payload.len(), %error, "ignoring message of unknown type {}", header.message_type ); None } }; Ok(Some(Incoming { header, payload, message, })) } pub fn factory(&self) -> &Arc { &self.factory } pub async fn shutdown(mut self) { self.sender.take(); if let Some(task) = self.writer_task.take() { let _ = task.await; } } } impl Drop for Messaging { fn drop(&mut self) { if let Some(task) = self.writer_task.take() { task.abort(); } } } async fn write_loop( mut writer: OwnedWriteHalf, mut receiver: mpsc::Receiver, mut outbound: Box, ) { while let Some(message) = receiver.recv().await { let cipher_text = match outbound.encrypt(&message.payload) { Ok(bytes) => bytes, Err(error) => { tracing::error!(%error, "failed to encrypt an outbound payload"); break; } }; let frame = MessageFrame::new(message.message_type, message.message_version, cipher_text); let encoded = match FrameCodec::encode(&frame) { Ok(bytes) => bytes, Err(error) => { tracing::error!(%error, "failed to frame an outbound message"); break; } }; if writer.write_all(&encoded).await.is_err() || writer.flush().await.is_err() { break; } } let _ = writer.shutdown().await; }