use crate::cms::enveloped_data::{KeyAgreeRecipientInfo, OriginatorIdentifierOrKey, RecipientInfo};
use crate::constants::TIGHTBEAM_KARI_KDF_INFO;
use crate::crypto::profiles::{CryptoProvider, DefaultCryptoProvider};
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
use crate::crypto::sign::elliptic_curve::{AffinePoint, FieldBytesSize, PublicKey, SecretKey};
use crate::transport::handshake::error::HandshakeError;
use crate::transport::handshake::kari::kari_unwrap;
pub struct TightBeamKariRecipient<P>
where
P: CryptoProvider,
{
recipient_priv: SecretKey<P::Curve>,
kdf_info: &'static [u8],
provider: P,
}
impl<P> TightBeamKariRecipient<P>
where
P: CryptoProvider,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
FieldBytesSize<P::Curve>: ModulusSize,
{
pub fn new(provider: P, recipient_priv: SecretKey<P::Curve>) -> Self {
Self::with_kdf_info(provider, recipient_priv, TIGHTBEAM_KARI_KDF_INFO)
}
pub fn with_kdf_info(provider: P, recipient_priv: SecretKey<P::Curve>, kdf_info: &'static [u8]) -> Self {
Self { recipient_priv, kdf_info, provider }
}
pub fn process_kari(
&self,
kari: &KeyAgreeRecipientInfo,
recipient_index: usize,
) -> Result<Vec<u8>, HandshakeError> {
if recipient_index >= kari.recipient_enc_keys.len() {
return Err(HandshakeError::InvalidRecipientIndex);
}
let originator_pub = self.extract_originator_public_key(kari)?;
let ukm = kari.ukm.as_ref().ok_or(HandshakeError::MissingUkm)?;
let wrapped_key = kari.recipient_enc_keys[recipient_index].enc_key.as_bytes();
kari_unwrap(
&self.provider,
&self.recipient_priv,
&originator_pub,
ukm.as_bytes(),
self.kdf_info,
wrapped_key,
)
}
fn extract_originator_public_key(
&self,
kari: &KeyAgreeRecipientInfo,
) -> Result<PublicKey<P::Curve>, HandshakeError> {
match &kari.originator {
OriginatorIdentifierOrKey::OriginatorKey(orig_key) => {
let pub_key_bytes = orig_key.public_key.raw_bytes();
Ok(PublicKey::<P::Curve>::from_sec1_bytes(pub_key_bytes)?)
}
_ => Err(HandshakeError::UnsupportedOriginatorIdentifier),
}
}
}
impl TightBeamKariRecipient<DefaultCryptoProvider> {
pub fn with_defaults(recipient_priv: k256::SecretKey) -> Self {
Self::new(DefaultCryptoProvider::default(), recipient_priv)
}
}
impl<P> super::enveloped_data::RecipientProcessor for TightBeamKariRecipient<P>
where
P: CryptoProvider,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
FieldBytesSize<P::Curve>: ModulusSize,
{
fn process_recipient(&self, info: &RecipientInfo, recipient_index: usize) -> Result<Vec<u8>, HandshakeError> {
match info {
RecipientInfo::Kari(kari) => self.process_kari(kari, recipient_index),
_ => Err(HandshakeError::UnsupportedOriginatorIdentifier),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
mod recipient {
use super::*;
use crate::cms::builder::RecipientInfoBuilder;
use crate::cms::enveloped_data::{KeyAgreeRecipientIdentifier, RecipientInfo, UserKeyingMaterial};
use crate::crypto::sign::ecdsa::k256::SecretKey as K256SecretKey;
use crate::oids::AES_256_WRAP;
use crate::random::{generate_nonce, OsRng};
use crate::spki::{AlgorithmIdentifierOwned, SubjectPublicKeyInfoOwned};
use crate::transport::handshake::builders::kari::TightBeamKariBuilder;
#[test]
fn test_full_roundtrip() -> Result<(), Box<dyn std::error::Error>> {
let sender_key = K256SecretKey::random(&mut OsRng);
let sender_pubkey = sender_key.public_key();
let sender_spki = SubjectPublicKeyInfoOwned::from_key(sender_pubkey)?;
let recipient_key = K256SecretKey::random(&mut OsRng);
let recipient_pubkey = recipient_key.public_key();
let client_nonce = [0x01u8; 32];
let server_nonce = [0x02u8; 32];
let mut ukm_bytes = Vec::new();
ukm_bytes.extend_from_slice(&client_nonce);
ukm_bytes.extend_from_slice(&server_nonce);
let ukm = UserKeyingMaterial::new(ukm_bytes)?;
let rid = KeyAgreeRecipientIdentifier::IssuerAndSerialNumber(cms::cert::IssuerAndSerialNumber {
issuer: x509_cert::name::Name::default(),
serial_number: x509_cert::serial_number::SerialNumber::new(&[0x01])?,
});
let key_enc_alg = AlgorithmIdentifierOwned { oid: AES_256_WRAP, parameters: None };
let original_cek = [0x42u8; 32];
let mut builder = TightBeamKariBuilder::default()
.with_sender_priv(sender_key.clone())
.with_sender_pub_spki(sender_spki)
.with_recipient_pub(recipient_pubkey)
.with_recipient_rid(rid)
.with_ukm(ukm)
.with_key_enc_alg(key_enc_alg);
let recipient_info = builder.build(&original_cek).map_err(|e| format!("build failed: {e:?}"))?;
let kari = match recipient_info {
RecipientInfo::Kari(k) => k,
_ => panic!("Expected Kari variant"),
};
let recipient = TightBeamKariRecipient::with_defaults(recipient_key);
let extracted_cek = recipient.process_kari(&kari, 0)?;
assert_eq!(extracted_cek, original_cek);
Ok(())
}
#[test]
fn test_wrong_key() -> Result<(), Box<dyn std::error::Error>> {
let sender_key = K256SecretKey::random(&mut OsRng);
let sender_pubkey = sender_key.public_key();
let sender_spki = SubjectPublicKeyInfoOwned::from_key(sender_pubkey)?;
let recipient_key = K256SecretKey::random(&mut OsRng);
let recipient_pubkey = recipient_key.public_key();
let wrong_recipient_key = K256SecretKey::random(&mut OsRng);
let ukm_bytes = generate_nonce::<64>(None)?;
let ukm = UserKeyingMaterial::new(ukm_bytes.to_vec())?;
let rid = KeyAgreeRecipientIdentifier::IssuerAndSerialNumber(cms::cert::IssuerAndSerialNumber {
issuer: x509_cert::name::Name::default(),
serial_number: x509_cert::serial_number::SerialNumber::new(&[0x01])?,
});
let key_enc_alg = AlgorithmIdentifierOwned { oid: AES_256_WRAP, parameters: None };
let original_cek = [0x42u8; 32];
let mut builder = TightBeamKariBuilder::default()
.with_sender_priv(sender_key)
.with_sender_pub_spki(sender_spki)
.with_recipient_pub(recipient_pubkey)
.with_recipient_rid(rid)
.with_ukm(ukm)
.with_key_enc_alg(key_enc_alg);
let recipient_info = builder.build(&original_cek).map_err(|e| format!("build failed: {e:?}"))?;
let kari = match recipient_info {
RecipientInfo::Kari(k) => k,
_ => panic!("Expected Kari variant"),
};
let wrong_recipient = TightBeamKariRecipient::with_defaults(wrong_recipient_key);
let result = wrong_recipient.process_kari(&kari, 0);
assert!(result.is_err());
Ok(())
}
}
}