Skip to main content

rightkit_http/
client.rs

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    /// RFC 9110 idempotent methods (safe to retry after an ambiguous failure).
34    pub fn is_idempotent(self) -> bool {
35        !matches!(self, Method::Post | Method::Patch)
36    }
37}
38
39/// One retry decision, reported to [`ClientConfig::on_retry`].
40#[derive(Debug, Clone)]
41pub struct RetryEvent {
42    pub url: String,
43    /// Retry number about to run (1 = first retry).
44    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    /// Whole-request deadline for buffered calls (headers + body).
56    pub request_timeout: Duration,
57    /// Time to receive response headers; for streams this is the first-byte
58    /// (first-token) deadline.
59    pub response_timeout: Duration,
60    /// Optional total lifetime cap for a stream body.
61    pub stream_timeout: Option<Duration>,
62    /// Redirects followed (0 = never; credentials are never replayed to a new host).
63    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    /// Allow retrying a POST/PATCH after HTTP-status failures (default true for
94    /// status codes; transport errors on non-idempotent methods only retry in
95    /// the connect phase regardless).
96    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    /// `application/x-www-form-urlencoded` body.
135    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
153/// Percent-encode for `application/x-www-form-urlencoded` / URL query values.
154pub 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    /// Lower-cased header names.
205    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    /// Decode the body as Server-Sent Events.
237    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    /// Buffered request; non-2xx becomes [`HttpError::Status`].
280    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    /// Buffered request returning any final status (retryable statuses are still retried).
295    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    /// Streaming request; non-2xx becomes [`HttpError::Status`]. Retries happen
309    /// only before the response headers arrive (never mid-stream).
310    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                    // Drain so the connection is released.
352                    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}