pub mod challenges;
pub mod machines;
pub mod search;
pub mod seasons;
pub mod sherlocks;
pub mod user;
pub mod vpn;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::Duration;
use reqwest::header::HeaderMap;
use serde::de::DeserializeOwned;
use serde::Serialize;
use crate::error::{ApiErrorBody, HtbError};
const BASE_URL: &str = "https://labs.hackthebox.com";
const USER_AGENT: &str = concat!("htb-cli/", env!("CARGO_PKG_VERSION"));
#[derive(Clone)]
pub struct HtbClient {
http: reqwest::Client,
base_url: String,
token: String,
rate_limit: Arc<RateLimitState>,
}
struct RateLimitState {
remaining: AtomicU32,
limit: AtomicU32,
}
impl RateLimitState {
fn new() -> Self {
Self {
remaining: AtomicU32::new(u32::MAX),
limit: AtomicU32::new(u32::MAX),
}
}
fn update(&self, headers: &HeaderMap) {
if let Some(limit) = headers
.get("x-ratelimit-limit")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u32>().ok())
{
self.limit.store(limit, Ordering::Relaxed);
}
if let Some(remaining) = headers
.get("x-ratelimit-remaining")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u32>().ok())
{
self.remaining.store(remaining, Ordering::Relaxed);
}
}
fn remaining(&self) -> u32 {
self.remaining.load(Ordering::Relaxed)
}
fn limit(&self) -> u32 {
self.limit.load(Ordering::Relaxed)
}
}
impl HtbClient {
pub fn new(token: String) -> Self {
let http = reqwest::Client::builder()
.user_agent(USER_AGENT)
.timeout(Duration::from_secs(30))
.build()
.expect("failed to build HTTP client");
Self {
http,
base_url: BASE_URL.to_string(),
token,
rate_limit: Arc::new(RateLimitState::new()),
}
}
#[cfg(test)]
pub fn with_base_url(token: String, base_url: String) -> Self {
let http = reqwest::Client::builder()
.user_agent(USER_AGENT)
.timeout(Duration::from_secs(30))
.build()
.expect("failed to build HTTP client");
Self {
http,
base_url,
token,
rate_limit: Arc::new(RateLimitState::new()),
}
}
pub async fn get<T: DeserializeOwned>(&self, path: &str) -> Result<T, HtbError> {
self.wait_for_rate_limit().await;
let url = format!("{}{}", self.base_url, path);
tracing::debug!(url = %url, "GET");
let resp = self.http.get(&url).bearer_auth(&self.token).send().await?;
self.handle_response(resp).await
}
pub async fn post<B: Serialize, T: DeserializeOwned>(
&self,
path: &str,
body: &B,
) -> Result<T, HtbError> {
self.wait_for_rate_limit().await;
let url = format!("{}{}", self.base_url, path);
tracing::debug!(url = %url, "POST");
let resp = self
.http
.post(&url)
.bearer_auth(&self.token)
.json(body)
.send()
.await?;
self.handle_response(resp).await
}
pub async fn get_bytes(&self, url_or_path: &str) -> Result<Vec<u8>, HtbError> {
self.wait_for_rate_limit().await;
let is_absolute = url_or_path.starts_with("http://") || url_or_path.starts_with("https://");
let url = if is_absolute {
url_or_path.to_string()
} else {
format!("{}{}", self.base_url, url_or_path)
};
tracing::debug!(url = %url, "GET (bytes)");
let req = self.http.get(&url);
let req = if is_absolute {
req
} else {
req.bearer_auth(&self.token)
};
let resp = req.send().await?;
self.rate_limit.update(resp.headers());
self.log_rate_limit();
let status = resp.status();
if status == 401 {
return Err(HtbError::NotAuthenticated);
}
if status == 429 {
return Err(HtbError::RateLimited);
}
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
return Err(HtbError::Api {
status: status.as_u16(),
message: body,
});
}
Ok(resp.bytes().await?.to_vec())
}
async fn handle_response<T: DeserializeOwned>(
&self,
resp: reqwest::Response,
) -> Result<T, HtbError> {
self.rate_limit.update(resp.headers());
self.log_rate_limit();
let status = resp.status();
if status == 401 {
return Err(HtbError::NotAuthenticated);
}
if status == 429 {
return Err(HtbError::RateLimited);
}
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
let message = serde_json::from_str::<ApiErrorBody>(&body)
.map(|e| e.message)
.unwrap_or(body);
return Err(HtbError::Api {
status: status.as_u16(),
message,
});
}
let body = resp.text().await?;
let parsed = serde_json::from_str(&body)?;
Ok(parsed)
}
async fn wait_for_rate_limit(&self) {
let remaining = self.rate_limit.remaining();
if remaining == 0 && self.rate_limit.limit() != u32::MAX {
tracing::warn!("Rate limit exhausted, waiting 5s before next request");
tokio::time::sleep(Duration::from_secs(5)).await;
}
}
fn log_rate_limit(&self) {
let remaining = self.rate_limit.remaining();
let limit = self.rate_limit.limit();
if limit != u32::MAX {
tracing::debug!(remaining, limit, "rate limit");
}
}
pub fn user(&self) -> user::UserApi<'_> {
user::UserApi(self)
}
pub fn machines(&self) -> machines::MachineApi<'_> {
machines::MachineApi(self)
}
pub fn challenges(&self) -> challenges::ChallengeApi<'_> {
challenges::ChallengeApi(self)
}
pub fn sherlocks(&self) -> sherlocks::SherlockApi<'_> {
sherlocks::SherlockApi(self)
}
pub fn seasons(&self) -> seasons::SeasonApi<'_> {
seasons::SeasonApi(self)
}
pub fn vpn(&self) -> vpn::VpnApi<'_> {
vpn::VpnApi(self)
}
pub fn search(&self) -> search::SearchApi<'_> {
search::SearchApi(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rate_limit_state_parses_headers() {
let state = RateLimitState::new();
let mut headers = HeaderMap::new();
headers.insert("x-ratelimit-limit", "25".parse().unwrap());
headers.insert("x-ratelimit-remaining", "14".parse().unwrap());
state.update(&headers);
assert_eq!(state.limit(), 25);
assert_eq!(state.remaining(), 14);
}
#[test]
fn rate_limit_state_ignores_missing_headers() {
let state = RateLimitState::new();
let headers = HeaderMap::new();
state.update(&headers);
assert_eq!(state.limit(), u32::MAX);
assert_eq!(state.remaining(), u32::MAX);
}
#[test]
fn rate_limit_state_ignores_garbage() {
let state = RateLimitState::new();
let mut headers = HeaderMap::new();
headers.insert("x-ratelimit-limit", "not-a-number".parse().unwrap());
headers.insert("x-ratelimit-remaining", "".parse().unwrap());
state.update(&headers);
assert_eq!(state.limit(), u32::MAX);
assert_eq!(state.remaining(), u32::MAX);
}
}