use serde::de::{
self, Deserialize, DeserializeSeed, Deserializer, EnumAccess, MapAccess, SeqAccess,
VariantAccess, Visitor,
};
pub const DEFAULT_DESERIALIZE_DEPTH: usize = 128;
fn too_deep<E: de::Error>() -> E {
E::custom("deserialization exceeded the configured recursion-depth limit")
}
pub struct DepthLimited<D> {
inner: D,
budget: usize,
}
impl<D> DepthLimited<D> {
pub fn new(inner: D, max_depth: usize) -> Self {
Self {
inner,
budget: max_depth,
}
}
}
pub fn from_deserializer<'de, T, D>(deserializer: D, max_depth: usize) -> Result<T, D::Error>
where
T: Deserialize<'de>,
D: Deserializer<'de>,
{
T::deserialize(DepthLimited::new(deserializer, max_depth))
}
macro_rules! forward_scalar {
($($method:ident),* $(,)?) => {
$(
fn $method<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
self.inner.$method(visitor)
}
)*
};
}
impl<'de, D> Deserializer<'de> for DepthLimited<D>
where
D: Deserializer<'de>,
{
type Error = D::Error;
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_unit, deserialize_identifier,
}
fn deserialize_unit_struct<V>(
self,
name: &'static str,
visitor: V,
) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
self.inner.deserialize_unit_struct(name, visitor)
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget;
self.inner.deserialize_option(DepthVisitor {
inner: visitor,
budget,
})
}
fn deserialize_newtype_struct<V>(
self,
name: &'static str,
visitor: V,
) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget;
self.inner.deserialize_newtype_struct(
name,
DepthVisitor {
inner: visitor,
budget,
},
)
}
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_seq(DepthVisitor {
inner: visitor,
budget,
})
}
fn deserialize_tuple<V>(self, len: usize, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_tuple(
len,
DepthVisitor {
inner: visitor,
budget,
},
)
}
fn deserialize_tuple_struct<V>(
self,
name: &'static str,
len: usize,
visitor: V,
) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_tuple_struct(
name,
len,
DepthVisitor {
inner: visitor,
budget,
},
)
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_map(DepthVisitor {
inner: visitor,
budget,
})
}
fn deserialize_struct<V>(
self,
name: &'static str,
fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_struct(
name,
fields,
DepthVisitor {
inner: visitor,
budget,
},
)
}
fn deserialize_enum<V>(
self,
name: &'static str,
variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_enum(
name,
variants,
DepthVisitor {
inner: visitor,
budget,
},
)
}
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_any(DepthVisitor {
inner: visitor,
budget,
})
}
fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value, D::Error>
where
V: Visitor<'de>,
{
let budget = self.budget.checked_sub(1).ok_or_else(too_deep)?;
self.inner.deserialize_ignored_any(DepthVisitor {
inner: visitor,
budget,
})
}
fn is_human_readable(&self) -> bool {
self.inner.is_human_readable()
}
}
struct DepthVisitor<V> {
inner: V,
budget: usize,
}
impl<'de, V> Visitor<'de> for DepthVisitor<V>
where
V: Visitor<'de>,
{
type Value = V::Value;
fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
self.inner.expecting(formatter)
}
fn visit_bool<E: de::Error>(self, v: bool) -> Result<Self::Value, E> {
self.inner.visit_bool(v)
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<Self::Value, E> {
self.inner.visit_i64(v)
}
fn visit_i128<E: de::Error>(self, v: i128) -> Result<Self::Value, E> {
self.inner.visit_i128(v)
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<Self::Value, E> {
self.inner.visit_u64(v)
}
fn visit_u128<E: de::Error>(self, v: u128) -> Result<Self::Value, E> {
self.inner.visit_u128(v)
}
fn visit_f64<E: de::Error>(self, v: f64) -> Result<Self::Value, E> {
self.inner.visit_f64(v)
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
self.inner.visit_str(v)
}
fn visit_borrowed_str<E: de::Error>(self, v: &'de str) -> Result<Self::Value, E> {
self.inner.visit_borrowed_str(v)
}
fn visit_string<E: de::Error>(self, v: String) -> Result<Self::Value, E> {
self.inner.visit_string(v)
}
fn visit_bytes<E: de::Error>(self, v: &[u8]) -> Result<Self::Value, E> {
self.inner.visit_bytes(v)
}
fn visit_borrowed_bytes<E: de::Error>(self, v: &'de [u8]) -> Result<Self::Value, E> {
self.inner.visit_borrowed_bytes(v)
}
fn visit_byte_buf<E: de::Error>(self, v: Vec<u8>) -> Result<Self::Value, E> {
self.inner.visit_byte_buf(v)
}
fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
self.inner.visit_none()
}
fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
self.inner.visit_unit()
}
fn visit_some<D2>(self, deserializer: D2) -> Result<Self::Value, D2::Error>
where
D2: Deserializer<'de>,
{
self.inner.visit_some(DepthLimited {
inner: deserializer,
budget: self.budget,
})
}
fn visit_newtype_struct<D2>(self, deserializer: D2) -> Result<Self::Value, D2::Error>
where
D2: Deserializer<'de>,
{
self.inner.visit_newtype_struct(DepthLimited {
inner: deserializer,
budget: self.budget,
})
}
fn visit_seq<A>(self, seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
self.inner.visit_seq(DepthSeqAccess {
inner: seq,
budget: self.budget,
})
}
fn visit_map<A>(self, map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
self.inner.visit_map(DepthMapAccess {
inner: map,
budget: self.budget,
})
}
fn visit_enum<A>(self, data: A) -> Result<Self::Value, A::Error>
where
A: EnumAccess<'de>,
{
self.inner.visit_enum(DepthEnumAccess {
inner: data,
budget: self.budget,
})
}
}
struct DepthSeqAccess<A> {
inner: A,
budget: usize,
}
impl<'de, A> SeqAccess<'de> for DepthSeqAccess<A>
where
A: SeqAccess<'de>,
{
type Error = A::Error;
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>, A::Error>
where
T: DeserializeSeed<'de>,
{
self.inner.next_element_seed(DepthSeed {
inner: seed,
budget: self.budget,
})
}
fn size_hint(&self) -> Option<usize> {
self.inner.size_hint()
}
}
struct DepthMapAccess<A> {
inner: A,
budget: usize,
}
impl<'de, A> MapAccess<'de> for DepthMapAccess<A>
where
A: MapAccess<'de>,
{
type Error = A::Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>, A::Error>
where
K: DeserializeSeed<'de>,
{
self.inner.next_key_seed(DepthSeed {
inner: seed,
budget: self.budget,
})
}
fn next_value_seed<Vv>(&mut self, seed: Vv) -> Result<Vv::Value, A::Error>
where
Vv: DeserializeSeed<'de>,
{
self.inner.next_value_seed(DepthSeed {
inner: seed,
budget: self.budget,
})
}
fn size_hint(&self) -> Option<usize> {
self.inner.size_hint()
}
}
struct DepthEnumAccess<A> {
inner: A,
budget: usize,
}
impl<'de, A> EnumAccess<'de> for DepthEnumAccess<A>
where
A: EnumAccess<'de>,
{
type Error = A::Error;
type Variant = DepthVariantAccess<A::Variant>;
fn variant_seed<Vs>(self, seed: Vs) -> Result<(Vs::Value, Self::Variant), A::Error>
where
Vs: DeserializeSeed<'de>,
{
let budget = self.budget;
let (value, variant) = self.inner.variant_seed(DepthSeed {
inner: seed,
budget,
})?;
Ok((
value,
DepthVariantAccess {
inner: variant,
budget,
},
))
}
}
struct DepthVariantAccess<A> {
inner: A,
budget: usize,
}
impl<'de, A> VariantAccess<'de> for DepthVariantAccess<A>
where
A: VariantAccess<'de>,
{
type Error = A::Error;
fn unit_variant(self) -> Result<(), A::Error> {
self.inner.unit_variant()
}
fn newtype_variant_seed<T>(self, seed: T) -> Result<T::Value, A::Error>
where
T: DeserializeSeed<'de>,
{
self.inner.newtype_variant_seed(DepthSeed {
inner: seed,
budget: self.budget,
})
}
fn tuple_variant<Vv>(self, len: usize, visitor: Vv) -> Result<Vv::Value, A::Error>
where
Vv: Visitor<'de>,
{
self.inner.tuple_variant(
len,
DepthVisitor {
inner: visitor,
budget: self.budget,
},
)
}
fn struct_variant<Vv>(
self,
fields: &'static [&'static str],
visitor: Vv,
) -> Result<Vv::Value, A::Error>
where
Vv: Visitor<'de>,
{
self.inner.struct_variant(
fields,
DepthVisitor {
inner: visitor,
budget: self.budget,
},
)
}
}
struct DepthSeed<S> {
inner: S,
budget: usize,
}
impl<'de, S> DeserializeSeed<'de> for DepthSeed<S>
where
S: DeserializeSeed<'de>,
{
type Value = S::Value;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
self.inner.deserialize(DepthLimited {
inner: deserializer,
budget: self.budget,
})
}
}