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}