use super::error::{Tr31CryptoError, Tr31Error};
use super::key_block_header::KeyBlockHeader;
use super::key_derivations::derive_keys_version_d;
use super::payload::{construct_payload, extract_key_from_payload};
use crate::SecretKey;
use zeroize::Zeroizing;
use paysec_crypto::{AesCbc, AesCmac, AesCmacKeyDerivation, AesKeySize};
const TR31_D_MAC_LEN: usize = 16;
const TR31_D_BLOCK_LEN: usize = 16;
const TR31_MAX_KEY_BLOCK_LENGTH: usize = 9999;
pub fn tr31_wrap<P, K: ?Sized>(
provider: &P,
kbpk: &K,
kbpk_size: AesKeySize,
mut header: KeyBlockHeader,
key: &[u8],
masked_key_len: usize,
random_seed: &[u8],
) -> Result<String, Tr31CryptoError<P::Error>>
where
P: AesCmacKeyDerivation<K>
+ AesCbc<<P as AesCmacKeyDerivation<K>>::DerivedKey>
+ AesCmac<<P as AesCmacKeyDerivation<K>>::DerivedKey>,
{
if header.version_id() != "D" {
return Err(Tr31Error::UnsupportedVersion(header.version_id().to_string()).into());
}
let (kbek, kbak) =
derive_keys_version_d(provider, kbpk, kbpk_size).map_err(Tr31CryptoError::Crypto)?;
let payload = Zeroizing::new(construct_payload(
key,
masked_key_len,
TR31_D_BLOCK_LEN,
random_seed,
)?);
let total_block_length = header.len() + (payload.len() * 2) + (TR31_D_MAC_LEN * 2);
if total_block_length % TR31_D_BLOCK_LEN != 0 {
return Err(Tr31Error::TotalBlockLengthNotMultiple {
block_length: TR31_D_BLOCK_LEN,
actual: total_block_length,
}
.into());
}
if total_block_length > TR31_MAX_KEY_BLOCK_LENGTH {
return Err(Tr31Error::KeyBlockLengthTooLarge {
maximum: TR31_MAX_KEY_BLOCK_LENGTH,
actual: total_block_length,
}
.into());
}
let encoded_block_length =
u16::try_from(total_block_length).map_err(|_| Tr31Error::KeyBlockLengthTooLarge {
maximum: TR31_MAX_KEY_BLOCK_LENGTH,
actual: total_block_length,
})?;
header.set_kb_length(encoded_block_length)?;
let header_str = header.export_str()?;
let mut mac_input = Zeroizing::new(header_str.as_bytes().to_vec());
mac_input.extend_from_slice(payload.as_slice());
let mac = provider
.calculate_cmac(&kbak, mac_input.as_slice())
.map_err(Tr31CryptoError::Crypto)?;
let iv = mac;
let encrypted_payload = provider
.encrypt_cbc(&kbek, &iv, payload.as_slice())
.map_err(Tr31CryptoError::Crypto)?;
let encrypted_payload_hex = hex::encode_upper(&encrypted_payload);
let mac_hex = hex::encode_upper(mac);
Ok(format!("{header_str}{encrypted_payload_hex}{mac_hex}"))
}
pub fn tr31_wrap_with_header_string<P, K: ?Sized>(
provider: &P,
kbpk: &K,
kbpk_size: AesKeySize,
header_str: &str,
key: &[u8],
masked_key_len: usize,
random_seed: &[u8],
) -> Result<String, Tr31CryptoError<P::Error>>
where
P: AesCmacKeyDerivation<K>
+ AesCbc<<P as AesCmacKeyDerivation<K>>::DerivedKey>
+ AesCmac<<P as AesCmacKeyDerivation<K>>::DerivedKey>,
{
let header = KeyBlockHeader::new_from_str(header_str)?;
tr31_wrap(
provider,
kbpk,
kbpk_size,
header,
key,
masked_key_len,
random_seed,
)
}
pub fn tr31_unwrap<P, K: ?Sized>(
provider: &P,
kbpk: &K,
kbpk_size: AesKeySize,
key_block: &str,
) -> Result<(KeyBlockHeader, SecretKey), Tr31CryptoError<P::Error>>
where
P: AesCmacKeyDerivation<K>
+ AesCbc<<P as AesCmacKeyDerivation<K>>::DerivedKey>
+ AesCmac<<P as AesCmacKeyDerivation<K>>::DerivedKey>,
{
let header = KeyBlockHeader::new_from_str(key_block)?;
let header_len = header.len();
let key_block_len = key_block.len();
let encoded_key_block_len = header.kb_length() as usize;
if key_block_len != encoded_key_block_len {
return Err(Tr31Error::KeyBlockLengthMismatch {
expected: encoded_key_block_len,
actual: key_block_len,
}
.into());
}
let min_key_block_len = 16 + (2 * TR31_D_BLOCK_LEN) + (2 * TR31_D_MAC_LEN);
if key_block_len < min_key_block_len {
return Err(Tr31Error::KeyBlockBelowMinimum {
minimum: min_key_block_len,
actual: key_block_len,
}
.into());
}
if header.version_id() != "D" {
return Err(Tr31Error::UnsupportedVersion(header.version_id().to_string()).into());
}
let mac_hex_len = TR31_D_MAC_LEN * 2;
let encrypted_payload_hex = &key_block[header_len..key_block_len - mac_hex_len];
let mac_hex = &key_block[key_block_len - mac_hex_len..];
let (kbek, kbak) =
derive_keys_version_d(provider, kbpk, kbpk_size).map_err(Tr31CryptoError::Crypto)?;
let encrypted_payload = hex::decode(encrypted_payload_hex)?;
let mac = hex::decode(mac_hex)?;
let iv: [u8; TR31_D_MAC_LEN] = mac.as_slice().try_into().map_err(|_| {
Tr31CryptoError::Tr31(Tr31Error::InvalidMacLength {
expected: TR31_D_MAC_LEN,
actual: mac.len(),
})
})?;
let decrypted_payload = Zeroizing::new(
provider
.decrypt_cbc(&kbek, &iv, &encrypted_payload)
.map_err(Tr31CryptoError::Crypto)?,
);
let mut mac_input = Zeroizing::new(key_block[..header_len].as_bytes().to_vec());
mac_input.extend_from_slice(decrypted_payload.as_slice());
let calculated_mac = provider
.calculate_cmac(&kbak, mac_input.as_slice())
.map_err(Tr31CryptoError::Crypto)?;
if mac.as_slice() != calculated_mac.as_slice() {
return Err(Tr31Error::MacVerificationFailed.into());
}
let key = extract_key_from_payload(decrypted_payload.as_slice())?;
Ok((header, SecretKey::new(key)))
}