use crate::error::{Credential, Error, Result};
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Method {
Get,
Post,
Put,
Patch,
Delete,
}
impl Method {
pub fn as_str(self) -> &'static str {
match self {
Self::Get => "GET",
Self::Post => "POST",
Self::Put => "PUT",
Self::Patch => "PATCH",
Self::Delete => "DELETE",
}
}
}
#[derive(Debug, Clone)]
pub struct Request {
pub method: Method,
pub path: String,
pub body: Option<serde_json::Value>,
pub timeout: Option<Duration>,
}
impl Request {
pub fn new(method: Method, path: impl Into<String>) -> Self {
Self {
method,
path: path.into(),
body: None,
timeout: None,
}
}
pub fn json(mut self, body: serde_json::Value) -> Self {
self.body = Some(body);
self
}
pub fn timeout(mut self, d: Duration) -> Self {
self.timeout = Some(d);
self
}
}
#[derive(Debug, Clone)]
pub struct Response {
pub status: u16,
pub body: String,
}
impl Response {
pub fn is_success(&self) -> bool {
(200..300).contains(&self.status)
}
}
pub trait Transport: Send + Sync + std::fmt::Debug {
fn send(&self, req: Request) -> BoxFuture<'_, Result<Response>>;
fn credential(&self) -> Credential;
}
#[derive(Debug, Clone, Default)]
pub enum Auth {
#[default]
None,
Basic { user: String, password: String },
Token(String),
}
impl Auth {
pub fn header(&self) -> Option<String> {
use base64::Engine;
match self {
Self::None => None,
Self::Basic { user, password } => Some(format!(
"Basic {}",
base64::engine::general_purpose::STANDARD.encode(format!("{user}:{password}"))
)),
Self::Token(t) => Some(format!("Bearer {t}")),
}
}
pub fn credential(&self) -> Credential {
match self {
Self::None => Credential::None,
Self::Basic { .. } => Credential::Basic,
Self::Token(_) => Credential::Token,
}
}
}
pub struct HttpTransport {
client: reqwest::Client,
base: String,
auth: Auth,
}
impl std::fmt::Debug for HttpTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpTransport")
.field("base", &self.base)
.field("auth", &self.auth.credential())
.finish()
}
}
impl HttpTransport {
pub fn new(base: impl Into<String>, auth: Auth, timeout: Duration, insecure: bool) -> Result<Self> {
let client = reqwest::Client::builder()
.timeout(timeout)
.danger_accept_invalid_certs(insecure)
.build()
.map_err(Error::Transport)?;
Ok(Self {
client,
base: normalize_base(&base.into()),
auth,
})
}
pub fn base(&self) -> &str {
&self.base
}
pub fn auth(&self) -> &Auth {
&self.auth
}
}
impl Transport for HttpTransport {
fn send(&self, req: Request) -> BoxFuture<'_, Result<Response>> {
Box::pin(async move {
let url = format!("{}{}", self.base, req.path);
let mut b = self
.client
.request(
reqwest::Method::from_bytes(req.method.as_str().as_bytes())
.expect("a fixed set of valid methods"),
&url,
);
if let Some(h) = self.auth.header() {
b = b.header(reqwest::header::AUTHORIZATION, h);
}
if let Some(body) = &req.body {
b = b.json(body);
}
if let Some(t) = req.timeout {
b = b.timeout(t);
}
let r = b.send().await.map_err(Error::Transport)?;
let status = r.status().as_u16();
let body = r.text().await.map_err(Error::Transport)?;
Ok(Response { status, body })
})
}
fn credential(&self) -> Credential {
self.auth.credential()
}
}
pub fn normalize_base(server: &str) -> String {
let s = server.trim().trim_end_matches('/');
if s.starts_with("http://") || s.starts_with("https://") {
s.to_string()
} else {
format!("http://{s}")
}
}
#[cfg(any(test, feature = "test-util"))]
pub mod stub {
use super::*;
use std::sync::Mutex;
#[derive(Debug, Clone, PartialEq)]
pub struct Seen {
pub method: Method,
pub path: String,
pub body: Option<serde_json::Value>,
}
#[derive(Debug)]
pub struct Stub {
queued: Mutex<std::collections::VecDeque<Response>>,
pub seen: Mutex<Vec<Seen>>,
credential: Credential,
}
impl Default for Stub {
fn default() -> Self {
Self::new()
}
}
impl Stub {
pub fn new() -> Self {
Self {
queued: Mutex::new(Default::default()),
seen: Mutex::new(Vec::new()),
credential: Credential::Token,
}
}
pub fn reply(self, status: u16, body: impl Into<String>) -> Self {
self.queued.lock().unwrap().push_back(Response {
status,
body: body.into(),
});
self
}
pub fn json(self, status: u16, body: serde_json::Value) -> Self {
self.reply(status, body.to_string())
}
pub fn calls(&self) -> Vec<Seen> {
self.seen.lock().unwrap().clone()
}
pub fn call_count(&self) -> usize {
self.seen.lock().unwrap().len()
}
}
impl Transport for Stub {
fn send(&self, req: Request) -> BoxFuture<'_, Result<Response>> {
self.seen.lock().unwrap().push(Seen {
method: req.method,
path: req.path.clone(),
body: req.body.clone(),
});
let next = self.queued.lock().unwrap().pop_front();
Box::pin(async move {
next.ok_or_else(|| {
Error::Invalid(format!(
"the stub had no reply queued for {} {}",
req.method.as_str(),
req.path
))
})
})
}
fn credential(&self) -> Credential {
self.credential
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_bare_host_and_port_becomes_a_url() {
assert_eq!(normalize_base("127.0.0.1:9090"), "http://127.0.0.1:9090");
assert_eq!(normalize_base("http://x:1/"), "http://x:1");
assert_eq!(normalize_base("https://x:1///"), "https://x:1");
assert_eq!(normalize_base(" x:1 "), "http://x:1");
}
#[test]
fn the_basic_header_is_byte_exact() {
let a = Auth::Basic {
user: "admin".into(),
password: "hunter2".into(),
};
assert_eq!(a.header().unwrap(), "Basic YWRtaW46aHVudGVyMg==");
assert!(a.header().unwrap().starts_with("Basic "), "one space, that case");
}
#[test]
fn a_password_containing_a_colon_still_round_trips() {
use base64::Engine;
let a = Auth::Basic {
user: "admin".into(),
password: "pass:word".into(),
};
let encoded = a.header().unwrap();
let raw = base64::engine::general_purpose::STANDARD
.decode(encoded.strip_prefix("Basic ").unwrap())
.unwrap();
assert_eq!(String::from_utf8(raw).unwrap(), "admin:pass:word");
}
#[test]
fn a_token_rides_as_a_bearer() {
let a = Auth::Token("applb_7f3a9c2b1e4d_secret".into());
assert_eq!(a.header().unwrap(), "Bearer applb_7f3a9c2b1e4d_secret");
assert_eq!(a.credential(), Credential::Token);
}
#[test]
fn no_auth_sends_no_header() {
assert!(Auth::None.header().is_none());
assert_eq!(Auth::None.credential(), Credential::None);
}
#[test]
fn debug_output_never_contains_a_credential() {
let t = HttpTransport::new(
"127.0.0.1:9090",
Auth::Basic {
user: "admin".into(),
password: "hunter2".into(),
},
Duration::from_secs(1),
false,
)
.unwrap();
let rendered = format!("{t:?}");
assert!(!rendered.contains("hunter2"), "{rendered}");
assert!(!rendered.contains("YWRtaW4"), "{rendered}");
let t = HttpTransport::new(
"127.0.0.1:9090",
Auth::Token("applb_abc_supersecret".into()),
Duration::from_secs(1),
false,
)
.unwrap();
let rendered = format!("{t:?}");
assert!(!rendered.contains("supersecret"), "{rendered}");
}
}