Skip to main content

rskit_messaging/
handler.rs

1//! Message handler trait and middleware chain for consumed messages.
2
3use std::sync::Arc;
4
5use async_trait::async_trait;
6use rskit_errors::AppResult;
7
8use crate::message::Message;
9
10/// Handler for processing consumed messages.
11///
12/// Implement this trait to define how incoming messages are processed. For simple cases,
13/// use [`FnHandler`] to wrap a closure.
14#[async_trait]
15pub trait MessageHandler<T: Send + Sync + 'static>: Send + Sync + 'static {
16    /// Process a single message.
17    async fn handle(&self, msg: Message<T>) -> AppResult<()>;
18}
19
20/// Middleware that wraps a handler with cross-cutting concerns.
21///
22/// Each middleware receives the next handler in the chain
23/// and returns a new handler that adds behaviour around it (logging, metrics, error recovery, etc.).
24pub trait HandlerMiddleware<T: Send + Sync + 'static>: Send + Sync + 'static {
25    /// Wrap the given handler, returning a new handler that adds middleware behaviour.
26    fn wrap(&self, next: Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>>;
27}
28
29/// Chains middleware around a base handler.
30///
31/// Middlewares are applied in order: the first middleware in the slice becomes the outermost wrapper.
32///
33/// ```text
34/// chain_handlers(base, [mw_a, mw_b])
35///   => mw_a(mw_b(base))
36/// ```
37pub fn chain_handlers<T: Send + Sync + 'static>(
38    base: Arc<dyn MessageHandler<T>>,
39    middlewares: &[Arc<dyn HandlerMiddleware<T>>],
40) -> Arc<dyn MessageHandler<T>> {
41    let mut handler = base;
42    for mw in middlewares.iter().rev() {
43        handler = mw.wrap(handler);
44    }
45    handler
46}
47
48/// Adapter that turns a closure into a [`HandlerMiddleware`].
49///
50/// Useful for simple one-off middleware that do not need their own struct.
51///
52/// # Example
53///
54/// ```rust,ignore
55/// use rskit_messaging::handler::middleware_fn;
56/// use std::sync::Arc;
57///
58/// let mw = middleware_fn(|next| {
59///     // wrap next handler with custom behaviour
60///     next
61/// });
62/// ```
63pub fn middleware_fn<T, F>(f: F) -> impl HandlerMiddleware<T>
64where
65    T: Send + Sync + 'static,
66    F: Fn(Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>> + Send + Sync + 'static,
67{
68    FnMiddleware {
69        func: f,
70        _marker: std::marker::PhantomData,
71    }
72}
73
74struct FnMiddleware<T, F> {
75    func: F,
76    _marker: std::marker::PhantomData<T>,
77}
78
79impl<T, F> HandlerMiddleware<T> for FnMiddleware<T, F>
80where
81    T: Send + Sync + 'static,
82    F: Fn(Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>> + Send + Sync + 'static,
83{
84    fn wrap(&self, next: Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>> {
85        (self.func)(next)
86    }
87}
88
89/// Adapter that turns an async closure into a [`MessageHandler`].
90///
91/// # Example
92///
93/// ```rust,ignore
94/// use rskit_messaging::handler::FnHandler;
95///
96/// let handler = FnHandler::new(|msg| async move {
97///     println!("got: {:?}", msg.payload);
98///     Ok(())
99/// });
100/// ```
101pub struct FnHandler<T, F, Fut>
102where
103    T: Send + Sync + 'static,
104    F: Fn(Message<T>) -> Fut + Send + Sync + 'static,
105    Fut: std::future::Future<Output = AppResult<()>> + Send + 'static,
106{
107    func: F,
108    _marker: std::marker::PhantomData<T>,
109}
110
111impl<T, F, Fut> FnHandler<T, F, Fut>
112where
113    T: Send + Sync + 'static,
114    F: Fn(Message<T>) -> Fut + Send + Sync + 'static,
115    Fut: std::future::Future<Output = AppResult<()>> + Send + 'static,
116{
117    /// Create a new function handler from the given closure.
118    pub const fn new(func: F) -> Self {
119        Self {
120            func,
121            _marker: std::marker::PhantomData,
122        }
123    }
124}
125
126#[async_trait]
127impl<T, F, Fut> MessageHandler<T> for FnHandler<T, F, Fut>
128where
129    T: Send + Sync + 'static,
130    F: Fn(Message<T>) -> Fut + Send + Sync + 'static,
131    Fut: std::future::Future<Output = AppResult<()>> + Send + 'static,
132{
133    async fn handle(&self, msg: Message<T>) -> AppResult<()> {
134        (self.func)(msg).await
135    }
136}
137
138#[cfg(test)]
139mod tests {
140    use std::sync::atomic::{AtomicU32, Ordering};
141
142    use super::*;
143
144    #[tokio::test]
145    async fn fn_handler_processes_message() {
146        let counter = Arc::new(AtomicU32::new(0));
147        let c = counter.clone();
148        let handler = FnHandler::new(move |_msg: Message<String>| {
149            let c = c.clone();
150            async move {
151                c.fetch_add(1, Ordering::SeqCst);
152                Ok(())
153            }
154        });
155
156        let msg = Message::new("t", "hello".to_string());
157        handler.handle(msg).await.unwrap();
158        assert_eq!(counter.load(Ordering::SeqCst), 1);
159    }
160
161    /// A middleware that increments a counter before delegating.
162    struct CountingMiddleware {
163        counter: Arc<AtomicU32>,
164    }
165
166    struct CountingHandler<T: Send + Sync + 'static> {
167        counter: Arc<AtomicU32>,
168        next: Arc<dyn MessageHandler<T>>,
169    }
170
171    #[async_trait]
172    impl<T: Send + Sync + 'static> MessageHandler<T> for CountingHandler<T> {
173        async fn handle(&self, msg: Message<T>) -> AppResult<()> {
174            self.counter.fetch_add(1, Ordering::SeqCst);
175            self.next.handle(msg).await
176        }
177    }
178
179    impl<T: Send + Sync + 'static> HandlerMiddleware<T> for CountingMiddleware {
180        fn wrap(&self, next: Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>> {
181            Arc::new(CountingHandler {
182                counter: self.counter.clone(),
183                next,
184            })
185        }
186    }
187
188    #[tokio::test]
189    async fn middleware_chain_applies_in_order() {
190        let mw_counter = Arc::new(AtomicU32::new(0));
191        let base_counter = Arc::new(AtomicU32::new(0));
192
193        let bc = base_counter.clone();
194        let base: Arc<dyn MessageHandler<String>> =
195            Arc::new(FnHandler::new(move |_msg: Message<String>| {
196                let bc = bc.clone();
197                async move {
198                    bc.fetch_add(1, Ordering::SeqCst);
199                    Ok(())
200                }
201            }));
202
203        let mw: Arc<dyn HandlerMiddleware<String>> = Arc::new(CountingMiddleware {
204            counter: mw_counter.clone(),
205        });
206
207        let chained = chain_handlers(base, &[mw]);
208        let msg = Message::new("t", "data".to_string());
209        chained.handle(msg).await.unwrap();
210
211        assert_eq!(mw_counter.load(Ordering::SeqCst), 1);
212        assert_eq!(base_counter.load(Ordering::SeqCst), 1);
213    }
214}