diff --git a/Cargo.toml b/Cargo.toml index 563998d..64aabe2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "curve25519-dalek" -version = "1.2.3" +version = "2.0.0-alpha.0" authors = ["Isis Lovecruft ", "Henry de Valence "] readme = "README.md" diff --git a/src/edwards.rs b/src/edwards.rs index d80392d..998af8d 100644 --- a/src/edwards.rs +++ b/src/edwards.rs @@ -217,7 +217,12 @@ impl Serialize for EdwardsPoint { fn serialize(&self, serializer: S) -> Result where S: Serializer { - serializer.serialize_bytes(self.compress().as_bytes()) + use serde::ser::SerializeTuple; + let mut tup = serializer.serialize_tuple(32)?; + for byte in self.compress().as_bytes().iter() { + tup.serialize_element(byte)?; + } + tup.end() } } @@ -226,7 +231,12 @@ impl Serialize for CompressedEdwardsY { fn serialize(&self, serializer: S) -> Result where S: Serializer { - serializer.serialize_bytes(self.as_bytes()) + use serde::ser::SerializeTuple; + let mut tup = serializer.serialize_tuple(32)?; + for byte in self.as_bytes().iter() { + tup.serialize_element(byte)?; + } + tup.end() } } @@ -244,22 +254,21 @@ impl<'de> Deserialize<'de> for EdwardsPoint { formatter.write_str("a valid point in Edwards y + sign format") } - fn visit_bytes(self, v: &[u8]) -> Result - where E: serde::de::Error + fn visit_seq(self, mut seq: A) -> Result + where A: serde::de::SeqAccess<'de> { - if v.len() == 32 { - let mut arr32 = [0u8; 32]; - arr32[0..32].copy_from_slice(v); - CompressedEdwardsY(arr32) - .decompress() - .ok_or(serde::de::Error::custom("decompression failed")) - } else { - Err(serde::de::Error::invalid_length(v.len(), &self)) + let mut bytes = [0u8; 32]; + for i in 0..32 { + bytes[i] = seq.next_element()? + .ok_or(serde::de::Error::invalid_length(i, &"expected 32 bytes"))?; } + CompressedEdwardsY(bytes) + .decompress() + .ok_or(serde::de::Error::custom("decompression failed")) } } - deserializer.deserialize_bytes(EdwardsPointVisitor) + deserializer.deserialize_tuple(32, EdwardsPointVisitor) } } @@ -277,20 +286,19 @@ impl<'de> Deserialize<'de> for CompressedEdwardsY { formatter.write_str("32 bytes of data") } - fn visit_bytes(self, v: &[u8]) -> Result - where E: serde::de::Error + fn visit_seq(self, mut seq: A) -> Result + where A: serde::de::SeqAccess<'de> { - if v.len() == 32 { - let mut arr32 = [0u8; 32]; - arr32[0..32].copy_from_slice(v); - Ok(CompressedEdwardsY(arr32)) - } else { - Err(serde::de::Error::invalid_length(v.len(), &self)) + let mut bytes = [0u8; 32]; + for i in 0..32 { + bytes[i] = seq.next_element()? + .ok_or(serde::de::Error::invalid_length(i, &"expected 32 bytes"))?; } + Ok(CompressedEdwardsY(bytes)) } } - deserializer.deserialize_bytes(CompressedEdwardsYVisitor) + deserializer.deserialize_tuple(32, CompressedEdwardsYVisitor) } } diff --git a/src/ristretto.rs b/src/ristretto.rs index 5d0be03..6d53e89 100644 --- a/src/ristretto.rs +++ b/src/ristretto.rs @@ -334,7 +334,12 @@ impl Serialize for RistrettoPoint { fn serialize(&self, serializer: S) -> Result where S: Serializer { - serializer.serialize_bytes(self.compress().as_bytes()) + use serde::ser::SerializeTuple; + let mut tup = serializer.serialize_tuple(32)?; + for byte in self.compress().as_bytes().iter() { + tup.serialize_element(byte)?; + } + tup.end() } } @@ -343,7 +348,12 @@ impl Serialize for CompressedRistretto { fn serialize(&self, serializer: S) -> Result where S: Serializer { - serializer.serialize_bytes(self.as_bytes()) + use serde::ser::SerializeTuple; + let mut tup = serializer.serialize_tuple(32)?; + for byte in self.as_bytes().iter() { + tup.serialize_element(byte)?; + } + tup.end() } } @@ -361,22 +371,21 @@ impl<'de> Deserialize<'de> for RistrettoPoint { formatter.write_str("a valid point in Ristretto format") } - fn visit_bytes(self, v: &[u8]) -> Result - where E: serde::de::Error + fn visit_seq(self, mut seq: A) -> Result + where A: serde::de::SeqAccess<'de> { - if v.len() == 32 { - let mut arr32 = [0u8; 32]; - arr32[0..32].copy_from_slice(v); - CompressedRistretto(arr32) - .decompress() - .ok_or(serde::de::Error::custom("decompression failed")) - } else { - Err(serde::de::Error::invalid_length(v.len(), &self)) + let mut bytes = [0u8; 32]; + for i in 0..32 { + bytes[i] = seq.next_element()? + .ok_or(serde::de::Error::invalid_length(i, &"expected 32 bytes"))?; } + CompressedRistretto(bytes) + .decompress() + .ok_or(serde::de::Error::custom("decompression failed")) } } - deserializer.deserialize_bytes(RistrettoPointVisitor) + deserializer.deserialize_tuple(32, RistrettoPointVisitor) } } @@ -394,20 +403,19 @@ impl<'de> Deserialize<'de> for CompressedRistretto { formatter.write_str("32 bytes of data") } - fn visit_bytes(self, v: &[u8]) -> Result - where E: serde::de::Error + fn visit_seq(self, mut seq: A) -> Result + where A: serde::de::SeqAccess<'de> { - if v.len() == 32 { - let mut arr32 = [0u8; 32]; - arr32[0..32].copy_from_slice(v); - Ok(CompressedRistretto(arr32)) - } else { - Err(serde::de::Error::invalid_length(v.len(), &self)) + let mut bytes = [0u8; 32]; + for i in 0..32 { + bytes[i] = seq.next_element()? + .ok_or(serde::de::Error::invalid_length(i, &"expected 32 bytes"))?; } + Ok(CompressedRistretto(bytes)) } } - deserializer.deserialize_bytes(CompressedRistrettoVisitor) + deserializer.deserialize_tuple(32, CompressedRistrettoVisitor) } } diff --git a/src/scalar.rs b/src/scalar.rs index 86085ac..3b252e2 100644 --- a/src/scalar.rs +++ b/src/scalar.rs @@ -385,7 +385,12 @@ impl Serialize for Scalar { fn serialize(&self, serializer: S) -> Result where S: Serializer { - serializer.serialize_bytes(self.reduce().as_bytes()) + use serde::ser::SerializeTuple; + let mut tup = serializer.serialize_tuple(32)?; + for byte in self.as_bytes().iter() { + tup.serialize_element(byte)?; + } + tup.end() } } @@ -400,32 +405,25 @@ impl<'de> Deserialize<'de> for Scalar { type Value = Scalar; fn expecting(&self, formatter: &mut ::core::fmt::Formatter) -> ::core::fmt::Result { - formatter.write_str("a canonically-encoded 32-byte scalar value") + formatter.write_str("a valid point in Edwards y + sign format") } - fn visit_bytes(self, v: &[u8]) -> Result - where E: serde::de::Error + fn visit_seq(self, mut seq: A) -> Result + where A: serde::de::SeqAccess<'de> { - if v.len() == 32 { - let mut bytes = [0u8;32]; - bytes.copy_from_slice(v); - - static ERRMSG: &'static str = "encoding was not canonical"; - - Scalar::from_canonical_bytes(bytes) - .ok_or( - serde::de::Error::invalid_value( - serde::de::Unexpected::Bytes(v), - &ERRMSG, - ) - ) - } else { - Err(serde::de::Error::invalid_length(v.len(), &self)) + let mut bytes = [0u8; 32]; + for i in 0..32 { + bytes[i] = seq.next_element()? + .ok_or(serde::de::Error::invalid_length(i, &"expected 32 bytes"))?; } + Scalar::from_canonical_bytes(bytes) + .ok_or(serde::de::Error::custom( + &"scalar was not canonically encoded" + )) } } - deserializer.deserialize_bytes(ScalarVisitor) + deserializer.deserialize_tuple(32, ScalarVisitor) } }