#[inline(always)]
fn quarter_round(state: &mut [u32], a: usize, b: usize, c: usize, d: usize) {
state[a] = state[a].wrapping_add(state[b]);
state[d] ^= state[a];
state[d] = state[d].rotate_left(16);
state[c] = state[c].wrapping_add(state[d]);
state[b] ^= state[c];
state[b] = state[b].rotate_left(12);
state[a] = state[a].wrapping_add(state[b]);
state[d] ^= state[a];
state[d] = state[d].rotate_left(8);
state[c] = state[c].wrapping_add(state[d]);
state[b] ^= state[c];
state[b] = state[b].rotate_left(7);
}
pub fn xor(
data: &mut [u8],
keystream_only: bool,
key: &[u32; 8],
counter: &[u32; 4],
double_rounds: usize,
) {
debug_assert!(data.len() <= 64, "Data length must not exceed 64 bytes");
if keystream_only {
debug_assert_eq!(data.len(), 64);
}
let state = [
0x61707865, 0x3320646e, 0x79622d32, 0x6b206574, key[0], key[1], key[2], key[3], key[4], key[5], key[6], key[7], counter[0], counter[1], counter[2], counter[3], ];
let mut working = state;
for _ in 0..double_rounds {
quarter_round(&mut working, 0, 4, 8, 12);
quarter_round(&mut working, 1, 5, 9, 13);
quarter_round(&mut working, 2, 6, 10, 14);
quarter_round(&mut working, 3, 7, 11, 15);
quarter_round(&mut working, 0, 5, 10, 15);
quarter_round(&mut working, 1, 6, 11, 12);
quarter_round(&mut working, 2, 7, 8, 13);
quarter_round(&mut working, 3, 4, 9, 14);
}
if keystream_only {
let mut i = 0;
while i < data.len() {
let word = working[i / 4].wrapping_add(state[i / 4]).to_le_bytes();
let j = i % 4;
data[i] = word[j];
i += 1;
}
} else {
let mut i = 0;
while i < data.len() {
let word = working[i / 4].wrapping_add(state[i / 4]).to_le_bytes();
let j = i % 4;
data[i] ^= word[j];
i += 1;
}
}
}
#[cfg(test)]
mod tests {
use crate::fallback_chacha20::xor;
#[test]
fn test_chacha20_xor_encrypt_decrypt() {
use core::slice;
let key = [0u8; 32];
let nonce = [0u8; 12];
let mut key_words = [0u32; 8];
for i in 0..8 {
key_words[i] =
u32::from_le_bytes([key[i * 4], key[i * 4 + 1], key[i * 4 + 2], key[i * 4 + 3]]);
}
let mut counter = [0u32; 4];
for i in 0..3 {
counter[i + 1] = u32::from_le_bytes([
nonce[i * 4],
nonce[i * 4 + 1],
nonce[i * 4 + 2],
nonce[i * 4 + 3],
]);
}
let plaintext = b"Hello, ChaCha20 fallback test!";
let mut buffer = [0u8; 32];
buffer[..plaintext.len()].copy_from_slice(plaintext);
xor(&mut buffer[..plaintext.len()], false, &key_words, &counter, 10);
let mut ciphertext = [0u8; 64];
ciphertext[..plaintext.len()].copy_from_slice(&buffer[..plaintext.len()]);
xor(&mut buffer[..plaintext.len()], false, &key_words, &counter, 10);
assert_eq!(&buffer[..plaintext.len()], plaintext);
assert_ne!(&ciphertext[..plaintext.len()], plaintext);
}
}