Skip to main content

rskit_messaging/middleware/
dedup.rs

1//! Message deduplication middleware based on the `message-id` header.
2
3use 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/// Configuration for the deduplication middleware.
14#[derive(Debug, Clone)]
15pub struct DedupConfig {
16    /// Maximum number of message IDs to track.
17    pub window_size: usize,
18    /// Time-to-live for tracked IDs; entries older than this are purged.
19    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
31/// Create a deduplication middleware.
32///
33/// Messages that carry a `message-id` header are tracked.
34/// If a duplicate ID arrives within the configured TTL window it is silently dropped.
35/// Messages without a `message-id` header are always forwarded.
36pub 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            // Purge expired entries.
72            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            // Enforce window size by evicting the oldest entry.
80            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        // Fill window with id-1 and id-2.
195        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        // id-3 should evict id-1 (the oldest).
205        handler
206            .handle(Message::new("t", "c".to_string()).with_header("message-id", "id-3"))
207            .await
208            .unwrap();
209
210        // id-1 should now be accepted again because it was evicted.
211        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}