use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Semaphore;
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError;
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
use crate::config::AppConfig;
use crate::core::Broker;
use crate::core::Context;
use crate::core::Envelope;
use crate::core::Handler;
use crate::core::Outcome;
use crate::runtime::backoff::Backoff;
#[derive(Debug, Clone)]
pub(crate) struct HandlerTask<Tok> {
pub envelope: Envelope,
pub token: Tok,
}
pub(crate) struct HandlerPool<B, H>
where
B: Broker,
H: Handler<B>,
{
context: Context<B>,
handler: Arc<H>,
max_in_flight: usize,
max_attempts: u32,
backoff: Arc<Backoff>,
input_rx: mpsc::Receiver<HandlerTask<B::AckToken>>,
ack_retry_tx: mpsc::Sender<B::AckToken>,
}
impl<B, H> HandlerPool<B, H>
where
B: Broker,
H: Handler<B>,
{
pub(crate) fn new(
handler: H, context: Context<B>, config: &AppConfig, backoff: Arc<Backoff>,
ack_retry_tx: mpsc::Sender<B::AckToken>,
) -> (mpsc::Sender<HandlerTask<B::AckToken>>, Self) {
let max_in_flight = config.max_in_flight_per_handler.min(H::MAX_IN_FLIGHT).max(1);
let (input_tx, input_rx) = mpsc::channel(config.handler_queue_capacity.max(1));
let pool = Self {
context,
handler: Arc::new(handler),
max_in_flight,
max_attempts: config.handler_max_attempts.max(1),
backoff,
input_rx,
ack_retry_tx,
};
(input_tx, pool)
}
pub(crate) async fn run(mut self, cancellation: CancellationToken) {
let semaphore = Arc::new(Semaphore::new(self.max_in_flight));
let mut in_flight: JoinSet<()> = JoinSet::new();
loop {
while in_flight.try_join_next().is_some() {}
let task = tokio::select! {
biased;
_ = cancellation.cancelled() => break,
maybe = self.input_rx.recv() => match maybe {
Some(task) => task,
None => break,
},
};
let permit = tokio::select! {
_ = cancellation.cancelled() => break,
permit = Arc::clone(&semaphore).acquire_owned() => match permit {
Ok(permit) => permit,
Err(_) => break,
},
};
let handler = Arc::clone(&self.handler);
let context = self.context.clone();
let ack_retry_tx = self.ack_retry_tx.clone();
let backoff = Arc::clone(&self.backoff);
let max_attempts = self.max_attempts;
let task_cancellation = cancellation.clone();
in_flight.spawn(async move {
process_item::<B, H>(
&handler,
&context,
task,
max_attempts,
&backoff,
&ack_retry_tx,
&task_cancellation,
)
.await;
drop(permit);
});
}
while in_flight.join_next().await.is_some() {}
}
}
async fn process_item<B, H>(
handler: &H, context: &Context<B>, task: HandlerTask<B::AckToken>, max_attempts: u32, backoff: &Backoff,
ack_retry_tx: &mpsc::Sender<B::AckToken>, cancellation: &CancellationToken,
) where
B: Broker,
H: Handler<B>,
{
let HandlerTask { envelope, token } = task;
let mut attempt: u32 = 1;
loop {
let outcome = match handler.handle(context, &envelope).await {
Ok(outcome) => outcome,
Err(error) => {
tracing::warn!(id = %envelope.id, %error, "handler returned an error; scheduling retry");
Outcome::Retry { after_ms: 0 }
}
};
match outcome {
Outcome::Ack | Outcome::Drop => {
finish_ack::<B>(context, token, ack_retry_tx).await;
return;
}
Outcome::DeadLetter { reason } => {
dead_letter::<B>(context, &envelope, &reason, token, ack_retry_tx).await;
return;
}
Outcome::Retry { after_ms } => {
if attempt >= max_attempts {
tracing::warn!(id = %envelope.id, attempt, "exhausted handler attempts; dead-lettering");
dead_letter::<B>(context, &envelope, "max retry attempts exceeded", token, ack_retry_tx).await;
return;
}
let delay = if after_ms > 0 { Duration::from_millis(after_ms) } else { backoff.delay(attempt) };
attempt += 1;
tokio::select! {
_ = cancellation.cancelled() => return,
_ = tokio::time::sleep(delay) => {}
}
}
}
}
}
async fn finish_ack<B: Broker>(context: &Context<B>, token: B::AckToken, ack_retry_tx: &mpsc::Sender<B::AckToken>) {
let retry_token = token.clone();
if let Err(error) = context.broker().ack(token).await {
tracing::warn!(%error, "ack failed; queueing for retry");
enqueue_failed_ack(ack_retry_tx, retry_token);
}
}
async fn dead_letter<B: Broker>(
context: &Context<B>, envelope: &Envelope, reason: &str, token: B::AckToken,
ack_retry_tx: &mpsc::Sender<B::AckToken>,
) {
match context.broker().dead_letter(envelope, reason).await {
Ok(()) => finish_ack::<B>(context, token, ack_retry_tx).await,
Err(error) => {
tracing::error!(id = %envelope.id, %error, "dead-letter write failed; leaving message pending");
}
}
}
fn enqueue_failed_ack<Tok>(ack_retry_tx: &mpsc::Sender<Tok>, token: Tok) {
match ack_retry_tx.try_send(token) {
Ok(()) => {}
Err(TrySendError::Full(_)) => {
tracing::error!("ack retry queue full; dropped failed ack token (message may be redelivered)");
}
Err(TrySendError::Closed(_)) => {
tracing::error!("ack retry queue closed; dropped failed ack token");
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tokio::sync::mpsc;
use tokio::time::timeout;
use tokio_util::sync::CancellationToken;
use super::HandlerPool;
use super::HandlerTask;
use crate::config::AppConfig;
use crate::core::Context;
use crate::core::Envelope;
use crate::core::Handler;
use crate::core::HandlerError;
use crate::core::Outcome;
use crate::runtime::backoff::Backoff;
use crate::test_support::MockBroker;
use crate::test_support::make_envelope;
use crate::test_support::token;
fn pool_for<H: Handler<MockBroker>>(
broker: MockBroker, handler: H, config: &AppConfig,
) -> (mpsc::Sender<HandlerTask<String>>, HandlerPool<MockBroker, H>, mpsc::Receiver<String>) {
let (ack_retry_tx, ack_retry_rx) = mpsc::channel(16);
let backoff = Arc::new(Backoff::new(config.backoff.clone()));
let (tx, pool) = HandlerPool::new(handler, Context::new(broker), config, backoff, ack_retry_tx);
(tx, pool, ack_retry_rx)
}
struct FixedOutcome(Outcome);
impl Handler<MockBroker> for FixedOutcome {
const STREAMS: &'static [&'static str] = &["orders"];
async fn handle(&self, _ctx: &Context<MockBroker>, _msg: &Envelope) -> Result<Outcome, HandlerError> {
Ok(self.0.clone())
}
}
async fn drive(handler: impl Handler<MockBroker>, ids: &[&str]) -> MockBroker {
let broker = MockBroker::default();
let config = AppConfig { handler_max_attempts: 3, ..AppConfig::default() };
let (tx, pool, _ack_rx) = pool_for(broker.clone(), handler, &config);
for id in ids {
tx.send(HandlerTask { envelope: make_envelope("orders", id), token: token(id) }).await.unwrap();
}
drop(tx);
timeout(Duration::from_secs(1), pool.run(CancellationToken::new())).await.expect("pool should drain");
broker
}
#[tokio::test]
async fn ack_and_drop_acknowledge() {
let acked = drive(FixedOutcome(Outcome::Ack), &["1-1"]).await;
assert_eq!(acked.ack_count(), 1);
let dropped = drive(FixedOutcome(Outcome::Drop), &["1-1"]).await;
assert_eq!(dropped.ack_count(), 1);
}
#[tokio::test]
async fn dead_letter_writes_then_acks() {
let broker = drive(FixedOutcome(Outcome::DeadLetter { reason: "poison".to_string() }), &["1-1"]).await;
assert_eq!(broker.ack_count(), 1);
let dlq = broker.dead_letters();
assert_eq!(dlq.len(), 1);
assert_eq!(dlq[0].1, "poison");
}
#[tokio::test]
async fn retry_exhaustion_dead_letters() {
let broker = drive(FixedOutcome(Outcome::Retry { after_ms: 0 }), &["1-1"]).await;
assert_eq!(broker.ack_count(), 1);
let dlq = broker.dead_letters();
assert_eq!(dlq.len(), 1);
assert_eq!(dlq[0].1, "max retry attempts exceeded");
}
#[tokio::test]
async fn retry_then_success_acks_without_dead_letter() {
struct FailThenOk(AtomicUsize);
impl Handler<MockBroker> for FailThenOk {
const STREAMS: &'static [&'static str] = &["orders"];
async fn handle(&self, _ctx: &Context<MockBroker>, _msg: &Envelope) -> Result<Outcome, HandlerError> {
if self.0.fetch_add(1, Ordering::SeqCst) == 0 {
Ok(Outcome::Retry { after_ms: 0 })
} else {
Ok(Outcome::Ack)
}
}
}
let broker = drive(FailThenOk(AtomicUsize::new(0)), &["1-1"]).await;
assert_eq!(broker.ack_count(), 1);
assert!(broker.dead_letters().is_empty());
}
#[tokio::test]
async fn failed_ack_is_enqueued_for_retry() {
let broker = MockBroker::default();
broker.fail_next_acks(1);
let config = AppConfig::default();
let (tx, pool, mut ack_rx) = pool_for(broker.clone(), FixedOutcome(Outcome::Ack), &config);
tx.send(HandlerTask { envelope: make_envelope("orders", "1-1"), token: token("1-1") }).await.unwrap();
drop(tx);
timeout(Duration::from_secs(1), pool.run(CancellationToken::new())).await.unwrap();
let requeued = ack_rx.try_recv().expect("failed ack should be queued for retry");
assert_eq!(requeued, "1-1");
}
}