Skip to main content

htb_cli/api/
mod.rs

1pub mod challenges;
2pub mod ctf;
3pub mod machines;
4pub mod search;
5pub mod seasons;
6pub mod sherlocks;
7pub mod user;
8pub mod vpn;
9
10pub fn encode_path(s: &str) -> String {
11    s.bytes()
12        .flat_map(|b| {
13            if b.is_ascii_alphanumeric() || b == b'-' || b == b'_' || b == b'.' || b == b'~' {
14                vec![b as char]
15            } else {
16                format!("%{b:02X}").chars().collect()
17            }
18        })
19        .collect()
20}
21
22use std::sync::atomic::{AtomicU32, Ordering};
23use std::sync::Arc;
24use std::time::Duration;
25
26use reqwest::header::HeaderMap;
27use serde::de::DeserializeOwned;
28use serde::Serialize;
29
30use crate::cache::Cache;
31use crate::error::{ApiErrorBody, HtbError};
32
33const BASE_URL: &str = "https://labs.hackthebox.com";
34const USER_AGENT: &str = concat!("htb-cli/", env!("CARGO_PKG_VERSION"));
35
36#[derive(Clone)]
37pub struct HtbClient {
38    http: reqwest::Client,
39    base_url: String,
40    token: String,
41    rate_limit: Arc<RateLimitState>,
42    cache: Option<Arc<Cache>>,
43}
44
45struct RateLimitState {
46    remaining: AtomicU32,
47    limit: AtomicU32,
48}
49
50impl RateLimitState {
51    fn new() -> Self {
52        Self {
53            remaining: AtomicU32::new(u32::MAX),
54            limit: AtomicU32::new(u32::MAX),
55        }
56    }
57
58    fn update(&self, headers: &HeaderMap) {
59        if let Some(limit) = headers
60            .get("x-ratelimit-limit")
61            .and_then(|v| v.to_str().ok())
62            .and_then(|v| v.parse::<u32>().ok())
63        {
64            self.limit.store(limit, Ordering::Relaxed);
65        }
66
67        if let Some(remaining) = headers
68            .get("x-ratelimit-remaining")
69            .and_then(|v| v.to_str().ok())
70            .and_then(|v| v.parse::<u32>().ok())
71        {
72            self.remaining.store(remaining, Ordering::Relaxed);
73        }
74    }
75
76    fn remaining(&self) -> u32 {
77        self.remaining.load(Ordering::Relaxed)
78    }
79
80    fn limit(&self) -> u32 {
81        self.limit.load(Ordering::Relaxed)
82    }
83}
84
85impl HtbClient {
86    pub fn new(token: String) -> Self {
87        Self::build(token, BASE_URL.to_string(), None)
88    }
89
90    pub fn with_cache_arc(token: String, cache: Arc<Cache>) -> Self {
91        Self::build(token, BASE_URL.to_string(), Some(cache))
92    }
93
94    pub fn with_base_url(token: String, base_url: String) -> Self {
95        Self::build(token, base_url, None)
96    }
97
98    pub fn with_base_url_and_cache(token: String, base_url: String, cache: Arc<Cache>) -> Self {
99        Self::build(token, base_url, Some(cache))
100    }
101
102    fn build(token: String, base_url: String, cache: Option<Arc<Cache>>) -> Self {
103        let http = reqwest::Client::builder()
104            .user_agent(USER_AGENT)
105            .timeout(Duration::from_secs(30))
106            .build()
107            .expect("failed to build HTTP client");
108
109        Self {
110            http,
111            base_url,
112            token,
113            rate_limit: Arc::new(RateLimitState::new()),
114            cache,
115        }
116    }
117
118    pub async fn get<T: DeserializeOwned>(&self, path: &str) -> Result<T, HtbError> {
119        let url = format!("{}{}", self.base_url, path);
120        let ttl = self.ttl_for_path(path);
121
122        if let Some(max_age) = ttl {
123            if let Some(cache) = &self.cache {
124                if let Some(body) = cache.get(&url, max_age) {
125                    match serde_json::from_str(&body) {
126                        Ok(parsed) => return Ok(parsed),
127                        Err(e) => {
128                            tracing::debug!(
129                                "cached response failed to deserialize, refetching: {e}"
130                            );
131                        }
132                    }
133                }
134            }
135        }
136
137        self.wait_for_rate_limit().await;
138        tracing::debug!(url = %url, "GET");
139        let resp = self.http.get(&url).bearer_auth(&self.token).send().await?;
140        let body = self.handle_response_raw(resp).await?;
141
142        if let Some(max_age) = ttl {
143            if max_age > Duration::ZERO {
144                if let Some(cache) = &self.cache {
145                    cache.set(&url, &body);
146                }
147            }
148        }
149
150        Ok(serde_json::from_str(&body)?)
151    }
152
153    pub async fn post<B: Serialize, T: DeserializeOwned>(
154        &self,
155        path: &str,
156        body: &B,
157    ) -> Result<T, HtbError> {
158        self.wait_for_rate_limit().await;
159
160        let url = format!("{}{}", self.base_url, path);
161        tracing::debug!(url = %url, "POST");
162
163        let resp = self
164            .http
165            .post(&url)
166            .bearer_auth(&self.token)
167            .json(body)
168            .send()
169            .await?;
170
171        let result = self.handle_response(resp).await?;
172        self.invalidate_after_post(path);
173        Ok(result)
174    }
175
176    pub async fn post_no_content<B: Serialize>(
177        &self,
178        path: &str,
179        body: &B,
180    ) -> Result<(), HtbError> {
181        self.wait_for_rate_limit().await;
182
183        let url = format!("{}{}", self.base_url, path);
184        tracing::debug!(url = %url, "POST (no content)");
185
186        let resp = self
187            .http
188            .post(&url)
189            .bearer_auth(&self.token)
190            .json(body)
191            .send()
192            .await?;
193
194        self.rate_limit.update(resp.headers());
195        self.log_rate_limit();
196
197        let status = resp.status();
198        if status == 401 {
199            return Err(HtbError::NotAuthenticated);
200        }
201        if status == 429 {
202            return Err(HtbError::RateLimited);
203        }
204        if !status.is_success() {
205            let body = resp.text().await.unwrap_or_default();
206            let message = serde_json::from_str::<ApiErrorBody>(&body)
207                .map(|e| e.message)
208                .unwrap_or(body);
209            return Err(HtbError::Api {
210                status: status.as_u16(),
211                message,
212            });
213        }
214
215        self.invalidate_after_post(path);
216        Ok(())
217    }
218
219    fn invalidate_after_post(&self, path: &str) {
220        let Some(cache) = &self.cache else { return };
221        if path.contains("/vm/spawn")
222            || path.contains("/vm/terminate")
223            || path.contains("/vm/reset")
224            || path.contains("/machine/own")
225            || path.contains("/machine/todo")
226        {
227            cache.invalidate_pattern("api_v4_machine");
228            cache.invalidate_pattern("api_v5_machine");
229        }
230        if path.contains("/container/start")
231            || path.contains("/container/stop")
232            || path.contains("/challenge/own")
233        {
234            cache.invalidate_pattern("api_v4_challenge");
235        }
236        if path.contains("/sherlocks/") && path.contains("/flag") {
237            cache.invalidate_pattern("api_v4_sherlock");
238        }
239        // CTF mutations
240        if path.contains("/flags/own")
241            || path.contains("/challenges/containers/start")
242            || path.contains("/challenges/containers/stop")
243        {
244            cache.invalidate_pattern("ctf.hackthebox.com_api_ctfs_");
245            cache.invalidate_pattern("ctf.hackthebox.com_api_challenges_");
246        }
247    }
248
249    pub async fn get_bytes(&self, url_or_path: &str) -> Result<Vec<u8>, HtbError> {
250        self.wait_for_rate_limit().await;
251
252        let is_absolute = url_or_path.starts_with("http://") || url_or_path.starts_with("https://");
253        let url = if is_absolute {
254            url_or_path.to_string()
255        } else {
256            format!("{}{}", self.base_url, url_or_path)
257        };
258        tracing::debug!(url = %url, "GET (bytes)");
259
260        let same_origin = url.starts_with(&format!("{}/", self.base_url));
261        let req = self.http.get(&url);
262        let req = if same_origin {
263            req.bearer_auth(&self.token)
264        } else {
265            req
266        };
267        let resp = req.send().await?;
268
269        self.rate_limit.update(resp.headers());
270        self.log_rate_limit();
271
272        let status = resp.status();
273        if status == 401 {
274            return Err(HtbError::NotAuthenticated);
275        }
276        if status == 429 {
277            return Err(HtbError::RateLimited);
278        }
279        if !status.is_success() {
280            let body = resp.text().await.unwrap_or_default();
281            return Err(HtbError::Api {
282                status: status.as_u16(),
283                message: body,
284            });
285        }
286
287        Ok(resp.bytes().await?.to_vec())
288    }
289
290    async fn handle_response<T: DeserializeOwned>(
291        &self,
292        resp: reqwest::Response,
293    ) -> Result<T, HtbError> {
294        let body = self.handle_response_raw(resp).await?;
295        Ok(serde_json::from_str(&body)?)
296    }
297
298    async fn handle_response_raw(&self, resp: reqwest::Response) -> Result<String, HtbError> {
299        self.rate_limit.update(resp.headers());
300        self.log_rate_limit();
301
302        let status = resp.status();
303
304        if status == 401 {
305            return Err(HtbError::NotAuthenticated);
306        }
307
308        if status == 429 {
309            return Err(HtbError::RateLimited);
310        }
311
312        if !status.is_success() {
313            let body = resp.text().await.unwrap_or_default();
314            let message = serde_json::from_str::<ApiErrorBody>(&body)
315                .map(|e| e.message)
316                .unwrap_or(body);
317
318            return Err(HtbError::Api {
319                status: status.as_u16(),
320                message,
321            });
322        }
323
324        Ok(resp.text().await?)
325    }
326
327    async fn wait_for_rate_limit(&self) {
328        let remaining = self.rate_limit.remaining();
329        if remaining == 0 && self.rate_limit.limit() != u32::MAX {
330            // Single-threaded runtime; no concurrent task will update the atomic
331            // while we sleep. Wait once and let the next response refresh the state.
332            tracing::warn!("Rate limit exhausted, waiting 5s before next request");
333            tokio::time::sleep(Duration::from_secs(5)).await;
334        }
335    }
336
337    fn log_rate_limit(&self) {
338        let remaining = self.rate_limit.remaining();
339        let limit = self.rate_limit.limit();
340        if limit != u32::MAX {
341            tracing::debug!(remaining, limit, "rate limit");
342        }
343    }
344
345    pub fn ctf(&self) -> ctf::CtfApi<'_> {
346        ctf::CtfApi(self)
347    }
348
349    pub fn user(&self) -> user::UserApi<'_> {
350        user::UserApi(self)
351    }
352
353    pub fn machines(&self) -> machines::MachineApi<'_> {
354        machines::MachineApi(self)
355    }
356
357    pub fn challenges(&self) -> challenges::ChallengeApi<'_> {
358        challenges::ChallengeApi(self)
359    }
360
361    pub fn sherlocks(&self) -> sherlocks::SherlockApi<'_> {
362        sherlocks::SherlockApi(self)
363    }
364
365    pub fn seasons(&self) -> seasons::SeasonApi<'_> {
366        seasons::SeasonApi(self)
367    }
368
369    pub fn vpn(&self) -> vpn::VpnApi<'_> {
370        vpn::VpnApi(self)
371    }
372
373    pub fn search(&self) -> search::SearchApi<'_> {
374        search::SearchApi(self)
375    }
376
377    fn ttl_for_path(&self, path: &str) -> Option<Duration> {
378        if path.contains("/download") {
379            return None;
380        }
381
382        let is_ctf = self.base_url.contains("ctf.hackthebox.com");
383        if is_ctf {
384            return self.ttl_for_ctf_path(path);
385        }
386
387        // Labs: reference data (30 min)
388        if path.contains("/categories/list")
389            || path.contains("/season/list")
390            || path.contains("/tags/list")
391        {
392            return Some(Duration::from_secs(1800));
393        }
394        // Labs: lists (5 min)
395        if path.starts_with("/api/v5/machines")
396            || path.starts_with("/api/v4/challenges?")
397            || path.starts_with("/api/v4/sherlocks?")
398        {
399            return Some(Duration::from_secs(300));
400        }
401        // Labs: challenge/machine/sherlock details (60 min, mostly static)
402        if path.contains("/challenge/info/")
403            || path.contains("/machine/profile/")
404            || path.contains("/sherlocks/")
405        {
406            return Some(Duration::from_secs(3600));
407        }
408        // Labs: user profiles (2 min, points/rank change after submissions)
409        if path.contains("/user/info") || path.contains("/user/profile/") {
410            return Some(Duration::from_secs(120));
411        }
412        None
413    }
414
415    fn ttl_for_ctf_path(&self, path: &str) -> Option<Duration> {
416        // Reference data (30 min)
417        if path.starts_with("/api/public/challenge-categories") {
418            return Some(Duration::from_secs(1800));
419        }
420        // Event list and details (5 min)
421        if path == "/api/ctfs" || path.starts_with("/api/ctfs/details/") {
422            return Some(Duration::from_secs(300));
423        }
424        // Profiles (2 min)
425        if path.starts_with("/api/users/profile") {
426            return Some(Duration::from_secs(120));
427        }
428        // Live data: challenges, scoreboard, solves (30 s)
429        if path.starts_with("/api/ctfs/scores/")
430            || path.starts_with("/api/ctfs/solves/")
431            || path.starts_with("/api/ctfs/score-charts/")
432            || path.starts_with("/api/challenges/")
433        {
434            return Some(Duration::from_secs(30));
435        }
436        // Event data with challenges (30 s)
437        if path.starts_with("/api/ctfs/") {
438            return Some(Duration::from_secs(30));
439        }
440        None
441    }
442}
443
444#[cfg(test)]
445mod tests {
446    use super::*;
447
448    #[test]
449    fn rate_limit_state_parses_headers() {
450        let state = RateLimitState::new();
451        let mut headers = HeaderMap::new();
452        headers.insert("x-ratelimit-limit", "25".parse().unwrap());
453        headers.insert("x-ratelimit-remaining", "14".parse().unwrap());
454
455        state.update(&headers);
456        assert_eq!(state.limit(), 25);
457        assert_eq!(state.remaining(), 14);
458    }
459
460    #[test]
461    fn rate_limit_state_ignores_missing_headers() {
462        let state = RateLimitState::new();
463        let headers = HeaderMap::new();
464
465        state.update(&headers);
466        assert_eq!(state.limit(), u32::MAX);
467        assert_eq!(state.remaining(), u32::MAX);
468    }
469
470    #[test]
471    fn rate_limit_state_ignores_garbage() {
472        let state = RateLimitState::new();
473        let mut headers = HeaderMap::new();
474        headers.insert("x-ratelimit-limit", "not-a-number".parse().unwrap());
475        headers.insert("x-ratelimit-remaining", "".parse().unwrap());
476
477        state.update(&headers);
478        assert_eq!(state.limit(), u32::MAX);
479        assert_eq!(state.remaining(), u32::MAX);
480    }
481}