use aes::cipher::KeyInit;
use aes::Aes256;
use serpent::Serpent;
use twofish::Twofish;
use xts_mode::Xts128;
use crate::error::{Result, VeraError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Prf {
Sha512,
Sha256,
Whirlpool,
Streebog,
Ripemd160,
}
impl Prf {
#[must_use]
pub fn all() -> [Prf; 5] {
[
Prf::Sha512,
Prf::Sha256,
Prf::Whirlpool,
Prf::Streebog,
Prf::Ripemd160,
]
}
#[must_use]
pub fn name(self) -> &'static str {
match self {
Prf::Sha512 => "sha512",
Prf::Sha256 => "sha256",
Prf::Whirlpool => "whirlpool",
Prf::Streebog => "streebog",
Prf::Ripemd160 => "ripemd160",
}
}
#[must_use]
pub fn iterations(self) -> u32 {
match self {
Prf::Ripemd160 => 655_331,
_ => 500_000,
}
}
#[must_use]
pub fn iterations_pim(self, pim: u32) -> u32 {
if pim == 0 {
return self.iterations();
}
match self {
Prf::Ripemd160 => pim.saturating_mul(2048),
_ => 15_000u32.saturating_add(pim.saturating_mul(1000)),
}
}
pub fn derive(self, password: &[u8], salt: &[u8], iterations: u32, out_len: usize) -> Vec<u8> {
let mut out = vec![0u8; out_len];
let it = iterations.max(1);
match self {
Prf::Sha512 => pbkdf2::pbkdf2_hmac::<sha2::Sha512>(password, salt, it, &mut out),
Prf::Sha256 => pbkdf2::pbkdf2_hmac::<sha2::Sha256>(password, salt, it, &mut out),
Prf::Whirlpool => {
pbkdf2::pbkdf2_hmac::<whirlpool::Whirlpool>(password, salt, it, &mut out);
}
Prf::Streebog => {
pbkdf2::pbkdf2_hmac::<streebog::Streebog512>(password, salt, it, &mut out);
}
Prf::Ripemd160 => {
pbkdf2::pbkdf2_hmac::<ripemd::Ripemd160>(password, salt, it, &mut out);
}
}
out
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Cipher {
Aes,
Serpent,
Twofish,
}
impl Cipher {
#[must_use]
pub fn all() -> [Cipher; 3] {
[Cipher::Aes, Cipher::Serpent, Cipher::Twofish]
}
#[must_use]
pub fn name(self) -> &'static str {
match self {
Cipher::Aes => "aes",
Cipher::Serpent => "serpent",
Cipher::Twofish => "twofish",
}
}
#[must_use]
pub fn key_len(self) -> usize {
64
}
}
#[must_use]
pub fn cipher_chains() -> Vec<Vec<Cipher>> {
use Cipher::{Aes, Serpent, Twofish};
vec![
vec![Aes],
vec![Serpent],
vec![Twofish],
vec![Twofish, Aes], vec![Serpent, Twofish, Aes], vec![Aes, Serpent], vec![Aes, Twofish, Serpent], vec![Serpent, Twofish], ]
}
pub const MAX_CHAIN_KEY_LEN: usize = 3 * 64;
pub fn xts_decrypt_chain(
ciphers: &[Cipher],
key: &[u8],
buffer: &mut [u8],
unit_size: usize,
base_unit: u128,
) -> Result<()> {
let n = ciphers.len();
if n == 0 || key.len() < 64 * n {
return Err(VeraError::Crypto {
what: "cascade key too short",
});
}
let mut subkey = [0u8; 64];
for j in (0..n).rev() {
subkey[..32].copy_from_slice(&key[32 * j..32 * j + 32]);
subkey[32..].copy_from_slice(&key[32 * (n + j)..32 * (n + j) + 32]);
xts_decrypt(ciphers[j], &subkey, buffer, unit_size, base_unit)?;
}
Ok(())
}
pub fn xts_decrypt(
cipher: Cipher,
key: &[u8],
buffer: &mut [u8],
unit_size: usize,
base_unit: u128,
) -> Result<()> {
if key.len() != 64 {
return Err(VeraError::Crypto {
what: "xts key must be 64 bytes",
});
}
let (k1, k2) = key.split_at(32);
match cipher {
Cipher::Aes => decrypt_units(
&Xts128::new(Aes256::new(k1.into()), Aes256::new(k2.into())),
buffer,
unit_size,
base_unit,
),
Cipher::Serpent => {
let c1 = Serpent::new_from_slice(k1).map_err(|_| VeraError::Crypto {
what: "serpent 256-bit key",
})?; let c2 = Serpent::new_from_slice(k2).map_err(|_| VeraError::Crypto {
what: "serpent 256-bit key",
})?; decrypt_units(&Xts128::new(c1, c2), buffer, unit_size, base_unit);
}
Cipher::Twofish => decrypt_units(
&Xts128::new(Twofish::new(k1.into()), Twofish::new(k2.into())),
buffer,
unit_size,
base_unit,
),
}
Ok(())
}
fn decrypt_units<C>(xts: &Xts128<C>, buffer: &mut [u8], unit_size: usize, base: u128)
where
C: aes::cipher::BlockCipher + aes::cipher::BlockEncrypt + aes::cipher::BlockDecrypt,
{
for (u, chunk) in buffer.chunks_mut(unit_size).enumerate() {
if chunk.len() < 16 {
continue; }
let tweak = (base + u as u128).to_le_bytes();
xts.decrypt_sector(chunk, tweak);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn prf_iterations_and_names() {
assert_eq!(Prf::Sha512.iterations(), 500_000);
assert_eq!(Prf::Ripemd160.iterations(), 655_331);
assert_eq!(Prf::Sha512.iterations_pim(0), 500_000);
assert_eq!(Prf::Sha512.iterations_pim(10), 25_000);
assert_eq!(Prf::Ripemd160.iterations_pim(10), 20_480);
assert_eq!(Prf::all().len(), 5);
assert_eq!(
Cipher::all().map(Cipher::name),
["aes", "serpent", "twofish"]
);
}
#[test]
fn every_prf_has_a_name_and_derives() {
assert_eq!(Prf::Sha256.name(), "sha256");
assert_eq!(Prf::Whirlpool.name(), "whirlpool");
assert_eq!(Prf::Streebog.name(), "streebog");
assert_eq!(Prf::Ripemd160.name(), "ripemd160");
for prf in Prf::all() {
let k = prf.derive(b"password", b"salt", 1, 64);
assert_eq!(k.len(), 64, "prf {} derive length", prf.name());
}
}
#[test]
fn cipher_key_len_is_two_256bit_subkeys() {
assert_eq!(Cipher::Aes.key_len(), 64);
assert_eq!(Cipher::Twofish.key_len(), 64);
}
#[test]
fn pbkdf2_sha512_matches_python() {
let k = Prf::Sha512.derive(b"password", b"salt", 1, 32);
assert_eq!(
hex(&k),
"867f70cf1ade02cff3752599a3a53dc4af34c7a669815ae5d513554e1c8cf252"
);
}
#[test]
fn xts_roundtrips_for_each_cipher() {
for cipher in Cipher::all() {
let key = [0x24u8; 64];
let mut buf = vec![0u8; 512];
for (i, b) in buf.iter_mut().enumerate() {
*b = (i as u8) ^ 0x3c;
}
let plain = buf.clone();
let (k1, k2) = key.split_at(32);
match cipher {
Cipher::Aes => encrypt_one(
&Xts128::new(Aes256::new(k1.into()), Aes256::new(k2.into())),
&mut buf,
256,
),
Cipher::Serpent => encrypt_one(
&Xts128::new(
Serpent::new_from_slice(k1).unwrap(),
Serpent::new_from_slice(k2).unwrap(),
),
&mut buf,
256,
),
Cipher::Twofish => encrypt_one(
&Xts128::new(Twofish::new(k1.into()), Twofish::new(k2.into())),
&mut buf,
256,
),
}
xts_decrypt(cipher, &key, &mut buf, 512, 256).unwrap();
assert_eq!(buf, plain, "cipher {}", cipher.name());
}
}
#[test]
fn xts_rejects_bad_key_len() {
let mut b = [0u8; 512];
assert!(matches!(
xts_decrypt(Cipher::Aes, &[0u8; 48], &mut b, 512, 0),
Err(VeraError::Crypto { .. })
));
}
#[test]
fn cipher_chains_are_the_eight_veracrypt_chains() {
let chains = cipher_chains();
assert_eq!(chains.len(), 8);
assert_eq!(chains[0], vec![Cipher::Aes]);
assert_eq!(chains[3], vec![Cipher::Twofish, Cipher::Aes]);
assert_eq!(
chains[6],
vec![Cipher::Aes, Cipher::Twofish, Cipher::Serpent]
);
assert!(chains.iter().all(|c| c.len() <= 3));
}
#[test]
fn xts_decrypt_chain_rejects_empty_and_short_key() {
let mut b = [0u8; 512];
assert!(matches!(
xts_decrypt_chain(&[], &[0u8; 64], &mut b, 512, 0),
Err(VeraError::Crypto { .. })
));
assert!(matches!(
xts_decrypt_chain(&[Cipher::Aes, Cipher::Twofish], &[0u8; 100], &mut b, 512, 0),
Err(VeraError::Crypto { .. })
));
}
#[test]
fn xts_decrypt_chain_roundtrips_a_three_cipher_cascade() {
let ciphers = [Cipher::Aes, Cipher::Twofish, Cipher::Serpent];
let n = ciphers.len();
let mut key = [0u8; 192];
for (i, b) in key.iter_mut().enumerate() {
*b = (i as u8).wrapping_mul(7) ^ 0x5a;
}
let mut buf = vec![0u8; 512];
for (i, b) in buf.iter_mut().enumerate() {
*b = (i as u8) ^ 0x91;
}
let plain = buf.clone();
for j in 0..n {
let mut k1 = [0u8; 32];
let mut k2 = [0u8; 32];
k1.copy_from_slice(&key[32 * j..32 * j + 32]);
k2.copy_from_slice(&key[32 * (n + j)..32 * (n + j) + 32]);
match ciphers[j] {
Cipher::Aes => encrypt_one(
&Xts128::new(Aes256::new((&k1).into()), Aes256::new((&k2).into())),
&mut buf,
256,
),
Cipher::Twofish => encrypt_one(
&Xts128::new(Twofish::new((&k1).into()), Twofish::new((&k2).into())),
&mut buf,
256,
),
Cipher::Serpent => encrypt_one(
&Xts128::new(
Serpent::new_from_slice(&k1).unwrap(),
Serpent::new_from_slice(&k2).unwrap(),
),
&mut buf,
256,
),
}
}
assert_ne!(buf, plain, "cascade must actually encrypt");
xts_decrypt_chain(&ciphers, &key, &mut buf, 512, 256).unwrap();
assert_eq!(buf, plain, "two-cipher cascade round-trip");
}
fn encrypt_one<C>(xts: &Xts128<C>, buf: &mut [u8], unit: u128)
where
C: aes::cipher::BlockCipher + aes::cipher::BlockEncrypt + aes::cipher::BlockDecrypt,
{
xts.encrypt_sector(buf, unit.to_le_bytes());
}
fn hex(b: &[u8]) -> String {
use std::fmt::Write;
b.iter().fold(String::new(), |mut s, x| {
let _ = write!(s, "{x:02x}");
s
})
}
}