use std::{
pin::Pin,
task::{Context, Poll},
};
use futures::{ready, SinkExt};
use nym_client_core::client::inbound_messages::InputMessage;
use nym_sphinx::{
addressing::Recipient, anonymous_replies::requests::AnonymousSenderTag, params::PacketType,
};
use nym_task::connections::TransmissionLane;
use tokio::{io::AsyncWrite, sync::mpsc, task::JoinHandle};
use tokio_util::sync::PollSender;
use crate::Error;
use super::{IncludedSurbs, MixnetMessageSender};
const SINK_BUFFER_SIZE_IN_MESSAGES: usize = 8;
pub trait MixnetMessageSinkTranslator: Unpin {
#[allow(clippy::result_large_err)]
fn to_input_message(&self, bytes: &[u8]) -> Result<InputMessage, Error>;
}
#[derive(Clone, Debug)]
pub struct DefaultMixnetMessageSinkTranslator {
destination: SinkDestination,
lane: TransmissionLane,
packet_type: Option<PacketType>,
}
#[derive(Clone, Debug)]
enum SinkDestination {
Recipient {
recipient: Box<Recipient>,
surbs: IncludedSurbs,
},
Reply(AnonymousSenderTag),
}
impl MixnetMessageSinkTranslator for DefaultMixnetMessageSinkTranslator {
fn to_input_message(&self, bytes: &[u8]) -> Result<InputMessage, Error> {
let bytes = bytes.to_vec();
match &self.destination {
SinkDestination::Recipient { recipient, surbs } => match surbs {
IncludedSurbs::ExposeSelfAddress => Ok(InputMessage::new_regular(
**recipient,
bytes,
self.lane,
self.packet_type,
)),
IncludedSurbs::Amount(surbs) => Ok(InputMessage::new_anonymous(
**recipient,
bytes,
*surbs,
self.lane,
self.packet_type,
)),
},
SinkDestination::Reply(tag) => Ok(InputMessage::new_reply(
*tag,
bytes,
self.lane,
self.packet_type,
)),
}
}
}
pub struct MixnetMessageSink<F>
where
F: MixnetMessageSinkTranslator,
{
message_translator: F,
tx: PollSender<InputMessage>,
send_task: JoinHandle<()>,
}
impl MixnetMessageSink<DefaultMixnetMessageSinkTranslator> {
pub fn new_recipient_sink<Sender>(
mixnet_client_sender: Sender,
recipient: Recipient,
surbs: IncludedSurbs,
) -> Self
where
Sender: MixnetMessageSender + Send + 'static,
{
let destination = SinkDestination::Recipient {
recipient: Box::new(recipient),
surbs,
};
let translator = DefaultMixnetMessageSinkTranslator {
destination,
lane: TransmissionLane::General,
packet_type: None,
};
Self::new_with_custom_translator(mixnet_client_sender, translator)
}
pub fn new_reply_sink<Sender>(
mixnet_client_sender: Sender,
recipient_tag: AnonymousSenderTag,
) -> Self
where
Sender: MixnetMessageSender + Send + 'static,
{
let destination = SinkDestination::Reply(recipient_tag);
let translator = DefaultMixnetMessageSinkTranslator {
destination,
lane: TransmissionLane::General,
packet_type: None,
};
Self::new_with_custom_translator(mixnet_client_sender, translator)
}
}
impl<F> MixnetMessageSink<F>
where
F: MixnetMessageSinkTranslator,
{
pub fn new_with_custom_translator<Sender>(
mixnet_client_sender: Sender,
message_translator: F,
) -> Self
where
Sender: MixnetMessageSender + Send + 'static,
{
let (tx, send_task) = Self::start_sender_task(mixnet_client_sender);
let tx = PollSender::new(tx);
MixnetMessageSink {
message_translator,
tx,
send_task,
}
}
fn start_sender_task<Sender>(
mixnet_client_sender: Sender,
) -> (mpsc::Sender<InputMessage>, JoinHandle<()>)
where
Sender: MixnetMessageSender + Send + 'static,
{
let (tx, mut rx) = mpsc::channel(SINK_BUFFER_SIZE_IN_MESSAGES);
let send_task = tokio::spawn(async move {
while let Some(input_message) = rx.recv().await {
if let Err(err) = mixnet_client_sender.send(input_message).await {
log::error!("failed to send packet to mixnet: {err}");
}
}
});
(tx, send_task)
}
}
impl<F> Drop for MixnetMessageSink<F>
where
F: MixnetMessageSinkTranslator,
{
fn drop(&mut self) {
self.send_task.abort();
}
}
impl<F> AsyncWrite for MixnetMessageSink<F>
where
F: MixnetMessageSinkTranslator,
{
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::result::Result<usize, std::io::Error>> {
ready!(self.tx.poll_ready_unpin(cx))
.map_err(|_| std::io::Error::other("failed to send packet to mixnet"))?;
let input_message = self
.message_translator
.to_input_message(buf)
.map_err(std::io::Error::other)?;
self.tx
.start_send_unpin(input_message)
.map_err(|_| std::io::Error::other("failed to send packet to mixnet"))?;
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(
mut self: Pin<&mut Self>,
cx: &mut Context,
) -> Poll<std::result::Result<(), std::io::Error>> {
ready!(self.tx.poll_flush_unpin(cx))
.map_err(|_| std::io::Error::other("failed to send packet to mixnet"))?;
Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context,
) -> Poll<std::result::Result<(), std::io::Error>> {
self.poll_flush(cx)
}
}