unb-server 2.0.0

unb inbound server: Node, request/subscribe handlers, catalog, relay orchestration, accept
Documentation
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");
    }
}