#![cfg(all(feature = "macros", feature = "json"))]
use std::collections::BTreeSet;
use kynos::{
extract::params::header::HeaderParams,
http::StatusCode,
middleware::{
cors::Cors,
request_id::{RequestId, XRequestId},
},
};
#[path = "support/mod.rs"]
mod support;
use support::{App, get};
async fn baseline() -> BTreeSet<String> {
fields(&get(&support::service(), "/users/1").call().await)
}
fn fields(reply: &support::Reply) -> BTreeSet<String> {
reply
.headers
.keys()
.map(|name| name.as_str().to_owned())
.collect()
}
fn declared<H: HeaderParams>() -> BTreeSet<String> {
H::NAMES
.iter()
.map(|name| name.to_ascii_lowercase())
.collect()
}
#[tokio::test]
async fn a_request_id_sets_the_name_its_group_declares_and_no_other() {
let service = support::router()
.intercept(RequestId::new())
.build(App::new())
.expect("a describable router");
let reply = get(&service, "/users/1").call().await;
assert_eq!(reply.status, StatusCode::OK);
let added: BTreeSet<String> = fields(&reply)
.difference(&baseline().await)
.cloned()
.collect();
assert_eq!(added, declared::<XRequestId>());
assert!(
reply
.field("x-request-id")
.is_some_and(|value| !value.is_empty()),
"the declared name was set to nothing"
);
}
#[tokio::test]
async fn a_service_without_a_request_id_sets_no_such_name() {
let reply = get(&support::service(), "/users/1").call().await;
assert!(reply.field("x-request-id").is_none());
}
#[tokio::test]
async fn a_client_supplied_id_is_used_only_when_it_was_trusted() {
let trusting = support::router()
.intercept(RequestId::new().trust_client(true))
.build(App::new())
.expect("a describable router");
let trusted = get(&trusting, "/users/1")
.header("x-request-id", "from-the-client")
.call()
.await;
assert_eq!(
trusted.field("x-request-id").as_deref(),
Some("from-the-client")
);
let ignored = get(&support::service(), "/users/1")
.header("x-request-id", "from-the-client")
.call()
.await;
assert!(ignored.field("x-request-id").is_none());
let untrusting = support::router()
.intercept(RequestId::new())
.build(App::new())
.expect("a describable router");
let replaced = get(&untrusting, "/users/1")
.header("x-request-id", "from-the-client")
.call()
.await;
assert_ne!(
replaced.field("x-request-id").as_deref(),
Some("from-the-client"),
"an untrusted client id was echoed back"
);
}
#[tokio::test]
async fn cors_adds_only_the_names_it_declares() {
let service = support::router()
.intercept(Cors::new().allow_origins(["https://app.example.com"]))
.build(App::new())
.expect("a describable router");
let reply = get(&service, "/users/1")
.header("origin", "https://app.example.com")
.call()
.await;
let added: BTreeSet<String> = fields(&reply)
.difference(&baseline().await)
.cloned()
.collect();
assert_eq!(
added,
["access-control-allow-origin", "vary"]
.into_iter()
.map(str::to_owned)
.collect::<BTreeSet<_>>()
);
}
#[cfg(feature = "compression")]
#[tokio::test]
async fn two_interceptors_contributing_vary_both_appear_in_it() {
use kynos::middleware::compression::Compression;
let service = support::router()
.intercept(Cors::new().allow_origins(["https://app.example.com"]))
.intercept(Compression::new())
.build(App::new())
.expect("a describable router");
let reply = get(&service, "/users/1")
.header("origin", "https://app.example.com")
.header("accept-encoding", "gzip")
.call()
.await;
let vary = reply.field("vary").expect("a Vary field");
let names: BTreeSet<String> = vary
.split(',')
.map(|name| name.trim().to_ascii_lowercase())
.collect();
assert!(names.contains("origin"), "{vary}");
assert!(names.contains("accept-encoding"), "{vary}");
}
#[test]
fn every_interceptor_kynos_ships_is_accounted_for() {
const WITNESSED: &[&str] = &[
"BodySize",
"BodyTimeout",
"Cache",
"Compression",
"Concurrency",
"Conditional",
"Cors",
"Csrf",
"Decompression",
"RateLimit",
"RequestId",
"SetCookies",
"Timeout",
];
let declared = implementors_of("> Interceptor<C> for ");
assert_eq!(
declared,
WITNESSED
.iter()
.map(|name| (*name).to_owned())
.collect::<BTreeSet<_>>(),
"`middleware/` implements `Interceptor` for a different set than this suite accounts \
for; an interceptor added without a case is one whose declaration nothing reads"
);
}
#[test]
fn every_observer_kynos_ships_is_accounted_for() {
const WITNESSED: &[&str] = &["Trace"];
let declared = implementors_of("> Observer<C> for ");
assert_eq!(
declared,
WITNESSED
.iter()
.map(|name| (*name).to_owned())
.collect::<BTreeSet<_>>(),
"`middleware/` implements `Observer` for a different set than this suite accounts for"
);
}
fn implementors_of(marker: &str) -> BTreeSet<String> {
let mut sources = Vec::new();
collect_sources(
std::path::Path::new(concat!(env!("CARGO_MANIFEST_DIR"), "/src/middleware")),
&mut sources,
);
assert!(
!sources.is_empty(),
"no sources found under `src/middleware/`"
);
sources
.iter()
.flat_map(|source| {
source
.match_indices(marker)
.map(|(at, _)| {
source[at + marker.len()..]
.chars()
.take_while(|character| character.is_alphanumeric() || *character == '_')
.collect::<String>()
})
.collect::<Vec<_>>()
})
.collect()
}
fn collect_sources(directory: &std::path::Path, into: &mut Vec<String>) {
let mut entries: Vec<_> = std::fs::read_dir(directory)
.unwrap_or_else(|error| panic!("read `{}`: {error}", directory.display()))
.map(|entry| entry.expect("read a directory entry").path())
.collect();
entries.sort();
for path in entries {
if path.is_dir() {
collect_sources(&path, into);
} else if path.extension().is_some_and(|extension| extension == "rs")
&& path.file_name().is_some_and(|name| name != "tests.rs")
{
into.push(
std::fs::read_to_string(&path)
.unwrap_or_else(|error| panic!("read `{}`: {error}", path.display())),
);
}
}
}
const SHORT_CIRCUITS: &[&str] = &[
"AtCapacity",
"BodySizeExceeded",
"CrossSite",
"Infallible",
"NotAcceptable",
"NotModified",
"RateLimited",
"RateLimitedFields",
"TimedOut",
"Undecodable",
];
const UNCONSTRUCTIBLE: &[&str] = &["Infallible"];
const UNCONSTRUCTIBLE_STATUSES: &[(&str, &[u16])] = &[(
"Infallible",
<std::convert::Infallible as kynos::response::ShortCircuit>::STATUSES,
)];
#[test]
fn every_unconstructible_short_circuit_declares_no_status() {
assert_eq!(
UNCONSTRUCTIBLE_STATUSES
.iter()
.map(|(name, _)| *name)
.collect::<BTreeSet<_>>(),
UNCONSTRUCTIBLE.iter().copied().collect::<BTreeSet<_>>(),
"a name is excluded from the sweep with nothing here to prove it earned"
);
for (name, statuses) in UNCONSTRUCTIBLE_STATUSES {
assert!(
statuses.is_empty(),
"`{name}` is excluded from the sweep and declares {statuses:?}, which no test reads"
);
}
}
#[test]
fn every_short_circuit_kynos_ships_is_accounted_for() {
let declared = impls_of("ShortCircuit");
assert_eq!(
declared,
SHORT_CIRCUITS
.iter()
.map(|name| (*name).to_owned())
.collect::<BTreeSet<_>>(),
"`src/` implements `ShortCircuit` for a different set than the sweep accounts for; a \
short circuit added without a case is one whose description nothing reads"
);
}
fn impls_of(trait_name: &str) -> BTreeSet<String> {
let mut sources = Vec::new();
collect_sources(
std::path::Path::new(concat!(env!("CARGO_MANIFEST_DIR"), "/src")),
&mut sources,
);
assert!(!sources.is_empty(), "no sources found under `src/`");
names_implementing(trait_name, &sources)
}
fn names_implementing(trait_name: &str, sources: &[String]) -> BTreeSet<String> {
let marker = format!("{trait_name} for ");
sources
.iter()
.flat_map(|source| source.lines())
.filter_map(|line| {
let code = line.trim_start();
if !code.starts_with("impl") {
return None;
}
let at = code.match_indices(&marker).map(|(at, _)| at).find(|at| {
code[..*at]
.chars()
.next_back()
.is_none_or(|character| !character.is_alphanumeric() && character != '_')
})?;
Some(
code[at + marker.len()..]
.chars()
.take_while(|character| character.is_alphanumeric() || *character == '_')
.collect::<String>(),
)
})
.collect()
}
#[test]
fn the_scan_reads_a_generic_head_and_anchors_the_trait_name() {
let sources = [
"impl<const CODE: u16> ShortCircuit for Redirected<CODE> {".to_owned(),
"impl ShortCircuit for Plain {".to_owned(),
" impl crate::response::ShortCircuit for Qualified {".to_owned(),
"impl MyShortCircuit for NotOurs {".to_owned(),
"/// # impl ShortCircuit for InADocExample {".to_owned(),
];
assert_eq!(
names_implementing("ShortCircuit", &sources),
["Plain", "Qualified", "Redirected"]
.into_iter()
.map(str::to_owned)
.collect::<BTreeSet<_>>()
);
}
struct Case {
name: &'static str,
claimed: &'static [u16],
status: u16,
media_type: Option<String>,
body_len: usize,
declared: kynos::openapi::Responses,
}
async fn case<S: kynos::response::ShortCircuit>(
registry: &mut kynos::schema::registry::Registry,
value: S,
) -> Case {
use http_body_util::BodyExt;
let declared = S::responses(registry);
let (parts, body) = value.into_response().into_parts();
let media_type = parts
.headers
.get(kynos::http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.map(|value| {
let (media_type, _) = value.split_once(';').unwrap_or((value, ""));
media_type.trim().to_ascii_lowercase()
});
let body = body.collect().await.expect("a readable body").to_bytes();
Case {
name: std::any::type_name::<S>()
.split('<')
.next()
.expect("a type name is not empty")
.rsplit("::")
.next()
.expect("a type name has a last segment"),
claimed: S::STATUSES,
status: parts.status.as_u16(),
media_type,
body_len: body.len(),
declared,
}
}
async fn every_case() -> Vec<Case> {
use std::time::Duration;
use kynos::middleware::{
csrf::CrossSite,
limits::{AtCapacity, BodySizeExceeded, TimedOut},
rate_limit::refusal::{RateLimited, RateLimitedFields},
};
let registry = &mut kynos::schema::registry::Registry::new();
let mut cases = Vec::new();
cases.push(case(registry, BodySizeExceeded::<()>::new(64)).await);
cases.push(case(registry, TimedOut::<()>::new(Duration::from_secs(1))).await);
cases.push(
case(
registry,
AtCapacity::<()>::new(Some(Duration::from_secs(1))),
)
.await,
);
cases.push(case(registry, CrossSite::<()>::new()).await);
cases.push(case(registry, RateLimited::<()>::new(Duration::from_secs(1), 10)).await);
cases.push(
case(
registry,
RateLimitedFields::<()>::new(Duration::from_secs(1), Vec::new(), Vec::new()),
)
.await,
);
#[cfg(feature = "compression")]
{
use kynos::middleware::{compression::NotAcceptable, decompression::Undecodable};
cases.push(case(registry, NotAcceptable::<()>::new()).await);
cases.push(case(registry, Undecodable::<(), (), ()>::unsupported_coding()).await);
cases.push(case(registry, Undecodable::<(), (), ()>::malformed()).await);
cases.push(case(registry, Undecodable::<(), (), ()>::too_large(64)).await);
}
#[cfg(feature = "cache")]
{
use kynos::{http::HeaderMap, middleware::conditional::NotModified};
cases.push(case(registry, NotModified::from_headers(&HeaderMap::new())).await);
}
cases
}
fn expected_names() -> BTreeSet<&'static str> {
const ABSENT: &[&str] = &[
#[cfg(not(feature = "compression"))]
"NotAcceptable",
#[cfg(not(feature = "compression"))]
"Undecodable",
#[cfg(not(feature = "cache"))]
"NotModified",
];
SHORT_CIRCUITS
.iter()
.copied()
.filter(|name| !UNCONSTRUCTIBLE.contains(name) && !ABSENT.contains(name))
.collect()
}
#[tokio::test]
async fn every_short_circuit_declares_the_content_it_sends() {
let cases = every_case().await;
let driven: BTreeSet<&str> = cases.iter().map(|case| case.name).collect();
assert_eq!(
driven,
expected_names(),
"a short circuit the sweep does not drive"
);
for name in &driven {
let claimed: BTreeSet<u16> = cases
.iter()
.find(|case| &case.name == name)
.expect("a driven case")
.claimed
.iter()
.copied()
.collect();
let reached: BTreeSet<u16> = cases
.iter()
.filter(|case| &case.name == name)
.map(|case| case.status)
.collect();
assert_eq!(
reached, claimed,
"`{name}` declares statuses the sweep does not drive a value to"
);
}
for case in &cases {
let key = kynos::openapi::StatusPattern::Code(case.status).to_string();
let Some(kynos::openapi::RefOr::Item(declared)) = case.declared.responses.get(&key) else {
panic!(
"`{}` answers {} and its description declares no such response",
case.name, case.status
);
};
if let Some(media_type) = case.media_type.as_deref() {
assert!(
declared.content.contains_key(media_type),
"`{}`'s {} sends a `{media_type}` body the description does not declare",
case.name,
case.status
);
} else {
assert!(
case.body_len == 0,
"`{}`'s {} sends {} bytes with no `Content-Type`",
case.name,
case.status,
case.body_len
);
assert!(
declared.content.is_empty(),
"`{}`'s {} declares content it does not send",
case.name,
case.status
);
}
}
}
struct Unsendable(std::marker::PhantomData<*const ()>);
impl kynos::error::problem::ProblemType for Unsendable {
const TYPE_URI: Option<&'static str> = Some("https://errors.example.com/unsendable");
}
struct Bare;
impl kynos::error::problem::ProblemType for Bare {
const TYPE_URI: Option<&'static str> = Some("https://errors.example.com/bare");
}
fn assert_send_sync<T: Send + Sync>() {}
fn assert_refusal_traits<T: Clone + std::fmt::Debug + Eq>() {}
fn assert_clone_and_debug<T: Clone + std::fmt::Debug>() {}
#[test]
fn a_refusal_is_send_and_sync_whatever_marker_names_it() {
use kynos::middleware::{
csrf::CrossSite,
limits::{AtCapacity, BodySizeExceeded, TimedOut},
};
assert_send_sync::<BodySizeExceeded<Unsendable>>();
assert_send_sync::<TimedOut<Unsendable>>();
assert_send_sync::<AtCapacity<Unsendable>>();
assert_send_sync::<CrossSite<Unsendable>>();
#[cfg(feature = "compression")]
{
use kynos::middleware::{compression::NotAcceptable, decompression::Undecodable};
assert_send_sync::<NotAcceptable<Unsendable>>();
assert_send_sync::<Undecodable<Unsendable, Unsendable, Unsendable>>();
}
}
#[test]
fn naming_a_problem_type_costs_the_marker_no_derives() {
use kynos::middleware::{
csrf::{CrossSite, Csrf},
limits::{AtCapacity, BodySize, BodySizeExceeded, Concurrency, TimedOut},
};
assert_refusal_traits::<BodySizeExceeded<Bare>>();
assert_refusal_traits::<TimedOut<Bare>>();
assert_refusal_traits::<AtCapacity<Bare>>();
assert_refusal_traits::<CrossSite<Bare>>();
assert_clone_and_debug::<BodySize<Bare>>();
assert_clone_and_debug::<Concurrency<Bare>>();
assert_clone_and_debug::<Csrf<Bare>>();
#[cfg(feature = "compression")]
{
use kynos::middleware::{
compression::{Compression, NotAcceptable},
decompression::{Decompression, Undecodable},
};
assert_refusal_traits::<NotAcceptable<Bare>>();
assert_refusal_traits::<Undecodable<Bare, Bare, Bare>>();
assert_clone_and_debug::<Compression<Bare>>();
assert_clone_and_debug::<Decompression<Bare, Bare, Bare>>();
}
}