Skip to main content

ntex/web/
app_service.rs

1use std::{cell::RefCell, marker::PhantomData, mem, rc::Rc};
2
3use crate::error::Failure;
4use crate::http::{Request, Response};
5use crate::router::{Path, ResourceDef, ResourceId, Router};
6use crate::service::{ServiceChainFactory, boxed, cfg::Cfg, cfg::Configuration};
7use crate::util::HashMap;
8use crate::{Ctx, Middleware, Service, ServiceFactory, factory};
9
10use super::config::WebAppConfig;
11use super::guard::Guard;
12use super::rmap::ResourceMap;
13use super::service::{AppServiceFactory, WebServiceConfig};
14use super::{HttpHandler, HttpRequest, HttpService, State, WebError, WebRequest, WebResponse};
15
16type Guards = Vec<Box<dyn Guard>>;
17
18/// Service factory to convert `Request` to a `WebRequest`.
19/// It also executes state factories.
20#[derive(derive_more::Debug)]
21#[debug("AppFactory")]
22pub struct AppFactory<St, In, Out, M, F>
23where
24    St: State,
25    F: ServiceFactory<
26            St,
27            WebRequest<In>,
28            Res = WebRequest<Out>,
29            Error = WebError<St, St::Error>,
30            InitError = Failure,
31        >,
32    M: Middleware<WebServiceRouter<St, In, Out, F::Service>, St> + 'static,
33    M::Service: Service<St, WebRequest<()>, Res = WebResponse, Error = WebError<St, St::Error>>,
34{
35    middleware: M,
36    filter: ServiceChainFactory<F, St, WebRequest<In>>,
37    rmap: Rc<ResourceMap>,
38    router: Rc<Router<HttpService<St, Out>, Guards>>,
39    default: HttpService<St, Out>,
40    config: Option<Cfg<WebAppConfig>>,
41}
42
43impl<St, In, Out, M, F> AppFactory<St, In, Out, M, F>
44where
45    St: State,
46    In: 'static,
47    Out: 'static,
48    F: ServiceFactory<
49            St,
50            WebRequest<In>,
51            Res = WebRequest<Out>,
52            Error = WebError<St, St::Error>,
53            InitError = Failure,
54        >,
55    M: Middleware<WebServiceRouter<St, In, Out, F::Service>, St> + 'static,
56    M::Service: Service<St, WebRequest<()>, Res = WebResponse, Error = WebError<St, St::Error>>,
57{
58    pub(super) fn new(
59        middleware: M,
60        filter: ServiceChainFactory<F, St, WebRequest<In>>,
61        services: Vec<Box<dyn AppServiceFactory<St, Out>>>,
62        default: Option<HttpService<St, Out>>,
63        config: Option<Cfg<WebAppConfig>>,
64        external: Vec<ResourceDef>,
65        case_insensitive: bool,
66    ) -> Self {
67        // Default service
68        let default = default.unwrap_or_else(|| {
69            boxed::factory(
70                factory(async move |req: WebRequest<Out>| {
71                    Ok(req.into_response(Response::NotFound().build()))
72                })
73                .map_init_err(|_| unreachable!()),
74            )
75        });
76
77        // Web app config
78        let mut cfg = WebServiceConfig::new();
79
80        // register services
81        for mut srv in services {
82            srv.register(&mut cfg);
83        }
84
85        // ResourceMap tree
86        let mut rmap = ResourceMap::new(ResourceDef::new(""));
87        for mut rdef in external {
88            rmap.add(&mut rdef, None);
89        }
90
91        // Complete pipeline creation
92        let services = cfg.into_services();
93        let services: Vec<_> = services
94            .into_iter()
95            .map(|(mut rdef, srv, guards, nested)| {
96                rmap.add(&mut rdef, nested);
97                (rdef, srv, RefCell::new(guards))
98            })
99            .collect();
100
101        // complete ResourceMap tree
102        let rmap = Rc::new(rmap);
103        rmap.build(&rmap);
104
105        // Create router
106        let mut router = Router::builder();
107        if case_insensitive {
108            router.case_insensitive();
109        }
110        for (path, factory, guards) in services {
111            router
112                .resource(path.clone(), factory)
113                .set_check_value(guards.borrow_mut().take());
114        }
115
116        Self {
117            rmap,
118            filter,
119            default,
120            config,
121            middleware,
122            router: Rc::new(router.build()),
123        }
124    }
125}
126
127impl<St, In, Out, M, F> ServiceFactory<St, Request> for AppFactory<St, In, Out, M, F>
128where
129    St: State,
130    In: 'static,
131    Out: 'static,
132    F: ServiceFactory<
133            St,
134            WebRequest<In>,
135            Res = WebRequest<Out>,
136            Error = WebError<St, St::Error>,
137            InitError = Failure,
138        >,
139    M: Middleware<WebServiceRouter<St, In, Out, F::Service>, St> + 'static,
140    M::Service: Service<St, WebRequest<()>, Res = WebResponse, Error = WebError<St, St::Error>>,
141{
142    type Res = Response;
143    type Error = WebError<St, St::Error>;
144
145    type Service = AppService<St, M::Service>;
146    type InitError = Failure;
147
148    async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
149        // main service
150        let service = WebServiceRouter::new(
151            self.filter.create(st).await?,
152            self.router.clone(),
153            self.default.clone(),
154        );
155
156        Ok(AppService {
157            service: self.middleware.create(st, service),
158            rmap: self.rmap.clone(),
159            config: self.config.clone(),
160            _t: PhantomData,
161        })
162    }
163}
164
165/// Service to convert `Request` to a `WebRequest`
166#[derive(derive_more::Debug)]
167#[debug("AppService")]
168pub struct AppService<St, F>
169where
170    St: State,
171    F: Service<St, WebRequest<()>, Res = WebResponse, Error = WebError<St, St::Error>>,
172{
173    service: F,
174    rmap: Rc<ResourceMap>,
175    config: Option<Cfg<WebAppConfig>>,
176    _t: PhantomData<St>,
177}
178
179impl<St, F> Service<St, Request> for AppService<St, F>
180where
181    St: State,
182    F: Service<St, WebRequest<()>, Res = WebResponse, Error = WebError<St, St::Error>>,
183{
184    type Res = Response;
185    type Error = F::Error;
186
187    crate::forward_ready!(St, service);
188    crate::forward_shutdown!(St, service);
189
190    async fn call(&self, req: Request, ctx: Ctx<'_, Self, St>) -> Result<Self::Res, F::Error> {
191        let config: Cfg<WebAppConfig> = if let Some(cfg) = &self.config {
192            cfg.clone()
193        } else if let Some(io) = req.io() {
194            io.cfg().ctx().get()
195        } else {
196            Cfg::<WebAppConfig>::default()
197        };
198
199        let (head, payload) = req.into_parts();
200
201        let req = if let Some(mut req) = config.get_request() {
202            let inner = Rc::get_mut(&mut req.0).unwrap();
203            inner.path.set(head.uri.clone());
204            inner.head = head;
205            // the pool is shared by all applications that use this config
206            if !Rc::ptr_eq(&inner.rmap, &self.rmap) {
207                inner.rmap = self.rmap.clone();
208            }
209            req
210        } else {
211            HttpRequest::new(Path::new(head.uri.clone()), head, self.rmap.clone(), config)
212        };
213        match ctx
214            .call(&self.service, WebRequest::new(req, payload, ()))
215            .await
216        {
217            Ok(r) => Ok(r.into()),
218            Err(e) => Ok(e.0.error_response(ctx.st())),
219        }
220    }
221}
222
223/// Web app service.
224#[derive(derive_more::Debug)]
225#[debug("Router")]
226pub struct WebServiceRouter<St: State, In, Out, F> {
227    filter: F,
228    router: Rc<Router<HttpService<St, Out>, Guards>>,
229    default: HttpService<St, Out>,
230    cache: RefCell<HashMap<ResourceId, HttpHandler<St, Out>>>,
231    cache_default: RefCell<Option<HttpHandler<St, Out>>>,
232    ph: PhantomData<In>,
233}
234
235impl<St: State, In, Out, F> WebServiceRouter<St, In, Out, F> {
236    pub fn new(
237        filter: F,
238        router: Rc<Router<HttpService<St, Out>, Guards>>,
239        default: HttpService<St, Out>,
240    ) -> Self {
241        Self {
242            filter,
243            router,
244            default,
245            cache: RefCell::new(HashMap::default()),
246            cache_default: RefCell::new(None),
247            ph: PhantomData,
248        }
249    }
250}
251
252impl<St, In, Out, F> Service<St, WebRequest<In>> for WebServiceRouter<St, In, Out, F>
253where
254    St: State,
255    Out: 'static,
256    F: Service<St, WebRequest<In>, Res = WebRequest<Out>, Error = WebError<St, St::Error>>,
257{
258    type Res = WebResponse;
259    type Error = WebError<St, St::Error>;
260
261    #[inline]
262    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Self::Error> {
263        ctx.ready(&self.filter).await
264    }
265
266    async fn call(
267        &self,
268        req: WebRequest<In>,
269        ctx: Ctx<'_, Self, St>,
270    ) -> Result<Self::Res, Self::Error> {
271        let mut req = ctx.call(&self.filter, req).await?;
272        let res = self.router.recognize_checked(&mut req, |req, guards| {
273            if let Some(guards) = guards {
274                for f in guards {
275                    if !f.check(req.head()) {
276                        return false;
277                    }
278                }
279            }
280            true
281        });
282
283        let svc = if let Some((sf, id)) = res {
284            if let Some(svc) = self.cache.borrow().get(&id) {
285                svc.clone()
286            } else if let Ok(svc) = sf.create(ctx.st()).await {
287                self.cache.borrow_mut().insert(id, svc.clone());
288                svc
289            } else {
290                return Ok(req.into_response(Response::InternalServerError().build()));
291            }
292        } else {
293            if let Some(svc) = &*self.cache_default.borrow() {
294                svc.clone()
295            } else if let Ok(svc) = self.default.create(ctx.st()).await {
296                *self.cache_default.borrow_mut() = Some(svc.clone());
297                svc
298            } else {
299                return Ok(req.into_response(Response::InternalServerError().build()));
300            }
301        };
302        ctx.call(&svc, req).await
303    }
304
305    async fn shutdown(&self, ctx: Ctx<'_, Self, St>) {
306        ctx.shutdown(&self.filter).await;
307
308        let svc = self.cache_default.borrow_mut().take();
309        if let Some(svc) = svc {
310            ctx.shutdown(&svc).await;
311        }
312
313        let services = mem::take(&mut *self.cache.borrow_mut());
314        for (_, svc) in services {
315            ctx.shutdown(&svc).await;
316        }
317    }
318}