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 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 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 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 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 if path.contains("/challenge/info/")
403 || path.contains("/machine/profile/")
404 || path.contains("/sherlocks/")
405 {
406 return Some(Duration::from_secs(3600));
407 }
408 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 if path.starts_with("/api/public/challenge-categories") {
418 return Some(Duration::from_secs(1800));
419 }
420 if path == "/api/ctfs" || path.starts_with("/api/ctfs/details/") {
422 return Some(Duration::from_secs(300));
423 }
424 if path.starts_with("/api/users/profile") {
426 return Some(Duration::from_secs(120));
427 }
428 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 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}