Move Serde De/Serialization to Group Implementation (#380)

* Move Serde De/Serialization to `Group` implementation

* Clean up Serde formats

* Fix typo
This commit is contained in:
daxpedda
2025-07-14 19:54:43 -07:00
committed by GitHub
parent 5b492aedbf
commit 1b6633cdcc
12 changed files with 415 additions and 220 deletions
+1
View File
@@ -24,6 +24,7 @@ serde = [
"curve25519-dalek?/serde", "curve25519-dalek?/serde",
"ecdsa?/serde", "ecdsa?/serde",
"ed25519-dalek?/serde", "ed25519-dalek?/serde",
"elliptic-curve/serde",
"generic-array/serde", "generic-array/serde",
"voprf/serde", "voprf/serde",
"zeroize/serde", "zeroize/serde",
+76 -23
View File
@@ -11,10 +11,11 @@
pub use curve25519_dalek; pub use curve25519_dalek;
use curve25519_dalek::montgomery::MontgomeryPoint; use curve25519_dalek::montgomery::MontgomeryPoint;
use curve25519_dalek::scalar; use curve25519_dalek::scalar;
use curve25519_dalek::traits::Identity; use curve25519_dalek::traits::IsIdentity;
use generic_array::GenericArray; use generic_array::GenericArray;
use generic_array::typenum::U32; use generic_array::typenum::U32;
use rand::{CryptoRng, RngCore}; use rand::{CryptoRng, RngCore};
use subtle::ConstantTimeEq;
use zeroize::Zeroize; use zeroize::Zeroize;
use super::Group; use super::Group;
@@ -27,22 +28,19 @@ pub struct Curve25519;
/// The implementation of such a subgroup for Curve25519 /// The implementation of such a subgroup for Curve25519
impl Group for Curve25519 { impl Group for Curve25519 {
type Pk = MontgomeryPoint; type Pk = NonIdentity;
type PkLen = U32; type PkLen = U32;
type Sk = Scalar; type Sk = Scalar;
type SkLen = U32; type SkLen = U32;
fn serialize_pk(pk: Self::Pk) -> GenericArray<u8, Self::PkLen> { fn serialize_pk(pk: Self::Pk) -> GenericArray<u8, Self::PkLen> {
pk.to_bytes().into() pk.0.to_bytes().into()
} }
fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> { fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> {
bytes bytes
.take_array::<U32>("public key") .take_array::<U32>("public key")
.ok() .and_then(|bytes| NonIdentity::from_bytes(bytes.into()))
.map(|array| MontgomeryPoint(array.into()))
.filter(|pk| pk != &MontgomeryPoint::identity())
.ok_or(ProtocolError::SerializationError)
} }
fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk { fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk {
@@ -59,7 +57,7 @@ impl Group for Curve25519 {
} }
fn public_key(sk: Self::Sk) -> Self::Pk { fn public_key(sk: Self::Sk) -> Self::Pk {
MontgomeryPoint::mul_base_clamped(sk.0) NonIdentity(MontgomeryPoint::mul_base_clamped(sk.0))
} }
fn serialize_sk(sk: Self::Sk) -> GenericArray<u8, Self::SkLen> { fn serialize_sk(sk: Self::Sk) -> GenericArray<u8, Self::SkLen> {
@@ -69,26 +67,81 @@ impl Group for Curve25519 {
fn deserialize_take_sk(bytes: &mut &[u8]) -> Result<Self::Sk, ProtocolError> { fn deserialize_take_sk(bytes: &mut &[u8]) -> Result<Self::Sk, ProtocolError> {
bytes bytes
.take_array::<U32>("secret key") .take_array::<U32>("secret key")
.ok() .and_then(|bytes| Scalar::from_bytes(bytes.into()))
.and_then(|bytes| {
let scalar = scalar::clamp_integer(bytes.into());
(scalar == *bytes).then_some(scalar)
})
.map(Scalar)
.ok_or(ProtocolError::SerializationError)
} }
} }
/// Curve25519 scalar.
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Zeroize)]
pub struct Scalar([u8; 32]);
impl DiffieHellman<Curve25519> for Scalar { impl DiffieHellman<Curve25519> for Scalar {
fn diffie_hellman(self, pk: MontgomeryPoint) -> GenericArray<u8, U32> { fn diffie_hellman(self, pk: NonIdentity) -> GenericArray<u8, U32> {
Curve25519::serialize_pk(pk.mul_clamped(self.0)) Curve25519::serialize_pk(NonIdentity(pk.0.mul_clamped(self.0)))
} }
} }
/// Non-identity point wrapper for [`MontgomeryPoint`].
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Zeroize)]
pub struct NonIdentity(
#[cfg_attr(feature = "serde", serde(deserialize_with = "serde_deserialize_pk"))]
MontgomeryPoint,
);
impl NonIdentity {
fn from_bytes(bytes: [u8; 32]) -> Result<Self, ProtocolError> {
let point = MontgomeryPoint(bytes);
if point.is_identity() {
Err(ProtocolError::SerializationError)
} else {
Ok(NonIdentity(point))
}
}
}
#[cfg(feature = "serde")]
fn serde_deserialize_pk<'de, D>(deserializer: D) -> Result<MontgomeryPoint, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::{Deserialize, Error};
let point = MontgomeryPoint::deserialize(deserializer)?;
NonIdentity::from_bytes(point.0)
.map(|point| point.0)
.map_err(Error::custom)
}
/// Curve25519 scalar.
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Zeroize)]
pub struct Scalar(
#[cfg_attr(feature = "serde", serde(deserialize_with = "serde_deserialize_sk"))] [u8; 32],
);
impl Scalar {
fn from_bytes(bytes: [u8; 32]) -> Result<Self, ProtocolError> {
let scalar = scalar::clamp_integer(bytes);
if scalar.ct_eq(&bytes).into() {
Ok(Self(scalar))
} else {
Err(ProtocolError::SerializationError)
}
}
}
#[cfg(feature = "serde")]
fn serde_deserialize_sk<'de, D>(deserializer: D) -> Result<[u8; 32], D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::{Deserialize, Error};
Scalar::from_bytes(<[u8; 32]>::deserialize(deserializer)?)
.map(|scalar| scalar.0)
.map_err(D::Error::custom)
}
////////////////////////// //////////////////////////
// Test Implementations // // Test Implementations //
//===================== // //===================== //
@@ -98,9 +151,9 @@ impl DiffieHellman<Curve25519> for Scalar {
use crate::serialization::AssertZeroized; use crate::serialization::AssertZeroized;
#[cfg(test)] #[cfg(test)]
impl AssertZeroized for MontgomeryPoint { impl AssertZeroized for NonIdentity {
fn assert_zeroized(&self) { fn assert_zeroized(&self) {
assert_eq!(*self, MontgomeryPoint::default()); assert_eq!(self.0, MontgomeryPoint::default());
} }
} }
+169 -83
View File
@@ -46,15 +46,9 @@ impl Group for Ed25519 {
} }
fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> { fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> {
let compressed = bytes let bytes = bytes.take_array("public key")?;
.take_array("public key")
.map(|bytes| CompressedEdwardsY(bytes.into()))?;
if let Some(point) = compressed.decompress().filter(|point| !point.is_identity()) { VerifyingKey::from_bytes(bytes.into())
Ok(VerifyingKey { point, compressed })
} else {
Err(ProtocolError::SerializationError)
}
} }
fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk { fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk {
@@ -83,56 +77,6 @@ impl Group for Ed25519 {
} }
} }
/// Ed25519 verifying key.
// `ed25519_dalek::VerifyingKey` doesn't implement `Zeroize`.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/747.
// Required for manual implementation of EdDSA.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/556.
#[derive(Clone, Copy, Debug, Eq, PartialEq, Zeroize)]
pub struct VerifyingKey {
point: EdwardsPoint,
compressed: CompressedEdwardsY,
}
/// Ed25519 siging key.
// We store the `ExpandedSecret` in memory to avoid computing it on demand and then discarding it
// again.
#[derive(Clone, Copy, Debug, Eq, PartialEq, Zeroize)]
pub struct SigningKey {
// `ed25519_dalek::SigningKey` doesn't implement `Zeroize`. See
// https://github.com/dalek-cryptography/curve25519-dalek/pull/747
// Required for manual implementation of EdDSA.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/556.
sk: SecretKey,
verifying_key: VerifyingKey,
// `ed25519_dalek::ExpandedSecret` doesn't implement traits we need. See
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/748 and
// https://github.com/dalek-cryptography/curve25519-dalek/pull/747.
scalar: Scalar,
hash_prefix: [u8; 32],
}
impl SigningKey {
fn from_bytes(sk: [u8; 32]) -> Self {
let ExpandedSecretKey {
scalar,
hash_prefix,
} = ExpandedSecretKey::from(&sk);
let point = EdwardsPoint::mul_base(&scalar);
let verifying_key = VerifyingKey {
point,
compressed: point.compress(),
};
SigningKey {
sk,
verifying_key,
scalar,
hash_prefix,
}
}
}
impl PureEddsaImpl for Ed25519 { impl PureEddsaImpl for Ed25519 {
type Signature = Signature; type Signature = Signature;
type SignatureLen = U64; type SignatureLen = U64;
@@ -278,8 +222,175 @@ fn verify<'a>(
} }
} }
/// Ed25519 verifying key.
// `ed25519_dalek::VerifyingKey` doesn't implement `Zeroize`.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/747.
// Required for manual implementation of EdDSA.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/556.
#[derive(Clone, Copy, Debug, Eq, PartialEq, Zeroize)]
pub struct VerifyingKey {
point: EdwardsPoint,
compressed: CompressedEdwardsY,
}
impl VerifyingKey {
fn from_bytes(bytes: [u8; 32]) -> Result<Self, ProtocolError> {
let compressed = CompressedEdwardsY(bytes);
if let Some(point) = compressed.decompress().filter(|point| !point.is_identity()) {
Ok(Self { point, compressed })
} else {
Err(ProtocolError::SerializationError)
}
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for VerifyingKey {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use core::fmt::{self, Formatter};
use serde::de::{Deserialize, Deserializer, Error, SeqAccess, Visitor};
struct VerifyingKeyVisitor;
impl<'de> Visitor<'de> for VerifyingKeyVisitor {
type Value = VerifyingKey;
fn expecting(&self, formatter: &mut Formatter) -> fmt::Result {
Formatter::write_str(formatter, "tuple struct VerifyingKey")
}
fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
let compressed = CompressedEdwardsY::deserialize(deserializer)?;
VerifyingKey::from_bytes(compressed.0).map_err(Error::custom)
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let compressed: CompressedEdwardsY = seq.next_element()?.ok_or_else(|| {
Error::invalid_length(0, &"tuple struct VerifyingKey with 1 element")
})?;
VerifyingKey::from_bytes(compressed.0).map_err(Error::custom)
}
}
deserializer.deserialize_newtype_struct("VerifyingKey", VerifyingKeyVisitor)
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for VerifyingKey {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_newtype_struct("VerifyingKey", &self.compressed)
}
}
/// Ed25519 signing key.
// We store the `ExpandedSecret` in memory to avoid computing it on demand and then discarding it
// again.
#[derive(Clone, Copy, Debug, Eq, PartialEq, Zeroize)]
pub struct SigningKey {
// `ed25519_dalek::SigningKey` doesn't implement `Zeroize`. See
// https://github.com/dalek-cryptography/curve25519-dalek/pull/747
// Required for manual implementation of EdDSA.
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/556.
sk: SecretKey,
verifying_key: VerifyingKey,
// `ed25519_dalek::ExpandedSecret` doesn't implement traits we need. See
// TODO: remove after https://github.com/dalek-cryptography/curve25519-dalek/pull/748 and
// https://github.com/dalek-cryptography/curve25519-dalek/pull/747.
scalar: Scalar,
hash_prefix: [u8; 32],
}
impl SigningKey {
fn from_bytes(sk: [u8; 32]) -> Self {
let ExpandedSecretKey {
scalar,
hash_prefix,
} = ExpandedSecretKey::from(&sk);
let point = EdwardsPoint::mul_base(&scalar);
let verifying_key = VerifyingKey {
point,
compressed: point.compress(),
};
SigningKey {
sk,
verifying_key,
scalar,
hash_prefix,
}
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for SigningKey {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use core::fmt::{self, Formatter};
use serde::de::{Deserialize, Deserializer, Error, SeqAccess, Visitor};
struct SigningKeyVisitor;
impl<'de> Visitor<'de> for SigningKeyVisitor {
type Value = SigningKey;
fn expecting(&self, formatter: &mut Formatter) -> fmt::Result {
Formatter::write_str(formatter, "tuple struct SigningKey")
}
fn visit_newtype_struct<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: Deserializer<'de>,
{
let sk = Scalar::deserialize(deserializer)?;
Ok(SigningKey::from_bytes(sk.to_bytes()))
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let sk: Scalar = seq.next_element()?.ok_or_else(|| {
Error::invalid_length(0, &"tuple struct SigningKey with 1 element")
})?;
Ok(SigningKey::from_bytes(sk.to_bytes()))
}
}
deserializer.deserialize_newtype_struct("SigningKey", SigningKeyVisitor)
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for SigningKey {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_newtype_struct("SigningKey", &self.sk)
}
}
/// Ed25519 Signature. /// Ed25519 Signature.
// `ed25519_dalek::Signature` doesn't implement validation with Serde de/serialization. // `ed25519_dalek::Signature` doesn't implement validation with Serde de/serialization.
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[derive(Clone, Copy, Debug, Eq, PartialEq)] #[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[allow(non_snake_case)] #[allow(non_snake_case)]
pub struct Signature { pub struct Signature {
@@ -310,31 +421,6 @@ impl Signature {
} }
} }
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for Signature {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error;
Signature::deserialize_take(
&mut (GenericArray::<_, U64>::deserialize(deserializer)?.as_slice()),
)
.map_err(D::Error::custom)
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for Signature {
fn serialize<SK>(&self, serializer: SK) -> Result<SK::Ok, SK::Error>
where
SK: serde::Serializer,
{
self.serialize().serialize(serializer)
}
}
impl Zeroize for Signature { impl Zeroize for Signature {
fn zeroize(&mut self) { fn zeroize(&mut self) {
self.R.0 = [0; 32]; self.R.0 = [0; 32];
+23 -20
View File
@@ -86,16 +86,35 @@ where
} }
} }
impl<G> DiffieHellman<G> for NonZeroScalar<G>
where
G: CurveArithmetic + voprf::CipherSuite<Group = G> + voprf::Group<Scalar = Scalar<G>>,
FieldBytesSize<G>: ModulusSize,
ProjectivePoint<G>: GroupEncoding<
Repr = GenericArray<u8, <FieldBytesSize<G> as ModulusSize>::CompressedPointSize>,
> + ToEncodedPoint<G>,
{
fn diffie_hellman(
self,
pk: NonIdentity<G>,
) -> GenericArray<u8, <FieldBytesSize<G> as ModulusSize>::CompressedPointSize> {
GenericArray::clone_from_slice((pk.0 * self).to_encoded_point(true).as_bytes())
}
}
/// Wrapper around [`NonIdentity`](point::NonIdentity) to implement [`Zeroize`]. /// Wrapper around [`NonIdentity`](point::NonIdentity) to implement [`Zeroize`].
// TODO: remove after https://github.com/RustCrypto/traits/pull/1832. // TODO: remove after https://github.com/RustCrypto/traits/pull/1832.
#[derive_where(Clone, Copy)] #[derive_where(Clone, Copy)]
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound( serde(
deserialize = "point::NonIdentity<ProjectivePoint<G>>: serde::Deserialize<'de>", bound(
serialize = "point::NonIdentity<ProjectivePoint<G>>: serde::Serialize" deserialize = "point::NonIdentity<ProjectivePoint<G>>: serde::Deserialize<'de>",
)) serialize = "point::NonIdentity<ProjectivePoint<G>>: serde::Serialize"
),
transparent
)
)] )]
pub struct NonIdentity<G: CurveArithmetic>(pub point::NonIdentity<ProjectivePoint<G>>); pub struct NonIdentity<G: CurveArithmetic>(pub point::NonIdentity<ProjectivePoint<G>>);
@@ -121,22 +140,6 @@ impl<G: CurveArithmetic> Zeroize for NonIdentity<G> {
} }
} }
impl<G> DiffieHellman<G> for NonZeroScalar<G>
where
G: CurveArithmetic + voprf::CipherSuite<Group = G> + voprf::Group<Scalar = Scalar<G>>,
FieldBytesSize<G>: ModulusSize,
ProjectivePoint<G>: GroupEncoding<
Repr = GenericArray<u8, <FieldBytesSize<G> as ModulusSize>::CompressedPointSize>,
> + ToEncodedPoint<G>,
{
fn diffie_hellman(
self,
pk: NonIdentity<G>,
) -> GenericArray<u8, <FieldBytesSize<G> as ModulusSize>::CompressedPointSize> {
GenericArray::clone_from_slice((pk.0 * self).to_encoded_point(true).as_bytes())
}
}
////////////////////////// //////////////////////////
// Test Implementations // // Test Implementations //
//===================== // //===================== //
+86 -25
View File
@@ -12,13 +12,14 @@ pub use curve25519_dalek;
use curve25519_dalek::constants::RISTRETTO_BASEPOINT_POINT; use curve25519_dalek::constants::RISTRETTO_BASEPOINT_POINT;
use curve25519_dalek::ristretto::{CompressedRistretto, RistrettoPoint}; use curve25519_dalek::ristretto::{CompressedRistretto, RistrettoPoint};
use curve25519_dalek::scalar::Scalar; use curve25519_dalek::scalar::Scalar;
use curve25519_dalek::traits::Identity; use curve25519_dalek::traits::IsIdentity;
use digest::core_api::BlockSizeUser; use digest::core_api::BlockSizeUser;
use digest::{FixedOutput, HashMarker}; use digest::{FixedOutput, HashMarker};
use generic_array::GenericArray; use generic_array::GenericArray;
use generic_array::typenum::{IsLess, IsLessOrEqual, U32, U256}; use generic_array::typenum::{IsLess, IsLessOrEqual, U32, U256};
use rand::{CryptoRng, RngCore}; use rand::{CryptoRng, RngCore};
use voprf::Mode; use voprf::Mode;
use zeroize::Zeroize;
use super::{Group, STR_OPAQUE_DERIVE_AUTH_KEY_PAIR}; use super::{Group, STR_OPAQUE_DERIVE_AUTH_KEY_PAIR};
use crate::errors::{InternalError, ProtocolError}; use crate::errors::{InternalError, ProtocolError};
@@ -31,21 +32,20 @@ use crate::serialization::SliceExt;
pub struct Ristretto255; pub struct Ristretto255;
impl Group for Ristretto255 { impl Group for Ristretto255 {
type Pk = RistrettoPoint; type Pk = NonIdentity;
type PkLen = U32; type PkLen = U32;
type Sk = Scalar; type Sk = NonZeroScalar;
type SkLen = U32; type SkLen = U32;
fn serialize_pk(pk: Self::Pk) -> GenericArray<u8, Self::PkLen> { fn serialize_pk(pk: Self::Pk) -> GenericArray<u8, Self::PkLen> {
pk.compress().to_bytes().into() pk.0.compress().to_bytes().into()
} }
fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> { fn deserialize_take_pk(bytes: &mut &[u8]) -> Result<Self::Pk, ProtocolError> {
CompressedRistretto::from_slice(&bytes.take_array::<U32>("public key")?) CompressedRistretto(bytes.take_array("public key")?.into())
.map_err(|_| ProtocolError::SerializationError)?
.decompress() .decompress()
.filter(|point| point != &RistrettoPoint::identity())
.ok_or(ProtocolError::SerializationError) .ok_or(ProtocolError::SerializationError)
.and_then(NonIdentity::from_point)
} }
fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk { fn random_sk<R: RngCore + CryptoRng>(rng: &mut R) -> Self::Sk {
@@ -53,34 +53,101 @@ impl Group for Ristretto255 {
let scalar = Scalar::random(rng); let scalar = Scalar::random(rng);
if scalar != Scalar::ZERO { if scalar != Scalar::ZERO {
break scalar; break NonZeroScalar(scalar);
} }
} }
} }
fn derive_scalar(seed: GenericArray<u8, Self::SkLen>) -> Result<Self::Sk, InternalError> { fn derive_scalar(seed: GenericArray<u8, Self::SkLen>) -> Result<Self::Sk, InternalError> {
voprf::derive_key::<Self>(&seed, &STR_OPAQUE_DERIVE_AUTH_KEY_PAIR, Mode::Oprf) voprf::derive_key::<Self>(&seed, &STR_OPAQUE_DERIVE_AUTH_KEY_PAIR, Mode::Oprf)
.map(NonZeroScalar)
.map_err(InternalError::from) .map_err(InternalError::from)
} }
fn public_key(sk: Self::Sk) -> Self::Pk { fn public_key(sk: Self::Sk) -> Self::Pk {
RISTRETTO_BASEPOINT_POINT * sk NonIdentity(RISTRETTO_BASEPOINT_POINT * sk.0)
} }
fn serialize_sk(sk: Self::Sk) -> GenericArray<u8, Self::SkLen> { fn serialize_sk(sk: Self::Sk) -> GenericArray<u8, Self::SkLen> {
sk.to_bytes().into() sk.0.to_bytes().into()
} }
fn deserialize_take_sk(bytes: &mut &[u8]) -> Result<Self::Sk, ProtocolError> { fn deserialize_take_sk(bytes: &mut &[u8]) -> Result<Self::Sk, ProtocolError> {
bytes Scalar::from_canonical_bytes(bytes.take_array("secret key")?.into())
.take_array::<U32>("secret key") .into_option()
.ok()
.and_then(|bytes| Scalar::from_canonical_bytes(bytes.into()).into())
.filter(|scalar| scalar != &Scalar::ZERO)
.ok_or(ProtocolError::SerializationError) .ok_or(ProtocolError::SerializationError)
.and_then(NonZeroScalar::from_scalar)
} }
} }
impl DiffieHellman<Ristretto255> for NonZeroScalar {
fn diffie_hellman(self, pk: NonIdentity) -> GenericArray<u8, U32> {
Ristretto255::serialize_pk(NonIdentity(pk.0 * self.0))
}
}
/// Non-identity point wrapper for [`RistrettoPoint`].
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Zeroize)]
pub struct NonIdentity(
#[cfg_attr(feature = "serde", serde(deserialize_with = "serde_deserialize_pk"))] RistrettoPoint,
);
impl NonIdentity {
fn from_point(point: RistrettoPoint) -> Result<Self, ProtocolError> {
if point.is_identity() {
Err(ProtocolError::SerializationError)
} else {
Ok(NonIdentity(point))
}
}
}
#[cfg(feature = "serde")]
fn serde_deserialize_pk<'de, D>(deserializer: D) -> Result<RistrettoPoint, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::{Deserialize, Error};
let point = RistrettoPoint::deserialize(deserializer)?;
NonIdentity::from_point(point)
.map(|point| point.0)
.map_err(Error::custom)
}
/// Non-zero scalar wrapper for [`Scalar`]
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Zeroize)]
pub struct NonZeroScalar(
#[cfg_attr(feature = "serde", serde(deserialize_with = "serde_deserialize_sk"))] Scalar,
);
impl NonZeroScalar {
fn from_scalar(scalar: Scalar) -> Result<Self, ProtocolError> {
if scalar == Scalar::ZERO {
Err(ProtocolError::SerializationError)
} else {
Ok(Self(scalar))
}
}
}
#[cfg(feature = "serde")]
fn serde_deserialize_sk<'de, D>(deserializer: D) -> Result<Scalar, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::{Deserialize, Error};
let scalar = Scalar::deserialize(deserializer)?;
NonZeroScalar::from_scalar(scalar)
.map(|scalar| scalar.0)
.map_err(Error::custom)
}
impl voprf::CipherSuite for Ristretto255 { impl voprf::CipherSuite for Ristretto255 {
const ID: &'static str = voprf::Ristretto255::ID; const ID: &'static str = voprf::Ristretto255::ID;
@@ -157,12 +224,6 @@ impl voprf::Group for Ristretto255 {
} }
} }
impl DiffieHellman<Ristretto255> for Scalar {
fn diffie_hellman(self, pk: RistrettoPoint) -> GenericArray<u8, U32> {
Ristretto255::serialize_pk(pk * self)
}
}
////////////////////////// //////////////////////////
// Test Implementations // // Test Implementations //
//===================== // //===================== //
@@ -172,15 +233,15 @@ impl DiffieHellman<Ristretto255> for Scalar {
use crate::serialization::AssertZeroized; use crate::serialization::AssertZeroized;
#[cfg(test)] #[cfg(test)]
impl AssertZeroized for RistrettoPoint { impl AssertZeroized for NonIdentity {
fn assert_zeroized(&self) { fn assert_zeroized(&self) {
assert_eq!(*self, RistrettoPoint::default()); assert_eq!(self.0, RistrettoPoint::default());
} }
} }
#[cfg(test)] #[cfg(test)]
impl AssertZeroized for Scalar { impl AssertZeroized for NonZeroScalar {
fn assert_zeroized(&self) { fn assert_zeroized(&self) {
assert_eq!(*self, Scalar::default()); assert_eq!(self.0, Scalar::default());
} }
} }
+8 -2
View File
@@ -57,7 +57,10 @@ pub trait DiffieHellman<G: Group> {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "G::Sk: serde::Deserialize<'de>",
serialize = "G::Sk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)]
@@ -70,7 +73,10 @@ pub struct Ke1State<G: Group> {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "G::Pk: serde::Deserialize<'de>",
serialize = "G::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
+1 -1
View File
@@ -136,7 +136,7 @@ where
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound = "", transparent)
)] )]
pub struct Signature<G: CurveArithmetic + PrimeCurve>(pub ecdsa::Signature<G>) pub struct Signature<G: CurveArithmetic + PrimeCurve>(pub ecdsa::Signature<G>)
where where
+12 -6
View File
@@ -144,10 +144,14 @@ pub trait SignatureProtocol {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound(deserialize = "'de: 'a", serialize = "")) serde(bound(
deserialize = "'de: 'a, <KeGroup<CS> as Group>::Pk: serde::Deserialize<'de>, KE::Pk: \
serde::Deserialize<'de>",
serialize = "<KeGroup<CS> as Group>::Pk: serde::Serialize, KE::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, PartialEq; PublicKey<KeGroup<CS>>, PublicKey<KE>)] #[derive_where(Debug, Eq, Hash, PartialEq; <KeGroup<CS> as Group>::Pk, KE::Pk)]
pub struct Ke2Builder<'a, CS: CipherSuite, KE: Group> { pub struct Ke2Builder<'a, CS: CipherSuite, KE: Group> {
transcript: Message<'a, CS, KE>, transcript: Message<'a, CS, KE>,
server_nonce: GenericArray<u8, NonceLen>, server_nonce: GenericArray<u8, NonceLen>,
@@ -166,8 +170,10 @@ pub struct Ke2Builder<'a, CS: CipherSuite, KE: Group> {
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound( serde(bound(
deserialize = "SIG::VerifyState<CS, KE>: serde::Deserialize<'de>", deserialize = "<SIG::Group as Group>::Pk: serde::Deserialize<'de>, SIG::VerifyState<CS, \
serialize = "SIG::VerifyState<CS, KE>: serde::Serialize" KE>: serde::Deserialize<'de>",
serialize = "<SIG::Group as Group>::Pk: serde::Serialize, SIG::VerifyState<CS, KE>: \
serde::Serialize"
)) ))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
@@ -184,8 +190,8 @@ pub struct Ke2State<CS: CipherSuite, SIG: SignatureProtocol, KE: Group> {
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound( serde(bound(
deserialize = "SIG::Signature: serde::Deserialize<'de>", deserialize = "KE::Pk: serde::Deserialize<'de>, SIG::Signature: serde::Deserialize<'de>",
serialize = "SIG::Signature: serde::Serialize" serialize = "KE::Pk: serde::Serialize, SIG::Signature: serde::Serialize"
)) ))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
+4 -1
View File
@@ -95,7 +95,10 @@ where
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "G::Pk: serde::Deserialize<'de>",
serialize = "G::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
+18 -54
View File
@@ -27,8 +27,8 @@ use crate::serialization::SliceExt;
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound( serde(bound(
deserialize = "SK: serde::Deserialize<'de>", deserialize = "G::Pk: serde::Deserialize<'de>, SK: serde::Deserialize<'de>",
serialize = "SK: serde::Serialize" serialize = "G::Pk: serde::Serialize, SK: serde::Serialize"
)) ))
)] )]
#[derive_where(Clone)] #[derive_where(Clone)]
@@ -83,6 +83,14 @@ impl<G: Group> KeyPair<G> {
} }
/// Wrapper around a Key to enforce that it's a private one. /// Wrapper around a Key to enforce that it's a private one.
#[cfg_attr(
feature = "serde",
derive(serde::Deserialize, serde::Serialize),
serde(bound(
deserialize = "G::Sk: serde::Deserialize<'de>",
serialize = "G::Sk: serde::Serialize"
))
)]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)]
pub struct PrivateKey<G: Group>(G::Sk); pub struct PrivateKey<G: Group>(G::Sk);
@@ -173,33 +181,15 @@ impl<G: Group> PrivateKeySerialization<G> for PrivateKey<G> {
} }
} }
#[cfg(feature = "serde")]
impl<'de, G: Group> serde::Deserialize<'de> for PrivateKey<G> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error;
G::deserialize_take_sk(
&mut (GenericArray::<_, G::SkLen>::deserialize(deserializer)?.as_slice()),
)
.map(Self)
.map_err(D::Error::custom)
}
}
#[cfg(feature = "serde")]
impl<G: Group> serde::Serialize for PrivateKey<G> {
fn serialize<SK>(&self, serializer: SK) -> Result<SK::Ok, SK::Error>
where
SK: serde::Serializer,
{
G::serialize_sk(self.0).serialize(serializer)
}
}
/// Wrapper around a Key to enforce that it's a public one. /// Wrapper around a Key to enforce that it's a public one.
#[cfg_attr(
feature = "serde",
derive(serde::Deserialize, serde::Serialize),
serde(bound(
deserialize = "G::Pk: serde::Deserialize<'de>",
serialize = "G::Pk: serde::Serialize"
))
)]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
pub struct PublicKey<G: Group + ?Sized>(G::Pk); pub struct PublicKey<G: Group + ?Sized>(G::Pk);
@@ -237,32 +227,6 @@ impl<G: Group> PublicKey<G> {
} }
} }
#[cfg(feature = "serde")]
impl<'de, G: Group> serde::Deserialize<'de> for PublicKey<G> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::de::Error;
G::deserialize_take_pk(
&mut (GenericArray::<_, G::PkLen>::deserialize(deserializer)?.as_slice()),
)
.map(Self)
.map_err(D::Error::custom)
}
}
#[cfg(feature = "serde")]
impl<G: Group> serde::Serialize for PublicKey<G> {
fn serialize<SK>(&self, serializer: SK) -> Result<SK::Ok, SK::Error>
where
SK: serde::Serializer,
{
G::serialize_pk(self.0).serialize(serializer)
}
}
/// Default OPRF seed container. /// Default OPRF seed container.
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
+8 -2
View File
@@ -58,7 +58,10 @@ pub struct RegistrationRequest<CS: CipherSuite> {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "<KeGroup<CS> as Group>::Pk: serde::Deserialize<'de>",
serialize = "<KeGroup<CS> as Group>::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone)] #[derive_where(Clone)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; voprf::EvaluationElement<CS::OprfCs>, <KeGroup<CS> as Group>::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; voprf::EvaluationElement<CS::OprfCs>, <KeGroup<CS> as Group>::Pk)]
@@ -74,7 +77,10 @@ pub struct RegistrationResponse<CS: CipherSuite> {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "<KeGroup<CS> as Group>::Pk: serde::Deserialize<'de>",
serialize = "<KeGroup<CS> as Group>::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <KeGroup<CS> as Group>::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <KeGroup<CS> as Group>::Pk)]
+9 -3
View File
@@ -62,8 +62,11 @@ const STR_OPAQUE_DERIVE_KEY_PAIR: &[u8; 20] = b"OPAQUE-DeriveKeyPair";
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound( serde(bound(
deserialize = "SK: serde::Deserialize<'de>, OS: serde::Deserialize<'de>", deserialize = "<KeGroup<CS> as Group>::Pk: serde::Deserialize<'de>, <KeGroup<CS> as \
serialize = "SK: serde::Serialize, OS: serde::Serialize" Group>::Sk: serde::Deserialize<'de>, SK: serde::Deserialize<'de>, OS: \
serde::Deserialize<'de>",
serialize = "<KeGroup<CS> as Group>::Pk: serde::Serialize, <KeGroup<CS> as Group>::Sk: \
serde::Serialize, SK: serde::Serialize, OS: serde::Serialize"
)) ))
)] )]
#[derive_where(Clone)] #[derive_where(Clone)]
@@ -99,7 +102,10 @@ pub struct ClientRegistration<CS: CipherSuite> {
#[cfg_attr( #[cfg_attr(
feature = "serde", feature = "serde",
derive(serde::Deserialize, serde::Serialize), derive(serde::Deserialize, serde::Serialize),
serde(bound = "") serde(bound(
deserialize = "<KeGroup<CS> as Group>::Pk: serde::Deserialize<'de>",
serialize = "<KeGroup<CS> as Group>::Pk: serde::Serialize"
))
)] )]
#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Clone, ZeroizeOnDrop)]
#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <KeGroup<CS> as Group>::Pk)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <KeGroup<CS> as Group>::Pk)]