#![cfg_attr(
not(test),
deny(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::unreachable,
clippy::todo,
clippy::unimplemented,
clippy::indexing_slicing,
clippy::string_slice,
clippy::arithmetic_side_effects,
)
)]
use std::collections::BTreeMap;
use std::fmt;
use serde::de::{
self, DeserializeOwned, DeserializeSeed, EnumAccess, MapAccess, SeqAccess, VariantAccess,
Visitor,
};
pub const MAX_DEPTH: usize = 16;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct QueryError {
path: Vec<String>,
message: String,
}
impl QueryError {
fn msg(message: impl Into<String>) -> Self {
Self {
path: Vec::new(),
message: message.into(),
}
}
fn with_segment(mut self, segment: impl Into<String>) -> Self {
self.path.insert(0, segment.into());
self
}
}
impl fmt::Display for QueryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.path.is_empty() {
f.write_str(&self.message)
} else {
write!(f, "{}: {}", self.path.join("."), self.message)
}
}
}
impl std::error::Error for QueryError {}
impl de::Error for QueryError {
fn custom<T: fmt::Display>(_msg: T) -> Self {
Self::msg("a value was rejected by a custom deserializer".to_owned())
}
fn invalid_type(unexp: de::Unexpected<'_>, exp: &dyn de::Expected) -> Self {
Self::msg(format!(
"invalid type: {}, expected {exp}",
unexpected_shape(&unexp)
))
}
fn invalid_value(unexp: de::Unexpected<'_>, exp: &dyn de::Expected) -> Self {
Self::msg(format!(
"invalid value: {}, expected {exp}",
unexpected_shape(&unexp)
))
}
fn unknown_variant(variant: &str, expected: &'static [&'static str]) -> Self {
let _ = variant;
Self::msg(format!(
"unknown variant, expected one of {}",
quoted_list(expected)
))
}
fn unknown_field(field: &str, expected: &'static [&'static str]) -> Self {
let _ = field;
Self::msg(format!(
"unknown field, expected one of {}",
quoted_list(expected)
))
}
}
const fn unexpected_shape(unexp: &de::Unexpected<'_>) -> &'static str {
match unexp {
de::Unexpected::Bool(_) => "a boolean",
de::Unexpected::Unsigned(_) | de::Unexpected::Signed(_) => "an integer",
de::Unexpected::Float(_) => "a float",
de::Unexpected::Char(_) => "a character",
de::Unexpected::Str(_) => "a string",
de::Unexpected::Bytes(_) => "a byte string",
de::Unexpected::Unit => "a unit value",
de::Unexpected::Option => "an optional value",
de::Unexpected::NewtypeStruct => "a newtype struct",
de::Unexpected::Seq => "a sequence",
de::Unexpected::Map => "an object",
de::Unexpected::Enum => "an enum",
de::Unexpected::UnitVariant => "a unit variant",
de::Unexpected::NewtypeVariant => "a newtype variant",
de::Unexpected::TupleVariant => "a tuple variant",
de::Unexpected::StructVariant => "a struct variant",
de::Unexpected::Other(_) => "a value",
}
}
fn quoted_list(expected: &'static [&'static str]) -> String {
if expected.is_empty() {
return "no variants".to_owned();
}
expected
.iter()
.map(|name| format!("`{name}`"))
.collect::<Vec<_>>()
.join(", ")
}
fn sanitize_key(raw: &str) -> String {
const MAX: usize = 48;
let mut out: String = raw
.chars()
.take(MAX)
.map(|c| if c.is_control() { '.' } else { c })
.collect();
if raw.chars().nth(MAX).is_some() {
out.push('…');
}
out
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum Segment<'a> {
Append,
Index { position: usize, raw: &'a str },
Key(&'a str),
}
fn parse_key(key: &str) -> Option<(&str, Vec<Segment<'_>>)> {
let open = key.find('[')?;
let (base, mut rest) = key.split_at(open);
let mut segments = Vec::new();
while !rest.is_empty() {
let (inner, remainder) = rest.strip_prefix('[')?.split_once(']')?;
if inner.contains('[') {
return None;
}
segments.push(classify_segment(inner));
rest = remainder;
}
Some((base, segments))
}
fn classify_segment(inner: &str) -> Segment<'_> {
if inner.is_empty() {
return Segment::Append;
}
if inner.bytes().all(|b| b.is_ascii_digit()) {
return Segment::Index {
position: inner.parse().unwrap_or(usize::MAX),
raw: inner,
};
}
Segment::Key(inner)
}
#[derive(Debug)]
enum Node {
Scalar(Vec<String>),
Seq(BTreeMap<SeqKey, Self>),
Map(BTreeMap<String, Self>),
Conflict(String),
}
impl Node {
const fn kind(&self) -> &'static str {
match self {
Self::Scalar(_) => "a value",
Self::Seq(_) => "a sequence",
Self::Map(_) => "an object",
Self::Conflict(_) => "an unresolvable key",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct SeqKey {
position: usize,
raw: String,
}
impl SeqKey {
fn synthesized(position: usize) -> Self {
Self {
position,
raw: position.to_string(),
}
}
}
fn key_path(base: &str, segments: &[Segment<'_>], upto: usize) -> String {
let mut out = base.to_owned();
for segment in segments.iter().take(upto) {
match segment {
Segment::Key(name) => {
out.push('[');
out.push_str(name);
out.push(']');
}
Segment::Index { raw, .. } => {
out.push('[');
out.push_str(raw);
out.push(']');
}
Segment::Append => out.push_str("[]"),
}
}
out
}
fn insert(root: &mut BTreeMap<String, Node>, base: &str, segments: &[Segment<'_>], value: String) {
if segments.len() >= MAX_DEPTH {
let node = root
.entry(base.to_owned())
.or_insert_with(|| Node::Scalar(Vec::new()));
poison(
node,
format!(
"query key `{}` nests deeper than the maximum of {MAX_DEPTH}",
sanitize_key(base)
),
);
return;
}
let mut node = root
.entry(base.to_owned())
.or_insert_with(|| empty_for(segments.first()));
for (position, segment) in segments.iter().enumerate() {
let next = segments.get(position.saturating_add(1));
match segment {
Segment::Key(name) => {
promote_to_map(node);
let Node::Map(entries) = node else {
return poison(node, conflict(base, segments, position, "an object"));
};
node = entries
.entry((*name).to_owned())
.or_insert_with(|| empty_for(next));
}
Segment::Index {
position: index,
raw,
} => match node {
Node::Seq(entries) => {
node = entries
.entry(SeqKey {
position: *index,
raw: (*raw).to_owned(),
})
.or_insert_with(|| empty_for(next));
}
Node::Map(entries) => {
node = entries
.entry((*raw).to_owned())
.or_insert_with(|| empty_for(next));
}
other => {
return poison(other, conflict(base, segments, position, "a sequence"));
}
},
Segment::Append => match node {
Node::Seq(entries) => {
let index = entries
.keys()
.next_back()
.map_or(0, |last| last.position.saturating_add(1));
node = entries
.entry(SeqKey::synthesized(index))
.or_insert_with(|| empty_for(next));
}
Node::Map(entries) => {
let index = entries
.keys()
.filter_map(|key| key.parse::<usize>().ok())
.max()
.map_or(0, |last| last.saturating_add(1));
node = entries
.entry(index.to_string())
.or_insert_with(|| empty_for(next));
}
other => {
return poison(other, conflict(base, segments, position, "a sequence"));
}
},
}
}
match node {
Node::Scalar(values) => values.push(value),
_ => poison(node, conflict(base, segments, segments.len(), "a value")),
}
}
fn promote_to_map(node: &mut Node) {
if let Node::Seq(entries) = node {
let promoted = std::mem::take(entries)
.into_iter()
.map(|(key, child)| (key.raw, child))
.collect();
*node = Node::Map(promoted);
}
}
fn poison(node: &mut Node, message: String) {
if !matches!(node, Node::Conflict(_)) {
*node = Node::Conflict(message);
}
}
const fn empty_for(next: Option<&Segment<'_>>) -> Node {
match next {
None => Node::Scalar(Vec::new()),
Some(Segment::Key(_)) => Node::Map(BTreeMap::new()),
Some(Segment::Append | Segment::Index { .. }) => Node::Seq(BTreeMap::new()),
}
}
fn conflict(base: &str, segments: &[Segment<'_>], upto: usize, wanted: &str) -> String {
format!(
"query key `{}` is used as both a container and {wanted}",
sanitize_key(&key_path(base, segments, upto))
)
}
fn build_tree<'a>(
pairs: impl Iterator<Item = (std::borrow::Cow<'a, str>, std::borrow::Cow<'a, str>)>,
) -> BTreeMap<String, Node> {
let mut root = BTreeMap::new();
for (key, value) in pairs {
match parse_key(&key) {
Some((base, segments)) => insert(&mut root, base, &segments, value.into_owned()),
None => insert(&mut root, &key, &[], value.into_owned()),
}
}
root
}
pub fn from_query_str<T: DeserializeOwned>(query: &str) -> Result<T, QueryError> {
let root = build_tree(url::form_urlencoded::parse(query.as_bytes()));
T::deserialize(NodeDeserializer {
node: &Node::Map(root),
})
}
struct ValueDeserializer<'a> {
value: &'a str,
}
macro_rules! parse_scalar {
($($method:ident => $visit:ident : $ty:ty),* $(,)?) => {
$(
fn $method<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
let parsed: $ty = self.value.parse().map_err(|_| {
QueryError::msg(concat!("invalid ", stringify!($ty), " value"))
})?;
visitor.$visit(parsed)
}
)*
};
}
impl<'de> de::Deserializer<'de> for ValueDeserializer<'_> {
type Error = QueryError;
fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_str(self.value)
}
parse_scalar! {
deserialize_bool => visit_bool: bool,
deserialize_i8 => visit_i8: i8,
deserialize_i16 => visit_i16: i16,
deserialize_i32 => visit_i32: i32,
deserialize_i64 => visit_i64: i64,
deserialize_i128 => visit_i128: i128,
deserialize_u8 => visit_u8: u8,
deserialize_u16 => visit_u16: u16,
deserialize_u32 => visit_u32: u32,
deserialize_u64 => visit_u64: u64,
deserialize_u128 => visit_u128: u128,
deserialize_f32 => visit_f32: f32,
deserialize_f64 => visit_f64: f64,
deserialize_char => visit_char: char,
}
fn deserialize_str<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_str(self.value)
}
fn deserialize_string<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_str(self.value)
}
fn deserialize_bytes<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_bytes(self.value.as_bytes())
}
fn deserialize_byte_buf<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_bytes(self.value.as_bytes())
}
fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_some(self)
}
fn deserialize_unit<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_unit()
}
fn deserialize_unit_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, QueryError> {
visitor.visit_unit()
}
fn deserialize_newtype_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, QueryError> {
visitor.visit_newtype_struct(self)
}
fn deserialize_seq<V: Visitor<'de>>(self, _visitor: V) -> Result<V::Value, QueryError> {
Err(QueryError::msg("expected a sequence, found a single value"))
}
fn deserialize_tuple<V: Visitor<'de>>(
self,
_len: usize,
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_seq(visitor)
}
fn deserialize_tuple_struct<V: Visitor<'de>>(
self,
_name: &'static str,
_len: usize,
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_seq(visitor)
}
fn deserialize_map<V: Visitor<'de>>(self, _visitor: V) -> Result<V::Value, QueryError> {
Err(QueryError::msg("expected an object, found a single value"))
}
fn deserialize_struct<V: Visitor<'de>>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_map(visitor)
}
fn deserialize_enum<V: Visitor<'de>>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, QueryError> {
visitor.visit_enum(UnitVariant { value: self.value })
}
fn deserialize_identifier<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_str(self.value)
}
fn deserialize_ignored_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_unit()
}
}
struct UnitVariant<'a> {
value: &'a str,
}
impl<'de> EnumAccess<'de> for UnitVariant<'_> {
type Error = QueryError;
type Variant = Self;
fn variant_seed<V: DeserializeSeed<'de>>(
self,
seed: V,
) -> Result<(V::Value, Self), QueryError> {
let variant = seed.deserialize(ValueDeserializer { value: self.value })?;
Ok((variant, self))
}
}
impl<'de> VariantAccess<'de> for UnitVariant<'_> {
type Error = QueryError;
fn unit_variant(self) -> Result<(), QueryError> {
Ok(())
}
fn newtype_variant_seed<T: DeserializeSeed<'de>>(
self,
_seed: T,
) -> Result<T::Value, QueryError> {
Err(QueryError::msg(
"expected a unit variant; a data-carrying variant needs the bracketed form \
`key[variant][field]=…`",
))
}
fn tuple_variant<V: Visitor<'de>>(
self,
_len: usize,
_visitor: V,
) -> Result<V::Value, QueryError> {
Err(QueryError::msg("expected a unit variant"))
}
fn struct_variant<V: Visitor<'de>>(
self,
_fields: &'static [&'static str],
_visitor: V,
) -> Result<V::Value, QueryError> {
Err(QueryError::msg("expected a unit variant"))
}
}
struct NodeDeserializer<'a> {
node: &'a Node,
}
impl<'a> NodeDeserializer<'a> {
fn check(&self) -> Result<(), QueryError> {
match self.node {
Node::Conflict(message) => Err(QueryError::msg(message.clone())),
_ => Ok(()),
}
}
fn as_value(&self, wanted: &str) -> Result<ValueDeserializer<'a>, QueryError> {
self.check()?;
match self.node {
Node::Scalar(values) if values.len() > 1 => Err(QueryError::msg(
"duplicate query parameter: more than one value was submitted for a \
single-valued field",
)),
Node::Scalar(values) => Ok(ValueDeserializer {
value: values.first().map_or("", String::as_str),
}),
other => Err(QueryError::msg(format!(
"expected {wanted}, found {} — the request used the bracketed form \
(`key[…]=`) for this parameter, so give the field a nested type or \
accept it as `serde_json::Value`",
other.kind()
))),
}
}
}
macro_rules! forward_scalar {
($($method:ident),* $(,)?) => {
$(
fn $method<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
self.as_value("a value")?.$method(visitor)
}
)*
};
}
impl<'de> de::Deserializer<'de> for NodeDeserializer<'_> {
type Error = QueryError;
fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
self.check()?;
match self.node {
Node::Scalar(values) if values.len() == 1 => {
visitor.visit_str(values.first().map_or("", String::as_str))
}
Node::Scalar(_) | Node::Seq(_) => self.deserialize_seq(visitor),
Node::Map(_) | Node::Conflict(_) => self.deserialize_map(visitor),
}
}
forward_scalar! {
deserialize_bool,
deserialize_i8, deserialize_i16, deserialize_i32, deserialize_i64, deserialize_i128,
deserialize_u8, deserialize_u16, deserialize_u32, deserialize_u64, deserialize_u128,
deserialize_f32, deserialize_f64,
deserialize_char, deserialize_str, deserialize_string,
deserialize_bytes, deserialize_byte_buf,
deserialize_identifier,
}
fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_some(self)
}
fn deserialize_unit<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
self.as_value("a value")?;
visitor.visit_unit()
}
fn deserialize_unit_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, QueryError> {
self.as_value("a value")?;
visitor.visit_unit()
}
fn deserialize_newtype_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, QueryError> {
visitor.visit_newtype_struct(self)
}
fn deserialize_seq<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
self.check()?;
match self.node {
Node::Scalar(values) => visitor.visit_seq(ValueSeq {
values: values.iter(),
}),
Node::Seq(entries) => visitor.visit_seq(NodeSeq {
entries: entries.iter(),
}),
Node::Map(entries) => visitor.visit_seq(PairSeq {
pairs: flatten_pairs(entries).into_iter(),
}),
Node::Conflict(message) => Err(QueryError::msg(message.clone())),
}
}
fn deserialize_tuple<V: Visitor<'de>>(
self,
_len: usize,
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_seq(visitor)
}
fn deserialize_tuple_struct<V: Visitor<'de>>(
self,
_name: &'static str,
_len: usize,
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_seq(visitor)
}
fn deserialize_map<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
self.check()?;
match self.node {
Node::Map(entries) => visitor.visit_map(NodeMap {
entries: entries.iter(),
value: None,
key: String::new(),
}),
Node::Seq(entries) => visitor.visit_map(IndexMap {
entries: entries.iter(),
value: None,
key: String::new(),
}),
other => Err(QueryError::msg(format!(
"expected an object, found {}",
other.kind()
))),
}
}
fn deserialize_struct<V: Visitor<'de>>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, QueryError> {
self.deserialize_map(visitor)
}
fn deserialize_enum<V: Visitor<'de>>(
self,
name: &'static str,
variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, QueryError> {
self.check()?;
match self.node {
Node::Scalar(_) => self
.as_value("a value")?
.deserialize_enum(name, variants, visitor),
Node::Map(entries) if entries.len() == 1 => match entries.iter().next() {
Some((variant, content)) => visitor.visit_enum(NodeVariant { variant, content }),
None => Err(QueryError::msg("expected an enum variant name")),
},
other => Err(QueryError::msg(format!(
"expected an enum variant name or a single-entry object, found {}",
other.kind()
))),
}
}
fn deserialize_ignored_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_unit()
}
}
enum PairValue<'a> {
Value(&'a str),
Node(&'a Node),
}
fn flatten_pairs(entries: &BTreeMap<String, Node>) -> Vec<(&str, PairValue<'_>)> {
let mut out = Vec::new();
for (key, node) in entries {
match node {
Node::Scalar(values) => {
out.extend(values.iter().map(|v| (key.as_str(), PairValue::Value(v))));
}
other => out.push((key.as_str(), PairValue::Node(other))),
}
}
out
}
struct ValueSeq<'a> {
values: std::slice::Iter<'a, String>,
}
impl<'de> SeqAccess<'de> for ValueSeq<'_> {
type Error = QueryError;
fn next_element_seed<T: DeserializeSeed<'de>>(
&mut self,
seed: T,
) -> Result<Option<T::Value>, QueryError> {
self.values
.next()
.map(|value| seed.deserialize(ValueDeserializer { value }))
.transpose()
}
fn size_hint(&self) -> Option<usize> {
Some(self.values.len())
}
}
struct NodeSeq<'a> {
entries: std::collections::btree_map::Iter<'a, SeqKey, Node>,
}
impl<'de> SeqAccess<'de> for NodeSeq<'_> {
type Error = QueryError;
fn next_element_seed<T: DeserializeSeed<'de>>(
&mut self,
seed: T,
) -> Result<Option<T::Value>, QueryError> {
let Some((key, node)) = self.entries.next() else {
return Ok(None);
};
seed.deserialize(NodeDeserializer { node })
.map(Some)
.map_err(|err| err.with_segment(sanitize_key(&key.raw)))
}
fn size_hint(&self) -> Option<usize> {
Some(self.entries.len())
}
}
struct PairSeq<'a> {
pairs: std::vec::IntoIter<(&'a str, PairValue<'a>)>,
}
impl<'de> SeqAccess<'de> for PairSeq<'_> {
type Error = QueryError;
fn next_element_seed<T: DeserializeSeed<'de>>(
&mut self,
seed: T,
) -> Result<Option<T::Value>, QueryError> {
let Some((key, value)) = self.pairs.next() else {
return Ok(None);
};
seed.deserialize(PairDeserializer { key, value })
.map(Some)
.map_err(|err| err.with_segment(sanitize_key(key)))
}
fn size_hint(&self) -> Option<usize> {
Some(self.pairs.len())
}
}
struct PairDeserializer<'a> {
key: &'a str,
value: PairValue<'a>,
}
impl<'de> de::Deserializer<'de> for PairDeserializer<'_> {
type Error = QueryError;
fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, QueryError> {
visitor.visit_seq(PairElements {
key: Some(self.key),
value: Some(self.value),
})
}
serde::forward_to_deserialize_any! {
bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string
bytes byte_buf option unit unit_struct newtype_struct seq tuple
tuple_struct map struct enum identifier ignored_any
}
}
struct PairElements<'a> {
key: Option<&'a str>,
value: Option<PairValue<'a>>,
}
impl<'de> SeqAccess<'de> for PairElements<'_> {
type Error = QueryError;
fn next_element_seed<T: DeserializeSeed<'de>>(
&mut self,
seed: T,
) -> Result<Option<T::Value>, QueryError> {
if let Some(key) = self.key.take() {
return seed.deserialize(ValueDeserializer { value: key }).map(Some);
}
match self.value.take() {
Some(PairValue::Value(value)) => {
seed.deserialize(ValueDeserializer { value }).map(Some)
}
Some(PairValue::Node(node)) => seed.deserialize(NodeDeserializer { node }).map(Some),
None => Ok(None),
}
}
}
struct NodeMap<'a> {
entries: std::collections::btree_map::Iter<'a, String, Node>,
value: Option<&'a Node>,
key: String,
}
impl<'de> MapAccess<'de> for NodeMap<'_> {
type Error = QueryError;
fn next_key_seed<K: DeserializeSeed<'de>>(
&mut self,
seed: K,
) -> Result<Option<K::Value>, QueryError> {
let Some((key, node)) = self.entries.next() else {
return Ok(None);
};
self.value = Some(node);
self.key = key.clone();
seed.deserialize(ValueDeserializer { value: key }).map(Some)
}
fn next_value_seed<V: DeserializeSeed<'de>>(
&mut self,
seed: V,
) -> Result<V::Value, QueryError> {
let node = self.value.take().ok_or_else(|| {
QueryError::msg("internal error: query map value requested before its key")
})?;
seed.deserialize(NodeDeserializer { node })
.map_err(|err| err.with_segment(sanitize_key(&self.key)))
}
fn size_hint(&self) -> Option<usize> {
Some(self.entries.len())
}
}
struct IndexMap<'a> {
entries: std::collections::btree_map::Iter<'a, SeqKey, Node>,
value: Option<&'a Node>,
key: String,
}
impl<'de> MapAccess<'de> for IndexMap<'_> {
type Error = QueryError;
fn next_key_seed<K: DeserializeSeed<'de>>(
&mut self,
seed: K,
) -> Result<Option<K::Value>, QueryError> {
let Some((key, node)) = self.entries.next() else {
return Ok(None);
};
self.value = Some(node);
self.key = key.raw.clone();
seed.deserialize(ValueDeserializer { value: &self.key })
.map(Some)
}
fn next_value_seed<V: DeserializeSeed<'de>>(
&mut self,
seed: V,
) -> Result<V::Value, QueryError> {
let node = self.value.take().ok_or_else(|| {
QueryError::msg("internal error: query map value requested before its key")
})?;
let key = sanitize_key(&self.key);
seed.deserialize(NodeDeserializer { node })
.map_err(|err| err.with_segment(key))
}
fn size_hint(&self) -> Option<usize> {
Some(self.entries.len())
}
}
struct NodeVariant<'a> {
variant: &'a str,
content: &'a Node,
}
impl<'de, 'a> EnumAccess<'de> for NodeVariant<'a> {
type Error = QueryError;
type Variant = NodeContent<'a>;
fn variant_seed<V: DeserializeSeed<'de>>(
self,
seed: V,
) -> Result<(V::Value, NodeContent<'a>), QueryError> {
let variant = seed.deserialize(ValueDeserializer {
value: self.variant,
})?;
Ok((
variant,
NodeContent {
content: self.content,
},
))
}
}
struct NodeContent<'a> {
content: &'a Node,
}
impl<'de> VariantAccess<'de> for NodeContent<'_> {
type Error = QueryError;
fn unit_variant(self) -> Result<(), QueryError> {
NodeDeserializer { node: self.content }.as_value("a value")?;
Ok(())
}
fn newtype_variant_seed<T: DeserializeSeed<'de>>(
self,
seed: T,
) -> Result<T::Value, QueryError> {
seed.deserialize(NodeDeserializer { node: self.content })
}
fn tuple_variant<V: Visitor<'de>>(
self,
_len: usize,
visitor: V,
) -> Result<V::Value, QueryError> {
de::Deserializer::deserialize_seq(NodeDeserializer { node: self.content }, visitor)
}
fn struct_variant<V: Visitor<'de>>(
self,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, QueryError> {
de::Deserializer::deserialize_map(NodeDeserializer { node: self.content }, visitor)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
#[derive(Debug, Deserialize, PartialEq)]
struct Filter {
status: String,
limit: Option<u32>,
}
#[derive(Debug, Deserialize, PartialEq)]
struct Item {
sku: String,
qty: u32,
}
#[derive(Debug, Deserialize, PartialEq, Default)]
struct Args {
q: Option<String>,
page: Option<u32>,
flag: Option<bool>,
tags: Option<Vec<String>>,
filter: Option<Filter>,
items: Option<Vec<Item>>,
}
fn args(query: &str) -> Args {
from_query_str(query).expect("decodes")
}
#[test]
fn flat_scalars_match_the_previous_urlencoded_behaviour() {
let out = args("q=foo&page=2&flag=true");
assert_eq!(out.q.as_deref(), Some("foo"));
assert_eq!(out.page, Some(2));
assert_eq!(out.flag, Some(true));
}
#[test]
fn absent_keys_are_none_and_an_empty_query_decodes() {
assert_eq!(args(""), Args::default());
}
#[test]
fn a_present_but_empty_optional_still_visits_some() {
assert_eq!(args("q=").q.as_deref(), Some(""));
assert!(
from_query_str::<Args>("page=").is_err(),
"an empty value for a numeric field is still a coercion failure"
);
}
#[test]
fn repeated_keys_fill_a_sequence_but_a_duplicated_scalar_is_rejected() {
assert_eq!(
args("tags=a&tags=b").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned()]
);
let err = from_query_str::<Args>("q=first&q=second").expect_err("duplicate scalar");
assert!(
err.to_string().contains("duplicate query parameter"),
"{err}"
);
}
#[test]
fn distinct_index_spellings_stay_distinct() {
#[derive(Debug, Deserialize)]
struct Counts {
counts: std::collections::HashMap<String, u32>,
}
assert_eq!(
args("tags[0]=a&tags[00]=b").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned()]
);
let out: Counts = from_query_str("counts[00]=5&counts[0]=6").expect("decodes");
assert_eq!(
out.counts.get("00"),
Some(&5),
"spelling preserved: {out:?}"
);
assert_eq!(out.counts.get("0"), Some(&6));
let out: Counts = from_query_str("counts[99999999999999999999]=7").expect("decodes");
assert_eq!(out.counts.get("99999999999999999999"), Some(&7));
assert!(from_query_str::<Args>("tags[0]=a&tags[0]=b").is_err());
}
#[test]
fn a_custom_deserializer_message_is_discarded_not_merely_bounded() {
fn picky<'de, D>(deserializer: D) -> Result<String, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw = String::deserialize(deserializer)?;
Err(serde::de::Error::custom(format!("invalid token {raw}")))
}
#[derive(Debug, Deserialize)]
struct Guarded {
#[serde(deserialize_with = "picky")]
#[allow(
dead_code,
reason = "the helper always errors; the field is never built"
)]
token: String,
}
let err = from_query_str::<Guarded>("token=SUPERSECRET").expect_err("helper rejects");
let rendered = err.to_string();
assert!(!rendered.contains("SUPERSECRET"), "leaked: {rendered}");
assert!(
!rendered.contains("invalid token"),
"retained content: {rendered}"
);
assert!(rendered.contains("token"), "{rendered}");
let long = "S3CRET".repeat(60);
let err = from_query_str::<Guarded>(&format!("token={long}")).expect_err("helper rejects");
assert!(!err.to_string().contains("S3CRET"), "{err}");
}
#[test]
fn a_unit_variant_still_validates_the_node_it_claims() {
#[derive(Debug, Deserialize, PartialEq)]
#[serde(rename_all = "lowercase")]
enum Mode {
Asc,
}
#[derive(Debug, Deserialize)]
struct Sorted {
mode: Mode,
}
assert_eq!(
from_query_str::<Sorted>("mode[asc]=1")
.expect("decodes")
.mode,
Mode::Asc
);
assert!(
from_query_str::<Sorted>("mode[asc][unexpected]=value").is_err(),
"an attached payload must not be silently discarded"
);
assert!(from_query_str::<Sorted>("mode=asc").is_ok());
}
#[test]
fn serde_generated_errors_never_echo_request_text() {
#[derive(Debug, Deserialize)]
#[serde(rename_all = "lowercase")]
enum Sort {
Asc,
}
#[derive(Debug, Deserialize)]
struct Sorted {
#[allow(dead_code)]
sort: Sort,
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Strict {
#[allow(dead_code)]
q: String,
}
#[derive(Debug, Deserialize)]
struct Nested {
#[allow(dead_code)]
filter: std::collections::HashMap<String, String>,
}
let err = from_query_str::<Sorted>("sort=SUPERSECRET").expect_err("unknown variant");
assert!(!err.to_string().contains("SUPERSECRET"), "{err}");
let err = from_query_str::<Strict>("q=x&SUPERSECRET=1").expect_err("unknown field");
assert!(!err.to_string().contains("SUPERSECRET"), "{err}");
let err =
from_query_str::<Nested>("filter[a]=x&filter=SUPERSECRET").expect_err("shape conflict");
assert!(!err.to_string().contains("SUPERSECRET"), "{err}");
}
#[test]
fn a_unit_field_still_validates_the_node_it_claims() {
#[derive(Debug, Deserialize)]
struct WithUnit {
#[allow(dead_code)]
flag: (),
}
assert!(
from_query_str::<WithUnit>("flag=x&flag[y]=z").is_err(),
"a claimed conflicting key must fail even for a unit field"
);
assert!(
from_query_str::<WithUnit>("flag[y]=z").is_err(),
"a bracketed container is not a value"
);
assert!(from_query_str::<WithUnit>("flag=x").is_ok());
}
#[test]
fn pair_sequence_errors_sanitize_their_key() {
let err = from_query_str::<Vec<(String, u32)>>("\nATTACKER=not-a-number")
.expect_err("coercion failure");
let rendered = err.to_string();
assert!(
!rendered.contains('\n'),
"control chars stripped: {rendered:?}"
);
let long = format!("{}=not-a-number", "k".repeat(500));
let err = from_query_str::<Vec<(String, u32)>>(&long).expect_err("coercion failure");
assert!(
err.to_string().chars().count() < 200,
"key bounded: {}",
err.to_string().len()
);
}
#[test]
fn a_container_may_mix_numeric_and_named_keys() {
#[derive(Debug, Deserialize)]
struct Dynamic {
filter: std::collections::HashMap<String, String>,
}
let out: Dynamic = from_query_str("filter[0]=zero&filter[name]=value").expect("decodes");
assert_eq!(out.filter.get("0").map(String::as_str), Some("zero"));
assert_eq!(out.filter.get("name").map(String::as_str), Some("value"));
let out: Dynamic = from_query_str("filter[name]=value&filter[0]=zero").expect("decodes");
assert_eq!(out.filter.get("0").map(String::as_str), Some("zero"));
assert_eq!(out.filter.get("name").map(String::as_str), Some("value"));
assert_eq!(
args("tags[0]=a&tags[1]=b").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned()]
);
}
#[test]
fn errors_never_echo_the_submitted_value() {
let err = from_query_str::<Args>("page=SUPERSECRET").expect_err("coercion failure");
let rendered = err.to_string();
assert!(
!rendered.contains("SUPERSECRET"),
"value leaked: {rendered}"
);
assert!(rendered.contains("page"), "field path kept: {rendered}");
}
#[test]
fn error_text_bounds_and_cleans_attacker_supplied_key_text() {
let long = "k".repeat(500);
let rendered = sanitize_key(&long);
assert!(
rendered.chars().count() <= 49,
"bounded: {}",
rendered.len()
);
assert_eq!(sanitize_key("a\nb"), "a.b");
}
#[test]
fn a_conflicting_key_the_target_ignores_is_still_ignored() {
let out = args("q=ok&utm=1&utm[src]=x");
assert_eq!(out.q.as_deref(), Some("ok"));
let deep = format!("q=ok&junk{}=1", "[x]".repeat(MAX_DEPTH));
assert_eq!(args(&deep).q.as_deref(), Some("ok"));
}
#[test]
fn an_all_digit_object_key_still_lands_in_a_map_target() {
let out: std::collections::HashMap<String, std::collections::HashMap<String, u32>> =
from_query_str("counts[2024]=5").expect("decodes");
assert_eq!(out["counts"]["2024"], 5);
}
#[test]
fn append_and_indexed_forms_fill_a_sequence() {
assert_eq!(
args("tags[]=a&tags[]=b").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned()]
);
assert_eq!(
args("tags[1]=b&tags[0]=a").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned()]
);
}
#[test]
fn gapped_indices_compact_in_ascending_order() {
assert_eq!(
args("tags[4]=c&tags[0]=a&tags[2]=b").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned(), "c".to_owned()]
);
}
#[test]
fn an_absurd_index_sorts_last_without_preallocating() {
assert_eq!(
args("tags[99999999999999999999]=last&tags[3]=first")
.tags
.unwrap(),
vec!["first".to_owned(), "last".to_owned()]
);
}
#[test]
fn appends_continue_after_the_highest_explicit_index() {
assert_eq!(
args("tags[]=a&tags[5]=b&tags[]=c").tags.unwrap(),
vec!["a".to_owned(), "b".to_owned(), "c".to_owned()]
);
}
#[test]
fn nested_objects_decode() {
assert_eq!(
args("filter[status]=open&filter[limit]=5").filter.unwrap(),
Filter {
status: "open".to_owned(),
limit: Some(5),
}
);
}
#[test]
fn arrays_of_objects_decode() {
let out = args("items[0][sku]=A-1&items[0][qty]=2&items[1][sku]=B-2&items[1][qty]=3");
assert_eq!(
out.items.unwrap(),
vec![
Item {
sku: "A-1".to_owned(),
qty: 2
},
Item {
sku: "B-2".to_owned(),
qty: 3
},
]
);
}
#[test]
fn percent_encoded_brackets_parse_identically() {
assert_eq!(
args("filter%5Bstatus%5D=open").filter.unwrap().status,
"open"
);
}
#[test]
fn plus_is_decoded_as_a_space() {
assert_eq!(args("q=hello+world").q.as_deref(), Some("hello world"));
}
#[test]
fn malformed_bracket_keys_stay_literal() {
assert_eq!(parse_key("weird[unclosed"), None);
assert_eq!(parse_key("a[b][c"), None);
assert_eq!(parse_key("a[[b]]"), None);
assert_eq!(parse_key("plain"), None);
assert_eq!(args("weird[unclosed=1&q=ok").q.as_deref(), Some("ok"));
}
#[test]
fn bracket_keys_parse_into_segments() {
assert_eq!(
parse_key("items[0][sku]"),
Some((
"items",
vec![
Segment::Index {
position: 0,
raw: "0"
},
Segment::Key("sku")
]
))
);
assert_eq!(parse_key("tags[]"), Some(("tags", vec![Segment::Append])));
}
#[test]
fn shape_conflicts_are_rejected() {
let scalar_then_object = from_query_str::<Args>("filter=flat&filter[status]=open");
assert!(scalar_then_object.is_err());
let object_then_scalar = from_query_str::<Args>("filter[status]=open&filter=flat");
assert!(object_then_scalar.is_err());
let seq_then_object = from_query_str::<Args>("tags[0]=a&tags[name]=b");
assert!(seq_then_object.is_err());
}
#[test]
fn nesting_deeper_than_the_cap_is_rejected_for_a_claimed_key() {
let deep = format!("filter{}=1", "[x]".repeat(MAX_DEPTH));
let err = from_query_str::<Args>(&deep).expect_err("depth-capped");
assert!(err.to_string().contains("deeper"), "{err}");
let ok = format!("q=x&junk{}=1", "[x]".repeat(MAX_DEPTH - 2));
assert!(from_query_str::<Args>(&ok).is_ok());
}
#[test]
fn errors_name_the_failing_field_path() {
let err = from_query_str::<Args>("filter[status]=open&filter[limit]=nope")
.expect_err("coercion failure");
assert!(err.to_string().starts_with("filter.limit:"), "{err}");
}
#[test]
fn shape_conflicts_surface_only_when_the_target_claims_the_key() {
let err = from_query_str::<Args>("filter=flat&filter[status]=open")
.expect_err("claimed conflicting key");
assert!(err.to_string().contains("used as both"), "{err}");
}
#[test]
fn untyped_targets_decode_through_deserialize_any() {
let out: serde_json::Value =
from_query_str("q=foo&tags=a&tags=b&filter[status]=open").expect("decodes");
assert_eq!(out["q"], "foo");
assert_eq!(out["tags"], serde_json::json!(["a", "b"]));
assert_eq!(out["filter"]["status"], "open");
}
#[test]
fn map_targets_keep_working() {
let out: std::collections::HashMap<String, String> =
from_query_str("a=1&b=2").expect("decodes");
assert_eq!(out.get("a").map(String::as_str), Some("1"));
assert_eq!(out.get("b").map(String::as_str), Some("2"));
}
#[test]
fn pair_sequence_targets_keep_working() {
let out: Vec<(String, String)> = from_query_str("a=1&a=2&b=3").expect("decodes");
assert_eq!(
out,
vec![
("a".to_owned(), "1".to_owned()),
("a".to_owned(), "2".to_owned()),
("b".to_owned(), "3".to_owned()),
]
);
}
#[test]
fn unit_enum_variants_decode_from_a_bare_value() {
#[derive(Debug, Deserialize, PartialEq)]
#[serde(rename_all = "lowercase")]
enum Sort {
Asc,
Desc,
}
#[derive(Debug, Deserialize)]
struct Sorted {
sort: Sort,
}
let out: Sorted = from_query_str("sort=desc").expect("decodes");
assert_eq!(out.sort, Sort::Desc);
}
#[test]
fn unknown_keys_are_ignored_including_nested_ones() {
let out = args("q=ok&unknown[deep][deeper]=1&other=2");
assert_eq!(out.q.as_deref(), Some("ok"));
}
#[test]
fn a_missing_required_field_still_fails() {
assert!(from_query_str::<Filter>("limit=5").is_err());
}
#[test]
fn serde_renames_are_honored_because_serde_resolves_the_key() {
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct Renamed {
word_count: u32,
}
let out: Renamed = from_query_str("wordCount=7").expect("decodes");
assert_eq!(out.word_count, 7);
}
}