Skip to main content

rama_http/service/web/endpoint/extract/body/
form.rs

1use rama_core::bytes::Bytes;
2
3use super::BytesRejection;
4use crate::body::util::BodyExt;
5use crate::service::web::extract::FromRequest;
6use crate::utils::macros::{composite_http_rejection, define_http_rejection};
7use crate::{Method, Request};
8
9pub use crate::service::web::endpoint::response::Form;
10
11define_http_rejection! {
12    #[status = UNSUPPORTED_MEDIA_TYPE]
13    #[body = "Form requests must have `Content-Type: application/x-www-form-urlencoded`"]
14    /// Rejection type for [`Form`]
15    /// used if the `Content-Type` header is missing
16    /// or its value is not `application/x-www-form-urlencoded`.
17    pub struct InvalidFormContentType;
18}
19
20define_http_rejection! {
21    #[status = BAD_REQUEST]
22    #[body = "Failed to deserialize form"]
23    /// Rejection type used if the [`Form`]
24    /// deserialize the form into the target type.
25    pub struct FailedToDeserializeForm(Error);
26}
27
28composite_http_rejection! {
29    /// Rejection used for [`Form`]
30    ///
31    /// Contains one variant for each way the [`Form`] extractor
32    /// can fail.
33    pub enum FormRejection {
34        InvalidFormContentType,
35        FailedToDeserializeForm,
36        BytesRejection,
37    }
38}
39
40impl<T> FromRequest for Form<T>
41where
42    T: serde::de::DeserializeOwned + Send + Sync + 'static,
43{
44    type Rejection = FormRejection;
45
46    async fn from_request(req: Request) -> Result<Self, Self::Rejection> {
47        // Extracted into separate fn so it's only compiled once for all T.
48        async fn extract_form_body_bytes(req: Request) -> Result<Bytes, FormRejection> {
49            if !crate::service::web::extract::has_any_content_type(
50                req.headers(),
51                &[&crate::mime::APPLICATION_WWW_FORM_URLENCODED],
52            ) {
53                return Err(InvalidFormContentType.into());
54            }
55
56            let body = req.into_body();
57            let bytes = body.collect().await.map_err(BytesRejection::from_err)?;
58
59            Ok(bytes.to_bytes())
60        }
61
62        if req.method() == Method::GET {
63            let value = match req.uri().query_params() {
64                Ok(value) => value,
65                Err(err) => return Err(FailedToDeserializeForm::from_err(err).into()),
66            };
67            Ok(Self(value))
68        } else {
69            let b = extract_form_body_bytes(req).await?;
70            Ok(Self(match serde_html_form::from_bytes(&b) {
71                Ok(value) => value,
72                Err(err) => return Err(FailedToDeserializeForm::from_err(err).into()),
73            }))
74        }
75    }
76}
77
78#[cfg(test)]
79mod test {
80    use super::*;
81    use crate::service::web::WebService;
82    use crate::{Body, Method, Request, StatusCode};
83    use rama_core::Service;
84
85    #[tokio::test]
86    async fn test_form_post_form_urlencoded() {
87        #[derive(Debug, serde::Deserialize)]
88        struct Input {
89            name: String,
90            age: u8,
91        }
92
93        let service = WebService::default().with_post("/", async |Form(body): Form<Input>| {
94            assert_eq!(body.name, "Devan");
95            assert_eq!(body.age, 29);
96        });
97
98        let req = Request::builder()
99            .uri("/")
100            .method(Method::POST)
101            .header("content-type", "application/x-www-form-urlencoded")
102            .body(r#"name=Devan&age=29"#.into())
103            .unwrap();
104        let resp = service.serve(req).await.unwrap();
105        assert_eq!(resp.status(), StatusCode::OK);
106    }
107
108    #[tokio::test]
109    async fn test_form_post_form_urlencoded_missing_data_fail() {
110        #[derive(Debug, serde::Deserialize)]
111        #[expect(dead_code)]
112        struct Input {
113            name: String,
114            age: u8,
115        }
116
117        let service =
118            WebService::default().with_post("/", async |Form(_): Form<Input>| StatusCode::OK);
119
120        let req = Request::builder()
121            .uri("/")
122            .method(Method::POST)
123            .header("content-type", "application/x-www-form-urlencoded")
124            .body(r#"age=29"#.into())
125            .unwrap();
126        let resp = service.serve(req).await.unwrap();
127        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
128    }
129
130    #[tokio::test]
131    async fn test_form_get_form_urlencoded_fail() {
132        #[derive(Debug, serde::Deserialize)]
133        #[expect(dead_code)]
134        struct Input {
135            name: String,
136            age: u8,
137        }
138
139        let service =
140            WebService::default().with_get("/", async |Form(_): Form<Input>| StatusCode::OK);
141
142        let req = Request::builder()
143            .uri("/")
144            .method(Method::GET)
145            .header("content-type", "application/x-www-form-urlencoded")
146            .body(r#"name=Devan&age=29"#.into())
147            .unwrap();
148        let resp = service.serve(req).await.unwrap();
149        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
150    }
151
152    #[tokio::test]
153    async fn test_form_get() {
154        #[derive(Debug, serde::Deserialize)]
155        struct Input {
156            name: String,
157            age: u8,
158        }
159
160        let service = WebService::default().with_get("/", async |Form(body): Form<Input>| {
161            assert_eq!(body.name, "Devan");
162            assert_eq!(body.age, 29);
163        });
164
165        let req = Request::builder()
166            .uri("/?name=Devan&age=29")
167            .method(Method::GET)
168            .body(Body::empty())
169            .unwrap();
170        let resp = service.serve(req).await.unwrap();
171        assert_eq!(resp.status(), StatusCode::OK);
172    }
173
174    #[tokio::test]
175    async fn test_form_get_fail_missing_data() {
176        #[derive(Debug, serde::Deserialize)]
177        #[expect(dead_code)]
178        struct Input {
179            name: String,
180            age: u8,
181        }
182
183        let service =
184            WebService::default().with_get("/", async |Form(_): Form<Input>| StatusCode::OK);
185
186        let req = Request::builder()
187            .uri("/?name=Devan")
188            .method(Method::GET)
189            .body(Body::empty())
190            .unwrap();
191        let resp = service.serve(req).await.unwrap();
192        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
193    }
194}