Skip to main content

tokio_prompt_orchestrator/
request_dedup.rs

1//! # Request Deduplication
2//!
3//! Coalesces identical in-flight requests so that only one backend call is
4//! made when multiple callers submit the same prompt concurrently.
5//!
6//! ## Design
7//!
8//! Each call to [`RequestDeduplicator::submit`] computes a SHA-256 key over
9//! `model_id + prompt_text`.  If a request with that key is already in-flight,
10//! the caller receives a [`DedupDecision::Waiting`] containing a
11//! [`tokio::sync::oneshot::Receiver`] that resolves when the original request
12//! completes.  The first caller for a key receives
13//! [`DedupDecision::Original`] and is responsible for calling
14//! [`RequestDeduplicator::complete`] when the result is ready.
15//!
16//! ## TTL pruning
17//!
18//! Entries older than `RequestDeduplicator::TTL_SECS` (30 s) are pruned on
19//! every call to `submit()` to prevent unbounded memory growth if the original
20//! caller crashes without completing.
21
22use sha2::{Digest, Sha256};
23use std::collections::HashMap;
24use std::time::{Duration, Instant};
25use tokio::sync::oneshot;
26
27/// TTL after which an in-flight dedup entry is considered stale and pruned.
28const TTL_SECS: u64 = 30;
29
30/// A unique identifier for an original (non-deduplicated) request.
31#[derive(Debug, Clone, PartialEq, Eq, Hash)]
32pub struct RequestId(pub String);
33
34impl RequestId {
35    /// Wrap a string as a `RequestId`.
36    pub fn new(id: impl Into<String>) -> Self {
37        Self(id.into())
38    }
39
40    /// Borrow the inner string.
41    pub fn as_str(&self) -> &str {
42        &self.0
43    }
44}
45
46impl std::fmt::Display for RequestId {
47    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48        f.write_str(&self.0)
49    }
50}
51
52/// Result of submitting a request to the deduplicator.
53pub enum DedupDecision {
54    /// This is the first request with this key; the caller should perform
55    /// the actual inference and call [`RequestDeduplicator::complete`].
56    Original(RequestId),
57    /// An identical request is already in-flight.  The receiver will resolve
58    /// with the response tokens when the original completes.
59    Waiting(oneshot::Receiver<Vec<String>>),
60}
61
62/// Snapshot of deduplicator statistics.
63#[derive(Debug, Clone)]
64pub struct DedupStats {
65    /// Total number of requests submitted (including duplicates).
66    pub total_submitted: u64,
67    /// Number of requests that were deduplicated (returned `Waiting`).
68    pub deduplicated: u64,
69    /// Number of requests currently in-flight.
70    pub active_requests: usize,
71    /// Fraction of requests that were deduplicated (`deduplicated / total_submitted`).
72    pub dedup_rate: f64,
73}
74
75struct InFlightEntry {
76    /// Original request ID (kept for debugging; read in tests).
77    #[allow(dead_code)]
78    id: RequestId,
79    /// Channels waiting for this request to complete.
80    waiters: Vec<oneshot::Sender<Vec<String>>>,
81    /// When the entry was created (for TTL pruning).
82    created_at: Instant,
83}
84
85/// Deduplicates identical in-flight LLM requests by their content hash.
86///
87/// All methods take `&mut self`, so wrap in a `Mutex` or `RwLock` for
88/// concurrent access.
89pub struct RequestDeduplicator {
90    /// In-flight entries keyed by SHA-256 hash of `model_id + prompt_text`.
91    in_flight: HashMap<String, InFlightEntry>,
92    total_submitted: u64,
93    deduplicated: u64,
94}
95
96impl RequestDeduplicator {
97    /// TTL for in-flight entries.
98    pub const TTL: Duration = Duration::from_secs(TTL_SECS);
99
100    /// Create a new, empty deduplicator.
101    #[must_use]
102    pub fn new() -> Self {
103        Self {
104            in_flight: HashMap::new(),
105            total_submitted: 0,
106            deduplicated: 0,
107        }
108    }
109
110    /// Compute the SHA-256 deduplication key for a given model ID and prompt.
111    pub fn hash_key(model_id: &str, prompt_text: &str) -> String {
112        let mut hasher = Sha256::new();
113        hasher.update(model_id.as_bytes());
114        hasher.update(b"\x00");
115        hasher.update(prompt_text.as_bytes());
116        hex::encode(hasher.finalize())
117    }
118
119    /// Submit a request and decide whether it is the original or a duplicate.
120    ///
121    /// Prunes stale entries (older than 30 s) before processing the new
122    /// request so that a crashed original does not block future requests
123    /// with the same key forever.
124    pub fn submit(
125        &mut self,
126        request_id: impl Into<String>,
127        model_id: &str,
128        prompt_text: &str,
129    ) -> DedupDecision {
130        // Prune stale entries first.
131        self.prune_stale();
132
133        let key = Self::hash_key(model_id, prompt_text);
134        self.total_submitted += 1;
135
136        if let Some(entry) = self.in_flight.get_mut(&key) {
137            // Duplicate: register a waiter.
138            self.deduplicated += 1;
139            let (tx, rx) = oneshot::channel();
140            entry.waiters.push(tx);
141            DedupDecision::Waiting(rx)
142        } else {
143            // Original: insert the entry.
144            let id = RequestId::new(request_id);
145            self.in_flight.insert(
146                key,
147                InFlightEntry {
148                    id: id.clone(),
149                    waiters: Vec::new(),
150                    created_at: Instant::now(),
151                },
152            );
153            DedupDecision::Original(id)
154        }
155    }
156
157    /// Signal that the original request identified by `model_id` +
158    /// `prompt_text` has completed with the given `result`.
159    ///
160    /// All registered waiters are notified.  Returns the number of waiters
161    /// that were fanned out to.
162    ///
163    /// If no in-flight entry is found (e.g., it was already pruned), returns
164    /// `0`.
165    pub fn complete(&mut self, model_id: &str, prompt_text: &str, result: Vec<String>) -> usize {
166        let key = Self::hash_key(model_id, prompt_text);
167        let Some(entry) = self.in_flight.remove(&key) else {
168            return 0;
169        };
170        let waiter_count = entry.waiters.len();
171        for tx in entry.waiters {
172            // Ignore send errors: receiver may have been dropped.
173            let _ = tx.send(result.clone());
174        }
175        waiter_count
176    }
177
178    /// Return current statistics.
179    pub fn stats(&self) -> DedupStats {
180        let total = self.total_submitted;
181        let dedup = self.deduplicated;
182        DedupStats {
183            total_submitted: total,
184            deduplicated: dedup,
185            active_requests: self.in_flight.len(),
186            dedup_rate: if total == 0 {
187                0.0
188            } else {
189                dedup as f64 / total as f64
190            },
191        }
192    }
193
194    /// Prune in-flight entries older than [`Self::TTL`].
195    ///
196    /// Called automatically by `submit`; can also be called manually.
197    pub fn prune_stale(&mut self) {
198        let ttl = Self::TTL;
199        self.in_flight
200            .retain(|_, entry| entry.created_at.elapsed() < ttl);
201    }
202
203    /// Return the number of currently active (in-flight) requests.
204    pub fn active_count(&self) -> usize {
205        self.in_flight.len()
206    }
207}
208
209impl Default for RequestDeduplicator {
210    fn default() -> Self {
211        Self::new()
212    }
213}
214
215// ============================================================================
216// Tests (15+)
217// ============================================================================
218
219#[cfg(test)]
220mod tests {
221    use super::*;
222
223    // Helper to build a DedupDecision::Original's RequestId if it matches.
224    fn unwrap_original(d: DedupDecision) -> RequestId {
225        match d {
226            DedupDecision::Original(id) => id,
227            DedupDecision::Waiting(_) => panic!("expected Original, got Waiting"),
228        }
229    }
230
231    fn unwrap_waiting(d: DedupDecision) -> oneshot::Receiver<Vec<String>> {
232        match d {
233            DedupDecision::Waiting(rx) => rx,
234            DedupDecision::Original(_) => panic!("expected Waiting, got Original"),
235        }
236    }
237
238    #[test]
239    fn first_submission_is_original() {
240        let mut dedup = RequestDeduplicator::new();
241        let decision = dedup.submit("req-1", "gpt-4", "hello");
242        let id = unwrap_original(decision);
243        assert_eq!(id.as_str(), "req-1");
244    }
245
246    #[test]
247    fn duplicate_submission_is_waiting() {
248        let mut dedup = RequestDeduplicator::new();
249        dedup.submit("req-1", "gpt-4", "hello");
250        let decision = dedup.submit("req-2", "gpt-4", "hello");
251        let _ = unwrap_waiting(decision);
252    }
253
254    #[test]
255    fn different_prompts_are_both_original() {
256        let mut dedup = RequestDeduplicator::new();
257        let d1 = dedup.submit("req-1", "gpt-4", "hello");
258        let d2 = dedup.submit("req-2", "gpt-4", "world");
259        unwrap_original(d1);
260        unwrap_original(d2);
261    }
262
263    #[test]
264    fn different_models_same_prompt_are_both_original() {
265        let mut dedup = RequestDeduplicator::new();
266        let d1 = dedup.submit("req-1", "gpt-4", "hello");
267        let d2 = dedup.submit("req-2", "claude-3", "hello");
268        unwrap_original(d1);
269        unwrap_original(d2);
270    }
271
272    #[tokio::test]
273    async fn complete_fans_out_to_waiters() {
274        let mut dedup = RequestDeduplicator::new();
275        dedup.submit("req-1", "gpt-4", "hello");
276        let rx1 = unwrap_waiting(dedup.submit("req-2", "gpt-4", "hello"));
277        let rx2 = unwrap_waiting(dedup.submit("req-3", "gpt-4", "hello"));
278        let n = dedup.complete("gpt-4", "hello", vec!["token".to_string()]);
279        assert_eq!(n, 2);
280        let r1 = rx1.await.expect("should receive");
281        let r2 = rx2.await.expect("should receive");
282        assert_eq!(r1, vec!["token"]);
283        assert_eq!(r2, vec!["token"]);
284    }
285
286    #[tokio::test]
287    async fn complete_removes_entry_so_next_is_original() {
288        let mut dedup = RequestDeduplicator::new();
289        dedup.submit("req-1", "gpt-4", "hello");
290        dedup.complete("gpt-4", "hello", vec![]);
291        let d = dedup.submit("req-3", "gpt-4", "hello");
292        unwrap_original(d);
293    }
294
295    #[test]
296    fn stats_initial_zeros() {
297        let dedup = RequestDeduplicator::new();
298        let s = dedup.stats();
299        assert_eq!(s.total_submitted, 0);
300        assert_eq!(s.deduplicated, 0);
301        assert_eq!(s.active_requests, 0);
302        assert!((s.dedup_rate - 0.0).abs() < f64::EPSILON);
303    }
304
305    #[test]
306    fn stats_after_submissions() {
307        let mut dedup = RequestDeduplicator::new();
308        dedup.submit("r1", "m", "p1");
309        dedup.submit("r2", "m", "p1"); // dedup
310        dedup.submit("r3", "m", "p2");
311        let s = dedup.stats();
312        assert_eq!(s.total_submitted, 3);
313        assert_eq!(s.deduplicated, 1);
314        assert_eq!(s.active_requests, 2);
315        let expected_rate = 1.0 / 3.0;
316        assert!((s.dedup_rate - expected_rate).abs() < 1e-10);
317    }
318
319    #[test]
320    fn hash_key_is_deterministic() {
321        let k1 = RequestDeduplicator::hash_key("gpt-4", "hello world");
322        let k2 = RequestDeduplicator::hash_key("gpt-4", "hello world");
323        assert_eq!(k1, k2);
324    }
325
326    #[test]
327    fn hash_key_differs_for_different_inputs() {
328        let k1 = RequestDeduplicator::hash_key("gpt-4", "hello");
329        let k2 = RequestDeduplicator::hash_key("gpt-4", "world");
330        assert_ne!(k1, k2);
331    }
332
333    #[test]
334    fn hash_key_differs_for_different_models() {
335        let k1 = RequestDeduplicator::hash_key("gpt-4", "hello");
336        let k2 = RequestDeduplicator::hash_key("claude-3", "hello");
337        assert_ne!(k1, k2);
338    }
339
340    #[test]
341    fn complete_returns_zero_when_no_entry() {
342        let mut dedup = RequestDeduplicator::new();
343        let n = dedup.complete("gpt-4", "nonexistent", vec![]);
344        assert_eq!(n, 0);
345    }
346
347    #[test]
348    fn active_count_decreases_after_complete() {
349        let mut dedup = RequestDeduplicator::new();
350        dedup.submit("r1", "m", "p");
351        assert_eq!(dedup.active_count(), 1);
352        dedup.complete("m", "p", vec![]);
353        assert_eq!(dedup.active_count(), 0);
354    }
355
356    #[test]
357    fn prune_stale_removes_old_entries() {
358        let mut dedup = RequestDeduplicator::new();
359        // Manually insert a stale entry.
360        let key = RequestDeduplicator::hash_key("m", "p");
361        dedup.in_flight.insert(
362            key,
363            super::InFlightEntry {
364                id: RequestId::new("old"),
365                waiters: Vec::new(),
366                created_at: Instant::now()
367                    .checked_sub(Duration::from_secs(60))
368                    .unwrap_or_else(Instant::now),
369            },
370        );
371        assert_eq!(dedup.active_count(), 1);
372        dedup.prune_stale();
373        assert_eq!(dedup.active_count(), 0);
374    }
375
376    #[test]
377    fn dedup_rate_zero_when_no_submissions() {
378        let dedup = RequestDeduplicator::new();
379        assert!((dedup.stats().dedup_rate - 0.0).abs() < f64::EPSILON);
380    }
381
382    #[test]
383    fn request_id_display() {
384        let id = RequestId::new("my-id");
385        assert_eq!(id.to_string(), "my-id");
386    }
387
388    #[test]
389    fn multiple_waiters_all_notified() {
390        let mut dedup = RequestDeduplicator::new();
391        dedup.submit("r1", "m", "prompt");
392        // Register 5 waiters.
393        let mut rxs = Vec::new();
394        for i in 2..=6 {
395            let rx = unwrap_waiting(dedup.submit(&format!("r{i}"), "m", "prompt"));
396            rxs.push(rx);
397        }
398        let n = dedup.complete("m", "prompt", vec!["out".to_string()]);
399        assert_eq!(n, 5);
400        // All receivers have a value pending (non-async check via try_recv).
401        for mut rx in rxs {
402            assert!(rx.try_recv().is_ok());
403        }
404    }
405}