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#[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 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 let mut cfg = WebServiceConfig::new();
79
80 for mut srv in services {
82 srv.register(&mut cfg);
83 }
84
85 let mut rmap = ResourceMap::new(ResourceDef::new(""));
87 for mut rdef in external {
88 rmap.add(&mut rdef, None);
89 }
90
91 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 let rmap = Rc::new(rmap);
103 rmap.build(&rmap);
104
105 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 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#[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 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#[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}