From d324584d792fa03ea31344c7ffa7323b48eecb64 Mon Sep 17 00:00:00 2001 From: daxpedda Date: Tue, 22 Apr 2025 09:08:17 +0200 Subject: [PATCH] Rework `SecretKey` API to facilitate async (#371) * Rework `SecretKey` API to facilitate async * Remove left-over constraints --- .github/workflows/main.yml | 11 + Cargo.toml | 14 +- deny.toml | 3 +- src/ciphersuite.rs | 2 +- src/envelope.rs | 6 +- src/errors.rs | 114 +++---- src/key_exchange/group/curve25519.rs | 45 +-- src/key_exchange/group/elliptic_curve.rs | 31 +- src/key_exchange/group/mod.rs | 9 +- src/key_exchange/group/ristretto255.rs | 23 +- src/key_exchange/traits.rs | 34 ++- src/key_exchange/tripledh.rs | 218 ++++++++----- src/keypair.rs | 157 +++++----- src/lib.rs | 84 +++-- src/messages.rs | 66 +++- src/opaque.rs | 135 +++++--- src/serialization/tests.rs | 2 +- src/tests/full_test.rs | 5 +- src/tests/mod.rs | 2 + src/tests/remote_key.rs | 373 +++++++++++++++++++++++ 20 files changed, 952 insertions(+), 382 deletions(-) create mode 100644 src/tests/remote_key.rs diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index a1063c6..5678a6e 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -25,6 +25,10 @@ jobs: - stable - 1.83.0 name: test + env: + PKCS11_MODULE: /usr/lib/softhsm/libsofthsm2.so + SOFTHSM2_CONF: /tmp/softhsm2.conf + RUSTFLAGS: --cfg test_hsm steps: - name: Checkout sources uses: actions/checkout@main @@ -36,6 +40,13 @@ jobs: toolchain: ${{ matrix.toolchain }} override: true + - name: Install SoftHSM + run: | + sudo apt-get update -y -qq && + sudo apt-get install -y -qq libsofthsm2 && + mkdir /tmp/tokens + echo "directories.tokendir = /tmp/tokens" > /tmp/softhsm2.conf + - name: Run cargo test uses: actions-rs/cargo@v1 with: diff --git a/Cargo.toml b/Cargo.toml index 52b3b19..f5118e5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -17,7 +17,7 @@ curve25519 = ["dep:curve25519-dalek"] default = ["ristretto255-voprf", "serde"] ristretto255 = ["dep:curve25519-dalek", "voprf/ristretto255"] ristretto255-voprf = ["ristretto255", "voprf/ristretto255-ciphersuite"] -serde = ["dep:serde", "generic-array/serde", "voprf/serde"] +serde = ["dep:serde", "generic-array/serde", "voprf/serde", "zeroize/serde"] std = ["dep:getrandom"] [dependencies] @@ -27,7 +27,7 @@ argon2 = { version = "0.5", default-features = false, features = [ curve25519-dalek = { version = "4", default-features = false, features = [ "zeroize", ], optional = true } -derive-where = { version = "1", features = ["zeroize-on-drop"] } +derive-where = { version = "1.3", features = ["zeroize-on-drop"] } digest = "0.10" displaydoc = { version = "0.2", default-features = false } elliptic-curve = { version = "0.13", features = ["hash2curve", "sec1"] } @@ -46,25 +46,32 @@ zeroize = { version = "1.8", features = ["zeroize_derive"] } getrandom = { version = "0.2", features = ["js"], optional = true } [dev-dependencies] +anyhow = "1" bincode = "1" chacha20poly1305 = "0.10" criterion = "0.5" +cryptoki = "0.9" +elliptic-curve = { version = "0.13", features = ["alloc", "pkcs8"] } hex = "0.4" p256 = { version = "0.13", default-features = false, features = [ "hash2curve", + "pkcs8", "voprf", ] } p384 = { version = "0.13", default-features = false, features = [ "hash2curve", + "pkcs8", "voprf", ] } p521 = { version = "0.13.3", default-features = false, features = [ "hash2curve", + "pkcs8", "voprf", ] } proptest = "1" rand = "0.8" regex = "1" +thiserror = "2" # MSRV rustyline = "15" scrypt = "0.11" @@ -82,3 +89,6 @@ targets = [] [[example]] name = "simple_login" required-features = ["argon2"] + +[lints.rust] +unexpected_cfgs = { level = "warn", check-cfg = ['cfg(test_hsm)'] } diff --git a/deny.toml b/deny.toml index 2c6e70a..34e7817 100644 --- a/deny.toml +++ b/deny.toml @@ -42,7 +42,8 @@ yanked = "warn" # A list of advisory IDs to ignore. Note that ignored advisories will still # output a note when they are encountered. ignore = [ - #"RUSTSEC-0000-0000", + # dev-dependency + "RUSTSEC-2024-0436", ] # Threshold for security vulnerabilities, any vulnerability with a CVSS score # lower than the range specified will be ignored. Note that ignored advisories diff --git a/src/ciphersuite.rs b/src/ciphersuite.rs index 69c9ede..5ea8787 100644 --- a/src/ciphersuite.rs +++ b/src/ciphersuite.rs @@ -36,7 +36,7 @@ where /// A VOPRF ciphersuite, see [`voprf::CipherSuite`]. type OprfCs: voprf::CipherSuite; /// A `Group` used for the `KeyExchange`. - type KeGroup: KeGroup; + type KeGroup: 'static + KeGroup; /// A key exchange protocol type KeyExchange: KeyExchange, Self::KeGroup>; /// A key stretching function, typically used for password hashing diff --git a/src/envelope.rs b/src/envelope.rs index 79230ba..b117b20 100644 --- a/src/envelope.rs +++ b/src/envelope.rs @@ -24,7 +24,7 @@ use crate::errors::utils::check_slice_size; use crate::errors::{InternalError, ProtocolError}; use crate::hash::OutputSize; use crate::key_exchange::group::KeGroup; -use crate::keypair::{KeyPair, PublicKey}; +use crate::keypair::{KeyPair, PrivateKey, PrivateKeySerialization, PublicKey}; use crate::opaque::{bytestrings_from_identifiers, Identifiers}; use crate::serialization::{Input, MacExt}; @@ -303,7 +303,7 @@ fn build_inner_envelope_internal( .expand(&nonce.concat(STR_PRIVATE_KEY.into()), &mut keypair_seed) .map_err(|_| InternalError::HkdfError)?; let client_static_keypair = - KeyPair::::from_private_key_slice(&CS::KeGroup::serialize_sk( + PrivateKey::::deserialize_key_pair(&CS::KeGroup::serialize_sk( CS::KeGroup::derive_auth_keypair::(keypair_seed)?, ))?; @@ -319,7 +319,7 @@ fn recover_keys_internal( .expand(&nonce.concat(STR_PRIVATE_KEY.into()), &mut keypair_seed) .map_err(|_| InternalError::HkdfError)?; let client_static_keypair = - KeyPair::::from_private_key_slice(&CS::KeGroup::serialize_sk( + PrivateKey::::deserialize_key_pair(&CS::KeGroup::serialize_sk( CS::KeGroup::derive_auth_keypair::(keypair_seed)?, ))?; diff --git a/src/errors.rs b/src/errors.rs index d9387ef..4c3e125 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -15,23 +15,7 @@ use displaydoc::Display; /// Represents an error in the manipulation of internal cryptographic data #[derive(Clone, Copy, Display, Eq, Hash, Ord, PartialEq, PartialOrd)] -pub enum InternalError { - /// Custom [`SecretKey`](crate::keypair::SecretKey) error type - Custom(T), - /// Deserializing from a byte sequence failed - InvalidByteSequence, - #[allow(clippy::doc_markdown)] - /// Invalid length for {name}: expected {len}, but is actually {actual_len}. - SizeError { - /// name - name: &'static str, - /// length - len: usize, - /// actual - actual_len: usize, - }, - /// Could not decompress point. - PointError, +pub enum InternalError { /// Size of input is empty or longer then [`u16::MAX`]. HashToScalar, /// Computing HKDF failed while deriving subkeys @@ -52,22 +36,9 @@ pub enum InternalError { OprfInternalError(voprf::InternalError), } -impl Debug for InternalError { +impl Debug for InternalError { fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { match self { - Self::Custom(custom) => f.debug_tuple("InvalidByteSequence").field(custom).finish(), - Self::InvalidByteSequence => f.debug_tuple("InvalidByteSequence").finish(), - Self::SizeError { - name, - len, - actual_len, - } => f - .debug_struct("SizeError") - .field("name", name) - .field("len", len) - .field("actual_len", actual_len) - .finish(), - Self::PointError => f.debug_tuple("PointError").finish(), Self::HashToScalar => f.debug_tuple("HashToScalar").finish(), Self::HkdfError => f.debug_tuple("HkdfError").finish(), Self::HmacError => f.debug_tuple("HmacError").finish(), @@ -84,35 +55,7 @@ impl Debug for InternalError { } } -impl Error for InternalError {} - -impl InternalError { - /// Convert `InternalError` into `InternalError` - pub fn into_custom(self) -> InternalError { - match self { - Self::Custom(_) => unreachable!(), - Self::InvalidByteSequence => InternalError::InvalidByteSequence, - Self::SizeError { - name, - len, - actual_len, - } => InternalError::SizeError { - name, - len, - actual_len, - }, - Self::PointError => InternalError::PointError, - Self::HashToScalar => InternalError::HashToScalar, - Self::HkdfError => InternalError::HkdfError, - Self::HmacError => InternalError::HmacError, - Self::KsfError => InternalError::KsfError, - Self::SealOpenHmacError => InternalError::SealOpenHmacError, - Self::IncompatibleEnvelopeModeError => InternalError::IncompatibleEnvelopeModeError, - Self::OprfError(error) => InternalError::OprfError(error), - Self::OprfInternalError(error) => InternalError::OprfInternalError(error), - } - } -} +impl Error for InternalError {} impl From for InternalError { fn from(voprf_error: voprf::Error) -> Self { @@ -136,17 +79,28 @@ impl From for ProtocolError { #[derive(Clone, Copy, Display, Eq, Hash, Ord, PartialEq, PartialOrd)] pub enum ProtocolError { /// Internal error encountered - LibraryError(InternalError), + LibraryError(InternalError), /// Error in validating credentials InvalidLoginError, /// Error with serializing / deserializing protocol messages SerializationError, + /// Invalid length for `name`: expected `len`, but is actually `actual_len`. + SizeError { + /// name + name: &'static str, + /// length + len: usize, + /// actual + actual_len: usize, + }, /** This error occurs when the client detects that the server has reflected the OPRF value (beta == alpha) */ ReflectedValueError, /** Identity group element was encountered during deserialization, which is invalid */ IdentityGroupElementError, + /// Custom [`SecretKey`](crate::keypair::PrivateKeySerialization) error type + Custom(T), } impl Debug for ProtocolError { @@ -157,8 +111,19 @@ impl Debug for ProtocolError { } Self::InvalidLoginError => f.debug_tuple("InvalidLoginError").finish(), Self::SerializationError => f.debug_tuple("SerializationError").finish(), + Self::SizeError { + name, + len, + actual_len, + } => f + .debug_struct("SizeError") + .field("name", name) + .field("len", len) + .field("actual_len", actual_len) + .finish(), Self::ReflectedValueError => f.debug_tuple("ReflectedValueError").finish(), Self::IdentityGroupElementError => f.debug_tuple("IdentityGroupElementError").finish(), + Self::Custom(custom) => f.debug_tuple("Custom").field(custom).finish(), } } } @@ -167,8 +132,8 @@ impl Error for ProtocolError {} // This is meant to express future(ly) non-trivial ways of converting the // internal error into a ProtocolError -impl From> for ProtocolError { - fn from(e: InternalError) -> ProtocolError { +impl From for ProtocolError { + fn from(e: InternalError) -> ProtocolError { Self::LibraryError(e) } } @@ -186,11 +151,18 @@ impl ProtocolError { /// Convert `ProtocolError` into `ProtocolError` pub fn into_custom(self) -> ProtocolError { match self { - Self::LibraryError(internal_error) => { - ProtocolError::LibraryError(internal_error.into_custom()) - } + Self::LibraryError(internal_error) => ProtocolError::LibraryError(internal_error), Self::InvalidLoginError => ProtocolError::InvalidLoginError, Self::SerializationError => ProtocolError::SerializationError, + Self::SizeError { + name, + len, + actual_len, + } => ProtocolError::SizeError { + name, + len, + actual_len, + }, Self::ReflectedValueError => ProtocolError::ReflectedValueError, Self::IdentityGroupElementError => ProtocolError::IdentityGroupElementError, } @@ -200,13 +172,13 @@ impl ProtocolError { pub(crate) mod utils { use super::*; - pub fn check_slice_size<'a, T>( + pub fn check_slice_size<'a>( slice: &'a [u8], expected_len: usize, arg_name: &'static str, - ) -> Result<&'a [u8], InternalError> { + ) -> Result<&'a [u8], ProtocolError> { if slice.len() != expected_len { - return Err(InternalError::SizeError { + return Err(ProtocolError::SizeError { name: arg_name, len: expected_len, actual_len: slice.len(), @@ -219,9 +191,9 @@ pub(crate) mod utils { slice: &'a [u8], expected_len: usize, arg_name: &'static str, - ) -> Result<&'a [u8], InternalError> { + ) -> Result<&'a [u8], ProtocolError> { if slice.len() < expected_len { - return Err(InternalError::SizeError { + return Err(ProtocolError::SizeError { name: arg_name, len: expected_len, actual_len: slice.len(), diff --git a/src/key_exchange/group/curve25519.rs b/src/key_exchange/group/curve25519.rs index 25a8b55..97221a8 100644 --- a/src/key_exchange/group/curve25519.rs +++ b/src/key_exchange/group/curve25519.rs @@ -9,7 +9,7 @@ //! Key Exchange group implementation for Curve25519 use curve25519_dalek::montgomery::MontgomeryPoint; -use curve25519_dalek::scalar::{self, Scalar}; +use curve25519_dalek::scalar; use curve25519_dalek::traits::Identity; use digest::core_api::BlockSizeUser; use digest::{FixedOutput, HashMarker, OutputSizeUser}; @@ -17,9 +17,11 @@ use generic_array::typenum::{IsLess, IsLessOrEqual, U256, U32}; use generic_array::GenericArray; use rand::{CryptoRng, RngCore}; use subtle::ConstantTimeEq; +use zeroize::Zeroize; use super::KeGroup; -use crate::errors::InternalError; +use crate::errors::{InternalError, ProtocolError}; +use crate::key_exchange::tripledh::DiffieHellman; /// Implementation for Curve25519. pub struct Curve25519; @@ -28,20 +30,20 @@ pub struct Curve25519; impl KeGroup for Curve25519 { type Pk = MontgomeryPoint; type PkLen = U32; - type Sk = [u8; 32]; + type Sk = Scalar; type SkLen = U32; fn serialize_pk(pk: Self::Pk) -> GenericArray { pk.to_bytes().into() } - fn deserialize_pk(bytes: &[u8]) -> Result { + fn deserialize_pk(bytes: &[u8]) -> Result { bytes .try_into() .ok() .map(MontgomeryPoint) .filter(|pk| pk != &MontgomeryPoint::identity()) - .ok_or(InternalError::PointError) + .ok_or(ProtocolError::SerializationError) } fn random_sk(rng: &mut R) -> Self::Sk { @@ -51,8 +53,8 @@ impl KeGroup for Curve25519 { rng.fill_bytes(&mut scalar_bytes); let scalar = scalar::clamp_integer(scalar_bytes); - if scalar != Scalar::ZERO.to_bytes() { - break scalar; + if scalar != curve25519_dalek::Scalar::ZERO.to_bytes() { + break Scalar(scalar); } } } @@ -72,26 +74,22 @@ impl KeGroup for Curve25519 { ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { - Ok(scalar::clamp_integer(seed.into())) + Ok(Scalar(scalar::clamp_integer(seed.into()))) } fn is_zero_scalar(scalar: Self::Sk) -> subtle::Choice { - scalar.ct_eq(&Scalar::ZERO.to_bytes()) + scalar.0.ct_eq(&curve25519_dalek::Scalar::ZERO.to_bytes()) } fn public_key(sk: Self::Sk) -> Self::Pk { - MontgomeryPoint::mul_base_clamped(sk) - } - - fn diffie_hellman(pk: Self::Pk, sk: Self::Sk) -> GenericArray { - Self::serialize_pk(pk.mul_clamped(sk)) + MontgomeryPoint::mul_base_clamped(sk.0) } fn serialize_sk(sk: Self::Sk) -> GenericArray { - sk.into() + sk.0.into() } - fn deserialize_sk(bytes: &[u8]) -> Result { + fn deserialize_sk(bytes: &[u8]) -> Result { bytes .try_into() .ok() @@ -99,7 +97,18 @@ impl KeGroup for Curve25519 { let scalar = scalar::clamp_integer(bytes); (scalar == bytes).then_some(scalar) }) - .filter(|scalar| scalar != &Scalar::ZERO.to_bytes()) - .ok_or(InternalError::PointError) + .filter(|scalar| scalar != &curve25519_dalek::Scalar::ZERO.to_bytes()) + .map(Scalar) + .ok_or(ProtocolError::SerializationError) + } +} + +/// Curve25519 scalar. +#[derive(Clone, Copy, Zeroize)] +pub struct Scalar([u8; 32]); + +impl DiffieHellman for Scalar { + fn diffie_hellman(self, pk: MontgomeryPoint) -> GenericArray { + Curve25519::serialize_pk(pk.mul_clamped(self.0)) } } diff --git a/src/key_exchange/group/elliptic_curve.rs b/src/key_exchange/group/elliptic_curve.rs index 782f53b..36f5c6e 100644 --- a/src/key_exchange/group/elliptic_curve.rs +++ b/src/key_exchange/group/elliptic_curve.rs @@ -19,7 +19,8 @@ use generic_array::GenericArray; use rand::{CryptoRng, RngCore}; use super::KeGroup; -use crate::errors::InternalError; +use crate::errors::{InternalError, ProtocolError}; +use crate::key_exchange::tripledh::DiffieHellman; impl KeGroup for G where @@ -41,10 +42,10 @@ where GenericArray::clone_from_slice(pk.to_encoded_point(true).as_bytes()) } - fn deserialize_pk(bytes: &[u8]) -> Result { + fn deserialize_pk(bytes: &[u8]) -> Result { PublicKey::::from_sec1_bytes(bytes) .map(|public_key| public_key.to_projective()) - .map_err(|_| InternalError::PointError) + .map_err(|_| ProtocolError::SerializationError) } fn random_sk(rng: &mut R) -> Self::Sk { @@ -77,17 +78,29 @@ where scalar.is_zero() } - fn diffie_hellman(pk: Self::Pk, sk: Self::Sk) -> GenericArray { - Self::serialize_pk(pk * sk) - } - fn serialize_sk(sk: Self::Sk) -> GenericArray { sk.into() } - fn deserialize_sk(bytes: &[u8]) -> Result { + fn deserialize_sk(bytes: &[u8]) -> Result { SecretKey::::from_slice(bytes) .map(|secret_key| *secret_key.to_nonzero_scalar()) - .map_err(|_| InternalError::PointError) + .map_err(|_| ProtocolError::SerializationError) + } +} + +impl DiffieHellman for Scalar +where + G: GroupDigest, + FieldBytesSize: ModulusSize, + AffinePoint: FromEncodedPoint + ToEncodedPoint, + ProjectivePoint: CofactorGroup + ToEncodedPoint, + Scalar: FromOkm, +{ + fn diffie_hellman( + self, + pk: ProjectivePoint, + ) -> GenericArray as ModulusSize>::CompressedPointSize> { + GenericArray::clone_from_slice((pk * self).to_encoded_point(true).as_bytes()) } } diff --git a/src/key_exchange/group/mod.rs b/src/key_exchange/group/mod.rs index 12c6de4..e665aea 100644 --- a/src/key_exchange/group/mod.rs +++ b/src/key_exchange/group/mod.rs @@ -22,7 +22,7 @@ use generic_array::{ArrayLength, GenericArray}; use rand::{CryptoRng, RngCore}; use zeroize::Zeroize; -use crate::errors::InternalError; +use crate::errors::{InternalError, ProtocolError}; const STR_OPAQUE_DERIVE_AUTH_KEY_PAIR: [u8; 33] = *b"OPAQUE-DeriveDiffieHellmanKeyPair"; @@ -41,7 +41,7 @@ pub trait KeGroup { fn serialize_pk(pk: Self::Pk) -> GenericArray; /// Return a public key from its fixed-length bytes representation - fn deserialize_pk(bytes: &[u8]) -> Result; + fn deserialize_pk(bytes: &[u8]) -> Result; /// Generate a random secret key fn random_sk(rng: &mut R) -> Self::Sk; @@ -104,14 +104,11 @@ pub trait KeGroup { /// Return a public key from its secret key fn public_key(sk: Self::Sk) -> Self::Pk; - /// Diffie-Hellman key exchange - fn diffie_hellman(pk: Self::Pk, sk: Self::Sk) -> GenericArray; - /// Serializes `self` fn serialize_sk(sk: Self::Sk) -> GenericArray; /// Return a public key from its fixed-length bytes representation - fn deserialize_sk(bytes: &[u8]) -> Result; + fn deserialize_sk(bytes: &[u8]) -> Result; } // Helper functions used to compute DeriveAuthKeyPair() (taken from the voprf diff --git a/src/key_exchange/group/ristretto255.rs b/src/key_exchange/group/ristretto255.rs index b5d9194..55a4367 100644 --- a/src/key_exchange/group/ristretto255.rs +++ b/src/key_exchange/group/ristretto255.rs @@ -21,7 +21,8 @@ use subtle::ConstantTimeEq; use voprf::Group; use super::KeGroup; -use crate::errors::InternalError; +use crate::errors::{InternalError, ProtocolError}; +use crate::key_exchange::tripledh::DiffieHellman; /// Implementation for Ristretto255. // This is necessary because Rust lacks specialization, otherwise we could @@ -38,12 +39,12 @@ impl KeGroup for Ristretto255 { pk.compress().to_bytes().into() } - fn deserialize_pk(bytes: &[u8]) -> Result { + fn deserialize_pk(bytes: &[u8]) -> Result { CompressedRistretto::from_slice(bytes) - .map_err(|_| InternalError::PointError)? + .map_err(|_| ProtocolError::SerializationError)? .decompress() .filter(|point| point != &RistrettoPoint::identity()) - .ok_or(InternalError::PointError) + .ok_or(ProtocolError::SerializationError) } fn random_sk(rng: &mut R) -> Self::Sk { @@ -89,21 +90,17 @@ impl KeGroup for Ristretto255 { RISTRETTO_BASEPOINT_POINT * sk } - fn diffie_hellman(pk: Self::Pk, sk: Self::Sk) -> GenericArray { - Self::serialize_pk(pk * sk) - } - fn serialize_sk(sk: Self::Sk) -> GenericArray { sk.to_bytes().into() } - fn deserialize_sk(bytes: &[u8]) -> Result { + fn deserialize_sk(bytes: &[u8]) -> Result { bytes .try_into() .ok() .and_then(|bytes| Scalar::from_canonical_bytes(bytes).into()) .filter(|scalar| scalar != &Scalar::ZERO) - .ok_or(InternalError::PointError) + .ok_or(ProtocolError::SerializationError) } } @@ -183,3 +180,9 @@ impl Group for Ristretto255 { ::deserialize_scalar(scalar_bits) } } + +impl DiffieHellman for Scalar { + fn diffie_hellman(self, pk: RistrettoPoint) -> GenericArray { + Ristretto255::serialize_pk(pk * self) + } +} diff --git a/src/key_exchange/traits.rs b/src/key_exchange/traits.rs index 9d7ffc3..f886756 100644 --- a/src/key_exchange/traits.rs +++ b/src/key_exchange/traits.rs @@ -17,7 +17,7 @@ use crate::ciphersuite::{CipherSuite, OprfHash}; use crate::errors::ProtocolError; use crate::hash::{Hash, ProxyHash}; use crate::key_exchange::group::KeGroup; -use crate::keypair::{PrivateKey, PublicKey, SecretKey}; +use crate::keypair::{PrivateKey, PublicKey}; pub trait KeyExchange where @@ -28,6 +28,9 @@ where type KE1State: Deserialize + Serialize + ZeroizeOnDrop + Clone; type KE2State: Deserialize + Serialize + ZeroizeOnDrop + Clone; type KE1Message: Deserialize + Serialize + ZeroizeOnDrop + Clone; + type KE2Builder: ZeroizeOnDrop + Clone; + type KE2BuilderData<'a>; + type KE2BuilderInput; type KE2Message: Deserialize + Serialize + ZeroizeOnDrop + Clone; type KE3Message: Deserialize + Serialize + ZeroizeOnDrop + Clone; @@ -36,25 +39,28 @@ where ) -> Result<(Self::KE1State, Self::KE1Message), ProtocolError>; #[allow(clippy::too_many_arguments)] - fn generate_ke2< - 'a, - 'b, - 'c, - 'd, - OprfCs: voprf::CipherSuite, - R: RngCore + CryptoRng, - S: SecretKey, - >( + fn ke2_builder<'a, 'b, 'c, 'd, OprfCs: voprf::CipherSuite, R: RngCore + CryptoRng>( rng: &mut R, - l1_bytes: impl Iterator, - l2_bytes: impl Iterator, + serialized_credential_request: impl Iterator, + serialized_credential_response: impl Iterator, ke1_message: Self::KE1Message, client_s_pk: PublicKey, - server_s_sk: S, id_u: impl Iterator, id_s: impl Iterator, context: &[u8], - ) -> Result, ProtocolError>; + ) -> Result; + + fn ke2_builder_data(builder: &Self::KE2Builder) -> Self::KE2BuilderData<'_>; + + fn generate_ke2_input( + builder: &Self::KE2Builder, + server_s_sk: &PrivateKey, + ) -> Self::KE2BuilderInput; + + fn build_ke2( + builder: Self::KE2Builder, + input: Self::KE2BuilderInput, + ) -> Result, ProtocolError>; #[allow(clippy::too_many_arguments)] fn generate_ke3<'a, 'b, 'c, 'd>( diff --git a/src/key_exchange/tripledh.rs b/src/key_exchange/tripledh.rs index 954a5e1..26fdae2 100644 --- a/src/key_exchange/tripledh.rs +++ b/src/key_exchange/tripledh.rs @@ -19,6 +19,7 @@ use generic_array::{ArrayLength, GenericArray}; use hkdf::{Hkdf, HkdfExtract}; use hmac::{Hmac, Mac}; use rand::{CryptoRng, RngCore}; +use zeroize::{Zeroize, ZeroizeOnDrop}; use crate::errors::utils::{check_slice_size, check_slice_size_atleast}; use crate::errors::{InternalError, ProtocolError}; @@ -27,7 +28,7 @@ use crate::key_exchange::group::KeGroup; use crate::key_exchange::traits::{ Deserialize, GenerateKe2Result, GenerateKe3Result, KeyExchange, Serialize, }; -use crate::keypair::{KeyPair, PrivateKey, PublicKey, SecretKey}; +use crate::keypair::{KeyPair, PrivateKey, PublicKey}; use crate::serialization::{Input, UpdateExt}; /////////////// @@ -49,6 +50,14 @@ static STR_OPAQUE: &[u8] = b"OPAQUE-"; //////////////////////////// /// The Triple Diffie-Hellman key exchange implementation +/// +/// # Remote Key +/// +/// [`ServerLoginBuilder::data()`](crate::ServerLoginBuilder::data()) will +/// return the client's ephemeral public key. +/// [`ServerLoginBuilder::build()`](crate::ServerLoginBuilder::build()) expects +/// a shared secret computed through Diffie-Hellman from the server's private +/// key and the given public key. pub struct TripleDh; /// The client state produced after the first key exchange message @@ -95,6 +104,31 @@ where session_key: Output, } +/// Builder for the second key exchange message +#[cfg_attr( + feature = "serde", + derive(serde::Deserialize, serde::Serialize), + serde(bound( + deserialize = "D: serde::Deserialize<'de>, PublicKey: serde::Deserialize<'de>", + serialize = "D: serde::Serialize, PublicKey: serde::Serialize", + )) +)] +#[derive_where(Clone)] +#[derive_where(Debug, Eq, Hash, PartialEq; D, PublicKey)] +pub struct Ke2Builder +where + D::Core: ProxyHash, + ::BlockSize: IsLess, + Le<::BlockSize, U256>: NonZero, +{ + server_nonce: GenericArray, + transcript_hasher: D, + client_e_pk: PublicKey, + server_e_pk: PublicKey, + shared_secret_1: GenericArray, + shared_secret_3: GenericArray, +} + /// The second key exchange message #[cfg_attr( feature = "serde", @@ -130,13 +164,20 @@ where mac: Output, } +/// Trait required by [`KeGroup::Sk`] to be compatible with [`TripleDh`]. +pub trait DiffieHellman { + /// Diffie-Hellman key exchange. + fn diffie_hellman(self, pk: KG::Pk) -> GenericArray; +} + //////////////////////////////// // High-level Implementations // // ========================== // //////////////////////////////// -impl KeyExchange for TripleDh +impl KeyExchange for TripleDh where + KG::Sk: DiffieHellman, D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, @@ -158,6 +199,9 @@ where type KE1State = Ke1State; type KE2State = Ke2State; type KE1Message = Ke1Message; + type KE2Builder = Ke2Builder; + type KE2BuilderData<'a> = &'a PublicKey; + type KE2BuilderInput = GenericArray; type KE2Message = Ke2Message; type KE3Message = Ke3Message; @@ -181,71 +225,82 @@ where )) } - #[allow(clippy::type_complexity)] - fn generate_ke2< - 'a, - 'b, - 'c, - 'd, - OprfCs: voprf::CipherSuite, - R: RngCore + CryptoRng, - S: SecretKey, - >( + fn ke2_builder<'a, 'b, 'c, 'd, OprfCs: voprf::CipherSuite, R: RngCore + CryptoRng>( rng: &mut R, serialized_credential_request: impl Iterator, - l2_bytes: impl Iterator, + serialized_credential_response: impl Iterator, ke1_message: Self::KE1Message, client_s_pk: PublicKey, - server_s_sk: S, id_u: impl Iterator, id_s: impl Iterator, context: &[u8], - ) -> Result, ProtocolError> { - let server_e_kp = KeyPair::::generate_random::(rng); + ) -> Result { + let server_e = KeyPair::::generate_random::(rng); let server_nonce = generate_nonce::(rng); - let mut transcript_hasher = D::new() + let transcript_hasher = D::new() .chain(STR_CONTEXT) - .chain_iter( - Input::::from(context) - .map_err(ProtocolError::into_custom)? - .iter(), - ) + .chain_iter(Input::::from(context)?.iter()) .chain_iter(id_u.into_iter()) .chain_iter(serialized_credential_request) .chain_iter(id_s.into_iter()) - .chain_iter(l2_bytes) + .chain_iter(serialized_credential_response) .chain(server_nonce) - .chain(server_e_kp.public().serialize()); + .chain(server_e.public().serialize()); - let result = derive_3dh_keys::( - TripleDhComponents { - pk1: ke1_message.client_e_pk.clone(), - sk1: server_e_kp.private().clone(), - pk2: ke1_message.client_e_pk.clone(), - sk2: server_s_sk, - pk3: client_s_pk, - sk3: server_e_kp.private().clone(), - }, - &transcript_hasher.clone().finalize(), + let shared_secret_1 = server_e + .private() + .ke_diffie_hellman(&ke1_message.client_e_pk); + let shared_secret_3 = server_e.private().ke_diffie_hellman(&client_s_pk); + + Ok(Ke2Builder { + server_nonce, + transcript_hasher, + client_e_pk: ke1_message.client_e_pk.clone(), + server_e_pk: server_e.public().clone(), + shared_secret_1, + shared_secret_3, + }) + } + + fn ke2_builder_data(builder: &Self::KE2Builder) -> Self::KE2BuilderData<'_> { + &builder.client_e_pk + } + + fn generate_ke2_input( + builder: &Self::KE2Builder, + server_s_sk: &PrivateKey, + ) -> Self::KE2BuilderInput { + server_s_sk.ke_diffie_hellman(&builder.client_e_pk) + } + + fn build_ke2( + mut builder: Self::KE2Builder, + shared_secret_2: Self::KE2BuilderInput, + ) -> Result, ProtocolError> { + let result = derive_3dh_keys::( + builder.shared_secret_1.clone(), + shared_secret_2, + builder.shared_secret_3.clone(), + &builder.transcript_hasher.clone().finalize(), )?; let mut mac_hasher = Hmac::::new_from_slice(&result.1).map_err(|_| InternalError::HmacError)?; - mac_hasher.update(&transcript_hasher.clone().finalize()); + mac_hasher.update(&builder.transcript_hasher.clone().finalize()); let mac = mac_hasher.finalize().into_bytes(); - Digest::update(&mut transcript_hasher, &mac); + Digest::update(&mut builder.transcript_hasher, &mac); Ok(( Ke2State { km3: result.2, - hashed_transcript: transcript_hasher.finalize(), + hashed_transcript: builder.transcript_hasher.clone().finalize(), session_key: result.0, }, Ke2Message { - server_nonce, - server_e_pk: server_e_kp.public().clone(), + server_nonce: builder.server_nonce, + server_e_pk: builder.server_e_pk.clone(), mac, }, #[cfg(test)] @@ -276,15 +331,12 @@ where .chain_iter(l2_component) .chain(ke2_message.to_bytes_without_mac()); - let result = derive_3dh_keys::>( - TripleDhComponents { - pk1: ke2_message.server_e_pk.clone(), - sk1: ke1_state.client_e_sk.clone(), - pk2: server_s_pk, - sk2: ke1_state.client_e_sk.clone(), - pk3: ke2_message.server_e_pk.clone(), - sk3: client_s_sk, - }, + let result = derive_3dh_keys::( + ke1_state + .client_e_sk + .ke_diffie_hellman(&ke2_message.server_e_pk), + ke1_state.client_e_sk.ke_diffie_hellman(&server_s_pk), + client_s_sk.ke_diffie_hellman(&ke2_message.server_e_pk), &transcript_hasher.clone().finalize(), )?; @@ -335,16 +387,6 @@ where //==================== // ///////////////////////// -// The triple of public and private components used in the 3DH computation -struct TripleDhComponents> { - pk1: PublicKey, - sk1: PrivateKey, - pk2: PublicKey, - sk2: S, - pk3: PublicKey, - sk3: PrivateKey, -} - // Consists of a session key, followed by two mac keys: (session_key, km2, km3) #[cfg(not(test))] type TripleDhDerivationResult = (Output, Output, Output); @@ -361,10 +403,12 @@ type TripleDhDerivationResult = (Output, Output, Output, Output); // Internal function which takes the public and private components of the client // and server keypairs, along with some auxiliary metadata, to produce the // session key and two MAC keys -fn derive_3dh_keys>( - dh: TripleDhComponents, +fn derive_3dh_keys( + shared_secret_1: GenericArray, + shared_secret_2: GenericArray, + shared_secret_3: GenericArray, hashed_derivation_transcript: &[u8], -) -> Result, ProtocolError> +) -> Result, ProtocolError> where D::Core: ProxyHash, ::BlockSize: IsLess, @@ -372,36 +416,24 @@ where { let mut hkdf = HkdfExtract::::new(None); - hkdf.input_ikm( - &dh.sk1 - .diffie_hellman(dh.pk1) - .map_err(InternalError::into_custom)?, - ); - hkdf.input_ikm(&dh.sk2.diffie_hellman(dh.pk2)?); - hkdf.input_ikm( - &dh.sk3 - .diffie_hellman(dh.pk3) - .map_err(InternalError::into_custom)?, - ); + hkdf.input_ikm(&shared_secret_1); + hkdf.input_ikm(&shared_secret_2); + hkdf.input_ikm(&shared_secret_3); let (_, extracted_ikm) = hkdf.finalize(); let handshake_secret = derive_secrets::( &extracted_ikm, STR_HANDSHAKE_SECRET, hashed_derivation_transcript, - ) - .map_err(ProtocolError::into_custom)?; + )?; let session_key = derive_secrets::( &extracted_ikm, STR_SESSION_KEY, hashed_derivation_transcript, - ) - .map_err(ProtocolError::into_custom)?; + )?; - let km2 = hkdf_expand_label::(&handshake_secret, STR_SERVER_MAC, b"") - .map_err(ProtocolError::into_custom)?; - let km3 = hkdf_expand_label::(&handshake_secret, STR_CLIENT_MAC, b"") - .map_err(ProtocolError::into_custom)?; + let km2 = hkdf_expand_label::(&handshake_secret, STR_SERVER_MAC, b"")?; + let km3 = hkdf_expand_label::(&handshake_secret, STR_CLIENT_MAC, b"")?; Ok(( session_key, @@ -579,6 +611,32 @@ where } } +impl Drop for Ke2Builder +where + D::Core: ProxyHash, + ::BlockSize: IsLess, + Le<::BlockSize, U256>: NonZero, +{ + fn drop(&mut self) { + struct AssertZeroizeOnDrop<'a, T: ZeroizeOnDrop>(#[allow(unused)] &'a T); + + self.server_nonce.zeroize(); + self.transcript_hasher.reset(); + let _ = AssertZeroizeOnDrop(&self.client_e_pk); + let _ = AssertZeroizeOnDrop(&self.server_e_pk); + self.shared_secret_1.zeroize(); + self.shared_secret_3.zeroize(); + } +} + +impl ZeroizeOnDrop for Ke2Builder +where + D::Core: ProxyHash, + ::BlockSize: IsLess, + Le<::BlockSize, U256>: NonZero, +{ +} + impl Deserialize for Ke2Message where D::Core: ProxyHash, diff --git a/src/keypair.rs b/src/keypair.rs index a23ea34..a9d17c2 100644 --- a/src/keypair.rs +++ b/src/keypair.rs @@ -14,8 +14,9 @@ use derive_where::derive_where; use generic_array::{ArrayLength, GenericArray}; use rand::{CryptoRng, RngCore}; -use crate::errors::{InternalError, ProtocolError}; +use crate::errors::ProtocolError; use crate::key_exchange::group::KeGroup; +use crate::key_exchange::tripledh::DiffieHellman; /// A Keypair trait with public-private verification #[cfg_attr( @@ -28,12 +29,17 @@ use crate::key_exchange::group::KeGroup; )] #[derive_where(Clone)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; KG::Pk, S)] -pub struct KeyPair = PrivateKey> { +pub struct KeyPair> { pk: PublicKey, sk: S, } -impl> KeyPair { +impl KeyPair { + /// Creates a new [`KeyPair`] from the given keys. + pub fn new(sk: S, pk: PublicKey) -> Self { + Self { pk, sk } + } + /// The public key component pub fn public(&self) -> &PublicKey { &self.pk @@ -43,20 +49,6 @@ impl> KeyPair { pub fn private(&self) -> &S { &self.sk } - - /// Obtains a [`KeyPair`] from a slice representing the private key - pub fn from_private_key_slice(input: &[u8]) -> Result> { - Self::from_private_key(S::deserialize(input)?) - } - - /// Obtains a [`KeyPair`] from a private key - pub fn from_private_key(private_key: S) -> Result> { - let pk = private_key.public_key()?; - Ok(Self { - pk, - sk: private_key, - }) - } } impl KeyPair { @@ -106,50 +98,60 @@ where #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; KG::Sk)] pub struct PrivateKey(KG::Sk); -/// A trait specifying the requirements for a private key container -pub trait SecretKey: Clone + Sized { - /// Custom error type that can be passed down to `InternalError::Custom` +impl PrivateKey { + /// Returns public key from private key + pub fn public_key(&self) -> PublicKey { + PublicKey(KG::public_key(self.0)) + } + + pub(crate) fn serialize(&self) -> GenericArray { + KG::serialize_sk(self.0) + } + + /// Creates a [`PrivateKey`] from the given bytes. + pub fn deserialize(input: &[u8]) -> Result { + KG::deserialize_sk(input).map(Self) + } +} + +impl PrivateKey +where + KG::Sk: DiffieHellman, +{ + /// Diffie-Hellman key exchange implementation + pub(crate) fn ke_diffie_hellman(&self, pk: &PublicKey) -> GenericArray { + self.0.diffie_hellman(pk.0) + } +} + +/// A trait to facilitate +/// [`ServerSetup::de/serialize`](crate::ServerSetup::serialize). +pub trait PrivateKeySerialization: Clone { + /// Custom error type that can be passed down to `ProtocolError::Custom` type Error; /// Serialization size in bytes. type Len: ArrayLength; - /// Diffie-Hellman key exchange implementation - fn diffie_hellman( - &self, - pk: PublicKey, - ) -> Result, InternalError>; - - /// Returns public key from private key - fn public_key(&self) -> Result, InternalError>; - /// Serialization into bytes - fn serialize(&self) -> GenericArray; + fn serialize_key_pair(key_pair: &KeyPair) -> GenericArray; /// Deserialization from bytes - fn deserialize(input: &[u8]) -> Result>; + fn deserialize_key_pair(input: &[u8]) -> Result, ProtocolError>; } -impl SecretKey for PrivateKey { +impl PrivateKeySerialization for PrivateKey { type Error = core::convert::Infallible; type Len = KG::SkLen; - fn diffie_hellman( - &self, - pk: PublicKey, - ) -> Result, InternalError> { - Ok(KG::diffie_hellman(pk.0, self.0)) + fn serialize_key_pair(key_pair: &KeyPair) -> GenericArray { + key_pair.private().serialize() } - fn public_key(&self) -> Result, InternalError> { - Ok(PublicKey(KG::public_key(self.0))) - } + fn deserialize_key_pair(input: &[u8]) -> Result, ProtocolError> { + let sk = PrivateKey::deserialize(input)?; + let pk = sk.public_key(); - fn serialize(&self) -> GenericArray { - KG::serialize_sk(self.0) - } - - fn deserialize(input: &[u8]) -> Result { - KG::deserialize_sk(input).map(Self) + Ok(KeyPair::new(sk, pk)) } } @@ -184,7 +186,7 @@ pub struct PublicKey(KG::Pk); impl PublicKey { /// Convert from bytes - pub fn deserialize(key_bytes: &[u8]) -> Result { + pub fn deserialize(key_bytes: &[u8]) -> Result { KG::deserialize_pk(key_bytes).map(Self) } @@ -192,6 +194,11 @@ impl PublicKey { pub fn serialize(&self) -> GenericArray { KG::serialize_pk(self.0) } + + /// Returns the inner [`KeGroup::Pk`]. + pub fn to_group_type(&self) -> KG::Pk { + self.0 + } } #[cfg(feature = "serde")] @@ -223,7 +230,6 @@ mod tests { use rand::rngs::OsRng; use super::*; - use crate::errors::*; use crate::util; #[test] @@ -255,15 +261,15 @@ mod tests { fn pub_from_priv(kp in KeyPair::<$point>::uniform_keypair_strategy::<$point>()) { let pk = kp.public(); let sk = kp.private(); - prop_assert_eq!(&sk.public_key()?, pk); + prop_assert_eq!(&sk.public_key(), pk); } #[test] fn dh(kp1 in KeyPair::<$point>::uniform_keypair_strategy::<$point>(), kp2 in KeyPair::<$point>::uniform_keypair_strategy::<$point>()) { - let dh1 = kp2.private().diffie_hellman(kp1.public().clone())?; - let dh2 = kp1.private().diffie_hellman(kp2.public().clone())?; + let dh1 = kp2.private().ke_diffie_hellman(&kp1.public()); + let dh2 = kp1.private().ke_diffie_hellman(kp2.public()); prop_assert_eq!(dh1, dh2); } @@ -272,7 +278,7 @@ mod tests { fn private_key_slice(kp in KeyPair::<$point>::uniform_keypair_strategy::<$point>()) { let sk_bytes = kp.private().serialize().to_vec(); - let kp2 = KeyPair::<$point>::from_private_key_slice(&sk_bytes)?; + let kp2 = PrivateKey::<$point>::deserialize_key_pair(&sk_bytes)?; let kp2_private_bytes = kp2.private().serialize().to_vec(); prop_assert_eq!(sk_bytes, kp2_private_bytes); @@ -320,38 +326,15 @@ mod tests { #[derive(Clone)] struct RemoteKey(PrivateKey); - impl SecretKey for RemoteKey { - type Error = core::convert::Infallible; - type Len = ::SkLen; - - fn diffie_hellman( - &self, - pk: PublicKey, - ) -> Result::PkLen>, InternalError> - { - self.0.diffie_hellman(pk) - } - - fn public_key(&self) -> Result, InternalError> { - self.0.public_key() - } - - fn serialize(&self) -> GenericArray { - self.0.serialize() - } - - fn deserialize(input: &[u8]) -> Result> { - PrivateKey::deserialize(input).map(Self) - } - } - const PASSWORD: &str = "password"; - let sk = KeCurve::random_sk(&mut OsRng); - let sk = RemoteKey(PrivateKey(sk)); - let keypair = KeyPair::from_private_key(sk).unwrap(); + let sk = PrivateKey(KeCurve::random_sk(&mut OsRng)); + let pk = sk.public_key(); + let sk = RemoteKey(sk); + let keypair = KeyPair::new(sk, pk); - let server_setup = ServerSetup::::new_with_key(&mut OsRng, keypair); + let server_setup = + ServerSetup::::new_with_key_pair(&mut OsRng, keypair); let ClientRegistrationStartResult { message, @@ -373,11 +356,7 @@ mod tests { message, state: client, } = ClientLogin::::start(&mut OsRng, PASSWORD.as_bytes()).unwrap(); - let ServerLoginStartResult { - message, - state: server, - .. - } = ServerLogin::start( + let builder = ServerLogin::builder( &mut OsRng, &server_setup, Some(file), @@ -386,6 +365,12 @@ mod tests { ServerLoginStartParameters::default(), ) .unwrap(); + let shared_secret = builder.private_key().0.ke_diffie_hellman(builder.data()); + let ServerLoginStartResult { + message, + state: server, + .. + } = builder.build(shared_secret).unwrap(); let ClientLoginFinishResult { message, .. } = client .finish( PASSWORD.as_bytes(), diff --git a/src/lib.rs b/src/lib.rs index f1c225c..3543c99 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -936,14 +936,12 @@ //! ## Remote Private Keys //! //! Servers that want to store their private key in an external location (e.g. -//! in an HSM or vault) can do so with the [`SecretKey`](keypair::SecretKey`) -//! trait. This allows [`ServerSetup`] to be constructed using an existing -//! keypair without exposing the bytes of the private key to this library. +//! in an HSM or vault) can do so with [`ServerLogin::builder()`] without +//! exposing the bytes of the private key to this library. //! ``` //! # use generic_array::{GenericArray, typenum::U0}; -//! # use opaque_ke::{CipherSuite, errors::{InternalError}, key_exchange::group::KeGroup, keypair::{KeyPair, PrivateKey, PublicKey, SecretKey}, ServerSetup}; +//! # use opaque_ke::{CipherSuite, ClientLogin, ClientRegistration, ClientRegistrationFinishParameters, ServerRegistration, errors::ProtocolError, keypair::PrivateKey, key_exchange::{group::KeGroup as OKeGroup, tripledh::DiffieHellman}}; //! # use rand::rngs::OsRng; -//! # use zeroize::Zeroize; //! # struct Default; //! # #[cfg(feature = "ristretto255")] //! # impl CipherSuite for Default { @@ -959,45 +957,66 @@ //! # type KeyExchange = opaque_ke::key_exchange::tripledh::TripleDh; //! # type Ksf = opaque_ke::ksf::Identity; //! # } -//! # #[derive(Debug)] +//! # type KeGroup = ::KeGroup; +//! # #[derive(Debug, thiserror::Error)] +//! # #[error("test error")] //! # struct YourRemoteKeyError; //! # #[derive(Clone)] -//! # struct YourRemoteKey(<::KeGroup as KeGroup>::Sk); +//! # struct YourRemoteKey(::Sk); //! # impl YourRemoteKey { -//! # fn diffie_hellman(&self, pk: &[u8]) -> Result::KeGroup as KeGroup>::PkLen>, YourRemoteKeyError> { todo!() } -//! # fn public_key(&self) -> Result::KeGroup as KeGroup>::PkLen>, YourRemoteKeyError> { Ok(<::KeGroup>::serialize_pk(<::KeGroup>::public_key(self.0))) } +//! # fn diffie_hellman(&self, pk: &PublicKey) -> Result::PkLen>, YourRemoteKeyError> { +//! # Ok(<::Sk as DiffieHellman>::diffie_hellman(self.0, KeGroup::deserialize_pk(&pk.serialize()).unwrap())) +//! # } //! # } -//! impl SecretKey<::KeGroup> for YourRemoteKey { +//! use opaque_ke::{ServerLogin, ServerLoginStartParameters, ServerSetup}; +//! use opaque_ke::keypair::{KeyPair, PrivateKeySerialization, PublicKey}; +//! +//! // Implement if you intend to use `ServerSetup::de/serialize` instead of `serde`. +//! impl PrivateKeySerialization for YourRemoteKey { //! type Error = YourRemoteKeyError; //! type Len = U0; //! -//! fn diffie_hellman( -//! &self, -//! pk: PublicKey<::KeGroup>, -//! ) -> Result::KeGroup as KeGroup>::PkLen>, InternalError> { -//! YourRemoteKey::diffie_hellman(self, &pk.serialize()).map_err(InternalError::Custom) +//! fn serialize_key_pair(_: &KeyPair) -> GenericArray { +//! unimplemented!() //! } //! -//! fn public_key( -//! &self -//! ) -> Result::KeGroup>, InternalError> { -//! PublicKey::deserialize(&YourRemoteKey::public_key(self).map_err(InternalError::Custom)?).map_err(InternalError::into_custom) -//! } -//! -//! fn serialize(&self) -> GenericArray { -//! // if you use Serde and the "serde" crate feature, you won't need this -//! todo!() -//! } -//! -//! fn deserialize(input: &[u8]) -> Result> { -//! // if you use Serde and the "serde" crate feature, you won't need this -//! todo!() +//! fn deserialize_key_pair(input: &[u8]) -> Result, ProtocolError> { +//! unimplemented!() //! } //! } //! -//! # let remote_key = YourRemoteKey(<::KeGroup>::random_sk(&mut OsRng)); -//! let keypair = KeyPair::from_private_key(remote_key).unwrap(); -//! let server_setup = ServerSetup::::new_with_key(&mut OsRng, keypair); +//! # let sk = KeGroup::random_sk(&mut OsRng); +//! # let pk = KeGroup::public_key(sk); +//! # let pk = KeGroup::serialize_pk(pk); +//! # let public_key = PublicKey::deserialize(&pk).unwrap(); +//! # let remote_key = YourRemoteKey(sk); +//! # let mut server_rng = OsRng; +//! let keypair = KeyPair::new(remote_key, public_key); +//! let server_setup = ServerSetup::::new_with_key_pair(&mut server_rng, keypair); +//! # let client_registration_start_result = ClientRegistration::::start( +//! # &mut OsRng, +//! # b"password", +//! # )?; +//! # let server_registration_start_result = ServerRegistration::::start(&server_setup, client_registration_start_result.message, b"alice@example.com")?; +//! # let client_registration_finish_result = client_registration_start_result.state.finish(&mut OsRng, b"password", server_registration_start_result.message, ClientRegistrationFinishParameters::default())?; +//! # let password_file_bytes = ServerRegistration::::finish(client_registration_finish_result.message).serialize(); +//! # let client_login_start_result = ClientLogin::::start( +//! # &mut OsRng, +//! # b"password", +//! # )?; +//! let password_file = ServerRegistration::::deserialize(&password_file_bytes)?; +//! let server_login_builder = ServerLogin::builder( +//! &mut server_rng, +//! &server_setup, +//! Some(password_file), +//! client_login_start_result.message, +//! b"alice@example.com", +//! ServerLoginStartParameters::default(), +//! )?; +//! let client_e_public_key = server_login_builder.data(); +//! let shared_secret = server_login_builder.private_key().diffie_hellman(&client_e_public_key)?; +//! let server_login_start_result = server_login_builder.build(shared_secret)?; +//! # Ok::<(), anyhow::Error>(()) //! ``` //! //! ## Custom KSF and Parameters @@ -1168,6 +1187,7 @@ pub use crate::messages::{ CredentialFinalization, CredentialFinalizationLen, CredentialRequest, CredentialRequestLen, CredentialResponse, CredentialResponseLen, RegistrationRequest, RegistrationRequestLen, RegistrationResponse, RegistrationResponseLen, RegistrationUpload, RegistrationUploadLen, + ServerLoginBuilder, }; pub use crate::opaque::{ ClientLogin, ClientLoginFinishParameters, ClientLoginFinishResult, ClientLoginStartResult, diff --git a/src/messages.rs b/src/messages.rs index a44814e..38beab9 100644 --- a/src/messages.rs +++ b/src/messages.rs @@ -18,6 +18,7 @@ use generic_array::{ArrayLength, GenericArray}; use rand::{CryptoRng, RngCore}; use subtle::ConstantTimeEq; use voprf::Group; +use zeroize::Zeroizing; use crate::ciphersuite::{CipherSuite, OprfGroup, OprfHash}; use crate::envelope::{Envelope, EnvelopeLen}; @@ -29,8 +30,10 @@ use crate::key_exchange::traits::{ Deserialize, Ke1MessageLen, Ke2MessageLen, Ke3MessageLen, KeyExchange, Serialize, }; use crate::key_exchange::tripledh::NonceLen; -use crate::keypair::{PublicKey, SecretKey}; -use crate::opaque::{MaskedResponse, MaskedResponseLen, ServerSetup}; +use crate::keypair::PublicKey; +use crate::opaque::{ + MaskedResponse, MaskedResponseLen, ServerLogin, ServerLoginStartResult, ServerSetup, +}; //////////////////////////// // High-level API Structs // @@ -107,6 +110,63 @@ pub struct CredentialRequest { pub(crate) ke1_message: , CS::KeGroup>>::KE1Message, } +/// Builder for [`ServerLogin`](crate::ServerLogin) when using remote keys. +#[cfg_attr( + feature = "serde", + derive(serde::Deserialize, serde::Serialize), + serde(bound( + deserialize = "S: serde::Deserialize<'de>, , \ + CS::KeGroup>>::KE2Builder: serde::Deserialize<'de>", + serialize = "S: serde::Serialize, , \ + CS::KeGroup>>::KE2Builder: serde::Serialize" + )) +)] +#[derive_where(Clone)] +#[derive_where( + Debug, Eq, PartialEq; + S, + voprf::EvaluationElement, + , CS::KeGroup>>::KE2Builder, +)] +pub struct ServerLoginBuilder { + pub(crate) server_s_sk: S, + pub(crate) evaluation_element: voprf::EvaluationElement, + pub(crate) masking_nonce: Zeroizing>, + pub(crate) masked_response: MaskedResponse, + #[cfg(test)] + pub(crate) oprf_key: Zeroizing as Group>::ScalarLen>>, + pub(crate) ke2_builder: , CS::KeGroup>>::KE2Builder, +} + +impl ServerLoginBuilder { + /// The returned data here has to be processed and the result given as an + /// input to [`ServerLoginBuilder::build()`]. To understand what kind of + /// output is expected here and how to process it, refer to the + /// documentation of your chosen [`CipherSuite::KeyExchange`]. + pub fn data( + &self, + ) -> , CS::KeGroup>>::KE2BuilderData<'_> { + CS::KeyExchange::ke2_builder_data(&self.ke2_builder) + } + + /// The handle to the corresponding [`ServerSetup`]s private key. + pub fn private_key(&self) -> &S { + &self.server_s_sk + } + + /// Build [`ServerLogin`] after attaining the input for the key exchange. To + /// understand what kind of input is expected here, refer to the + /// documentation of your chosen [`CipherSuite::KeyExchange`]. + /// + /// See [`ServerLogin::start()`] for the regular path. + pub fn build( + self, + input: , CS::KeGroup>>::KE2BuilderInput, + ) -> Result, ProtocolError> { + ServerLogin::build(self, input) + } +} + /// The answer sent by the server to the user, upon reception of the login /// attempt #[cfg_attr( @@ -265,7 +325,7 @@ impl RegistrationUpload { } // Creates a dummy instance used for faking a [CredentialResponse] - pub(crate) fn dummy>( + pub(crate) fn dummy( rng: &mut R, server_setup: &ServerSetup, ) -> Self { diff --git a/src/opaque.rs b/src/opaque.rs index 77c3b04..5fa2f67 100644 --- a/src/opaque.rs +++ b/src/opaque.rs @@ -8,7 +8,7 @@ //! Provides the main OPAQUE API -use core::ops::Add; +use core::ops::{Add, Deref}; use derive_where::derive_where; use digest::Output; @@ -19,6 +19,7 @@ use hkdf::{Hkdf, HkdfExtract}; use rand::{CryptoRng, RngCore}; use subtle::ConstantTimeEq; use voprf::Group; +use zeroize::Zeroizing; use crate::ciphersuite::{CipherSuite, OprfGroup, OprfHash}; use crate::envelope::{Envelope, EnvelopeLen}; @@ -30,13 +31,13 @@ use crate::key_exchange::traits::{ Deserialize, Ke1MessageLen, Ke1StateLen, Ke2StateLen, KeyExchange, Serialize, }; use crate::key_exchange::tripledh::NonceLen; -use crate::keypair::{KeyPair, PrivateKey, PublicKey, SecretKey}; +use crate::keypair::{KeyPair, PrivateKey, PrivateKeySerialization, PublicKey}; use crate::ksf::Ksf; use crate::messages::{CredentialRequestLen, RegistrationUploadLen}; use crate::serialization::Input; use crate::{ CredentialFinalization, CredentialRequest, CredentialResponse, RegistrationRequest, - RegistrationResponse, RegistrationUpload, + RegistrationResponse, RegistrationUpload, ServerLoginBuilder, }; /////////////// @@ -64,12 +65,9 @@ const STR_OPAQUE_DERIVE_KEY_PAIR: &[u8; 20] = b"OPAQUE-DeriveKeyPair"; )) )] #[derive_where(Clone)] -#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; ::Pk, ::Sk, S)] -pub struct ServerSetup< - CS: CipherSuite, - S: SecretKey = PrivateKey<::KeGroup>, -> { - oprf_seed: Output>, +#[derive_where(Debug, Eq, PartialEq; ::Pk, ::Sk, S)] +pub struct ServerSetup::KeGroup>> { + oprf_seed: Zeroizing>>, keypair: KeyPair, pub(crate) fake_keypair: KeyPair, } @@ -159,21 +157,21 @@ impl ServerSetup> { /// Generate a new instance of server setup pub fn new(rng: &mut R) -> Self { let keypair = KeyPair::generate_random::(rng); - Self::new_with_key(rng, keypair) + Self::new_with_key_pair(rng, keypair) } } /// Length of [`ServerSetup`] in bytes for serialization. -pub type ServerSetupLen> = +pub type ServerSetupLen> = Sum>, S::Len>, ::SkLen>; -impl> ServerSetup { +impl ServerSetup { /// Create [`ServerSetup`] with the given keypair /// /// This function should not be used to restore a previously-existing /// instance of [`ServerSetup`]. Instead, use [`ServerSetup::serialize`] and /// [`ServerSetup::deserialize`] for this purpose. - pub fn new_with_key( + pub fn new_with_key_pair( rng: &mut R, keypair: KeyPair, ) -> Self { @@ -181,7 +179,7 @@ impl> ServerSetup { rng.fill_bytes(&mut oprf_seed); Self { - oprf_seed, + oprf_seed: Zeroizing::new(oprf_seed), keypair, fake_keypair: KeyPair::::generate_random::(rng), } @@ -190,6 +188,7 @@ impl> ServerSetup { /// Serialization into bytes pub fn serialize(&self) -> GenericArray> where + S: PrivateKeySerialization, // ServerSetup: Hash + KeSk + KeSk OutputSize>: Add, Sum>, S::Len>: @@ -197,21 +196,26 @@ impl> ServerSetup { ServerSetupLen: ArrayLength, { self.oprf_seed + .deref() .clone() - .concat(self.keypair.private().serialize()) + .concat(S::serialize_key_pair(&self.keypair)) .concat(self.fake_keypair.private().serialize()) } /// Deserialization from bytes - pub fn deserialize(input: &[u8]) -> Result> { + pub fn deserialize(input: &[u8]) -> Result> + where + S: PrivateKeySerialization, + { let seed_len = OutputSize::>::USIZE; let key_len = ::SkLen::USIZE; - let checked_slice = check_slice_size(input, seed_len + key_len + key_len, "server_setup")?; + let checked_slice = check_slice_size(input, seed_len + key_len + key_len, "server_setup") + .map_err(ProtocolError::into_custom)?; Ok(Self { - oprf_seed: GenericArray::clone_from_slice(&checked_slice[..seed_len]), - keypair: KeyPair::from_private_key_slice(&checked_slice[seed_len..seed_len + key_len])?, - fake_keypair: KeyPair::from_private_key_slice(&checked_slice[seed_len + key_len..]) + oprf_seed: Zeroizing::new(GenericArray::clone_from_slice(&checked_slice[..seed_len])), + keypair: S::deserialize_key_pair(&checked_slice[seed_len..seed_len + key_len])?, + fake_keypair: PrivateKey::deserialize_key_pair(&checked_slice[seed_len + key_len..]) .map_err(ProtocolError::into_custom)?, }) } @@ -368,7 +372,7 @@ impl ServerRegistration { /// From the client's "blinded" password, returns a response to be sent back /// to the client, as well as a [`ServerRegistration`] - pub fn start>( + pub fn start( server_setup: &ServerSetup, message: RegistrationRequest, credential_identifier: &[u8], @@ -395,7 +399,7 @@ impl ServerRegistration { } // Creates a dummy instance used for faking a [CredentialResponse] - pub(crate) fn dummy>( + pub(crate) fn dummy( rng: &mut R, server_setup: &ServerSetup, ) -> Self { @@ -594,9 +598,10 @@ impl ServerLogin { }) } - /// From the client's "blinded" password, returns a challenge to be sent - /// back to the client, as well as a [`ServerLogin`] - pub fn start>( + /// Create a [`ServerLoginBuilder`] to use with a remote private key. + /// + /// See [`ServerLogin::start()`] for the regular path. + pub fn builder( rng: &mut R, server_setup: &ServerSetup, password_file: Option>, @@ -606,7 +611,7 @@ impl ServerLogin { context, identifiers, }: ServerLoginStartParameters, - ) -> Result, ProtocolError> + ) -> Result, ProtocolError> where // MaskedResponse: (Nonce + Hash) + KePk NonceLen: Add>>, @@ -621,8 +626,7 @@ impl ServerLogin { let client_s_pk = record.0.client_s_pk.clone(); let context = context.unwrap_or(&[]); - let server_s_sk = server_setup.keypair.private(); - let server_s_pk = server_s_sk.public_key()?; + let server_s_pk = server_setup.keypair.public(); let mut masking_nonce = GenericArray::<_, NonceLen>::default(); rng.fill_bytes(&mut masking_nonce); @@ -630,17 +634,15 @@ impl ServerLogin { let masked_response = mask_response( &record.0.masking_key, masking_nonce.as_slice(), - &server_s_pk, + server_s_pk, &record.0.envelope, - ) - .map_err(ProtocolError::into_custom)?; + )?; let (id_u, id_s) = bytestrings_from_identifiers::( identifiers, client_s_pk.serialize(), server_s_pk.serialize(), - ) - .map_err(ProtocolError::into_custom)?; + )?; let blinded_element = OprfGroup::::serialize_elem(credential_request.blinded_element.value()); @@ -648,32 +650,46 @@ impl ServerLogin { let credential_request_bytes = CredentialRequest::::serialize_iter(&blinded_element, &ke1_message); - let oprf_key = oprf_key_from_seed::(&server_setup.oprf_seed, credential_identifier) - .map_err(ProtocolError::into_custom)?; - let server = voprf::OprfServer::new_with_key(&oprf_key) - .map_err(|e| ProtocolError::into_custom(e.into()))?; + let oprf_key = oprf_key_from_seed::(&server_setup.oprf_seed, credential_identifier)?; + let server = voprf::OprfServer::new_with_key(&oprf_key).map_err(ProtocolError::from)?; let evaluation_element = server.blind_evaluate(&credential_request.blinded_element); let beta = OprfGroup::::serialize_elem(evaluation_element.value()); let credential_response_component = CredentialResponse::::serialize_without_ke(&beta, &masking_nonce, &masked_response); - let result = CS::KeyExchange::generate_ke2::( + let ke2_builder = CS::KeyExchange::ke2_builder::( rng, credential_request_bytes, credential_response_component, credential_request.ke1_message.clone(), client_s_pk, - server_s_sk.clone(), id_u.iter(), id_s.iter(), context, )?; - let credential_response = CredentialResponse { + Ok(ServerLoginBuilder { + server_s_sk: server_setup.keypair().private().clone(), evaluation_element, - masking_nonce, + masking_nonce: Zeroizing::new(masking_nonce), masked_response, + #[cfg(test)] + oprf_key: Zeroizing::new(oprf_key), + ke2_builder, + }) + } + + pub(crate) fn build( + builder: ServerLoginBuilder, + input: , CS::KeGroup>>::KE2BuilderInput, + ) -> Result, ProtocolError> { + let result = CS::KeyExchange::build_ke2(builder.ke2_builder.clone(), input)?; + + let credential_response = CredentialResponse { + evaluation_element: builder.evaluation_element.clone(), + masking_nonce: *builder.masking_nonce.deref(), + masked_response: builder.masked_response.clone(), ke2_message: result.1, }; @@ -687,10 +703,43 @@ impl ServerLogin { #[cfg(test)] server_mac_key: result.3, #[cfg(test)] - oprf_key, + oprf_key: builder.oprf_key.deref().clone(), }) } + /// From the client's "blinded" password, returns a challenge to be sent + /// back to the client, as well as a [`ServerLogin`] + pub fn start( + rng: &mut R, + server_setup: &ServerSetup, + password_file: Option>, + credential_request: CredentialRequest, + credential_identifier: &[u8], + parameters: ServerLoginStartParameters, + ) -> Result, ProtocolError> + where + // MaskedResponse: (Nonce + Hash) + KePk + NonceLen: Add>>, + Sum>>: + ArrayLength + Add<::PkLen>, + MaskedResponseLen: ArrayLength, + { + let builder = Self::builder( + rng, + server_setup, + password_file, + credential_request, + credential_identifier, + parameters, + )?; + let input = CS::KeyExchange::generate_ke2_input( + &builder.ke2_builder, + server_setup.keypair.private(), + ); + + Self::build(builder, input) + } + /// From the client's second and final message, check the client's /// authentication and produce a message transport pub fn finish( @@ -942,7 +991,7 @@ fn oprf_key_from_seed( derive(serde::Deserialize, serde::Serialize), serde(bound = "") )] -#[derive_where(Clone)] +#[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, PartialEq)] pub(crate) struct MaskedResponse { pub(crate) nonce: GenericArray, diff --git a/src/serialization/tests.rs b/src/serialization/tests.rs index 33ac416..fc3e068 100644 --- a/src/serialization/tests.rs +++ b/src/serialization/tests.rs @@ -28,7 +28,7 @@ use crate::key_exchange::traits::{ Deserialize, Ke1MessageLen, Ke1StateLen, Ke2MessageLen, KeyExchange, Serialize, }; use crate::key_exchange::tripledh::{NonceLen, TripleDh}; -use crate::keypair::{KeyPair, SecretKey}; +use crate::keypair::KeyPair; use crate::messages::CredentialResponseWithoutKeLen; use crate::opaque::{ClientLoginLen, ClientRegistrationLen, MaskedResponseLen}; use crate::serialization::{i2osp, os2ip}; diff --git a/src/tests/full_test.rs b/src/tests/full_test.rs index d9a1f80..7493a4d 100644 --- a/src/tests/full_test.rs +++ b/src/tests/full_test.rs @@ -27,8 +27,7 @@ use crate::errors::*; use crate::hash::OutputSize; use crate::key_exchange::group::KeGroup; use crate::key_exchange::traits::{Ke1MessageLen, Ke1StateLen, Ke2MessageLen}; -use crate::key_exchange::tripledh::{NonceLen, TripleDh}; -use crate::keypair::SecretKey; +use crate::key_exchange::tripledh::{DiffieHellman, NonceLen, TripleDh}; use crate::ksf::Identity; use crate::messages::{ CredentialRequestLen, CredentialResponseLen, CredentialResponseWithoutKeLen, @@ -1497,6 +1496,7 @@ fn test_zeroize_client_login_start() -> Result<(), ProtocolError> { _test_vector: &str, ) -> Result<(), ProtocolError> where + ::Sk: DiffieHellman, // CredentialRequest: KgPk + Ke1Message as Group>::ElemLen: Add::PkLen>>, CredentialRequestLen: ArrayLength, @@ -1595,6 +1595,7 @@ fn test_zeroize_client_login_finish() -> Result<(), ProtocolError> { _test_vector: &str, ) -> Result<(), ProtocolError> where + ::Sk: DiffieHellman, // MaskedResponse: (Nonce + Hash) + KePk NonceLen: Add>>, Sum>>: diff --git a/src/tests/mod.rs b/src/tests/mod.rs index 45c96b6..2d25207 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -12,4 +12,6 @@ mod full_test_vectors; pub mod mock_rng; mod opaque_vectors; mod parser; +#[cfg(test_hsm)] +mod remote_key; mod test_opaque_vectors; diff --git a/src/tests/remote_key.rs b/src/tests/remote_key.rs new file mode 100644 index 0000000..f7dea39 --- /dev/null +++ b/src/tests/remote_key.rs @@ -0,0 +1,373 @@ +use std::env; +use std::ops::Add; +use std::sync::{LazyLock, Mutex}; +use std::vec::Vec; + +use cryptoki::context::{CInitializeArgs, Pkcs11}; +use cryptoki::mechanism::elliptic_curve::{EcKdf, Ecdh1DeriveParams}; +use cryptoki::mechanism::Mechanism; +use cryptoki::object::{Attribute, AttributeType, KeyType, ObjectClass, ObjectHandle}; +use cryptoki::session::{Session, UserType}; +use cryptoki::types::AuthPin; +use elliptic_curve::group::Curve; +use elliptic_curve::pkcs8::der::asn1::{OctetString, OctetStringRef}; +use elliptic_curve::pkcs8::der::{Decode, Encode}; +use elliptic_curve::pkcs8::{AssociatedOid, ObjectIdentifier}; +use elliptic_curve::point::{AffineCoordinates, DecompressPoint}; +use elliptic_curve::sec1::{ModulusSize, Tag, ToEncodedPoint}; +use elliptic_curve::{AffinePoint, CurveArithmetic, FieldBytesSize, Group, ProjectivePoint}; +use generic_array::typenum::Sum; +use generic_array::{ArrayLength, GenericArray}; +use p256::NistP256; +use p384::NistP384; +use p521::NistP521; +use rand::rngs::OsRng; +use subtle::{Choice, ConditionallySelectable, ConstantTimeEq}; + +use crate::ciphersuite::OprfHash; +use crate::envelope::NonceLen; +use crate::hash::OutputSize; +use crate::key_exchange::group::KeGroup; +use crate::key_exchange::tripledh::{DiffieHellman, TripleDh}; +use crate::keypair::{KeyPair, PublicKey}; +use crate::ksf::Identity; +use crate::opaque::MaskedResponseLen; +use crate::{ + CipherSuite, ClientLogin, ClientLoginFinishParameters, ClientLoginStartResult, + ClientRegistration, ClientRegistrationFinishParameters, ClientRegistrationStartResult, + ServerLogin, ServerLoginStartParameters, ServerLoginStartResult, ServerRegistration, + ServerSetup, +}; +#[cfg(all(feature = "curve25519", feature = "ristretto255"))] +use crate::{Curve25519, Ristretto255}; + +#[test] +fn p256() { + struct Suite; + + impl CipherSuite for Suite { + type OprfCs = NistP256; + type KeGroup = NistP256; + type KeyExchange = TripleDh; + type Ksf = Identity; + } + + test::(Mechanism::EccKeyPairGen, NistP256::OID); +} + +#[test] +fn p384() { + struct Suite; + + impl CipherSuite for Suite { + type OprfCs = NistP384; + type KeGroup = NistP384; + type KeyExchange = TripleDh; + type Ksf = Identity; + } + + test::(Mechanism::EccKeyPairGen, NistP384::OID); +} + +#[test] +fn p521() { + struct Suite; + + impl CipherSuite for Suite { + type OprfCs = NistP521; + type KeGroup = NistP521; + type KeyExchange = TripleDh; + type Ksf = Identity; + } + + test::(Mechanism::EccKeyPairGen, NistP521::OID); +} + +#[test] +#[cfg(all(feature = "curve25519", feature = "ristretto255"))] +fn curve25519() { + struct Suite; + + impl CipherSuite for Suite { + type OprfCs = Ristretto255; + type KeGroup = Curve25519; + type KeyExchange = TripleDh; + type Ksf = Identity; + } + + test::( + // This should be [`Mechanism::EccMontgomeryKeyPairGen`], but SoftHSM has an incorrect + // implementation. See https://github.com/softhsm/SoftHSMv2/issues/647. + Mechanism::EccEdwardsKeyPairGen, + ObjectIdentifier::new("1.3.101.110").unwrap(), + ); +} + +#[derive(Clone)] +struct RemoteKey(ObjectHandle); + +trait Pkcs11DiffieHellman { + fn pkcs11_diffie_hellman( + &self, + server_pk: &PublicKey, + client_pk: &PublicKey, + ) -> GenericArray; +} + +fn test>(mechanism: Mechanism, oid: ObjectIdentifier) +where + RemoteKey: Pkcs11DiffieHellman, + ::Sk: DiffieHellman, + // MaskedResponse: (Nonce + Hash) + KePk + NonceLen: Add>>, + Sum>>: ArrayLength + Add<::PkLen>, + MaskedResponseLen: ArrayLength, + // Ke1State: KeSk + Nonce + ::SkLen: Add, + Sum<::SkLen, NonceLen>: ArrayLength, + // Ke1Message: Nonce + KePk + NonceLen: Add<::PkLen>, + Sum::PkLen>: ArrayLength, + // Ke2State: (Hash + Hash) + Hash + OutputSize>: Add>>, + Sum>, OutputSize>>: + ArrayLength + Add>>, + Sum>, OutputSize>>, OutputSize>>: + ArrayLength, + // Ke2Message: (Nonce + KePk) + Hash + NonceLen: Add<::PkLen>, + Sum::PkLen>: ArrayLength + Add>>, + Sum::PkLen>, OutputSize>>: ArrayLength, +{ + let (remote_key, pk) = pkcs11_generate_key_pair(mechanism, oid); + + let keypair = KeyPair::new(RemoteKey(remote_key), pk); + let server_setup = ServerSetup::new_with_key_pair(&mut OsRng, keypair); + + const PASSWORD: &str = "password"; + + let ClientRegistrationStartResult { + message, + state: client, + } = ClientRegistration::::start(&mut OsRng, PASSWORD.as_bytes()).unwrap(); + let message = ServerRegistration::start(&server_setup, message, &[]) + .unwrap() + .message; + let message = client + .finish( + &mut OsRng, + PASSWORD.as_bytes(), + message, + ClientRegistrationFinishParameters::default(), + ) + .unwrap() + .message; + let file = ServerRegistration::finish(message); + + let ClientLoginStartResult { + message, + state: client, + } = ClientLogin::::start(&mut OsRng, PASSWORD.as_bytes()).unwrap(); + let builder = ServerLogin::builder( + &mut OsRng, + &server_setup, + Some(file), + message, + &[], + ServerLoginStartParameters::default(), + ) + .unwrap(); + let shared_secret = builder + .private_key() + .pkcs11_diffie_hellman(server_setup.keypair().public(), builder.data()); + + let ServerLoginStartResult { + message, + state: server, + .. + } = builder.clone().build(shared_secret).unwrap(); + + let message = client + .clone() + .finish( + PASSWORD.as_bytes(), + message, + ClientLoginFinishParameters::default(), + ) + .map(|result| result.message); + + message + .map(|message| server.finish(message).unwrap()) + .unwrap(); +} + +static SESSION: LazyLock> = LazyLock::new(|| { + let module = env::var("PKCS11_MODULE").expect("`PKCS11_MODULE` environment variable"); + let pkcs11 = Pkcs11::new(module).unwrap(); + pkcs11.initialize(CInitializeArgs::OsThreads).unwrap(); + + let slot = pkcs11.get_slots_with_token().unwrap()[0]; + + let so_pin = AuthPin::new("abcdef".into()); + pkcs11.init_token(slot, &so_pin, "Test Token").unwrap(); + + let user_pin = AuthPin::new("fedcba".into()); + + { + let session = pkcs11.open_rw_session(slot).unwrap(); + session.login(UserType::So, Some(&so_pin)).unwrap(); + session.init_pin(&user_pin).unwrap(); + } + + let session = pkcs11.open_rw_session(slot).unwrap(); + session.login(UserType::User, Some(&user_pin)).unwrap(); + + Mutex::new(session) +}); + +fn pkcs11_generate_key_pair( + mechanism: Mechanism, + oid: ObjectIdentifier, +) -> (ObjectHandle, PublicKey) { + let session = SESSION.lock().unwrap(); + let (pk, remote_key) = session + .generate_key_pair( + &mechanism, + &[ + Attribute::Token(false), + Attribute::EcParams(oid.to_der().unwrap()), + ], + &[Attribute::Token(false), Attribute::Derive(true)], + ) + .unwrap(); + + let Attribute::EcPoint(pk) = session + .get_attributes(pk, &[AttributeType::EcPoint]) + .unwrap() + .pop() + .unwrap() + else { + unreachable!() + }; + drop(session); + + let pk = OctetString::from_der(&pk).unwrap(); + let pk = PublicKey::deserialize(pk.as_bytes()).unwrap(); + + (remote_key, pk) +} + +impl Pkcs11DiffieHellman for RemoteKey { + fn pkcs11_diffie_hellman( + &self, + server_pk: &PublicKey, + client_pk: &PublicKey, + ) -> GenericArray::PkLen> { + ec_pkcs_11_derive_secret::(self.0, server_pk, client_pk) + } +} + +impl Pkcs11DiffieHellman for RemoteKey { + fn pkcs11_diffie_hellman( + &self, + server_pk: &PublicKey, + client_pk: &PublicKey, + ) -> GenericArray::PkLen> { + ec_pkcs_11_derive_secret::(self.0, server_pk, client_pk) + } +} + +impl Pkcs11DiffieHellman for RemoteKey { + fn pkcs11_diffie_hellman( + &self, + server_pk: &PublicKey, + client_pk: &PublicKey, + ) -> GenericArray::PkLen> { + ec_pkcs_11_derive_secret::(self.0, server_pk, client_pk) + } +} + +#[cfg(all(feature = "curve25519", feature = "ristretto255"))] +impl Pkcs11DiffieHellman for RemoteKey { + fn pkcs11_diffie_hellman( + &self, + _: &PublicKey, + pk: &PublicKey, + ) -> GenericArray::PkLen> { + let shared_secret = pkcs11_derive_secret(self.0, &pk.serialize()); + + GenericArray::clone_from_slice(&shared_secret) + } +} + +fn pkcs11_derive_secret(sk: ObjectHandle, pk: &[u8]) -> Vec { + let session = SESSION.lock().unwrap(); + let shared_secret = session + .derive_key( + &Mechanism::Ecdh1Derive(Ecdh1DeriveParams::new(EcKdf::null(), pk)), + sk, + &[ + Attribute::Token(false), + Attribute::KeyType(KeyType::GENERIC_SECRET), + Attribute::Class(ObjectClass::SECRET_KEY), + Attribute::Extractable(true), + ], + ) + .unwrap(); + + let Attribute::Value(shared_secret) = session + .get_attributes(shared_secret, &[AttributeType::Value]) + .unwrap() + .pop() + .unwrap() + else { + unreachable!() + }; + drop(session); + + shared_secret +} + +fn ec_pkcs_11_derive_secret( + server_sk: ObjectHandle, + server_pk: &PublicKey, + client_pk: &PublicKey, +) -> GenericArray::PkLen> +where + KG: KeGroup> + CurveArithmetic, + AffinePoint: DecompressPoint + ToEncodedPoint, + FieldBytesSize: ModulusSize, +{ + let client_pk_point = client_pk.to_group_type(); + let client_pk = client_pk.serialize(); + let client_pk = OctetStringRef::new(&client_pk).unwrap(); + let client_pk = client_pk.to_der().unwrap(); + + let shared_secret_bytes = pkcs11_derive_secret(server_sk, &client_pk); + let shared_secret_point = AffinePoint::::decompress( + &GenericArray::clone_from_slice(&shared_secret_bytes), + Choice::from(0), + ) + .unwrap(); + let mut shared_secret = GenericArray::default(); + shared_secret[1..].copy_from_slice(&shared_secret_bytes); + + let shifted_client_pk = client_pk_point + ProjectivePoint::::generator(); + let shifted_client_pk = shifted_client_pk.to_affine().to_encoded_point(true); + let shifted_client_pk = OctetStringRef::new(shifted_client_pk.as_bytes()).unwrap(); + let shifted_client_pk = shifted_client_pk.to_der().unwrap(); + + let check_point = pkcs11_derive_secret(server_sk, &shifted_client_pk); + + let shifted_server_pk = server_pk.to_group_type() + shared_secret_point; + let shifted_server_pk = shifted_server_pk.to_affine(); + + let tag = u8::conditional_select( + &(Tag::CompressedEvenY as u8), + &(Tag::CompressedOddY as u8), + check_point.ct_ne(&shifted_server_pk.x()), + ); + shared_secret[0] = tag; + + shared_secret +}