use tokio::sync::mpsc::UnboundedSender;
#[derive(Debug, PartialEq)]
pub(super) enum AckResult {
Ack(String),
Nack(String),
}
#[derive(Debug)]
#[non_exhaustive]
pub enum Handler {
AtLeastOnce(AtLeastOnce),
}
impl Handler {
pub fn ack(self) {
match self {
Handler::AtLeastOnce(h) => h.ack(),
}
}
}
#[derive(Debug)]
struct AtLeastOnceImpl {
ack_id: String,
ack_tx: UnboundedSender<AckResult>,
}
impl AtLeastOnceImpl {
fn ack(self) {
let _ = self.ack_tx.send(AckResult::Ack(self.ack_id));
}
fn nack(self) {
let _ = self.ack_tx.send(AckResult::Nack(self.ack_id));
}
}
#[derive(Debug)]
pub struct AtLeastOnce {
inner: Option<AtLeastOnceImpl>,
}
impl AtLeastOnce {
pub(super) fn new(ack_id: String, ack_tx: UnboundedSender<AckResult>) -> Self {
Self {
inner: Some(AtLeastOnceImpl { ack_id, ack_tx }),
}
}
pub fn ack(mut self) {
if let Some(inner) = self.inner.take() {
inner.ack();
}
}
#[cfg(test)]
pub(crate) fn ack_id(&self) -> &str {
self.inner
.as_ref()
.map(|i| i.ack_id.as_str())
.unwrap_or_default()
}
}
impl Drop for AtLeastOnce {
fn drop(&mut self) {
if let Some(inner) = self.inner.take() {
inner.nack();
}
}
}
#[cfg(test)]
mod tests {
use super::super::lease_state::tests::test_id;
use super::*;
use tokio::sync::mpsc::error::TryRecvError;
use tokio::sync::mpsc::unbounded_channel;
#[test]
fn handler_ack() -> anyhow::Result<()> {
let (ack_tx, mut ack_rx) = unbounded_channel();
let h = Handler::AtLeastOnce(AtLeastOnce::new(test_id(1), ack_tx));
assert_eq!(ack_rx.try_recv(), Err(TryRecvError::Empty));
h.ack();
let ack = ack_rx.try_recv()?;
assert_eq!(ack, AckResult::Ack(test_id(1)));
Ok(())
}
#[test]
fn handler_nack() -> anyhow::Result<()> {
let (ack_tx, mut ack_rx) = unbounded_channel();
let h = Handler::AtLeastOnce(AtLeastOnce::new(test_id(1), ack_tx));
assert_eq!(ack_rx.try_recv(), Err(TryRecvError::Empty));
drop(h);
let ack = ack_rx.try_recv()?;
assert_eq!(ack, AckResult::Nack(test_id(1)));
Ok(())
}
#[test]
fn at_least_once_ack() -> anyhow::Result<()> {
let (ack_tx, mut ack_rx) = unbounded_channel();
let h = AtLeastOnce::new(test_id(1), ack_tx);
assert_eq!(ack_rx.try_recv(), Err(TryRecvError::Empty));
h.ack();
let ack = ack_rx.try_recv()?;
assert_eq!(ack, AckResult::Ack(test_id(1)));
Ok(())
}
#[test]
fn at_least_once_nack() -> anyhow::Result<()> {
let (ack_tx, mut ack_rx) = unbounded_channel();
let h = AtLeastOnce::new(test_id(1), ack_tx);
assert_eq!(ack_rx.try_recv(), Err(TryRecvError::Empty));
drop(h);
let ack = ack_rx.try_recv()?;
assert_eq!(ack, AckResult::Nack(test_id(1)));
Ok(())
}
#[test]
fn at_least_once_drop_nacks() -> anyhow::Result<()> {
let (ack_tx, mut ack_rx) = unbounded_channel();
let h = AtLeastOnce::new(test_id(1), ack_tx);
assert_eq!(ack_rx.try_recv(), Err(TryRecvError::Empty));
drop(h);
let ack = ack_rx.try_recv()?;
assert_eq!(ack, AckResult::Nack(test_id(1)));
Ok(())
}
}