use std::mem;
use super::{master_key::MasterKey, AES_BLOCK, AES_BLOCK64};
pub trait Backing {
type Error;
fn read(&mut self, dst: &mut [u8], offset: u64) -> Result<(), Self::Error>;
fn write(&mut self, src: &[u8], offset: u64) -> Result<(), Self::Error>;
fn len(&mut self) -> Result<u64, Self::Error>;
fn encryption_error() -> Self::Error;
}
const GROUP_STRIDE: usize = 16;
pub struct Xex {
nonce: [u8; AES_BLOCK],
enc: openssl::symm::Crypter,
dec: openssl::symm::Crypter,
tweaker: Tweaker,
tmp_tweaks: [[u8; AES_BLOCK]; GROUP_STRIDE],
tmp_enc: [u8; AES_BLOCK * GROUP_STRIDE + AES_BLOCK],
tmp_write: Vec<u8>,
}
impl Xex {
pub fn new(
master: &MasterKey,
filename: &str,
) -> Result<Self, crate::support::error::Error> {
let key = master.xex_key(filename);
let nonce = master.xex_nonce(filename);
let mut enc = openssl::symm::Crypter::new(
openssl::symm::Cipher::aes_128_ecb(),
openssl::symm::Mode::Encrypt,
&key,
None,
)?;
enc.pad(false);
let mut dec = openssl::symm::Crypter::new(
openssl::symm::Cipher::aes_128_ecb(),
openssl::symm::Mode::Decrypt,
&key,
None,
)?;
dec.pad(false);
Ok(Self {
nonce,
enc,
dec,
tweaker: Tweaker::new(),
tmp_tweaks: Default::default(),
tmp_enc: [0u8; AES_BLOCK * GROUP_STRIDE + AES_BLOCK],
tmp_write: Default::default(),
})
}
pub fn read<B: Backing>(
&mut self,
backing: &mut B,
mut dst: &mut [u8],
mut offset: u64,
) -> Result<(), B::Error> {
if dst.is_empty() {
return Ok(());
}
let end = offset + dst.len() as u64;
if !end.is_multiple_of(AES_BLOCK64) {
let final_block_start = end / AES_BLOCK64 * AES_BLOCK64;
let final_block_len = end - final_block_start;
if final_block_len as usize >= dst.len() {
let mut block = [0u8; AES_BLOCK];
let block = &mut block[..final_block_len as usize];
self.read_incomplete_aligned_block(
backing,
block,
final_block_start,
)?;
let inner_off = (offset - final_block_start) as usize;
dst.copy_from_slice(&block[inner_off..][..dst.len()]);
return Ok(());
}
let dst_block_start = dst.len() - final_block_len as usize;
self.read_incomplete_aligned_block(
backing,
&mut dst[dst_block_start..],
final_block_start,
)?;
dst = &mut dst[..dst_block_start];
}
if !offset.is_multiple_of(AES_BLOCK64) {
let first_block_start = offset / AES_BLOCK64 * AES_BLOCK64;
let first_block_len = AES_BLOCK64 - (offset - first_block_start);
let mut block = [0u8; AES_BLOCK];
self.read_aligned_blocks(backing, &mut block, first_block_start)?;
dst[..first_block_len as usize]
.copy_from_slice(&block[(offset % AES_BLOCK64) as usize..]);
dst = &mut dst[first_block_len as usize..];
offset += first_block_len;
}
self.read_aligned_blocks(backing, dst, offset)
}
fn read_incomplete_aligned_block<B: Backing>(
&mut self,
backing: &mut B,
dst: &mut [u8],
offset: u64,
) -> Result<(), B::Error> {
debug_assert_eq!(0, offset % AES_BLOCK64);
debug_assert!(dst.len() < AES_BLOCK);
let mut whole = [0u8; AES_BLOCK];
self.read_aligned_blocks(backing, &mut whole, offset)?;
dst.copy_from_slice(&whole[..dst.len()]);
Ok(())
}
fn read_aligned_blocks<B: Backing>(
&mut self,
backing: &mut B,
dst: &mut [u8],
offset: u64,
) -> Result<(), B::Error> {
debug_assert_eq!(0, offset % AES_BLOCK64);
debug_assert_eq!(0, dst.len() % AES_BLOCK);
if dst.is_empty() {
return Ok(());
}
backing.read(dst, offset)?;
self.crypt::<B>(dst, offset, false)?;
Ok(())
}
pub fn write<B: Backing>(
&mut self,
backing: &mut B,
mut src: &[u8],
mut offset: u64,
) -> Result<(), B::Error> {
let end = offset + src.len() as u64;
let tail_block_start = end / AES_BLOCK64 * AES_BLOCK64;
let tail_block_len = end - tail_block_start;
if tail_block_len as usize >= src.len() {
return self
.write_incomplete_misaligned_block(backing, src, offset);
}
if !offset.is_multiple_of(AES_BLOCK64) {
let new_offset = offset.next_multiple_of(AES_BLOCK64);
let advance = (new_offset - offset) as usize;
self.write_misaligned_block(backing, &src[..advance], offset)?;
src = &src[advance..];
offset = new_offset;
}
self.write_aligned_blocks(
backing,
&src[..(tail_block_start - offset) as usize],
offset,
)?;
if tail_block_len > 0 {
self.write_incomplete_misaligned_block(
backing,
&src[(tail_block_start - offset) as usize..],
tail_block_start,
)?;
}
Ok(())
}
fn write_incomplete_misaligned_block<B: Backing>(
&mut self,
backing: &mut B,
src: &[u8],
offset: u64,
) -> Result<(), B::Error> {
debug_assert!(src.len() < AES_BLOCK);
let mut block = [0u8; AES_BLOCK];
let block_start = offset / AES_BLOCK64 * AES_BLOCK64;
let end = offset + src.len() as u64;
if backing.len()? >= end {
self.read_aligned_blocks(backing, &mut block, block_start)?;
}
block[(offset % AES_BLOCK64) as usize..][..src.len()]
.copy_from_slice(src);
self.write_aligned_blocks(backing, &block, block_start)
}
fn write_misaligned_block<B: Backing>(
&mut self,
backing: &mut B,
src: &[u8],
offset: u64,
) -> Result<(), B::Error> {
debug_assert!(src.len() < AES_BLOCK);
debug_assert_eq!(0, (offset + src.len() as u64) % AES_BLOCK64);
let mut block = [0u8; AES_BLOCK];
let block_start = offset / AES_BLOCK64 * AES_BLOCK64;
self.read_aligned_blocks(backing, &mut block, block_start)?;
block[AES_BLOCK - src.len()..].copy_from_slice(src);
self.write_aligned_blocks(backing, &block, block_start)
}
fn write_aligned_blocks<B: Backing>(
&mut self,
backing: &mut B,
src: &[u8],
offset: u64,
) -> Result<(), B::Error> {
assert_eq!(0, offset % AES_BLOCK64);
assert_eq!(0, src.len() % AES_BLOCK);
let mut tmp = mem::take(&mut self.tmp_write);
tmp.clear();
tmp.extend_from_slice(src);
let result = self.crypt::<B>(&mut tmp, offset, true);
self.tmp_write = tmp;
result?;
backing.write(&self.tmp_write, offset)
}
fn crypt<B: Backing>(
&mut self,
dst: &mut [u8],
offset: u64,
encrypt: bool,
) -> Result<(), B::Error> {
let tmp = &mut self.tmp_enc;
let tweaks = &mut self.tmp_tweaks;
for (block_group, offset) in dst
.chunks_mut(tmp.len() - AES_BLOCK)
.zip((offset..).step_by(AES_BLOCK * GROUP_STRIDE))
{
for ((block, tweak), offset) in block_group
.chunks_exact_mut(AES_BLOCK)
.zip(tweaks.iter_mut())
.zip((offset..).step_by(AES_BLOCK))
{
*tweak =
self.tweaker.gen_tweak(&mut self.enc, &self.nonce, offset);
for (b, t) in block.iter_mut().zip(tweak.iter().copied()) {
*b ^= t;
}
}
let crypt = if encrypt {
&mut self.enc
} else {
&mut self.dec
};
match crypt.update(
block_group,
&mut tmp[..block_group.len() + AES_BLOCK],
) {
Ok(len) => {
debug_assert_eq!(block_group.len(), len);
},
Err(_) => return Err(B::encryption_error()),
}
for ((block, tmp), tweak) in block_group
.chunks_exact_mut(AES_BLOCK)
.zip(tmp.chunks_exact(AES_BLOCK))
.zip(&*tweaks)
{
for ((dst, src), tw) in block
.iter_mut()
.zip(tmp.iter().copied())
.zip(tweak.iter().copied())
{
*dst = src ^ tw;
}
}
}
Ok(())
}
}
const TWEAK_SECTOR_SIZE_BLOCKS: u64 = 64;
struct Tweaker {
sector: u64,
base: u128,
}
impl Tweaker {
fn new() -> Self {
Self {
sector: u64::MAX,
base: 0,
}
}
fn gen_tweak(
&mut self,
enc: &mut openssl::symm::Crypter,
nonce: &[u8; AES_BLOCK],
file_offset: u64,
) -> [u8; AES_BLOCK] {
debug_assert_eq!(0, file_offset % AES_BLOCK64);
let block = file_offset / AES_BLOCK64;
let sector = block / TWEAK_SECTOR_SIZE_BLOCKS;
let block_in_sector = (block % TWEAK_SECTOR_SIZE_BLOCKS) as u32;
if sector != self.sector {
let mut enc_in = *nonce;
for (b, e) in sector.to_le_bytes().into_iter().zip(&mut enc_in) {
*e ^= b;
}
let mut enc_out = [0u8; AES_BLOCK * 2];
let n = enc.update(&enc_in, &mut enc_out).unwrap();
assert_eq!(AES_BLOCK, n);
self.sector = sector;
self.base =
u128::from_le_bytes(enc_out[..AES_BLOCK].try_into().unwrap());
}
gf128_mul_exp2(self.base, block_in_sector).to_le_bytes()
}
}
fn gf128_mul_exp2(n: u128, exp2: u32) -> u128 {
debug_assert!(exp2 <= 120);
if 0 == exp2 {
return n;
}
let lo_half = n << exp2;
let hi_half = n >> 128 - exp2;
lo_half ^ (hi_half << 7) ^ (hi_half << 2) ^ (hi_half << 1) ^ hi_half
}
#[cfg(test)]
mod test {
use std::convert::TryFrom;
use super::*;
impl Backing for Vec<u8> {
type Error = &'static str;
fn read(
&mut self,
dst: &mut [u8],
offset: u64,
) -> Result<(), &'static str> {
let offset =
usize::try_from(offset).map_err(|_| "offset too large")?;
let end = offset
.checked_add(dst.len())
.ok_or("end beyond usize::MAX")?;
if end > Vec::len(self) {
return Err("read short");
}
dst.copy_from_slice(&self[offset..end]);
Ok(())
}
fn write(
&mut self,
src: &[u8],
offset: u64,
) -> Result<(), &'static str> {
let offset =
usize::try_from(offset).map_err(|_| "offset too large")?;
let end = offset
.checked_add(src.len())
.ok_or("end beyond usize::MAX")?;
if end > Vec::len(self) {
self.resize(offset, 0);
self.extend_from_slice(src);
} else {
self[offset..end].copy_from_slice(src);
}
Ok(())
}
fn len(&mut self) -> Result<u64, &'static str> {
Ok(Vec::len(self) as u64)
}
fn encryption_error() -> &'static str {
"encryption error"
}
}
#[test]
fn basic() {
let mut backing = Vec::<u8>::new();
let master = MasterKey::new();
let mut xex = Xex::new(&master, "some-db").unwrap();
xex.write(&mut backing, b"hello world", 0).unwrap();
let mut read = [0u8; 11];
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(b"hello world", &read);
}
#[test]
fn random_write() {
let mut backing = Vec::<u8>::new();
let master = MasterKey::new();
let mut xex = Xex::new(&master, "some-db").unwrap();
let mut read = [0u8; 48];
xex.write(
&mut backing,
b"The quick brown fox jumps over the lazy dog.",
0,
)
.unwrap();
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(
b"The quick brown fox jumps over the lazy dog.\x00\x00\x00\x00",
&read,
);
xex.write(&mut backing, b"FOX JUMPS OVER T", 16).unwrap();
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(
b"The quick brown FOX JUMPS OVER The lazy dog.\x00\x00\x00\x00",
&read,
);
xex.write(&mut backing, b"HE", 1).unwrap();
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(
b"THE quick brown FOX JUMPS OVER The lazy dog.\x00\x00\x00\x00",
&read,
);
xex.write(&mut backing, b"another doge", 31).unwrap();
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(
b"THE quick brown FOX JUMPS OVER another doge.\x00\x00\x00\x00",
&read,
);
xex.write(&mut backing, b"lazy dog jumps over the quick brown fox", 4)
.unwrap();
xex.read(&mut backing, &mut read, 0).unwrap();
assert_eq!(
b"THE lazy dog jumps over the quick brown fox.\x00\x00\x00\x00",
&read,
);
}
#[test]
fn random_read() {
let mut backing = Vec::<u8>::new();
let master = MasterKey::new();
let mut xex = Xex::new(&master, "some-db").unwrap();
let mut read = [0u8; 48];
xex.write(
&mut backing,
b"The quick brown fox jumps over the lazy dog.",
0,
)
.unwrap();
xex.read(&mut backing, &mut read[..8], 16).unwrap();
assert_eq!(b"fox jump", &read[..8]);
xex.read(&mut backing, &mut read[..8], 24).unwrap();
assert_eq!(b"s over t", &read[..8]);
xex.read(&mut backing, &mut read[..5], 4).unwrap();
assert_eq!(b"quick", &read[..5]);
xex.read(&mut backing, &mut read[..9], 10).unwrap();
assert_eq!(b"brown fox", &read[..9]);
xex.read(&mut backing, &mut read[..34], 10).unwrap();
assert_eq!(b"brown fox jumps over the lazy dog.", &read[..34]);
}
#[test]
fn test_gf128_mul_exp2() {
fn gf128_mul_exp2_slow(mut n: u128, exp2: u32) -> u128 {
for _ in 0..exp2 {
let high_bit = 0 != n >> 127;
n <<= 1;
if high_bit {
n ^= 0b10000111;
}
}
n
}
for n in [
1u128,
2,
3,
4,
0x2ee4108684a71b2f89d0c73a9ef65217,
0xb596d9bfada34102f38f7a24fa107bd7,
0xacf0daba883f2031ccba47bf3947abfa,
0xbfc3ed42a15e8c4afc13c344692ae1e3,
0xd03cfd5db6e15417e9d3144450f4f5f1,
0xc2b8bd6e23b61e678140bafd3bb15307,
0x828197e375f4979e39c138ca9e4e8657,
0x2c7a7ad6a2709581124e17791209dc60,
] {
for exp2 in 0..=64 {
let slow = gf128_mul_exp2_slow(n, exp2);
let fast = gf128_mul_exp2(n, exp2);
assert_eq!(slow, fast, "for {n} * 2**{exp2}");
}
}
}
}