rskit_messaging/middleware/
dedup.rs1use std::collections::HashMap;
4use std::sync::Arc;
5use std::time::{Duration, Instant};
6
7use async_trait::async_trait;
8use rskit_errors::AppResult;
9
10use crate::handler::{HandlerMiddleware, MessageHandler};
11use crate::message::Message;
12
13#[derive(Debug, Clone)]
15pub struct DedupConfig {
16 pub window_size: usize,
18 pub ttl: Duration,
20}
21
22impl Default for DedupConfig {
23 fn default() -> Self {
24 Self {
25 window_size: 10_000,
26 ttl: Duration::from_mins(5),
27 }
28 }
29}
30
31pub fn dedup<T: Send + Sync + 'static>(config: DedupConfig) -> impl HandlerMiddleware<T> {
37 DedupMiddleware {
38 config,
39 seen: Arc::new(parking_lot::Mutex::new(HashMap::new())),
40 }
41}
42
43struct DedupMiddleware {
44 config: DedupConfig,
45 seen: Arc<parking_lot::Mutex<HashMap<String, Instant>>>,
46}
47
48impl<T: Send + Sync + 'static> HandlerMiddleware<T> for DedupMiddleware {
49 fn wrap(&self, next: Arc<dyn MessageHandler<T>>) -> Arc<dyn MessageHandler<T>> {
50 Arc::new(DedupHandler {
51 config: self.config.clone(),
52 seen: self.seen.clone(),
53 next,
54 })
55 }
56}
57
58struct DedupHandler<T: Send + Sync + 'static> {
59 config: DedupConfig,
60 seen: Arc<parking_lot::Mutex<HashMap<String, Instant>>>,
61 next: Arc<dyn MessageHandler<T>>,
62}
63
64#[async_trait]
65impl<T: Send + Sync + 'static> MessageHandler<T> for DedupHandler<T> {
66 async fn handle(&self, msg: Message<T>) -> AppResult<()> {
67 if let Some(id) = msg.headers.get("message-id") {
68 let mut seen = self.seen.lock();
69 let now = Instant::now();
70
71 seen.retain(|_, ts| now.duration_since(*ts) < self.config.ttl);
73
74 if seen.contains_key(id) {
75 ::tracing::debug!(message_id = %id, "duplicate message skipped");
76 return Ok(());
77 }
78
79 while seen.len() >= self.config.window_size {
81 let oldest = seen
82 .iter()
83 .min_by_key(|(_, ts)| *ts)
84 .map(|(k, _)| k.clone());
85 if let Some(key) = oldest {
86 seen.remove(&key);
87 } else {
88 break;
89 }
90 }
91
92 seen.insert(id.clone(), now);
93 }
94 self.next.handle(msg).await
95 }
96}
97
98#[cfg(test)]
99mod tests {
100 use std::sync::atomic::{AtomicU32, Ordering};
101
102 use super::*;
103 use crate::handler::{FnHandler, chain_handlers};
104
105 fn counting_handler(counter: &Arc<AtomicU32>) -> Arc<dyn MessageHandler<String>> {
106 let c = counter.clone();
107 Arc::new(FnHandler::new(move |_msg: Message<String>| {
108 let c = c.clone();
109 async move {
110 c.fetch_add(1, Ordering::SeqCst);
111 Ok(())
112 }
113 }))
114 }
115
116 #[tokio::test]
117 async fn duplicate_message_is_skipped() {
118 let counter = Arc::new(AtomicU32::new(0));
119 let mw = DedupMiddleware {
120 config: DedupConfig::default(),
121 seen: Arc::new(parking_lot::Mutex::new(HashMap::new())),
122 };
123 let handler = chain_handlers(
124 counting_handler(&counter),
125 &[Arc::new(mw) as Arc<dyn HandlerMiddleware<String>>],
126 );
127
128 let msg1 = Message::new("t", "a".to_string()).with_header("message-id", "id-1");
129 let msg2 = Message::new("t", "b".to_string()).with_header("message-id", "id-1");
130
131 handler.handle(msg1).await.unwrap();
132 handler.handle(msg2).await.unwrap();
133
134 assert_eq!(counter.load(Ordering::SeqCst), 1);
135 }
136
137 #[tokio::test]
138 async fn different_ids_are_processed() {
139 let counter = Arc::new(AtomicU32::new(0));
140 let mw = DedupMiddleware {
141 config: DedupConfig::default(),
142 seen: Arc::new(parking_lot::Mutex::new(HashMap::new())),
143 };
144 let handler = chain_handlers(
145 counting_handler(&counter),
146 &[Arc::new(mw) as Arc<dyn HandlerMiddleware<String>>],
147 );
148
149 let msg1 = Message::new("t", "a".to_string()).with_header("message-id", "id-1");
150 let msg2 = Message::new("t", "b".to_string()).with_header("message-id", "id-2");
151
152 handler.handle(msg1).await.unwrap();
153 handler.handle(msg2).await.unwrap();
154
155 assert_eq!(counter.load(Ordering::SeqCst), 2);
156 }
157
158 #[tokio::test]
159 async fn messages_without_id_always_processed() {
160 let counter = Arc::new(AtomicU32::new(0));
161 let mw = DedupMiddleware {
162 config: DedupConfig::default(),
163 seen: Arc::new(parking_lot::Mutex::new(HashMap::new())),
164 };
165 let handler = chain_handlers(
166 counting_handler(&counter),
167 &[Arc::new(mw) as Arc<dyn HandlerMiddleware<String>>],
168 );
169
170 let msg1 = Message::new("t", "a".to_string());
171 let msg2 = Message::new("t", "b".to_string());
172
173 handler.handle(msg1).await.unwrap();
174 handler.handle(msg2).await.unwrap();
175
176 assert_eq!(counter.load(Ordering::SeqCst), 2);
177 }
178
179 #[tokio::test]
180 async fn window_size_evicts_oldest() {
181 let counter = Arc::new(AtomicU32::new(0));
182 let mw = DedupMiddleware {
183 config: DedupConfig {
184 window_size: 2,
185 ttl: Duration::from_mins(5),
186 },
187 seen: Arc::new(parking_lot::Mutex::new(HashMap::new())),
188 };
189 let handler = chain_handlers(
190 counting_handler(&counter),
191 &[Arc::new(mw) as Arc<dyn HandlerMiddleware<String>>],
192 );
193
194 handler
196 .handle(Message::new("t", "a".to_string()).with_header("message-id", "id-1"))
197 .await
198 .unwrap();
199 handler
200 .handle(Message::new("t", "b".to_string()).with_header("message-id", "id-2"))
201 .await
202 .unwrap();
203
204 handler
206 .handle(Message::new("t", "c".to_string()).with_header("message-id", "id-3"))
207 .await
208 .unwrap();
209
210 handler
212 .handle(Message::new("t", "d".to_string()).with_header("message-id", "id-1"))
213 .await
214 .unwrap();
215
216 assert_eq!(counter.load(Ordering::SeqCst), 4);
217 }
218
219 #[tokio::test]
220 async fn public_dedup_constructor_tracks_duplicates() {
221 let counter = Arc::new(AtomicU32::new(0));
222 let handler = chain_handlers(
223 counting_handler(&counter),
224 &[Arc::new(dedup(DedupConfig {
225 window_size: 1,
226 ttl: Duration::from_millis(1),
227 })) as Arc<dyn HandlerMiddleware<String>>],
228 );
229
230 handler
231 .handle(Message::new("t", "a".to_string()).with_header("message-id", "id-1"))
232 .await
233 .unwrap();
234 handler
235 .handle(Message::new("t", "b".to_string()).with_header("message-id", "id-1"))
236 .await
237 .unwrap();
238
239 assert_eq!(counter.load(Ordering::SeqCst), 1);
240 }
241}