use core::arch::x86_64::*;
pub const LANES: usize = 8;
pub const STRIDE: usize = LANES * 64;
#[inline(always)]
fn rotl(x: __m256i, n: i32) -> __m256i {
match n {
16 => unsafe {
_mm256_shuffle_epi8(
x,
_mm256_set_epi8(
13, 12, 15, 14, 9, 8, 11, 10, 5, 4, 7, 6, 1, 0, 3, 2, 13, 12, 15, 14, 9, 8, 11,
10, 5, 4, 7, 6, 1, 0, 3, 2,
),
)
},
8 => unsafe {
_mm256_shuffle_epi8(
x,
_mm256_set_epi8(
14, 13, 12, 15, 10, 9, 8, 11, 6, 5, 4, 7, 2, 1, 0, 3, 14, 13, 12, 15, 10, 9, 8,
11, 6, 5, 4, 7, 2, 1, 0, 3,
),
)
},
12 => unsafe { _mm256_or_si256(_mm256_slli_epi32(x, 12), _mm256_srli_epi32(x, 20)) },
7 => unsafe { _mm256_or_si256(_mm256_slli_epi32(x, 7), _mm256_srli_epi32(x, 25)) },
_ => unreachable!(),
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn eight_blocks(key: &[u8; 32], nonce: &[u8; 12], counter: u32, data: &mut [u8]) {
debug_assert_eq!(data.len(), STRIDE);
let k = |i: usize| -> u32 {
u32::from_le_bytes([key[i * 4], key[i * 4 + 1], key[i * 4 + 2], key[i * 4 + 3]])
};
let n = |i: usize| -> u32 {
u32::from_le_bytes([
nonce[i * 4],
nonce[i * 4 + 1],
nonce[i * 4 + 2],
nonce[i * 4 + 3],
])
};
let splat = |v: u32| _mm256_set1_epi32(v as i32);
let start: [__m256i; 16] = [
splat(0x6170_7865),
splat(0x3320_646e),
splat(0x7962_2d32),
splat(0x6b20_6574),
splat(k(0)),
splat(k(1)),
splat(k(2)),
splat(k(3)),
splat(k(4)),
splat(k(5)),
splat(k(6)),
splat(k(7)),
{
_mm256_setr_epi32(
counter as i32,
counter.wrapping_add(1) as i32,
counter.wrapping_add(2) as i32,
counter.wrapping_add(3) as i32,
counter.wrapping_add(4) as i32,
counter.wrapping_add(5) as i32,
counter.wrapping_add(6) as i32,
counter.wrapping_add(7) as i32,
)
},
splat(n(0)),
splat(n(1)),
splat(n(2)),
];
let mut v = start;
macro_rules! quarter {
($a:expr, $b:expr, $c:expr, $d:expr) => {{
v[$a] = _mm256_add_epi32(v[$a], v[$b]);
v[$d] = rotl(_mm256_xor_si256(v[$d], v[$a]), 16);
v[$c] = _mm256_add_epi32(v[$c], v[$d]);
v[$b] = rotl(_mm256_xor_si256(v[$b], v[$c]), 12);
v[$a] = _mm256_add_epi32(v[$a], v[$b]);
v[$d] = rotl(_mm256_xor_si256(v[$d], v[$a]), 8);
v[$c] = _mm256_add_epi32(v[$c], v[$d]);
v[$b] = rotl(_mm256_xor_si256(v[$b], v[$c]), 7);
}};
}
for _ in 0..10 {
quarter!(0, 4, 8, 12);
quarter!(1, 5, 9, 13);
quarter!(2, 6, 10, 14);
quarter!(3, 7, 11, 15);
quarter!(0, 5, 10, 15);
quarter!(1, 6, 11, 12);
quarter!(2, 7, 8, 13);
quarter!(3, 4, 9, 14);
}
let mut words = [[0u32; LANES]; 16];
for i in 0..16 {
let sum = _mm256_add_epi32(v[i], start[i]);
unsafe { _mm256_storeu_si256(words[i].as_mut_ptr() as *mut __m256i, sum) };
}
for (lane, block) in data.chunks_exact_mut(64).enumerate() {
for (word, out) in block.chunks_exact_mut(4).enumerate() {
let ks = words[word][lane].to_le_bytes();
for (d, k) in out.iter_mut().zip(ks) {
*d ^= k;
}
}
}
}