diff --git a/src/de.rs b/src/de.rs index f5a1e17..0bece7a 100644 --- a/src/de.rs +++ b/src/de.rs @@ -1,40 +1,51 @@ use std::str::FromStr; -use serde::de::{Visitor, IntoDeserializer}; use serde::de::value::MapDeserializer; +use serde::de::{IntoDeserializer, Visitor}; use regex::Regex; use crate::error::*; -pub(crate) struct Deserializer<'de> { +pub struct Deserializer<'de> { input: &'de str, regex: Regex, } impl<'de> Deserializer<'de> { - pub fn new(input: &'de str, regex: Regex) -> Deserializer { - Deserializer { - input, - regex, - } + pub const fn new(input: &'de str, regex: Regex) -> Deserializer { + Deserializer { input, regex } } } -impl<'de, 'a> serde::Deserializer<'de> for &'a mut Deserializer<'de> { +impl<'de> serde::Deserializer<'de> for &'_ mut Deserializer<'de> { type Error = Error; - fn deserialize_any(self, visitor: V) -> Result where V: Visitor<'de> { + #[inline] + fn deserialize_any(self, visitor: V) -> Result + where + V: Visitor<'de>, + { self.deserialize_map(visitor) } - fn deserialize_map(self, visitor: V) -> Result where V: Visitor<'de> { - let caps = self.regex.captures(&self.input).ok_or_else(Error::NoMatch)?; + fn deserialize_map(self, visitor: V) -> Result + where + V: Visitor<'de>, + { + let caps = self.regex.captures(self.input).ok_or_else(Error::NoMatch)?; let items = self.regex.capture_names().filter_map(|n| { - n.map(|name| { - let value = &caps[name]; - (name.to_owned(), Value { name: name.to_owned(), value: value.to_owned() }) + n.and_then(|name| { + caps.name(name).map(|value| { + ( + name.to_owned(), + Value { + name: name.to_owned(), + value: value.as_str().to_owned(), + }, + ) + }) }) }); @@ -54,14 +65,17 @@ impl<'de, 'a> serde::Deserializer<'de> for &'a mut Deserializer<'de> { } } - struct Value { name: String, value: String, } impl Value { - fn parse(&self) -> Result where T: FromStr { + #[allow(clippy::map_err_ignore)] + fn parse(&self) -> Result + where + T: FromStr, + { self.value.parse().map_err(|_| self.get_parse_error()) } @@ -84,11 +98,17 @@ impl<'de> IntoDeserializer<'de, Error> for Value { impl<'de> serde::Deserializer<'de> for Value { type Error = Error; - fn deserialize_any(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_any(self, visitor: V) -> Result + where + V: Visitor<'de>, + { self.value.into_deserializer().deserialize_any(visitor) } - fn deserialize_bool(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_bool(self, visitor: V) -> Result + where + V: Visitor<'de>, + { if self.value.eq_ignore_ascii_case("true") { visitor.visit_bool(true) } else if self.value.eq_ignore_ascii_case("false") { @@ -98,47 +118,80 @@ impl<'de> serde::Deserializer<'de> for Value { } } - fn deserialize_i8(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_i8(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_i8(self.parse()?) } - fn deserialize_i16(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_i16(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_i16(self.parse()?) } - fn deserialize_i32(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_i32(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_i32(self.parse()?) } - fn deserialize_i64(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_i64(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_i64(self.parse()?) } - fn deserialize_u8(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_u8(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_u8(self.parse()?) } - fn deserialize_u16(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_u16(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_u16(self.parse()?) } - fn deserialize_u32(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_u32(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_u32(self.parse()?) } - fn deserialize_u64(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_u64(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_u64(self.parse()?) } - fn deserialize_f32(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_f32(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_f32(self.parse()?) } - fn deserialize_f64(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_f64(self, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_f64(self.parse()?) } - fn deserialize_option(self, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_option(self, visitor: V) -> Result + where + V: Visitor<'de>, + { if self.value.is_empty() { visitor.visit_none() } else { @@ -146,11 +199,22 @@ impl<'de> serde::Deserializer<'de> for Value { } } - fn deserialize_newtype_struct(self, _: &'static str, visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_newtype_struct(self, _: &'static str, visitor: V) -> Result + where + V: Visitor<'de>, + { visitor.visit_newtype_struct(self) } - fn deserialize_enum(self, _name: &'static str, _variants: &'static [&'static str], visitor: V) -> Result where V: Visitor<'de> { + fn deserialize_enum( + self, + _name: &'static str, + _variants: &'static [&'static str], + visitor: V, + ) -> Result + where + V: Visitor<'de>, + { visitor.visit_enum(self.value.into_deserializer()) } @@ -160,4 +224,4 @@ impl<'de> serde::Deserializer<'de> for Value { unit seq bytes byte_buf map unit_struct tuple_struct tuple ignored_any struct } -} \ No newline at end of file +} diff --git a/src/error.rs b/src/error.rs index b44f44d..81011a1 100644 --- a/src/error.rs +++ b/src/error.rs @@ -2,6 +2,7 @@ use std::fmt::{Display, Formatter}; /// An error that occurred during deserialization. #[derive(Debug)] +#[non_exhaustive] pub enum Error { /// An error occurred while parsing the regular expression BadRegex(regex::Error), @@ -23,7 +24,11 @@ pub enum Error { } impl serde::de::Error for Error { - fn custom(msg: T) -> Self where T: Display { + #[inline] + fn custom(msg: T) -> Self + where + T: Display, + { Self::Custom(msg.to_string()) } } @@ -31,17 +36,23 @@ impl serde::de::Error for Error { impl std::error::Error for Error {} impl Display for Error { + #[inline] fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { use Error::*; - match self { - BadRegex(err) => err.fmt(f), + match *self { + BadRegex(ref err) => err.fmt(f), NoMatch() => write!(f, "String doesn't match pattern"), - BadValue { name, value } => write!(f, "Unable to convert value for group {}: {}", name, value), - Custom(err) => write!(f, "{}", err), + BadValue { + ref name, + ref value, + } => { + write!(f, "Unable to convert value for group {}: {}", name, value) + } + Custom(ref err) => write!(f, "{}", err), } } } // Do not use this alias in public parts of the crate because // it would hide the direct link to the actual error type in rustdoc. -pub(crate) type Result = std::result::Result; +pub type Result = std::result::Result; diff --git a/src/lib.rs b/src/lib.rs index 7a51762..8e6f2ed 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -89,13 +89,13 @@ If your regular expression looks like a behemoth no mere mortal will ever unders * */ -mod error; mod de; +mod error; pub use error::Error; -use serde::Deserialize; use regex::Regex; +use serde::Deserialize; /// Deserialize an input string into a struct. /// @@ -120,9 +120,13 @@ use regex::Regex; /// # Ok(()) /// # } /// ``` -pub fn from_str<'a, T>(input: &'a str, regex: &str) -> std::result::Result where T: Deserialize<'a> { - let regex = Regex::new(®ex).map_err(Error::BadRegex)?; - from_str_regex(input, regex) +#[inline] +pub fn from_str<'input, T>(input: &'input str, regex: &str) -> std::result::Result +where + T: Deserialize<'input>, +{ + let rex = Regex::new(regex).map_err(Error::BadRegex)?; + from_str_regex(input, rex) } /// Deserialize an input string into a struct. @@ -149,15 +153,19 @@ pub fn from_str<'a, T>(input: &'a str, regex: &str) -> std::result::Result(input: &'a str, regex: Regex) -> std::result::Result where T: Deserialize<'a> { +#[inline] +pub fn from_str_regex<'input, T>(input: &'input str, regex: Regex) -> std::result::Result +where + T: Deserialize<'input>, +{ let mut deserializer = de::Deserializer::new(input, regex); T::deserialize(&mut deserializer) } #[cfg(test)] mod test { - use super::*; use super::error::Result; + use super::*; #[derive(Deserialize, PartialEq, Debug)] struct Test { @@ -233,20 +241,23 @@ mod test { let input = "true,1,2,3,4,-1,-2,-3,-4,1.0,-1.0,foobar"; let output: Test2 = from_str(input, TEST2_PATTERN).unwrap(); - assert_eq!(output, Test2 { - f_bool: true, - f_u8: 1, - f_u16: 2, - f_u32: 3, - f_u64: 4, - f_i8: -1, - f_i16: -2, - f_i32: -3, - f_i64: -4, - f_f32: 1.0, - f_f64: -1.0, - f_str: "foobar".to_owned(), - }); + assert_eq!( + output, + Test2 { + f_bool: true, + f_u8: 1, + f_u16: 2, + f_u32: 3, + f_u64: 4, + f_i8: -1, + f_i16: -2, + f_i32: -3, + f_i64: -4, + f_f32: 1.0, + f_f64: -1.0, + f_str: "foobar".to_owned(), + } + ); } #[derive(Deserialize, PartialEq, Debug)] @@ -261,7 +272,13 @@ mod test { let input = "1,-2"; let output: Test3 = from_str(input, regex).unwrap(); - assert_eq!(output, Test3 { foo: Some(1), bar: Some(-2) }); + assert_eq!( + output, + Test3 { + foo: Some(1), + bar: Some(-2) + } + ); } #[test] @@ -270,7 +287,43 @@ mod test { let input = ","; let output: Test3 = from_str(input, regex).unwrap(); - assert_eq!(output, Test3 { foo: None, bar: None }); + assert_eq!( + output, + Test3 { + foo: None, + bar: None + } + ); + } + + #[test] + fn test_option_present() { + let regex = r"^(?P\d*)(?:,(?P-?\d*))?$"; + let input = "1,-2"; + let output: Test3 = from_str(input, regex).unwrap(); + + assert_eq!( + output, + Test3 { + foo: Some(1), + bar: Some(-2) + } + ); + } + + #[test] + fn test_option_missing() { + let regex = r"^(?P\d*)(?:,(?P-?\d*))?$"; + let input = "1"; + let output: Test3 = from_str(input, regex).unwrap(); + + assert_eq!( + output, + Test3 { + foo: Some(1), + bar: None + } + ); } #[test] @@ -347,7 +400,11 @@ mod test { let input = "aaa1,-2"; let output: Result = from_str(input, regex); - assert!(matches!(output, Err(Error::BadValue{..})), "Expected Error::BadValue got {:?}", output); + assert!( + matches!(output, Err(Error::BadValue { .. })), + "Expected Error::BadValue got {:?}", + output + ); } #[test]