Move retry and rate limiting to a library

This commit is contained in:
2024-07-15 13:26:26 +01:00
parent d8206cd99b
commit 72aaf40f4b
3 changed files with 149 additions and 102 deletions
+34 -97
View File
@@ -1,120 +1,57 @@
use std::cmp::max;
use std::num::NonZero;
use chrono::Utc;
use std::time::Duration;
use axum::async_trait;
use governor::clock::DefaultClock;
use governor::state::{InMemoryState, NotKeyed};
use reqwest::{Client, Error, Method, Request, RequestBuilder, Response, StatusCode};
use tokio::sync::oneshot;
use reqwest_middleware::{ClientBuilder, ClientWithMiddleware};
use reqwest_retry::{RetryTransientMiddleware, policies::ExponentialBackoff, Jitter};
#[derive(Debug, thiserror::Error)]
enum TogglApiError {
#[error("Reqwest error: {0}")]
ReqwestError(#[from] reqwest::Error),
#[error("Retries exceeded")]
RetriesExceeded,
}
struct TogglApiRequest(
Request,
oneshot::Sender<Result<Response, TogglApiError>>,
);
struct TogglApiWorker {
struct ReqwestRateLimiter {
rate_limiter: governor::RateLimiter<NotKeyed, InMemoryState, DefaultClock>,
rx: tokio::sync::mpsc::Receiver<TogglApiRequest>,
client: Client,
max_retries: Option<u32>,
default_delay: f64,
}
impl TogglApiWorker {
fn new() -> (Self, TogglApi) {
let (tx, rx) = tokio::sync::mpsc::channel(100);
let client = Client::new();
let rate_limiter = governor::RateLimiter::direct(
governor::Quota::per_second(NonZero::new(1u32).unwrap())
);
(Self { rate_limiter, client, rx, max_retries: Some(3), default_delay: 3.14 }, TogglApi { tx })
}
pub async fn start(&mut self) {
loop {
// We limit ourselves to the recommended rate of 1 req/s
self.rate_limiter.until_ready().await;
let TogglApiRequest(request, tx) = self.rx.recv().await.unwrap();
let response = self.make_request(request).await;
tx.send(response).unwrap();
impl ReqwestRateLimiter {
fn new() -> Self {
Self {
rate_limiter: governor::RateLimiter::direct(
governor::Quota::per_second(NonZero::new(1u32).unwrap())
),
}
}
}
async fn make_request(&self, request: RequestBuilder) -> Result<Response, TogglApiError> {
let max_retries = self.max_retries.unwrap_or(1);
for _ in 0..max_retries {
let response = self.client.execute(request.clone()).await?;
if response.status().is_server_error() || response.status() == StatusCode::TOO_MANY_REQUESTS {
let delay = self.parse_retry_after_header(&response)
.unwrap_or(self.default_delay);
tokio::time::sleep(tokio::time::Duration::from_secs_f64(delay)).await;
} else {
return Ok(response);
}
}
Err(TogglApiError::RetriesExceeded)
}
fn parse_retry_after_header(&self, response: &Response) -> Option<f64> {
match response.headers().get("Retry-After") {
Some(retry_after) => {
let retry_after = retry_after.to_str()
.unwrap();
let a = retry_after.parse::<f64>()
.ok()
.or_else(|_| {
let date = chrono::NaiveDateTime::parse_from_str(
retry_after,
"%a, %d %b %Y %H:%M:%S GMT",
);
date.map(|date| date.and_utc())
.map(|date| Utc::now().signed_duration_since(date))
.map(|time_delta| time_delta.num_seconds() as f64)
.ok()
});
a
}
None => Some(self.default_delay),
}
#[async_trait]
impl reqwest_ratelimit::RateLimiter for ReqwestRateLimiter {
async fn acquire_permit(&self) {
self.rate_limiter.until_ready().await;
}
}
struct TogglApi {
tx: tokio::sync::mpsc::Sender<TogglApiRequest>,
client: ClientWithMiddleware,
api_key: String,
workspace_id: u32,
}
impl TogglApi {
pub async fn request(&self, request: Request) -> Result<Response, TogglApiError> {
let (tx, rx) = oneshot::channel();
self.tx.send(TogglApiRequest(request, tx)).await.expect("send request");
rx.await.unwrap()
fn new(api_key: String, workspace_id: u32) -> Self {
let rate_limiter = ReqwestRateLimiter::new();
let backoff = ExponentialBackoff::builder()
.retry_bounds(Duration::from_secs(1), Duration::from_secs(60))
.jitter(Jitter::Bounded)
.base(2)
.build_with_total_retry_duration(Duration::from_secs(24 * 60 * 60));
let client = ClientBuilder::new(reqwest::Client::new())
.with(reqwest_ratelimit::all(rate_limiter))
.with(RetryTransientMiddleware::new_with_policy(backoff))
.build();
Self { client, api_key, workspace_id }
}
}
#[tokio::main]
async fn main() {
let (mut worker, api_client) = TogglApiWorker::new();
tokio::spawn(async move { worker.start().await; });
dbg!(api_client.request(Request::new(Method::GET, "https://www.google.com".parse().unwrap())).await.unwrap());
let api = TogglApi::new("api_key".to_string(), 123);
}