Skip to main content

laminar_db/ai/backends/
rate_limited.rs

1//! Client-side token-bucket rate limiting for remote providers. One cell is
2//! acquired per input row before dispatch; the wait happens on Ring 1 only.
3
4use std::num::NonZeroU32;
5use std::sync::Arc;
6
7use async_trait::async_trait;
8use governor::{DefaultDirectRateLimiter, Quota, RateLimiter};
9
10use crate::ai::provider::{InferenceProvider, InferenceRequest, InferenceResponse, ProviderError};
11
12/// An [`InferenceProvider`] that paces calls to a steady rate.
13pub struct RateLimitedProvider {
14    inner: Arc<dyn InferenceProvider>,
15    limiter: DefaultDirectRateLimiter,
16}
17
18impl RateLimitedProvider {
19    /// Wrap `inner`, limiting to `requests_per_second`.
20    #[must_use]
21    pub fn new(inner: Arc<dyn InferenceProvider>, requests_per_second: NonZeroU32) -> Self {
22        Self {
23            inner,
24            limiter: RateLimiter::direct(Quota::per_second(requests_per_second)),
25        }
26    }
27}
28
29#[async_trait]
30impl InferenceProvider for RateLimitedProvider {
31    async fn infer_batch(
32        &self,
33        request: InferenceRequest,
34    ) -> Result<InferenceResponse, ProviderError> {
35        for _ in 0..request.inputs.len().max(1) {
36            self.limiter.until_ready().await;
37        }
38        self.inner.infer_batch(request).await
39    }
40
41    fn name(&self) -> &'static str {
42        self.inner.name()
43    }
44
45    fn intrinsic_labels(&self, model: &str) -> Option<Vec<String>> {
46        self.inner.intrinsic_labels(model)
47    }
48}
49
50#[cfg(test)]
51mod tests {
52    use super::*;
53    use crate::ai::provider::{InferenceOutputs, InferenceParams, Usage};
54    use crate::ai::registry::Task;
55    use std::time::Instant;
56
57    struct Echo;
58
59    #[async_trait]
60    impl InferenceProvider for Echo {
61        async fn infer_batch(
62            &self,
63            request: InferenceRequest,
64        ) -> Result<InferenceResponse, ProviderError> {
65            Ok(InferenceResponse {
66                outputs: InferenceOutputs::Text(request.inputs),
67                usage: Usage::ZERO,
68            })
69        }
70        fn name(&self) -> &'static str {
71            "echo"
72        }
73    }
74
75    fn request(rows: usize) -> InferenceRequest {
76        InferenceRequest {
77            task: Task::Sentiment,
78            model: "m".into(),
79            inputs: vec!["x".to_string(); rows],
80            params: InferenceParams::default(),
81        }
82    }
83
84    /// A burst beyond the per-second budget is delayed, not dropped or sent
85    /// unbounded; the output still passes through unchanged.
86    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
87    async fn burst_beyond_the_rate_is_paced() {
88        // 100 rps → ~10 ms/cell, burst 100. 130 rows = 30 over budget ⇒ ≥ ~300 ms.
89        let p = RateLimitedProvider::new(Arc::new(Echo), NonZeroU32::new(100).unwrap());
90        let start = Instant::now();
91        let resp = p.infer_batch(request(130)).await.unwrap();
92        assert_eq!(resp.outputs.len(), 130);
93        assert!(
94            start.elapsed() >= std::time::Duration::from_millis(200),
95            "burst was not paced: {:?}",
96            start.elapsed()
97        );
98    }
99
100    #[tokio::test]
101    async fn name_delegates_to_inner() {
102        let p = RateLimitedProvider::new(Arc::new(Echo), NonZeroU32::new(1000).unwrap());
103        assert_eq!(p.name(), "echo");
104    }
105}