use super::SendError;
pub fn broadcast_channel<T: Clone>(
capacity: usize,
) -> std::io::Result<(BroadcastSender<T>, BroadcastReceiver<T>)> {
if !capacity.is_power_of_two() || capacity > 1_048_576 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"broadcast capacity must be a power of two from 1 through 1048576",
));
}
let (sender, receiver) = tokio::sync::broadcast::channel(capacity);
Ok((
BroadcastSender { inner: sender },
BroadcastReceiver { inner: receiver },
))
}
#[derive(Debug)]
pub struct BroadcastSender<T> {
inner: tokio::sync::broadcast::Sender<T>,
}
impl<T> Clone for BroadcastSender<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<T: Clone> BroadcastSender<T> {
pub fn send(&self, value: T) -> Result<usize, SendError<T>> {
self.inner.send(value).map_err(|error| SendError(error.0))
}
pub fn subscribe(&self) -> BroadcastReceiver<T> {
BroadcastReceiver {
inner: self.inner.subscribe(),
}
}
}
#[derive(Debug)]
pub struct BroadcastReceiver<T> {
inner: tokio::sync::broadcast::Receiver<T>,
}
impl<T: Clone> BroadcastReceiver<T> {
pub async fn recv(&mut self) -> Result<T, BroadcastRecvError> {
self.inner.recv().await.map_err(|error| match error {
tokio::sync::broadcast::error::RecvError::Lagged(n) => BroadcastRecvError::Lagged(n),
tokio::sync::broadcast::error::RecvError::Closed => BroadcastRecvError::Closed,
})
}
pub fn try_recv(&mut self) -> Result<T, BroadcastTryRecvError> {
self.inner.try_recv().map_err(|error| match error {
tokio::sync::broadcast::error::TryRecvError::Empty => BroadcastTryRecvError::Empty,
tokio::sync::broadcast::error::TryRecvError::Lagged(n) => {
BroadcastTryRecvError::Lagged(n)
}
tokio::sync::broadcast::error::TryRecvError::Closed => BroadcastTryRecvError::Closed,
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum BroadcastRecvError {
#[error("broadcast receiver skipped {0} values")]
Lagged(u64),
#[error("broadcast channel is closed")]
Closed,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum BroadcastTryRecvError {
#[error("broadcast channel is empty")]
Empty,
#[error("broadcast receiver skipped {0} values")]
Lagged(u64),
#[error("broadcast channel is closed")]
Closed,
}
#[cfg(feature = "event-stream")]
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
#[error("broadcast receiver skipped {skipped} values")]
pub struct BroadcastLagged {
pub skipped: u64,
}
#[cfg(feature = "event-stream")]
pub struct BroadcastStream<T, F> {
inner: tokio_stream::wrappers::BroadcastStream<T>,
map: F,
}
#[cfg(feature = "event-stream")]
impl<T: Clone + Send + 'static> BroadcastReceiver<T> {
pub fn into_stream_with<U, F>(self, map: F) -> BroadcastStream<T, F>
where
F: FnMut(Result<T, BroadcastLagged>) -> Option<U> + Unpin,
{
BroadcastStream {
inner: tokio_stream::wrappers::BroadcastStream::new(self.inner),
map,
}
}
}
#[cfg(feature = "event-stream")]
impl<T, U, F> futures_core::Stream for BroadcastStream<T, F>
where
T: Clone + Send + 'static,
F: FnMut(Result<T, BroadcastLagged>) -> Option<U> + Unpin,
{
type Item = U;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<U>> {
let this = self.get_mut();
loop {
let result = match std::task::ready!(std::pin::Pin::new(&mut this.inner).poll_next(cx))
{
Some(result) => result,
None => return std::task::Poll::Ready(None),
};
let result = result.map_err(|error| match error {
tokio_stream::wrappers::errors::BroadcastStreamRecvError::Lagged(skipped) => {
BroadcastLagged { skipped }
}
});
if let Some(value) = (this.map)(result) {
return std::task::Poll::Ready(Some(value));
}
}
}
}