use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use crate::error::{ApiError, ApiErrorBody, ApifyClientError, ApifyClientResult};
const RATE_LIMIT_EXCEEDED_STATUS_CODE: u16 = 429;
const MIN_SERVER_ERROR_STATUS_CODE: u16 = 500;
const MAX_SUCCESS_STATUS_CODE: u16 = 300;
const BACKOFF_FACTOR: u32 = 2;
const MIN_COMPRESS_BYTES: usize = 1024;
const CONTENT_ENCODING_BROTLI: &str = "br";
const CONTENT_ENCODING_GZIP: &str = "gzip";
const BROTLI_QUALITY: u32 = 6;
const BROTLI_WINDOW_SIZE: u32 = 22;
const BROTLI_BUFFER_SIZE: usize = 4096;
const GZIP_COMPRESSION_LEVEL: u32 = 6;
const ALREADY_COMPRESSED_MEDIA_TYPE_PREFIXES: [&str; 3] = ["audio/", "image/", "video/"];
const ALREADY_COMPRESSED_MEDIA_TYPES: [&str; 19] = [
"application/epub+zip",
"application/gzip",
"application/java-archive",
"application/vnd.android.package-archive",
"application/vnd.openxmlformats-officedocument.presentationml.presentation",
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
"application/vnd.rar",
"application/x-7z-compressed",
"application/x-bzip",
"application/x-bzip2",
"application/x-gzip",
"application/x-rar-compressed",
"application/x-xz",
"application/x-zip-compressed",
"application/zip",
"application/zstd",
"font/woff",
"font/woff2",
];
const COMPRESSIBLE_MEDIA_TYPES: [&str; 16] = [
"audio/aiff",
"audio/basic",
"audio/l16",
"audio/l24",
"audio/midi",
"audio/vnd.wave",
"audio/wav",
"audio/wave",
"audio/x-aiff",
"audio/x-wav",
"image/bmp",
"image/tiff",
"image/vnd.adobe.photoshop",
"image/vnd.microsoft.icon",
"image/x-icon",
"image/x-ms-bmp",
];
const COMPRESSIBLE_MEDIA_TYPE_SUFFIXES: [&str; 2] = ["+json", "+xml"];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum RequestCompression {
#[default]
Brotli,
Gzip,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HttpMethod {
Get,
Post,
Put,
Delete,
Head,
}
impl HttpMethod {
pub fn as_str(&self) -> &'static str {
match self {
HttpMethod::Get => "GET",
HttpMethod::Post => "POST",
HttpMethod::Put => "PUT",
HttpMethod::Delete => "DELETE",
HttpMethod::Head => "HEAD",
}
}
}
#[derive(Debug, Clone)]
pub struct HttpRequest {
pub method: HttpMethod,
pub url: String,
pub headers: HashMap<String, String>,
pub body: Option<Vec<u8>>,
pub timeout: Duration,
}
#[derive(Debug, Clone)]
pub struct HttpResponse {
pub status: u16,
pub headers: HashMap<String, String>,
pub body: Vec<u8>,
}
impl HttpResponse {
pub fn header(&self, name: &str) -> Option<&str> {
let lower = name.to_ascii_lowercase();
self.headers
.iter()
.find(|(k, _)| k.to_ascii_lowercase() == lower)
.map(|(_, v)| v.as_str())
}
}
#[async_trait]
pub trait HttpBackend: Send + Sync + std::fmt::Debug {
async fn send(&self, request: HttpRequest) -> ApifyClientResult<HttpResponse>;
}
#[derive(Debug, Clone)]
pub struct ReqwestBackend {
client: reqwest::Client,
}
impl ReqwestBackend {
pub fn new() -> Self {
Self {
client: reqwest::Client::new(),
}
}
pub fn with_client(client: reqwest::Client) -> Self {
Self { client }
}
}
impl Default for ReqwestBackend {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl HttpBackend for ReqwestBackend {
async fn send(&self, request: HttpRequest) -> ApifyClientResult<HttpResponse> {
let method = match request.method {
HttpMethod::Get => reqwest::Method::GET,
HttpMethod::Post => reqwest::Method::POST,
HttpMethod::Put => reqwest::Method::PUT,
HttpMethod::Delete => reqwest::Method::DELETE,
HttpMethod::Head => reqwest::Method::HEAD,
};
let mut builder = self
.client
.request(method, &request.url)
.timeout(request.timeout);
for (key, value) in &request.headers {
builder = builder.header(key, value);
}
if let Some(body) = request.body {
builder = builder.body(body);
}
let response = builder.send().await?;
let status = response.status().as_u16();
let mut headers = HashMap::new();
for (name, value) in response.headers().iter() {
if let Ok(v) = value.to_str() {
headers.insert(name.as_str().to_string(), v.to_string());
}
}
let body = response.bytes().await?.to_vec();
Ok(HttpResponse {
status,
headers,
body,
})
}
}
#[derive(Debug, Clone)]
pub struct RetryConfig {
pub max_retries: u32,
pub min_delay_between_retries: Duration,
pub timeout: Duration,
}
#[derive(Debug, Clone)]
pub struct HttpClient {
backend: Arc<dyn HttpBackend>,
token: Option<String>,
user_agent: String,
retry: RetryConfig,
compression: RequestCompression,
}
impl HttpClient {
pub(crate) fn new(
backend: Arc<dyn HttpBackend>,
token: Option<String>,
user_agent: String,
retry: RetryConfig,
compression: RequestCompression,
) -> Self {
Self {
backend,
token,
user_agent,
retry,
compression,
}
}
pub async fn call(&self, mut request: HttpRequest) -> ApifyClientResult<HttpResponse> {
request
.headers
.insert("User-Agent".to_string(), self.user_agent.clone());
if let Some(token) = &self.token {
request
.headers
.insert("Authorization".to_string(), format!("Bearer {token}"));
}
maybe_compress_request(&mut request, self.compression);
let method_str = request.method.as_str().to_string();
let path = extract_path(&request.url);
let base_timeout = request.timeout;
let mut delay = self.retry.min_delay_between_retries;
let max_attempts = self.retry.max_retries.saturating_add(1);
let mut attempt = 1;
loop {
let mut attempt_request = request.clone();
attempt_request.timeout = self.attempt_timeout(base_timeout, attempt);
let outcome = match self.backend.send(attempt_request).await {
Ok(response) => {
if response.status < MAX_SUCCESS_STATUS_CODE {
return Ok(response);
}
let api_error = build_api_error(&response, attempt, &method_str, &path);
let retryable = is_status_retryable(response.status);
(ApifyClientError::from(api_error), retryable)
}
Err(err) => {
let retryable = is_error_retryable(&err);
(err, retryable)
}
};
let (error, retryable) = outcome;
if !retryable || attempt == max_attempts {
return Err(error);
}
sleep(randomized_delay(delay)).await;
delay = delay.saturating_mul(BACKOFF_FACTOR).min(self.retry.timeout);
attempt += 1;
}
}
fn attempt_timeout(&self, base: Duration, attempt: u32) -> Duration {
let scaled = base.saturating_mul(2u32.saturating_pow(attempt.saturating_sub(1)));
scaled.min(self.retry.timeout)
}
pub(crate) fn user_agent(&self) -> &str {
&self.user_agent
}
pub(crate) fn stream_credentials(&self) -> (Option<String>, String) {
(self.token.clone(), self.user_agent.clone())
}
}
fn maybe_compress_request(request: &mut HttpRequest, compression: RequestCompression) {
let Some(body) = request.body.as_ref() else {
return;
};
if body.len() < MIN_COMPRESS_BYTES {
return;
}
let already_encoded = request
.headers
.keys()
.any(|k| k.eq_ignore_ascii_case("Content-Encoding"));
if already_encoded {
return;
}
let content_type = request
.headers
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case("Content-Type"))
.map(|(_, v)| v.as_str());
if !is_compressible_content_type(content_type) {
return;
}
let (encoding, compressed) = match compression {
RequestCompression::Brotli => (CONTENT_ENCODING_BROTLI, brotli_compress(body)),
RequestCompression::Gzip => (CONTENT_ENCODING_GZIP, gzip_compress(body)),
};
request
.headers
.insert("Content-Encoding".to_string(), encoding.to_string());
request.body = Some(compressed);
}
fn is_compressible_content_type(content_type: Option<&str>) -> bool {
let Some(content_type) = content_type else {
return true;
};
let media_type = content_type
.split(';')
.next()
.unwrap_or(content_type)
.trim()
.to_ascii_lowercase();
if COMPRESSIBLE_MEDIA_TYPES.contains(&media_type.as_str()) {
return true;
}
if COMPRESSIBLE_MEDIA_TYPE_SUFFIXES
.iter()
.any(|suffix| media_type.ends_with(suffix))
{
return true;
}
if ALREADY_COMPRESSED_MEDIA_TYPES.contains(&media_type.as_str()) {
return false;
}
!ALREADY_COMPRESSED_MEDIA_TYPE_PREFIXES
.iter()
.any(|prefix| media_type.starts_with(prefix))
}
fn brotli_compress(data: &[u8]) -> Vec<u8> {
use std::io::Write;
let mut writer = brotli::CompressorWriter::new(
Vec::new(),
BROTLI_BUFFER_SIZE,
BROTLI_QUALITY,
BROTLI_WINDOW_SIZE,
);
writer
.write_all(data)
.expect("writing to an in-memory Vec never fails");
writer.into_inner()
}
fn gzip_compress(data: &[u8]) -> Vec<u8> {
use flate2::{write::GzEncoder, Compression};
use std::io::Write;
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(GZIP_COMPRESSION_LEVEL));
encoder
.write_all(data)
.expect("writing to an in-memory Vec never fails");
encoder
.finish()
.expect("finishing an in-memory Vec never fails")
}
pub(crate) fn extract_path(url: &str) -> Option<String> {
let after_scheme = url.split_once("://").map(|(_, rest)| rest).unwrap_or(url);
after_scheme
.find('/')
.map(|idx| after_scheme[idx..].to_string())
}
fn is_status_retryable(status: u16) -> bool {
status == RATE_LIMIT_EXCEEDED_STATUS_CODE || status >= MIN_SERVER_ERROR_STATUS_CODE
}
fn is_error_retryable(err: &ApifyClientError) -> bool {
matches!(err, ApifyClientError::Http(_) | ApifyClientError::Timeout)
}
pub(crate) fn build_api_error(
response: &HttpResponse,
attempt: u32,
method: &str,
path: &Option<String>,
) -> ApiError {
let parsed: Option<ApiErrorBody> = serde_json::from_slice(&response.body).ok();
let (error_type, message, data) = match parsed {
Some(body) => (
body.error.error_type,
body.error
.message
.unwrap_or_else(|| format!("Unexpected error with status {}", response.status)),
body.error.data,
),
None => {
let raw = String::from_utf8_lossy(&response.body);
let message = if raw.trim().is_empty() {
format!("Unexpected error with status {}", response.status)
} else {
format!("Unexpected error: {raw}")
};
(None, message, None)
}
};
ApiError {
status_code: response.status,
error_type,
message,
attempt,
http_method: Some(method.to_string()),
path: path.clone(),
data,
}
}
pub(crate) fn randomized_delay(delay: Duration) -> Duration {
let base = delay.as_millis() as u64;
if base == 0 {
return delay;
}
let extra = next_jitter() % base;
Duration::from_millis(base + extra)
}
fn next_jitter() -> u64 {
use std::sync::atomic::{AtomicU64, Ordering};
static STATE: AtomicU64 = AtomicU64::new(0);
const GOLDEN_GAMMA: u64 = 0x9E3779B97F4A7C15;
if STATE.load(Ordering::Relaxed) == 0 {
let seed = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(GOLDEN_GAMMA)
| 1;
let _ = STATE.compare_exchange(0, seed, Ordering::Relaxed, Ordering::Relaxed);
}
let mut z = STATE
.fetch_add(GOLDEN_GAMMA, Ordering::Relaxed)
.wrapping_add(GOLDEN_GAMMA);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
pub(crate) async fn sleep_public(duration: Duration) {
sleep(duration).await;
}
async fn sleep(duration: Duration) {
tokio::time::sleep(duration).await;
}