1use axum::{
4 body::Bytes,
5 extract::{FromRequest, Request},
6 http::{HeaderMap, Method, header},
7 response::{IntoResponse, Response},
8};
9use serde::{
10 Deserializer,
11 de::{self, DeserializeOwned, IntoDeserializer, Visitor, value::MapDeserializer, value::SeqDeserializer},
12};
13
14use crate::{ApiError, Error, Result};
15
16#[derive(Debug, Clone, Copy, Default)]
93pub struct NestedForm<T>(pub T);
94
95impl<T: DeserializeOwned> NestedForm<T> {
96 pub fn parse(input: &str) -> Result<T> {
114 let pairs: Vec<(String, String)> = serde_urlencoded::from_str(input).unwrap_or_default();
116 let mut root = Vec::new();
117 for (name, value) in pairs {
118 insert(&mut root, &segments(&name), value)?;
119 }
120 T::deserialize(Node::Map(root)).map_err(|err| Error::bad_request(format!("Invalid form data: {err}")))
121 }
122}
123
124impl<T: DeserializeOwned, S: Send + Sync> FromRequest<S> for NestedForm<T> {
125 type Rejection = Response;
126
127 async fn from_request(req: Request, state: &S) -> std::result::Result<Self, Self::Rejection> {
128 let html = wants_html(req.headers());
129 read(req, state).await.map(Self).map_err(|err| rejection(err, html))
130 }
131}
132
133async fn read<T: DeserializeOwned, S: Send + Sync>(req: Request, state: &S) -> Result<T> {
134 if req.method() == Method::GET || req.method() == Method::HEAD {
135 return NestedForm::parse(req.uri().query().unwrap_or(""));
136 }
137 let content_type = req.headers().get(header::CONTENT_TYPE).and_then(|value| value.to_str().ok()).unwrap_or("");
138 if !content_type.starts_with("application/x-www-form-urlencoded") {
139 return Err(Error::bad_request("Expected an application/x-www-form-urlencoded body"));
140 }
141 let body = Bytes::from_request(req, state).await.map_err(|err| Error::bad_request(err.body_text()))?;
142 NestedForm::parse(&String::from_utf8_lossy(&body))
143}
144
145fn wants_html(headers: &HeaderMap) -> bool {
147 headers.get(header::ACCEPT).and_then(|value| value.to_str().ok()).is_some_and(|accept| accept.contains("text/html"))
148}
149
150fn rejection(err: Error, html: bool) -> Response {
151 #[cfg(feature = "html")]
152 if html {
153 return err.into_response();
154 }
155 let _ = html;
156 ApiError(err).into_response()
157}
158
159#[derive(Debug, Clone, PartialEq)]
161enum Node {
162 Leaf(String),
163 Map(Vec<(String, Node)>),
164 List(Vec<Node>),
165}
166
167fn segments(name: &str) -> Vec<&str> {
169 let Some(open) = name.find('[').filter(|&open| open > 0 && name.ends_with(']')) else {
170 return vec![name];
171 };
172 let inner = &name[open + 1..name.len() - 1];
173 let parts: Vec<&str> = inner.split("][").collect();
174 if parts.iter().any(|part| part.contains(['[', ']'])) {
175 return vec![name];
176 }
177 std::iter::once(&name[..open]).chain(parts).collect()
178}
179
180fn conflict(key: &str) -> Error {
181 Error::bad_request(format!("Invalid form data: `{key}` is both a value and a group of fields"))
182}
183
184fn insert(map: &mut Vec<(String, Node)>, path: &[&str], value: String) -> Result<()> {
186 let (key, rest) = (path[0], &path[1..]);
187 let position = map.iter().position(|(name, _)| name == key);
188 let Some(next) = rest.first() else {
189 match position {
190 Some(index) => map[index].1 = Node::Leaf(value),
191 None => map.push((key.to_owned(), Node::Leaf(value))),
192 }
193 return Ok(());
194 };
195 let index = position.unwrap_or_else(|| {
196 let empty = if next.is_empty() { Node::List(Vec::new()) } else { Node::Map(Vec::new()) };
197 map.push((key.to_owned(), empty));
198 map.len() - 1
199 });
200 match (&mut map[index].1, next.is_empty()) {
201 (Node::Map(child), false) => insert(child, rest, value),
202 (Node::List(items), true) => push(items, &rest[1..], value),
203 _ => Err(conflict(key)),
204 }
205}
206
207fn push(items: &mut Vec<Node>, path: &[&str], value: String) -> Result<()> {
209 let Some(name) = path.first() else {
210 items.push(Node::Leaf(value));
211 return Ok(());
212 };
213 if name.is_empty() {
214 let mut inner = Vec::new();
216 let result = push(&mut inner, &path[1..], value);
217 items.push(Node::List(inner));
218 return result;
219 }
220 let mut fields = match items.pop() {
221 Some(Node::Map(fields)) if !fields.iter().any(|(field, _)| field == name) => fields,
222 Some(last) => {
223 items.push(last);
224 Vec::new()
225 }
226 None => Vec::new(),
227 };
228 let result = insert(&mut fields, path, value);
229 items.push(Node::Map(fields));
230 result
231}
232
233type DeError = de::value::Error;
234
235impl Node {
236 fn invalid(&self, expected: &str) -> DeError {
237 de::Error::custom(match self {
238 Node::Leaf(value) => format!("expected {expected}, found `{value}`"),
239 _ => format!("expected {expected}, found a group of fields"),
240 })
241 }
242
243 fn parse<T: std::str::FromStr>(&self, expected: &str) -> std::result::Result<T, DeError> {
244 match self {
245 Node::Leaf(value) => value.trim().parse().map_err(|_| self.invalid(expected)),
246 _ => Err(self.invalid(expected)),
247 }
248 }
249}
250
251impl<'de> IntoDeserializer<'de, DeError> for Node {
252 type Deserializer = Self;
253
254 fn into_deserializer(self) -> Self {
255 self
256 }
257}
258
259macro_rules! parse_number {
260 ($($method:ident => $visit:ident, $expected:literal;)*) => {
261 $(fn $method<V: Visitor<'de>>(self, visitor: V) -> std::result::Result<V::Value, DeError> {
262 visitor.$visit(self.parse($expected)?)
263 })*
264 };
265}
266
267impl<'de> Deserializer<'de> for Node {
268 type Error = DeError;
269
270 fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> std::result::Result<V::Value, DeError> {
271 match self {
272 Node::Leaf(value) => visitor.visit_string(value),
273 Node::Map(fields) => visitor.visit_map(MapDeserializer::new(fields.into_iter())),
274 Node::List(items) => visitor.visit_seq(SeqDeserializer::new(items.into_iter())),
275 }
276 }
277
278 parse_number! {
279 deserialize_i8 => visit_i8, "an integer";
280 deserialize_i16 => visit_i16, "an integer";
281 deserialize_i32 => visit_i32, "an integer";
282 deserialize_i64 => visit_i64, "an integer";
283 deserialize_u8 => visit_u8, "a positive integer";
284 deserialize_u16 => visit_u16, "a positive integer";
285 deserialize_u32 => visit_u32, "a positive integer";
286 deserialize_u64 => visit_u64, "a positive integer";
287 deserialize_f32 => visit_f32, "a number";
288 deserialize_f64 => visit_f64, "a number";
289 }
290
291 fn deserialize_bool<V: Visitor<'de>>(self, visitor: V) -> std::result::Result<V::Value, DeError> {
292 let value = match &self {
293 Node::Leaf(value) => match value.trim().to_ascii_lowercase().as_str() {
294 "1" | "true" | "on" | "yes" => Some(true),
295 "" | "0" | "false" | "off" | "no" => Some(false),
296 _ => None,
297 },
298 _ => None,
299 };
300 visitor.visit_bool(value.ok_or_else(|| self.invalid("true or false"))?)
301 }
302
303 fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> std::result::Result<V::Value, DeError> {
304 match &self {
305 Node::Leaf(value) if value.is_empty() => visitor.visit_none(),
306 _ => visitor.visit_some(self),
307 }
308 }
309
310 fn deserialize_seq<V: Visitor<'de>>(self, visitor: V) -> std::result::Result<V::Value, DeError> {
311 match self {
312 Node::List(items) => visitor.visit_seq(SeqDeserializer::new(items.into_iter())),
313 Node::Map(fields) => {
314 let mut indexed = Vec::with_capacity(fields.len());
316 for (name, node) in fields {
317 let index: u64 = name.parse().map_err(|_| Node::Map(Vec::new()).invalid("a list"))?;
318 indexed.push((index, node));
319 }
320 indexed.sort_by_key(|(index, _)| *index);
321 visitor.visit_seq(SeqDeserializer::new(indexed.into_iter().map(|(_, node)| node)))
322 }
323 leaf => visitor.visit_seq(SeqDeserializer::new(std::iter::once(leaf))),
324 }
325 }
326
327 fn deserialize_enum<V: Visitor<'de>>(
328 self,
329 _name: &'static str,
330 _variants: &'static [&'static str],
331 visitor: V,
332 ) -> std::result::Result<V::Value, DeError> {
333 match self {
334 Node::Leaf(value) => visitor.visit_enum(value.into_deserializer()),
335 other => Err(other.invalid("one of the choices")),
336 }
337 }
338
339 fn deserialize_newtype_struct<V: Visitor<'de>>(
340 self,
341 _name: &'static str,
342 visitor: V,
343 ) -> std::result::Result<V::Value, DeError> {
344 visitor.visit_newtype_struct(self)
345 }
346
347 serde::forward_to_deserialize_any! {
348 i128 u128 char str string bytes byte_buf unit unit_struct tuple
349 tuple_struct map struct identifier ignored_any
350 }
351}
352
353#[cfg(test)]
354#[path = "../tests/form.rs"]
355mod tests;