use crate::{
error::{RpcErrorExt, TransportError, TransportErrorKind},
TransportFut,
};
use alloy_json_rpc::{RequestPacket, ResponsePacket};
use core::fmt;
use std::{
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
task::{Context, Poll},
time::Duration,
};
use tower::{Layer, Service};
use tracing::trace;
#[cfg(all(target_family = "wasm", target_os = "unknown"))]
use wasmtimer::tokio::sleep;
#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
use tokio::time::sleep;
const DEFAULT_AVG_COST: u64 = 20u64;
#[derive(Debug, Clone)]
pub struct RetryBackoffLayer<P: RetryPolicy = RateLimitRetryPolicy> {
max_rate_limit_retries: u32,
initial_backoff: u64,
compute_units_per_second: u64,
avg_cost: u64,
policy: P,
}
impl RetryBackoffLayer {
pub const fn new(
max_rate_limit_retries: u32,
initial_backoff: u64,
compute_units_per_second: u64,
) -> Self {
Self {
max_rate_limit_retries,
initial_backoff,
compute_units_per_second,
avg_cost: DEFAULT_AVG_COST,
policy: RateLimitRetryPolicy,
}
}
pub const fn with_avg_unit_cost(mut self, avg_cost: u64) -> Self {
self.avg_cost = avg_cost;
self
}
}
impl<P: RetryPolicy> RetryBackoffLayer<P> {
pub const fn new_with_policy(
max_rate_limit_retries: u32,
initial_backoff: u64,
compute_units_per_second: u64,
policy: P,
) -> Self {
Self {
max_rate_limit_retries,
initial_backoff,
compute_units_per_second,
policy,
avg_cost: DEFAULT_AVG_COST,
}
}
}
#[derive(Debug, Copy, Clone, Default)]
#[non_exhaustive]
pub struct RateLimitRetryPolicy;
impl RateLimitRetryPolicy {
pub fn or<F>(self, f: F) -> OrRetryPolicyFn<Self>
where
F: Fn(&TransportError) -> bool + Send + Sync + 'static,
{
OrRetryPolicyFn::new(self, f)
}
}
pub trait RetryPolicy: Send + Sync + std::fmt::Debug {
fn should_retry(&self, error: &TransportError) -> bool;
fn backoff_hint(&self, error: &TransportError) -> Option<std::time::Duration>;
}
impl RetryPolicy for RateLimitRetryPolicy {
fn should_retry(&self, error: &TransportError) -> bool {
error.is_retryable()
}
fn backoff_hint(&self, error: &TransportError) -> Option<std::time::Duration> {
error.backoff_hint()
}
}
#[derive(Clone)]
pub struct OrRetryPolicyFn<P = RateLimitRetryPolicy> {
inner: Arc<dyn Fn(&TransportError) -> bool + Send + Sync>,
base: P,
}
impl<P> OrRetryPolicyFn<P> {
pub fn new<F>(base: P, or: F) -> Self
where
F: Fn(&TransportError) -> bool + Send + Sync + 'static,
{
Self { inner: Arc::new(or), base }
}
}
impl<P: RetryPolicy> RetryPolicy for OrRetryPolicyFn<P> {
fn should_retry(&self, error: &TransportError) -> bool {
self.inner.as_ref()(error) || self.base.should_retry(error)
}
fn backoff_hint(&self, error: &TransportError) -> Option<Duration> {
self.base.backoff_hint(error)
}
}
impl<P: fmt::Debug> fmt::Debug for OrRetryPolicyFn<P> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OrRetryPolicyFn")
.field("base", &self.base)
.field("inner", &"{{..}}")
.finish_non_exhaustive()
}
}
impl<S, P: RetryPolicy + Clone> Layer<S> for RetryBackoffLayer<P> {
type Service = RetryBackoffService<S, P>;
fn layer(&self, inner: S) -> Self::Service {
RetryBackoffService {
inner,
policy: self.policy.clone(),
max_rate_limit_retries: self.max_rate_limit_retries,
initial_backoff: self.initial_backoff,
compute_units_per_second: self.compute_units_per_second,
requests_enqueued: Arc::new(AtomicU32::new(0)),
avg_cost: self.avg_cost,
}
}
}
#[derive(Debug, Clone)]
pub struct RetryBackoffService<S, P: RetryPolicy = RateLimitRetryPolicy> {
inner: S,
policy: P,
max_rate_limit_retries: u32,
initial_backoff: u64,
compute_units_per_second: u64,
requests_enqueued: Arc<AtomicU32>,
avg_cost: u64,
}
impl<S, P: RetryPolicy> RetryBackoffService<S, P> {
const fn initial_backoff(&self) -> Duration {
Duration::from_millis(self.initial_backoff)
}
}
#[derive(Debug)]
struct QueuedRequest {
requests_enqueued: Arc<AtomicU32>,
}
impl QueuedRequest {
fn new(requests_enqueued: Arc<AtomicU32>) -> (Self, u64) {
let ahead_in_queue = requests_enqueued.fetch_add(1, Ordering::SeqCst) as u64;
(Self { requests_enqueued }, ahead_in_queue)
}
}
impl Drop for QueuedRequest {
fn drop(&mut self) {
self.requests_enqueued.fetch_sub(1, Ordering::SeqCst);
}
}
impl<S, P> Service<RequestPacket> for RetryBackoffService<S, P>
where
S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
+ Send
+ 'static
+ Clone,
P: RetryPolicy + Clone + 'static,
{
type Response = ResponsePacket;
type Error = TransportError;
type Future = TransportFut<'static>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: RequestPacket) -> Self::Future {
let inner = self.inner.clone();
let this = self.clone();
let mut inner = std::mem::replace(&mut self.inner, inner);
Box::pin(async move {
let (_queued_request, ahead_in_queue) =
QueuedRequest::new(this.requests_enqueued.clone());
let mut rate_limit_retry_number: u32 = 0;
loop {
let err;
let res = inner.call(request.clone()).await;
match res {
Ok(res) => {
if let Some(e) = res.as_error() {
err = TransportError::ErrorResp(e.clone())
} else {
return Ok(res);
}
}
Err(e) => err = e,
}
let should_retry = this.policy.should_retry(&err);
if should_retry {
rate_limit_retry_number += 1;
if rate_limit_retry_number > this.max_rate_limit_retries {
return Err(TransportErrorKind::custom_str(&format!(
"Max retries exceeded {err}"
)));
}
trace!(%err, "retrying request");
let current_queued_reqs = this.requests_enqueued.load(Ordering::SeqCst) as u64;
let backoff_hint = this.policy.backoff_hint(&err);
let next_backoff = backoff_hint.unwrap_or_else(|| this.initial_backoff());
let seconds_to_wait_for_compute_budget = compute_unit_offset_in_secs(
this.avg_cost,
this.compute_units_per_second,
current_queued_reqs,
ahead_in_queue,
);
let total_backoff = next_backoff.saturating_add(
std::time::Duration::from_secs(seconds_to_wait_for_compute_budget),
);
trace!(
total_backoff_millis = total_backoff.as_millis(),
budget_backoff_millis = seconds_to_wait_for_compute_budget * 1000,
default_backoff_millis = next_backoff.as_millis(),
backoff_hint_millis = backoff_hint.map(|d| d.as_millis()),
"(all in ms) backing off due to rate limit"
);
sleep(total_backoff).await;
} else {
return Err(err);
}
}
})
}
}
fn compute_unit_offset_in_secs(
avg_cost: u64,
compute_units_per_second: u64,
current_queued_requests: u64,
ahead_in_queue: u64,
) -> u64 {
let request_capacity_per_second = compute_units_per_second.saturating_div(avg_cost).max(1);
if current_queued_requests > request_capacity_per_second {
current_queued_requests.min(ahead_in_queue).saturating_div(request_capacity_per_second)
} else {
0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn request_queue_count_decrements_when_future_is_dropped() {
let pending_service =
tower::service_fn(|_| -> TransportFut<'static> { Box::pin(std::future::pending()) });
let mut service = RetryBackoffLayer::new(1, 1, 1).layer(pending_service);
let requests_enqueued = service.requests_enqueued.clone();
let request = RequestPacket::Single(
alloy_json_rpc::Request::new("test", alloy_json_rpc::Id::Number(1), ())
.serialize()
.unwrap(),
);
let mut future = service.call(request);
assert!(futures::poll!(&mut future).is_pending());
assert_eq!(requests_enqueued.load(Ordering::SeqCst), 1);
drop(future);
assert_eq!(requests_enqueued.load(Ordering::SeqCst), 0);
}
#[test]
fn test_compute_units_per_second() {
let offset = compute_unit_offset_in_secs(17, 10, 0, 0);
assert_eq!(offset, 0);
let offset = compute_unit_offset_in_secs(17, 10, 2, 2);
assert_eq!(offset, 2);
}
}