Fix algorithm and add quarter-round test #1

Merged
neemek merged 1 commit from better into main 2025-09-09 21:18:31 +00:00
Showing only changes of commit e6f0a1e7b5 - Show all commits

View file

@ -41,32 +41,16 @@ fn main() {
key[i] = u32::from_le_bytes(key_bytes[i * 4..(i + 1) * 4].try_into().unwrap()); key[i] = u32::from_le_bytes(key_bytes[i * 4..(i + 1) * 4].try_into().unwrap());
} }
let mut cipher = make_chacha_block(&key, 0, args.nonce);
let mut output = Vec::new();
let mut data_chunks = input.chunks_exact(64);
for c in &mut data_chunks {
cipher = chacha_rounds(cipher, args.rounds);
let data = u32x16::from(many_u8_to_few_u32(c.try_into().unwrap()));
cipher ^= data;
output.push(cipher);
}
let mut bytes: Vec<u8> = Vec::new(); let mut bytes: Vec<u8> = Vec::new();
for d in output { for (i, c) in input.chunks(64).enumerate() {
bytes.extend_from_slice(few_u32_to_many_u8(d.as_array()).as_slice()); let mut x = [0u8; 64];
} x[..c.len()].copy_from_slice(c);
let remainder = data_chunks.remainder(); let cipher = chacha_block(&key, i as u64, args.nonce, args.rounds);
let remaining = remainder.len(); let data = u32x16::from(many_u8_to_few_u32(x));
if remaining != 0 { let new = cipher ^ data;
let mut padded = [0u8; 64]; bytes.extend_from_slice(&few_u32_to_many_u8(new.as_array())[..c.len()]);
padded[0..remainder.len()].copy_from_slice(remainder);
cipher = chacha_rounds(cipher, 20);
let data = u32x16::from(many_u8_to_few_u32(padded));
cipher ^= data;
bytes.extend_from_slice(&few_u32_to_many_u8(cipher.as_array())[..remaining]);
} }
let stdout = stdout(); let stdout = stdout();
@ -74,6 +58,7 @@ fn main() {
out.write_all(bytes.as_slice()).unwrap(); out.write_all(bytes.as_slice()).unwrap();
} }
#[inline]
fn many_u8_to_few_u32(data: [u8; 64]) -> [u32; 16] { fn many_u8_to_few_u32(data: [u8; 64]) -> [u32; 16] {
let mut out = [0u32; 16]; let mut out = [0u32; 16];
@ -84,6 +69,7 @@ fn many_u8_to_few_u32(data: [u8; 64]) -> [u32; 16] {
out out
} }
#[inline]
fn few_u32_to_many_u8(data: &[u32; 16]) -> [u8; 64] { fn few_u32_to_many_u8(data: &[u32; 16]) -> [u8; 64] {
let mut out = [0u8; 64]; let mut out = [0u8; 64];
@ -99,59 +85,78 @@ fn make_chacha_block(key: &[u32; 8], counter: u64, nonce: u64) -> u32x16 {
let [k1, k2, k3, k4, k5, k6, k7, k8] = key; let [k1, k2, k3, k4, k5, k6, k7, k8] = key;
u32x16::from([ u32x16::from([
u32::from_ne_bytes(*b"expa"), u32::from_ne_bytes(*b"expa"), u32::from_ne_bytes(*b"nd 3"), u32::from_ne_bytes(*b"2-by"), u32::from_ne_bytes(*b"te k"),
u32::from_ne_bytes(*b"nd 3"), *k1, *k2, *k3, *k4,
u32::from_ne_bytes(*b"2-by"), *k5, *k6, *k7, *k8,
u32::from_ne_bytes(*b"te k"), (counter >> 32) as u32, (counter & 0xFFFFFFFF) as u32, (nonce >> 32) as u32, (nonce & 0xFFFFFFFF) as u32
*k1, *k2, *k3, *k4, *k5, *k6, *k7, *k8,
(counter >> 32) as u32, (counter & 0xFFFFFFFF) as u32,
(nonce >> 32) as u32, (nonce & 0xFFFFFFFF) as u32
]) ])
} }
fn chacha_rounds(input: u32x16, rounds: usize) -> u32x16 { fn chacha_rounds(input: u32x16, rounds: usize) -> u32x16 {
let x: u32x16 = input; let mut x: u32x16 = input;
for i in 0..rounds {
match i & 1 {
// odd rounds
0 => {
chacha_quarter_round(x, 0, 4, 8, 12);
chacha_quarter_round(x, 1, 5, 9, 13);
chacha_quarter_round(x, 2, 6, 10, 14);
chacha_quarter_round(x, 3, 7, 11, 15);
}
// even rounds
1 => {
chacha_quarter_round(x, 0, 5, 10, 15);
chacha_quarter_round(x, 1, 6, 11, 12);
chacha_quarter_round(x, 2, 7, 8, 13);
chacha_quarter_round(x, 3, 4, 9, 14);
}
_ => ()
}
for _ in (0..rounds).step_by(2) {
// odd rounds
chacha_quarter_round(&mut x, 0, 4, 8, 12);
chacha_quarter_round(&mut x, 1, 5, 9, 13);
chacha_quarter_round(&mut x, 2, 6, 10, 14);
chacha_quarter_round(&mut x, 3, 7, 11, 15);
// even rounds
chacha_quarter_round(&mut x, 0, 5, 10, 15);
chacha_quarter_round(&mut x, 1, 6, 11, 12);
chacha_quarter_round(&mut x, 2, 7, 8, 13);
chacha_quarter_round(&mut x, 3, 4, 9, 14);
} }
x + input x + input
} }
#[inline] #[inline]
fn chacha_quarter_round(mut x: u32x16, a: usize, b: usize, c: usize, d: usize) { fn chacha_block(key: &[u32; 8], counter: u64, nonce: u64, rounds: usize) -> u32x16 {
x[a] += x[b]; let block = make_chacha_block(key, counter, nonce);
chacha_rounds(block, rounds)
}
#[inline]
fn chacha_quarter_round(x: &mut u32x16, a: usize, b: usize, c: usize, d: usize) {
x[a] = x[a].wrapping_add(x[b]);
x[d] ^= x[a]; x[d] ^= x[a];
x[d] = x[d].rotate_left(16); x[d] = x[d].rotate_left(16);
x[c] += x[d]; x[c] = x[c].wrapping_add(x[d]);
x[b] ^= x[c]; x[b] ^= x[c];
x[b] = x[b].rotate_left(12); x[b] = x[b].rotate_left(12);
x[a] += x[b]; x[a] = x[a].wrapping_add(x[b]);
x[d] ^= x[a]; x[d] ^= x[a];
x[d] = x[d].rotate_left(8); x[d] = x[d].rotate_left(8);
x[c] += x[d]; x[c] = x[c].wrapping_add(x[d]);
x[b] ^= x[c]; x[b] ^= x[c];
x[b] = x[b].rotate_left(7); x[b] = x[b].rotate_left(7);
} }
#[cfg(test)]
mod tests {
use std::simd::u32x16;
use crate::chacha_quarter_round;
#[test]
fn test_chacha_quarter_round() {
let mut data = [0u32; 16];
data[0] = 0x11111111;
data[1] = 0x01020304;
data[2] = 0x9b8d6f43;
data[3] = 0x01234567;
let mut arr = u32x16::from(data);
chacha_quarter_round(&mut arr, 0, 1, 2, 3);
assert_eq!(arr[0], 0xea2a92f4);
assert_eq!(arr[1], 0xcb1cf8ce);
assert_eq!(arr[2], 0x4581472e);
assert_eq!(arr[3], 0x5881c4bb);
}
}