use std::time::Duration;
use reqwest::{Client, StatusCode};
use serde::de::DeserializeOwned;
use tracing::warn;
use crate::error::{ExchangeError, Result};
pub(crate) fn percent_encode(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for byte in s.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
out.push(byte as char);
}
b => {
out.push('%');
out.push(
char::from_digit(u32::from(b) >> 4, 16)
.unwrap()
.to_ascii_uppercase(),
);
out.push(
char::from_digit(u32::from(b) & 0xF, 16)
.unwrap()
.to_ascii_uppercase(),
);
}
}
}
out
}
pub(crate) fn build_query_string(params: &[(&str, &str)]) -> String {
if params.is_empty() {
return String::new();
}
let pairs: Vec<String> = params
.iter()
.map(|(k, v)| format!("{}={}", percent_encode(k), percent_encode(v)))
.collect();
format!("?{}", pairs.join("&"))
}
pub(crate) fn jitter_secs(base: f64) -> f64 {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.subsec_nanos();
let factor = (f64::from(nanos) / 1_000_000_000.0 - 0.5) * 0.5;
base * factor
}
pub(crate) const DEFAULT_RETRIES: u32 = 3;
pub(crate) const DEFAULT_BACKOFF: f64 = 1.5;
pub(crate) const MAX_RATE_LIMIT_RETRIES: u32 = 5;
const DEFAULT_HTTP_TIMEOUT_SECS: u64 = 10;
#[derive(Clone)]
pub struct PublicRestClient {
http: Client,
base_url: String,
}
impl PublicRestClient {
pub fn new(base_url: impl Into<String>) -> Result<Self> {
Self::with_timeout(base_url, Duration::from_secs(DEFAULT_HTTP_TIMEOUT_SECS))
}
pub fn with_timeout(base_url: impl Into<String>, timeout: Duration) -> Result<Self> {
let http = Client::builder()
.timeout(timeout)
.build()
.map_err(|e| ExchangeError::Config(format!("failed to build HTTP client: {e}")))?;
Ok(Self {
http,
base_url: base_url.into(),
})
}
#[must_use]
pub fn base_url(&self) -> &str {
&self.base_url
}
pub async fn get<T: DeserializeOwned>(&self, path: &str, params: &[(&str, &str)]) -> Result<T> {
let qs = build_query_string(params);
let url = format!("{}{path}{qs}", self.base_url);
let mut last_err: Option<ExchangeError> = None;
let mut rate_limit_hits: u32 = 0;
for attempt in 0..DEFAULT_RETRIES {
let send_result = self.http.get(&url).send().await;
let resp = match send_result {
Ok(r) => r,
Err(e) if attempt < DEFAULT_RETRIES - 1 => {
let base = DEFAULT_BACKOFF.powi(attempt.cast_signed() + 1);
let wait = (base + jitter_secs(base)).max(0.1);
warn!(
attempt,
path,
error = %e,
wait_secs = wait,
"public GET failed, retrying"
);
tokio::time::sleep(Duration::from_secs_f64(wait)).await;
last_err = Some(ExchangeError::Http(e));
continue;
}
Err(e) => return Err(ExchangeError::Http(e)),
};
if resp.status() == StatusCode::TOO_MANY_REQUESTS {
rate_limit_hits += 1;
if rate_limit_hits > MAX_RATE_LIMIT_RETRIES {
return Err(ExchangeError::Api {
code: "429".into(),
message: format!(
"GET {path} was rate-limited \
{MAX_RATE_LIMIT_RETRIES} times; giving up"
),
});
}
let wait = parse_retry_after(&resp).unwrap_or(Duration::from_secs(2));
warn!(
attempt,
path,
wait_ms = wait.as_millis(),
rate_limit_hits,
"public GET rate-limited — waiting before retry"
);
tokio::time::sleep(wait).await;
last_err = Some(ExchangeError::Api {
code: "429".into(),
message: "rate limited".into(),
});
continue;
}
if !resp.status().is_success() {
let code = resp.status().as_u16().to_string();
let message = resp
.text()
.await
.unwrap_or_else(|_| String::from("no body"));
return Err(ExchangeError::Api { code, message });
}
return Ok(resp.json::<T>().await?);
}
Err(last_err.unwrap_or_else(|| ExchangeError::Api {
code: "retry_exhausted".into(),
message: format!("GET {path} failed after {DEFAULT_RETRIES} attempts"),
}))
}
}
fn parse_retry_after(resp: &reqwest::Response) -> Option<Duration> {
resp.headers()
.get("retry-after")
.and_then(|h| h.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
.map(Duration::from_secs)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn percent_encode_leaves_unreserved_chars_unchanged() {
assert_eq!(percent_encode("XBTUSDTM"), "XBTUSDTM");
assert_eq!(percent_encode("abc-123_def.ghi~"), "abc-123_def.ghi~");
}
#[test]
fn percent_encode_encodes_special_chars() {
assert_eq!(percent_encode("a b"), "a%20b");
assert_eq!(percent_encode("a=b&c=d"), "a%3Db%26c%3Dd");
assert_eq!(percent_encode("a+b"), "a%2Bb");
}
#[test]
fn build_query_string_empty() {
assert_eq!(build_query_string(&[]), "");
}
#[test]
fn build_query_string_encodes_values() {
let qs = build_query_string(&[("symbol", "XBT USDT"), ("side", "buy&sell")]);
assert_eq!(qs, "?symbol=XBT%20USDT&side=buy%26sell");
}
#[test]
fn jitter_stays_within_25_percent() {
let base = 4.0_f64;
for _ in 0..100 {
let j = jitter_secs(base);
assert!(
j.abs() <= base.mul_add(0.25, 1e-9),
"jitter {j} exceeded ±25% of {base}"
);
}
}
}