Skip to main content

ntex/web/
guard.rs

1//! Route match guards.
2//!
3//! Guards are one of the ways how ntex web router chooses a
4//! handler service. In essence it is just a function that accepts a
5//! reference to a `RequestHead` instance and returns a boolean.
6//! It is possible to add guards to *scopes*, *resources*
7//! and *routes*. ntex provide several guards by default, like various
8//! http methods, header, etc. To become a guard, type must implement `Guard`
9//! trait. Simple functions can be used as guards as well, see [`fn_guard`].
10//!
11//! Multiple guards on the same route, resource, or scope must all match.
12//!
13//! Guards can not modify the request object. But it is possible
14//! to store extra attributes on a request by using the `Extensions` container.
15//! Extensions containers are available via the `RequestHead::extensions()` method.
16//!
17//! ```rust
18//! use ntex::web::{self, guard, App, HttpResponse};
19//!
20//! fn main() {
21//!     App::default().service(web::resource("/index.html").route(
22//!         web::route()
23//!              .guard(guard::Post())
24//!              .guard(guard::fn_guard(|head| head.headers.contains_key("x-api-key")))
25//!              .to(async || { HttpResponse::Ok() }))
26//!     );
27//! }
28//! ```
29#![allow(non_snake_case)]
30
31use std::{borrow::Cow, fmt};
32
33use urly::Url;
34
35use crate::http::{Method, RequestHead, header};
36
37/// Trait defines resource guards. Guards are used for route selection.
38///
39/// Guards can not modify the request object. But it is possible
40/// to store extra attributes on a request by using the `Extensions` container.
41/// Extensions containers are available via the `RequestHead::extensions()` method.
42pub trait Guard {
43    /// Check if request matches predicate
44    fn check(&self, request: &RequestHead) -> bool;
45
46    /// Debug format
47    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
48        f.debug_tuple("Guard").finish()
49    }
50}
51
52/// Create guard object for supplied function.
53///
54/// ```rust
55/// use ntex::web::{self, guard, App, HttpResponse};
56///
57/// fn main() {
58///     App::default().service(web::resource("/index.html").route(
59///         web::route()
60///             .guard(
61///                 guard::fn_guard(
62///                     |req| req.headers()
63///                              .contains_key("content-type")))
64///             .to(async || { HttpResponse::MethodNotAllowed() }))
65///     );
66/// }
67/// ```
68pub fn fn_guard<F>(f: F) -> impl Guard
69where
70    F: Fn(&RequestHead) -> bool,
71{
72    FnGuard(f)
73}
74
75struct FnGuard<F: Fn(&RequestHead) -> bool>(F);
76
77impl<F> Guard for FnGuard<F>
78where
79    F: Fn(&RequestHead) -> bool,
80{
81    fn check(&self, head: &RequestHead) -> bool {
82        (self.0)(head)
83    }
84
85    /// Debug format
86    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
87        f.debug_tuple("FnGuard")
88            .field(&std::any::type_name::<F>())
89            .finish()
90    }
91}
92
93impl<F> fmt::Debug for FnGuard<F>
94where
95    F: Fn(&RequestHead) -> bool,
96{
97    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
98        Guard::fmt(self, f)
99    }
100}
101
102impl<F> Guard for F
103where
104    F: Fn(&RequestHead) -> bool,
105{
106    fn check(&self, head: &RequestHead) -> bool {
107        (self)(head)
108    }
109}
110
111/// Return guard that matches if any of supplied guards.
112///
113/// ```rust
114/// use ntex::web::{self, guard, App, HttpResponse};
115///
116/// fn main() {
117///     App::default().service(web::resource("/index.html").route(
118///         web::route()
119///              .guard(guard::Any(guard::Get()).or(guard::Post()))
120///              .to(async || { HttpResponse::MethodNotAllowed() }))
121///     );
122/// }
123/// ```
124pub fn Any<F: Guard + 'static>(guard: F) -> AnyGuard {
125    AnyGuard(vec![Box::new(guard)])
126}
127
128#[derive(Default)]
129/// Matches any of supplied guards match.
130pub struct AnyGuard(pub Vec<Box<dyn Guard>>);
131
132impl AnyGuard {
133    #[must_use]
134    /// Add guard to a list of guards to check.
135    pub fn or<F: Guard + 'static>(mut self, guard: F) -> Self {
136        self.0.push(Box::new(guard));
137        self
138    }
139}
140
141impl Guard for AnyGuard {
142    fn check(&self, req: &RequestHead) -> bool {
143        for p in &self.0 {
144            if p.check(req) {
145                return true;
146            }
147        }
148        false
149    }
150
151    /// Debug format
152    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
153        write!(f, "AnyGuard(")?;
154        self.0.iter().for_each(|t| {
155            let _ = t.fmt(f);
156        });
157        write!(f, ")")
158    }
159}
160
161impl fmt::Debug for AnyGuard {
162    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
163        Guard::fmt(self, f)
164    }
165}
166
167/// Return guard that matches if all of the supplied guards.
168///
169/// ```rust
170/// use ntex::web::{self, guard, App, HttpResponse};
171///
172/// fn main() {
173///     App::default().service(web::resource("/index.html").route(
174///         web::route()
175///             .guard(
176///                 guard::All(guard::Get()).and(guard::Header("content-type", "text/plain")))
177///             .to(async || { HttpResponse::MethodNotAllowed() }))
178///     );
179/// }
180/// ```
181pub fn All<F: Guard + 'static>(guard: F) -> AllGuard {
182    AllGuard(vec![Box::new(guard)])
183}
184
185#[derive(Default)]
186/// Matches all of supplied guards..
187pub struct AllGuard(pub(super) Vec<Box<dyn Guard>>);
188
189impl AllGuard {
190    #[must_use]
191    /// Add new guard to the list of guards to check.
192    pub fn and<F: Guard + 'static>(mut self, guard: F) -> Self {
193        self.0.push(Box::new(guard));
194        self
195    }
196
197    /// Add guard to a list of guards to check.
198    pub fn add<F: Guard + 'static>(&mut self, guard: F) {
199        self.0.push(Box::new(guard));
200    }
201}
202
203impl Guard for AllGuard {
204    fn check(&self, request: &RequestHead) -> bool {
205        for p in &self.0 {
206            if !p.check(request) {
207                return false;
208            }
209        }
210        true
211    }
212
213    /// Debug format
214    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
215        write!(f, "AllGuard(")?;
216        self.0.iter().for_each(|t| {
217            let _ = t.fmt(f);
218        });
219        write!(f, ")")
220    }
221}
222
223impl fmt::Debug for AllGuard {
224    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
225        Guard::fmt(self, f)
226    }
227}
228
229/// Return guard that matches if supplied guard does not match.
230pub fn Not<F: Guard + 'static>(guard: F) -> NotGuard {
231    NotGuard(Box::new(guard))
232}
233
234#[doc(hidden)]
235pub struct NotGuard(Box<dyn Guard>);
236
237impl Guard for NotGuard {
238    fn check(&self, request: &RequestHead) -> bool {
239        !self.0.check(request)
240    }
241
242    /// Debug format
243    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
244        write!(f, "NotGuard(")?;
245        self.0.fmt(f)?;
246        write!(f, ")")
247    }
248}
249
250impl fmt::Debug for NotGuard {
251    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
252        Guard::fmt(self, f)
253    }
254}
255
256/// Http method guard
257#[doc(hidden)]
258#[derive(Debug)]
259pub struct MethodGuard(Method);
260
261impl Guard for MethodGuard {
262    fn check(&self, request: &RequestHead) -> bool {
263        request.method == self.0
264    }
265
266    /// Debug format
267    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
268        fmt::Debug::fmt(self, f)
269    }
270}
271
272/// Guard to match *GET* http method
273pub fn Get() -> MethodGuard {
274    MethodGuard(Method::GET)
275}
276
277/// Predicate to match *POST* http method
278pub fn Post() -> MethodGuard {
279    MethodGuard(Method::POST)
280}
281
282/// Predicate to match *PUT* http method
283pub fn Put() -> MethodGuard {
284    MethodGuard(Method::PUT)
285}
286
287/// Predicate to match *DELETE* http method
288pub fn Delete() -> MethodGuard {
289    MethodGuard(Method::DELETE)
290}
291
292/// Predicate to match *HEAD* http method
293pub fn Head() -> MethodGuard {
294    MethodGuard(Method::HEAD)
295}
296
297/// Predicate to match *OPTIONS* http method
298pub fn Options() -> MethodGuard {
299    MethodGuard(Method::OPTIONS)
300}
301
302/// Predicate to match *CONNECT* http method
303pub fn Connect() -> MethodGuard {
304    MethodGuard(Method::CONNECT)
305}
306
307/// Predicate to match *PATCH* http method
308pub fn Patch() -> MethodGuard {
309    MethodGuard(Method::PATCH)
310}
311
312/// Predicate to match *TRACE* http method
313pub fn Trace() -> MethodGuard {
314    MethodGuard(Method::TRACE)
315}
316
317/// Predicate to match *QUERY* http method
318pub fn Query() -> MethodGuard {
319    MethodGuard(Method::QUERY)
320}
321
322/// Predicate to match specified http method
323pub fn Method(method: Method) -> MethodGuard {
324    MethodGuard(method)
325}
326
327/// Return predicate that matches if request contains specified header and
328/// value.
329pub fn Header(name: &'static str, value: &'static str) -> HeaderGuard {
330    HeaderGuard(
331        header::HeaderName::try_from(name).unwrap(),
332        header::HeaderValue::from_static(value),
333    )
334}
335
336#[doc(hidden)]
337#[derive(Debug)]
338pub struct HeaderGuard(header::HeaderName, header::HeaderValue);
339
340impl Guard for HeaderGuard {
341    fn check(&self, req: &RequestHead) -> bool {
342        if let Some(val) = req.headers.get(&self.0) {
343            return val == self.1;
344        }
345        false
346    }
347
348    /// Debug format
349    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
350        fmt::Debug::fmt(self, f)
351    }
352}
353
354/// Return predicate that matches if request contains specified Host name.
355///
356/// ```rust
357/// use ntex::web::{self, guard::Host, App, HttpResponse};
358///
359/// fn main() {
360///     App::default().service(
361///         web::resource("/index.html")
362///             .guard(Host("www.rust-lang.org"))
363///             .to(async || { HttpResponse::MethodNotAllowed() })
364///     );
365/// }
366/// ```
367pub fn Host<H: AsRef<str>>(host: H) -> HostGuard {
368    HostGuard(host.as_ref().to_string(), None)
369}
370
371fn get_host_uri(req: &RequestHead) -> Option<Cow<'_, Url>> {
372    // the authority of an absolute-form target takes precedence over `Host`,
373    // see RFC 9112 section 3.2.2
374    if req.uri.authority().is_some() {
375        return Some(Cow::Borrowed(&req.uri));
376    }
377    let host = req.headers.get(header::HOST)?.to_str().ok()?;
378    if host.contains("://") {
379        Url::try_from(host).ok().map(Cow::Owned)
380    } else {
381        Url::parse(host).ok().map(Cow::Owned)
382    }
383}
384
385#[doc(hidden)]
386#[derive(Debug)]
387pub struct HostGuard(String, Option<String>);
388
389impl HostGuard {
390    #[must_use]
391    /// Set request scheme to match.
392    ///
393    /// The scheme is taken from the request uri if it is in absolute-form, or
394    /// from the `Host` header value otherwise. A `Host` header normally contains no
395    /// scheme; in that case the scheme is not checked and the guard matches on
396    /// the host name alone. Do not rely on this check for security decisions.
397    pub fn scheme<H: AsRef<str>>(mut self, scheme: H) -> HostGuard {
398        self.1 = Some(scheme.as_ref().to_string());
399        self
400    }
401}
402
403impl Guard for HostGuard {
404    fn check(&self, req: &RequestHead) -> bool {
405        let Some(req_host_uri) = get_host_uri(req) else {
406            return false;
407        };
408
409        if let Some(uri_host) = req_host_uri.host() {
410            if !self.0.eq_ignore_ascii_case(uri_host) {
411                return false;
412            }
413        } else {
414            return false;
415        }
416
417        if let Some(ref scheme) = self.1
418            && let Some(req_host_uri_scheme) = req_host_uri.scheme_str()
419        {
420            return scheme.eq_ignore_ascii_case(req_host_uri_scheme);
421        }
422
423        true
424    }
425
426    /// Debug format
427    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
428        fmt::Debug::fmt(self, f)
429    }
430}
431
432#[cfg(test)]
433mod tests {
434    use super::*;
435    use crate::web::test::TestRequest;
436
437    #[test]
438    fn test_header() {
439        let req = TestRequest::with_header(header::TRANSFER_ENCODING, "chunked").to_http_request();
440
441        let pred = Header("transfer-encoding", "chunked");
442        assert!(pred.check(req.head()));
443
444        let pred = Header("transfer-encoding", "other");
445        assert!(!pred.check(req.head()));
446
447        let pred = Header("content-type", "other");
448        assert!(!pred.check(req.head()));
449
450        assert!(format!("{pred:?}").contains("Header"));
451    }
452
453    #[test]
454    fn test_host() {
455        let req = TestRequest::default()
456            .header(
457                header::HOST,
458                header::HeaderValue::from_static("www.rust-lang.org"),
459            )
460            .to_http_request();
461
462        let pred = Host("www.rust-lang.org");
463        assert!(pred.check(req.head()));
464
465        let pred = Host("www.rust-lang.org").scheme("https");
466        assert!(pred.check(req.head()));
467
468        let pred = Host("blog.rust-lang.org");
469        assert!(!pred.check(req.head()));
470
471        let pred = Host("blog.rust-lang.org").scheme("https");
472        assert!(!pred.check(req.head()));
473
474        let pred = Host("crates.io");
475        assert!(!pred.check(req.head()));
476
477        let pred = Host("localhost");
478        assert!(!pred.check(req.head()));
479    }
480
481    #[test]
482    fn test_host_scheme() {
483        let req = TestRequest::default()
484            .header(
485                header::HOST,
486                header::HeaderValue::from_static("https://www.rust-lang.org"),
487            )
488            .to_http_request();
489
490        let pred = Host("www.rust-lang.org").scheme("https");
491        assert!(pred.check(req.head()));
492
493        let pred = Host("www.rust-lang.org");
494        assert!(pred.check(req.head()));
495
496        let pred = Host("www.rust-lang.org").scheme("http");
497        assert!(!pred.check(req.head()));
498
499        let pred = Host("blog.rust-lang.org");
500        assert!(!pred.check(req.head()));
501
502        let pred = Host("blog.rust-lang.org").scheme("https");
503        assert!(!pred.check(req.head()));
504
505        let pred = Host("crates.io").scheme("https");
506        assert!(!pred.check(req.head()));
507
508        let pred = Host("localhost");
509        assert!(!pred.check(req.head()));
510
511        assert!(format!("{pred:?}").contains("Host"));
512    }
513
514    #[test]
515    fn test_host_without_header() {
516        let req = TestRequest::default()
517            .uri("//www.rust-lang.org")
518            .to_http_request();
519
520        let pred = Host("www.rust-lang.org");
521        assert!(pred.check(req.head()));
522
523        let pred = Host("www.rust-lang.org").scheme("https");
524        assert!(pred.check(req.head()));
525
526        let pred = Host("blog.rust-lang.org");
527        assert!(!pred.check(req.head()));
528
529        let pred = Host("blog.rust-lang.org").scheme("https");
530        assert!(!pred.check(req.head()));
531
532        let pred = Host("crates.io");
533        assert!(!pred.check(req.head()));
534
535        let pred = Host("localhost");
536        assert!(!pred.check(req.head()));
537    }
538
539    #[test]
540    fn test_host_absolute_form() {
541        // the target authority takes precedence over `Host`
542        let req = TestRequest::with_uri("http://www.rust-lang.org/p")
543            .header(header::HOST, "crates.io")
544            .to_http_request();
545        assert!(Host("www.rust-lang.org").scheme("http").check(req.head()));
546        assert!(!Host("www.rust-lang.org").scheme("https").check(req.head()));
547        assert!(!Host("crates.io").check(req.head()));
548
549        let req = TestRequest::default()
550            .header(header::HOST, "www.rust-lang.org:8080")
551            .to_http_request();
552        assert!(Host("www.rust-lang.org").check(req.head()));
553
554        let req = TestRequest::default()
555            .header(header::HOST, "www.rust-lang.org/p")
556            .to_http_request();
557        assert!(!Host("www.rust-lang.org").check(req.head()));
558    }
559
560    #[test]
561    fn test_methods() {
562        let req = TestRequest::default().to_http_request();
563        let req2 = TestRequest::default()
564            .method(Method::POST)
565            .to_http_request();
566
567        assert!(Get().check(req.head()));
568        assert!(!Get().check(req2.head()));
569        assert!(Post().check(req2.head()));
570        assert!(!Post().check(req.head()));
571
572        let r = TestRequest::default().method(Method::PUT).to_http_request();
573        assert!(Put().check(r.head()));
574        assert!(!Put().check(req.head()));
575
576        let r = TestRequest::default()
577            .method(Method::DELETE)
578            .to_http_request();
579        assert!(Delete().check(r.head()));
580        assert!(!Delete().check(req.head()));
581
582        let r = TestRequest::default()
583            .method(Method::HEAD)
584            .to_http_request();
585        assert!(Head().check(r.head()));
586        assert!(!Head().check(req.head()));
587
588        let r = TestRequest::default()
589            .method(Method::OPTIONS)
590            .to_http_request();
591        assert!(Options().check(r.head()));
592        assert!(!Options().check(req.head()));
593
594        let r = TestRequest::default()
595            .method(Method::CONNECT)
596            .to_http_request();
597        assert!(Connect().check(r.head()));
598        assert!(!Connect().check(req.head()));
599
600        let r = TestRequest::default()
601            .method(Method::PATCH)
602            .to_http_request();
603        assert!(Patch().check(r.head()));
604        assert!(!Patch().check(req.head()));
605
606        let r = TestRequest::default()
607            .method(Method::TRACE)
608            .to_http_request();
609        assert!(Trace().check(r.head()));
610        assert!(!Trace().check(req.head()));
611        assert!(format!("{:?}", Trace()).contains("MethodGuard(TRACE)"));
612
613        let r = TestRequest::default()
614            .method(Method::QUERY)
615            .to_http_request();
616        assert!(Query().check(r.head()));
617        assert!(!Query().check(req.head()));
618    }
619
620    #[test]
621    fn test_preds() {
622        let r = TestRequest::default()
623            .method(Method::TRACE)
624            .to_http_request();
625
626        assert!(Not(Get()).check(r.head()));
627        assert!(!Not(Trace()).check(r.head()));
628        assert!(format!("{:?}", Not(Get())).contains("NotGuard"));
629
630        assert!(All(Trace()).and(Trace()).check(r.head()));
631        assert!(!All(Get()).and(Trace()).check(r.head()));
632        assert!(format!("{:?}", All(Get())).contains("AllGuard"));
633
634        assert!(Any(Get()).or(Trace()).check(r.head()));
635        assert!(!Any(Get()).or(Get()).check(r.head()));
636        assert!(format!("{:?}", Any(Get())).contains("AnyGuard"));
637    }
638
639    #[test]
640    fn test_fn_guard() {
641        let req = TestRequest::with_header(header::CONTENT_TYPE, "text/plain").to_http_request();
642
643        let g = fn_guard(|req| req.headers().contains_key("content-type"));
644        assert!(g.check(req.head()));
645        let g = FnGuard(|req| req.headers().contains_key("content-type"));
646        assert!(format!("{g:?}").contains("FnGuard"));
647
648        let g = |req: &RequestHead| req.headers().contains_key("content-type");
649        assert!(g.check(req.head()));
650    }
651
652    #[test]
653    fn test_guard_debug() {
654        struct Custom;
655
656        impl Guard for Custom {
657            fn check(&self, _: &RequestHead) -> bool {
658                true
659            }
660        }
661
662        let guard = Any(Header("content-type", "text/plain"))
663            .or(Host("localhost"))
664            .or(Custom);
665        let s = format!("{guard:?}");
666        assert!(s.contains("HeaderGuard"), "{s}");
667        assert!(s.contains("HostGuard"), "{s}");
668        assert!(s.contains("Guard"), "{s}");
669    }
670
671    #[test]
672    fn test_host_without_host_name() {
673        let req = TestRequest::default()
674            .header(header::HOST, header::HeaderValue::from_static("bad host"))
675            .to_http_request();
676        assert!(!Host("localhost").check(req.head()));
677
678        let req = TestRequest::default()
679            .header(header::HOST, header::HeaderValue::from_static("/path"))
680            .to_http_request();
681        assert!(!Host("localhost").check(req.head()));
682    }
683}