use crate::{Client, ClientError, PayloadStream, Result};
use bytes::Bytes;
use futures::StreamExt as _;
use loonfs_api::{ApiError, ErrorCode};
use reqwest::Method;
use std::time::Duration;
pub(crate) const MAX_TRANSIENT_ATTEMPTS: u32 = 4;
pub(crate) const INITIAL_TRANSIENT_RETRY_DELAY: Duration = Duration::from_millis(250);
pub(crate) const MAX_TRANSIENT_RETRY_DELAY: Duration = Duration::from_secs(2);
pub(crate) const IO_INACTIVITY_TIMEOUT: Duration = Duration::from_secs(60);
pub(crate) struct WireRequest {
method: Method,
url: String,
headers: Vec<(String, String)>,
authenticate: bool,
}
impl Client {
pub(crate) fn get(&self, url: &str) -> WireRequest {
WireRequest::to_server(Method::GET, url)
}
pub(crate) fn post(&self, url: &str) -> WireRequest {
WireRequest::to_server(Method::POST, url)
}
pub(crate) fn put(&self, url: &str) -> WireRequest {
WireRequest::to_server(Method::PUT, url)
}
pub(crate) fn delete(&self, url: &str) -> WireRequest {
WireRequest::to_server(Method::DELETE, url)
}
}
impl WireRequest {
fn to_server(method: Method, url: &str) -> Self {
Self {
method,
url: url.to_owned(),
headers: Vec::new(),
authenticate: true,
}
}
pub(crate) fn presigned(method: Method, url: &str) -> Self {
Self {
method,
url: url.to_owned(),
headers: Vec::new(),
authenticate: false,
}
}
pub(crate) fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.headers.push((name.into(), value.into()));
self
}
}
impl Client {
pub(crate) async fn request_json<Req, Resp>(
&self,
request: WireRequest,
body: Option<&Req>,
) -> Result<Resp>
where
Req: serde::Serialize,
Resp: serde::de::DeserializeOwned,
{
self.request_json_inner(request, body, true).await
}
pub(crate) async fn request_json_once<Req, Resp>(
&self,
request: WireRequest,
body: Option<&Req>,
) -> Result<Resp>
where
Req: serde::Serialize,
Resp: serde::de::DeserializeOwned,
{
self.request_json_inner(request, body, false).await
}
async fn request_json_inner<Req, Resp>(
&self,
request: WireRequest,
body: Option<&Req>,
retry: bool,
) -> Result<Resp>
where
Req: serde::Serialize,
Resp: serde::de::DeserializeOwned,
{
let body = match body {
Some(body) => {
Some(Bytes::from(
serde_json::to_vec(body).map_err(|err| ClientError::Json(err.to_string()))?,
))
}
None => None,
};
let request = match body {
Some(_) => request.header("content-type", "application/json"),
None => request,
};
let bytes = if retry {
self.call_with_transient_retry(&request, body.as_ref())
.await?
} else {
self.call_once(&request, body.as_ref()).await?
};
serde_json::from_slice(&bytes).map_err(|err| ClientError::Json(err.to_string()))
}
pub(crate) async fn request_bytes(&self, url: &str) -> Result<Vec<u8>> {
let request = self.get(url);
self.call_with_transient_retry(&request, None).await
}
pub(crate) async fn call_once(
&self,
request: &WireRequest,
body: Option<&Bytes>,
) -> Result<Vec<u8>> {
self.send(request, body)
.await
.map(|response| response.bytes)
.map_err(|attempt| attempt.error)
}
pub(crate) async fn call_streamed_once(
&self,
request: &WireRequest,
body: PayloadStream,
size_bytes: Option<u64>,
) -> Result<Vec<u8>> {
self.send_streamed(request, body, size_bytes)
.await
.map(|response| response.bytes)
.map_err(|attempt| attempt.error)
}
pub(crate) async fn call_for_response_stream(
&self,
request: &WireRequest,
) -> Result<PayloadStream> {
#[cfg(test)]
if let Some(outcome) = test_transport::next(request) {
return match outcome {
Ok(response) => {
Ok(
futures::stream::once(async move { Ok(Bytes::from(response.bytes)) })
.boxed(),
)
}
Err(attempt) => Err(attempt.error),
};
}
let response = self
.build(request)
.send()
.await
.map_err(|err| ClientError::Http(describe_send_error(&request.url, &err)))?;
let status = response.status();
if !status.is_success() {
let bytes = response
.bytes()
.await
.map_err(|err| ClientError::Http(describe_send_error(&request.url, &err)))?;
return Err(map_status_error(status.as_u16(), &bytes));
}
Ok(response
.bytes_stream()
.map(|chunk| {
chunk.map_err(|err| {
std::io::Error::other(format!("response body ended early: {err}"))
})
})
.boxed())
}
pub(crate) async fn call_with_transient_retry(
&self,
request: &WireRequest,
body: Option<&Bytes>,
) -> Result<Vec<u8>> {
self.call_with_transient_retry_headers(request, body)
.await
.map(|response| response.bytes)
}
pub(crate) async fn call_with_transient_retry_headers(
&self,
request: &WireRequest,
body: Option<&Bytes>,
) -> Result<WireResponse> {
let mut attempts = 0;
loop {
let attempt = match self.send(request, body).await {
Ok(response) => return Ok(response),
Err(attempt) => attempt,
};
attempts += 1;
let transient = transient_failure(attempt.transport, &attempt.error);
if !self.transient_retry || !transient || attempts >= MAX_TRANSIENT_ATTEMPTS {
return Err(attempt.error);
}
transient_retry_pause(transient_retry_backoff(attempts)).await;
}
}
async fn send(
&self,
request: &WireRequest,
body: Option<&Bytes>,
) -> std::result::Result<WireResponse, FailedAttempt> {
#[cfg(test)]
if let Some(outcome) = test_transport::next(request) {
return outcome;
}
let mut builder = self.build(request);
if let Some(bytes) = body {
builder = builder.body(bytes.clone());
}
self.dispatch(request, builder).await
}
async fn send_streamed(
&self,
request: &WireRequest,
body: PayloadStream,
size_bytes: Option<u64>,
) -> std::result::Result<WireResponse, FailedAttempt> {
#[cfg(test)]
if let Some(outcome) = test_transport::next(request) {
let mut body = body;
while futures::StreamExt::next(&mut body).await.is_some() {}
return outcome;
}
let mut builder = self.build(request).body(reqwest::Body::wrap_stream(body));
if let Some(size_bytes) = size_bytes {
builder = builder.header(http::header::CONTENT_LENGTH, size_bytes);
}
self.dispatch(request, builder).await
}
fn build(&self, request: &WireRequest) -> reqwest::RequestBuilder {
let mut builder = self.http.request(request.method.clone(), &request.url);
if request.authenticate {
if let Some(token) = &self.auth_token {
builder = builder.bearer_auth(token);
}
}
for (name, value) in &request.headers {
builder = builder.header(name.as_str(), value.as_str());
}
builder
}
async fn dispatch(
&self,
request: &WireRequest,
builder: reqwest::RequestBuilder,
) -> std::result::Result<WireResponse, FailedAttempt> {
let response = builder.send().await.map_err(|err| FailedAttempt {
transport: true,
error: ClientError::Http(describe_send_error(&request.url, &err)),
})?;
let status = response.status();
let headers = response.headers().clone();
let bytes = response.bytes().await.map_err(|err| FailedAttempt {
transport: true,
error: ClientError::Http(describe_send_error(&request.url, &err)),
})?;
if status.is_success() {
return Ok(WireResponse {
headers,
bytes: bytes.to_vec(),
});
}
Err(FailedAttempt {
transport: false,
error: map_status_error(status.as_u16(), &bytes),
})
}
}
pub(crate) struct WireResponse {
headers: reqwest::header::HeaderMap,
pub(crate) bytes: Vec<u8>,
}
impl WireResponse {
pub(crate) fn get(
&self,
name: reqwest::header::HeaderName,
) -> Option<&reqwest::header::HeaderValue> {
self.headers.get(name)
}
}
pub(crate) struct FailedAttempt {
pub(crate) transport: bool,
pub(crate) error: ClientError,
}
pub(crate) fn map_status_error(status: u16, body: &[u8]) -> ClientError {
match serde_json::from_slice::<ApiError>(body) {
Ok(body) => ClientError::Api {
status,
code: body.code,
feature: body.feature,
message: body.message,
request_id: body.request_id,
details: body.details,
},
Err(err) => ClientError::Http(format!(
"http status {status} with a non-envelope body: {err}"
)),
}
}
fn describe_send_error(url: &str, error: &reqwest::Error) -> String {
render_send_error(url, error, error.is_connect(), error.is_timeout())
}
fn render_send_error(
url: &str,
error: &(dyn std::error::Error + 'static),
connect_failure: bool,
timed_out: bool,
) -> String {
let mut detail = error.to_string();
let mut source = error.source();
while let Some(cause) = source {
let rendered = cause.to_string();
if !detail.contains(&rendered) {
detail.push_str(": ");
detail.push_str(&rendered);
}
source = cause.source();
}
if connect_failure {
format!(
"cannot connect to `{url}`: {detail}; check that the server is running and that the \
profile's `server_url` points at it"
)
} else if timed_out {
format!("request to `{url}` timed out: {detail}")
} else {
format!("request to `{url}` failed: {detail}")
}
}
pub(crate) fn transient_failure(transport: bool, error: &ClientError) -> bool {
transport
|| matches!(
error,
ClientError::Api { code, .. }
if code == ErrorCode::ServerBusy.as_str()
|| code == ErrorCode::CommitQueueFull.as_str()
|| code == ErrorCode::ShuttingDown.as_str()
)
}
fn transient_retry_backoff(attempt: u32) -> Duration {
let doublings = attempt.saturating_sub(1).min(16);
INITIAL_TRANSIENT_RETRY_DELAY
.saturating_mul(1u32 << doublings)
.min(MAX_TRANSIENT_RETRY_DELAY)
}
#[allow(clippy::disallowed_methods)]
async fn transient_retry_pause(backoff: Duration) {
tokio::time::sleep(backoff).await;
}
#[cfg(test)]
pub(crate) mod test_transport {
use super::{FailedAttempt, WireRequest, WireResponse};
use crate::ClientError;
use std::cell::RefCell;
use std::collections::VecDeque;
pub(crate) enum Outcome {
TransportFailure,
Success(Vec<u8>),
PartAccepted(String),
}
struct State {
outcomes: VecDeque<Outcome>,
attempts: usize,
sent: Vec<SentRequest>,
}
#[derive(Debug, Clone)]
pub(crate) struct SentRequest {
#[allow(dead_code, reason = "part of what a scripted attempt carried")]
pub url: String,
pub headers: Vec<(String, String)>,
}
impl SentRequest {
pub(crate) fn header(&self, name: &str) -> Option<&str> {
self.headers
.iter()
.find(|(header, _)| header.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
}
thread_local! {
static STATE: RefCell<Option<State>> = const { RefCell::new(None) };
}
pub(crate) struct Guard;
impl Guard {
pub(crate) fn attempts(&self) -> usize {
STATE.with(|state| {
state
.borrow()
.as_ref()
.expect("test transport should be installed")
.attempts
})
}
pub(crate) fn sent(&self) -> Vec<SentRequest> {
STATE.with(|state| {
state
.borrow()
.as_ref()
.expect("test transport should be installed")
.sent
.clone()
})
}
}
impl Drop for Guard {
fn drop(&mut self) {
STATE.with(|state| *state.borrow_mut() = None);
}
}
pub(crate) fn failures(count: usize) -> Guard {
install(std::iter::repeat_with(|| Outcome::TransportFailure).take(count))
}
pub(crate) fn failure_then_success(body: Vec<u8>) -> Guard {
install([Outcome::TransportFailure, Outcome::Success(body)])
}
pub(crate) fn script(outcomes: impl IntoIterator<Item = Outcome>) -> Guard {
install(outcomes)
}
fn install(outcomes: impl IntoIterator<Item = Outcome>) -> Guard {
STATE.with(|state| {
let replaced = state.borrow_mut().replace(State {
outcomes: outcomes.into_iter().collect(),
attempts: 0,
sent: Vec::new(),
});
assert!(replaced.is_none(), "test transport already installed");
});
Guard
}
pub(super) fn next(request: &WireRequest) -> Option<Result<WireResponse, FailedAttempt>> {
STATE.with(|state| {
let mut state = state.borrow_mut();
let state = state.as_mut()?;
state.attempts += 1;
state.sent.push(SentRequest {
url: request.url.clone(),
headers: request.headers.clone(),
});
let outcome = state
.outcomes
.pop_front()
.expect("test transport exhausted before client stopped sending");
Some(match outcome {
Outcome::TransportFailure => Err(FailedAttempt {
transport: true,
error: ClientError::Http("injected transport failure".to_owned()),
}),
Outcome::Success(bytes) => Ok(WireResponse {
headers: reqwest::header::HeaderMap::new(),
bytes,
}),
Outcome::PartAccepted(etag) => {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::ETAG,
etag.parse().expect("etag is a valid header value"),
);
Ok(WireResponse {
headers,
bytes: Vec::new(),
})
}
})
})
}
}
#[cfg(test)]
mod tests {
use super::render_send_error;
#[derive(Debug)]
struct Layered {
message: &'static str,
cause: Option<Box<Layered>>,
}
impl std::fmt::Display for Layered {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.message)
}
}
impl std::error::Error for Layered {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
self.cause
.as_deref()
.map(|cause| cause as &(dyn std::error::Error + 'static))
}
}
#[test]
fn send_errors_surface_the_root_cause_and_the_url() {
let error = Layered {
message: "error sending request",
cause: Some(Box::new(Layered {
message: "client error (Connect)",
cause: Some(Box::new(Layered {
message: "tcp connect error: Connection refused (os error 61)",
cause: None,
})),
})),
};
let connect = render_send_error("http://127.0.0.1:9/v0/namespaces", &error, true, false);
assert!(
connect.contains("cannot connect to `http://127.0.0.1:9/v0/namespaces`"),
"{connect}"
);
assert!(connect.contains("Connection refused"), "{connect}");
assert!(connect.contains("`server_url`"), "{connect}");
let timeout = render_send_error("http://h/v0", &error, false, true);
assert!(timeout.contains("timed out"), "{timeout}");
let repeated = Layered {
message: "outer: inner detail",
cause: Some(Box::new(Layered {
message: "inner detail",
cause: None,
})),
};
let rendered = render_send_error("http://h/v0", &repeated, false, false);
assert_eq!(
rendered,
"request to `http://h/v0` failed: outer: inner detail"
);
}
}