1use std::io::Read;
2use std::sync::Arc;
3use std::time::Duration;
4
5use serde::de::DeserializeOwned;
6use serde::Serialize;
7
8use crate::error::{HttpError, MAX_ERROR_BODY_BYTES};
9use crate::retry::{parse_retry_after, RetryPolicy};
10use crate::sse::SseReader;
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum Method {
14 Get,
15 Post,
16 Put,
17 Patch,
18 Delete,
19 Head,
20}
21
22impl Method {
23 pub fn as_str(self) -> &'static str {
24 match self {
25 Method::Get => "GET",
26 Method::Post => "POST",
27 Method::Put => "PUT",
28 Method::Patch => "PATCH",
29 Method::Delete => "DELETE",
30 Method::Head => "HEAD",
31 }
32 }
33 pub fn is_idempotent(self) -> bool {
35 !matches!(self, Method::Post | Method::Patch)
36 }
37}
38
39#[derive(Debug, Clone)]
41pub struct RetryEvent {
42 pub url: String,
43 pub retry: u32,
45 pub reason: String,
46 pub delay: Duration,
47}
48
49type RetryObserver = Arc<dyn Fn(&RetryEvent) + Send + Sync>;
50
51#[derive(Clone)]
52pub struct ClientConfig {
53 pub user_agent: String,
54 pub connect_timeout: Duration,
55 pub request_timeout: Duration,
57 pub response_timeout: Duration,
60 pub stream_timeout: Option<Duration>,
62 pub max_redirects: u32,
64 pub max_response_bytes: usize,
65 pub retry: RetryPolicy,
66 pub on_retry: Option<RetryObserver>,
67}
68
69impl Default for ClientConfig {
70 fn default() -> Self {
71 Self {
72 user_agent: concat!("rightkit-http/", env!("CARGO_PKG_VERSION")).to_string(),
73 connect_timeout: Duration::from_secs(10),
74 request_timeout: Duration::from_secs(60),
75 response_timeout: Duration::from_secs(60),
76 stream_timeout: None,
77 max_redirects: 0,
78 max_response_bytes: 32 * 1024 * 1024,
79 retry: RetryPolicy::default(),
80 on_retry: None,
81 }
82 }
83}
84
85#[derive(Debug, Clone)]
86pub struct Request {
87 pub method: Method,
88 pub url: String,
89 pub headers: Vec<(String, String)>,
90 pub body: Vec<u8>,
91 pub timeout: Option<Duration>,
92 pub retry: Option<RetryPolicy>,
93 pub retry_unsafe_statuses: bool,
97}
98
99impl Request {
100 pub fn new(method: Method, url: impl Into<String>) -> Self {
101 Self {
102 method,
103 url: url.into(),
104 headers: Vec::new(),
105 body: Vec::new(),
106 timeout: None,
107 retry: None,
108 retry_unsafe_statuses: true,
109 }
110 }
111 pub fn get(url: impl Into<String>) -> Self {
112 Self::new(Method::Get, url)
113 }
114 pub fn post(url: impl Into<String>) -> Self {
115 Self::new(Method::Post, url)
116 }
117 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
118 self.headers.push((name.into(), value.into()));
119 self
120 }
121 pub fn bearer(self, token: &str) -> Self {
122 self.header("Authorization", format!("Bearer {token}"))
123 }
124 pub fn body(mut self, body: impl Into<Vec<u8>>) -> Self {
125 self.body = body.into();
126 self
127 }
128 pub fn json<T: Serialize>(mut self, value: &T) -> Result<Self, HttpError> {
129 self.body = serde_json::to_vec(value).map_err(|e| HttpError::Json(e.to_string()))?;
130 self.headers
131 .push(("Content-Type".into(), "application/json".into()));
132 Ok(self)
133 }
134 pub fn form(mut self, pairs: &[(&str, &str)]) -> Self {
136 self.body = form_encode(pairs).into_bytes();
137 self.headers.push((
138 "Content-Type".into(),
139 "application/x-www-form-urlencoded".into(),
140 ));
141 self
142 }
143 pub fn timeout(mut self, d: Duration) -> Self {
144 self.timeout = Some(d);
145 self
146 }
147 pub fn retry(mut self, policy: RetryPolicy) -> Self {
148 self.retry = Some(policy);
149 self
150 }
151}
152
153pub fn percent_encode(input: &str) -> String {
155 let mut out = String::with_capacity(input.len());
156 for b in input.bytes() {
157 match b {
158 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
159 out.push(b as char)
160 }
161 b' ' => out.push('+'),
162 _ => out.push_str(&format!("%{b:02X}")),
163 }
164 }
165 out
166}
167
168pub fn percent_decode(input: &str) -> String {
169 let bytes = input.as_bytes();
170 let mut out = Vec::with_capacity(bytes.len());
171 let mut i = 0;
172 while i < bytes.len() {
173 match bytes[i] {
174 b'+' => out.push(b' '),
175 b'%' if i + 2 < bytes.len() => {
176 let hi = (bytes[i + 1] as char).to_digit(16);
177 let lo = (bytes[i + 2] as char).to_digit(16);
178 match (hi, lo) {
179 (Some(h), Some(l)) => {
180 out.push((h * 16 + l) as u8);
181 i += 2;
182 }
183 _ => out.push(b'%'),
184 }
185 }
186 b => out.push(b),
187 }
188 i += 1;
189 }
190 String::from_utf8_lossy(&out).into_owned()
191}
192
193pub fn form_encode(pairs: &[(&str, &str)]) -> String {
194 pairs
195 .iter()
196 .map(|(k, v)| format!("{}={}", percent_encode(k), percent_encode(v)))
197 .collect::<Vec<_>>()
198 .join("&")
199}
200
201#[derive(Debug, Clone)]
202pub struct Response {
203 pub status: u16,
204 pub headers: Vec<(String, String)>,
206 pub body: Vec<u8>,
207 pub attempts: u32,
208}
209
210impl Response {
211 pub fn header(&self, name: &str) -> Option<&str> {
212 header_of(&self.headers, name)
213 }
214 pub fn text(&self) -> String {
215 String::from_utf8_lossy(&self.body).into_owned()
216 }
217 pub fn json<T: DeserializeOwned>(&self) -> Result<T, HttpError> {
218 serde_json::from_slice(&self.body).map_err(|e| HttpError::Json(e.to_string()))
219 }
220}
221
222pub struct StreamResponse {
223 pub status: u16,
224 pub headers: Vec<(String, String)>,
225 pub attempts: u32,
226 reader: Box<dyn Read + Send>,
227}
228
229impl StreamResponse {
230 pub fn header(&self, name: &str) -> Option<&str> {
231 header_of(&self.headers, name)
232 }
233 pub fn into_reader(self) -> Box<dyn Read + Send> {
234 self.reader
235 }
236 pub fn sse(self) -> SseReader<Box<dyn Read + Send>> {
238 SseReader::new(self.reader)
239 }
240}
241
242fn header_of<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
243 headers
244 .iter()
245 .find(|(k, _)| k.eq_ignore_ascii_case(name))
246 .map(|(_, v)| v.as_str())
247}
248
249#[derive(Clone)]
250pub struct Client {
251 agent: ureq::Agent,
252 config: ClientConfig,
253}
254
255struct Raw {
256 status: u16,
257 headers: Vec<(String, String)>,
258 body: ureq::Body,
259 attempts: u32,
260}
261
262impl Client {
263 pub fn new(config: ClientConfig) -> Self {
264 let agent: ureq::Agent = ureq::Agent::config_builder()
265 .http_status_as_error(false)
266 .max_redirects(config.max_redirects)
267 .max_redirects_will_error(false)
268 .user_agent(config.user_agent.clone())
269 .timeout_connect(Some(config.connect_timeout))
270 .build()
271 .into();
272 Self { agent, config }
273 }
274
275 pub fn config(&self) -> &ClientConfig {
276 &self.config
277 }
278
279 pub fn send(&self, req: &Request) -> Result<Response, HttpError> {
281 let resp = self.send_any(req)?;
282 if (200..300).contains(&resp.status) {
283 Ok(resp)
284 } else {
285 Err(status_error(
286 resp.status,
287 &resp.headers,
288 resp.body,
289 resp.attempts,
290 ))
291 }
292 }
293
294 pub fn send_any(&self, req: &Request) -> Result<Response, HttpError> {
296 let timeout = req.timeout.unwrap_or(self.config.request_timeout);
297 let mut raw = self.execute(req, Some(timeout))?;
298 let mut reader = raw.body.as_reader();
299 let body = read_limited(&mut reader, self.config.max_response_bytes)?;
300 Ok(Response {
301 status: raw.status,
302 headers: raw.headers,
303 body,
304 attempts: raw.attempts,
305 })
306 }
307
308 pub fn stream(&self, req: &Request) -> Result<StreamResponse, HttpError> {
311 let mut raw = self.execute(req, self.config.stream_timeout)?;
312 if !(200..300).contains(&raw.status) {
313 let mut reader = raw.body.as_reader();
314 let body = read_limited(&mut reader, MAX_ERROR_BODY_BYTES).unwrap_or_default();
315 return Err(status_error(raw.status, &raw.headers, body, raw.attempts));
316 }
317 Ok(StreamResponse {
318 status: raw.status,
319 headers: raw.headers,
320 attempts: raw.attempts,
321 reader: Box::new(raw.body.into_reader()),
322 })
323 }
324
325 pub fn get_json<T: DeserializeOwned>(&self, url: &str) -> Result<T, HttpError> {
326 self.send(&Request::get(url).header("Accept", "application/json"))?
327 .json()
328 }
329
330 fn execute(&self, req: &Request, total: Option<Duration>) -> Result<Raw, HttpError> {
331 let policy = req
332 .retry
333 .clone()
334 .unwrap_or_else(|| self.config.retry.clone());
335 let max = policy.max_attempts.max(1);
336 let mut attempt = 0u32;
337 loop {
338 attempt += 1;
339 let outcome = self.once(req, total);
340 let (reason, retry_after) = match outcome {
341 Ok(raw) => {
342 let retryable = policy.is_retryable_status(raw.status)
343 && (req.method.is_idempotent() || req.retry_unsafe_statuses);
344 if !retryable || attempt >= max {
345 return Ok(Raw {
346 attempts: attempt,
347 ..raw
348 });
349 }
350 let ra = header_of(&raw.headers, "retry-after").and_then(parse_retry_after);
351 let mut body = raw.body;
353 let mut reader = body.as_reader();
354 let _ = read_limited(&mut reader, MAX_ERROR_BODY_BYTES);
355 (format!("HTTP {}", raw.status), ra)
356 }
357 Err(err) => {
358 let retryable = match &err {
359 HttpError::Transport { connect_phase, .. } => {
360 *connect_phase || req.method.is_idempotent()
361 }
362 HttpError::Timeout(_) => req.method.is_idempotent(),
363 _ => false,
364 };
365 if !retryable || attempt >= max {
366 return Err(err);
367 }
368 (err.to_string(), None)
369 }
370 };
371 let delay = policy.delay(attempt, retry_after);
372 if let Some(obs) = &self.config.on_retry {
373 obs(&RetryEvent {
374 url: req.url.clone(),
375 retry: attempt,
376 reason,
377 delay,
378 });
379 }
380 std::thread::sleep(delay);
381 }
382 }
383
384 fn once(&self, req: &Request, total: Option<Duration>) -> Result<Raw, HttpError> {
385 let mut builder = ureq::http::Request::builder()
386 .method(req.method.as_str())
387 .uri(&req.url);
388 for (k, v) in &req.headers {
389 builder = builder.header(k.as_str(), v.as_str());
390 }
391 let request = builder
392 .body(req.body.clone())
393 .map_err(|e| HttpError::InvalidRequest(e.to_string()))?;
394 let configured = self
395 .agent
396 .configure_request(request)
397 .timeout_global(total)
398 .timeout_recv_response(Some(self.config.response_timeout))
399 .build();
400 let response = self.agent.run(configured).map_err(map_ureq_error)?;
401 let status = response.status().as_u16();
402 let headers = response
403 .headers()
404 .iter()
405 .map(|(k, v)| {
406 (
407 k.as_str().to_ascii_lowercase(),
408 String::from_utf8_lossy(v.as_bytes()).into_owned(),
409 )
410 })
411 .collect();
412 Ok(Raw {
413 status,
414 headers,
415 body: response.into_body(),
416 attempts: 1,
417 })
418 }
419}
420
421fn status_error(
422 status: u16,
423 headers: &[(String, String)],
424 body: Vec<u8>,
425 attempts: u32,
426) -> HttpError {
427 let mut body = body;
428 body.truncate(MAX_ERROR_BODY_BYTES);
429 HttpError::Status {
430 status,
431 body: String::from_utf8_lossy(&body).into_owned(),
432 retry_after: header_of(headers, "retry-after").and_then(parse_retry_after),
433 attempts,
434 }
435}
436
437fn read_limited(reader: &mut dyn Read, limit: usize) -> Result<Vec<u8>, HttpError> {
438 let mut out = Vec::new();
439 reader
440 .take(limit as u64 + 1)
441 .read_to_end(&mut out)
442 .map_err(|e| match e.kind() {
443 std::io::ErrorKind::TimedOut => HttpError::Timeout(e.to_string()),
444 _ => HttpError::Transport {
445 message: e.to_string(),
446 connect_phase: false,
447 },
448 })?;
449 if out.len() > limit {
450 return Err(HttpError::BodyTooLarge { limit });
451 }
452 Ok(out)
453}
454
455fn map_ureq_error(err: ureq::Error) -> HttpError {
456 use ureq::Error as E;
457 match err {
458 E::Timeout(t) => HttpError::Timeout(format!("{t:?}")),
459 E::HostNotFound => HttpError::Transport {
460 message: "host not found".into(),
461 connect_phase: true,
462 },
463 E::ConnectionFailed => HttpError::Transport {
464 message: "connection failed".into(),
465 connect_phase: true,
466 },
467 E::Io(io)
468 if matches!(
469 io.kind(),
470 std::io::ErrorKind::ConnectionRefused | std::io::ErrorKind::NotFound
471 ) =>
472 {
473 HttpError::Transport {
474 message: io.to_string(),
475 connect_phase: true,
476 }
477 }
478 E::Io(io) => HttpError::Transport {
479 message: io.to_string(),
480 connect_phase: false,
481 },
482 E::BadUri(m) => HttpError::InvalidRequest(m),
483 E::Http(e) => HttpError::InvalidRequest(e.to_string()),
484 other => HttpError::Transport {
485 message: other.to_string(),
486 connect_phase: false,
487 },
488 }
489}