Skip to main content

alpaca_http/
client.rs

1use std::sync::Arc;
2use std::time::{Duration, Instant};
3
4use alpaca_core::BaseUrl;
5use reqwest::{
6    StatusCode,
7    header::{CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue},
8};
9use serde::de::DeserializeOwned;
10
11use crate::Error;
12use crate::auth::Authenticator;
13use crate::meta::{ErrorMeta, HttpResponse, ResponseMeta};
14use crate::observer::{
15    ErrorEvent, NoopObserver, RequestStart, ResponseEvent, RetryEvent, TransportObserver,
16};
17use crate::rate_limit::ConcurrencyLimit;
18use crate::request::{NoContent, RequestBody, RequestParts};
19use crate::retry::{RetryConfig, RetryDecision};
20
21#[derive(Clone)]
22pub struct HttpClient {
23    client: reqwest::Client,
24    default_headers: HeaderMap,
25    request_id_header_name: HeaderName,
26    retry_config: RetryConfig,
27    observer: Arc<dyn TransportObserver>,
28    concurrency_limit: ConcurrencyLimit,
29}
30
31#[derive(Clone)]
32pub struct HttpClientBuilder {
33    reqwest_client: Option<reqwest::Client>,
34    timeout: Duration,
35    default_headers: HeaderMap,
36    request_id_header_name: HeaderName,
37    retry_config: RetryConfig,
38    observer: Arc<dyn TransportObserver>,
39    concurrency_limit: ConcurrencyLimit,
40}
41
42struct ResponseParts {
43    meta: ResponseMeta,
44    body: String,
45}
46
47impl HttpClient {
48    #[must_use]
49    pub fn builder() -> HttpClientBuilder {
50        HttpClientBuilder::default()
51    }
52
53    pub async fn send_json<T>(
54        &self,
55        base_url: &BaseUrl,
56        request: RequestParts,
57        authenticator: Option<&dyn Authenticator>,
58    ) -> Result<HttpResponse<T>, Error>
59    where
60        T: DeserializeOwned,
61    {
62        let response = self.send(base_url, &request, authenticator).await?;
63        let parsed = serde_json::from_str(&response.body).map_err(|error| {
64            let meta = ErrorMeta::from_response_meta(response.meta.clone(), response.body.clone());
65            let error = Error::Deserialize {
66                message: error.to_string(),
67                meta: Some(meta.clone()),
68            };
69            self.observer.on_error(&ErrorEvent { meta: Some(meta) });
70            error
71        })?;
72
73        self.observer.on_response(&ResponseEvent {
74            meta: response.meta.clone(),
75        });
76        Ok(HttpResponse::new(parsed, response.meta))
77    }
78
79    pub async fn send_json_expected<T>(
80        &self,
81        base_url: &BaseUrl,
82        request: RequestParts,
83        expected_status: StatusCode,
84        authenticator: Option<&dyn Authenticator>,
85    ) -> Result<HttpResponse<T>, Error>
86    where
87        T: DeserializeOwned,
88    {
89        let response = self.send(base_url, &request, authenticator).await?;
90        let response = self.require_status(response, expected_status)?;
91        let parsed = serde_json::from_str(&response.body).map_err(|error| {
92            let meta = ErrorMeta::from_response_meta(response.meta.clone(), response.body.clone());
93            let error = Error::Deserialize {
94                message: error.to_string(),
95                meta: Some(meta.clone()),
96            };
97            self.observer.on_error(&ErrorEvent { meta: Some(meta) });
98            error
99        })?;
100
101        self.observer.on_response(&ResponseEvent {
102            meta: response.meta.clone(),
103        });
104        Ok(HttpResponse::new(parsed, response.meta))
105    }
106
107    pub async fn send_json_or_empty_expected<T>(
108        &self,
109        base_url: &BaseUrl,
110        request: RequestParts,
111        expected_status: StatusCode,
112        authenticator: Option<&dyn Authenticator>,
113    ) -> Result<HttpResponse<Option<T>>, Error>
114    where
115        T: DeserializeOwned,
116    {
117        let response = self.send(base_url, &request, authenticator).await?;
118        let response = self.require_status(response, expected_status)?;
119        let parsed = if response.body.is_empty() {
120            None
121        } else {
122            Some(serde_json::from_str(&response.body).map_err(|error| {
123                let meta =
124                    ErrorMeta::from_response_meta(response.meta.clone(), response.body.clone());
125                let error = Error::Deserialize {
126                    message: error.to_string(),
127                    meta: Some(meta.clone()),
128                };
129                self.observer.on_error(&ErrorEvent { meta: Some(meta) });
130                error
131            })?)
132        };
133
134        self.observer.on_response(&ResponseEvent {
135            meta: response.meta.clone(),
136        });
137        Ok(HttpResponse::new(parsed, response.meta))
138    }
139
140    pub async fn send_text(
141        &self,
142        base_url: &BaseUrl,
143        request: RequestParts,
144        authenticator: Option<&dyn Authenticator>,
145    ) -> Result<HttpResponse<String>, Error> {
146        let response = self.send(base_url, &request, authenticator).await?;
147        self.observer.on_response(&ResponseEvent {
148            meta: response.meta.clone(),
149        });
150        Ok(HttpResponse::new(response.body, response.meta))
151    }
152
153    pub async fn send_no_content(
154        &self,
155        base_url: &BaseUrl,
156        request: RequestParts,
157        authenticator: Option<&dyn Authenticator>,
158    ) -> Result<HttpResponse<NoContent>, Error> {
159        self.send_empty_expected(base_url, request, StatusCode::NO_CONTENT, authenticator)
160            .await
161    }
162
163    pub async fn send_empty_expected(
164        &self,
165        base_url: &BaseUrl,
166        request: RequestParts,
167        expected_status: StatusCode,
168        authenticator: Option<&dyn Authenticator>,
169    ) -> Result<HttpResponse<NoContent>, Error> {
170        let response = self.send(base_url, &request, authenticator).await?;
171        let response = self.require_status(response, expected_status)?;
172        if !response.body.is_empty() {
173            let meta = ErrorMeta::from_response_meta(response.meta, response.body);
174            let error = Error::Deserialize {
175                message: format!(
176                    "expected an empty response body for HTTP {}",
177                    expected_status.as_u16()
178                ),
179                meta: Some(meta.clone()),
180            };
181            self.observer.on_error(&ErrorEvent {
182                meta: Some(meta.clone()),
183            });
184            return Err(error);
185        }
186
187        self.observer.on_response(&ResponseEvent {
188            meta: response.meta.clone(),
189        });
190        Ok(HttpResponse::new(NoContent, response.meta))
191    }
192
193    fn require_status(
194        &self,
195        response: ResponseParts,
196        expected_status: StatusCode,
197    ) -> Result<ResponseParts, Error> {
198        if response.meta.status() == expected_status.as_u16() {
199            return Ok(response);
200        }
201
202        let meta = ErrorMeta::from_response_meta(response.meta, response.body);
203        let error = Error::HttpStatus(meta.clone());
204        self.observer.on_error(&ErrorEvent { meta: Some(meta) });
205        Err(error)
206    }
207
208    async fn send(
209        &self,
210        base_url: &BaseUrl,
211        request: &RequestParts,
212        authenticator: Option<&dyn Authenticator>,
213    ) -> Result<ResponseParts, Error> {
214        let _permit = self.concurrency_limit.acquire().await?;
215        let url = base_url.join_path(request.path());
216        let mut attempt = 0;
217        let started_at = Instant::now();
218
219        loop {
220            let observed_url = url_with_query(&url, request.query());
221            self.observer.on_request_start(&RequestStart {
222                operation: request.operation().map(ToOwned::to_owned),
223                method: request.method(),
224                url: observed_url,
225            });
226
227            let request_builder = self.build_request(&url, request, authenticator)?;
228            let response = match request_builder.send().await {
229                Ok(response) => response,
230                Err(error) => {
231                    match self.retry_config.classify_transport_error(
232                        &request.method(),
233                        attempt,
234                        started_at.elapsed(),
235                    ) {
236                        RetryDecision::RetryAfter(wait) => {
237                            self.observer.on_retry(&RetryEvent {
238                                operation: request.operation().map(ToOwned::to_owned),
239                                method: request.method(),
240                                url: url.clone(),
241                                attempt: attempt + 1,
242                                status: None,
243                                wait,
244                            });
245                            tokio::time::sleep(wait).await;
246                            attempt += 1;
247                            continue;
248                        }
249                        RetryDecision::DoNotRetry => {
250                            let error = Error::from_reqwest(error, None);
251                            self.observer.on_error(&ErrorEvent { meta: None });
252                            return Err(error);
253                        }
254                    }
255                }
256            };
257
258            let status = response.status();
259            let headers = response.headers().clone();
260            let meta = ResponseMeta::from_response_parts(
261                request.operation().map(ToOwned::to_owned),
262                url.clone(),
263                status,
264                &headers,
265                &self.request_id_header_name,
266                attempt + 1,
267                started_at.elapsed(),
268            );
269            let body = match response.text().await {
270                Ok(body) => body,
271                Err(error) => {
272                    match self.retry_config.classify_transport_error(
273                        &request.method(),
274                        attempt,
275                        started_at.elapsed(),
276                    ) {
277                        RetryDecision::RetryAfter(wait) => {
278                            self.observer.on_retry(&RetryEvent {
279                                operation: request.operation().map(ToOwned::to_owned),
280                                method: request.method(),
281                                url: url.clone(),
282                                attempt: attempt + 1,
283                                status: Some(status),
284                                wait,
285                            });
286                            tokio::time::sleep(wait).await;
287                            attempt += 1;
288                            continue;
289                        }
290                        RetryDecision::DoNotRetry => {
291                            let error_meta =
292                                ErrorMeta::from_response_meta(meta.clone(), String::new());
293                            let error = Error::from_reqwest(error, Some(error_meta.clone()));
294                            self.observer.on_error(&ErrorEvent {
295                                meta: Some(error_meta),
296                            });
297                            return Err(error);
298                        }
299                    }
300                }
301            };
302
303            match self.retry_config.classify_response(
304                &request.method(),
305                status,
306                attempt,
307                meta.retry_after(),
308                started_at.elapsed(),
309            ) {
310                RetryDecision::RetryAfter(wait) => {
311                    self.observer.on_retry(&RetryEvent {
312                        operation: request.operation().map(ToOwned::to_owned),
313                        method: request.method(),
314                        url: url.clone(),
315                        attempt: attempt + 1,
316                        status: Some(status),
317                        wait,
318                    });
319                    tokio::time::sleep(wait).await;
320                    attempt += 1;
321                    continue;
322                }
323                RetryDecision::DoNotRetry => {}
324            }
325
326            if status.is_success() {
327                return Ok(ResponseParts { meta, body });
328            }
329
330            let error_meta = ErrorMeta::from_response_meta(meta, body);
331            let error = if status == reqwest::StatusCode::TOO_MANY_REQUESTS {
332                Error::RateLimited(error_meta.clone())
333            } else {
334                Error::HttpStatus(error_meta.clone())
335            };
336            self.observer.on_error(&ErrorEvent {
337                meta: Some(error_meta),
338            });
339            return Err(error);
340        }
341    }
342
343    fn build_request(
344        &self,
345        url: &str,
346        request: &RequestParts,
347        authenticator: Option<&dyn Authenticator>,
348    ) -> Result<reqwest::RequestBuilder, Error> {
349        let mut headers = self.default_headers.clone();
350        headers.extend(request.headers().clone());
351        if let Some(authenticator) = authenticator {
352            authenticator.apply(&mut headers)?;
353        }
354
355        let mut builder = self
356            .client
357            .request(request.method(), url)
358            .headers(headers)
359            .query(request.query());
360
361        builder = match request.body() {
362            RequestBody::Empty => builder,
363            RequestBody::Json(value) => builder.json(value),
364            RequestBody::Text(value) => builder.body(value.clone()),
365            RequestBody::Bytes(value) => builder.body(value.clone()),
366        };
367
368        if matches!(request.body(), RequestBody::Text(_))
369            && !request.headers().contains_key(CONTENT_TYPE)
370            && !self.default_headers.contains_key(CONTENT_TYPE)
371        {
372            builder = builder.header(CONTENT_TYPE, HeaderValue::from_static("text/plain"));
373        }
374
375        Ok(builder)
376    }
377}
378
379fn url_with_query(url: &str, query: &[(String, String)]) -> String {
380    if query.is_empty() {
381        return url.to_owned();
382    }
383    let Ok(mut parsed) = reqwest::Url::parse(url) else {
384        return url.to_owned();
385    };
386    parsed
387        .query_pairs_mut()
388        .extend_pairs(query.iter().map(|(key, value)| (key, value)));
389    parsed.into()
390}
391
392impl Default for HttpClientBuilder {
393    fn default() -> Self {
394        Self {
395            reqwest_client: None,
396            timeout: Duration::from_secs(30),
397            default_headers: HeaderMap::new(),
398            request_id_header_name: HeaderName::from_static("x-request-id"),
399            retry_config: RetryConfig::default(),
400            observer: Arc::new(NoopObserver),
401            concurrency_limit: ConcurrencyLimit::default(),
402        }
403    }
404}
405
406impl HttpClientBuilder {
407    #[must_use]
408    pub fn timeout(mut self, timeout: Duration) -> Self {
409        self.timeout = timeout;
410        self
411    }
412
413    #[must_use]
414    pub fn reqwest_client(mut self, client: reqwest::Client) -> Self {
415        self.reqwest_client = Some(client);
416        self
417    }
418
419    pub fn default_header(mut self, name: &str, value: &str) -> Result<Self, Error> {
420        let name = HeaderName::from_bytes(name.as_bytes()).map_err(|error| {
421            Error::InvalidRequest(format!("invalid default header name: {error}"))
422        })?;
423        let value = HeaderValue::from_str(value).map_err(|error| {
424            Error::InvalidRequest(format!("invalid default header value: {error}"))
425        })?;
426        self.default_headers.insert(name, value);
427        Ok(self)
428    }
429
430    pub fn request_id_header_name(mut self, name: &str) -> Result<Self, Error> {
431        self.request_id_header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|error| {
432            Error::InvalidRequest(format!("invalid request id header name: {error}"))
433        })?;
434        Ok(self)
435    }
436
437    #[must_use]
438    pub fn retry_config(mut self, retry_config: RetryConfig) -> Self {
439        self.retry_config = retry_config;
440        self
441    }
442
443    #[must_use]
444    pub fn observer(mut self, observer: Arc<dyn TransportObserver>) -> Self {
445        self.observer = observer;
446        self
447    }
448
449    #[must_use]
450    pub fn concurrency_limit(mut self, concurrency_limit: ConcurrencyLimit) -> Self {
451        self.concurrency_limit = concurrency_limit;
452        self
453    }
454
455    pub fn build(self) -> Result<HttpClient, Error> {
456        let client = match self.reqwest_client {
457            Some(client) => client,
458            None => reqwest::Client::builder()
459                .timeout(self.timeout)
460                .build()
461                .map_err(|error| Error::from_reqwest(error, None))?,
462        };
463
464        Ok(HttpClient {
465            client,
466            default_headers: self.default_headers,
467            request_id_header_name: self.request_id_header_name,
468            retry_config: self.retry_config,
469            observer: self.observer,
470            concurrency_limit: self.concurrency_limit,
471        })
472    }
473}