1#![allow(non_snake_case)]
30
31use std::{borrow::Cow, fmt};
32
33use urly::Url;
34
35use crate::http::{Method, RequestHead, header};
36
37pub trait Guard {
43 fn check(&self, request: &RequestHead) -> bool;
45
46 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
48 f.debug_tuple("Guard").finish()
49 }
50}
51
52pub 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 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
111pub fn Any<F: Guard + 'static>(guard: F) -> AnyGuard {
125 AnyGuard(vec![Box::new(guard)])
126}
127
128#[derive(Default)]
129pub struct AnyGuard(pub Vec<Box<dyn Guard>>);
131
132impl AnyGuard {
133 #[must_use]
134 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 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
167pub fn All<F: Guard + 'static>(guard: F) -> AllGuard {
182 AllGuard(vec![Box::new(guard)])
183}
184
185#[derive(Default)]
186pub struct AllGuard(pub(super) Vec<Box<dyn Guard>>);
188
189impl AllGuard {
190 #[must_use]
191 pub fn and<F: Guard + 'static>(mut self, guard: F) -> Self {
193 self.0.push(Box::new(guard));
194 self
195 }
196
197 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 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
229pub 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 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#[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 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
268 fmt::Debug::fmt(self, f)
269 }
270}
271
272pub fn Get() -> MethodGuard {
274 MethodGuard(Method::GET)
275}
276
277pub fn Post() -> MethodGuard {
279 MethodGuard(Method::POST)
280}
281
282pub fn Put() -> MethodGuard {
284 MethodGuard(Method::PUT)
285}
286
287pub fn Delete() -> MethodGuard {
289 MethodGuard(Method::DELETE)
290}
291
292pub fn Head() -> MethodGuard {
294 MethodGuard(Method::HEAD)
295}
296
297pub fn Options() -> MethodGuard {
299 MethodGuard(Method::OPTIONS)
300}
301
302pub fn Connect() -> MethodGuard {
304 MethodGuard(Method::CONNECT)
305}
306
307pub fn Patch() -> MethodGuard {
309 MethodGuard(Method::PATCH)
310}
311
312pub fn Trace() -> MethodGuard {
314 MethodGuard(Method::TRACE)
315}
316
317pub fn Query() -> MethodGuard {
319 MethodGuard(Method::QUERY)
320}
321
322pub fn Method(method: Method) -> MethodGuard {
324 MethodGuard(method)
325}
326
327pub 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 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
350 fmt::Debug::fmt(self, f)
351 }
352}
353
354pub 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 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 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 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 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}