use aes::cipher::{KeyIvInit, StreamCipher};
use aes::{Aes128, Aes192, Aes256};
use aes_kw::{KekAes128, KekAes192, KekAes256};
use alloc::vec;
use alloc::vec::Vec;
use ctr::Ctr128BE;
use hmac::Hmac;
use pbkdf2::pbkdf2;
use sha1::Sha1;
use crate::error::{Error, Result};
use crate::packet::EncryptionKeyField;
pub const PBKDF2_ITERATIONS: u32 = 2048;
pub const SALT_LEN: usize = 16;
const PBKDF2_SALT_LEN: usize = 8;
const IV_SALT_LEN: usize = 14;
const AES_BLOCK_LEN: usize = 16;
fn invalid_key_length(what: &'static str) -> Error {
Error::InvalidField {
what,
reason: "length must be 16, 24, or 32 bytes (AES-128/192/256)",
}
}
pub fn derive_kek(passphrase: &[u8], salt: &[u8; SALT_LEN], klen: usize) -> Result<Vec<u8>> {
if !matches!(klen, 16 | 24 | 32) {
return Err(invalid_key_length("KLen"));
}
let pbkdf2_salt = &salt[SALT_LEN - PBKDF2_SALT_LEN..];
let mut kek = vec![0u8; klen];
pbkdf2::<Hmac<Sha1>>(passphrase, pbkdf2_salt, PBKDF2_ITERATIONS, &mut kek)
.expect("HMAC-SHA1 accepts any key length, so PBKDF2 cannot fail here");
Ok(kek)
}
pub fn wrap_sek(kek: &[u8], plaintext_keys: &[u8]) -> Result<([u8; 8], Vec<u8>)> {
let mut out = vec![0u8; plaintext_keys.len() + 8];
match kek.len() {
16 => KekAes128::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.wrap(plaintext_keys, &mut out),
24 => KekAes192::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.wrap(plaintext_keys, &mut out),
32 => KekAes256::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.wrap(plaintext_keys, &mut out),
_ => return Err(invalid_key_length("KEK")),
}
.map_err(|_| Error::InvalidField {
what: "SEK",
reason: "length must be a multiple of 8 bytes (RFC 3394 semiblocks)",
})?;
let mut icv = [0u8; 8];
icv.copy_from_slice(&out[..8]);
Ok((icv, out[8..].to_vec()))
}
pub fn unwrap_sek(kek: &[u8], icv: &[u8; 8], wrapped: &[u8]) -> Result<Vec<u8>> {
let mut input = Vec::with_capacity(8 + wrapped.len());
input.extend_from_slice(icv);
input.extend_from_slice(wrapped);
let mut out = vec![0u8; wrapped.len()];
let bad_wrap = || Error::InvalidField {
what: "AES key wrap",
reason: "integrity check failed (wrong KEK / passphrase, or corrupt wire data)",
};
match kek.len() {
16 => KekAes128::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.unwrap(&input, &mut out),
24 => KekAes192::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.unwrap(&input, &mut out),
32 => KekAes256::try_from(kek)
.map_err(|_| invalid_key_length("KEK"))?
.unwrap(&input, &mut out),
_ => return Err(invalid_key_length("KEK")),
}
.map_err(|_| bad_wrap())?;
Ok(out)
}
pub fn packet_counter(salt: &[u8; SALT_LEN], pkt_seq_no: u32) -> [u8; AES_BLOCK_LEN] {
let mut counter = [0u8; AES_BLOCK_LEN];
counter[10..14].copy_from_slice(&pkt_seq_no.to_be_bytes());
for i in 0..IV_SALT_LEN {
counter[i] ^= salt[i];
}
counter
}
pub fn aes_ctr_apply(
sek: &[u8],
salt: &[u8; SALT_LEN],
pkt_seq_no: u32,
data: &mut [u8],
) -> Result<()> {
let counter = packet_counter(salt, pkt_seq_no);
match sek.len() {
16 => {
let mut cipher =
Ctr128BE::<Aes128>::new_from_slices(sek, &counter).map_err(|_| bad_sek())?;
cipher.apply_keystream(data);
}
24 => {
let mut cipher =
Ctr128BE::<Aes192>::new_from_slices(sek, &counter).map_err(|_| bad_sek())?;
cipher.apply_keystream(data);
}
32 => {
let mut cipher =
Ctr128BE::<Aes256>::new_from_slices(sek, &counter).map_err(|_| bad_sek())?;
cipher.apply_keystream(data);
}
_ => return Err(bad_sek()),
}
Ok(())
}
fn bad_sek() -> Error {
invalid_key_length("SEK")
}
pub fn select_sek<'a>(
key_flag: EncryptionKeyField,
even: Option<&'a [u8]>,
odd: Option<&'a [u8]>,
) -> Result<&'a [u8]> {
let no_sek = |parity: &'static str| Error::InvalidField {
what: "SEK",
reason: match parity {
"even" => "even key not currently held",
_ => "odd key not currently held",
},
};
match key_flag {
EncryptionKeyField::Even => even.ok_or_else(|| no_sek("even")),
EncryptionKeyField::Odd => odd.ok_or_else(|| no_sek("odd")),
EncryptionKeyField::NotEncrypted => Err(Error::InvalidField {
what: "KK",
reason: "packet is not encrypted (KK=00b)",
}),
EncryptionKeyField::Reserved(_) => Err(Error::InvalidField {
what: "KK",
reason: "reserved value (11b) is control-packet-only, not valid on a data packet",
}),
}
}
#[cfg(test)]
mod tests {
use super::*;
const RFC3394_KEK_128: [u8; 16] = [
0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E,
0x0F,
];
const RFC3394_KEY_DATA_128: [u8; 16] = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, 0xAA, 0xBB, 0xCC, 0xDD, 0xEE,
0xFF,
];
const RFC3394_WRAPPED_128: [u8; 24] = [
0x1F, 0xA6, 0x8B, 0x0A, 0x81, 0x12, 0xB4, 0x47, 0xAE, 0xF3, 0x4B, 0xD8, 0xFB, 0x5A, 0x7B,
0x82, 0x9D, 0x3E, 0x86, 0x23, 0x71, 0xD2, 0xCF, 0xE5,
];
#[test]
fn rfc3394_wrap_matches_worked_vector() {
let (icv, wrapped) = wrap_sek(&RFC3394_KEK_128, &RFC3394_KEY_DATA_128).unwrap();
assert_eq!(&icv[..], &RFC3394_WRAPPED_128[..8]);
assert_eq!(wrapped.as_slice(), &RFC3394_WRAPPED_128[8..]);
}
#[test]
fn packet_counter_matches_libsrt_hcrypt_setctriv() {
let salt: [u8; SALT_LEN] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E,
0x0F, 0x10,
];
let expected: [u8; AES_BLOCK_LEN] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x1A, 0x2E, 0x3E, 0x4A,
0x00, 0x00,
];
let got = packet_counter(&salt, 0x1122_3344);
assert_eq!(got, expected, "counter must match hcrypt_SetCtrIV (no <<2)");
assert_eq!(got[0], 0x01, "byte 0 is pure salt, not the <<2 value 0x04");
}
#[test]
fn rfc3394_unwrap_matches_worked_vector() {
let mut icv = [0u8; 8];
icv.copy_from_slice(&RFC3394_WRAPPED_128[..8]);
let recovered = unwrap_sek(&RFC3394_KEK_128, &icv, &RFC3394_WRAPPED_128[8..]).unwrap();
assert_eq!(recovered.as_slice(), &RFC3394_KEY_DATA_128[..]);
}
#[test]
fn rfc3394_unwrap_rejects_wrong_kek() {
let mut icv = [0u8; 8];
icv.copy_from_slice(&RFC3394_WRAPPED_128[..8]);
let mut wrong_kek = RFC3394_KEK_128;
wrong_kek[0] ^= 0xFF;
assert!(unwrap_sek(&wrong_kek, &icv, &RFC3394_WRAPPED_128[8..]).is_err());
}
const NIST_F5_1_KEY: [u8; 16] = [
0x2B, 0x7E, 0x15, 0x16, 0x28, 0xAE, 0xD2, 0xA6, 0xAB, 0xF7, 0x15, 0x88, 0x09, 0xCF, 0x4F,
0x3C,
];
const NIST_F5_1_INIT_COUNTER: [u8; 16] = [
0xF0, 0xF1, 0xF2, 0xF3, 0xF4, 0xF5, 0xF6, 0xF7, 0xF8, 0xF9, 0xFA, 0xFB, 0xFC, 0xFD, 0xFE,
0xFF,
];
const NIST_F5_1_PLAINTEXT: [u8; 64] = [
0x6B, 0xC1, 0xBE, 0xE2, 0x2E, 0x40, 0x9F, 0x96, 0xE9, 0x3D, 0x7E, 0x11, 0x73, 0x93, 0x17,
0x2A, 0xAE, 0x2D, 0x8A, 0x57, 0x1E, 0x03, 0xAC, 0x9C, 0x9E, 0xB7, 0x6F, 0xAC, 0x45, 0xAF,
0x8E, 0x51, 0x30, 0xC8, 0x1C, 0x46, 0xA3, 0x5C, 0xE4, 0x11, 0xE5, 0xFB, 0xC1, 0x19, 0x1A,
0x0A, 0x52, 0xEF, 0xF6, 0x9F, 0x24, 0x45, 0xDF, 0x4F, 0x9B, 0x17, 0xAD, 0x2B, 0x41, 0x7B,
0xE6, 0x6C, 0x37, 0x10,
];
const NIST_F5_1_CIPHERTEXT: [u8; 64] = [
0x87, 0x4D, 0x61, 0x91, 0xB6, 0x20, 0xE3, 0x26, 0x1B, 0xEF, 0x68, 0x64, 0x99, 0x0D, 0xB6,
0xCE, 0x98, 0x06, 0xF6, 0x6B, 0x79, 0x70, 0xFD, 0xFF, 0x86, 0x17, 0x18, 0x7B, 0xB9, 0xFF,
0xFD, 0xFF, 0x5A, 0xE4, 0xDF, 0x3E, 0xDB, 0xD5, 0xD3, 0x5E, 0x5B, 0x4F, 0x09, 0x02, 0x0D,
0xB0, 0x3E, 0xAB, 0x1E, 0x03, 0x1D, 0xDA, 0x2F, 0xBE, 0x03, 0xD1, 0x79, 0x21, 0x70, 0xA0,
0xF3, 0x00, 0x9C, 0xEE,
];
#[test]
fn nist_sp800_38a_f5_1_ctr_aes128_encrypt() {
let mut buf = NIST_F5_1_PLAINTEXT;
let mut cipher =
Ctr128BE::<Aes128>::new_from_slices(&NIST_F5_1_KEY, &NIST_F5_1_INIT_COUNTER).unwrap();
cipher.apply_keystream(&mut buf);
assert_eq!(buf, NIST_F5_1_CIPHERTEXT);
}
#[test]
fn nist_sp800_38a_f5_1_ctr_aes128_decrypt() {
let mut buf = NIST_F5_1_CIPHERTEXT;
let mut cipher =
Ctr128BE::<Aes128>::new_from_slices(&NIST_F5_1_KEY, &NIST_F5_1_INIT_COUNTER).unwrap();
cipher.apply_keystream(&mut buf);
assert_eq!(buf, NIST_F5_1_PLAINTEXT);
}
#[test]
fn srt_payload_round_trips_and_wrong_sek_does_not_recover() {
let sek = [0x42u8; 16];
let salt = [0x99u8; SALT_LEN];
let pkt_seq_no = 0x0123_4567u32;
let plaintext = b"SRT payload encryption round trip test vector.".to_vec();
let mut encrypted = plaintext.clone();
aes_ctr_apply(&sek, &salt, pkt_seq_no, &mut encrypted).unwrap();
assert_ne!(encrypted, plaintext, "encryption must change the bytes");
let mut decrypted = encrypted.clone();
aes_ctr_apply(&sek, &salt, pkt_seq_no, &mut decrypted).unwrap();
assert_eq!(
decrypted, plaintext,
"correct SEK must recover the plaintext"
);
let wrong_sek = [0x43u8; 16];
let mut wrongly_decrypted = encrypted;
aes_ctr_apply(&wrong_sek, &salt, pkt_seq_no, &mut wrongly_decrypted).unwrap();
assert_ne!(
wrongly_decrypted, plaintext,
"wrong SEK must not recover the plaintext"
);
}
#[test]
fn different_seq_no_gives_different_keystream() {
let sek = [0x11u8; 24];
let salt = [0x22u8; SALT_LEN];
let plaintext = [0u8; 32];
let mut a = plaintext;
aes_ctr_apply(&sek, &salt, 1, &mut a).unwrap();
let mut b = plaintext;
aes_ctr_apply(&sek, &salt, 2, &mut b).unwrap();
assert_ne!(a, b);
}
#[test]
fn kek_derivation_all_sizes_and_deterministic() {
for klen in [16usize, 24, 32] {
let salt = [0xABu8; SALT_LEN];
let kek1 = derive_kek(b"correct horse battery staple", &salt, klen).unwrap();
let kek2 = derive_kek(b"correct horse battery staple", &salt, klen).unwrap();
assert_eq!(kek1.len(), klen);
assert_eq!(kek1, kek2, "PBKDF2 is deterministic for the same inputs");
let different_salt = [0xACu8; SALT_LEN];
let kek3 = derive_kek(b"correct horse battery staple", &different_salt, klen).unwrap();
assert_ne!(kek1, kek3, "different salt must give a different KEK");
}
}
#[test]
fn invalid_klen_errs_without_panic() {
let salt = [0u8; SALT_LEN];
assert!(derive_kek(b"pw", &salt, 20).is_err());
assert!(wrap_sek(&[0u8; 20], &[0u8; 16]).is_err());
assert!(aes_ctr_apply(&[0u8; 20], &salt, 0, &mut [0u8; 4]).is_err());
}
#[test]
fn select_sek_picks_correct_parity_and_rejects_bad_flags() {
let even = [1u8; 16];
let odd = [2u8; 16];
assert_eq!(
select_sek(EncryptionKeyField::Even, Some(&even), Some(&odd)).unwrap(),
&even[..]
);
assert_eq!(
select_sek(EncryptionKeyField::Odd, Some(&even), Some(&odd)).unwrap(),
&odd[..]
);
assert!(select_sek(EncryptionKeyField::Even, None, Some(&odd)).is_err());
assert!(select_sek(EncryptionKeyField::NotEncrypted, Some(&even), Some(&odd)).is_err());
assert!(select_sek(EncryptionKeyField::Reserved(0b11), Some(&even), Some(&odd)).is_err());
}
#[test]
fn both_seks_wrap_unwrap_round_trip() {
let kek = [0x77u8; 16];
let even_sek = [0xAAu8; 16];
let odd_sek = [0xBBu8; 16];
let mut plaintext = Vec::new();
plaintext.extend_from_slice(&even_sek);
plaintext.extend_from_slice(&odd_sek);
let (icv, wrapped) = wrap_sek(&kek, &plaintext).unwrap();
assert_eq!(wrapped.len(), 32);
let recovered = unwrap_sek(&kek, &icv, &wrapped).unwrap();
assert_eq!(recovered, plaintext);
assert_eq!(&recovered[..16], &even_sek[..]);
assert_eq!(&recovered[16..], &odd_sek[..]);
}
}