Skip to main content

headless_lms_server/controllers/helpers/
multi_query.rs

1//! Query-string extraction for endpoints whose filters accept several values.
2
3use std::borrow::Cow;
4use std::collections::BTreeMap;
5use std::future::{Ready, ready};
6use std::ops::Deref;
7
8use actix_web::{FromRequest, HttpRequest, dev::Payload};
9use serde::de::value::{CowStrDeserializer, Error as ValueError, MapDeserializer, SeqDeserializer};
10use serde::de::{self, DeserializeOwned, Deserializer, IntoDeserializer, Visitor};
11use serde::forward_to_deserialize_any;
12use url::form_urlencoded;
13
14use headless_lms_base::prelude_base_and_re_exports::BackendError;
15
16use crate::domain::error::{ControllerError, ControllerErrorType, controller_err};
17
18/// A handler's query parameters, where a parameter given more than once reads as a list.
19///
20/// Use this in place of [`actix_web::web::Query`] whenever the query struct has a `Vec` field.
21/// `web::Query` deserializes with `serde_urlencoded`, which rejects a repeated key
22/// (`?state=a&state=b`) as a duplicate map entry before any field deserializer runs, so a `Vec`
23/// field can never be filled over that transport.
24///
25/// A `Vec` field accepts `?state=a&state=b` and `?state=a,b` alike. A scalar field is never split
26/// on commas, so a free-text parameter may contain them; given twice, it takes the last value. A
27/// parameter present with an empty value reads as absent.
28///
29/// Answers 400 when the query string does not fit `T`.
30#[derive(Debug)]
31pub struct MultiQuery<T>(T);
32
33impl<T> MultiQuery<T> {
34    pub fn into_inner(self) -> T {
35        self.0
36    }
37}
38
39impl<T> Deref for MultiQuery<T> {
40    type Target = T;
41
42    fn deref(&self) -> &T {
43        &self.0
44    }
45}
46
47impl<T: DeserializeOwned> FromRequest for MultiQuery<T> {
48    type Error = ControllerError;
49    type Future = Ready<Result<Self, Self::Error>>;
50
51    fn from_request(req: &HttpRequest, _payload: &mut Payload) -> Self::Future {
52        ready(
53            from_query_string(req.query_string())
54                .map(MultiQuery)
55                .map_err(|err| controller_err!(BadRequest, format!("Query parse error: {err}"))),
56        )
57    }
58}
59
60/// Deserializes `T` from a raw query string under [`MultiQuery`]'s rules.
61pub fn from_query_string<T: DeserializeOwned>(query: &str) -> Result<T, ValueError> {
62    let mut fields: BTreeMap<Cow<'_, str>, Vec<Cow<'_, str>>> = BTreeMap::new();
63    for (key, value) in form_urlencoded::parse(query.as_bytes()) {
64        if value.is_empty() {
65            continue;
66        }
67        fields.entry(key).or_default().push(value);
68    }
69    T::deserialize(QueryDeserializer { fields })
70}
71
72/// Every parameter of one query string, keyed by name, in the order the values were given.
73struct QueryDeserializer<'q> {
74    fields: BTreeMap<Cow<'q, str>, Vec<Cow<'q, str>>>,
75}
76
77impl<'de> Deserializer<'de> for QueryDeserializer<'de> {
78    type Error = ValueError;
79
80    fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
81        MapDeserializer::new(
82            self.fields
83                .into_iter()
84                .map(|(key, values)| (key, Values(values))),
85        )
86        .deserialize_map(visitor)
87    }
88
89    forward_to_deserialize_any! {
90        bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string
91        bytes byte_buf option unit unit_struct newtype_struct seq tuple
92        tuple_struct map struct enum identifier ignored_any
93    }
94}
95
96/// The values one parameter was given, never empty.
97struct Values<'q>(Vec<Cow<'q, str>>);
98
99impl<'de> IntoDeserializer<'de, ValueError> for Values<'de> {
100    type Deserializer = Self;
101
102    fn into_deserializer(self) -> Self {
103        self
104    }
105}
106
107impl<'de> Values<'de> {
108    /// The last value given, as the deserializer a scalar field reads.
109    fn scalar(self) -> CowStrDeserializer<'de, ValueError> {
110        let mut values = self.0;
111        values
112            .pop()
113            .unwrap_or(Cow::Borrowed(""))
114            .into_deserializer()
115    }
116}
117
118/// Parses a scalar field's value out of its string, the way `serde_urlencoded` does: a query string
119/// carries no types, so `limit=50` has to reach a `u32` field as a number.
120macro_rules! deserialize_parsed {
121    ($($method:ident => $visit:ident => $target:ty,)*) => {
122        $(
123            fn $method<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
124                let value = self.scalar_value();
125                match value.parse::<$target>() {
126                    Ok(parsed) => visitor.$visit(parsed),
127                    Err(_) => Err(de::Error::invalid_value(
128                        de::Unexpected::Str(&value),
129                        &visitor,
130                    )),
131                }
132            }
133        )*
134    };
135}
136
137impl<'de> Values<'de> {
138    fn scalar_value(&self) -> Cow<'de, str> {
139        self.0.last().cloned().unwrap_or(Cow::Borrowed(""))
140    }
141}
142
143impl<'de> Deserializer<'de> for Values<'de> {
144    type Error = ValueError;
145
146    fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
147        if self.0.len() > 1 {
148            self.deserialize_seq(visitor)
149        } else {
150            self.scalar().deserialize_any(visitor)
151        }
152    }
153
154    /// A key only reaches here when it was given, so it is always `Some`.
155    fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
156        visitor.visit_some(self)
157    }
158
159    fn deserialize_seq<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
160        SeqDeserializer::new(self.0.into_iter().flat_map(comma_separated)).deserialize_seq(visitor)
161    }
162
163    fn deserialize_enum<V: Visitor<'de>>(
164        self,
165        name: &'static str,
166        variants: &'static [&'static str],
167        visitor: V,
168    ) -> Result<V::Value, ValueError> {
169        self.scalar().deserialize_enum(name, variants, visitor)
170    }
171
172    fn deserialize_newtype_struct<V: Visitor<'de>>(
173        self,
174        _name: &'static str,
175        visitor: V,
176    ) -> Result<V::Value, ValueError> {
177        visitor.visit_newtype_struct(self)
178    }
179
180    fn deserialize_str<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
181        self.scalar().deserialize_str(visitor)
182    }
183
184    fn deserialize_string<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
185        self.scalar().deserialize_string(visitor)
186    }
187
188    fn deserialize_char<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
189        self.scalar().deserialize_char(visitor)
190    }
191
192    fn deserialize_identifier<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, ValueError> {
193        self.scalar().deserialize_identifier(visitor)
194    }
195
196    deserialize_parsed! {
197        deserialize_bool => visit_bool => bool,
198        deserialize_i8 => visit_i8 => i8,
199        deserialize_i16 => visit_i16 => i16,
200        deserialize_i32 => visit_i32 => i32,
201        deserialize_i64 => visit_i64 => i64,
202        deserialize_u8 => visit_u8 => u8,
203        deserialize_u16 => visit_u16 => u16,
204        deserialize_u32 => visit_u32 => u32,
205        deserialize_u64 => visit_u64 => u64,
206        deserialize_f32 => visit_f32 => f32,
207        deserialize_f64 => visit_f64 => f64,
208    }
209
210    forward_to_deserialize_any! {
211        i128 u128 bytes byte_buf unit unit_struct tuple tuple_struct map struct ignored_any
212    }
213}
214
215/// Splits one raw value on commas, dropping empty parts, so `?state=a,b` and `?state=a&state=b`
216/// mean the same thing.
217fn comma_separated(value: Cow<'_, str>) -> Vec<Cow<'_, str>> {
218    match value {
219        Cow::Borrowed(value) => value
220            .split(',')
221            .filter(|part| !part.is_empty())
222            .map(Cow::Borrowed)
223            .collect(),
224        Cow::Owned(value) => value
225            .split(',')
226            .filter(|part| !part.is_empty())
227            .map(|part| Cow::Owned(part.to_owned()))
228            .collect(),
229    }
230}
231
232#[cfg(test)]
233mod tests {
234    use super::*;
235    use chrono::{DateTime, Utc};
236    use serde::Deserialize;
237    use uuid::Uuid;
238
239    #[derive(Debug, Deserialize, PartialEq)]
240    #[serde(rename_all = "snake_case")]
241    enum State {
242        Registered,
243        Blocked,
244    }
245
246    #[derive(Debug, Deserialize, PartialEq)]
247    struct Query {
248        page: Option<u32>,
249        state: Option<Vec<State>>,
250        course_id: Option<Uuid>,
251        submitted_after: Option<DateTime<Utc>>,
252        search: Option<String>,
253        include_superseded: Option<bool>,
254    }
255
256    fn parse(query: &str) -> Query {
257        from_query_string(query).expect("the query string should fit Query")
258    }
259
260    #[test]
261    fn reads_a_repeated_parameter_as_a_list() {
262        assert_eq!(
263            parse("state=registered&state=blocked").state,
264            Some(vec![State::Registered, State::Blocked])
265        );
266        assert_eq!(
267            parse("state=registered,blocked").state,
268            Some(vec![State::Registered, State::Blocked])
269        );
270        assert_eq!(
271            parse("state=registered").state,
272            Some(vec![State::Registered])
273        );
274        assert_eq!(parse("page=2").state, None);
275    }
276
277    #[test]
278    fn reads_the_scalar_parameters_beside_it() {
279        let query = parse(
280            "state=blocked&page=3&course_id=8e4aeba5-1958-49bc-9b40-3c76bb0d3ad4\
281             &submitted_after=2026-09-06T09:51:00Z&include_superseded=true&search=a,b",
282        );
283        assert_eq!(
284            query,
285            Query {
286                page: Some(3),
287                state: Some(vec![State::Blocked]),
288                course_id: Some(
289                    Uuid::parse_str("8e4aeba5-1958-49bc-9b40-3c76bb0d3ad4").expect("a valid uuid")
290                ),
291                submitted_after: Some(
292                    "2026-09-06T09:51:00Z"
293                        .parse::<DateTime<Utc>>()
294                        .expect("a valid timestamp")
295                ),
296                // Never comma-split: a search term may contain one.
297                search: Some("a,b".to_string()),
298                include_superseded: Some(true),
299            }
300        );
301    }
302
303    #[test]
304    fn reads_an_empty_value_as_absent() {
305        assert_eq!(parse("page=&state=&search="), parse(""));
306    }
307
308    #[test]
309    fn refuses_a_value_that_does_not_fit_the_field() {
310        assert!(from_query_string::<Query>("page=soon").is_err());
311        assert!(from_query_string::<Query>("state=elsewhere").is_err());
312    }
313}