use super::kari::TightBeamKariBuilder;
use crate::cms::builder::RecipientInfoBuilder;
use crate::cms::content_info::{CmsVersion, ContentInfo};
use crate::cms::enveloped_data::{EncryptedContentInfo, EnvelopedData, RecipientInfo, RecipientInfos};
use crate::crypto::aead::{AeadCore, Encryptor, KeyInit};
use crate::crypto::common::typenum::Unsigned;
use crate::crypto::profiles::{CryptoProvider, DefaultCryptoProvider};
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
use crate::crypto::sign::elliptic_curve::{AffinePoint, FieldBytesSize};
use crate::crypto::x509::attr::{Attribute, Attributes};
use crate::der::asn1::{Any, SetOfVec};
use crate::oids::{DATA, ENVELOPED_DATA};
use crate::transport::handshake::attributes::HandshakeAttribute;
use crate::transport::handshake::error::HandshakeError;
use crate::transport::handshake::utils::generate_cek;
pub struct TightBeamEnvelopedDataBuilder<P>
where
P: CryptoProvider,
{
kari_builder: Option<TightBeamKariBuilder<P>>,
unprotected_attrs: Vec<HandshakeAttribute>,
}
impl<P> TightBeamEnvelopedDataBuilder<P>
where
P: CryptoProvider,
P::AeadCipher: KeyInit,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
FieldBytesSize<P::Curve>: ModulusSize,
{
pub fn new(kari_builder: TightBeamKariBuilder<P>) -> Self {
Self { kari_builder: Some(kari_builder), unprotected_attrs: Vec::new() }
}
pub fn with_unprotected_attr(mut self, attr: HandshakeAttribute) -> Self {
self.unprotected_attrs.push(attr);
self
}
pub fn with_unprotected_attrs(mut self, attrs: Vec<HandshakeAttribute>) -> Self {
self.unprotected_attrs.extend(attrs);
self
}
fn validate_builder_state(&self) -> Result<(), HandshakeError> {
if self.kari_builder.is_none() {
Err(HandshakeError::KariBuilderConsumed)
} else {
Ok(())
}
}
fn build_kari_with_cek(&mut self, cek: &[u8]) -> Result<cms::enveloped_data::RecipientInfo, HandshakeError> {
let mut kari_builder = self.kari_builder.take().ok_or(HandshakeError::KariBuilderConsumed)?;
Ok(kari_builder.build(cek)?)
}
fn build_unprotected_attributes(&mut self) -> Result<Option<Attributes>, HandshakeError> {
if self.unprotected_attrs.is_empty() {
return Ok(None);
}
self.unprotected_attrs.sort();
let attrs = core::mem::take(&mut self.unprotected_attrs);
let x509_attrs: Result<Vec<_>, der::Error> = attrs
.into_iter()
.map(|attr| Ok(Attribute { oid: attr.attr_type, values: SetOfVec::try_from(attr.attr_values)? }))
.collect();
Ok(Some(SetOfVec::try_from(x509_attrs?)?))
}
fn build_recipient_infos(&self, recipient_info: RecipientInfo) -> Result<RecipientInfos, HandshakeError> {
Ok(RecipientInfos::try_from(vec![recipient_info])?)
}
fn generate_nonce() -> Vec<u8> {
use rand_core::RngCore;
let mut nonce_bytes = vec![0u8; <P::AeadCipher as AeadCore>::NonceSize::USIZE];
rand_core::OsRng.fill_bytes(&mut nonce_bytes);
nonce_bytes
}
fn create_cipher_from_cek(cek_bytes: &[u8]) -> Result<P::AeadCipher, HandshakeError> {
P::AeadCipher::new_from_slice(cek_bytes)
.map_err(|_| HandshakeError::InvalidKeySize { expected: 32, received: cek_bytes.len() })
}
fn encrypt_content_with_cipher(
cipher: &P::AeadCipher,
plaintext: &[u8],
nonce: &[u8],
) -> Result<EncryptedContentInfo, HandshakeError>
where
P::AeadCipher: Encryptor<P::AeadOid>,
{
Ok(cipher.encrypt_content(plaintext, nonce, Some(DATA))?)
}
pub fn build(mut self, plaintext: &[u8], _aad: Option<&[u8]>) -> Result<EnvelopedData, HandshakeError> {
self.validate_builder_state()?;
let cek = generate_cek()?;
let recipient_info = cek.with(|cek_bytes| self.build_kari_with_cek(cek_bytes))??;
let encrypted_content = cek.with(|cek_bytes| {
let cipher = Self::create_cipher_from_cek(cek_bytes)?;
let nonce = Self::generate_nonce();
Self::encrypt_content_with_cipher(&cipher, plaintext, &nonce)
})??;
let unprotected_attrs = self.build_unprotected_attributes()?;
let recip_infos = self.build_recipient_infos(recipient_info)?;
Ok(EnvelopedData {
version: CmsVersion::V3,
originator_info: None,
recip_infos,
encrypted_content,
unprotected_attrs,
})
}
pub fn build_content_info(self, plaintext: &[u8], aad: Option<&[u8]>) -> Result<ContentInfo, HandshakeError> {
let enveloped_data = self.build(plaintext, aad)?;
let content = Any::encode_from(&enveloped_data)?;
Ok(ContentInfo { content_type: ENVELOPED_DATA, content })
}
}
impl TightBeamEnvelopedDataBuilder<DefaultCryptoProvider> {
pub fn with_defaults(kari_builder: TightBeamKariBuilder<DefaultCryptoProvider>) -> Self {
Self::new(kari_builder)
}
}
#[cfg(test)]
mod tests {
use super::*;
mod enveloped_data {
use super::*;
use crate::der::{Decode, Encode};
use crate::transport::handshake::attributes::{encode_client_nonce, encode_server_nonce};
use crate::transport::handshake::tests::{
create_test_key_enc_alg, create_test_keypair, create_test_recipient_id, create_test_ukm,
};
fn create_test_kari_builder() -> TightBeamKariBuilder<DefaultCryptoProvider> {
let (sender_key, sender_spki, _recipient_key, recipient_pubkey) = create_test_keypair();
let ukm = create_test_ukm();
let rid = create_test_recipient_id();
let key_enc_alg = create_test_key_enc_alg();
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)
}
#[test]
fn test_basic_enveloped_data() {
let kari_builder = create_test_kari_builder();
let plaintext = b"Hello, TightBeam!";
let builder = TightBeamEnvelopedDataBuilder::with_defaults(kari_builder);
let enveloped_data = builder.build(plaintext, None).unwrap();
assert_eq!(enveloped_data.version, CmsVersion::V3);
assert_eq!(enveloped_data.recip_infos.0.len(), 1);
assert!(enveloped_data.encrypted_content.encrypted_content.is_some());
assert_eq!(enveloped_data.encrypted_content.content_type, DATA);
}
#[test]
fn test_with_unprotected_attributes() {
let kari_builder = create_test_kari_builder();
let client_nonce = [0x11u8; 32];
let server_nonce = [0x22u8; 32];
let attr1 = encode_client_nonce(&client_nonce).unwrap();
let attr2 = encode_server_nonce(&server_nonce).unwrap();
let plaintext = b"Authenticated message";
let builder = TightBeamEnvelopedDataBuilder::with_defaults(kari_builder)
.with_unprotected_attr(attr1)
.with_unprotected_attr(attr2);
let enveloped_data = builder.build(plaintext, None).unwrap();
assert!(enveloped_data.unprotected_attrs.is_some());
let attrs = enveloped_data.unprotected_attrs.unwrap();
assert_eq!(attrs.len(), 2);
}
#[test]
fn test_der_encoding() {
let kari_builder = create_test_kari_builder();
let plaintext = b"DER encoding test";
let builder = TightBeamEnvelopedDataBuilder::with_defaults(kari_builder);
let built = builder.build(plaintext, None).unwrap();
let der_bytes = built.to_der().unwrap();
let decoded = EnvelopedData::from_der(&der_bytes).unwrap();
assert_eq!(decoded.version, CmsVersion::V3);
}
#[test]
fn test_content_info_wrapper() {
let kari_builder = create_test_kari_builder();
let plaintext = b"ContentInfo wrapper test";
let builder = TightBeamEnvelopedDataBuilder::with_defaults(kari_builder);
let content_info = builder.build_content_info(plaintext, None).unwrap();
assert_eq!(content_info.content_type, ENVELOPED_DATA);
let enveloped_data: EnvelopedData = content_info.content.decode_as().unwrap();
assert_eq!(enveloped_data.version, CmsVersion::V3);
}
}
}