use std::{borrow::Cow, sync::Arc, time::Duration};
use anyhow::Context as _;
use bytes::Bytes;
use reqwest::header::{self, HeaderMap, HeaderValue};
use thiserror::Error;
use tracing::Instrument;
use crate::{
error::CrpcError,
token_source::{TokenSource, TokenSourceError},
};
#[derive(Debug, Error)]
pub enum CrpcClientError {
#[error("connection error {context}: {source:#?}")]
ConnectionError {
context: Cow<'static, str>,
source: Box<dyn std::error::Error + Send + Sync + 'static>,
},
#[error("server returned an error: {0:#?}")]
CrpcError(CrpcError),
#[error("failed to decode response body: {context}: {source:#?}")]
DecodeError {
context: Cow<'static, str>,
source: Option<Box<dyn std::error::Error + Send + Sync + 'static>>,
body: Option<Bytes>,
},
#[error("failed to retrieve token: {0}")]
TokenSourceError(#[from] TokenSourceError),
}
const APPLICATION_PROTO: &str = "application/proto";
pub struct CrpcClient {
http_client: reqwest::Client,
base_url: url::Url,
token_source: Option<Arc<dyn TokenSource>>,
user_agent: HeaderValue,
}
impl CrpcClient {
pub fn new(base_url: &url::Url) -> anyhow::Result<Self> {
let http_client = reqwest::ClientBuilder::new()
.timeout(Duration::from_secs(30))
.build()
.context("error creating HTTP client")?;
Self::new_with_client(base_url, http_client)
}
pub fn new_with_client(
base_url: &url::Url,
http_client: reqwest::Client,
) -> anyhow::Result<Self> {
let user_agent =
HeaderValue::from_str(&format!("reqwest-crpc {}", env!("CARGO_PKG_VERSION")))
.context("error creating user agent header")?;
Ok(CrpcClient {
http_client,
base_url: base_url.clone(),
token_source: None,
user_agent,
})
}
pub fn use_token_source(&mut self, token_source: Arc<dyn TokenSource>) -> &mut Self {
self.token_source = Some(token_source);
self
}
pub fn use_user_agent(&mut self, user_agent: &str) -> anyhow::Result<&mut Self> {
self.user_agent = HeaderValue::from_str(user_agent)
.with_context(|| format!("error creating user agent header from {user_agent}"))?;
Ok(self)
}
pub async fn unary_request<Req, Res>(
&self,
path: &str,
req: &Req,
) -> Result<Res, CrpcClientError>
where
Req: prost::Message,
Res: prost::Message + Default,
{
self.do_unary_request(path, req)
.instrument(tracing::info_span!("request", %path, id = rand::random::<u16>()))
.await
}
async fn do_unary_request<Req, Res>(
&self,
path: &str,
req: &Req,
) -> Result<Res, CrpcClientError>
where
Req: prost::Message,
Res: prost::Message + Default,
{
let url = self.base_url.join(path).map_err(|e| {
CrpcClientError::ConnectionError {
context: "error joining base URL and path".into(),
source: e.into(),
}
})?;
let mut headers = HeaderMap::with_capacity(3);
headers.insert(
header::CONTENT_TYPE,
header::HeaderValue::from_static(APPLICATION_PROTO),
);
headers.insert(header::USER_AGENT, self.user_agent.clone());
tracing::trace!(?url, ?headers, "Sending crpc unary request");
if let Some(token_source) = &self.token_source {
let token = token_source.get_token().await?;
let token_header = header::HeaderValue::from_str(&token_source.format_header(token))
.map_err(|e| {
CrpcClientError::TokenSourceError(
format!("error formatting token as header value: {e:?}").into(),
)
})?;
headers.insert(header::AUTHORIZATION, token_header);
}
let body = req.encode_to_vec();
let response = self
.http_client
.post(url)
.body(reqwest::Body::from(body))
.headers(headers)
.send()
.await
.map_err(|e| {
CrpcClientError::ConnectionError {
context: "error sending request".into(),
source: e.into(),
}
})?;
tracing::trace!(status=%response.status(), body_len=%response.content_length().unwrap_or(0), "Received crpc unary response");
let status = response.status();
if !status.is_success() {
let response_raw = response
.text()
.await
.unwrap_or_else(|_| "<failed to read body>".to_string());
match serde_json::from_str::<CrpcError>(&response_raw) {
Ok(crpc_err) => {
return Err(CrpcClientError::CrpcError(crpc_err));
}
Err(_) => {
return Err(CrpcClientError::CrpcError(CrpcError::new(
status.into(),
response_raw,
)));
}
}
}
let body = response.bytes().await.map_err(|e| {
CrpcClientError::DecodeError {
context: "error reading response body".into(),
source: Some(e.into()),
body: None,
}
})?;
Res::decode(&body[..]).map_err(|e| {
CrpcClientError::DecodeError {
context: "error decoding response body".into(),
source: Some(e.into()),
body: Some(body.clone()),
}
})
}
}