use crate::{
block_cipher::BlockCipher,
error::{CryptoError, Result},
};
pub trait XtsCipher {
fn encrypt(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()>;
fn decrypt(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()>;
fn key_size() -> usize;
}
#[allow(clippy::needless_range_loop)]
fn gf128_mul(x: &mut [u8; 16]) {
let mut feedback: u8 = 0;
for i in 0..16 {
let tmp = x[i];
x[i] = (tmp << 1) | feedback;
feedback = tmp >> 7;
}
if feedback != 0 {
x[0] ^= 0x87;
}
}
fn xts_encrypt<B: BlockCipher>(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()> {
let half = B::key_size();
if key.len() != half * 2 {
return Err(CryptoError::InvalidKey);
}
if tweak.len() != 16 {
return Err(CryptoError::InvalidLength);
}
if data.len() < 16 {
return Err(CryptoError::InvalidInput);
}
let data_key = &key[..half];
let tweak_key = &key[half..];
let mut xts_tweak: [u8; 16] = tweak.try_into().map_err(|_| CryptoError::InvalidLength)?;
B::encrypt(tweak_key, &mut xts_tweak)?;
let full_blocks = data.len() / 16;
let remainder = data.len() % 16;
for i in 0..full_blocks {
let start = i * 16;
let block = &mut data[start..start + 16];
for j in 0..16 {
block[j] ^= xts_tweak[j];
}
B::encrypt(data_key, block)?;
for j in 0..16 {
block[j] ^= xts_tweak[j];
}
gf128_mul(&mut xts_tweak);
}
if remainder > 0 {
let last_full = full_blocks - 1;
let last_full_start = last_full * 16;
let mut cc: [u8; 16] = [0u8; 16];
cc.copy_from_slice(&data[last_full_start..last_full_start + 16]);
let tail_start = full_blocks * 16;
let tmp = data[tail_start..tail_start + remainder].to_vec();
data[last_full_start..last_full_start + remainder].copy_from_slice(&tmp);
data[last_full_start + remainder..last_full_start + 16].copy_from_slice(&cc[remainder..16]);
let combined_start = last_full_start;
for j in 0..16 {
data[combined_start + j] ^= xts_tweak[j];
}
B::encrypt(data_key, &mut data[combined_start..combined_start + 16])?;
for j in 0..16 {
data[combined_start + j] ^= xts_tweak[j];
}
data[tail_start..tail_start + remainder].copy_from_slice(&cc[..remainder]);
}
Ok(())
}
fn xts_decrypt<B: BlockCipher>(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()> {
let half = B::key_size();
if key.len() != half * 2 {
return Err(CryptoError::InvalidKey);
}
if tweak.len() != 16 {
return Err(CryptoError::InvalidLength);
}
if data.len() < 16 {
return Err(CryptoError::InvalidInput);
}
let data_key = &key[..half];
let tweak_key = &key[half..];
let mut xts_tweak: [u8; 16] = tweak.try_into().map_err(|_| CryptoError::InvalidLength)?;
B::encrypt(tweak_key, &mut xts_tweak)?;
let full_blocks = data.len() / 16;
let remainder = data.len() % 16;
if remainder == 0 {
for i in 0..full_blocks {
let start = i * 16;
let block = &mut data[start..start + 16];
for j in 0..16 {
block[j] ^= xts_tweak[j];
}
B::decrypt(data_key, block)?;
for j in 0..16 {
block[j] ^= xts_tweak[j];
}
gf128_mul(&mut xts_tweak);
}
} else {
let mut tweak_nm1: [u8; 16] = xts_tweak;
for _ in 0..full_blocks - 1 {
gf128_mul(&mut tweak_nm1);
}
let mut tweak_n: [u8; 16] = tweak_nm1;
gf128_mul(&mut tweak_n);
let last_full = full_blocks - 1;
let last_full_start = last_full * 16;
let mut pp: [u8; 16] = [0u8; 16];
pp.copy_from_slice(&data[last_full_start..last_full_start + 16]);
for j in 0..16 {
pp[j] ^= tweak_n[j];
}
{
let mut block_arr = pp;
B::decrypt(data_key, &mut block_arr)?;
pp = block_arr;
}
for j in 0..16 {
pp[j] ^= tweak_n[j];
}
let tail_start = full_blocks * 16;
let mut cc: [u8; 16] = [0u8; 16];
cc[..remainder].copy_from_slice(&data[tail_start..tail_start + remainder]);
cc[remainder..16].copy_from_slice(&pp[remainder..16]);
for j in 0..16 {
cc[j] ^= tweak_nm1[j];
}
B::decrypt(data_key, &mut cc)?;
for j in 0..16 {
cc[j] ^= tweak_nm1[j];
}
data[last_full_start..last_full_start + 16].copy_from_slice(&cc);
data[tail_start..tail_start + remainder].copy_from_slice(&pp[..remainder]);
let mut tw: [u8; 16] = tweak.try_into().map_err(|_| CryptoError::InvalidLength)?;
B::encrypt(tweak_key, &mut tw)?;
for i in 0..last_full {
let start = i * 16;
let block = &mut data[start..start + 16];
for j in 0..16 {
block[j] ^= tw[j];
}
B::decrypt(data_key, block)?;
for j in 0..16 {
block[j] ^= tw[j];
}
gf128_mul(&mut tw);
}
}
Ok(())
}
pub struct Aes128Xts;
pub struct Aes256Xts;
pub struct Sm4Xts;
macro_rules! impl_xts {
($wrapper:ident, $block:ty, $single_key:expr) => {
impl XtsCipher for $wrapper {
fn encrypt(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()> {
xts_encrypt::<$block>(key, tweak, data)
}
fn decrypt(key: &[u8], tweak: &[u8], data: &mut [u8]) -> Result<()> {
xts_decrypt::<$block>(key, tweak, data)
}
fn key_size() -> usize {
$single_key * 2
}
}
};
}
impl_xts!(Aes128Xts, crate::block_cipher::Aes128Ecb, 16);
impl_xts!(Aes256Xts, crate::block_cipher::Aes256Ecb, 32);
impl_xts!(Sm4Xts, crate::block_cipher::Sm4Ecb, 16);