use crate::crypto::profiles::CryptoProvider;
use crate::crypto::secret::Secret;
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
use crate::crypto::sign::elliptic_curve::subtle::ConstantTimeEq;
use crate::crypto::sign::elliptic_curve::{AffinePoint, Curve, CurveArithmetic, PublicKey, SecretKey};
use crate::transport::handshake::error::HandshakeError;
#[cfg(feature = "ecdh")]
use crate::crypto::sign::elliptic_curve::ecdh::diffie_hellman;
#[cfg(feature = "zeroize")]
use crate::zeroize::Zeroize;
fn derive_shared_secret<C>(priv_key: &SecretKey<C>, peer_pub: &PublicKey<C>) -> Result<Secret<Vec<u8>>, HandshakeError>
where
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
let shared = diffie_hellman(priv_key.to_nonzero_scalar(), peer_pub.as_affine());
let vec = shared.raw_secret_bytes().as_ref().to_vec();
Ok(vec.into())
}
fn derive_kek<P>(
provider: &P,
shared_secret: &Secret<Vec<u8>>,
ukm: &[u8],
kdf_info: &[u8],
) -> Result<[u8; 32], HandshakeError>
where
P: CryptoProvider,
{
if ukm.is_empty() {
return Err(HandshakeError::MissingUkm);
}
let kdf = provider.as_key_deriver::<HandshakeError, 32>();
shared_secret.with(|ss| kdf(ss.as_ref(), ukm, kdf_info))?
}
pub fn kari_wrap<P, C>(
provider: &P,
sender_priv: &SecretKey<C>,
recipient_pub: &PublicKey<C>,
ukm: &[u8],
kdf_info: &[u8],
cek: &[u8],
) -> Result<Vec<u8>, HandshakeError>
where
P: CryptoProvider,
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
let shared_secret = derive_shared_secret(sender_priv, recipient_pub)?;
let mut kek = derive_kek(provider, &shared_secret, ukm, kdf_info)?;
let wrapper = provider.as_key_wrapper_32::<HandshakeError>();
let wrapped = wrapper(cek, &kek)?;
#[cfg(feature = "zeroize")]
kek.zeroize();
Ok(wrapped)
}
pub fn kari_unwrap<P, C>(
provider: &P,
recipient_priv: &SecretKey<C>,
originator_pub: &PublicKey<C>,
ukm: &[u8],
kdf_info: &[u8],
wrapped: &[u8],
) -> Result<Vec<u8>, HandshakeError>
where
P: CryptoProvider,
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
let shared_secret = derive_shared_secret(recipient_priv, originator_pub)?;
let mut kek = derive_kek(provider, &shared_secret, ukm, kdf_info)?;
let unwrapper = provider.as_key_unwrapper_32::<HandshakeError>();
let cek = unwrapper(wrapped, &kek)?;
let wrapper = provider.as_key_wrapper_32::<HandshakeError>();
let rewrapped = wrapper(&cek, &kek)?;
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
let valid: bool = rewrapped.as_slice().ct_eq(wrapped).into();
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
#[cfg(feature = "zeroize")]
kek.zeroize();
if !valid {
return Err(HandshakeError::AesKeyWrap(
crate::crypto::aead::aes_kw::Error::IntegrityCheckFailed,
));
}
Ok(cek)
}
#[cfg(feature = "kem")]
pub fn kari_wrap_hybrid<P, C>(
provider: &P,
sender_ec_priv: &SecretKey<C>,
recipient_ec_pub: &PublicKey<C>,
kem_shared_secret: &[u8],
ukm: &[u8],
kdf_info: &[u8],
cek: &[u8],
) -> Result<Vec<u8>, HandshakeError>
where
P: CryptoProvider,
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
use crate::transport::handshake::primitives::multi_input_kdf;
let ecdh_secret = derive_shared_secret(sender_ec_priv, recipient_ec_pub)?;
let combined_key = ecdh_secret.with(|ecdh| multi_input_kdf::<P>(&[ecdh, kem_shared_secret], ukm, kdf_info))??;
let wrapper = provider.as_key_wrapper_32::<HandshakeError>();
let wrapped = wrapper(cek, &combined_key)?;
Ok(wrapped)
}
#[cfg(feature = "kem")]
pub fn kari_unwrap_hybrid<P, C>(
provider: &P,
recipient_ec_priv: &SecretKey<C>,
originator_ec_pub: &PublicKey<C>,
kem_shared_secret: &[u8],
ukm: &[u8],
kdf_info: &[u8],
wrapped: &[u8],
) -> Result<Vec<u8>, HandshakeError>
where
P: CryptoProvider,
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
use crate::transport::handshake::primitives::multi_input_kdf;
let ecdh_secret = derive_shared_secret(recipient_ec_priv, originator_ec_pub)?;
let mut combined_key =
ecdh_secret.with(|ecdh| multi_input_kdf::<P>(&[ecdh, kem_shared_secret], ukm, kdf_info))??;
let unwrapper = provider.as_key_unwrapper_32::<HandshakeError>();
let cek = unwrapper(wrapped, &combined_key)?;
let wrapper = provider.as_key_wrapper_32::<HandshakeError>();
let rewrapped = wrapper(&cek, &combined_key)?;
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
let valid: bool = rewrapped.as_slice().ct_eq(wrapped).into();
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
#[cfg(feature = "zeroize")]
combined_key.zeroize();
if !valid {
return Err(HandshakeError::HybridKariIntegrityCheckFailed);
}
Ok(cek)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::constants::TIGHTBEAM_KARI_KDF_INFO;
use crate::crypto::profiles::DefaultCryptoProvider;
use crate::crypto::sign::ecdsa::k256::SecretKey as K256SecretKey;
use crate::random::OsRng;
#[test]
fn wrap_unwrap_roundtrip() -> Result<(), Box<dyn std::error::Error>> {
let provider = DefaultCryptoProvider::default();
let sender = K256SecretKey::random(&mut OsRng);
let recipient = K256SecretKey::random(&mut OsRng);
let recipient_pub = recipient.public_key();
let sender_pub = sender.public_key();
let ukm = [0x55u8; 64];
let cek = [0x42u8; 32];
let wrapped = kari_wrap(&provider, &sender, &recipient_pub, &ukm, TIGHTBEAM_KARI_KDF_INFO, &cek)?;
assert!(wrapped.len() > cek.len());
let unwrapped = kari_unwrap(&provider, &recipient, &sender_pub, &ukm, TIGHTBEAM_KARI_KDF_INFO, &wrapped)?;
assert_eq!(unwrapped, cek);
Ok(())
}
#[test]
fn unwrap_fail_with_wrong_key() -> Result<(), Box<dyn std::error::Error>> {
let provider = DefaultCryptoProvider::default();
let sender = K256SecretKey::random(&mut OsRng);
let recipient = K256SecretKey::random(&mut OsRng);
let wrong_recipient = K256SecretKey::random(&mut OsRng);
let recipient_pub = recipient.public_key();
let sender_pub = sender.public_key();
let ukm = [0x33u8; 64];
let cek = [0xABu8; 32];
let wrapped = kari_wrap(
&provider,
&sender,
&recipient_pub,
&ukm,
crate::constants::TIGHTBEAM_KARI_KDF_INFO,
&cek,
)?;
let bad = kari_unwrap(
&provider,
&wrong_recipient,
&sender_pub,
&ukm,
crate::constants::TIGHTBEAM_KARI_KDF_INFO,
&wrapped,
);
assert!(bad.is_err());
Ok(())
}
}