tokio_prompt_orchestrator/
request_dedup.rs1use sha2::{Digest, Sha256};
23use std::collections::HashMap;
24use std::time::{Duration, Instant};
25use tokio::sync::oneshot;
26
27const TTL_SECS: u64 = 30;
29
30#[derive(Debug, Clone, PartialEq, Eq, Hash)]
32pub struct RequestId(pub String);
33
34impl RequestId {
35 pub fn new(id: impl Into<String>) -> Self {
37 Self(id.into())
38 }
39
40 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
52pub enum DedupDecision {
54 Original(RequestId),
57 Waiting(oneshot::Receiver<Vec<String>>),
60}
61
62#[derive(Debug, Clone)]
64pub struct DedupStats {
65 pub total_submitted: u64,
67 pub deduplicated: u64,
69 pub active_requests: usize,
71 pub dedup_rate: f64,
73}
74
75struct InFlightEntry {
76 #[allow(dead_code)]
78 id: RequestId,
79 waiters: Vec<oneshot::Sender<Vec<String>>>,
81 created_at: Instant,
83}
84
85pub struct RequestDeduplicator {
90 in_flight: HashMap<String, InFlightEntry>,
92 total_submitted: u64,
93 deduplicated: u64,
94}
95
96impl RequestDeduplicator {
97 pub const TTL: Duration = Duration::from_secs(TTL_SECS);
99
100 #[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 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 pub fn submit(
125 &mut self,
126 request_id: impl Into<String>,
127 model_id: &str,
128 prompt_text: &str,
129 ) -> DedupDecision {
130 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 self.deduplicated += 1;
139 let (tx, rx) = oneshot::channel();
140 entry.waiters.push(tx);
141 DedupDecision::Waiting(rx)
142 } else {
143 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 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 let _ = tx.send(result.clone());
174 }
175 waiter_count
176 }
177
178 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 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 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#[cfg(test)]
220mod tests {
221 use super::*;
222
223 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.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 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 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 for mut rx in rxs {
402 assert!(rx.try_recv().is_ok());
403 }
404 }
405}