Skip to main content

haste_fhir_client/
middleware.rs

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    /// Create a new [`Middleware`] execution chain.
54    ///
55    /// # Panics
56    ///
57    /// This function will panic if the provided `middleware` vector is empty,
58    /// as an empty chain cannot be unwrapped into an execution target.
59    #[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    /// Executes the middleware chain for the given request.
84    ///
85    /// # Errors
86    ///
87    /// Returns an error if any `MiddlewareChain` in the execution pipeline fails
88    /// while processing the state or request context.
89    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}