scroll.server/crates/titan/src/net/messaging.rs

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;
}