headless_lms_server/controllers/helpers/
multi_query.rs1use 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#[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
60pub 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
72struct 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
96struct 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 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
118macro_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 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
215fn 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 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}