use serde::de::{
DeserializeOwned, DeserializeSeed, EnumAccess, Error, Expected, MapAccess, SeqAccess,
VariantAccess, Visitor,
};
use serde::{Deserialize, Deserializer};
use std::slice::Iter;
use yaml_rust2::{Yaml, YamlLoader};
pub fn from_yaml<'de, T: Deserialize<'de>>(yaml: &'de Yaml) -> Result<T, serde::de::value::Error> {
let deserializer = &mut YamlDeserializer::new(yaml);
return T::deserialize(deserializer);
}
pub fn from_str<T: DeserializeOwned>(data: &str) -> Result<Vec<T>, serde::de::value::Error> {
let mut ret = Vec::new();
let yamls = YamlLoader::load_from_str(data).map_err(serde::de::Error::custom)?;
for yaml in &yamls {
ret.push(from_yaml::<T>(yaml)?);
}
return Ok(ret);
}
fn invalid(yaml: &Yaml, expected: &dyn Expected) -> serde::de::value::Error {
use serde::de::Unexpected;
return serde::de::value::Error::invalid_type(
match yaml {
Yaml::Real(s) | Yaml::String(s) => Unexpected::Str(s),
Yaml::Integer(i) => Unexpected::Signed(*i),
Yaml::Boolean(b) => Unexpected::Bool(*b),
Yaml::Array(_) => Unexpected::Seq,
Yaml::Hash(_) => Unexpected::Map,
Yaml::Alias(_) => Unexpected::Other("alias"),
Yaml::Null => Unexpected::Unit,
Yaml::BadValue => Unexpected::Other("bad value"),
},
expected,
);
}
struct SeqAccesser<'de> {
seq: Iter<'de, Yaml>,
}
impl<'de> SeqAccess<'de> for SeqAccesser<'de> {
type Error = serde::de::value::Error;
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>, Self::Error>
where
T: DeserializeSeed<'de>,
{
let Some(yaml) = self.seq.next() else {
return Ok(None);
};
return seed.deserialize(&mut YamlDeserializer { yaml }).map(Some);
}
}
pub struct MapAcceser<'de> {
keys: hashlink::linked_hash_map::Keys<'de, Yaml, Yaml>,
values: hashlink::linked_hash_map::Values<'de, Yaml, Yaml>,
}
impl<'de, 'a> MapAccess<'de> for MapAcceser<'de> {
type Error = serde::de::value::Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>, Self::Error>
where
K: DeserializeSeed<'de>,
{
let Some(yaml) = self.keys.next() else {
return Ok(None);
};
return seed.deserialize(&mut YamlDeserializer { yaml }).map(Some);
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value, Self::Error>
where
V: DeserializeSeed<'de>,
{
let Some(yaml) = self.values.next() else {
return Err(serde::de::Error::custom("invalid call to next_value_seed"));
};
return seed.deserialize(&mut YamlDeserializer { yaml });
}
}
pub struct EnumAccesser<'de> {
tag: &'de Yaml,
value: &'de Yaml,
}
impl<'de> EnumAccess<'de> for EnumAccesser<'de> {
type Error = serde::de::value::Error;
type Variant = Self;
fn variant_seed<V>(self, seed: V) -> Result<(V::Value, Self::Variant), Self::Error>
where
V: DeserializeSeed<'de>,
{
return Ok((
seed.deserialize(&mut YamlDeserializer { yaml: &self.tag })?,
self,
));
}
}
impl<'de> VariantAccess<'de> for EnumAccesser<'de> {
type Error = serde::de::value::Error;
fn unit_variant(self) -> Result<(), Self::Error> {
let Yaml::Null = self.value else {
return Err(invalid(self.value, &"null"));
};
return Ok(());
}
fn newtype_variant_seed<T>(self, seed: T) -> Result<T::Value, Self::Error>
where
T: DeserializeSeed<'de>,
{
return seed.deserialize(&mut YamlDeserializer { yaml: self.value });
}
fn tuple_variant<V>(self, len: usize, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return YamlDeserializer { yaml: self.value }.deserialize_tuple(len, visitor);
}
fn struct_variant<V>(
self,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return YamlDeserializer { yaml: self.value }.deserialize_map(visitor);
}
}
pub struct YamlDeserializer<'de> {
yaml: &'de Yaml,
}
impl<'de> YamlDeserializer<'de> {
#[allow(clippy::should_implement_trait)]
pub fn new(yaml: &'de Yaml) -> Self {
return Self { yaml };
}
}
impl<'de> Deserializer<'de> for &mut YamlDeserializer<'de> {
type Error = serde::de::value::Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return match self.yaml {
Yaml::Real(v) => visitor.visit_f64(v.parse().map_err(Self::Error::custom)?),
Yaml::Hash(v) => visitor.visit_map(MapAcceser {
keys: v.keys(),
values: v.values(),
}),
Yaml::Array(v) => visitor.visit_seq(SeqAccesser { seq: v.iter() }),
Yaml::Integer(v) => visitor.visit_i64(v.clone()),
Yaml::String(v) => visitor.visit_string(v.clone()),
Yaml::Boolean(v) => visitor.visit_bool(v.clone()),
Yaml::Null => visitor.visit_none(),
_ => Err(Self::Error::custom("Unexpected yaml node type")),
};
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(b) = self.yaml.as_bool() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_bool(b);
}
fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_i8(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_i16(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_i32(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_i64(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_u8(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_u16(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_u32(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Some(i) = self.yaml.as_i64() else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_u64(i.try_into().map_err(Self::Error::custom)?);
}
fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let i: f32 = match &self.yaml {
Yaml::Real(v) => v.parse().map_err(Self::Error::custom)?,
Yaml::Integer(v) => *v as f32,
yaml => return Err(invalid(yaml, &visitor)),
};
return visitor.visit_f32(i);
}
fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let i: f64 = match &self.yaml {
Yaml::Real(v) => v.parse().map_err(Self::Error::custom)?,
Yaml::Integer(v) => *v as f64,
_ => return Err(invalid(self.yaml, &visitor)),
};
return visitor.visit_f64(i);
}
fn deserialize_char<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Yaml::String(s) = &self.yaml else {
return Err(invalid(self.yaml, &visitor));
};
if s.len() != 1 {
return Err(invalid(self.yaml, &visitor));
}
return visitor.visit_char(s.chars().next().unwrap());
}
fn deserialize_str<V>(self, _visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
unimplemented!("Deserialization of &str is not supported");
}
fn deserialize_string<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let s = match &self.yaml {
Yaml::Integer(i) => i.to_string(),
Yaml::Real(s) | Yaml::String(s) => s.clone(),
Yaml::Boolean(b) => b.to_string(),
yaml => return Err(invalid(yaml, &visitor)),
};
return visitor.visit_string(s);
}
fn deserialize_bytes<V>(self, _visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
unimplemented!("Deserialization of bytes is not supported");
}
fn deserialize_byte_buf<V>(self, _visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
unimplemented!("Deserialization of byte buffer is not supported");
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return match self.yaml {
Yaml::Null => visitor.visit_none(),
yaml => visitor.visit_some(&mut YamlDeserializer { yaml }),
};
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Yaml::Null = &self.yaml else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_unit();
}
fn deserialize_unit_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return self.deserialize_unit(visitor);
}
fn deserialize_newtype_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return visitor.visit_newtype_struct(self);
}
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Yaml::Array(v) = &self.yaml else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_seq(SeqAccesser { seq: v.iter() });
}
fn deserialize_tuple<V>(self, _visitor: usize, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return self.deserialize_seq(visitor);
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
_len: usize,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return self.deserialize_seq(visitor);
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Yaml::Hash(v) = &self.yaml else {
return Err(invalid(self.yaml, &visitor));
};
return visitor.visit_map(MapAcceser {
keys: v.keys(),
values: v.values(),
});
}
fn deserialize_struct<V>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
self.deserialize_map(visitor)
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
let Yaml::Hash(v) = &self.yaml else {
return Err(invalid(self.yaml, &visitor));
};
if v.len() != 1 {
return Err(invalid(self.yaml, &visitor));
}
let (tag, value) = v.into_iter().next().unwrap();
return visitor.visit_enum(EnumAccesser { tag, value });
}
fn deserialize_identifier<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return self.deserialize_string(visitor);
}
fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
return self.deserialize_any(visitor);
}
}
#[cfg(test)]
mod tests {
use crate::Yaml as YamlWrapper;
use crate::de::from_str;
use serde::Deserialize;
use yaml_rust2::Yaml;
macro_rules! test {
($type:ty, $expected:expr, $data:literal) => {{
let result: $type = from_str($data).unwrap().remove(0);
assert_eq!($expected, result);
}};
}
#[test]
fn should_work() {
test!((), (), "null");
test!((), (), "~");
test!(char, 'a', "a");
test!(i8, 123, "123");
test!(i16, 123, "123");
test!(i32, 123, "123");
test!(i64, 123, "123");
test!(u8, 123, "123");
test!(u16, 123, "123");
test!(u32, 123, "123");
test!(u64, 123, "123");
test!(i64, -123, "-123");
test!(f64, 123.341, "123.341");
test!(f64, -123.341, "-123.341");
test!(f32, 123.341, "123.341");
test!(f32, -123.341, "-123.341");
test!(f64, 0.0, "0");
test!(f64, 0.0, "-0");
test!(bool, true, "true");
test!(bool, false, "false");
test!(String, "Hello double quotes", r#""Hello double quotes""#);
test!(String, "Hello single quotes", r#"'Hello single quotes'"#);
test!(String, "Hello no quotes", r#"Hello no quotes"#);
test!(Option<i32>, Some(32), "32");
test!(Option<i32>, None, "~");
test!(Option<i32>, None, "null");
test!(
String,
"This is multiline string\n",
r#"
>
This is
multiline
string
"#
);
test!(
String,
"This is\nmultiline\nstring\n",
r#"
|
This is
multiline
string
"#
);
test!(
String,
"This is multiline string",
r#"
>-
This is
multiline
string
"#
);
test!(
String,
"This is\nmultiline\nstring",
r#"
|-
This is
multiline
string
"#
);
test!(Vec<i32>, Vec::<i32>::from([1, 2, 3]), "[1,2,3]");
test!(
Vec<Option<i32>>,
Vec::<Option<i32>>::from([Some(1), Some(2), Some(3)]),
"[1,2,3]"
);
test!(
Vec<Option<i32>>,
Vec::<Option<i32>>::from([Some(1), None, Some(3)]),
"[1,~,3]"
);
test!(Vec<i32>, Vec::<i32>::new(), "[]");
test!(
Vec<i32>,
Vec::<i32>::from([1, 2, 3]),
r#"
- 1
- 2
- 3
"#
);
test!((i32, i8, i64), (1, 2, 3), "[1,2,3]");
test!((), (), "~");
test!((), (), "null");
test!(
(i32, i8, i64),
(321, 12, -4123),
r#"
- 321
- 12
- -4123
"#
);
#[derive(Deserialize, Debug, PartialEq)]
struct TestUnitStruct;
test!(TestUnitStruct, TestUnitStruct, "null");
test!(TestUnitStruct, TestUnitStruct, "~");
#[derive(Deserialize, Debug, PartialEq)]
struct TestEmptyTupleStruct();
test!(TestEmptyTupleStruct, TestEmptyTupleStruct(), "[]");
#[derive(Deserialize, Debug, PartialEq)]
struct TestTupleStruct(i32, String, bool, f64);
test!(
TestTupleStruct,
TestTupleStruct(32, String::from("Hello string"), true, 45.0),
r#"[32, "Hello string", true, 45.0]"#
);
test!(
TestTupleStruct,
TestTupleStruct(123, String::from("Hello world"), true, 123.0),
r#"
- 123
- Hello world
- true
- 123.00
"#
);
#[derive(Deserialize, Debug, PartialEq)]
struct TestStruct {
x: i32,
y: String,
}
test!(
TestStruct,
TestStruct {
x: 3123,
y: String::from("Hello world")
},
r#"
x: 3123
y: Hello world
"#
);
#[derive(Deserialize, Debug, PartialEq)]
enum TestEnum {
VariantA,
VariantB(),
VariantC(i32, String),
VariantD(TestStruct),
}
test!(TestEnum, TestEnum::VariantA, r#"VariantA: ~"#);
test!(TestEnum, TestEnum::VariantB(), r#"VariantB: []"#);
test!(
TestEnum,
TestEnum::VariantC(12, String::from("Hello world")),
r#"VariantC: [12, 'Hello world']"#
);
test!(
TestEnum,
TestEnum::VariantD(TestStruct {
x: 12,
y: String::from("Hello world")
}),
r#"
VariantD:
x: 12
y: Hello world
"#
);
{
type Map = std::collections::HashMap<String, String>;
test!(
Map,
Map::from([(String::from("foo"), String::from("321"))]),
r#"foo: 321"#
);
test!(
Map,
Map::from([(String::from("foo"), String::from("321"))]),
r#"foo: '321'"#
);
test!(
Map,
Map::from([(String::from("foo"), String::from("321"))]),
r#"foo: "321""#
);
}
{
type Map = std::collections::HashMap<String, Option<String>>;
test!(Map, Map::from([(String::from("foo"), None)]), r#"foo: ~"#);
}
#[derive(Deserialize, Debug, PartialEq)]
struct TestStructWithWrapper {
kind: String,
data: YamlWrapper,
}
test!(
TestStructWithWrapper,
TestStructWithWrapper {
kind: String::from("Test"),
data: YamlWrapper::new(Yaml::Array(vec![
Yaml::String("Hello".to_owned()),
Yaml::String("world".to_owned())
])),
},
"kind: Test\ndata: [ 'Hello', 'world' ]"
);
}
}