use core::arch::x86_64::*;
use super::K512;
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn ror256<const R: i32, const L: i32>(x: __m256i) -> __m256i {
_mm256_or_si256(_mm256_srli_epi64::<R>(x), _mm256_slli_epi64::<L>(x))
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn ror128<const R: i32, const L: i32>(x: __m128i) -> __m128i {
_mm_or_si128(_mm_srli_epi64::<R>(x), _mm_slli_epi64::<L>(x))
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn sigma0(x: __m256i) -> __m256i {
_mm256_xor_si256(
_mm256_xor_si256(ror256::<1, 63>(x), ror256::<8, 56>(x)),
_mm256_srli_epi64::<7>(x),
)
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn sigma1_128(x: __m128i) -> __m128i {
_mm_xor_si128(
_mm_xor_si128(ror128::<19, 45>(x), ror128::<61, 3>(x)),
_mm_srli_epi64::<6>(x),
)
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn load_block(w: &mut [u64; 80], block: &[u8]) {
let bswap = _mm256_set_epi8(
8, 9, 10, 11, 12, 13, 14, 15, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 0, 1,
2, 3, 4, 5, 6, 7,
);
for i in 0..4 {
let v = _mm256_loadu_si256(block.as_ptr().add(i * 32) as *const __m256i);
_mm256_storeu_si256(
w.as_mut_ptr().add(i * 4) as *mut __m256i,
_mm256_shuffle_epi8(v, bswap),
);
}
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn extend4(w: &mut [u64; 80], at: usize) {
let w16 = _mm256_loadu_si256(w.as_ptr().add(at - 16) as *const __m256i);
let w15 = _mm256_loadu_si256(w.as_ptr().add(at - 15) as *const __m256i);
let w7 = _mm256_loadu_si256(w.as_ptr().add(at - 7) as *const __m256i);
let t = _mm256_add_epi64(_mm256_add_epi64(w16, sigma0(w15)), w7);
let lo = _mm256_castsi256_si128(t);
let w2 = _mm_loadu_si128(w.as_ptr().add(at - 2) as *const __m128i);
let res_lo = _mm_add_epi64(lo, sigma1_128(w2));
_mm_storeu_si128(w.as_mut_ptr().add(at) as *mut __m128i, res_lo);
let hi = _mm256_extracti128_si256::<1>(t);
let res_hi = _mm_add_epi64(hi, sigma1_128(res_lo));
_mm_storeu_si128(w.as_mut_ptr().add(at + 2) as *mut __m128i, res_hi);
}
#[cfg(test)]
#[target_feature(enable = "avx2")]
pub unsafe fn schedule(block: &[u8], w: &mut [u64; 80]) {
debug_assert_eq!(block.len(), 128);
load_block(w, block);
let mut i = 16;
while i < 80 {
extend4(w, i);
i += 4;
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn compress(h: &mut [u64; 8], w: &mut [u64; 80], block: &[u8]) {
debug_assert_eq!(block.len(), 128);
load_block(w, block);
let [mut a, mut b, mut c, mut d, mut e, mut f, mut g, mut hh] = *h;
macro_rules! round {
($i:expr) => {{
let s1 = e.rotate_right(14) ^ e.rotate_right(18) ^ e.rotate_right(41);
let ch = (e & f) ^ ((!e) & g);
let t1 = hh
.wrapping_add(s1)
.wrapping_add(ch)
.wrapping_add(K512[$i])
.wrapping_add(w[$i]);
let s0 = a.rotate_right(28) ^ a.rotate_right(34) ^ a.rotate_right(39);
let maj = (a & b) ^ (a & c) ^ (b & c);
let t2 = s0.wrapping_add(maj);
hh = g;
g = f;
f = e;
e = d.wrapping_add(t1);
d = c;
c = b;
b = a;
a = t1.wrapping_add(t2);
}};
}
let mut i = 0;
while i < 80 {
extend4(w, 16 + (i / 5) * 4);
round!(i);
round!(i + 1);
round!(i + 2);
round!(i + 3);
round!(i + 4);
i += 5;
}
h[0] = h[0].wrapping_add(a);
h[1] = h[1].wrapping_add(b);
h[2] = h[2].wrapping_add(c);
h[3] = h[3].wrapping_add(d);
h[4] = h[4].wrapping_add(e);
h[5] = h[5].wrapping_add(f);
h[6] = h[6].wrapping_add(g);
h[7] = h[7].wrapping_add(hh);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_rotations_use_complementary_shifts() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let vals: [u64; 4] = [
0x0123_4567_89ab_cdef,
0xffff_ffff_ffff_ffff,
1,
0x8000_0000_0000_0000,
];
unsafe {
let v = _mm256_loadu_si256(vals.as_ptr() as *const __m256i);
let mut out = [0u64; 4];
macro_rules! check256 {
($r:literal, $l:literal) => {{
_mm256_storeu_si256(out.as_mut_ptr() as *mut __m256i, ror256::<$r, $l>(v));
for (o, i) in out.iter().zip(vals.iter()) {
assert_eq!(*o, i.rotate_right($r), "ror256::<{}, {}>", $r, $l);
}
}};
}
check256!(1, 63);
check256!(8, 56);
let w = _mm_loadu_si128(vals.as_ptr() as *const __m128i);
let mut out2 = [0u64; 2];
macro_rules! check128 {
($r:literal, $l:literal) => {{
_mm_storeu_si128(out2.as_mut_ptr() as *mut __m128i, ror128::<$r, $l>(w));
for (o, i) in out2.iter().zip(vals.iter()) {
assert_eq!(*o, i.rotate_right($r), "ror128::<{}, {}>", $r, $l);
}
}};
}
check128!(19, 45);
check128!(61, 3);
}
}
#[test]
fn the_avx2_compress_agrees_with_the_portable_one() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let mut state = 0x5151_2345_9876_abcdu64;
let mut next = || {
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
state.wrapping_mul(0x2545_f491_4f6c_dd1d)
};
let iv = [
0x6a09_e667_f3bc_c908u64,
0xbb67_ae85_84ca_a73b,
0x3c6e_f372_fe94_f82b,
0xa54f_f53a_5f1d_36f1,
0x510e_527f_ade6_82d1,
0x9b05_688c_2b3e_6c1f,
0x1f83_d9ab_fb41_bd6b,
0x5be0_cd19_137e_2179,
];
let mut h_v = iv;
let mut h_s = iv;
let mut w_v = [0u64; 80];
let mut w_s = [0u64; 80];
for case in 0..300 {
let mut block = [0u8; 128];
match case {
0 => {}
1 => block = [0xff; 128],
_ => {
for chunk in block.chunks_exact_mut(8) {
chunk.copy_from_slice(&next().to_le_bytes());
}
}
}
unsafe { compress(&mut h_v, &mut w_v, &block) };
super::super::Core512::schedule(&mut w_s, &block);
super::super::Core512::rounds(&mut h_s, &w_s);
assert_eq!(h_v, h_s, "case {case}, chained");
}
}
#[test]
fn the_avx2_schedule_agrees_with_the_portable_one() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let mut state = 0xdead_beef_0bad_f00du64;
let mut next = || {
state ^= state >> 12;
state ^= state << 25;
state ^= state >> 27;
state.wrapping_mul(0x2545_f491_4f6c_dd1d)
};
for case in 0..500 {
let mut block = [0u8; 128];
match case {
0 => {}
1 => block = [0xff; 128],
2 => {
for (i, b) in block.iter_mut().enumerate() {
*b = i as u8;
}
}
_ => {
for chunk in block.chunks_exact_mut(8) {
chunk.copy_from_slice(&next().to_le_bytes());
}
}
}
let mut want = [0u64; 80];
super::super::Core512::schedule(&mut want, &block);
let mut got = [0u64; 80];
unsafe { schedule(&block, &mut got) };
assert_eq!(got, want, "case {case}");
}
}
}
#[cfg(test)]
mod bench {
use super::*;
use std::time::Instant;
#[test]
#[ignore = "diagnostic, not a test"]
fn schedule_ab() {
if !std::is_x86_feature_detected!("avx2") {
println!("no avx2 here");
return;
}
let mut block = [0u8; 128];
for (i, b) in block.iter_mut().enumerate() {
*b = (i as u8).wrapping_mul(31).wrapping_add(7);
}
let n = 100_000;
let (mut best_v, mut best_s) = (f64::INFINITY, f64::INFINITY);
let mut w = [0u64; 80];
let mut h = [1u64, 2, 3, 4, 5, 6, 7, 8];
for _ in 0..30 {
let t = Instant::now();
for _ in 0..n {
unsafe { compress(&mut h, &mut w, core::hint::black_box(&block)) };
}
best_v = best_v.min(t.elapsed().as_secs_f64() / n as f64 * 1e9);
let t = Instant::now();
for _ in 0..n {
super::super::Core512::schedule(&mut w, core::hint::black_box(&block));
super::super::Core512::rounds(&mut h, &w);
}
best_s = best_s.min(t.elapsed().as_secs_f64() / n as f64 * 1e9);
}
println!(
"
compress, avx2 interleaved {best_v:>8.1} ns/block"
);
println!(" compress, scalar {best_s:>8.1} ns/block");
println!(" ratio {:>8.2}x", best_s / best_v);
}
}