228 lines
7.2 KiB
Rust
228 lines
7.2 KiB
Rust
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<u8>,
|
|
}
|
|
impl OutboundMessage {
|
|
pub fn of(message: &dyn Message) -> Result<Self, MessagingError> {
|
|
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<u8>) -> Self {
|
|
Self {
|
|
message_type,
|
|
message_version,
|
|
payload,
|
|
}
|
|
}
|
|
}
|
|
pub struct Incoming {
|
|
pub header: FrameHeader,
|
|
pub payload: Vec<u8>,
|
|
pub message: Option<Box<dyn Message>>,
|
|
}
|
|
impl Incoming {
|
|
pub fn message_type(&self) -> u16 {
|
|
self.header.message_type
|
|
}
|
|
pub fn downcast<T: 'static>(&self) -> Option<&T> {
|
|
self.message
|
|
.as_ref()
|
|
.and_then(|message| message.as_any().downcast_ref::<T>())
|
|
}
|
|
}
|
|
#[derive(Clone)]
|
|
pub struct MessagingSender {
|
|
sender: mpsc::Sender<OutboundMessage>,
|
|
}
|
|
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<dyn StreamEncrypter>,
|
|
factory: Arc<dyn MessageFactory>,
|
|
config: MessagingConfig,
|
|
sender: Option<MessagingSender>,
|
|
writer_task: Option<JoinHandle<()>>,
|
|
}
|
|
impl Messaging {
|
|
pub fn attach(
|
|
stream: TcpStream,
|
|
cipher: SessionCipher,
|
|
factory: Arc<dyn MessageFactory>,
|
|
config: MessagingConfig,
|
|
) -> std::io::Result<Self> {
|
|
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<Option<Incoming>, 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<dyn MessageFactory> {
|
|
&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<OutboundMessage>,
|
|
mut outbound: Box<dyn StreamEncrypter>,
|
|
) {
|
|
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;
|
|
}
|