use serde_json::Value;
use snafu::ResultExt;
use url::Url;
use crate::cassettes;
use crate::error::{Error, Result, error};
use crate::path::{PathMode, call_url};
use crate::transport::{
Call, SpecFetch, SpecTransport, StreamingTransport, TapesTransport, TransportError,
WireRequest, WireResponse,
};
const MAX_ATTEMPTS: u32 = 4;
const UNAUTHORIZED: u16 = 401;
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub struct Rejected<'a> {
pub status: u16,
pub endpoint: &'a str,
pub attempt: u32,
}
#[derive(Debug)]
#[non_exhaustive]
pub enum Unauthorized {
Retry,
Surface,
Fail(TransportError),
}
pub trait HttpAuth {
fn authorize(
&self,
request: &WireRequest<'_>,
attempt: u32,
) -> impl Future<Output = std::result::Result<Vec<(String, String)>, TransportError>>;
fn on_unauthorized(&self, rejected: Rejected<'_>) -> impl Future<Output = Unauthorized> {
let _ = rejected;
async { Unauthorized::Surface }
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct NoAuth;
impl HttpAuth for NoAuth {
async fn authorize(
&self,
_request: &WireRequest<'_>,
_attempt: u32,
) -> std::result::Result<Vec<(String, String)>, TransportError> {
Ok(Vec::new())
}
}
#[derive(Clone)]
pub struct HttpEngine<A> {
http: Option<reqwest::Client>,
base: Url,
mode: PathMode,
auth: A,
}
impl<A> std::fmt::Debug for HttpEngine<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpEngine")
.field("base", &self.base)
.field("mode", &self.mode)
.finish_non_exhaustive()
}
}
pub type DirectHttp = HttpEngine<NoAuth>;
impl HttpEngine<NoAuth> {
#[must_use]
pub fn new(base: Url) -> Self {
Self::with_auth(base, NoAuth)
}
}
impl<A> HttpEngine<A> {
#[must_use]
pub fn with_auth(base: Url, auth: A) -> Self {
let http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.ok();
Self {
http,
base,
mode: PathMode::Direct,
auth,
}
}
#[must_use]
pub fn under_base(mut self) -> Self {
self.mode = PathMode::UnderBase;
self
}
#[must_use]
pub fn with_client(mut self, http: reqwest::Client) -> Self {
self.http = Some(http);
self
}
#[must_use]
pub fn base(&self) -> &Url {
&self.base
}
#[must_use]
pub fn auth(&self) -> &A {
&self.auth
}
fn http(&self) -> Result<&reqwest::Client> {
self.http
.as_ref()
.ok_or_else(|| error::ClientInitSnafu.build())
}
fn refuse_moved(
&self,
response: &reqwest::Response,
) -> std::result::Result<(), TransportError> {
if response.status().is_redirection()
&& response.status() != reqwest::StatusCode::NOT_MODIFIED
{
return Err(TransportError::new(
"the server answered with a redirect; this client does not follow them",
));
}
if response.url().origin() != self.base.origin() {
return Err(TransportError::new(
"the response came from a different origin than the configured server",
));
}
Ok(())
}
}
impl<A: HttpAuth> HttpEngine<A> {
async fn attempt(
&self,
call: &WireRequest<'_>,
url: &Url,
attempt: u32,
) -> std::result::Result<reqwest::Response, TransportError> {
let method = reqwest::Method::from_bytes(call.method.as_bytes())
.map_err(|_| TransportError::new(format!("unusable HTTP method {:?}", call.method)))?;
let client = self
.http()
.map_err(|error| TransportError::new(error.to_string()))?;
let mut request = client.request(method, url.clone());
for (name, value) in &call.headers {
request = request.header(name, value);
}
for (name, value) in self.auth.authorize(call, attempt).await? {
request = request.header(name, value);
}
if let Some(body) = &call.body {
request = request
.header(http::header::CONTENT_TYPE, "application/json")
.body(body.clone());
}
let response = request.send().await.map_err(|source| {
TransportError::with_source("could not reach the tapes API", source)
})?;
self.refuse_moved(&response)?;
Ok(response)
}
async fn send_call(
&self,
call: &WireRequest<'_>,
) -> std::result::Result<(reqwest::Response, Url), TransportError> {
let url = call_url(&self.base, call, self.mode)
.map_err(|error| TransportError::new(error.to_string()))?;
let endpoint = url.to_string();
for attempt in 1..=MAX_ATTEMPTS {
let response = self.attempt(call, &url, attempt).await?;
if response.status().as_u16() != UNAUTHORIZED {
return Ok((response, url));
}
let decision = self
.auth
.on_unauthorized(Rejected {
status: UNAUTHORIZED,
endpoint: &endpoint,
attempt,
})
.await;
match decision {
Unauthorized::Retry if attempt < MAX_ATTEMPTS => {
drop(response);
}
Unauthorized::Retry => {
return Err(TransportError::new(format!(
"the credential was refused {MAX_ATTEMPTS} times running",
)));
}
Unauthorized::Surface => return Ok((response, url)),
Unauthorized::Fail(error) => return Err(error),
}
}
Err(TransportError::new("no attempt was made"))
}
pub async fn fetch_discovery(&self) -> Result<Value> {
cassettes::fetch_discovery(self).await
}
pub async fn fetch_spec(&self, path: &str, etag: Option<&str>) -> Result<SpecFetch> {
cassettes::fetch_spec(self, path, etag).await
}
pub async fn execute(&self, call: &Call<'_>) -> Result<Value> {
cassettes::invoke(self, call).await
}
pub async fn execute_stream(&self, call: &Call<'_>) -> Result<reqwest::Response> {
let (response, url) = self.send_call(call).await.context(error::TransportSnafu)?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
return Err(Error::ApiStatus {
status: status.as_u16(),
endpoint: url.to_string(),
body,
});
}
Ok(response)
}
}
impl<A: HttpAuth> TapesTransport for HttpEngine<A> {
async fn send(
&self,
request: &WireRequest<'_>,
) -> std::result::Result<WireResponse, TransportError> {
let (response, url) = self.send_call(request).await?;
let status = response.status().as_u16();
let headers = response
.headers()
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_owned(), value.to_owned()))
})
.collect();
let body = response
.bytes()
.await
.map_err(|source| TransportError::with_source("could not read the response", source))?
.to_vec();
Ok(WireResponse::new(status, url.to_string(), headers, body))
}
}
impl<A: HttpAuth> StreamingTransport for HttpEngine<A> {
type Body = reqwest::Response;
async fn send_stream(&self, request: &WireRequest<'_>) -> Result<Self::Body> {
self.execute_stream(request).await
}
}
impl<A: HttpAuth> SpecTransport for HttpEngine<A> {
type Error = Error;
async fn fetch_discovery(&self) -> Result<Value> {
Self::fetch_discovery(self).await
}
async fn fetch_spec(&self, path: &str, etag: Option<&str>) -> Result<SpecFetch> {
Self::fetch_spec(self, path, etag).await
}
async fn execute(&self, call: &Call<'_>) -> Result<Value> {
Self::execute(self, call).await
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn base(raw: &str) -> DirectHttp {
DirectHttp::new(Url::parse(raw).unwrap())
}
struct Minting {
mints: Arc<AtomicU32>,
retries: u32,
}
impl HttpAuth for Minting {
async fn authorize(
&self,
_request: &WireRequest<'_>,
attempt: u32,
) -> std::result::Result<Vec<(String, String)>, TransportError> {
self.mints.fetch_add(1, Ordering::SeqCst);
Ok(vec![(
"x-tapes-auth".to_owned(),
format!("Bearer token-{attempt}"),
)])
}
async fn on_unauthorized(&self, rejected: Rejected<'_>) -> Unauthorized {
if rejected.attempt <= self.retries {
Unauthorized::Retry
} else {
Unauthorized::Fail(TransportError::new(
"not authenticated; run the login command",
))
}
}
}
struct Forever;
impl HttpAuth for Forever {
async fn authorize(
&self,
_request: &WireRequest<'_>,
_attempt: u32,
) -> std::result::Result<Vec<(String, String)>, TransportError> {
Ok(Vec::new())
}
async fn on_unauthorized(&self, _rejected: Rejected<'_>) -> Unauthorized {
Unauthorized::Retry
}
}
fn minting(server: &MockServer, retries: u32) -> (HttpEngine<Minting>, Arc<AtomicU32>) {
let mints = Arc::new(AtomicU32::new(0));
let engine = HttpEngine::with_auth(
Url::parse(&server.uri()).unwrap(),
Minting {
mints: Arc::clone(&mints),
retries,
},
);
(engine, mints)
}
#[tokio::test]
async fn a_hooks_headers_ride_the_request_the_engine_built() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.and(header("x-tapes-auth", "Bearer token-1"))
.respond_with(ResponseTemplate::new(200).set_body_string("{}"))
.mount(&server)
.await;
let (engine, mints) = minting(&server, 1);
let response = engine
.send(&WireRequest {
method: "GET",
path: "/v1/sessions",
..Default::default()
})
.await
.unwrap();
assert_eq!(response.status, 200);
assert_eq!(mints.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_401_is_retried_once_with_a_freshly_authorised_request() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(401))
.up_to_n_times(1)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(200).set_body_string("{}"))
.mount(&server)
.await;
let (engine, mints) = minting(&server, 1);
let response = engine
.send(&WireRequest {
method: "GET",
path: "/v1/sessions",
..Default::default()
})
.await
.unwrap();
assert_eq!(response.status, 200);
assert_eq!(
mints.load(Ordering::SeqCst),
2,
"the retry was not authorised again"
);
}
#[tokio::test]
async fn an_exhausted_retry_fails_with_the_hooks_own_words() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(401))
.mount(&server)
.await;
let (engine, _) = minting(&server, 1);
let err = engine
.send(&WireRequest {
method: "GET",
path: "/v1/sessions",
..Default::default()
})
.await
.unwrap_err();
assert!(err.to_string().contains("login"), "got: {err}");
}
#[tokio::test]
async fn a_hook_that_never_gives_up_is_stopped_by_the_engine() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(401))
.mount(&server)
.await;
let engine = HttpEngine::with_auth(Url::parse(&server.uri()).unwrap(), Forever);
let err = engine
.send(&WireRequest {
method: "GET",
path: "/v1/sessions",
..Default::default()
})
.await
.unwrap_err();
assert!(err.to_string().contains("times running"), "got: {err}");
assert_eq!(
server.received_requests().await.unwrap_or_default().len(),
MAX_ATTEMPTS as usize,
);
}
#[tokio::test]
async fn a_surfaced_401_reaches_the_caller_with_the_body_that_explains_it() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(401).set_body_string(r#"{"error":"no tenant"}"#))
.mount(&server)
.await;
let response = base(&server.uri())
.send(&WireRequest {
method: "GET",
path: "/v1/sessions",
..Default::default()
})
.await
.unwrap();
assert_eq!(response.status, 401);
assert_eq!(response.body, br#"{"error":"no tenant"}"#.to_vec());
}
#[tokio::test]
async fn an_under_base_join_lands_beneath_the_gateway_prefix() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/primary/tapes/v1/sessions/s-1/traces"))
.respond_with(ResponseTemplate::new(200).set_body_string("{}"))
.mount(&server)
.await;
let engine =
DirectHttp::new(Url::parse(&format!("{}/primary/tapes/", server.uri())).unwrap())
.under_base();
let response = engine
.send(&WireRequest {
method: "GET",
path: "/v1/sessions/{id}/traces",
path_params: vec![("id".to_owned(), "s-1".to_owned())],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.status, 200);
}
#[tokio::test]
async fn an_injected_client_is_the_one_that_sends() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/cassettes"))
.respond_with(ResponseTemplate::new(200).set_body_string(r#"{"cassettes":[]}"#))
.mount(&server)
.await;
let client = reqwest::Client::builder().no_proxy().build().unwrap();
let engine = DirectHttp::new(Url::parse(&server.uri()).unwrap()).with_client(client);
assert_eq!(
engine.fetch_discovery().await.unwrap()["cassettes"],
serde_json::json!([]),
);
}
#[tokio::test]
async fn a_spec_path_may_not_change_the_request_authority() {
let client = base("http://tapes.local:8081");
for path in ["//evil.example/spec.json", "relative/spec.json", ""] {
let err = client.fetch_spec(path, None).await.unwrap_err();
assert!(
err.to_string().contains("non-relative OpenAPI path"),
"{path:?} produced the wrong error: {err}",
);
}
}
#[tokio::test]
async fn a_redirected_spec_fetch_may_not_leave_the_configured_origin() {
let elsewhere = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/spec.json"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"openapi": "3.1.0"
})))
.mount(&elsewhere)
.await;
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/cassettes/x/openapi.json"))
.respond_with(ResponseTemplate::new(302).insert_header(
"location",
format!("{}/spec.json", elsewhere.uri()).as_str(),
))
.mount(&server)
.await;
let client = base(&server.uri());
let err = client
.fetch_spec("/v1/cassettes/x/openapi.json", None)
.await
.unwrap_err();
assert!(err.to_string().contains("redirect"), "wrong error: {err}");
assert!(
elsewhere
.received_requests()
.await
.unwrap_or_default()
.is_empty(),
"the foreign host must never see a request",
);
}
#[tokio::test]
async fn a_matched_validator_reads_as_unchanged_over_real_http() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/cassettes/x/openapi.json"))
.and(header("if-none-match", "\"sha256:abc\""))
.respond_with(ResponseTemplate::new(304))
.mount(&server)
.await;
let got = base(&server.uri())
.fetch_spec("/v1/cassettes/x/openapi.json", Some("\"sha256:abc\""))
.await
.unwrap();
assert!(matches!(got, SpecFetch::Unchanged), "got: {got:?}");
}
#[tokio::test]
async fn an_error_body_is_surfaced_with_the_status() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/cassettes"))
.respond_with(
ResponseTemplate::new(400).set_body_string(r#"{"error":"invalid cursor"}"#),
)
.mount(&server)
.await;
let err = base(&server.uri()).fetch_discovery().await.unwrap_err();
let rendered = format!("{err}");
assert!(rendered.contains("400"), "got: {rendered}");
assert!(rendered.contains("invalid cursor"), "got: {rendered}");
}
#[tokio::test]
async fn a_successful_empty_body_is_null_not_a_decode_failure() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/cassettes"))
.respond_with(ResponseTemplate::new(204))
.mount(&server)
.await;
let got = base(&server.uri()).fetch_discovery().await.unwrap();
assert_eq!(got, Value::Null);
}
#[tokio::test]
async fn a_stream_refuses_a_non_success_status() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions/s-1/export"))
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
.mount(&server)
.await;
let err = base(&server.uri())
.execute_stream(&WireRequest {
method: "GET",
path: "/v1/sessions/{id}/export",
path_params: vec![("id".to_owned(), "s-1".to_owned())],
..Default::default()
})
.await
.unwrap_err();
assert!(
matches!(err, Error::ApiStatus { status: 500, .. }),
"got {err:?}",
);
}
#[tokio::test]
async fn the_sealed_surface_rides_the_same_transport() {
use crate::core::CoreClient;
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/sessions"))
.respond_with(ResponseTemplate::new(200).set_body_string(r#"{"items":[{"id":"s1"}]}"#))
.mount(&server)
.await;
let client = CoreClient::new(base(&server.uri()));
let got: Value = client.call("listSessions", Vec::new()).await.unwrap();
assert_eq!(got["items"][0]["id"], "s1");
}
}