use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use bytes::Bytes;
use crate::handler::{EventStream, HandlerError};
use crate::peer::VerifiedPeer;
#[derive(Clone, Debug)]
pub enum Origin {
Client { session: String },
Peer { session: String, peer: VerifiedPeer },
Local,
Nested,
}
pub enum ServiceBody {
Unary(Bytes),
Stream(EventStream),
}
pub type ErasedCall = Arc<
dyn Fn(
http::Request<Bytes>,
) -> Pin<
Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send>,
> + Send
+ Sync,
>;
pub trait Layer: Send + Sync + 'static {
fn call(
&self,
request: http::Request<Bytes>,
next: Next,
) -> Pin<Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + '_>>;
}
pub struct Next {
layers: Arc<[Arc<dyn Layer>]>,
index: usize,
terminal: ErasedCall,
}
impl Next {
pub(crate) fn root(layers: Arc<[Arc<dyn Layer>]>, terminal: ErasedCall) -> Next {
Next {
layers,
index: 0,
terminal,
}
}
pub async fn run(
mut self,
request: http::Request<Bytes>,
) -> Result<http::Response<ServiceBody>, HandlerError> {
if self.index < self.layers.len() {
let layer = self.layers[self.index].clone();
self.index += 1;
layer.call(request, self).await
} else {
(self.terminal)(request).await
}
}
}
pub struct LayerFn<F>(F);
impl<F, Fut> Layer for LayerFn<F>
where
F: Fn(http::Request<Bytes>, Next) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
{
fn call(
&self,
request: http::Request<Bytes>,
next: Next,
) -> Pin<Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + '_>>
{
Box::pin((self.0)(request, next))
}
}
pub fn layer_fn<F, Fut>(f: F) -> LayerFn<F>
where
F: Fn(http::Request<Bytes>, Next) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
{
LayerFn(f)
}
#[cfg(test)]
mod tests {
use serde_json::{json, Value};
use unb_core::{Envelope, ErrorCode};
use super::*;
fn request() -> http::Request<Bytes> {
let mut request = http::Request::builder()
.method("POST")
.uri("/probe")
.body(Bytes::from_static(b"{}"))
.expect("test request is well formed");
request.extensions_mut().insert(Origin::Local);
request
}
fn unary(value: Value) -> http::Response<ServiceBody> {
http::Response::builder()
.body(ServiceBody::Unary(Envelope::encode_payload(&value)))
.expect("test response is well formed")
}
fn terminal(trace: Arc<parking_lot::Mutex<Vec<&'static str>>>) -> ErasedCall {
Arc::new(move |_request| {
let trace = trace.clone();
Box::pin(async move {
trace.lock().push("handler");
Ok(unary(json!({ "done": true })))
})
})
}
fn tracing_layer(
trace: Arc<parking_lot::Mutex<Vec<&'static str>>>,
before: &'static str,
after: &'static str,
) -> Arc<dyn Layer> {
Arc::new(layer_fn(move |request, next: Next| {
let trace = trace.clone();
async move {
trace.lock().push(before);
let response = next.run(request).await;
trace.lock().push(after);
response
}
}))
}
#[tokio::test]
async fn ordered_layers_wrap_the_terminal_and_unwind_in_reverse() {
let trace = Arc::new(parking_lot::Mutex::new(Vec::new()));
let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![
tracing_layer(trace.clone(), "a-before", "a-after"),
tracing_layer(trace.clone(), "b-before", "b-after"),
]);
let response = Next::root(layers, terminal(trace.clone()))
.run(request())
.await
.unwrap_or_else(|error| panic!("chain failed: {error}"));
let ServiceBody::Unary(payload) = response.into_body() else {
panic!("expected a unary response");
};
let value: Value = serde_json::from_slice(&payload).expect("unary payload is json");
assert_eq!(value["done"], true);
assert_eq!(
*trace.lock(),
vec!["a-before", "b-before", "handler", "b-after", "a-after"]
);
}
#[tokio::test]
async fn a_rejecting_layer_stops_the_chain_before_the_terminal() {
let trace = Arc::new(parking_lot::Mutex::new(Vec::new()));
let reject: Arc<dyn Layer> = Arc::new(layer_fn(|_request, _next: Next| async move {
Err(HandlerError::new(ErrorCode::Unauthorized, "no entry"))
}));
let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![reject]);
let error = match Next::root(layers, terminal(trace.clone()))
.run(request())
.await
{
Err(error) => error,
Ok(_) => panic!("the chain must reject"),
};
assert_eq!(error.code, ErrorCode::Unauthorized);
assert!(trace.lock().is_empty(), "the handler must not run");
}
#[tokio::test]
async fn a_typed_extension_flows_downstream_within_one_request() {
#[derive(Clone, PartialEq, Debug)]
struct Who(&'static str);
let enrich: Arc<dyn Layer> =
Arc::new(layer_fn(|mut request: http::Request<Bytes>, next: Next| {
request.extensions_mut().insert(Who("verified"));
async move { next.run(request).await }
}));
let observe: Arc<dyn Layer> =
Arc::new(layer_fn(|request: http::Request<Bytes>, next: Next| {
let who = request.extensions().get::<Who>().cloned();
async move {
assert_eq!(who, Some(Who("verified")));
next.run(request).await
}
}));
let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![enrich, observe]);
let terminal: ErasedCall = Arc::new(|request| {
Box::pin(async move {
assert_eq!(
request.extensions().get::<Who>(),
Some(&Who("verified")),
"the handler-facing request keeps the typed fact"
);
Ok(http::Response::builder()
.body(ServiceBody::Unary(Bytes::new()))
.expect("test response is well formed"))
})
});
Next::root(layers, terminal)
.run(request())
.await
.unwrap_or_else(|error| panic!("chain failed: {error}"));
}
#[tokio::test]
async fn a_layer_sets_response_headers_the_caller_observes() {
let stamp: Arc<dyn Layer> = Arc::new(layer_fn(|request, next: Next| async move {
let mut response = next.run(request).await?;
response
.headers_mut()
.insert("x-served-by", http::HeaderValue::from_static("layer"));
Ok(response)
}));
let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![stamp]);
let terminal: ErasedCall =
Arc::new(|_request| Box::pin(async move { Ok(unary(json!({ "done": true }))) }));
let response = Next::root(layers, terminal)
.run(request())
.await
.unwrap_or_else(|error| panic!("chain failed: {error}"));
assert_eq!(response.headers()["x-served-by"], "layer");
}
}