use cosmian_kms_crypto::reexport::cosmian_crypto_core::CsRng;
use openssl::{
cipher::{Cipher, CipherRef},
cipher_ctx::CipherCtx,
encrypt::Decrypter,
hash::MessageDigest,
pkey::{PKey, Private, Public},
rsa::{Padding, Rsa},
};
use sha2::digest::crypto_common::rand_core::{RngCore, SeedableRng};
use crate::error::{KmsCliError, result::KmsCliResult};
pub(crate) fn generate_rsa_keypair() -> KmsCliResult<(PKey<Private>, PKey<Public>)> {
let key_sizes = [2048, 3072, 4096];
let mut rng = CsRng::from_entropy();
let bits = key_sizes[(rng.next_u32() as usize) % key_sizes.len()];
let rsa = Rsa::generate(bits)
.map_err(|e| KmsCliError::Default(format!("Failed to generate RSA key: {e}")))?;
let private_key = PKey::from_rsa(rsa.clone())
.map_err(|e| KmsCliError::Default(format!("Failed to build private key: {e}")))?;
let public_key = PKey::from_rsa(
Rsa::from_public_components(
rsa.n()
.to_owned()
.map_err(|e| KmsCliError::Default(format!("Failed to clone modulus: {e}")))?,
rsa.e()
.to_owned()
.map_err(|e| KmsCliError::Default(format!("Failed to clone exponent: {e}")))?,
)
.map_err(|e| KmsCliError::Default(format!("Failed to build public RSA key: {e}")))?,
)
.map_err(|e| KmsCliError::Default(format!("Failed to build public key: {e}")))?;
Ok((private_key, public_key))
}
pub(crate) fn rsa_aes_key_wrap_sha1_unwrap(
ciphertext: &[u8],
private_key: &PKey<Private>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let rsa_key_size = private_key.size();
if ciphertext.len() <= rsa_key_size {
return Err("Ciphertext too short for RSA_AES_KEY_WRAP".into());
}
let (encrypted_aes_key, wrapped_key_material) = ciphertext.split_at(rsa_key_size);
let aes_key = rsaes_oaep_sha1_unwrap(encrypted_aes_key, private_key)?;
let unwrapped_key = aes_key_unwrap(wrapped_key_material, &aes_key)?;
Ok(unwrapped_key)
}
pub(crate) fn rsaes_oaep_sha256_unwrap(
ciphertext: &[u8],
private_key: &PKey<Private>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let mut decrypter = Decrypter::new(private_key)?;
decrypter.set_rsa_padding(Padding::PKCS1_OAEP)?;
decrypter.set_rsa_oaep_md(MessageDigest::sha256())?;
decrypter.set_rsa_mgf1_md(MessageDigest::sha256())?;
let buffer_len = decrypter.decrypt_len(ciphertext)?;
let mut decrypted = vec![0_u8; buffer_len];
let decrypted_len = decrypter.decrypt(ciphertext, &mut decrypted)?;
decrypted.truncate(decrypted_len);
Ok(decrypted)
}
pub(crate) fn rsaes_oaep_sha1_unwrap(
ciphertext: &[u8],
private_key: &PKey<Private>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let mut decrypter = Decrypter::new(private_key)?;
decrypter.set_rsa_padding(Padding::PKCS1_OAEP)?;
decrypter.set_rsa_oaep_md(MessageDigest::sha1())?;
decrypter.set_rsa_mgf1_md(MessageDigest::sha1())?;
let buffer_len = decrypter.decrypt_len(ciphertext)?;
let mut decrypted = vec![0_u8; buffer_len];
let decrypted_len = decrypter.decrypt(ciphertext, &mut decrypted)?;
decrypted.truncate(decrypted_len);
Ok(decrypted)
}
fn aes_key_unwrap(ciphertext: &[u8], kek: &[u8]) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
const AES_WRAP_BLOCK_SIZE: usize = 8;
if ciphertext.len() < 16 || !ciphertext.len().is_multiple_of(AES_WRAP_BLOCK_SIZE) {
return Err("Invalid ciphertext size for AES Key Unwrap".into());
}
let cipher: &CipherRef = match kek.len() {
16 => Cipher::aes_128_wrap_pad(),
24 => Cipher::aes_192_wrap_pad(),
32 => Cipher::aes_256_wrap_pad(),
_ => {
return Err(format!(
"Invalid KEK size: {} bytes. Expected 16, 24, or 32",
kek.len()
)
.into());
}
};
let mut ctx = CipherCtx::new()?;
ctx.decrypt_init(Some(cipher), Some(kek), None)?;
let mut plaintext = vec![0_u8; ciphertext.len() + 16];
let mut written = ctx.cipher_update(ciphertext, Some(&mut plaintext))?;
written += ctx.cipher_final(&mut plaintext[written..])?;
plaintext.truncate(written);
Ok(plaintext)
}
pub(crate) fn rsa_aes_key_wrap_sha256_unwrap(
ciphertext: &[u8],
private_key: &PKey<Private>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let rsa_key_size = private_key.size();
if ciphertext.len() <= rsa_key_size {
return Err("Ciphertext too short for RSA_AES_KEY_WRAP".into());
}
let (encrypted_aes_key, wrapped_key_material) = ciphertext.split_at(rsa_key_size);
let aes_key = rsaes_oaep_sha256_unwrap(encrypted_aes_key, private_key)?;
let unwrapped_key = aes_key_unwrap(wrapped_key_material, &aes_key)?;
Ok(unwrapped_key)
}