1use std::{fmt::Debug, pin::Pin, sync::Arc};
2
3#[derive(Debug)]
4pub struct Context<CTX: Debug, Request: Debug, Response: Debug> {
5 pub ctx: CTX,
6 pub request: Request,
7 pub response: Option<Response>,
8}
9
10pub type MiddlewareOutput<Context, Error> =
11 Pin<Box<dyn Future<Output = Result<Context, Error>> + Send>>;
12pub type Next<State, Context, Error> =
13 Box<dyn Fn(State, Context) -> MiddlewareOutput<Context, Error> + Send + Sync>;
14
15pub type MiddlewareChainOld<State, CTX, Request, Response, Error> = Box<
16 dyn Fn(
17 State,
18 Context<CTX, Request, Response>,
19 Option<Arc<Next<State, Context<CTX, Request, Response>, Error>>>,
20 ) -> MiddlewareOutput<Context<CTX, Request, Response>, Error>
21 + Send
22 + Sync,
23>;
24
25pub type MiddlewareNext<'a, State, CTX, Request, Response, Error> =
26 Arc<Next<State, Context<CTX, Request, Response>, Error>>;
27
28pub trait MiddlewareChain<State, CTX: Debug, Request: Debug, Response: Debug, Error>:
29 Send + Sync
30{
31 fn call(
32 &self,
33 state: State,
34 ctx: Context<CTX, Request, Response>,
35 next: Option<MiddlewareNext<State, CTX, Request, Response, Error>>,
36 ) -> MiddlewareOutput<Context<CTX, Request, Response>, Error>;
37}
38
39pub struct Middleware<State, CTX: Debug, Request: Debug, Response: Debug, Error> {
40 _state: std::marker::PhantomData<State>,
41 _phantom: std::marker::PhantomData<CTX>,
42 execute: Arc<Next<State, Context<CTX, Request, Response>, Error>>,
43}
44
45impl<
46 State: 'static + Send + Sync,
47 CTX: 'static + Send + Sync + Debug,
48 Request: 'static + Send + Sync + Debug,
49 Response: 'static + Send + Sync + Debug,
50 Error: 'static + Send + Sync,
51> Middleware<State, CTX, Request, Response, Error>
52{
53 #[must_use]
60 pub fn new(
61 mut middleware: Vec<Box<dyn MiddlewareChain<State, CTX, Request, Response, Error>>>,
62 ) -> Self {
63 middleware.reverse();
64 let next: Option<MiddlewareNext<State, CTX, Request, Response, Error>> = middleware
65 .into_iter()
66 .fold(
67 None,
68 |prev_next: Option<MiddlewareNext<State, CTX, Request, Response, Error>>,
69 middleware: Box<dyn MiddlewareChain<State, CTX, Request, Response, Error>>| {
70 Some(Arc::new(Box::new(move |state, ctx| {
71 middleware.call(state, ctx, prev_next.clone())
72 })))
73 },
74 );
75
76 Middleware {
77 _state: std::marker::PhantomData,
78 _phantom: std::marker::PhantomData,
79 execute: next.unwrap(),
80 }
81 }
82
83 pub async fn call(
90 &self,
91 state: State,
92 ctx: CTX,
93 request: Request,
94 ) -> Result<Context<CTX, Request, Response>, Error> {
95 (self.execute)(
96 state,
97 Context {
98 ctx,
99 request,
100 response: None,
101 },
102 )
103 .await
104 }
105}
106
107#[cfg(test)]
108mod test {
109 use super::*;
110
111 struct MiddlewareChain1;
112 impl MiddlewareChain<(), (), usize, usize, String> for MiddlewareChain1 {
113 fn call(
114 &self,
115 _state: (),
116 x: Context<(), usize, usize>,
117 next: Option<Arc<Next<(), Context<(), usize, usize>, String>>>,
118 ) -> Pin<Box<dyn Future<Output = Result<Context<(), usize, usize>, String>> + Send>>
119 {
120 Box::pin(async move {
121 let mut x = if let Some(next) = next {
122 next((), x).await
123 } else {
124 Ok(x)
125 }?;
126 println!("Middleware 1 executed");
127 x.response = x.response.map(|r| r + 1);
128 Ok(x)
129 })
130 }
131 }
132
133 struct MiddlewareChain2;
134 impl MiddlewareChain<(), (), usize, usize, String> for MiddlewareChain2 {
135 fn call(
136 &self,
137 _state: (),
138 x: Context<(), usize, usize>,
139 next: Option<Arc<Next<(), Context<(), usize, usize>, String>>>,
140 ) -> Pin<Box<dyn Future<Output = Result<Context<(), usize, usize>, String>> + Send>>
141 {
142 Box::pin(async move {
143 let mut x = if let Some(next) = next {
144 next((), x).await
145 } else {
146 Ok(x)
147 }?;
148
149 println!("Middleware 2 executed {:?}", x.response);
150 x.response = x.response.map(|r| r + 2);
151 Ok(x)
152 })
153 }
154 }
155
156 struct MiddlewareChain3;
157 impl MiddlewareChain<(), (), usize, usize, String> for MiddlewareChain3 {
158 fn call(
159 &self,
160 _state: (),
161 x: Context<(), usize, usize>,
162 next: Option<Arc<Next<(), Context<(), usize, usize>, String>>>,
163 ) -> Pin<Box<dyn Future<Output = Result<Context<(), usize, usize>, String>> + Send>>
164 {
165 Box::pin(async move {
166 let mut x = if let Some(next) = next {
167 next((), x).await
168 } else {
169 Ok(x)
170 }?;
171
172 x.response = x.response.map_or(Some(x.request + 3), |r| Some(r + 3));
173 Ok(x)
174 })
175 }
176 }
177
178 #[tokio::test]
179 async fn test_middleware() {
180 let test = Middleware::new(vec![
181 Box::new(MiddlewareChain1 {}),
182 Box::new(MiddlewareChain2 {}),
183 Box::new(MiddlewareChain3 {}),
184 ]);
185
186 let ret = test.call((), (), 42).await;
187 assert_eq!(Some(48), ret.unwrap().response);
188 }
189}