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 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 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 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 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 if path.contains("/challenge/info/")
419 || path.contains("/machine/profile/")
420 || path.contains("/sherlocks/")
421 {
422 return Some(Duration::from_secs(3600));
423 }
424 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 if path.starts_with("/api/public/challenge-categories") {
434 return Some(Duration::from_secs(1800));
435 }
436 if path == "/api/ctfs" || path.starts_with("/api/ctfs/details/") {
438 return Some(Duration::from_secs(300));
439 }
440 if path.starts_with("/api/users/profile") {
442 return Some(Duration::from_secs(120));
443 }
444 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 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}