1use crate::error::{Error, Result};
4use crate::request::Request;
5use bytes::Bytes;
6use serde::de::DeserializeOwned;
7use std::collections::HashMap;
8use std::path::{Component, Path, PathBuf};
9
10#[derive(Debug, Clone)]
15pub struct Upload {
16 pub field: String,
17 pub filename: Option<String>,
18 pub content_type: Option<String>,
19 pub data: Bytes,
20}
21
22impl Upload {
23 pub fn size(&self) -> usize {
25 self.data.len()
26 }
27
28 pub fn extension(&self) -> Option<String> {
30 self.filename
31 .as_deref()
32 .and_then(|n| Path::new(n).extension())
33 .and_then(|e| e.to_str())
34 .map(|e| e.to_ascii_lowercase())
35 }
36
37 pub fn mime(&self) -> Option<&str> {
39 self.content_type
40 .as_deref()
41 .map(str::trim)
42 .filter(|s| !s.is_empty())
43 }
44
45 pub fn mime_type(&self) -> Option<String> {
47 self.mime()
48 .map(|m| {
49 m.split(';')
50 .next()
51 .unwrap_or(m)
52 .trim()
53 .to_ascii_lowercase()
54 })
55 .filter(|s| !s.is_empty())
56 }
57
58 pub fn validate(&self, rules: &UploadRules) -> Result<()> {
60 rules.check(self)
61 }
62
63 pub async fn save(&self, path: impl AsRef<Path>) -> Result<()> {
65 let path = path.as_ref();
66 if let Some(parent) = path.parent() {
67 if !parent.as_os_str().is_empty() {
68 tokio::fs::create_dir_all(parent)
69 .await
70 .map_err(|e| Error::Internal(e.to_string()))?;
71 }
72 }
73 tokio::fs::write(path, &self.data)
74 .await
75 .map_err(|e| Error::Internal(e.to_string()))
76 }
77
78 pub async fn save_in(&self, dir: impl AsRef<Path>, filename: &str) -> Result<PathBuf> {
80 let name = Path::new(filename);
81 if !is_safe_relative(name) {
82 return Err(Error::BadRequest("unsafe upload filename".into()));
83 }
84 let dest = dir.as_ref().join(name);
85 self.save(&dest).await?;
86 Ok(dest)
87 }
88
89 pub fn suggested_name(&self) -> &str {
91 self.filename
92 .as_deref()
93 .filter(|s| !s.is_empty())
94 .unwrap_or(self.field.as_str())
95 }
96}
97
98#[derive(Debug, Clone, Default)]
100pub struct UploadRules {
101 max_bytes: Option<usize>,
102 extensions: Vec<String>,
103 mimes: Vec<String>,
104}
105
106impl UploadRules {
107 pub fn new() -> Self {
108 Self::default()
109 }
110
111 pub fn max_bytes(mut self, n: usize) -> Self {
112 self.max_bytes = Some(n);
113 self
114 }
115
116 pub fn extensions<I, S>(mut self, exts: I) -> Self
117 where
118 I: IntoIterator<Item = S>,
119 S: AsRef<str>,
120 {
121 self.extensions = exts
122 .into_iter()
123 .map(|s| s.as_ref().trim_start_matches('.').to_ascii_lowercase())
124 .filter(|s| !s.is_empty())
125 .collect();
126 self
127 }
128
129 pub fn mimes<I, S>(mut self, mimes: I) -> Self
130 where
131 I: IntoIterator<Item = S>,
132 S: AsRef<str>,
133 {
134 self.mimes = mimes
135 .into_iter()
136 .map(|s| s.as_ref().trim().to_ascii_lowercase())
137 .filter(|s| !s.is_empty())
138 .collect();
139 self
140 }
141
142 fn check(&self, upload: &Upload) -> Result<()> {
143 if upload.size() == 0 {
144 return Err(Error::BadRequest("empty file".into()));
145 }
146 if let Some(max) = self.max_bytes {
147 if upload.size() > max {
148 return Err(Error::BadRequest(format!(
149 "file too large (max {max} bytes)"
150 )));
151 }
152 }
153 if !self.extensions.is_empty() {
154 let ext = upload.extension().ok_or_else(|| {
155 Error::BadRequest("file extension required".into())
156 })?;
157 if !self.extensions.iter().any(|e| e == &ext) {
158 return Err(Error::BadRequest(format!(
159 "invalid file extension `{ext}`"
160 )));
161 }
162 }
163 if !self.mimes.is_empty() {
164 let mime = upload.mime_type().ok_or_else(|| {
165 Error::BadRequest("file content-type required".into())
166 })?;
167 if !self.mimes.iter().any(|m| m == &mime) {
168 return Err(Error::BadRequest(format!(
169 "invalid content-type `{mime}`"
170 )));
171 }
172 }
173 Ok(())
174 }
175}
176
177#[derive(Debug, Clone, Default)]
179pub struct FormData {
180 texts: HashMap<String, Vec<String>>,
181 files: HashMap<String, Vec<Upload>>,
182}
183
184impl FormData {
185 pub fn get(&self, name: &str) -> Option<&str> {
186 self.texts.get(name)?.first().map(String::as_str)
187 }
188
189 pub fn get_all(&self, name: &str) -> &[String] {
190 self.texts
191 .get(name)
192 .map(Vec::as_slice)
193 .unwrap_or(&[])
194 }
195
196 pub fn file(&self, name: &str) -> Option<&Upload> {
197 self.files.get(name)?.first()
198 }
199
200 pub fn files(&self, name: &str) -> &[Upload] {
201 self.files
202 .get(name)
203 .map(Vec::as_slice)
204 .unwrap_or(&[])
205 }
206
207 pub fn text_map(&self) -> &HashMap<String, Vec<String>> {
208 &self.texts
209 }
210
211 pub fn file_map(&self) -> &HashMap<String, Vec<Upload>> {
212 &self.files
213 }
214
215 fn push_text(&mut self, name: String, value: String) {
216 self.texts.entry(name).or_default().push(value);
217 }
218
219 #[cfg(feature = "multipart")]
220 fn push_file(&mut self, upload: Upload) {
221 self.files
222 .entry(upload.field.clone())
223 .or_default()
224 .push(upload);
225 }
226
227 fn first_values(&self) -> HashMap<String, String> {
229 self.texts
230 .iter()
231 .filter_map(|(k, v)| v.first().cloned().map(|val| (k.clone(), val)))
232 .collect()
233 }
234}
235
236impl Request {
237 pub async fn input(&mut self) -> Result<&FormData> {
239 if self.get::<FormData>().is_some() {
240 return Ok(self.get::<FormData>().expect("FormData"));
241 }
242 let parsed = parse_form_data(self).await?;
243 self.set(parsed);
244 Ok(self.get::<FormData>().expect("FormData"))
245 }
246
247 pub async fn form<T: DeserializeOwned>(&mut self) -> Result<T> {
249 let data = self.input().await?;
250 let map = data.first_values();
251 let encoded = serde_urlencoded::to_string(&map)
252 .map_err(|e| Error::BadRequest(format!("form encode: {e}")))?;
253 serde_urlencoded::from_str(&encoded)
254 .map_err(|e| Error::BadRequest(format!("form error: {e}")))
255 }
256}
257
258async fn parse_form_data(req: &mut Request) -> Result<FormData> {
259 let ct = req.content_type().unwrap_or("").to_ascii_lowercase();
260 if ct.starts_with("multipart/") {
261 #[cfg(feature = "multipart")]
262 {
263 return parse_multipart(req).await;
264 }
265 #[cfg(not(feature = "multipart"))]
266 {
267 return Err(Error::BadRequest(
268 "multipart body requires the `multipart` feature".into(),
269 ));
270 }
271 }
272
273 let bytes = req.collect_body("form").await?;
275 let mut data = FormData::default();
276 if bytes.is_empty() {
277 return Ok(data);
278 }
279 let pairs: Vec<(String, String)> = serde_urlencoded::from_bytes(&bytes)
280 .map_err(|e| Error::BadRequest(format!("form error: {e}")))?;
281 for (k, v) in pairs {
282 data.push_text(k, v);
283 }
284 Ok(data)
285}
286
287#[cfg(feature = "multipart")]
288async fn parse_multipart(req: &mut Request) -> Result<FormData> {
289 use bytes::BytesMut;
290 use futures_util::stream;
291 use http_body_util::BodyExt;
292 use multer::Multipart;
293
294 let ct = req
295 .header("content-type")
296 .ok_or_else(|| Error::BadRequest("missing content-type".into()))?
297 .to_string();
298 let boundary = multer::parse_boundary(&ct)
299 .map_err(|e| Error::BadRequest(format!("multipart boundary: {e}")))?;
300
301 let limit = req.body_limit();
302 let mut body = req.into_body_stream_as("multipart")?;
303 let mut collected = BytesMut::new();
304 while let Some(frame) = body.frame().await {
305 let frame = frame.map_err(|e| Error::BadRequest(format!("multipart: {e}")))?;
306 if let Ok(chunk) = frame.into_data() {
307 if collected.len().saturating_add(chunk.len()) > limit {
308 return Err(Error::PayloadTooLarge);
309 }
310 collected.extend_from_slice(&chunk);
311 }
312 }
313 let bytes = collected.freeze();
314 req.body = crate::request::ReqBody::Bytes(bytes.clone());
316
317 let stream = stream::once(async move { Ok::<_, std::io::Error>(bytes) });
318 let mut mp = Multipart::new(stream, boundary);
319 let mut data = FormData::default();
320 while let Some(field) = mp
321 .next_field()
322 .await
323 .map_err(|e| Error::BadRequest(format!("multipart: {e}")))?
324 {
325 let name = field.name().unwrap_or("").to_string();
326 let filename = field.file_name().map(str::to_string);
327 let content_type = field.content_type().map(|m| m.to_string());
328 let part = field
329 .bytes()
330 .await
331 .map_err(|e| Error::BadRequest(format!("multipart field: {e}")))?;
332 if filename.is_some() {
333 data.push_file(Upload {
334 field: name,
335 filename,
336 content_type,
337 data: part,
338 });
339 } else {
340 let s = String::from_utf8_lossy(&part).into_owned();
341 data.push_text(name, s);
342 }
343 }
344 Ok(data)
345}
346
347fn is_safe_relative(path: &Path) -> bool {
348 !path.as_os_str().is_empty()
349 && path
350 .components()
351 .all(|c| matches!(c, Component::Normal(_)))
352}
353
354#[cfg(test)]
355mod upload_rules_tests {
356 use super::*;
357 use bytes::Bytes;
358
359 fn upload(name: &str, ct: Option<&str>, data: &'static [u8]) -> Upload {
360 Upload {
361 field: "file".into(),
362 filename: Some(name.into()),
363 content_type: ct.map(str::to_owned),
364 data: Bytes::from_static(data),
365 }
366 }
367
368 #[test]
369 fn helpers_extension_mime_size() {
370 let u = upload("Photo.PNG", Some("image/png; charset=binary"), b"abc");
371 assert_eq!(u.size(), 3);
372 assert_eq!(u.extension().as_deref(), Some("png"));
373 assert_eq!(u.mime_type().as_deref(), Some("image/png"));
374 }
375
376 #[test]
377 fn rejects_empty_and_oversized() {
378 let empty = upload("a.txt", None, b"");
379 assert!(empty.validate(&UploadRules::new()).is_err());
380
381 let big = upload("a.txt", None, b"hello");
382 assert!(big
383 .validate(&UploadRules::new().max_bytes(4))
384 .is_err());
385 assert!(big
386 .validate(&UploadRules::new().max_bytes(5))
387 .is_ok());
388 }
389
390 #[test]
391 fn extensions_and_mimes() {
392 let u = upload("a.JPG", Some("image/jpeg"), b"x");
393 assert!(u
394 .validate(&UploadRules::new().extensions(["png", "jpg"]))
395 .is_ok());
396 assert!(u
397 .validate(&UploadRules::new().extensions(["png"]))
398 .is_err());
399 assert!(u
400 .validate(&UploadRules::new().mimes(["image/jpeg"]))
401 .is_ok());
402 assert!(u
403 .validate(&UploadRules::new().mimes(["image/png"]))
404 .is_err());
405 }
406}
407
408#[cfg(all(test, feature = "multipart"))]
409mod tests {
410 use super::*;
411 use crate::Request;
412 use bytes::Bytes;
413 use http::Method;
414
415 fn multipart_body(boundary: &str, parts: &str) -> Bytes {
416 Bytes::from(format!("--{boundary}\r\n{parts}--{boundary}--\r\n"))
417 }
418
419 fn multipart_req(boundary: &str, parts: &str) -> Request {
420 Request::builder()
421 .method(Method::POST)
422 .path("/upload")
423 .header(
424 "content-type",
425 format!("multipart/form-data; boundary={boundary}"),
426 )
427 .body(multipart_body(boundary, parts))
428 .build()
429 }
430
431 #[tokio::test]
432 async fn parses_text_and_file_fields() {
433 let boundary = "----sovaBound";
434 let parts = concat!(
435 "Content-Disposition: form-data; name=\"title\"\r\n\r\n",
436 "hello\r\n",
437 "------sovaBound\r\n",
438 "Content-Disposition: form-data; name=\"file\"; filename=\"a.txt\"\r\n",
439 "Content-Type: text/plain\r\n\r\n",
440 "file-bytes\r\n",
441 );
442 let mut req = multipart_req(boundary, parts);
443 let data = req.input().await.unwrap();
444 assert_eq!(data.get("title"), Some("hello"));
445 let file = data.file("file").unwrap();
446 assert_eq!(file.filename.as_deref(), Some("a.txt"));
447 assert_eq!(file.data.as_ref(), b"file-bytes");
448 }
449
450 #[tokio::test]
451 async fn urlencoded_form_via_input() {
452 let mut req = Request::builder()
453 .method(Method::POST)
454 .path("/")
455 .header("content-type", "application/x-www-form-urlencoded")
456 .body("name=Ada&age=1")
457 .build();
458 #[derive(serde::Deserialize, Debug, PartialEq)]
459 struct Body {
460 name: String,
461 age: u32,
462 }
463 let body: Body = req.form().await.unwrap();
464 assert_eq!(
465 body,
466 Body {
467 name: "Ada".into(),
468 age: 1
469 }
470 );
471 }
472
473 #[tokio::test]
474 async fn missing_boundary_is_bad_request() {
475 let mut req = Request::builder()
476 .method(Method::POST)
477 .path("/")
478 .header("content-type", "multipart/form-data")
479 .body("x")
480 .build();
481 let err = req.input().await.unwrap_err();
482 assert!(matches!(err, Error::BadRequest(_)));
483 }
484
485 #[tokio::test]
486 async fn oversize_body_is_413() {
487 let boundary = "b";
488 let big = "x".repeat(64);
489 let parts = format!("Content-Disposition: form-data; name=\"f\"\r\n\r\n{big}\r\n");
490 let mut req = Request::builder()
491 .method(Method::POST)
492 .path("/")
493 .header(
494 "content-type",
495 format!("multipart/form-data; boundary={boundary}"),
496 )
497 .body(multipart_body(boundary, &parts))
498 .body_limit(16)
499 .build();
500 let err = req.input().await.unwrap_err();
501 assert!(matches!(err, Error::PayloadTooLarge), "got {err:?}");
502 }
503
504 #[tokio::test]
505 async fn broken_delimiter_is_bad_request() {
506 let mut req = Request::builder()
507 .method(Method::POST)
508 .path("/")
509 .header("content-type", "multipart/form-data; boundary=abc")
510 .body("not-a-multipart-body")
511 .build();
512 let err = req.input().await.unwrap_err();
513 assert!(matches!(err, Error::BadRequest(_)), "got {err:?}");
514 }
515}