pub mod assertions;
pub mod reporters;
use anyhow::Result;
use handlebars::Handlebars;
use httpstat::{httpstat, Config as HttpstatConfig};
use httpstat::{Header, StatResult as HttpstatResult, Timing as HttpstatTiming};
use serde::de::value::SeqAccessDeserializer;
use serde::de::{Deserializer, IntoDeserializer, SeqAccess, Visitor};
use serde::ser::{SerializeMap, Serializer};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use std::collections::{HashMap, HashSet};
use std::fmt::{self, Formatter};
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::Duration;
use std::{env, fs, str};
use thiserror::Error;
use url::Url;
use assertions::{AssertionConfig, AssertionResult, Equal, RequestAssertionConfig};
use reporters::Reporter;
pub use serde_yaml::Value;
mod handlebars_helpers {
use handlebars::handlebars_helper;
handlebars_helper!(eqi: |x: String, y: String| x.to_lowercase() == y.to_lowercase());
handlebars_helper!(nei: |x: String, y: String| x.to_lowercase() != y.to_lowercase());
}
#[derive(Debug, Clone)]
pub enum HeaderValue {
List(Vec<Header>),
Template(String),
}
impl<'de> Deserialize<'de> for HeaderValue {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct HeaderValueVisitor;
impl<'de> Visitor<'de> for HeaderValueVisitor {
type Value = HeaderValue;
fn expecting(&self, formatter: &mut Formatter) -> fmt::Result {
formatter.write_str("Header template or list of headers")
}
fn visit_str<E>(self, value: &str) -> Result<HeaderValue, E>
where
E: serde::de::Error,
{
let template: String = Deserialize::deserialize(value.into_deserializer())?;
Ok(HeaderValue::Template(template))
}
fn visit_seq<V>(self, seq: V) -> Result<HeaderValue, V::Error>
where
V: SeqAccess<'de>,
{
let headers: Vec<Header> =
Deserialize::deserialize(SeqAccessDeserializer::new(seq))?;
Ok(HeaderValue::List(headers))
}
}
deserializer.deserialize_any(HeaderValueVisitor)
}
}
#[derive(Deserialize, Debug, Clone)]
pub enum RequestData {
StringValue(String),
FilePath(String),
}
#[derive(Deserialize, Debug, Clone)]
#[serde(default)]
pub struct RequestConfig {
#[serde(deserialize_with = "ensure_not_empty")]
pub name: Option<String>,
pub requires: Vec<String>,
pub request_method: String,
#[serde(deserialize_with = "request_data")]
pub data: Option<RequestData>,
pub headers: Option<HeaderValue>,
pub url: String,
pub assertions: Vec<AssertionConfig>,
}
impl Default for RequestConfig {
fn default() -> Self {
Self {
name: Default::default(),
requires: Default::default(),
request_method: "GET".into(),
data: Default::default(),
headers: Default::default(),
url: Default::default(),
assertions: vec![AssertionConfig::Equal(RequestAssertionConfig {
skip: None,
path: Some(".\"response_code\"".into()),
assertion: Equal::new(200),
})],
}
}
}
pub(crate) fn ensure_not_empty<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
where
D: Deserializer<'de>,
{
let result: Result<Option<String>, _> = Option::deserialize(deserializer);
if let Ok(Some(value)) = result {
if !value.trim().is_empty() {
return Ok(Some(value));
}
}
Ok(None)
}
pub(crate) fn request_data<'de, D>(deserializer: D) -> Result<Option<RequestData>, D::Error>
where
D: Deserializer<'de>,
{
let result: Result<Option<String>, _> = Option::deserialize(deserializer);
if let Ok(Some(value)) = result {
if let Some(path) = value.strip_prefix('@') {
return Ok(Some(RequestData::FilePath(path.into())));
} else {
return Ok(Some(RequestData::StringValue(value)));
}
}
Ok(None)
}
#[derive(Deserialize, Default, Debug, Clone)]
pub struct Config {
pub env_var_prefix: Option<String>,
pub location: bool,
pub insecure: bool,
pub client_cert: Option<String>,
pub client_key: Option<String>,
pub ca_cert: Option<String>,
pub connect_timeout: Option<u64>,
pub verbose: bool,
pub fail_request: bool,
pub max_response_size: Option<usize>,
pub requests: Vec<RequestConfig>,
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct Timing {
pub namelookup: u64,
pub connect: u64,
pub pretransfer: u64,
pub starttransfer: u64,
pub total: u64,
pub dns_resolution: u64,
pub tcp_connection: u64,
pub tls_connection: u64,
pub server_processing: u64,
pub content_transfer: u64,
}
impl From<HttpstatTiming> for Timing {
fn from(timing: HttpstatTiming) -> Self {
Self {
namelookup: timing.namelookup.as_millis() as u64,
connect: timing.connect.as_millis() as u64,
pretransfer: timing.pretransfer.as_millis() as u64,
starttransfer: timing.starttransfer.as_millis() as u64,
total: timing.total.as_millis() as u64,
dns_resolution: timing.dns_resolution.as_millis() as u64,
tcp_connection: timing.tcp_connection.as_millis() as u64,
tls_connection: timing.tls_connection.as_millis() as u64,
server_processing: timing.server_processing.as_millis() as u64,
content_transfer: timing.content_transfer.as_millis() as u64,
}
}
}
#[derive(Deserialize, Serialize, Debug, Clone)]
#[serde(untagged)]
pub enum ResponseContent {
NoContent,
Json(JsonValue),
Other(String),
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct StatResult {
pub http_version: String,
pub response_code: i32,
pub response_message: Option<String>,
pub headers: Vec<Header>,
pub timing: Timing,
pub content: ResponseContent,
}
impl From<HttpstatResult> for StatResult {
fn from(result: HttpstatResult) -> Self {
let mut stat_result = Self {
http_version: result.http_version,
response_code: result.response_code,
response_message: result.response_message,
headers: result.headers,
timing: result.timing.into(),
content: ResponseContent::NoContent,
};
if let Ok(body) = str::from_utf8(&result.body[..]) {
if let Ok(json_content) = serde_json::from_str(body) {
stat_result.content = ResponseContent::Json(json_content);
} else {
stat_result.content = ResponseContent::Other(body.into());
}
}
stat_result
}
}
#[derive(Debug, Clone)]
pub struct ResponseData {
inner: Arc<StatResult>,
}
impl From<Arc<StatResult>> for ResponseData {
fn from(result: Arc<StatResult>) -> Self {
Self { inner: result }
}
}
impl Serialize for ResponseData {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut map = serializer.serialize_map(Some(2))?;
map.serialize_entry("headers", &self.inner.headers)?;
map.serialize_entry("content", &self.inner.content)?;
map.end()
}
}
#[derive(Debug, Error)]
pub enum Error {
#[error("Request failed: {1:?}")]
RequestError(Arc<RequestConfig>, String),
#[error("Request dependencies failed: {1:?}")]
DependencyError(Arc<RequestConfig>, String),
}
enum State<F>
where
F: Future<Output = RequestResult>,
{
Wait(Arc<RequestConfig>),
Future(F),
Done(Box<RequestResult>),
}
#[derive(Debug, Clone, Eq, PartialEq, Hash)]
enum StateKey {
Name(String),
Idx(usize),
}
type RequestResult = Result<(Arc<RequestConfig>, Arc<StatResult>)>;
type RequestState = State<Pin<Box<dyn Future<Output = RequestResult>>>>;
struct RequestsFuture<C>
where
C: Serialize,
{
config: Arc<Config>,
context: Arc<Mutex<TemplateContext<C>>>,
states: HashMap<StateKey, RequestState>,
}
impl<C> RequestsFuture<C>
where
C: Serialize,
{
fn new(mut config: Config, context: Option<C>) -> Self {
let mut request_config_map = HashMap::new();
let requests = std::mem::take(&mut config.requests);
for (idx, request_config) in requests.into_iter().enumerate() {
let state_key = match request_config.name {
Some(ref key) => StateKey::Name(key.clone()),
None => StateKey::Idx(idx),
};
request_config_map.insert(state_key, State::Wait(Arc::new(request_config)));
}
let envvars = env::vars()
.filter(|(key, _value)| match config.env_var_prefix {
Some(ref env_var_prefix) => key.starts_with(env_var_prefix),
None => true,
})
.collect();
Self {
context: Arc::new(Mutex::new(TemplateContext {
user: context,
env: envvars,
requests: HashMap::new(),
})),
config: Arc::new(config),
states: request_config_map,
}
}
}
impl<C> Future for RequestsFuture<C>
where
C: Serialize + 'static,
{
type Output = Result<Vec<RequestResult>>;
fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
let mut this = unsafe { self.get_unchecked_mut() };
mark_ready_states(&mut this);
if poll_states(&mut this, context) {
let states = std::mem::take(&mut this.states);
let results = states
.into_values()
.map(|state| match state {
State::Done(result) => *result,
_ => unreachable!(),
})
.collect();
Poll::Ready(Ok(results))
} else {
context.waker().wake_by_ref();
Poll::Pending
}
}
}
#[derive(Debug)]
enum RequirementState {
Wait,
Success,
Error,
}
impl<F> From<&State<F>> for RequirementState
where
F: Future<Output = RequestResult>,
{
fn from(state: &State<F>) -> Self {
match state {
State::Done(result) => match **result {
Ok(_) => Self::Success,
Err(_) => Self::Error,
},
_ => Self::Wait,
}
}
}
fn mark_ready_states<C>(requests_future: &mut RequestsFuture<C>)
where
C: Serialize + 'static,
{
let states = &mut requests_future.states;
let all_map: HashMap<String, RequirementState> = states
.iter()
.filter_map(|(ref state_key, state)| match state_key {
StateKey::Name(key) => Some((key.clone(), RequirementState::from(state))),
_ => None,
})
.collect();
let mut error_set: HashSet<&str> = HashSet::new();
let mut success_set: HashSet<&str> = HashSet::new();
for (key, state) in all_map.iter() {
match state {
RequirementState::Success => {
success_set.insert(key);
}
RequirementState::Error => {
error_set.insert(key);
}
_ => {}
};
}
for (state_key, state) in states.iter_mut() {
if let State::Wait(request_config) = state {
let requirements_validated = if !request_config.requires.is_empty() {
let exists = request_config
.requires
.iter()
.all(|r| all_map.keys().any(|key| r == key));
let requires_self = request_config.requires.iter().any(|r| match state_key {
StateKey::Name(ref key) => r == key,
_ => false,
});
exists && !requires_self
} else {
true
};
if requirements_validated {
let requires_set: HashSet<_> =
request_config.requires.iter().map(|s| s.as_str()).collect();
let success_intersection: HashSet<_> =
success_set.intersection(&requires_set).collect();
let error_intersection: HashSet<_> =
error_set.intersection(&requires_set).collect();
if !error_intersection.is_empty() {
*state = State::Done(Box::new(Err(Error::DependencyError(
request_config.clone(),
error_intersection
.into_iter()
.cloned()
.collect::<Vec<&str>>()
.join(", "),
)
.into())));
} else if success_intersection.len() == requires_set.len() {
*state = State::Future(Box::pin(run_request(
requests_future.config.clone(),
request_config.clone(),
requests_future.context.clone(),
)));
}
} else {
*state = State::Done(Box::new(Err(Error::DependencyError(
request_config.clone(),
"Invalid requirements".into(),
)
.into())));
}
}
}
}
fn poll_states<C>(requests_future: &mut RequestsFuture<C>, context: &mut Context<'_>) -> bool
where
C: Serialize + 'static,
{
let states = &mut requests_future.states;
let mut all_ready = true;
for (state_key, state) in states.iter_mut() {
match state {
State::Future(future) => {
match Pin::new(future).poll(context) {
Poll::Ready(result) => {
let result = Box::new(result);
if let Ok((_, ref stat_result)) = *result {
if let StateKey::Name(name) = state_key {
let mut context = requests_future.context.lock().unwrap();
context
.requests
.insert(name.clone(), ResponseData::from(stat_result.clone()));
}
}
*state = State::Done(result);
}
Poll::Pending => {
all_ready = false;
continue;
}
}
}
State::Wait(_) => {
all_ready = false;
continue;
}
State::Done(_) => continue, }
}
all_ready
}
async fn run_request<C>(
config: Arc<Config>,
request_config: Arc<RequestConfig>,
context: Arc<Mutex<TemplateContext<C>>>,
) -> RequestResult
where
C: Serialize,
{
let mut handlebars = Handlebars::new();
handlebars.register_helper("eqi", Box::new(handlebars_helpers::eqi));
handlebars.register_helper("nei", Box::new(handlebars_helpers::nei));
let headers = match request_config.headers {
Some(HeaderValue::List(ref headers)) => {
let mut rendered_headers = Vec::new();
for header in headers {
rendered_headers.push(Header {
name: header.name.clone(),
value: handlebars.render_template(&header.value, &*context.lock().unwrap())?,
});
}
rendered_headers
}
Some(HeaderValue::Template(ref template)) => {
let rendered_template =
handlebars.render_template(template, &*context.lock().unwrap())?;
let mut rendered_headers = Vec::new();
for line in rendered_template.lines() {
if let Some((name, value)) = line.split_once(':') {
rendered_headers.push(Header {
name: name.into(),
value: value.trim_start().into(),
});
}
}
rendered_headers
}
None => Vec::new(),
};
let url =
Url::parse(&handlebars.render_template(&request_config.url, &*context.lock().unwrap())?)?;
let data = match request_config.data {
Some(RequestData::StringValue(ref data)) => {
Some(handlebars.render_template(data, &*context.lock().unwrap())?)
}
Some(RequestData::FilePath(ref path)) => Some(
handlebars.render_template(&fs::read_to_string(path)?, &*context.lock().unwrap())?,
),
None => None,
};
let fail_request = config.fail_request;
let httpstat_config = HttpstatConfig {
location: config.location,
insecure: config.insecure,
client_cert: config.client_cert.clone(),
client_key: config.client_key.clone(),
ca_cert: config.ca_cert.clone(),
verbose: config.verbose,
connect_timeout: config.connect_timeout.map(Duration::from_millis),
max_response_size: config.max_response_size,
request_method: request_config.request_method.clone().into(),
url: url.into(),
data,
headers,
};
match httpstat(&httpstat_config).await {
Ok(httpstat_result) => {
if fail_request && (400..600).contains(&httpstat_result.response_code) {
Err(Error::RequestError(
request_config,
format!("HTTP {}", httpstat_result.response_code),
)
.into())
} else {
Ok((request_config, Arc::new(httpstat_result.into())))
}
}
Err(error) => Err(Error::RequestError(request_config, error.to_string()).into()),
}
}
#[derive(Serialize, Debug, Clone)]
struct TemplateContext<C>
where
C: Serialize,
{
user: Option<C>,
env: HashMap<String, String>,
requests: HashMap<String, ResponseData>,
}
pub enum UpcakeResult {
Success,
Failures(usize),
}
pub async fn upcake<C, R>(
config: Config,
context: Option<C>,
reporter: &mut R,
) -> Result<UpcakeResult>
where
C: Serialize + 'static,
R: Reporter,
{
reporter.start();
let results = RequestsFuture::new(config, context).await?;
let mut failure_count = 0;
for result in results {
match result {
Ok((request_config, stat_result)) => {
reporter.step_suite(&request_config);
let result = serde_yaml::to_value(&*stat_result)?;
for assertion in request_config.assertions.iter() {
let assertion_result = assertion.assert(&result)?;
match assertion_result {
AssertionResult::Failure(_, _) | AssertionResult::FailureOther(_, _) => {
failure_count += 1;
}
_ => {}
}
reporter.step_result(assertion_result);
}
}
Err(error) => {
failure_count += 1;
match error.downcast_ref::<Error>() {
Some(error) => {
let request_config = match error {
Error::RequestError(request_config, _) => request_config,
Error::DependencyError(request_config, _) => request_config,
};
reporter.step_suite(request_config);
reporter
.step_result(AssertionResult::FailureOther(None, error.to_string()));
}
None => {
reporter
.step_result(AssertionResult::FailureOther(None, format!("{}", error)));
}
};
}
}
}
reporter.end();
if failure_count == 0 {
Ok(UpcakeResult::Success)
} else {
Ok(UpcakeResult::Failures(failure_count))
}
}