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