laminar_db/ai/backends/
rate_limited.rs1use 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
12pub struct RateLimitedProvider {
14 inner: Arc<dyn InferenceProvider>,
15 limiter: DefaultDirectRateLimiter,
16}
17
18impl RateLimitedProvider {
19 #[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 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
87 async fn burst_beyond_the_rate_is_paced() {
88 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}