use std::str::FromStr;
use serde::de::{self, DeserializeSeed, Deserializer, EnumAccess, VariantAccess, Visitor};
use crate::error::{Error, Result};
use crate::value::Scalar;
#[derive(Clone)]
enum KeyText<'de> {
Borrowed(&'de str),
Local(Scalar),
}
pub(crate) struct KeyDeserializer<'de> {
text: KeyText<'de>,
}
impl<'de> KeyDeserializer<'de> {
pub(crate) fn borrowed(value: &'de str) -> Self {
Self {
text: KeyText::Borrowed(value),
}
}
pub(crate) fn local(value: Scalar) -> Self {
Self {
text: KeyText::Local(value),
}
}
fn text(&self) -> &str {
match &self.text {
KeyText::Borrowed(s) => s,
KeyText::Local(s) => s.as_str(),
}
}
fn visit_text<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
match self.text {
KeyText::Borrowed(s) => visitor.visit_borrowed_str(s),
KeyText::Local(s) => visitor.visit_str(&s),
}
}
fn parse<T: FromStr>(&self, type_name: &'static str) -> Result<T> {
let s = self.text();
s.parse::<T>().map_err(|_| {
<Error as de::Error>::custom(format!("failed to parse map key '{s}' as {type_name}"))
})
}
}
impl<'de> Deserializer<'de> for KeyDeserializer<'de> {
type Error = Error;
fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
self.visit_text(visitor)
}
fn deserialize_str<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
match self.text {
KeyText::Borrowed(s) => visitor.visit_borrowed_str(s),
KeyText::Local(s) => {
if s.is_heap_allocated() {
visitor.visit_string(s.into_string())
} else {
visitor.visit_str(&s)
}
}
}
}
fn deserialize_string<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
match self.text {
KeyText::Borrowed(s) => visitor.visit_borrowed_str(s),
KeyText::Local(s) => visitor.visit_string(s.into_string()),
}
}
fn deserialize_bool<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_bool(self.parse("bool")?)
}
fn deserialize_i8<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i8(self.parse("i8")?)
}
fn deserialize_i16<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i16(self.parse("i16")?)
}
fn deserialize_i32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i32(self.parse("i32")?)
}
fn deserialize_i64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i64(self.parse("i64")?)
}
fn deserialize_i128<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_i128(self.parse("i128")?)
}
fn deserialize_u8<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u8(self.parse("u8")?)
}
fn deserialize_u16<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u16(self.parse("u16")?)
}
fn deserialize_u32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u32(self.parse("u32")?)
}
fn deserialize_u64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u64(self.parse("u64")?)
}
fn deserialize_u128<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_u128(self.parse("u128")?)
}
fn deserialize_f32<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_f32(self.parse("f32")?)
}
fn deserialize_f64<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_f64(self.parse("f64")?)
}
fn deserialize_char<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
let s = self.text();
let mut chars = s.chars();
match (chars.next(), chars.next()) {
(Some(c), None) => visitor.visit_char(c),
_ => Err(<Error as de::Error>::custom(format!(
"expected single character map key, got '{s}'"
))),
}
}
fn deserialize_newtype_struct<V: Visitor<'de>>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value> {
visitor.visit_newtype_struct(self)
}
fn deserialize_enum<V: Visitor<'de>>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value> {
visitor.visit_enum(KeyEnum { text: self.text })
}
fn deserialize_identifier<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
self.visit_text(visitor)
}
fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value> {
visitor.visit_some(self)
}
serde::forward_to_deserialize_any! {
bytes byte_buf unit unit_struct seq tuple tuple_struct map
struct ignored_any
}
}
struct KeyEnum<'de> {
text: KeyText<'de>,
}
impl<'de> EnumAccess<'de> for KeyEnum<'de> {
type Error = Error;
type Variant = KeyUnitVariant;
fn variant_seed<V: DeserializeSeed<'de>>(self, seed: V) -> Result<(V::Value, Self::Variant)> {
let variant = seed.deserialize(KeyDeserializer { text: self.text })?;
Ok((variant, KeyUnitVariant))
}
}
struct KeyUnitVariant;
impl<'de> VariantAccess<'de> for KeyUnitVariant {
type Error = Error;
fn unit_variant(self) -> Result<()> {
Ok(())
}
fn newtype_variant_seed<T: DeserializeSeed<'de>>(self, _seed: T) -> Result<T::Value> {
Err(<Error as de::Error>::invalid_type(
de::Unexpected::UnitVariant,
&"newtype variant",
))
}
fn tuple_variant<V: Visitor<'de>>(self, _len: usize, _visitor: V) -> Result<V::Value> {
Err(<Error as de::Error>::invalid_type(
de::Unexpected::UnitVariant,
&"tuple variant",
))
}
fn struct_variant<V: Visitor<'de>>(
self,
_fields: &'static [&'static str],
_visitor: V,
) -> Result<V::Value> {
Err(<Error as de::Error>::invalid_type(
de::Unexpected::UnitVariant,
&"struct variant",
))
}
}