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}