#[cfg(not(feature = "std"))]
extern crate alloc;
#[cfg(not(feature = "std"))]
use alloc::{boxed::Box, vec::Vec};
use crate::crypto::secret::Secret;
use crate::spki::AlgorithmIdentifierOwned;
use crate::transport::handshake::error::HandshakeError;
#[cfg(feature = "transport-ecies")]
use crate::asn1::OctetString;
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
use crate::crypto::sign::elliptic_curve::{AffinePoint, Curve, CurveArithmetic, PublicKey};
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
use crate::x509::Certificate;
pub fn generate_cek() -> Result<Secret<[u8; 32]>, HandshakeError> {
use rand_core::RngCore;
let mut cek = [0u8; 32];
rand_core::OsRng
.try_fill_bytes(&mut cek)
.map_err(|_| HandshakeError::RandomGenerationFailed)?;
Ok(Secret::from(Box::new(cek)))
}
pub fn aes_gcm_encrypt(key: &[u8], plaintext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>, HandshakeError> {
use aes_gcm::aead::{Aead, Payload};
use aes_gcm::{Aes256Gcm, KeyInit, Nonce};
use rand_core::RngCore;
if key.len() != 32 {
return Err(HandshakeError::InvalidKeySize { expected: 32, received: key.len() });
}
let cipher = Aes256Gcm::new_from_slice(key)?;
let mut nonce_bytes = [0u8; 12];
rand_core::OsRng
.try_fill_bytes(&mut nonce_bytes)
.map_err(|_| HandshakeError::RandomGenerationFailed)?;
let nonce = Nonce::from_slice(&nonce_bytes);
let ciphertext = cipher.encrypt(nonce, Payload { msg: plaintext, aad: aad.unwrap_or(&[]) })?;
let mut result = Vec::with_capacity(12 + ciphertext.len());
result.extend_from_slice(&nonce_bytes);
result.extend_from_slice(&ciphertext);
Ok(result)
}
pub fn aes_gcm_decrypt(key: &[u8], ciphertext: &[u8], aad: Option<&[u8]>) -> Result<Vec<u8>, HandshakeError> {
use aes_gcm::aead::{Aead, Payload};
use aes_gcm::{Aes256Gcm, KeyInit, Nonce};
if key.len() != 32 {
return Err(HandshakeError::InvalidKeySize { expected: 32, received: key.len() });
}
if ciphertext.len() < 12 + 16 {
return Err(HandshakeError::CiphertextTooShort { minimum: 28, received: ciphertext.len() });
}
let cipher = Aes256Gcm::new_from_slice(key)?;
let nonce = Nonce::from_slice(&ciphertext[..12]);
let ct = &ciphertext[12..];
let plaintext = cipher.decrypt(nonce, Payload { msg: ct, aad: aad.unwrap_or(&[]) })?;
Ok(plaintext)
}
pub fn aes_256_gcm_algorithm() -> AlgorithmIdentifierOwned {
use crate::oids::AES_256_GCM;
AlgorithmIdentifierOwned { oid: AES_256_GCM, parameters: None }
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
#[inline]
pub fn validate_state<S: PartialEq>(current: S, expected: S) -> Result<(), HandshakeError> {
if current != expected {
Err(HandshakeError::InvalidState)
} else {
Ok(())
}
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
pub fn extract_verifying_key_from_cert<C>(cert: &Certificate) -> Result<PublicKey<C>, HandshakeError>
where
C: Curve + CurveArithmetic,
<C as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<C>: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
let pubkey_bytes = crate::crypto::x509::utils::extract_verifying_key_bytes(cert);
Ok(PublicKey::<C>::from_sec1_bytes(pubkey_bytes)?)
}
#[cfg(feature = "transport-ecies")]
pub fn octet_string_to_32_byte_array(octet_string: &OctetString) -> Result<[u8; 32], HandshakeError> {
let bytes = octet_string.as_bytes();
if bytes.len() != 32 {
return Err(HandshakeError::OctetStringLengthError((bytes.len(), 32).into()));
}
let mut out = [0u8; 32];
out.copy_from_slice(bytes);
Ok(out)
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
pub fn compute_transcript_digest<D>(data: &[u8]) -> Result<[u8; 32], HandshakeError>
where
D: crate::crypto::hash::Digest,
{
crate::transport::handshake::primitives::transcript::digest_output_to_array(&D::digest(data))
}
#[cfg(feature = "transport-ecies")]
pub fn compute_client_auth_digest<D>(
transcript_hash: &[u8; 32],
encrypted_data: &[u8],
client_cert_der: &[u8],
) -> Result<[u8; 32], HandshakeError>
where
D: crate::crypto::hash::Digest,
{
let mut data = Vec::with_capacity(32 + encrypted_data.len() + client_cert_der.len());
data.extend_from_slice(transcript_hash);
data.extend_from_slice(encrypted_data);
data.extend_from_slice(client_cert_der);
compute_transcript_digest::<D>(&data)
}
#[cfg(feature = "transport-ecies")]
pub fn clear_session_randoms(
base_session_key: &mut Option<[u8; 32]>,
client_random: &mut Option<[u8; 32]>,
server_random: &mut Option<[u8; 32]>,
) {
use crate::zeroize::Zeroize;
base_session_key.zeroize();
client_random.zeroize();
server_random.zeroize();
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_aes_gcm_roundtrip() -> Result<(), HandshakeError> {
let key = [0x42u8; 32];
let plaintext = b"Hello, World!";
let aad = Some(b"additional data".as_slice());
let ciphertext = aes_gcm_encrypt(&key, plaintext, aad)?;
let decrypted = aes_gcm_decrypt(&key, &ciphertext, aad)?;
assert_eq!(plaintext, decrypted.as_slice());
Ok(())
}
#[test]
fn test_aes_gcm_wrong_key() -> Result<(), HandshakeError> {
let key = [0x42u8; 32];
let wrong_key = [0x43u8; 32];
let plaintext = b"Secret message";
let ciphertext = aes_gcm_encrypt(&key, plaintext, None)?;
let result = aes_gcm_decrypt(&wrong_key, &ciphertext, None);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_aes_gcm_wrong_aad() -> Result<(), HandshakeError> {
let key = [0x42u8; 32];
let plaintext = b"Authenticated data";
let aad = Some(b"correct aad".as_slice());
let wrong_aad = Some(b"wrong aad".as_slice());
let ciphertext = aes_gcm_encrypt(&key, plaintext, aad)?;
let result = aes_gcm_decrypt(&key, &ciphertext, wrong_aad);
assert!(result.is_err());
Ok(())
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
#[test]
fn test_compute_transcript_digest_widths() -> Result<(), HandshakeError> {
use crate::crypto::hash::{Digest, Sha3_256, Sha3_512};
let digest = compute_transcript_digest::<Sha3_256>(b"transcript")?;
assert_eq!(digest.len(), 32);
let wide = compute_transcript_digest::<Sha3_512>(b"transcript")?;
assert_eq!(wide.as_slice(), &Sha3_512::digest(b"transcript")[..32]);
Ok(())
}
#[test]
fn test_cek_generation() -> Result<(), Box<dyn std::error::Error>> {
let cek1 = generate_cek()?;
let cek2 = generate_cek()?;
cek1.with(|bytes| assert_eq!(bytes.len(), 32))?;
cek2.with(|bytes| assert_eq!(bytes.len(), 32))?;
let bytes1 = cek1.with(|b| *b)?;
let bytes2 = cek2.with(|b| *b)?;
assert_ne!(bytes1, bytes2);
Ok(())
}
}