1use std::{fmt, mem, rc::Rc};
2
3use crate::error::Failure;
4use crate::http::Method;
5use crate::service::{Ctx, Service, ServiceFactory};
6
7use super::error::{WebError, WebResponseError};
8use super::guard::{self, AllGuard, Guard};
9use super::handler::{Handler, HandlerFn, HandlerSt, HandlerStWrapper, HandlerWrapper};
10use super::{FromRequest, HttpResponse, State, WebRequest, WebResponse};
11
12pub struct Route<St: State, In = ()> {
49 handler: Rc<dyn HandlerFn<St, In>>,
50 methods: Vec<Method>,
51 guards: Rc<AllGuard>,
52}
53
54impl<St: State, In: 'static> Route<St, In> {
55 pub fn new() -> Route<St, In> {
57 Route {
58 handler: HandlerWrapper::<St, In, _, ()>::create(async || HttpResponse::NotFound()),
59 methods: Vec::new(),
60 guards: Rc::default(),
61 }
62 }
63
64 pub(super) fn take_guards(&mut self) -> Vec<Box<dyn Guard>> {
65 for m in &self.methods {
66 Rc::get_mut(&mut self.guards)
67 .unwrap()
68 .add(guard::Method(m.clone()));
69 }
70
71 mem::take(&mut Rc::get_mut(&mut self.guards).unwrap().0)
72 }
73
74 pub(super) fn service(&self) -> RouteService<St, In> {
75 RouteService {
76 handler: self.handler.clone(),
77 guards: self.guards.clone(),
78 methods: self.methods.clone(),
79 }
80 }
81}
82
83impl<St: State, In: 'static> Default for Route<St, In> {
84 fn default() -> Self {
85 Self::new()
86 }
87}
88
89impl<St: State, In: 'static> ServiceFactory<St, WebRequest<In>> for Route<St, In> {
90 type Res = WebResponse;
91 type Error = WebError<St, St::Error>;
92
93 type Service = RouteService<St, In>;
94 type InitError = Failure;
95
96 async fn create(&self, _: &St) -> Result<Self::Service, Self::InitError> {
97 Ok(self.service())
98 }
99}
100
101impl<St: State, In: 'static> Route<St, In> {
102 #[must_use]
116 pub fn method(mut self, method: Method) -> Self {
117 self.methods.push(method);
118 self
119 }
120
121 #[must_use]
145 pub fn guard<F: Guard + 'static>(mut self, f: F) -> Self {
146 Rc::get_mut(&mut self.guards).unwrap().add(f);
147 self
148 }
149
150 #[must_use]
184 pub fn to<H, Args>(mut self, handler: H) -> Self
185 where
186 H: Handler<St, Args> + 'static,
187 Args: FromRequest<St> + 'static,
188 Args::Error: WebResponseError<St, St::Error>,
189 {
190 self.handler = HandlerWrapper::create(handler);
191 self
192 }
193
194 #[must_use]
224 pub fn to_with_state<H, Args>(mut self, handler: H) -> Self
225 where
226 H: HandlerSt<St, In, Args> + 'static,
227 Args: FromRequest<St> + 'static,
228 Args::Error: WebResponseError<St, St::Error>,
229 {
230 self.handler = HandlerStWrapper::create(handler);
231 self
232 }
233}
234
235impl<St: State, In> fmt::Debug for Route<St, In> {
236 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
237 f.debug_struct("Route")
238 .field("handler", &self.handler)
239 .field("methods", &self.methods)
240 .field("guards", &self.guards)
241 .finish()
242 }
243}
244
245pub struct RouteService<St: State, In> {
246 handler: Rc<dyn HandlerFn<St, In>>,
247 methods: Vec<Method>,
248 guards: Rc<AllGuard>,
249}
250
251impl<St: State, In> RouteService<St, In> {
252 pub fn check(&self, req: &mut WebRequest<In>) -> bool {
253 if !self.methods.is_empty() && !self.methods.contains(&req.head().method) {
254 return false;
255 }
256
257 self.guards.check(req.head())
258 }
259}
260
261impl<St: State, In> fmt::Debug for RouteService<St, In> {
262 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
263 f.debug_struct("RouteService")
264 .field("handler", &self.handler)
265 .field("methods", &self.methods)
266 .field("guards", &self.guards)
267 .finish()
268 }
269}
270
271impl<St: State, In> Service<St, WebRequest<In>> for RouteService<St, In> {
272 type Res = WebResponse;
273 type Error = WebError<St, St::Error>;
274
275 async fn call(
276 &self,
277 req: WebRequest<In>,
278 ctx: Ctx<'_, Self, St>,
279 ) -> Result<Self::Res, Self::Error> {
280 Ok(self.handler.call(ctx.st(), req).await)
281 }
282}
283
284pub trait IntoRoutes<St: State, In> {
286 fn routes(self) -> Vec<Route<St, In>>;
288}
289
290impl<St: State, In> IntoRoutes<St, In> for Route<St, In> {
291 fn routes(self) -> Vec<Route<St, In>> {
292 vec![self]
293 }
294}
295
296impl<St: State, In> IntoRoutes<St, In> for Vec<Route<St, In>> {
297 fn routes(self) -> Vec<Route<St, In>> {
298 self
299 }
300}
301
302macro_rules! tuple_routes(
303 {$(#[$meta:meta])* $(($n:tt, $T:ident)),+} => {
304 $(#[$meta])*
305 #[allow(unused_parens)]
306 impl<St: State, U, $($T,)+> IntoRoutes<St, U> for ($($T,)+)
307 where
308 $($T: Into<Route<St, U>> + 'static,)+ {
309 fn routes(self) -> Vec<Route<St, U>> {
310 vec![$(self.$n.into(),)+]
311 }
312 }
313 }
314);
315
316impl<St: State, In, T, const N: usize> IntoRoutes<St, In> for [T; N]
317where
318 T: Into<Route<St, In>>,
319{
320 fn routes(self) -> Vec<Route<St, In>> {
321 let mut routes = Vec::with_capacity(N);
322 for route in self {
323 routes.push(route.into());
324 }
325 routes
326 }
327}
328
329#[allow(clippy::wildcard_imports)]
330#[rustfmt::skip]
331mod m {
332 use variadics_please::all_tuples_enumerated;
333
334 use super::*;
335
336 all_tuples_enumerated!(#[doc(fake_variadic)] tuple_routes, 1, 12, T);
337}
338
339#[cfg(test)]
340mod tests {
341 use crate::http::{Method, StatusCode, header};
342 use crate::time::{Millis, sleep};
343 use crate::web::test::{TestRequest, call_service, init_service, read_body};
344 use crate::web::{self, App, HttpResponse, error, guard};
345 use crate::{ServiceFactory, util::Bytes};
346
347 #[derive(serde::Serialize, PartialEq, Debug)]
348 struct MyObject {
349 name: String,
350 }
351
352 #[crate::rt_test]
353 async fn test_route() {
354 let srv = init_service(
355 App::new()
356 .service(web::resource("/test").route(vec![
357 web::get().to(async || { HttpResponse::Ok() }),
358 web::put().to(async || {
359 Err::<HttpResponse, _>(
360 error::ErrorBadRequest::<_>("err"),
361 )
362 }),
363 web::post().to(async || {
364 sleep(Millis(100)).await;
365 HttpResponse::Created()
366 }),
367 web::patch()
368 .guard(guard::fn_guard(|req|
369 req.headers().contains_key("content-type")
370 ))
371 .to(async || { HttpResponse::Conflict() }),
372 web::delete().to(async || {
373 sleep(Millis(100)).await;
374 Err::<HttpResponse, _>(error::ErrorBadRequest("err"))
375 }),
376 ]))
377 .service(web::resource("/json").route(web::get().to(async || {
378 sleep(Millis(25)).await;
379 web::types::Json(MyObject {
380 name: "test".to_string(),
381 })
382 }))),
383 )
384 .await;
385
386 let req = TestRequest::with_uri("/test")
387 .method(Method::GET)
388 .to_request();
389 let resp = call_service(&srv, req).await;
390 assert_eq!(resp.status(), StatusCode::OK);
391
392 let req = TestRequest::with_uri("/test")
393 .method(Method::POST)
394 .to_request();
395 let resp = call_service(&srv, req).await;
396 assert_eq!(resp.status(), StatusCode::CREATED);
397
398 let req = TestRequest::with_uri("/test")
399 .method(Method::PUT)
400 .to_request();
401 let resp = call_service(&srv, req).await;
402 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
403
404 let req = TestRequest::with_uri("/test")
405 .method(Method::PATCH)
406 .to_request();
407 let resp = call_service(&srv, req).await;
408 assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED);
409
410 let req = TestRequest::with_uri("/test")
411 .method(Method::PATCH)
412 .header(header::CONTENT_TYPE, "text/plain")
413 .to_request();
414 let resp = call_service(&srv, req).await;
415 assert_eq!(resp.status(), StatusCode::CONFLICT);
416
417 let req = TestRequest::with_uri("/test")
418 .method(Method::DELETE)
419 .to_request();
420 let resp = call_service(&srv, req).await;
421 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
422
423 let req = TestRequest::with_uri("/test")
424 .method(Method::HEAD)
425 .to_request();
426 let resp = call_service(&srv, req).await;
427 assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED);
428
429 let req = TestRequest::with_uri("/json").to_request();
430 let resp = call_service(&srv, req).await;
431 assert_eq!(resp.status(), StatusCode::OK);
432
433 let body = read_body(resp).await;
434 assert_eq!(body, Bytes::from_static(b"{\"name\":\"test\"}"));
435
436 let route: web::Route<(), ()> = web::get();
437 let repr = format!("{route:?}");
438 assert!(repr.contains("Route"), "{}", repr);
439 assert!(
440 repr.contains(
441 "handler: HandlerNoState(\"ntex::web::route::Route<()>::new::{{closure}}\")"
442 ),
443 "{}",
444 repr
445 );
446 assert!(repr.contains("methods: [GET]"), "{}", repr);
447 assert!(repr.contains("guards: AllGuard()"), "{}", repr);
448
449 assert!(route.create(&()).await.is_ok());
450
451 let route_service = route.service();
452 let repr = format!("{route_service:?}");
453 assert!(repr.contains("RouteService"));
454 assert!(repr.contains(
455 "handler: HandlerNoState(\"ntex::web::route::Route<()>::new::{{closure}}\")"
456 ));
457 assert!(repr.contains("methods: [GET]"));
458 assert!(repr.contains("guards: AllGuard()"));
459 }
460
461 #[crate::rt_test]
462 async fn test_route_array_and_extractor_error() {
463 let srv = init_service(App::new().service(web::resource("/test/{id}").route([
464 web::get().to(async |p: web::types::Path<u32>| {
465 HttpResponse::Ok().body(format!("{}", p.into_inner()))
466 }),
467 web::post().to(async || HttpResponse::Created()),
468 ])))
469 .await;
470
471 let req = TestRequest::with_uri("/test/10").to_request();
472 let resp = call_service(&srv, req).await;
473 assert_eq!(resp.status(), StatusCode::OK);
474 assert_eq!(read_body(resp).await, Bytes::from_static(b"10"));
475
476 let req = TestRequest::with_uri("/test/abc").to_request();
477 let resp = call_service(&srv, req).await;
478 assert_eq!(resp.status(), StatusCode::NOT_FOUND);
479
480 let req = TestRequest::with_uri("/test/abc")
481 .method(Method::POST)
482 .to_request();
483 let resp = call_service(&srv, req).await;
484 assert_eq!(resp.status(), StatusCode::CREATED);
485 }
486
487 #[test]
488 fn test_route_debug() {
489 let route: web::Route<(), ()> = web::get().to(async || HttpResponse::Ok());
490 let s = format!("{route:?}");
491 assert!(s.contains("HandlerNoState"), "{s}");
492
493 let route =
494 web::Route::<(), ()>::new().to_with_state(async |(): &(), (): ()| HttpResponse::Ok());
495 let s = format!("{route:?}");
496 assert!(s.contains("HandlerSt("), "{s}");
497 }
498}