use tokio::sync::mpsc;
pub struct MessageSender<M> {
inner: mpsc::Sender<M>,
}
impl<M> MessageSender<M> {
pub(crate) fn new(inner: mpsc::Sender<M>) -> Self {
Self { inner }
}
pub async fn send(&self, msg: M) -> Result<(), MessageSendError<M>> {
self.inner
.send(msg)
.await
.map_err(|e| MessageSendError(e.0))
}
pub fn try_send(&self, msg: M) -> Result<(), TrySendError<M>> {
self.inner.try_send(msg).map_err(TrySendError::from_tokio)
}
pub fn is_closed(&self) -> bool {
self.inner.is_closed()
}
pub fn capacity(&self) -> usize {
self.inner.capacity()
}
pub fn max_capacity(&self) -> usize {
self.inner.max_capacity()
}
pub fn into_inner(self) -> mpsc::Sender<M> {
self.inner
}
}
impl<M> Clone for MessageSender<M> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<M> std::fmt::Debug for MessageSender<M> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MessageSender").finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct MessageSendError<T>(pub T);
impl<T> std::fmt::Display for MessageSendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "message sender: receiver dropped")
}
}
impl<T: std::fmt::Debug> std::error::Error for MessageSendError<T> {}
#[derive(Debug)]
pub enum TrySendError<T> {
Full(T),
Closed(T),
}
impl<T> TrySendError<T> {
pub fn into_inner(self) -> T {
match self {
Self::Full(t) | Self::Closed(t) => t,
}
}
fn from_tokio(err: mpsc::error::TrySendError<T>) -> Self {
match err {
mpsc::error::TrySendError::Full(t) => TrySendError::Full(t),
mpsc::error::TrySendError::Closed(t) => TrySendError::Closed(t),
}
}
}
impl<T> std::fmt::Display for TrySendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Full(_) => write!(f, "message sender: channel full"),
Self::Closed(_) => write!(f, "message sender: receiver dropped"),
}
}
}
impl<T: std::fmt::Debug> std::error::Error for TrySendError<T> {}
#[cfg(test)]
mod tests {
use super::*;
fn _assert_send_sync_clone<T: Send + Sync + Clone>() {}
fn _compile_assertions() {
_assert_send_sync_clone::<MessageSender<u32>>();
_assert_send_sync_clone::<MessageSender<String>>();
}
#[tokio::test]
async fn test_send_round_trip() {
let (tx, mut rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
sender.send(42).await.expect("send succeeds");
let received = rx.recv().await;
assert_eq!(received, Some(42));
}
#[tokio::test]
async fn test_try_send_full() {
let (tx, _rx) = mpsc::channel::<u32>(1);
let sender = MessageSender::new(tx);
sender.try_send(1).expect("first try_send succeeds");
let err = sender.try_send(2).expect_err("second try_send is Full");
assert!(matches!(err, TrySendError::Full(2)));
}
#[tokio::test]
async fn test_send_closed() {
let (tx, rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
drop(rx);
let err = sender.send(99).await.expect_err("send fails");
assert_eq!(err.0, 99);
}
#[tokio::test]
async fn test_try_send_closed() {
let (tx, rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
drop(rx);
let err = sender.try_send(99).expect_err("try_send fails");
assert!(matches!(err, TrySendError::Closed(99)));
}
#[tokio::test]
async fn test_is_closed() {
let (tx, rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
assert!(!sender.is_closed());
drop(rx);
assert!(sender.is_closed());
}
#[tokio::test]
async fn test_capacity_and_max_capacity() {
let (tx, _rx) = mpsc::channel::<u32>(4);
let sender = MessageSender::new(tx);
assert_eq!(sender.max_capacity(), 4);
assert_eq!(sender.capacity(), 4);
sender.try_send(1).expect("first try_send");
assert_eq!(sender.capacity(), 3);
assert_eq!(sender.max_capacity(), 4); }
#[tokio::test]
async fn test_clone_shares_channel() {
let (tx, mut rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
let cloned = sender.clone();
sender.send(1).await.expect("original send");
cloned.send(2).await.expect("cloned send");
assert_eq!(rx.recv().await, Some(1));
assert_eq!(rx.recv().await, Some(2));
}
#[tokio::test]
async fn test_into_inner_escape_hatch() {
let (tx, mut rx) = mpsc::channel::<u32>(16);
let sender = MessageSender::new(tx);
let inner: mpsc::Sender<u32> = sender.into_inner();
inner.send(42).await.expect("inner tokio send succeeds");
assert_eq!(rx.recv().await, Some(42));
}
#[test]
fn test_try_send_error_into_inner() {
let err = TrySendError::Full(42u32);
assert_eq!(err.into_inner(), 42);
let err = TrySendError::Closed(99u32);
assert_eq!(err.into_inner(), 99);
}
#[test]
fn test_message_send_error_message_recovered() {
let err = MessageSendError(42u32);
assert_eq!(err.0, 42);
}
#[test]
fn test_display_impls() {
let msg_err: MessageSendError<u32> = MessageSendError(42);
assert_eq!(msg_err.to_string(), "message sender: receiver dropped");
let full: TrySendError<u32> = TrySendError::Full(1);
assert_eq!(full.to_string(), "message sender: channel full");
let closed: TrySendError<u32> = TrySendError::Closed(2);
assert_eq!(closed.to_string(), "message sender: receiver dropped");
}
}