mod block;
#[cfg(test)]
mod circuit_tests;
mod circuits;
mod stream;
use crate::key::ControlWord;
use block::BitslicedBlock;
use stream::BitslicedStream;
type Word = u64;
pub const LANES: usize = Word::BITS as usize;
const BLOCK_BYTES: usize = 8;
const BITS_PER_BYTE: usize = 8;
const BLOCK_BITS: usize = BLOCK_BYTES * BITS_PER_BYTE;
const _: () = assert!(BLOCK_BITS == LANES, "the transpose is square");
fn transpose(m: &mut [Word; LANES]) {
let mut j = LANES / 2;
let mut mask: Word = !0 >> (LANES / 2);
while j != 0 {
let mut k = 0;
while k < LANES {
let t = ((m[k] >> j) ^ m[k | j]) & mask;
m[k] ^= t << j;
m[k | j] ^= t;
k = ((k | j) + 1) & !j;
}
j >>= 1;
mask ^= mask << j;
}
}
struct Group {
len: [usize; LANES],
blocks: [usize; LANES],
max_blocks: usize,
max_stream: usize,
}
impl Group {
fn new(payloads: &[&mut [u8]]) -> Self {
let mut g = Self {
len: [0; LANES],
blocks: [0; LANES],
max_blocks: 0,
max_stream: 0,
};
for (lane, p) in payloads.iter().enumerate() {
if p.len() < BLOCK_BYTES {
continue;
}
g.len[lane] = p.len();
g.blocks[lane] = p.len() / BLOCK_BYTES;
g.max_blocks = g.max_blocks.max(g.blocks[lane]);
g.max_stream = g.max_stream.max(p.len() - BLOCK_BYTES);
}
g
}
}
fn gather(payloads: &[&mut [u8]], offsets: &[Option<usize>; LANES], m: &mut [Word; LANES]) {
*m = [0; LANES];
for (lane, off) in offsets.iter().enumerate() {
if let Some(off) = *off {
let bytes: [u8; BLOCK_BYTES] = payloads[lane][off..off + BLOCK_BYTES]
.try_into()
.expect("slice is exactly one block");
m[lane] = Word::from_le_bytes(bytes);
}
}
}
fn scatter(payloads: &mut [&mut [u8]], offsets: &[Option<usize>; LANES], m: &[Word; LANES]) {
for (lane, off) in offsets.iter().enumerate() {
if let Some(off) = *off {
payloads[lane][off..off + BLOCK_BYTES].copy_from_slice(&m[lane].to_le_bytes());
}
}
}
fn block_offsets(g: &Group, index: usize) -> [Option<usize>; LANES] {
let mut o = [None; LANES];
for (slot, &blocks) in o.iter_mut().zip(g.blocks.iter()) {
if index < blocks {
*slot = Some(index * BLOCK_BYTES);
}
}
o
}
fn block_offsets_from_end(g: &Group, from_end: usize) -> [Option<usize>; LANES] {
let mut o = [None; LANES];
for (slot, &blocks) in o.iter_mut().zip(g.blocks.iter()) {
if from_end < blocks {
*slot = Some((blocks - 1 - from_end) * BLOCK_BYTES);
}
}
o
}
pub fn scramble_batch(cw: &ControlWord, payloads: &mut [&mut [u8]]) {
for group in payloads.chunks_mut(LANES) {
scramble_group(cw, group);
}
}
pub fn descramble_batch(cw: &ControlWord, payloads: &mut [&mut [u8]]) {
for group in payloads.chunks_mut(LANES) {
descramble_group(cw, group);
}
}
fn scramble_group(cw: &ControlWord, payloads: &mut [&mut [u8]]) {
let g = Group::new(payloads);
if g.max_blocks == 0 {
return;
}
let bc = BitslicedBlock::new(cw.expand_block());
let mut m = [0 as Word; LANES];
for from_end in 0..g.max_blocks {
let here = block_offsets_from_end(&g, from_end);
if from_end > 0 {
let next = block_offsets_from_end(&g, from_end - 1);
for lane in 0..LANES {
if let (Some(h), Some(n)) = (here[lane], next[lane]) {
let following: [u8; BLOCK_BYTES] = payloads[lane][n..n + BLOCK_BYTES]
.try_into()
.expect("slice is exactly one block");
for (dst, src) in payloads[lane][h..h + BLOCK_BYTES].iter_mut().zip(following) {
*dst ^= src;
}
}
}
}
gather(payloads, &here, &mut m);
transpose(&mut m);
bc.encrypt(&mut m);
transpose(&mut m);
scatter(payloads, &here, &m);
}
stream_xor(cw, payloads, &g);
}
fn descramble_group(cw: &ControlWord, payloads: &mut [&mut [u8]]) {
let g = Group::new(payloads);
if g.max_blocks == 0 {
return;
}
let bc = BitslicedBlock::new(cw.expand_block());
stream_xor(cw, payloads, &g);
let mut m = [0 as Word; LANES];
let mut cipher = [0 as Word; LANES];
for index in 0..g.max_blocks {
let here = block_offsets(&g, index);
gather(payloads, &here, &mut m);
cipher.copy_from_slice(&m);
transpose(&mut m);
bc.decrypt(&mut m);
transpose(&mut m);
if index > 0 {
let prev = block_offsets(&g, index - 1);
for lane in 0..LANES {
if let (Some(p), Some(_)) = (prev[lane], here[lane]) {
let c = cipher[lane].to_le_bytes();
for (dst, src) in payloads[lane][p..p + BLOCK_BYTES].iter_mut().zip(c) {
*dst ^= src;
}
}
}
}
scatter(payloads, &here, &m);
}
}
fn stream_xor(cw: &ControlWord, payloads: &mut [&mut [u8]], g: &Group) {
if g.max_stream == 0 {
return;
}
let mut iv = [0 as Word; LANES];
let mut first = [None; LANES];
for (slot, &len) in first.iter_mut().zip(g.len.iter()) {
if len >= BLOCK_BYTES {
*slot = Some(0);
}
}
gather(payloads, &first, &mut iv);
transpose(&mut iv);
let mut sc = BitslicedStream::new(&cw.expand_stream(), &iv);
let mut ks = [0 as Word; LANES];
let mut done = 0;
while done < g.max_stream {
for byte in 0..BLOCK_BYTES {
let bits = sc.keystream_byte();
for bit in 0..BITS_PER_BYTE {
ks[byte * BITS_PER_BYTE + bit] = bits[bit];
}
}
transpose(&mut ks);
for lane in 0..LANES {
let base = BLOCK_BYTES + done;
if base >= g.len[lane] {
continue;
}
let bytes = ks[lane].to_le_bytes();
let n = (g.len[lane] - base).min(BLOCK_BYTES);
for (j, b) in bytes.iter().take(n).enumerate() {
payloads[lane][base + j] ^= b;
}
}
done += BLOCK_BYTES;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn transpose_is_an_involution() {
let mut m = [0 as Word; LANES];
for (i, w) in m.iter_mut().enumerate() {
*w = (i as Word).wrapping_mul(0x9E37_79B9_7F4A_7C15);
}
let original = m;
transpose(&mut m);
assert_ne!(
m, original,
"transpose of an asymmetric matrix is not a no-op"
);
transpose(&mut m);
assert_eq!(m, original);
}
#[test]
fn transpose_moves_row_bits_to_column_bits() {
let mut m = [0 as Word; LANES];
m[3] = 1 << 5;
transpose(&mut m);
for (i, w) in m.iter().enumerate() {
let want: Word = if i == 5 { 1 << 3 } else { 0 };
assert_eq!(*w, want, "row {i}");
}
}
#[test]
fn short_payloads_pass_through() {
let cw = ControlWord::from_bytes([1, 2, 3, 4, 5, 6, 7, 8]);
let mut a = [0xAAu8; 7];
let mut b = [0xBBu8; 0];
let mut batch: [&mut [u8]; 2] = [&mut a, &mut b];
scramble_batch(&cw, &mut batch);
descramble_batch(&cw, &mut batch);
assert_eq!(a, [0xAAu8; 7]);
}
#[test]
fn empty_batch_is_a_no_op() {
let cw = ControlWord::from_bytes([1, 2, 3, 4, 5, 6, 7, 8]);
scramble_batch(&cw, &mut []);
descramble_batch(&cw, &mut []);
}
}