Skip to main content

ocre/
form.rs

1//! Forms whose field names use brackets, Rails-style: `post[title]`, `tag_ids[]`, `items[0][name]`.
2
3use 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/// Form extractor for bracketed field names, like Rails' `params` (`post[title]`, `tag_ids[]`, `lines[0][qty]`).
17///
18/// axum's `Form` only knows flat names (`title=...`); use it for ordinary
19/// forms. `NestedForm` is for forms that send a list or nested records:
20///
21/// | Field name | Becomes |
22/// |---|---|
23/// | `title` | the field `title` |
24/// | `post[title]` | the field `title` of the struct in field `post` |
25/// | `tag_ids[]` (repeated) | a `Vec` with one item per value, in order |
26/// | `lines[0][qty]`, `lines[1][qty]` | a `Vec` of structs, sorted by index (Rails' `fields_for` / nested attributes) |
27/// | `lines[][qty]`, `lines[][price]` | a `Vec` of structs; a name already set in the last item starts a new one |
28///
29/// A name sent twice keeps the last value (so a hidden `done=0` before the
30/// checkbox `done=1` works as in Rails). Values are text: numbers and
31/// booleans are parsed from it (`1`/`true`/`on`/`yes` are `true`, `0`,
32/// `false`, `off`, `no` and an empty value are `false`), an empty value is
33/// `None` for an `Option`, and a single value fills a one-item `Vec`. Enums
34/// with unit variants read their variant name.
35///
36/// On `GET` and `HEAD` it reads the query string; otherwise the body, which
37/// must be `application/x-www-form-urlencoded` (2 MB at most, axum's
38/// default limit). A body that does not decode into `T` (a missing field, a
39/// number that is not one, `a=1&a[b]=2`) is a 400, as an HTML page for
40/// browsers (feature `html`) or JSON for API clients. Pure CPU, no binding
41/// call.
42///
43/// # Examples
44///
45/// ```
46/// use ocre::NestedForm;
47/// use serde::Deserialize;
48///
49/// #[derive(Debug, Deserialize, PartialEq)]
50/// struct Order {
51///     customer: Customer,
52///     #[serde(default)]
53///     tag_ids: Vec<i64>,
54///     lines: Vec<Line>,
55/// }
56///
57/// #[derive(Debug, Deserialize, PartialEq)]
58/// struct Customer {
59///     name: String,
60/// }
61///
62/// #[derive(Debug, Deserialize, PartialEq)]
63/// struct Line {
64///     product: String,
65///     qty: u32,
66///     #[serde(default)]
67///     remove: bool,
68/// }
69///
70/// let body = "customer[name]=Ada&tag_ids[]=3&tag_ids[]=7\
71///             &lines[0][product]=Tea&lines[0][qty]=2\
72///             &lines[1][product]=Cake&lines[1][qty]=1&lines[1][remove]=0&lines[1][remove]=1";
73/// let order: Order = NestedForm::parse(body).unwrap();
74/// assert_eq!(order.customer.name, "Ada");
75/// assert_eq!(order.tag_ids, [3, 7]);
76/// assert_eq!(order.lines[1], Line { product: "Cake".into(), qty: 1, remove: true });
77/// ```
78///
79/// In a handler, like axum's `Form`:
80///
81/// ```no_run
82/// use axum::response::Redirect;
83/// use ocre::{NestedForm, Result};
84/// # #[derive(serde::Deserialize)] struct Order { lines: Vec<String> }
85///
86/// async fn create(NestedForm(order): NestedForm<Order>) -> Result<Redirect> {
87///     # let _ = order.lines;
88///     Ok(Redirect::to("/orders"))
89/// }
90/// # let _ = create;
91/// ```
92#[derive(Debug, Clone, Copy, Default)]
93pub struct NestedForm<T>(pub T);
94
95impl<T: DeserializeOwned> NestedForm<T> {
96    /// Decodes an `application/x-www-form-urlencoded` string (a body or a query string) with bracketed names.
97    ///
98    /// The rules are those of [`NestedForm`]; errors are
99    /// [`Error::BadRequest`] explaining which field failed.
100    ///
101    /// # Examples
102    ///
103    /// ```
104    /// use ocre::NestedForm;
105    /// use std::collections::BTreeMap;
106    ///
107    /// // `?filter[status]=open&filter[author]=ada` on an index page.
108    /// let filter: BTreeMap<String, BTreeMap<String, String>> =
109    ///     NestedForm::parse("filter[status]=open&filter[author]=ada").unwrap();
110    /// assert_eq!(filter["filter"]["author"], "ada");
111    /// assert!(NestedForm::<BTreeMap<String, u32>>::parse("n=many").is_err());
112    /// ```
113    pub fn parse(input: &str) -> Result<T> {
114        // Decoding pairs cannot fail: invalid percent-escapes and UTF-8 are kept as text.
115        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
145/// Browsers ask for HTML; API clients (`curl`, `fetch`) do not.
146fn 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/// A decoded form: text values in maps and lists.
160#[derive(Debug, Clone, PartialEq)]
161enum Node {
162    Leaf(String),
163    Map(Vec<(String, Node)>),
164    List(Vec<Node>),
165}
166
167/// `a[b][]` is `["a", "b", ""]`; a name that is not well bracketed is one segment.
168fn 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
184/// Puts `value` at `path` in `map`, creating groups on the way.
185fn 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
207/// `key[]=v` appends `v`; `key[][name]=v` sets `name` in the last item, or in a new one when it has it.
208fn 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        // `key[][]=v`: each value is a new one-item list.
215        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                // `lines[0][qty]`: a group whose names are all indexes is a list, in index order.
315                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;