use crate::Error as CrateError;
use crate::adaptive_concurrency::http::HttpError as GenericHttpError;
use futures::future::BoxFuture;
use bytes::Bytes;
use http::{Request as HttpRequest, StatusCode};
use reqwest;
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
use tower::Service;
#[derive(Clone)]
pub struct ReqwestService {
client: reqwest::Client,
}
impl ReqwestService {
pub fn new() -> Self {
Self {
client: reqwest::Client::new(),
}
}
pub fn new_with_client(client: reqwest::Client) -> Self {
Self { client }
}
}
impl Default for ReqwestService {
fn default() -> Self {
Self::new()
}
}
impl Service<HttpRequest<Option<Bytes>>> for ReqwestService {
type Response = reqwest::Response;
type Error = GenericHttpError;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, http_request: HttpRequest<Option<Bytes>>) -> Self::Future {
let (parts, body_option) = http_request.into_parts();
let url_str = parts.uri.to_string();
let url = match reqwest::Url::parse(&url_str) {
Ok(u) => u,
Err(parse_err) => {
let ge = GenericHttpError::InvalidRequest {
details: format!("Invalid URL '{}': {}", url_str, parse_err),
};
return Box::pin(async move { Err(ge) });
}
};
let mut request_builder = self.client.request(parts.method, url);
for (header_name, header_value) in parts.headers.iter() {
request_builder = request_builder.header(header_name, header_value);
}
if let Some(bytes_body) = body_option {
let reqwest_body = reqwest::Body::from(bytes_body);
request_builder = request_builder.body(reqwest_body);
}
let request_future = request_builder.send();
Box::pin(async move {
match request_future.await {
Ok(response) => {
let status = response.status();
if status.is_success() {
Ok(response)
} else {
let error_body = response
.text()
.await
.unwrap_or_else(|_| "Could not read error body".to_string());
if status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS {
warn!(
status = %status,
error_body = %error_body,
"Server error or rate limited"
);
} else if status.is_client_error() {
error!(
status = %status,
error_body = %error_body,
"Client error"
);
}
Err(GenericHttpError::ServerError {
status: status.as_u16(),
body: error_body,
})
}
}
Err(e) => {
if e.is_timeout() {
warn!(error = %e, "Request timed out");
Err(GenericHttpError::Timeout)
} else if e.is_connect() {
error!(error = %e, "Connection error");
Err(GenericHttpError::Transport {
source: Box::new(e),
})
} else {
error!(error = %e, "Other reqwest error");
Err(GenericHttpError::ClientError {
source: Box::new(e),
})
}
}
}
})
}
}