rskit_messaging/
handler.rs1use std::sync::Arc;
4
5use async_trait::async_trait;
6use rskit_errors::AppResult;
7
8use crate::message::Message;
9
10#[async_trait]
15pub trait MessageHandler<T: Send + Sync + 'static>: Send + Sync + 'static {
16 async fn handle(&self, msg: Message<T>) -> AppResult<()>;
18}
19
20pub trait HandlerMiddleware<T: Send + Sync + 'static>: Send + Sync + 'static {
25 fn wrap(&self, next: Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>>;
27}
28
29pub 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
48pub 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
89pub 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 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 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}