use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use reqwest::{Client, RequestBuilder, Response};
use serde::{de::DeserializeOwned, Serialize};
use tokio::sync::RwLock;
use tokio::time::Instant;
use crate::{
config::HighLevelConfig,
error::{Error, Result},
};
const DEFAULT_RATE: u32 = 95;
const DEFAULT_WINDOW_SECS: u64 = 10;
const MAX_RETRIES: u32 = 3;
const CIRCUIT_BREAK_THRESHOLD: u32 = 3;
const CIRCUIT_BREAK_RESET_SECS: u64 = 30;
struct TokenBucket {
capacity: u32,
tokens: f64,
refill_per_sec: f64,
last_refill: Instant,
}
impl TokenBucket {
fn new(capacity: u32, window_secs: u64) -> Self {
Self {
capacity,
tokens: capacity as f64,
refill_per_sec: capacity as f64 / window_secs as f64,
last_refill: Instant::now(),
}
}
fn try_acquire(&mut self) -> bool {
self.refill();
if self.tokens >= 1.0 {
self.tokens -= 1.0;
true
} else {
false
}
}
fn refill(&mut self) {
let elapsed = self.last_refill.elapsed().as_secs_f64();
self.tokens = (self.tokens + elapsed * self.refill_per_sec).min(self.capacity as f64);
self.last_refill = Instant::now();
}
}
struct CircuitBreaker {
consecutive_429s: AtomicU32,
open_until: RwLock<Option<Instant>>,
}
impl CircuitBreaker {
fn new() -> Self {
Self {
consecutive_429s: AtomicU32::new(0),
open_until: RwLock::new(None),
}
}
async fn is_open(&self) -> bool {
if let Some(until) = *self.open_until.read().await {
if Instant::now() < until {
return true;
}
self.open_until.write().await.take();
self.consecutive_429s.store(0, Ordering::SeqCst);
}
false
}
async fn record_429(&self) {
let count = self.consecutive_429s.fetch_add(1, Ordering::SeqCst) + 1;
if count >= CIRCUIT_BREAK_THRESHOLD {
*self.open_until.write().await =
Some(Instant::now() + Duration::from_secs(CIRCUIT_BREAK_RESET_SECS));
}
}
async fn record_success(&self) {
self.consecutive_429s.store(0, Ordering::SeqCst);
self.open_until.write().await.take();
}
}
pub struct HttpClient {
inner: Client,
pub base_url: String,
pub api_version: String,
token: RwLock<Option<String>>,
bucket: RwLock<TokenBucket>,
breaker: CircuitBreaker,
}
impl HttpClient {
pub fn new(config: &HighLevelConfig) -> Self {
Self {
inner: Client::new(),
base_url: config.base_url.clone(),
api_version: config.api_version.clone(),
token: RwLock::new(None),
bucket: RwLock::new(TokenBucket::new(DEFAULT_RATE, DEFAULT_WINDOW_SECS)),
breaker: CircuitBreaker::new(),
}
}
pub async fn set_token(&self, token: String) {
*self.token.write().await = Some(token);
}
pub async fn get_token(&self) -> Option<String> {
self.token.read().await.clone()
}
fn url(&self, path: &str) -> String {
format!("{}{}", self.base_url, path)
}
async fn acquire_token(&self) -> Result<()> {
if self.breaker.is_open().await {
return Err(Error::RateLimited {
retry_after: Some(CIRCUIT_BREAK_RESET_SECS),
circuit_breaker: true,
});
}
loop {
if self.bucket.write().await.try_acquire() {
return Ok(());
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
async fn execute_with_retry<F>(&self, build: F) -> Result<Response>
where
F: Fn(&Client) -> RequestBuilder,
{
self.acquire_token().await?;
let token = self.token.read().await.clone();
let mut builder = build(&self.inner);
if let Some(ref t) = token {
builder = builder.bearer_auth(t);
}
let req = builder.build()?;
let mut response = self.inner.execute(req).await?;
let mut retries = 0u32;
while response.status() == 429 && retries < MAX_RETRIES {
self.breaker.record_429().await;
let delay = exponential_backoff(retries) + jitter();
tokio::time::sleep(delay).await;
let token = self.token.read().await.clone();
let mut builder = build(&self.inner);
if let Some(ref t) = token {
builder = builder.bearer_auth(t);
}
let req = builder.build()?;
response = self.inner.execute(req).await?;
retries += 1;
}
if response.status() == 429 {
let retry_after = response
.headers()
.get("Retry-After")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok());
return Err(Error::RateLimited {
retry_after,
circuit_breaker: self.breaker.is_open().await,
});
}
self.breaker.record_success().await;
Ok(response)
}
async fn handle_response<T: DeserializeOwned>(&self, resp: Response) -> Result<T> {
let status = resp.status();
if status.is_success() {
Ok(resp.json::<T>().await?)
} else {
let code = status.as_u16();
let body = resp.text().await.unwrap_or_default();
let message = serde_json::from_str::<serde_json::Value>(&body)
.ok()
.and_then(|v| {
v.get("message")
.and_then(|m| m.as_str())
.map(|s| s.to_string())
})
.unwrap_or(body);
Err(Error::Api {
status: code,
message,
})
}
}
async fn handle_no_content(&self, resp: Response) -> Result<()> {
let status = resp.status();
if status.is_success() {
Ok(())
} else {
let code = status.as_u16();
let message = resp.text().await.unwrap_or_else(|_| status.to_string());
Err(Error::Api {
status: code,
message,
})
}
}
pub async fn get<T: DeserializeOwned>(&self, path: &str) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let resp = self
.execute_with_retry(|client| {
client
.get(u.as_str())
.header("Version", &v)
.header("Content-Type", "application/json")
})
.await?;
self.handle_response(resp).await
}
pub async fn get_with_query<T: DeserializeOwned, Q: Serialize + ?Sized>(
&self,
path: &str,
query: &Q,
) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let qv = serde_json::to_value(query).map_err(Error::Json)?;
let resp = self
.execute_with_retry(move |client| {
client
.get(u.as_str())
.header("Version", &v)
.header("Content-Type", "application/json")
.query(&qv)
})
.await?;
self.handle_response(resp).await
}
pub async fn post<T: DeserializeOwned, B: Serialize>(&self, path: &str, body: &B) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let body_json = serde_json::to_value(body).map_err(Error::Json)?;
let resp = self
.execute_with_retry(move |client| {
client
.post(u.as_str())
.header("Version", &v)
.json(&body_json)
})
.await?;
self.handle_response(resp).await
}
pub async fn put<T: DeserializeOwned, B: Serialize>(&self, path: &str, body: &B) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let body_json = serde_json::to_value(body).map_err(Error::Json)?;
let resp = self
.execute_with_retry(move |client| {
client
.put(u.as_str())
.header("Version", &v)
.json(&body_json)
})
.await?;
self.handle_response(resp).await
}
pub async fn patch<T: DeserializeOwned, B: Serialize>(
&self,
path: &str,
body: &B,
) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let body_json = serde_json::to_value(body).map_err(Error::Json)?;
let resp = self
.execute_with_retry(move |client| {
client
.patch(u.as_str())
.header("Version", &v)
.json(&body_json)
})
.await?;
self.handle_response(resp).await
}
pub async fn delete<T: DeserializeOwned>(&self, path: &str) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let resp = self
.execute_with_retry(|client| client.delete(u.as_str()).header("Version", &v))
.await?;
self.handle_response(resp).await
}
pub async fn delete_with_body<T: DeserializeOwned, B: Serialize>(
&self,
path: &str,
body: &B,
) -> Result<T> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let body_json = serde_json::to_value(body).map_err(Error::Json)?;
let resp = self
.execute_with_retry(move |client| {
client
.delete(u.as_str())
.header("Version", &v)
.json(&body_json)
})
.await?;
self.handle_response(resp).await
}
pub async fn delete_no_content(&self, path: &str) -> Result<()> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let resp = self
.execute_with_retry(|client| client.delete(u.as_str()).header("Version", &v))
.await?;
self.handle_no_content(resp).await
}
pub async fn get_raw(&self, path: &str) -> Result<Response> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
self.execute_with_retry(|client| {
client
.get(u.as_str())
.header("Version", &v)
.header("Content-Type", "application/json")
})
.await
}
pub async fn post_raw<B: Serialize>(&self, path: &str, body: &B) -> Result<Response> {
let u = Arc::new(self.url(path));
let v = self.api_version.clone();
let body_json = serde_json::to_value(body).map_err(Error::Json)?;
self.execute_with_retry(move |client| {
client
.post(u.as_str())
.header("Version", &v)
.json(&body_json)
})
.await
}
}
fn exponential_backoff(retry: u32) -> Duration {
let base_ms = 200u64 * 2u64.pow(retry);
Duration::from_millis(base_ms.min(10_000))
}
fn jitter() -> Duration {
let t = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let hash = t
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
Duration::from_millis(hash % 200)
}