From db3981d60856cc3f8d44f96863fc08678abfed97 Mon Sep 17 00:00:00 2001 From: Monica Moniot Date: Mon, 11 Sep 2023 12:47:02 -0400 Subject: [PATCH] added serde impl --- performance/Cargo.toml | 1 + performance/src/ratchet_state.rs | 144 ++++++++++++++++++++++++++++++- 2 files changed, 144 insertions(+), 1 deletion(-) diff --git a/performance/Cargo.toml b/performance/Cargo.toml index 957fc19..6c93f5b 100644 --- a/performance/Cargo.toml +++ b/performance/Cargo.toml @@ -22,6 +22,7 @@ p384 = { version = "0.13.0", default-features = false, features = ["ecdh"], opti sha2 = { version = "0.10.7", default-features = false, optional = true } hmac = { version = "0.12.1", default-features = false, optional = true } openssl-sys = { version = "0.9.91", default-features = false, optional = true } +serde = { version = "1.0", default-features = false, optional = true } [features] default = ["debug", "default-crypto"] diff --git a/performance/src/ratchet_state.rs b/performance/src/ratchet_state.rs index 72a4ca9..5c5bdb0 100644 --- a/performance/src/ratchet_state.rs +++ b/performance/src/ratchet_state.rs @@ -13,9 +13,9 @@ use crate::proto::*; /// Corresponds to the Ratchet Key and Ratchet Fingerprint described in Section 3. #[derive(Clone, Eq)] pub struct RatchetState { + pub chain_len: u64, pub key: Zeroizing<[u8; RATCHET_SIZE]>, pub fingerprint: Option>, - pub chain_len: u64, } impl PartialEq for RatchetState { fn eq(&self, other: &Self) -> bool { @@ -163,3 +163,145 @@ impl<'a> RatchetUpdate<'a> { } } } + +#[cfg(feature = "serde")] +impl serde::Serialize for RatchetState { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + use serde::ser::SerializeSeq; + if let Some(rf) = &self.fingerprint { + let mut seq = serializer.serialize_seq(Some(4))?; + seq.serialize_element(&0u8)?; + seq.serialize_element(&self.chain_len)?; + seq.serialize_element(self.key.as_ref())?; + seq.serialize_element(rf.as_ref())?; + seq.end() + } else { + let mut seq = serializer.serialize_seq(Some(3))?; + seq.serialize_element(&0u8)?; + seq.serialize_element(&self.chain_len)?; + seq.serialize_element(self.key.as_ref())?; + seq.end() + } + } +} +#[cfg(feature = "serde")] +impl<'de> serde::Deserialize<'de> for RatchetState { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + struct Visitor; + + impl<'de> serde::de::Visitor<'de> for Visitor { + type Value = RatchetState; + + fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { + formatter.write_str("a ratchet state sequence") + } + + fn visit_seq(self, mut seq: A) -> Result + where + A: serde::de::SeqAccess<'de>, + { + use serde::de::Error; + if Some(0u8) == seq.next_element()? { + if let Some(chain_len) = seq.next_element()? { + if let Some(rk) = seq.next_element::<&[u8]>()? { + if rk.len() == RATCHET_SIZE { + let mut key: Zeroizing<[u8; RATCHET_SIZE]> = Zeroizing::default(); + key.copy_from_slice(rk); + let fingerprint = if let Some(rf) = seq.next_element::<&[u8]>()? { + if rf.len() == RATCHET_SIZE { + let mut fingerprint: Zeroizing<[u8; RATCHET_SIZE]> = Zeroizing::default(); + fingerprint.copy_from_slice(rf); + Some(fingerprint) + } else { + return Err(A::Error::invalid_length( + rf.len(), + &"a ratchet fingerprint of length 32", + )); + } + } else { + None + }; + + Ok(RatchetState { chain_len, key, fingerprint }) + } else { + Err(A::Error::invalid_length( + rk.len(), + &"a ratchet key of length 32", + )) + } + } else { + Err(A::Error::custom("expected a ratchet key")) + } + } else { + Err(A::Error::custom("expected an unsigned integer")) + } + } else { + Err(A::Error::custom("invalid version byte")) + } + } + } + deserializer.deserialize_seq(Visitor) + } +} + +#[cfg(feature = "serde")] +impl serde::Serialize for RatchetStates { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + use serde::ser::SerializeSeq; + if let Some(state2) = &self.state2 { + let mut seq = serializer.serialize_seq(Some(3))?; + seq.serialize_element(&0u8)?; + seq.serialize_element(&self.state1)?; + seq.serialize_element(state2)?; + seq.end() + } else { + let mut seq = serializer.serialize_seq(Some(2))?; + seq.serialize_element(&0u8)?; + seq.serialize_element(&self.state1)?; + seq.end() + } + } +} +#[cfg(feature = "serde")] +impl<'de> serde::Deserialize<'de> for RatchetStates { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + struct Visitor; + + impl<'de> serde::de::Visitor<'de> for Visitor { + type Value = RatchetStates; + + fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result { + formatter.write_str("a pair of ratchet states") + } + + fn visit_seq(self, mut seq: A) -> Result + where + A: serde::de::SeqAccess<'de>, + { + use serde::de::Error; + if Some(0u8) == seq.next_element()? { + if let Some(state1) = seq.next_element()? { + Ok(RatchetStates { state1, state2: seq.next_element()? }) + } else { + Err(A::Error::custom("expected a ratchet state")) + } + } else { + Err(A::Error::custom("invalid version byte")) + } + } + } + deserializer.deserialize_seq(Visitor) + } +}