mod body_limit;
mod cors;
mod error_handling;
#[cfg(not(target_arch = "wasm32"))]
mod timeout;
pub mod auth;
pub mod compression;
use std::fmt::{self, Debug};
use http_kit::{Endpoint, Request, Response};
use skyzen_core::{
middleware::{boxed, BoxFuture, Dispatch},
Error,
};
pub use body_limit::BodyLimit;
pub use compression::{CompressionEncoding, CompressionLevel, CompressionMiddleware};
pub use cors::{AllowOrigin, Cors, CorsConfigError};
pub use error_handling::ErrorHandlingMiddleware;
#[doc(inline)]
pub use skyzen_core::middleware::{
apply, from_fn, BoxMiddleware, FromFn, Middleware, MiddlewareFn, Next,
};
#[cfg(not(target_arch = "wasm32"))]
pub use timeout::Timeout;
pub struct Layered<E> {
endpoint: E,
middleware: [BoxMiddleware; 1],
}
impl<E: Debug> Debug for Layered<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Layered")
.field("endpoint", &self.endpoint)
.field("middleware", &self.middleware[0].middleware_name())
.finish()
}
}
impl<E: Clone> Clone for Layered<E> {
fn clone(&self) -> Self {
Self {
endpoint: self.endpoint.clone(),
middleware: self.middleware.clone(),
}
}
}
impl<E> Layered<E> {
pub const fn endpoint(&self) -> &E {
&self.endpoint
}
}
pub fn layer<E, M>(endpoint: E, middleware: M) -> Layered<E>
where
E: Endpoint + Clone + Send + Sync + 'static,
M: Middleware,
{
Layered {
endpoint,
middleware: [boxed(middleware)],
}
}
impl<E> Endpoint for Layered<E>
where
E: Endpoint + Clone + Send + Sync + 'static,
{
type Error = http_kit::error::BoxHttpError;
async fn respond(&mut self, request: &mut Request) -> Result<Response, Self::Error> {
let terminal = CloneEndpoint(&self.endpoint);
Next::new(&self.middleware, &terminal)
.run(request)
.await
.map_err(Error::into_boxed_http_error)
}
}
struct CloneEndpoint<'e, E>(&'e E);
impl<E> Dispatch for CloneEndpoint<'_, E>
where
E: Endpoint + Clone + Send + Sync + 'static,
{
fn dispatch<'a>(&'a self, request: &'a mut Request) -> BoxFuture<'a, Result<Response, Error>> {
Box::pin(async move {
let mut endpoint = self.0.clone();
endpoint.respond(request).await.map_err(Error::from)
})
}
}
#[cfg(test)]
mod tests {
use super::{layer, Middleware, Next};
use crate::{
routing::{CreateRouteNode, Route},
Body, Endpoint, Error, Request, Response, Result,
};
use http_kit::header::{HeaderName, HeaderValue};
#[derive(Debug)]
struct Stamp(&'static str);
impl Middleware for Stamp {
async fn handle(
&self,
request: &mut Request,
next: Next<'_>,
) -> std::result::Result<Response, Error> {
let mut response = next.run(request).await?;
response.headers_mut().append(
HeaderName::from_static("x-stamp"),
HeaderValue::from_static(self.0),
);
Ok(response)
}
}
#[tokio::test]
async fn layering_an_endpoint_applies_the_last_wrapper_outermost() {
let router = Route::new(("/ping".at(|| async { Result::Ok("pong") }),)).build();
let mut endpoint = layer(layer(router, Stamp("inner")), Stamp("outer"));
let mut request = Request::new(Body::empty());
*request.uri_mut() = "/ping".parse().expect("valid path");
let response = endpoint.respond(&mut request).await.unwrap();
let stamps: Vec<&str> = response
.headers()
.get_all("x-stamp")
.iter()
.map(|value| value.to_str().unwrap())
.collect();
assert_eq!(stamps, ["inner", "outer"]);
}
}