1use crate::{Error, HttpRequest};
38use bytes::Bytes;
39use serde::de::DeserializeOwned;
40use std::ops::Deref;
41use std::sync::Arc;
42
43pub trait FromRequest: Sized {
45 fn from_request(request: &HttpRequest) -> Result<Self, Error>;
47}
48
49#[derive(Debug)]
93pub struct State<T: Send + Sync + 'static>(pub Arc<T>);
94
95impl<T: Send + Sync + 'static> State<T> {
96 #[inline]
98 pub fn new(value: Arc<T>) -> Self {
99 Self(value)
100 }
101
102 #[inline]
104 pub fn into_inner(self) -> Arc<T> {
105 self.0
106 }
107}
108
109impl<T: Send + Sync + 'static> Clone for State<T> {
110 #[inline]
111 fn clone(&self) -> Self {
112 Self(Arc::clone(&self.0))
113 }
114}
115
116impl<T: Send + Sync + 'static> Deref for State<T> {
117 type Target = T;
118
119 #[inline]
120 fn deref(&self) -> &Self::Target {
121 &self.0
122 }
123}
124
125impl<T: Send + Sync + 'static> AsRef<T> for State<T> {
126 #[inline]
127 fn as_ref(&self) -> &T {
128 &self.0
129 }
130}
131
132impl<T: Send + Sync + 'static> FromRequest for State<T> {
133 #[inline]
140 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
141 request.extensions.get_arc::<T>().map(State).ok_or_else(|| {
142 Error::ProviderNotFound(format!(
143 "State<{}> not found in request extensions. \
144 Did you forget to register it with `app.with_state()`?",
145 std::any::type_name::<T>()
146 ))
147 })
148 }
149}
150
151pub trait FromRequestNamed: Sized {
153 fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error>;
155}
156
157#[derive(Debug, Clone)]
174pub struct Body<T>(pub T);
175
176impl<T> Body<T> {
177 pub fn new(value: T) -> Self {
179 Self(value)
180 }
181
182 pub fn into_inner(self) -> T {
184 self.0
185 }
186}
187
188impl<T> Deref for Body<T> {
189 type Target = T;
190
191 fn deref(&self) -> &Self::Target {
192 &self.0
193 }
194}
195
196impl<T: DeserializeOwned> FromRequest for Body<T> {
197 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
198 let value: T = request.json()?;
199 Ok(Body(value))
200 }
201}
202
203#[derive(Debug, Clone)]
221pub struct Query<T>(pub T);
222
223impl<T> Query<T> {
224 pub fn new(value: T) -> Self {
226 Self(value)
227 }
228
229 pub fn into_inner(self) -> T {
231 self.0
232 }
233}
234
235impl<T> Deref for Query<T> {
236 type Target = T;
237
238 fn deref(&self) -> &Self::Target {
239 &self.0
240 }
241}
242
243impl<T: DeserializeOwned> FromRequest for Query<T> {
244 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
245 let value: T = serde_urlencoded::from_str(request.query_string().unwrap_or(""))
249 .map_err(|e| Error::Validation(format!("Invalid query parameters: {}", e)))?;
250
251 Ok(Query(value))
252 }
253}
254
255#[derive(Debug, Clone)]
267pub struct Path<T>(pub T);
268
269impl<T> Path<T> {
270 pub fn new(value: T) -> Self {
272 Self(value)
273 }
274
275 pub fn into_inner(self) -> T {
277 self.0
278 }
279}
280
281impl<T> Deref for Path<T> {
282 type Target = T;
283
284 fn deref(&self) -> &Self::Target {
285 &self.0
286 }
287}
288
289impl<T: std::str::FromStr> FromRequestNamed for Path<T>
290where
291 T::Err: std::fmt::Display,
292{
293 fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error> {
294 let value_str = request
295 .param(name)
296 .ok_or_else(|| Error::Validation(format!("Missing path parameter: {}", name)))?;
297
298 let value: T = value_str.parse().map_err(|e: T::Err| {
299 Error::Validation(format!("Invalid path parameter '{}': {}", name, e))
300 })?;
301
302 Ok(Path(value))
303 }
304}
305
306#[derive(Debug, Clone)]
323pub struct PathParams<T>(pub T);
324
325impl<T> PathParams<T> {
326 pub fn new(value: T) -> Self {
328 Self(value)
329 }
330
331 pub fn into_inner(self) -> T {
333 self.0
334 }
335}
336
337impl<T> Deref for PathParams<T> {
338 type Target = T;
339
340 fn deref(&self) -> &Self::Target {
341 &self.0
342 }
343}
344
345impl<T: DeserializeOwned> FromRequest for PathParams<T> {
346 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
347 let pairs: Vec<(&str, &str)> = request
350 .path_params
351 .iter()
352 .filter_map(|(k, v)| std::str::from_utf8(v).ok().map(|v| (*k, v)))
353 .collect();
354 let params_string = serde_urlencoded::to_string(&pairs)
355 .map_err(|e| Error::Validation(format!("Invalid path parameters: {}", e)))?;
356
357 let value: T = serde_urlencoded::from_str(¶ms_string)
358 .map_err(|e| Error::Validation(format!("Invalid path parameters: {}", e)))?;
359
360 Ok(PathParams(value))
361 }
362}
363
364#[derive(Debug, Clone)]
378pub struct Header {
379 name: String,
380 value: String,
381}
382
383impl Header {
384 pub fn new(name: impl Into<String>, value: impl Into<String>) -> Self {
386 Self {
387 name: name.into(),
388 value: value.into(),
389 }
390 }
391
392 pub fn name(&self) -> &str {
394 &self.name
395 }
396
397 pub fn value(&self) -> &str {
399 &self.value
400 }
401
402 pub fn into_value(self) -> String {
404 self.value
405 }
406
407 pub fn optional(request: &HttpRequest, name: &str) -> Option<Self> {
409 request.headers.get(name).map(|v| Header::new(name, v))
412 }
413}
414
415impl FromRequestNamed for Header {
416 fn from_request(request: &HttpRequest, name: &str) -> Result<Self, Error> {
417 let value = request
418 .headers
419 .get(name)
420 .ok_or_else(|| Error::Validation(format!("Missing header: {}", name)))?;
421
422 Ok(Header::new(name, value))
423 }
424}
425
426impl Deref for Header {
427 type Target = str;
428
429 fn deref(&self) -> &Self::Target {
430 &self.value
431 }
432}
433
434#[derive(Debug, Clone)]
438pub struct Headers(pub std::collections::HashMap<String, String>);
439
440impl Headers {
441 pub fn get(&self, name: &str) -> Option<&String> {
443 if let Some(value) = self.0.get(name) {
446 return Some(value);
447 }
448 self.0
455 .iter()
456 .find(|(k, _)| k.eq_ignore_ascii_case(name))
457 .map(|(_, v)| v)
458 }
459
460 pub fn contains(&self, name: &str) -> bool {
462 self.get(name).is_some()
463 }
464
465 pub fn iter(&self) -> impl Iterator<Item = (&String, &String)> {
467 self.0.iter()
468 }
469}
470
471impl FromRequest for Headers {
472 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
473 Ok(Headers(request.headers.clone().into()))
475 }
476}
477
478impl Deref for Headers {
479 type Target = std::collections::HashMap<String, String>;
480
481 fn deref(&self) -> &Self::Target {
482 &self.0
483 }
484}
485
486#[derive(Debug, Clone)]
497pub struct RawBody(pub Bytes);
498
499impl RawBody {
500 pub fn new(data: impl Into<Bytes>) -> Self {
502 Self(data.into())
503 }
504
505 pub fn len(&self) -> usize {
507 self.0.len()
508 }
509
510 pub fn is_empty(&self) -> bool {
512 self.0.is_empty()
513 }
514
515 pub fn to_string_lossy(&self) -> String {
517 String::from_utf8_lossy(&self.0).to_string()
518 }
519
520 pub fn to_string(&self) -> Result<String, std::string::FromUtf8Error> {
522 String::from_utf8(self.0.to_vec())
523 }
524
525 pub fn into_inner(self) -> Bytes {
527 self.0
528 }
529}
530
531impl FromRequest for RawBody {
532 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
533 Ok(RawBody(request.body.clone()))
534 }
535}
536
537impl Deref for RawBody {
538 type Target = [u8];
539
540 fn deref(&self) -> &Self::Target {
541 &self.0
542 }
543}
544
545#[derive(Debug, Clone)]
561pub struct Form<T>(pub T);
562
563impl<T> Form<T> {
564 pub fn new(value: T) -> Self {
566 Self(value)
567 }
568
569 pub fn into_inner(self) -> T {
571 self.0
572 }
573}
574
575impl<T> Deref for Form<T> {
576 type Target = T;
577
578 fn deref(&self) -> &Self::Target {
579 &self.0
580 }
581}
582
583impl<T: DeserializeOwned> FromRequest for Form<T> {
584 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
585 let value: T = request.form()?;
586 Ok(Form(value))
587 }
588}
589
590#[derive(Debug, Clone)]
594pub struct ContentType(pub String);
595
596impl ContentType {
597 pub fn is_json(&self) -> bool {
599 self.0.contains("application/json")
600 }
601
602 pub fn is_form(&self) -> bool {
604 self.0.contains("application/x-www-form-urlencoded")
605 }
606
607 pub fn is_multipart(&self) -> bool {
609 self.0.contains("multipart/form-data")
610 }
611
612 pub fn into_inner(self) -> String {
614 self.0
615 }
616}
617
618impl FromRequest for ContentType {
619 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
620 let value = request
621 .headers
622 .get("content-type")
623 .map(str::to_owned)
624 .unwrap_or_default();
625
626 Ok(ContentType(value))
627 }
628}
629
630impl Deref for ContentType {
631 type Target = str;
632
633 fn deref(&self) -> &Self::Target {
634 &self.0
635 }
636}
637
638#[derive(Debug, Clone)]
646pub struct MethodExtractor(pub crate::Method);
647
648impl MethodExtractor {
649 pub fn is_get(&self) -> bool {
651 self.0 == "GET"
652 }
653
654 pub fn is_post(&self) -> bool {
656 self.0 == "POST"
657 }
658
659 pub fn is_put(&self) -> bool {
661 self.0 == "PUT"
662 }
663
664 pub fn is_delete(&self) -> bool {
666 self.0 == "DELETE"
667 }
668
669 pub fn is_patch(&self) -> bool {
671 self.0 == "PATCH"
672 }
673}
674
675impl FromRequest for MethodExtractor {
676 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
677 Ok(MethodExtractor(request.method.clone()))
678 }
679}
680
681impl Deref for MethodExtractor {
682 type Target = crate::Method;
683
684 fn deref(&self) -> &Self::Target {
685 &self.0
686 }
687}
688
689impl FromRequest for HttpRequest {
692 fn from_request(request: &HttpRequest) -> Result<Self, Error> {
693 Ok(request.clone())
694 }
695}
696
697#[macro_export]
709macro_rules! body {
710 ($request:expr, $type:ty) => {
711 <$crate::extractors::Body<$type> as $crate::extractors::FromRequest>::from_request(
712 &$request,
713 )
714 .map(|b| b.into_inner())
715 };
716}
717
718#[macro_export]
726macro_rules! query {
727 ($request:expr, $type:ty) => {
728 <$crate::extractors::Query<$type> as $crate::extractors::FromRequest>::from_request(
729 &$request,
730 )
731 .map(|q| q.into_inner())
732 };
733}
734
735#[macro_export]
743macro_rules! path {
744 ($request:expr, $name:expr, $type:ty) => {
745 <$crate::extractors::Path<$type> as $crate::extractors::FromRequestNamed>::from_request(
746 &$request, $name,
747 )
748 .map(|p| p.into_inner())
749 };
750}
751
752#[macro_export]
760macro_rules! header {
761 ($request:expr, $name:expr) => {
762 <$crate::extractors::Header as $crate::extractors::FromRequestNamed>::from_request(
763 &$request, $name,
764 )
765 .map(|h| h.into_value())
766 };
767}
768
769#[cfg(test)]
770mod tests {
771 use super::*;
772 use serde::Deserialize;
773
774 fn create_request() -> HttpRequest {
775 let mut req = HttpRequest::new("GET", "/users/123?page=1&limit=10");
776 req.push_param("id", "123");
777 req.headers
778 .insert("Authorization", "Bearer token123".to_string());
779 req.headers
780 .insert("Content-Type", "application/json".to_string());
781 req
782 }
783
784 #[test]
785 fn test_path_extraction() {
786 let request = create_request();
787 let id: Path<u32> = Path::from_request(&request, "id").unwrap();
788 assert_eq!(*id, 123);
789 }
790
791 #[test]
792 fn test_path_missing() {
793 let request = create_request();
794 let result: Result<Path<u32>, _> = Path::from_request(&request, "missing");
795 assert!(result.is_err());
796 }
797
798 #[test]
799 fn test_header_extraction() {
800 let request = create_request();
801 let auth: Header = Header::from_request(&request, "Authorization").unwrap();
802 assert_eq!(auth.value(), "Bearer token123");
803 }
804
805 #[test]
806 fn test_header_optional() {
807 let request = create_request();
808
809 let auth = Header::optional(&request, "Authorization");
810 assert!(auth.is_some());
811
812 let missing = Header::optional(&request, "X-Missing");
813 assert!(missing.is_none());
814 }
815
816 #[test]
817 fn test_headers_extraction() {
818 let request = create_request();
819 let headers: Headers = Headers::from_request(&request).unwrap();
820
821 assert!(headers.contains("Authorization"));
822 assert!(headers.contains("Content-Type"));
823 assert!(!headers.contains("X-Missing"));
824 }
825
826 #[test]
827 fn test_query_extraction() {
828 let request = create_request();
829
830 #[derive(Debug, Deserialize, PartialEq)]
831 struct Pagination {
832 page: u32,
833 limit: u32,
834 }
835
836 let query: Query<Pagination> = Query::from_request(&request).unwrap();
837 assert_eq!(query.page, 1);
838 assert_eq!(query.limit, 10);
839 }
840
841 #[test]
842 fn test_query_extraction_with_ampersand_and_equals_in_value() {
843 let request = HttpRequest::new("GET", "/items?note=1%26b%3D2&name=a%3Db%26c");
846
847 #[derive(Debug, Deserialize, PartialEq)]
848 struct Filters {
849 note: String,
850 name: String,
851 }
852
853 let query: Query<Filters> = Query::from_request(&request).unwrap();
854 assert_eq!(query.note, "1&b=2");
855 assert_eq!(query.name, "a=b&c");
856 }
857
858 #[test]
859 fn test_path_params_extraction_with_ampersand_and_equals_in_value() {
860 let mut request = HttpRequest::new("GET", "/items/x");
863 request.push_param("slug", Bytes::from_static(b"a&b=c"));
864
865 #[derive(Debug, Deserialize, PartialEq)]
866 struct Params {
867 slug: String,
868 }
869
870 let params: PathParams<Params> = PathParams::from_request(&request).unwrap();
871 assert_eq!(params.slug, "a&b=c");
872 }
873
874 #[test]
875 fn test_body_extraction() {
876 let mut request = create_request();
877 request.body = Bytes::from(
878 serde_json::to_vec(&serde_json::json!({
879 "name": "Test",
880 "email": "test@example.com"
881 }))
882 .unwrap(),
883 );
884
885 #[derive(Debug, Deserialize)]
886 struct CreateUser {
887 name: String,
888 email: String,
889 }
890
891 let body: Body<CreateUser> = Body::from_request(&request).unwrap();
892 assert_eq!(body.name, "Test");
893 assert_eq!(body.email, "test@example.com");
894 }
895
896 #[test]
897 fn test_raw_body() {
898 let mut request = create_request();
899 request.body = Bytes::from_static(b"raw content");
900
901 let raw: RawBody = RawBody::from_request(&request).unwrap();
902 assert_eq!(raw.len(), 11);
903 assert_eq!(raw.to_string_lossy(), "raw content");
904 }
905
906 #[test]
907 fn test_content_type() {
908 let request = create_request();
909 let ct: ContentType = ContentType::from_request(&request).unwrap();
910
911 assert!(ct.is_json());
912 assert!(!ct.is_form());
913 assert!(!ct.is_multipart());
914 }
915
916 #[test]
917 fn test_method() {
918 let request = create_request();
919 let method: MethodExtractor = MethodExtractor::from_request(&request).unwrap();
920
921 assert!(method.is_get());
922 assert!(!method.is_post());
923 }
924
925 #[test]
926 fn test_request_extraction() {
927 let request = create_request();
928 let extracted: HttpRequest = HttpRequest::from_request(&request).unwrap();
929
930 assert_eq!(extracted.method, request.method);
931 assert_eq!(extracted.path, request.path);
932 }
933}