Skip to main content

doiget_core/
rate_limiter.rs

1//! Process-wide rate limiter for HTTP fetches across all `Source` impls.
2//!
3//! See `docs/SECURITY.md` (per-session fetch flood mitigation) and
4//! `docs/SOURCES.md` §6 (Politeness defaults). The constants enforced here
5//! are the load-bearing safeguards from `docs/LEGAL.md` §6 safeguard 8.
6
7use std::collections::HashMap;
8use std::sync::Arc;
9use std::time::Duration;
10
11use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore};
12use tokio::time::{sleep_until, Instant};
13
14use crate::RateLimits;
15
16/// Process-wide async rate limiter.
17///
18/// Enforces three invariants on every [`acquire`](RateLimiter::acquire):
19///   1. **Global concurrency** — at most
20///      [`RateLimits::max_concurrent_fetches`](crate::RateLimits::max_concurrent_fetches)
21///      in flight at once.
22///   2. **Global rate** — at most
23///      [`RateLimits::max_fetches_per_second`](crate::RateLimits::max_fetches_per_second)
24///      starts in any rolling one-second window.
25///   3. **Per-source backoff** — at least
26///      [`RateLimits::per_source_backoff_ms`](crate::RateLimits::per_source_backoff_ms)
27///      between consecutive starts to the same source name.
28///
29/// The returned [`Permit`] holds the concurrency slot for the lifetime of the
30/// value; drop it when the fetch is done.
31///
32/// 429 / `Retry-After` handling is split: the limiter only exposes the admin
33/// hook [`sleep_for`](RateLimiter::sleep_for); the actual `Retry-After`
34/// header parse and call lives at the `Source::fetch` call site, per
35/// `docs/SOURCES.md` §6.
36#[derive(Debug)]
37pub struct RateLimiter {
38    limits: RateLimits,
39    sem: Arc<Semaphore>,
40    // Global rolling-second window: timestamps of starts within the last second.
41    global_starts: Arc<Mutex<Vec<Instant>>>,
42    // Earliest-allowed start time per source name.
43    per_source_next: Arc<Mutex<HashMap<String, Instant>>>,
44    // #493: per-source concurrency, for sources whose terms cap it below
45    // the global semaphore. Created lazily so a source with no override
46    // costs nothing.
47    per_source_sem: Arc<Mutex<HashMap<String, Arc<Semaphore>>>>,
48}
49
50/// Held while a fetch is in flight; releases the concurrency slot on drop.
51#[derive(Debug)]
52pub struct Permit {
53    _slot: OwnedSemaphorePermit,
54    // #493: held for the same lifetime as the global slot when the source
55    // has a stricter concurrency cap.
56    _source_slot: Option<OwnedSemaphorePermit>,
57}
58
59impl RateLimiter {
60    /// Construct from the hard-coded [`RateLimits`] (the only public path).
61    pub fn new(limits: RateLimits) -> Self {
62        let max = limits.max_concurrent_fetches() as usize;
63        Self {
64            limits,
65            sem: Arc::new(Semaphore::new(max)),
66            global_starts: Arc::new(Mutex::new(Vec::new())),
67            per_source_next: Arc::new(Mutex::new(HashMap::new())),
68            per_source_sem: Arc::new(Mutex::new(HashMap::new())),
69        }
70    }
71
72    /// Wait out the per-source interval for one ADDITIONAL request inside a
73    /// fetch that already holds a [`Permit`].
74    ///
75    /// #493. arXiv's terms cap *requests*, and one arXiv attempt issues two
76    /// of them -- the Atom feed, then the PDF -- under a single `acquire`.
77    /// The second was previously unpaced, so even a perfectly serialised
78    /// caller sent two requests back to back.
79    ///
80    /// Does not touch the concurrency slots: the caller already holds them,
81    /// and taking them again would deadlock at `max_concurrent_for == 1`,
82    /// which is exactly the arXiv case.
83    pub async fn pace(&self, source: &str) {
84        // The global cap admits this request too. Review of #493 caught the
85        // first draft pushing a start into `global_starts` WITHOUT waiting
86        // on the window -- which inflates the window for every other source
87        // while never being bounded by it. It cannot bite at arXiv's 3 s
88        // interval, and "it cannot bite in practice" is the reasoning that
89        // produced #493 in the first place.
90        self.await_global_rate_window().await;
91
92        let backoff = Duration::from_millis(self.limits.backoff_ms_for(source));
93        let mut next_map = self.per_source_next.lock().await;
94        if let Some(&next) = next_map.get(source) {
95            if Instant::now() < next {
96                drop(next_map);
97                sleep_until(next).await;
98                next_map = self.per_source_next.lock().await;
99            }
100        }
101        let start = Instant::now();
102        next_map.insert(source.to_string(), start + backoff);
103        drop(next_map);
104
105        let mut starts = self.global_starts.lock().await;
106        starts.push(start);
107    }
108
109    /// Block until a slot is available, then return a [`Permit`].
110    ///
111    /// Order of waits, in this exact sequence:
112    ///   1. global concurrency (semaphore acquire),
113    ///   2. global rate cap (sleep if the rolling-second window is full),
114    ///   3. per-source backoff (sleep until the source's `next` time).
115    ///
116    /// Lock-acquisition order is always `global_starts` first, THEN
117    /// `per_source_next`. Any future call site that needs both locks MUST
118    /// follow the same order to keep the system deadlock-free.
119    pub async fn acquire(&self, source: &str) -> Permit {
120        // Step 1: global concurrency — bounded by Semaphore::new(max).
121        // `acquire_owned` only errors when the semaphore is closed; this
122        // type never closes it (no `close()` call exists), so the Err arm
123        // is structurally unreachable. The local `allow` is the documented
124        // exception to the workspace `expect_used` lint.
125        #[allow(clippy::expect_used)]
126        let slot = self
127            .sem
128            .clone()
129            .acquire_owned()
130            .await
131            .expect("rate-limiter semaphore is never closed");
132
133        // Step 2: global rate cap.
134        self.await_global_rate_window().await;
135
136        // Step 2b: per-source concurrency (#493). After the global
137        // semaphore, so the global cap is still the ceiling, and before the
138        // interval wait, so a source capped at one connection serialises
139        // rather than piling up sleepers.
140        let source_slot = {
141            let cap = self.limits.max_concurrent_for(source) as usize;
142            if cap >= self.limits.max_concurrent_fetches() as usize {
143                None
144            } else {
145                let sem = {
146                    let mut map = self.per_source_sem.lock().await;
147                    Arc::clone(
148                        map.entry(source.to_string())
149                            .or_insert_with(|| Arc::new(Semaphore::new(cap))),
150                    )
151                };
152                #[allow(clippy::expect_used)]
153                Some(
154                    sem.acquire_owned()
155                        .await
156                        .expect("per-source semaphore is never closed"),
157                )
158            }
159        };
160
161        // Step 3: per-source backoff. Acquire `per_source_next` strictly
162        // after dropping `global_starts` above (lock order documented).
163        //
164        // #493: `backoff_ms_for`, not `per_source_backoff_ms` -- the
165        // vendor's published guideline when it is stricter than the global
166        // 200 ms floor.
167        let backoff = Duration::from_millis(self.limits.backoff_ms_for(source));
168        let mut next_map = self.per_source_next.lock().await;
169        let now = Instant::now();
170        if let Some(&next) = next_map.get(source) {
171            if now < next {
172                drop(next_map);
173                sleep_until(next).await;
174                next_map = self.per_source_next.lock().await;
175            }
176        }
177        // Record this start in both ledgers. We re-read `Instant::now()`
178        // because we may have slept in step 2 or step 3.
179        let start = Instant::now();
180        next_map.insert(source.to_string(), start + backoff);
181        drop(next_map);
182
183        // Push the start timestamp into the global window. Done AFTER
184        // releasing per_source_next to keep the documented lock order
185        // (global → per-source) on every code path.
186        let mut starts = self.global_starts.lock().await;
187        starts.push(start);
188        drop(starts);
189
190        Permit {
191            _slot: slot,
192            _source_slot: source_slot,
193        }
194    }
195
196    /// Block until the global rolling-second window has room.
197    ///
198    /// Loops because another task may take the slot between the wake and
199    /// the re-check. Holds only `global_starts`, and never across a sleep,
200    /// so it composes with the documented lock order.
201    async fn await_global_rate_window(&self) {
202        let max_per_sec = self.limits.max_fetches_per_second() as usize;
203        let one_sec = Duration::from_secs(1);
204        loop {
205            let mut starts = self.global_starts.lock().await;
206            let now = Instant::now();
207            // Prune entries older than 1 s. `starts` is FIFO, so this is a
208            // contiguous prefix.
209            let cutoff = now.checked_sub(one_sec).unwrap_or(now);
210            let drop_count = starts.iter().take_while(|t| **t <= cutoff).count();
211            if drop_count > 0 {
212                starts.drain(..drop_count);
213            }
214            if starts.len() < max_per_sec {
215                return;
216            }
217            // Window is full -- wake when the oldest entry ages out.
218            // `starts.len() >= max_per_sec >= 1` here, so `[0]` is safe.
219            let wake = starts[0] + one_sec;
220            drop(starts);
221            sleep_until(wake).await;
222        }
223    }
224
225    /// Tell the limiter to delay further starts to `source` by at least
226    /// `dur`. Used when the source returns 429 with `Retry-After`.
227    pub async fn sleep_for(&self, source: &str, dur: Duration) {
228        let mut next_map = self.per_source_next.lock().await;
229        let target = Instant::now() + dur;
230        let entry = next_map.entry(source.to_string()).or_insert(target);
231        if *entry < target {
232            *entry = target;
233        }
234    }
235}
236
237// ---------------------------------------------------------------------------
238// Tests
239// ---------------------------------------------------------------------------
240
241#[cfg(test)]
242#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
243mod tests {
244    use super::*;
245    use std::sync::atomic::{AtomicUsize, Ordering};
246
247    use crate::{RateLimits, MAX_CONCURRENT_FETCHES, MAX_FETCHES_PER_SECOND};
248
249    /// Convenience: shared `Arc<RateLimiter>` initialized from
250    /// `RateLimits::HARD_CODED`.
251    fn limiter() -> Arc<RateLimiter> {
252        Arc::new(RateLimiter::new(RateLimits::HARD_CODED))
253    }
254
255    #[tokio::test(flavor = "current_thread", start_paused = true)]
256    async fn concurrent_acquires_respect_max_concurrency() {
257        // Spawn 10 tasks racing to acquire; assert the live count never
258        // exceeds MAX_CONCURRENT_FETCHES.
259        let rl = limiter();
260        let live = Arc::new(AtomicUsize::new(0));
261        let max_seen = Arc::new(AtomicUsize::new(0));
262        let mut handles = Vec::new();
263        for i in 0..10u32 {
264            let rl = rl.clone();
265            let live = live.clone();
266            let max_seen = max_seen.clone();
267            let src = format!("src-{}", i);
268            handles.push(tokio::spawn(async move {
269                let permit = rl.acquire(&src).await;
270                let now = live.fetch_add(1, Ordering::SeqCst) + 1;
271                max_seen.fetch_max(now, Ordering::SeqCst);
272                // Hold the permit briefly so peers contend.
273                tokio::time::sleep(Duration::from_millis(50)).await;
274                live.fetch_sub(1, Ordering::SeqCst);
275                drop(permit);
276            }));
277        }
278        for h in handles {
279            h.await.expect("task ok");
280        }
281        let max = max_seen.load(Ordering::SeqCst);
282        assert!(
283            max <= MAX_CONCURRENT_FETCHES as usize,
284            "max concurrent live = {}, expected <= {}",
285            max,
286            MAX_CONCURRENT_FETCHES
287        );
288        assert!(max > 0, "at least one acquire should succeed");
289    }
290
291    #[tokio::test(flavor = "current_thread", start_paused = true)]
292    async fn same_source_starts_separated_by_backoff() {
293        // Two acquires for the same source must be at least
294        // per_source_backoff_ms apart.
295        let rl = limiter();
296        let backoff_ms = RateLimits::HARD_CODED.per_source_backoff_ms();
297
298        let t0 = Instant::now();
299        let p0 = rl.acquire("crossref").await;
300        drop(p0);
301        let _p1 = rl.acquire("crossref").await;
302        let elapsed = Instant::now().duration_since(t0);
303
304        assert!(
305            elapsed >= Duration::from_millis(backoff_ms),
306            "elapsed {:?} < backoff {} ms",
307            elapsed,
308            backoff_ms
309        );
310    }
311
312    #[tokio::test(flavor = "current_thread", start_paused = true)]
313    async fn different_sources_no_per_source_wait() {
314        // Acquire source A, then source B back-to-back: per-source backoff
315        // must not apply between distinct sources. (Global rate still
316        // applies; with only two starts it does not bind.)
317        let rl = limiter();
318        let backoff = Duration::from_millis(RateLimits::HARD_CODED.per_source_backoff_ms());
319
320        let t0 = Instant::now();
321        let _p_a = rl.acquire("source-a").await;
322        let _p_b = rl.acquire("source-b").await;
323        let elapsed = Instant::now().duration_since(t0);
324
325        assert!(
326            elapsed < backoff,
327            "elapsed {:?} should be well under per-source backoff {:?}",
328            elapsed,
329            backoff
330        );
331    }
332
333    #[tokio::test(flavor = "current_thread", start_paused = true)]
334    async fn global_rate_caps_starts_per_second() {
335        // Acquire 10 distinct sources back-to-back, dropping each permit
336        // immediately so the concurrency cap (5) does not collide with the
337        // rate cap we're trying to observe. Only MAX_FETCHES_PER_SECOND
338        // starts may complete in the first second; the remainder must wait
339        // for the rolling-second window to free.
340        let rl = limiter();
341        let max_per_sec = MAX_FETCHES_PER_SECOND as usize;
342
343        let t0 = Instant::now();
344        let mut completion_offsets: Vec<Duration> = Vec::with_capacity(10);
345        for i in 0..10u32 {
346            let src = format!("src-{}", i);
347            let p = rl.acquire(&src).await;
348            completion_offsets.push(Instant::now().duration_since(t0));
349            drop(p); // release immediately — we are testing rate, not concurrency.
350        }
351
352        // Within the first second from t0, at most max_per_sec acquires
353        // should have completed.
354        let in_first_sec = completion_offsets
355            .iter()
356            .filter(|d| **d < Duration::from_secs(1))
357            .count();
358        assert!(
359            in_first_sec <= max_per_sec,
360            "{} starts completed in first second, expected <= {}",
361            in_first_sec,
362            max_per_sec
363        );
364    }
365
366    #[tokio::test(flavor = "current_thread", start_paused = true)]
367    async fn sleep_for_delays_target_source() {
368        // sleep_for("X", 500ms) then acquire("X") must take at least 500
369        // ms; acquire("Y") in the same window must NOT be delayed by it.
370        let rl = limiter();
371        let delay = Duration::from_millis(500);
372        rl.sleep_for("X", delay).await;
373
374        // Y is unaffected.
375        let t_y = Instant::now();
376        let _p_y = rl.acquire("Y").await;
377        let elapsed_y = Instant::now().duration_since(t_y);
378        assert!(
379            elapsed_y < delay,
380            "Y elapsed {:?} should be far less than {:?}",
381            elapsed_y,
382            delay
383        );
384
385        // X is delayed by at least `delay`.
386        let t_x = Instant::now();
387        let _p_x = rl.acquire("X").await;
388        let elapsed_x = Instant::now().duration_since(t_x);
389        assert!(
390            elapsed_x >= delay,
391            "X elapsed {:?} < requested delay {:?}",
392            elapsed_x,
393            delay
394        );
395    }
396
397    // ---- #493: a vendor guideline stricter than the global cap ---------
398
399    /// arXiv publishes one request every three seconds. The global cap is
400    /// 5/s, so before #493 two consecutive arXiv requests were 200 ms
401    /// apart -- 15x the permitted rate -- while three places in the tree
402    /// asserted the global cap "comfortably respects" the guideline.
403    #[tokio::test(flavor = "current_thread", start_paused = true)]
404    async fn arxiv_requests_are_three_seconds_apart() {
405        let rl = RateLimiter::new(RateLimits::HARD_CODED);
406        let t0 = Instant::now();
407        drop(rl.acquire("arxiv").await);
408        drop(rl.acquire("arxiv").await);
409        let elapsed = Instant::now() - t0;
410        assert!(
411            elapsed >= Duration::from_millis(3_000),
412            "arXiv requests must be >= 3 s apart; got {elapsed:?}"
413        );
414    }
415
416    /// #500: NCBI's keyless 3 requests a second. Four acquires span at
417    /// least a second, and two are not the arXiv 3 s apart.
418    #[tokio::test(flavor = "current_thread", start_paused = true)]
419    async fn ncbi_requests_stay_under_three_a_second() {
420        let rl = RateLimiter::new(RateLimits::HARD_CODED);
421        let t0 = Instant::now();
422        for _ in 0..4 {
423            drop(rl.acquire(crate::pubmed::NCBI).await);
424        }
425        let elapsed = Instant::now() - t0;
426        assert!(
427            elapsed >= Duration::from_millis(1_000),
428            "four NCBI requests inside a second: {elapsed:?}"
429        );
430        assert!(
431            elapsed < Duration::from_millis(1_500),
432            "tighter than asked: {elapsed:?}"
433        );
434    }
435
436    /// The table only ever tightens. A source with no entry keeps the
437    /// global 200 ms floor, so the fix cannot have slowed everything else
438    /// down by accident.
439    #[tokio::test(flavor = "current_thread", start_paused = true)]
440    async fn a_source_without_an_override_keeps_the_global_backoff() {
441        let rl = RateLimiter::new(RateLimits::HARD_CODED);
442        let t0 = Instant::now();
443        drop(rl.acquire("crossref").await);
444        drop(rl.acquire("crossref").await);
445        let elapsed = Instant::now() - t0;
446        assert!(
447            elapsed >= Duration::from_millis(200),
448            "the global floor still applies; got {elapsed:?}"
449        );
450        assert!(
451            elapsed < Duration::from_millis(3_000),
452            "crossref has no override and must not inherit arXiv's; got {elapsed:?}"
453        );
454    }
455
456    /// The second request of ONE arXiv attempt is paced too. arXiv caps
457    /// requests, not attempts, and an attempt issues two -- the Atom feed
458    /// and the PDF -- under a single permit (#493).
459    #[tokio::test(flavor = "current_thread", start_paused = true)]
460    async fn pace_spaces_a_second_request_inside_one_attempt() {
461        let rl = RateLimiter::new(RateLimits::HARD_CODED);
462        let permit = rl.acquire("arxiv").await;
463        let t0 = Instant::now();
464        // Deliberately while the permit is still held: `pace` must not
465        // touch the concurrency slots, or arXiv's cap of one connection
466        // would deadlock against itself.
467        rl.pace("arxiv").await;
468        let elapsed = Instant::now() - t0;
469        drop(permit);
470        assert!(
471            elapsed >= Duration::from_millis(3_000),
472            "the second leg must wait out the interval; got {elapsed:?}"
473        );
474    }
475
476    /// `pace` is admitted by the global window, not merely recorded in it.
477    ///
478    /// Found by reading the diff, not by a failing test -- the first draft
479    /// pushed a start into `global_starts` without ever waiting on the
480    /// window, so an extra request inflated the cap for every other source
481    /// while being bounded by none of it. Unreachable at arXiv's 3 s
482    /// interval, which is precisely the argument that let #493 ship.
483    #[tokio::test(flavor = "current_thread", start_paused = true)]
484    async fn pace_is_admitted_by_the_global_window_too() {
485        let rl = RateLimiter::new(RateLimits::HARD_CODED);
486        // Fill the rolling second to the global maximum, on distinct
487        // sources so no per-source interval is in play.
488        for s in ["a", "b", "c", "d", "e"] {
489            drop(rl.acquire(s).await);
490        }
491        let t0 = Instant::now();
492        rl.pace("f").await;
493        let elapsed = Instant::now() - t0;
494        assert!(
495            elapsed >= Duration::from_millis(900),
496            "pace must wait for the global window exactly as acquire does; got {elapsed:?}"
497        );
498    }
499
500    /// The table is a ceiling-tightener, not a general knob: an entry that
501    /// tried to be *looser* than the global settings must not take effect.
502    #[test]
503    fn an_override_can_only_tighten() {
504        let l = RateLimits::HARD_CODED;
505        assert_eq!(l.backoff_ms_for("arxiv"), 3_000);
506        assert_eq!(l.backoff_ms_for("crossref"), l.per_source_backoff_ms());
507        assert_eq!(l.max_concurrent_for("arxiv"), 1);
508        assert_eq!(l.max_concurrent_for("crossref"), l.max_concurrent_fetches());
509        // Every entry is at least as strict as the global cap on both axes.
510        // `backoff_ms_for` / `max_concurrent_for` enforce it at call time;
511        // this pins the TABLE, so a future entry cannot be added in the
512        // belief that it relaxes something and then silently do nothing.
513        for (name, r) in crate::SOURCE_RATE_OVERRIDES {
514            assert!(
515                r.min_interval_ms >= l.per_source_backoff_ms(),
516                "{name}: an override looser than the global floor is silently ignored"
517            );
518            assert!(
519                r.max_concurrent <= l.max_concurrent_fetches(),
520                "{name}: an override cannot raise the global concurrency cap"
521            );
522        }
523    }
524}