Skip to main content

mq_bridge/
event_handler.rs

1//  mq-bridge
2//  © Copyright 2025, by Marco Mengelkoch
3//  Licensed under MIT License, see License file for more details
4//  git clone https://github.com/marcomq/mq-bridge
5
6use crate::errors::PublisherError;
7use crate::traits::{Handler, HandlerError, MessagePublisher};
8use crate::CanonicalMessage;
9use async_trait::async_trait;
10use std::any::Any;
11use std::sync::Arc;
12
13use crate::traits::{Sent, SentBatch};
14
15/// A publisher middleware that intercepts messages and passes them to a `Handler`.
16/// This middleware is terminal; it consumes the message and does not pass it to an inner publisher.
17pub struct EventPublisher {
18    handler: Arc<dyn Handler>,
19}
20
21impl EventPublisher {
22    pub fn new(handler: impl Handler + 'static) -> Self {
23        Self {
24            handler: Arc::new(handler),
25        }
26    }
27}
28
29#[async_trait]
30impl MessagePublisher for EventPublisher {
31    async fn send(&self, message: CanonicalMessage) -> Result<Sent, PublisherError> {
32        match self.handler.handle(message).await {
33            Ok(_) => Ok(Sent::Ack), // Ignore result (Ack or Publish), just Ack.
34            Err(e) => Err(e),       // Converts HandlerError to PublisherError
35        }
36    }
37
38    async fn send_batch(
39        &self,
40        messages: Vec<CanonicalMessage>,
41    ) -> Result<SentBatch, PublisherError> {
42        let results = self.handler.handle_many(messages.clone()).await;
43        if results.len() != messages.len() {
44            return Err(PublisherError::NonRetryable(anyhow::anyhow!(
45                "handler returned {} results for {} messages",
46                results.len(),
47                messages.len()
48            )));
49        }
50
51        let mut failed = Vec::new();
52        let mut iter = messages.into_iter().zip(results);
53        while let Some((message, result)) = iter.next() {
54            match result {
55                Ok(_) => {}
56                Err(HandlerError::NonRetryable(err)) => {
57                    failed.push((message, PublisherError::NonRetryable(err)));
58                }
59                Err(HandlerError::Retryable(err)) => {
60                    failed.push((message, PublisherError::Retryable(err)));
61                    for (remaining, _) in iter {
62                        failed.push((
63                            remaining,
64                            PublisherError::Retryable(anyhow::anyhow!(
65                                "Batch aborted due to previous error"
66                            )),
67                        ));
68                    }
69                    break;
70                }
71                Err(HandlerError::Connection(err)) => {
72                    failed.push((message, PublisherError::Connection(err)));
73                    for (remaining, _) in iter {
74                        failed.push((
75                            remaining,
76                            PublisherError::Connection(anyhow::anyhow!(
77                                "Batch aborted due to previous connection error"
78                            )),
79                        ));
80                    }
81                    break;
82                }
83            }
84        }
85
86        if failed.is_empty() {
87            Ok(SentBatch::Ack)
88        } else {
89            Ok(SentBatch::Partial {
90                responses: None,
91                failed,
92            })
93        }
94    }
95
96    fn as_any(&self) -> &dyn Any {
97        self
98    }
99}
100
101#[cfg(test)]
102mod tests {
103    use super::*;
104    use crate::traits::Handled;
105    use std::sync::atomic::{AtomicBool, Ordering};
106
107    #[tokio::test]
108    async fn test_event_handler() {
109        let event_handled = Arc::new(AtomicBool::new(false));
110        let handler = Arc::new({
111            let flag = event_handled.clone();
112            move |_msg: CanonicalMessage| {
113                let flag_clone = flag.clone();
114                async move {
115                    flag_clone.store(true, Ordering::SeqCst);
116                    Ok(Handled::Ack)
117                }
118            }
119        });
120        let publisher = EventPublisher::new(handler);
121        publisher
122            .send(CanonicalMessage::new(b"event1".to_vec(), None))
123            .await
124            .unwrap();
125        assert!(event_handled.load(Ordering::SeqCst));
126    }
127
128    #[tokio::test]
129    async fn test_event_handler_send_batch_retryable_error_aborts_remainder() {
130        struct BatchHandler;
131
132        #[async_trait]
133        impl Handler for BatchHandler {
134            async fn handle(&self, _msg: CanonicalMessage) -> Result<Handled, HandlerError> {
135                unreachable!("send_batch should use handle_many")
136            }
137
138            async fn handle_many(
139                &self,
140                msgs: Vec<CanonicalMessage>,
141            ) -> Vec<Result<Handled, HandlerError>> {
142                msgs.into_iter()
143                    .map(|msg| {
144                        if msg.get_payload_str() == "two" {
145                            Err(HandlerError::Retryable(anyhow::anyhow!(
146                                "temporary failure"
147                            )))
148                        } else {
149                            Ok(Handled::Ack)
150                        }
151                    })
152                    .collect()
153            }
154        }
155
156        let publisher = EventPublisher::new(BatchHandler);
157        let result = publisher
158            .send_batch(vec!["one".into(), "two".into(), "three".into()])
159            .await
160            .unwrap();
161
162        match result {
163            SentBatch::Partial { responses, failed } => {
164                assert!(responses.is_none());
165                assert_eq!(failed.len(), 2);
166                assert_eq!(failed[0].0.get_payload_str(), "two");
167                assert_eq!(failed[1].0.get_payload_str(), "three");
168                assert!(matches!(failed[0].1, PublisherError::Retryable(_)));
169                assert!(matches!(failed[1].1, PublisherError::Retryable(_)));
170            }
171            other => panic!("expected partial failure, got {other:?}"),
172        }
173    }
174}