Skip to main content

alloy_transport/layers/
throttle.rs

1use crate::{TransportError, TransportFut};
2use alloy_json_rpc::{RequestPacket, ResponsePacket};
3use governor::{
4    clock::{QuantaClock, QuantaInstant},
5    middleware::NoOpMiddleware,
6    state::{InMemoryState, NotKeyed},
7    Quota, RateLimiter,
8};
9use std::{
10    num::NonZeroU32,
11    sync::Arc,
12    task::{Context, Poll},
13};
14use tower::{Layer, Service};
15
16/// A rate limiter for throttling RPC requests.
17type Throttle = RateLimiter<NotKeyed, InMemoryState, QuantaClock, NoOpMiddleware<QuantaInstant>>;
18
19/// A Transport Layer responsible for throttling RPC requests.
20#[derive(Debug)]
21pub struct ThrottleLayer {
22    /// Rate limiter used to throttle requests.
23    pub throttle: Arc<Throttle>,
24}
25
26impl ThrottleLayer {
27    /// Creates a new throttle layer with the specified requests per second.
28    ///
29    /// # Panics
30    ///
31    /// Panics if `requests_per_second` is 0.
32    pub fn new(requests_per_second: u32) -> Self {
33        Self::new_with_burst(requests_per_second, NonZeroU32::new(1).unwrap())
34    }
35
36    /// Creates a new throttle layer with the specified requests per second and burst size.
37    ///
38    /// # Panics
39    ///
40    /// Panics if `requests_per_second` is 0.
41    pub fn new_with_burst(requests_per_second: u32, burst: NonZeroU32) -> Self {
42        let quota = Quota::per_second(
43            NonZeroU32::new(requests_per_second)
44                .expect("Request per second must be greater than 0"),
45        )
46        .allow_burst(burst);
47        let throttle = Arc::new(RateLimiter::direct(quota));
48
49        Self { throttle }
50    }
51}
52
53/// A Tower Service used by the ThrottleLayer that is responsible for throttling rpc requests.
54#[derive(Debug, Clone)]
55pub struct ThrottleService<S> {
56    /// The inner service
57    inner: S,
58    throttle: Arc<Throttle>,
59}
60
61impl<S> Layer<S> for ThrottleLayer {
62    type Service = ThrottleService<S>;
63
64    fn layer(&self, inner: S) -> Self::Service {
65        ThrottleService { inner, throttle: self.throttle.clone() }
66    }
67}
68
69impl<S> Service<RequestPacket> for ThrottleService<S>
70where
71    S: Service<RequestPacket, Response = ResponsePacket, Error = TransportError>
72        + Send
73        + 'static
74        + Clone,
75    S::Future: Send + 'static,
76{
77    type Response = ResponsePacket;
78    type Error = TransportError;
79    type Future = TransportFut<'static>;
80
81    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
82        self.inner.poll_ready(cx)
83    }
84
85    fn call(&mut self, request: RequestPacket) -> Self::Future {
86        let throttle = self.throttle.clone();
87        let mut inner = self.inner.clone();
88
89        Box::pin(async move {
90            throttle.until_ready().await;
91            inner.call(request).await
92        })
93    }
94}