diff --git a/src/scalar.rs b/src/scalar.rs index 34a45ba..7f90624 100644 --- a/src/scalar.rs +++ b/src/scalar.rs @@ -185,6 +185,50 @@ impl CTAssignable for Scalar { } } +#[cfg(feature = "serde")] +use serde::{self, Serialize, Deserialize, Serializer, Deserializer}; +#[cfg(feature = "serde")] +use serde::de::Visitor; + +#[cfg(feature = "serde")] +impl Serialize for Scalar { + fn serialize(&self, serializer: S) -> Result + where S: Serializer + { + serializer.serialize_bytes(self.as_bytes()) + } +} + +#[cfg(feature = "serde")] +impl<'de> Deserialize<'de> for Scalar { + fn deserialize(deserializer: D) -> Result + where D: Deserializer<'de> + { + struct ScalarVisitor; + + impl<'de> Visitor<'de> for ScalarVisitor { + type Value = Scalar; + + fn expecting(&self, formatter: &mut ::core::fmt::Formatter) -> ::core::fmt::Result { + formatter.write_str("a 32-byte scalar value") + } + + fn visit_bytes(self, v: &[u8]) -> Result + where E: serde::de::Error + { + if v.len() == 32 { + // array_ref turns &[u8] into &[u8;32] + Ok(Scalar(*array_ref!(v,0,32))) + } else { + Err(serde::de::Error::invalid_length(v.len(), &self)) + } + } + } + + deserializer.deserialize_bytes(ScalarVisitor) + } +} + impl Scalar { /// Return a `Scalar` chosen uniformly at random using a user-provided RNG. /// @@ -827,6 +871,17 @@ mod test { assert_eq!(should_be_X, X); } + + #[cfg(feature = "serde")] + use serde_cbor; + + #[test] + #[cfg(feature = "serde")] + fn serde_cbor_scalar_roundtrip() { + let output = serde_cbor::to_vec(&X).unwrap(); + let parsed: Scalar = serde_cbor::from_slice(&output).unwrap(); + assert_eq!(parsed, X); + } } #[cfg(all(test, feature = "bench"))]