Skip to main content

unb_server/
layer.rs

1use std::future::Future;
2use std::pin::Pin;
3use std::sync::Arc;
4
5use bytes::Bytes;
6
7use crate::handler::{EventStream, HandlerError};
8use crate::peer::VerifiedPeer;
9
10#[derive(Clone, Debug)]
11pub enum Origin {
12    Client { session: String },
13    Peer { session: String, peer: VerifiedPeer },
14    Local,
15    Nested,
16}
17
18pub enum ServiceBody {
19    Unary(Bytes),
20    Stream(EventStream),
21}
22
23pub type ErasedCall = Arc<
24    dyn Fn(
25            http::Request<Bytes>,
26        ) -> Pin<
27            Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send>,
28        > + Send
29        + Sync,
30>;
31
32pub trait Layer: Send + Sync + 'static {
33    fn call(
34        &self,
35        request: http::Request<Bytes>,
36        next: Next,
37    ) -> Pin<Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + '_>>;
38}
39
40pub struct Next {
41    layers: Arc<[Arc<dyn Layer>]>,
42    index: usize,
43    terminal: ErasedCall,
44}
45
46impl Next {
47    pub(crate) fn root(layers: Arc<[Arc<dyn Layer>]>, terminal: ErasedCall) -> Next {
48        Next {
49            layers,
50            index: 0,
51            terminal,
52        }
53    }
54
55    pub async fn run(
56        mut self,
57        request: http::Request<Bytes>,
58    ) -> Result<http::Response<ServiceBody>, HandlerError> {
59        if self.index < self.layers.len() {
60            let layer = self.layers[self.index].clone();
61            self.index += 1;
62            layer.call(request, self).await
63        } else {
64            (self.terminal)(request).await
65        }
66    }
67}
68
69pub struct LayerFn<F>(F);
70
71impl<F, Fut> Layer for LayerFn<F>
72where
73    F: Fn(http::Request<Bytes>, Next) -> Fut + Send + Sync + 'static,
74    Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
75{
76    fn call(
77        &self,
78        request: http::Request<Bytes>,
79        next: Next,
80    ) -> Pin<Box<dyn Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + '_>>
81    {
82        Box::pin((self.0)(request, next))
83    }
84}
85
86pub fn layer_fn<F, Fut>(f: F) -> LayerFn<F>
87where
88    F: Fn(http::Request<Bytes>, Next) -> Fut + Send + Sync + 'static,
89    Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
90{
91    LayerFn(f)
92}
93
94#[cfg(test)]
95mod tests {
96    use serde_json::{json, Value};
97    use unb_core::{Envelope, ErrorCode};
98
99    use super::*;
100
101    fn request() -> http::Request<Bytes> {
102        let mut request = http::Request::builder()
103            .method("POST")
104            .uri("/probe")
105            .body(Bytes::from_static(b"{}"))
106            .expect("test request is well formed");
107        request.extensions_mut().insert(Origin::Local);
108        request
109    }
110
111    fn unary(value: Value) -> http::Response<ServiceBody> {
112        http::Response::builder()
113            .body(ServiceBody::Unary(Envelope::encode_payload(&value)))
114            .expect("test response is well formed")
115    }
116
117    fn terminal(trace: Arc<parking_lot::Mutex<Vec<&'static str>>>) -> ErasedCall {
118        Arc::new(move |_request| {
119            let trace = trace.clone();
120            Box::pin(async move {
121                trace.lock().push("handler");
122                Ok(unary(json!({ "done": true })))
123            })
124        })
125    }
126
127    fn tracing_layer(
128        trace: Arc<parking_lot::Mutex<Vec<&'static str>>>,
129        before: &'static str,
130        after: &'static str,
131    ) -> Arc<dyn Layer> {
132        Arc::new(layer_fn(move |request, next: Next| {
133            let trace = trace.clone();
134            async move {
135                trace.lock().push(before);
136                let response = next.run(request).await;
137                trace.lock().push(after);
138                response
139            }
140        }))
141    }
142
143    #[tokio::test]
144    async fn ordered_layers_wrap_the_terminal_and_unwind_in_reverse() {
145        let trace = Arc::new(parking_lot::Mutex::new(Vec::new()));
146        let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![
147            tracing_layer(trace.clone(), "a-before", "a-after"),
148            tracing_layer(trace.clone(), "b-before", "b-after"),
149        ]);
150        let response = Next::root(layers, terminal(trace.clone()))
151            .run(request())
152            .await
153            .unwrap_or_else(|error| panic!("chain failed: {error}"));
154        let ServiceBody::Unary(payload) = response.into_body() else {
155            panic!("expected a unary response");
156        };
157        let value: Value = serde_json::from_slice(&payload).expect("unary payload is json");
158        assert_eq!(value["done"], true);
159        assert_eq!(
160            *trace.lock(),
161            vec!["a-before", "b-before", "handler", "b-after", "a-after"]
162        );
163    }
164
165    #[tokio::test]
166    async fn a_rejecting_layer_stops_the_chain_before_the_terminal() {
167        let trace = Arc::new(parking_lot::Mutex::new(Vec::new()));
168        let reject: Arc<dyn Layer> = Arc::new(layer_fn(|_request, _next: Next| async move {
169            Err(HandlerError::new(ErrorCode::Unauthorized, "no entry"))
170        }));
171        let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![reject]);
172        let error = match Next::root(layers, terminal(trace.clone()))
173            .run(request())
174            .await
175        {
176            Err(error) => error,
177            Ok(_) => panic!("the chain must reject"),
178        };
179        assert_eq!(error.code, ErrorCode::Unauthorized);
180        assert!(trace.lock().is_empty(), "the handler must not run");
181    }
182
183    #[tokio::test]
184    async fn a_typed_extension_flows_downstream_within_one_request() {
185        #[derive(Clone, PartialEq, Debug)]
186        struct Who(&'static str);
187        let enrich: Arc<dyn Layer> =
188            Arc::new(layer_fn(|mut request: http::Request<Bytes>, next: Next| {
189                request.extensions_mut().insert(Who("verified"));
190                async move { next.run(request).await }
191            }));
192        let observe: Arc<dyn Layer> =
193            Arc::new(layer_fn(|request: http::Request<Bytes>, next: Next| {
194                let who = request.extensions().get::<Who>().cloned();
195                async move {
196                    assert_eq!(who, Some(Who("verified")));
197                    next.run(request).await
198                }
199            }));
200        let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![enrich, observe]);
201        let terminal: ErasedCall = Arc::new(|request| {
202            Box::pin(async move {
203                assert_eq!(
204                    request.extensions().get::<Who>(),
205                    Some(&Who("verified")),
206                    "the handler-facing request keeps the typed fact"
207                );
208                Ok(http::Response::builder()
209                    .body(ServiceBody::Unary(Bytes::new()))
210                    .expect("test response is well formed"))
211            })
212        });
213        Next::root(layers, terminal)
214            .run(request())
215            .await
216            .unwrap_or_else(|error| panic!("chain failed: {error}"));
217    }
218
219    #[tokio::test]
220    async fn a_layer_sets_response_headers_the_caller_observes() {
221        let stamp: Arc<dyn Layer> = Arc::new(layer_fn(|request, next: Next| async move {
222            let mut response = next.run(request).await?;
223            response
224                .headers_mut()
225                .insert("x-served-by", http::HeaderValue::from_static("layer"));
226            Ok(response)
227        }));
228        let layers: Arc<[Arc<dyn Layer>]> = Arc::from(vec![stamp]);
229        let terminal: ErasedCall =
230            Arc::new(|_request| Box::pin(async move { Ok(unary(json!({ "done": true }))) }));
231        let response = Next::root(layers, terminal)
232            .run(request())
233            .await
234            .unwrap_or_else(|error| panic!("chain failed: {error}"));
235        assert_eq!(response.headers()["x-served-by"], "layer");
236    }
237}