use crate::{Error, error::Result};
use serde_core::{
Deserialize,
de::{self, IntoDeserializer},
};
use std::collections::HashMap;
use std::str::FromStr;
pub struct Deserializer {
sections: HashMap<String, HashMap<String, String>>,
}
pub fn from_str<'a, T>(s: &'a str) -> Result<T>
where
T: Deserialize<'a>,
{
let mut deserializer = Deserializer::from_str(s)?;
let t = T::deserialize(&mut deserializer)?;
Ok(t)
}
impl Deserializer {
fn from_str(input: &str) -> Result<Self> {
let mut sections = HashMap::new();
let mut current_section = String::new();
sections.insert(current_section.clone(), HashMap::new());
for line in input.lines() {
let line = line.trim();
if line.is_empty() || line.starts_with(';') || line.starts_with('#') {
continue;
}
if line.starts_with('[') && line.ends_with(']') {
current_section = line[1..line.len() - 1].to_string();
sections.insert(current_section.clone(), HashMap::new());
continue;
}
if let Some(eq_pos) = line.find('=') {
let key = line[..eq_pos].trim().to_string();
let value = Self::unescape_value(line[eq_pos + 1..].trim())?;
if let Some(section) = sections.get_mut(¤t_section) {
section.insert(key, value);
}
}
}
Ok(Deserializer { sections })
}
fn unescape_value(value: &str) -> Result<String> {
let mut unescaped = String::with_capacity(value.len());
let mut iter = value.chars();
while let Some(c) = iter.next() {
match c {
'\\' => match iter
.next()
.ok_or(Error::Unescape("Backslash without escape".to_string()))?
{
'0' => unescaped.push('\0'),
'a' => unescaped.push('\x07'),
'b' => unescaped.push('\x08'),
't' => unescaped.push('\t'),
'r' => unescaped.push('\r'),
'n' => unescaped.push('\n'),
'x' => {
let s: String = iter.by_ref().take(4).collect();
let codepoint = u32::from_str_radix(&s, 16)
.map_err(|e| Error::Unescape(e.to_string()))?;
let char = char::from_u32(codepoint).ok_or_else(|| {
Error::Unescape(format!("{s:?} is not a valid unicode codepoint"))
})?;
unescaped.push(char);
}
escapee => unescaped.push(escapee),
},
c => unescaped.push(c),
}
}
Ok(unescaped)
}
}
impl<'de> de::Deserializer<'de> for &mut Deserializer {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
self.deserialize_map(visitor)
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_bool(true)
}
fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i8(0)
}
fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i16(0)
}
fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i32(0)
}
fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i64(0)
}
fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u8(0)
}
fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u16(0)
}
fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u32(0)
}
fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u64(0)
}
fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_f32(0.0)
}
fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_f64(0.0)
}
fn deserialize_char<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_char('a')
}
fn deserialize_str<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_borrowed_str("")
}
fn deserialize_string<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_string(String::new())
}
fn deserialize_bytes<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_bytes(b"")
}
fn deserialize_byte_buf<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_byte_buf(Vec::new())
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_some(self)
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_unit_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_newtype_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_newtype_struct(self)
}
fn deserialize_seq<V>(self, _visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("sequences".to_string()))
}
fn deserialize_tuple<V>(self, _len: usize, _visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("tuples".to_string()))
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
_len: usize,
_visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("tuple structs".to_string()))
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_map(MapAccess::new(self))
}
fn deserialize_struct<V>(
self,
name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
if name.is_empty() || self.sections.contains_key(name) {
if name.is_empty() {
visitor.visit_map(RootStructAccess::new(self))
} else {
visitor.visit_map(StructAccess::new(self, name))
}
} else {
if self.sections.len() > 1
|| (self.sections.len() == 1 && !self.sections.contains_key(""))
{
visitor.visit_map(RootStructAccess::new(self))
} else {
visitor.visit_map(StructAccess::new(self, ""))
}
}
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
_visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("enums".to_string()))
}
fn deserialize_identifier<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_str("")
}
fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
}
struct MapAccess<'a> {
de: &'a mut Deserializer,
sections: Vec<String>,
index: usize,
}
impl<'a> MapAccess<'a> {
fn new(de: &'a mut Deserializer) -> Self {
let sections: Vec<String> = de.sections.keys().cloned().collect();
MapAccess {
de,
sections,
index: 0,
}
}
}
enum FieldSource {
Root(String),
Section,
}
struct RootStructAccess<'a> {
de: &'a mut Deserializer,
fields: Vec<(String, FieldSource)>,
index: usize,
}
impl<'a> RootStructAccess<'a> {
fn new(de: &'a mut Deserializer) -> Self {
let mut fields = Vec::new();
if let Some(root_section) = de.sections.get("") {
for (key, value) in root_section {
if de.sections.contains_key(key) {
fields.push((key.clone(), FieldSource::Section));
} else {
fields.push((key.clone(), FieldSource::Root(value.clone())));
}
}
}
for section_name in de.sections.keys() {
if !section_name.is_empty() {
if !fields.iter().any(|(name, _)| name == section_name) {
fields.push((section_name.clone(), FieldSource::Section));
}
}
}
RootStructAccess {
de,
fields,
index: 0,
}
}
}
impl<'de> de::MapAccess<'de> for RootStructAccess<'_> {
type Error = Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
where
K: de::DeserializeSeed<'de>,
{
if self.index < self.fields.len() {
let (key, _) = &self.fields[self.index];
self.index += 1;
seed.deserialize(key.as_str().into_deserializer()).map(Some)
} else {
Ok(None)
}
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: de::DeserializeSeed<'de>,
{
let (key, source) = &self.fields[self.index - 1];
match source {
FieldSource::Root(value) => seed.deserialize(ValueDeserializer::new(value)),
FieldSource::Section => seed.deserialize(&mut SectionDeserializer::new(self.de, key)),
}
}
}
impl<'de> de::MapAccess<'de> for MapAccess<'_> {
type Error = Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
where
K: de::DeserializeSeed<'de>,
{
if self.index < self.sections.len() {
let key = &self.sections[self.index];
self.index += 1;
seed.deserialize(key.as_str().into_deserializer()).map(Some)
} else {
Ok(None)
}
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: de::DeserializeSeed<'de>,
{
let section = &self.sections[self.index - 1];
seed.deserialize(&mut SectionDeserializer::new(self.de, section))
}
}
struct StructAccess {
fields: Vec<(String, String)>,
index: usize,
}
impl StructAccess {
fn new(de: &mut Deserializer, section: &str) -> Self {
let fields = if let Some(section_map) = de.sections.get(section) {
section_map
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect()
} else {
Vec::new()
};
StructAccess { fields, index: 0 }
}
}
struct SectionDeserializer<'a> {
de: &'a mut Deserializer,
section: String,
}
impl<'a> SectionDeserializer<'a> {
fn new(de: &'a mut Deserializer, section: &str) -> Self {
SectionDeserializer {
de,
section: section.to_string(),
}
}
}
impl<'de, 'a> de::Deserializer<'de> for &'a mut SectionDeserializer<'a> {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
self.deserialize_struct("", &[], visitor)
}
fn deserialize_struct<V>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_map(StructAccess::new(self.de, &self.section))
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_some(self)
}
serde_core::forward_to_deserialize_any! {
bool i8 i16 i32 i64 u8 u16 u32 u64 f32 f64 char str string
bytes byte_buf unit unit_struct newtype_struct seq tuple
tuple_struct map enum identifier ignored_any
}
}
impl<'de> de::MapAccess<'de> for StructAccess {
type Error = Error;
fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
where
K: de::DeserializeSeed<'de>,
{
if self.index < self.fields.len() {
let (key, _) = &self.fields[self.index];
self.index += 1;
seed.deserialize(key.as_str().into_deserializer()).map(Some)
} else {
Ok(None)
}
}
fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
where
V: de::DeserializeSeed<'de>,
{
let (_, value) = &self.fields[self.index - 1];
seed.deserialize(ValueDeserializer::new(value))
}
}
struct ValueDeserializer {
value: String,
}
impl ValueDeserializer {
fn new(value: &str) -> Self {
ValueDeserializer {
value: value.to_string(),
}
}
}
impl<'de> de::Deserializer<'de> for ValueDeserializer {
type Error = Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
self.deserialize_str(visitor)
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
match self.value.as_str() {
"true" => visitor.visit_bool(true),
"false" => visitor.visit_bool(false),
_ => Err(Error::InvalidValue {
typ: "bool".to_string(),
value: self.value,
}),
}
}
fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i8(i8::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "i8".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i16(i16::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "i16".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i32(i32::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "i32".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_i64(i64::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "i64".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u8(u8::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "u8".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u16(u16::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "u16".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u32(u32::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "u32".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_u64(u64::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "u64".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_f32(f32::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "f32".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_f64(f64::from_str(&self.value).map_err(|_| Error::InvalidValue {
typ: "f64".to_string(),
value: self.value.clone(),
})?)
}
fn deserialize_char<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
if self.value.len() == 1 {
visitor.visit_char(self.value.chars().next().unwrap())
} else {
Err(Error::InvalidValue {
typ: "char".to_string(),
value: self.value,
})
}
}
fn deserialize_str<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_string(self.value)
}
fn deserialize_string<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_string(self.value)
}
fn deserialize_bytes<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_bytes(self.value.as_bytes())
}
fn deserialize_byte_buf<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_byte_buf(self.value.into_bytes())
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_some(self)
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_unit_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_newtype_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_newtype_struct(self)
}
fn deserialize_seq<V>(self, _visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("sequences".to_string()))
}
fn deserialize_tuple<V>(self, _len: usize, _visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("tuples".to_string()))
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
_len: usize,
_visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("tuple structs".to_string()))
}
fn deserialize_map<V>(self, _visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("maps in values".to_string()))
}
fn deserialize_struct<V>(
self,
_name: &'static str,
_fields: &'static [&'static str],
_visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("structs in values".to_string()))
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
_visitor: V,
) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
Err(Error::UnsupportedFeature("enums".to_string()))
}
fn deserialize_identifier<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
self.deserialize_str(visitor)
}
fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value>
where
V: de::Visitor<'de>,
{
visitor.visit_unit()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn unescape_value() {
for (escaped, should_unescaped) in [
(r#"\\\\"#, r#"\\"#),
(r#"\\n"#, r#"\n"#),
(r#"\\r"#, r#"\r"#),
(r#"\\t"#, r#"\t"#),
(r#"\\\""#, r#"\""#),
(r#"\\\;"#, r#"\;"#),
(r#"\\\#"#, r#"\#"#),
] {
assert_eq!(
should_unescaped.to_string(),
Deserializer::unescape_value(escaped).unwrap()
);
}
}
}