mod add_data;
mod catch_panic;
#[cfg(feature = "compression")]
mod compression;
#[cfg(feature = "cookie")]
mod cookie_jar_manager;
mod cors;
#[cfg(feature = "csrf")]
mod csrf;
mod force_https;
mod normalize_path;
#[cfg(feature = "opentelemetry")]
mod opentelemetry_metrics;
#[cfg(feature = "opentelemetry")]
mod opentelemetry_tracing;
mod propagate_header;
mod sensitive_header;
mod set_header;
mod size_limit;
#[cfg(feature = "tokio-metrics")]
mod tokio_metrics_mw;
#[cfg(feature = "tower-compat")]
mod tower_compat;
mod tracing_mw;
#[cfg(feature = "compression")]
pub use self::compression::{Compression, CompressionEndpoint};
#[cfg(feature = "cookie")]
pub use self::cookie_jar_manager::{CookieJarManager, CookieJarManagerEndpoint};
#[cfg(feature = "csrf")]
pub use self::csrf::{Csrf, CsrfEndpoint};
#[cfg(feature = "opentelemetry")]
pub use self::opentelemetry_metrics::{OpenTelemetryMetrics, OpenTelemetryMetricsEndpoint};
#[cfg(feature = "opentelemetry")]
pub use self::opentelemetry_tracing::{OpenTelemetryTracing, OpenTelemetryTracingEndpoint};
#[cfg(feature = "tokio-metrics")]
pub use self::tokio_metrics_mw::{TokioMetrics, TokioMetricsEndpoint};
#[cfg(feature = "tower-compat")]
pub use self::tower_compat::TowerLayerCompatExt;
pub use self::{
add_data::{AddData, AddDataEndpoint},
catch_panic::{CatchPanic, CatchPanicEndpoint, PanicHandler},
cors::{Cors, CorsEndpoint},
force_https::ForceHttps,
normalize_path::{NormalizePath, NormalizePathEndpoint, TrailingSlash},
propagate_header::{PropagateHeader, PropagateHeaderEndpoint},
sensitive_header::{SensitiveHeader, SensitiveHeaderEndpoint},
set_header::{SetHeader, SetHeaderEndpoint},
size_limit::{SizeLimit, SizeLimitEndpoint},
tracing_mw::{Tracing, TracingEndpoint},
};
use crate::endpoint::Endpoint;
pub trait Middleware<E: Endpoint> {
type Output: Endpoint;
fn transform(&self, ep: E) -> Self::Output;
}
poem_derive::generate_implement_middlewares!();
pub struct FnMiddleware<T>(T);
impl<T, E, E2> Middleware<E> for FnMiddleware<T>
where
T: Fn(E) -> E2,
E: Endpoint,
E2: Endpoint,
{
type Output = E2;
fn transform(&self, ep: E) -> Self::Output {
(self.0)(ep)
}
}
pub fn make<T>(f: T) -> FnMiddleware<T> {
FnMiddleware(f)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
handler,
http::{header::HeaderName, HeaderValue},
test::TestClient,
web::Data,
EndpointExt, IntoResponse, Request, Response, Result,
};
#[tokio::test]
async fn test_make() {
#[handler(internal)]
fn index() -> &'static str {
"abc"
}
struct AddHeader<E> {
ep: E,
header: HeaderName,
value: HeaderValue,
}
#[async_trait::async_trait]
impl<E: Endpoint> Endpoint for AddHeader<E> {
type Output = Response;
async fn call(&self, req: Request) -> Result<Self::Output> {
let mut resp = self.ep.call(req).await?.into_response();
resp.headers_mut()
.insert(self.header.clone(), self.value.clone());
Ok(resp)
}
}
let ep = index.with(make(|ep| AddHeader {
ep,
header: HeaderName::from_static("hello"),
value: HeaderValue::from_static("world"),
}));
let cli = TestClient::new(ep);
let resp = cli.get("/").send().await;
resp.assert_header("hello", "world");
resp.assert_text("abc").await;
}
#[tokio::test]
async fn test_with_multiple_middlewares() {
#[handler(internal)]
fn index(data: Data<&i32>) -> String {
data.0.to_string()
}
let ep = index.with((
AddData::new(10),
SetHeader::new().appending("myheader-1", "a"),
SetHeader::new().appending("myheader-2", "b"),
));
let cli = TestClient::new(ep);
let resp = cli.get("/").send().await;
resp.assert_status_is_ok();
resp.assert_header("myheader-1", "a");
resp.assert_header("myheader-2", "b");
resp.assert_text("10").await;
}
}