#[cfg(any(
test,
not(target_arch = "x86_64"),
feature = "std",
not(target_feature = "avx512f")
))]
use core::array;
#[cfg(target_arch = "x86_64")]
mod row_x86;
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[path = "blake3_schedule/wasm32_simd128.rs"]
mod wasm32_simd128;
pub(super) const IV: [u32; 8] = [
0x6a09_e667,
0xbb67_ae85,
0x3c6e_f372,
0xa54f_f53a,
0x510e_527f,
0x9b05_688c,
0x1f83_d9ab,
0x5be0_cd19,
];
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
const ROUNDS: usize = 7;
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
const MSG_SCHEDULE: [[usize; 16]; ROUNDS] = [
[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15],
[2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8],
[3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1],
[10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6],
[12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4],
[9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7],
[11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13],
];
pub(super) const PACKED_LANES: usize = 16;
#[cfg(any(
test,
all(target_arch = "x86_64", feature = "std"),
all(target_arch = "x86_64", not(target_feature = "avx512f")),
all(target_arch = "aarch64", target_feature = "neon"),
all(target_arch = "wasm32", target_feature = "simd128"),
))]
#[inline]
fn compress_via_sub_batches<const W: usize>(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
f: impl Fn([[u32; W]; 8], [[u32; W]; 16]) -> [[u32; W]; 8],
) -> [[u32; PACKED_LANES]; 8] {
let mut out = [[0u32; PACKED_LANES]; 8];
for chunk in 0..(PACKED_LANES / W) {
let base = chunk * W;
let sub_cv: [[u32; W]; 8] =
array::from_fn(|word| array::from_fn(|lane| cv[word][base + lane]));
let sub_block: [[u32; W]; 16] =
array::from_fn(|word| array::from_fn(|lane| block[word][base + lane]));
let sub_out = f(sub_cv, sub_block);
for (word, sub_word) in out.iter_mut().zip(sub_out) {
word[base..base + W].copy_from_slice(&sub_word);
}
}
out
}
#[cfg(all(target_arch = "x86_64", feature = "std"))]
pub(crate) mod cpu {
#[inline]
pub(crate) fn has_sse41() -> bool {
std::is_x86_feature_detected!("sse4.1")
}
#[inline]
pub(crate) fn has_avx512f() -> bool {
std::is_x86_feature_detected!("avx512f")
}
#[inline]
pub(crate) fn has_avx512vl() -> bool {
has_avx512f() && std::is_x86_feature_detected!("avx512vl")
}
#[inline]
pub(crate) fn has_avx2() -> bool {
std::is_x86_feature_detected!("avx2")
}
}
#[cfg(all(target_arch = "x86_64", feature = "std"))]
mod native_backend {
use super::{
PACKED_LANES, compress_via_sub_batches, cpu, x86_64_avx2, x86_64_avx512, x86_64_sse2,
};
#[inline]
pub(super) fn compress(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
if cpu::has_avx512f() {
unsafe { x86_64_avx512::compress_packed_16(*cv, *block) }
} else if cpu::has_avx2() {
compress_via_sub_batches::<8>(cv, block, |cv, block| {
unsafe { x86_64_avx2::compress_packed_8(cv, block) }
})
} else {
compress_via_sub_batches::<4>(cv, block, |cv, block| {
unsafe { x86_64_sse2::compress_packed_4(cv, block) }
})
}
}
}
#[cfg(all(target_arch = "x86_64", not(feature = "std"), target_feature = "avx512f"))]
mod native_backend {
use super::PACKED_LANES;
#[inline(always)]
pub(super) fn compress(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
unsafe { super::x86_64_avx512::compress_packed_16(*cv, *block) }
}
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
target_feature = "avx2",
not(target_feature = "avx512f")
))]
mod native_backend {
use super::{PACKED_LANES, compress_via_sub_batches};
#[inline(always)]
pub(super) fn compress(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
compress_via_sub_batches::<8>(cv, block, |cv, block| {
unsafe { super::x86_64_avx2::compress_packed_8(cv, block) }
})
}
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
not(target_feature = "avx2"),
not(target_feature = "avx512f")
))]
mod native_backend {
use super::{PACKED_LANES, compress_via_sub_batches};
#[inline(always)]
pub(super) fn compress(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
compress_via_sub_batches::<4>(cv, block, |cv, block| {
unsafe { super::x86_64_sse2::compress_packed_4(cv, block) }
})
}
}
#[cfg(not(target_arch = "x86_64"))]
mod native_backend {
use super::PACKED_LANES;
#[inline(always)]
pub(super) fn compress(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
super::compress_via_sub_batches::<4>(
cv,
block,
super::wasm32_simd128::compress_packed_4,
)
}
#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
{
super::compress_via_sub_batches::<4>(cv, block, super::neon::compress_packed_4)
}
#[cfg(not(any(
all(target_arch = "wasm32", target_feature = "simd128"),
all(target_arch = "aarch64", target_feature = "neon")
)))]
{
super::compress_packed(*cv, *block)
}
}
}
#[cfg(all(target_arch = "x86_64", feature = "std"))]
mod row_dispatch {
use super::{cpu, row_x86};
#[inline]
pub(super) fn compress_raw(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 8] {
if cpu::has_avx512vl() {
unsafe { row_x86::compress_raw_avx512vl(cv, block) }
} else if cpu::has_sse41() {
unsafe { row_x86::compress_raw_sse41(cv, block) }
} else {
row_x86::compress_raw(cv, block)
}
}
#[inline]
pub(super) fn compress_raw_xof(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 16] {
if cpu::has_avx512vl() {
unsafe { row_x86::compress_raw_xof_avx512vl(cv, block) }
} else if cpu::has_sse41() {
unsafe { row_x86::compress_raw_xof_sse41(cv, block) }
} else {
row_x86::compress_raw_xof(cv, block)
}
}
}
#[inline(always)]
#[cfg(any(test, not(target_arch = "x86_64")))]
fn g(v: &mut [u32; 16], a: usize, b: usize, c: usize, d: usize, x: u32, y: u32) {
v[a] = v[a].wrapping_add(v[b]).wrapping_add(x);
v[d] = (v[d] ^ v[a]).rotate_right(16);
v[c] = v[c].wrapping_add(v[d]);
v[b] = (v[b] ^ v[c]).rotate_right(12);
v[a] = v[a].wrapping_add(v[b]).wrapping_add(y);
v[d] = (v[d] ^ v[a]).rotate_right(8);
v[c] = v[c].wrapping_add(v[d]);
v[b] = (v[b] ^ v[c]).rotate_right(7);
}
#[inline(always)]
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
fn add_packed<const LANES: usize>(a: [u32; LANES], b: [u32; LANES]) -> [u32; LANES] {
array::from_fn(|i| a[i].wrapping_add(b[i]))
}
#[inline(always)]
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
fn xor_packed<const LANES: usize>(a: [u32; LANES], b: [u32; LANES]) -> [u32; LANES] {
array::from_fn(|i| a[i] ^ b[i])
}
#[inline(always)]
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
fn rotr_packed<const LANES: usize>(a: [u32; LANES], n: u32) -> [u32; LANES] {
array::from_fn(|i| a[i].rotate_right(n))
}
#[inline(always)]
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
fn g_packed<const LANES: usize>(
v: &mut [[u32; LANES]; 16],
a: usize,
b: usize,
c: usize,
d: usize,
x: [u32; LANES],
y: [u32; LANES],
) {
v[a] = add_packed(add_packed(v[a], v[b]), x);
v[d] = rotr_packed(xor_packed(v[d], v[a]), 16);
v[c] = add_packed(v[c], v[d]);
v[b] = rotr_packed(xor_packed(v[b], v[c]), 12);
v[a] = add_packed(add_packed(v[a], v[b]), y);
v[d] = rotr_packed(xor_packed(v[d], v[a]), 8);
v[c] = add_packed(v[c], v[d]);
v[b] = rotr_packed(xor_packed(v[b], v[c]), 7);
}
#[inline(always)]
#[cfg(any(test, not(target_arch = "x86_64")))]
fn permuted_state_with_parameter_words(
cv: [u32; 8],
block: [u32; 16],
parameter_words: [u32; 4],
) -> [u32; 16] {
let mut v = [0u32; 16];
v[..8].copy_from_slice(&cv);
v[8..12].copy_from_slice(&IV[..4]);
v[12..16].copy_from_slice(¶meter_words);
macro_rules! round {
($($s:literal),*) => {{
let s = [$($s),*];
g(&mut v, 0, 4, 8, 12, block[s[0]], block[s[1]]);
g(&mut v, 1, 5, 9, 13, block[s[2]], block[s[3]]);
g(&mut v, 2, 6, 10, 14, block[s[4]], block[s[5]]);
g(&mut v, 3, 7, 11, 15, block[s[6]], block[s[7]]);
g(&mut v, 0, 5, 10, 15, block[s[8]], block[s[9]]);
g(&mut v, 1, 6, 11, 12, block[s[10]], block[s[11]]);
g(&mut v, 2, 7, 8, 13, block[s[12]], block[s[13]]);
g(&mut v, 3, 4, 9, 14, block[s[14]], block[s[15]]);
}};
}
round!(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
round!(2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8);
round!(3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1);
round!(10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6);
round!(12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4);
round!(9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7);
round!(11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13);
v
}
pub(super) fn compress_raw(cv: [u32; 8], block: [u32; 16]) -> [u32; 8] {
#[cfg(all(target_arch = "x86_64", feature = "std"))]
{
row_dispatch::compress_raw(&cv, &block)
}
#[cfg(all(target_arch = "x86_64", not(feature = "std"), target_feature = "avx512vl"))]
{
unsafe { row_x86::compress_raw_avx512vl(&cv, &block) }
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
target_feature = "sse4.1",
not(target_feature = "avx512vl")
))]
{
unsafe { row_x86::compress_raw_sse41(&cv, &block) }
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
not(target_feature = "sse4.1"),
not(target_feature = "avx512vl")
))]
{
row_x86::compress_raw(&cv, &block)
}
#[cfg(not(target_arch = "x86_64"))]
{
let v = permuted_state_with_parameter_words(cv, block, [IV[4], IV[5], IV[6], IV[7]]);
array::from_fn(|i| v[i] ^ v[i + 8])
}
}
pub(super) fn compress_raw_xof(cv: [u32; 8], block: [u32; 16]) -> [u32; 16] {
#[cfg(all(target_arch = "x86_64", feature = "std"))]
{
row_dispatch::compress_raw_xof(&cv, &block)
}
#[cfg(all(target_arch = "x86_64", not(feature = "std"), target_feature = "avx512vl"))]
{
unsafe { row_x86::compress_raw_xof_avx512vl(&cv, &block) }
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
target_feature = "sse4.1",
not(target_feature = "avx512vl")
))]
{
unsafe { row_x86::compress_raw_xof_sse41(&cv, &block) }
}
#[cfg(all(
target_arch = "x86_64",
not(feature = "std"),
not(target_feature = "sse4.1"),
not(target_feature = "avx512vl")
))]
{
row_x86::compress_raw_xof(&cv, &block)
}
#[cfg(not(target_arch = "x86_64"))]
{
let v = permuted_state_with_parameter_words(cv, block, [IV[4], IV[5], IV[6], IV[7]]);
array::from_fn(|i| if i < 8 { v[i] ^ v[i + 8] } else { v[i] ^ cv[i - 8] })
}
}
#[cfg(test)]
pub(super) fn compress_raw_with_parameter_words(
cv: [u32; 8],
block: [u32; 16],
parameter_words: [u32; 4],
) -> [u32; 8] {
let v = permuted_state_with_parameter_words(cv, block, parameter_words);
array::from_fn(|i| v[i] ^ v[i + 8])
}
#[cfg(test)]
pub(super) fn compress_raw_xof_with_parameter_words(
cv: [u32; 8],
block: [u32; 16],
parameter_words: [u32; 4],
) -> [u32; 16] {
let v = permuted_state_with_parameter_words(cv, block, parameter_words);
array::from_fn(|i| if i < 8 { v[i] ^ v[i + 8] } else { v[i] ^ cv[i - 8] })
}
#[cfg(any(
test,
all(target_arch = "aarch64", not(target_feature = "neon")),
all(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
not(all(target_arch = "wasm32", target_feature = "simd128")),
),
))]
pub(super) fn compress_packed<const LANES: usize>(
cv: [[u32; LANES]; 8],
block: [[u32; LANES]; 16],
) -> [[u32; LANES]; 8] {
let mut v = [[0u32; LANES]; 16];
v[..8].copy_from_slice(&cv);
for i in 0..8 {
v[8 + i] = [IV[i]; LANES];
}
for s in MSG_SCHEDULE.iter() {
g_packed(&mut v, 0, 4, 8, 12, block[s[0]], block[s[1]]);
g_packed(&mut v, 1, 5, 9, 13, block[s[2]], block[s[3]]);
g_packed(&mut v, 2, 6, 10, 14, block[s[4]], block[s[5]]);
g_packed(&mut v, 3, 7, 11, 15, block[s[6]], block[s[7]]);
g_packed(&mut v, 0, 5, 10, 15, block[s[8]], block[s[9]]);
g_packed(&mut v, 1, 6, 11, 12, block[s[10]], block[s[11]]);
g_packed(&mut v, 2, 7, 8, 13, block[s[12]], block[s[13]]);
g_packed(&mut v, 3, 4, 9, 14, block[s[14]], block[s[15]]);
}
array::from_fn(|i| xor_packed(v[i], v[i + 8]))
}
#[inline]
pub(super) fn compress_packed_native(
cv: &[[u32; PACKED_LANES]; 8],
block: &[[u32; PACKED_LANES]; 16],
) -> [[u32; PACKED_LANES]; 8] {
native_backend::compress(cv, block)
}
#[cfg(target_arch = "x86_64")]
#[rustfmt::skip]
macro_rules! define_x86_packed_compress {
($(#[$attr:meta])* $name:ident, $lanes:literal) => {
$(#[$attr])*
#[inline]
pub(super) unsafe fn $name(
cv: [[u32; $lanes]; 8],
block: [[u32; $lanes]; 16],
) -> [[u32; $lanes]; 8] {
let mut v0 = load(&cv[0]);
let mut v1 = load(&cv[1]);
let mut v2 = load(&cv[2]);
let mut v3 = load(&cv[3]);
let mut v4 = load(&cv[4]);
let mut v5 = load(&cv[5]);
let mut v6 = load(&cv[6]);
let mut v7 = load(&cv[7]);
let mut v8 = splat(IV[0]);
let mut v9 = splat(IV[1]);
let mut v10 = splat(IV[2]);
let mut v11 = splat(IV[3]);
let mut v12 = splat(IV[4]);
let mut v13 = splat(IV[5]);
let mut v14 = splat(IV[6]);
let mut v15 = splat(IV[7]);
macro_rules! round {
(
$m0:literal,
$m1:literal,
$m2:literal,
$m3:literal,
$m4:literal,
$m5:literal,
$m6:literal,
$m7:literal,
$m8:literal,
$m9:literal,
$m10:literal,
$m11:literal,
$m12:literal,
$m13:literal,
$m14:literal,
$m15:literal
) => {{
let m0 = load(&block[$m0]);
let m1 = load(&block[$m1]);
let m2 = load(&block[$m2]);
let m3 = load(&block[$m3]);
let m4 = load(&block[$m4]);
let m5 = load(&block[$m5]);
let m6 = load(&block[$m6]);
let m7 = load(&block[$m7]);
let m8 = load(&block[$m8]);
let m9 = load(&block[$m9]);
let m10 = load(&block[$m10]);
let m11 = load(&block[$m11]);
let m12 = load(&block[$m12]);
let m13 = load(&block[$m13]);
let m14 = load(&block[$m14]);
let m15 = load(&block[$m15]);
v0 = add(add(v0, v4), m0);
v1 = add(add(v1, v5), m2);
v2 = add(add(v2, v6), m4);
v3 = add(add(v3, v7), m6);
v12 = rotr16(xor(v12, v0));
v13 = rotr16(xor(v13, v1));
v14 = rotr16(xor(v14, v2));
v15 = rotr16(xor(v15, v3));
v8 = add(v8, v12);
v9 = add(v9, v13);
v10 = add(v10, v14);
v11 = add(v11, v15);
v4 = rotr12(xor(v4, v8));
v5 = rotr12(xor(v5, v9));
v6 = rotr12(xor(v6, v10));
v7 = rotr12(xor(v7, v11));
v0 = add(add(v0, v4), m1);
v1 = add(add(v1, v5), m3);
v2 = add(add(v2, v6), m5);
v3 = add(add(v3, v7), m7);
v12 = rotr8(xor(v12, v0));
v13 = rotr8(xor(v13, v1));
v14 = rotr8(xor(v14, v2));
v15 = rotr8(xor(v15, v3));
v8 = add(v8, v12);
v9 = add(v9, v13);
v10 = add(v10, v14);
v11 = add(v11, v15);
v4 = rotr7(xor(v4, v8));
v5 = rotr7(xor(v5, v9));
v6 = rotr7(xor(v6, v10));
v7 = rotr7(xor(v7, v11));
v0 = add(add(v0, v5), m8);
v1 = add(add(v1, v6), m10);
v2 = add(add(v2, v7), m12);
v3 = add(add(v3, v4), m14);
v15 = rotr16(xor(v15, v0));
v12 = rotr16(xor(v12, v1));
v13 = rotr16(xor(v13, v2));
v14 = rotr16(xor(v14, v3));
v10 = add(v10, v15);
v11 = add(v11, v12);
v8 = add(v8, v13);
v9 = add(v9, v14);
v5 = rotr12(xor(v5, v10));
v6 = rotr12(xor(v6, v11));
v7 = rotr12(xor(v7, v8));
v4 = rotr12(xor(v4, v9));
v0 = add(add(v0, v5), m9);
v1 = add(add(v1, v6), m11);
v2 = add(add(v2, v7), m13);
v3 = add(add(v3, v4), m15);
v15 = rotr8(xor(v15, v0));
v12 = rotr8(xor(v12, v1));
v13 = rotr8(xor(v13, v2));
v14 = rotr8(xor(v14, v3));
v10 = add(v10, v15);
v11 = add(v11, v12);
v8 = add(v8, v13);
v9 = add(v9, v14);
v5 = rotr7(xor(v5, v10));
v6 = rotr7(xor(v6, v11));
v7 = rotr7(xor(v7, v8));
v4 = rotr7(xor(v4, v9));
}};
}
round!(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
round!(2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8);
round!(3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1);
round!(10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6);
round!(12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4);
round!(9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7);
round!(11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13);
[
store(xor(v0, v8)),
store(xor(v1, v9)),
store(xor(v2, v10)),
store(xor(v3, v11)),
store(xor(v4, v12)),
store(xor(v5, v13)),
store(xor(v6, v14)),
store(xor(v7, v15)),
]
}
};
}
#[cfg(all(
target_arch = "x86_64",
any(feature = "std", not(any(target_feature = "avx2", target_feature = "avx512f")))
))]
mod x86_64_sse2 {
use core::arch::x86_64::*;
use super::IV;
#[inline(always)]
fn load(xs: &[u32; 4]) -> __m128i {
unsafe { _mm_loadu_si128(xs.as_ptr().cast()) }
}
#[inline(always)]
fn store(x: __m128i) -> [u32; 4] {
let mut out = [0u32; 4];
unsafe { _mm_storeu_si128(out.as_mut_ptr().cast(), x) };
out
}
#[inline(always)]
fn splat(x: u32) -> __m128i {
unsafe { _mm_set1_epi32(x as i32) }
}
#[inline(always)]
fn add(a: __m128i, b: __m128i) -> __m128i {
unsafe { _mm_add_epi32(a, b) }
}
#[inline(always)]
fn xor(a: __m128i, b: __m128i) -> __m128i {
unsafe { _mm_xor_si128(a, b) }
}
#[inline(always)]
fn rotr16(x: __m128i) -> __m128i {
unsafe { _mm_or_si128(_mm_srli_epi32::<16>(x), _mm_slli_epi32::<16>(x)) }
}
#[inline(always)]
fn rotr12(x: __m128i) -> __m128i {
unsafe { _mm_or_si128(_mm_srli_epi32::<12>(x), _mm_slli_epi32::<20>(x)) }
}
#[inline(always)]
fn rotr8(x: __m128i) -> __m128i {
unsafe { _mm_or_si128(_mm_srli_epi32::<8>(x), _mm_slli_epi32::<24>(x)) }
}
#[inline(always)]
fn rotr7(x: __m128i) -> __m128i {
unsafe { _mm_or_si128(_mm_srli_epi32::<7>(x), _mm_slli_epi32::<25>(x)) }
}
define_x86_packed_compress!(compress_packed_4, 4);
}
#[cfg(all(
target_arch = "x86_64",
any(feature = "std", all(target_feature = "avx2", not(target_feature = "avx512f")))
))]
mod x86_64_avx2 {
use core::arch::x86_64::*;
use super::IV;
#[inline(always)]
fn load(xs: &[u32; 8]) -> __m256i {
unsafe { _mm256_loadu_si256(xs.as_ptr().cast()) }
}
#[inline(always)]
fn store(x: __m256i) -> [u32; 8] {
let mut out = [0u32; 8];
unsafe { _mm256_storeu_si256(out.as_mut_ptr().cast(), x) };
out
}
#[inline(always)]
fn splat(x: u32) -> __m256i {
unsafe { _mm256_set1_epi32(x as i32) }
}
#[inline(always)]
fn add(a: __m256i, b: __m256i) -> __m256i {
unsafe { _mm256_add_epi32(a, b) }
}
#[inline(always)]
fn xor(a: __m256i, b: __m256i) -> __m256i {
unsafe { _mm256_xor_si256(a, b) }
}
#[inline(always)]
fn rotr16(x: __m256i) -> __m256i {
unsafe {
let mask = _mm256_setr_epi8(
2, 3, 0, 1, 6, 7, 4, 5, 10, 11, 8, 9, 14, 15, 12, 13, 2, 3, 0, 1, 6, 7, 4, 5, 10,
11, 8, 9, 14, 15, 12, 13,
);
_mm256_shuffle_epi8(x, mask)
}
}
#[inline(always)]
fn rotr12(x: __m256i) -> __m256i {
unsafe { _mm256_or_si256(_mm256_srli_epi32::<12>(x), _mm256_slli_epi32::<20>(x)) }
}
#[inline(always)]
fn rotr8(x: __m256i) -> __m256i {
unsafe {
let mask = _mm256_setr_epi8(
1, 2, 3, 0, 5, 6, 7, 4, 9, 10, 11, 8, 13, 14, 15, 12, 1, 2, 3, 0, 5, 6, 7, 4, 9,
10, 11, 8, 13, 14, 15, 12,
);
_mm256_shuffle_epi8(x, mask)
}
}
#[inline(always)]
fn rotr7(x: __m256i) -> __m256i {
unsafe { _mm256_or_si256(_mm256_srli_epi32::<7>(x), _mm256_slli_epi32::<25>(x)) }
}
define_x86_packed_compress!(
#[target_feature(enable = "avx2")]
compress_packed_8,
8
);
}
#[cfg(all(target_arch = "x86_64", any(feature = "std", target_feature = "avx512f")))]
mod x86_64_avx512 {
use core::arch::x86_64::*;
use super::IV;
#[inline(always)]
fn load(xs: &[u32; 16]) -> __m512i {
unsafe { _mm512_loadu_si512(xs.as_ptr().cast()) }
}
#[inline(always)]
fn store(x: __m512i) -> [u32; 16] {
let mut out = [0u32; 16];
unsafe { _mm512_storeu_si512(out.as_mut_ptr().cast(), x) };
out
}
#[inline(always)]
fn splat(x: u32) -> __m512i {
unsafe { _mm512_set1_epi32(x as i32) }
}
#[inline(always)]
fn add(a: __m512i, b: __m512i) -> __m512i {
unsafe { _mm512_add_epi32(a, b) }
}
#[inline(always)]
fn xor(a: __m512i, b: __m512i) -> __m512i {
unsafe { _mm512_xor_si512(a, b) }
}
#[inline(always)]
fn rotr16(x: __m512i) -> __m512i {
unsafe { _mm512_ror_epi32::<16>(x) }
}
#[inline(always)]
fn rotr12(x: __m512i) -> __m512i {
unsafe { _mm512_ror_epi32::<12>(x) }
}
#[inline(always)]
fn rotr8(x: __m512i) -> __m512i {
unsafe { _mm512_ror_epi32::<8>(x) }
}
#[inline(always)]
fn rotr7(x: __m512i) -> __m512i {
unsafe { _mm512_ror_epi32::<7>(x) }
}
define_x86_packed_compress!(
#[target_feature(enable = "avx512f")]
compress_packed_16,
16
);
}
#[cfg(all(target_arch = "aarch64", target_feature = "neon"))]
mod neon {
use core::arch::aarch64::*;
use super::IV;
#[inline(always)]
fn load(xs: &[u32; 4]) -> uint32x4_t {
unsafe { vld1q_u32(xs.as_ptr()) }
}
#[inline(always)]
fn store(x: uint32x4_t) -> [u32; 4] {
let mut out = [0u32; 4];
unsafe { vst1q_u32(out.as_mut_ptr(), x) };
out
}
#[inline(always)]
fn splat(x: u32) -> uint32x4_t {
unsafe { vdupq_n_u32(x) }
}
#[inline(always)]
fn add(a: uint32x4_t, b: uint32x4_t) -> uint32x4_t {
unsafe { vaddq_u32(a, b) }
}
#[inline(always)]
fn xor(a: uint32x4_t, b: uint32x4_t) -> uint32x4_t {
unsafe { veorq_u32(a, b) }
}
#[inline(always)]
fn rotr16(x: uint32x4_t) -> uint32x4_t {
unsafe { vreinterpretq_u32_u16(vrev32q_u16(vreinterpretq_u16_u32(x))) }
}
#[inline(always)]
fn rotr12(x: uint32x4_t) -> uint32x4_t {
unsafe { vsriq_n_u32::<12>(vshlq_n_u32::<20>(x), x) }
}
#[inline(always)]
fn rotr8(x: uint32x4_t) -> uint32x4_t {
unsafe {
let mask = vld1q_u8([1u8, 2, 3, 0, 5, 6, 7, 4, 9, 10, 11, 8, 13, 14, 15, 12].as_ptr());
vreinterpretq_u32_u8(vqtbl1q_u8(vreinterpretq_u8_u32(x), mask))
}
}
#[inline(always)]
fn rotr7(x: uint32x4_t) -> uint32x4_t {
unsafe { vsriq_n_u32::<7>(vshlq_n_u32::<25>(x), x) }
}
#[inline(always)]
pub(super) fn compress_packed_4(cv: [[u32; 4]; 8], block: [[u32; 4]; 16]) -> [[u32; 4]; 8] {
let mut v0 = load(&cv[0]);
let mut v1 = load(&cv[1]);
let mut v2 = load(&cv[2]);
let mut v3 = load(&cv[3]);
let mut v4 = load(&cv[4]);
let mut v5 = load(&cv[5]);
let mut v6 = load(&cv[6]);
let mut v7 = load(&cv[7]);
let mut v8 = splat(IV[0]);
let mut v9 = splat(IV[1]);
let mut v10 = splat(IV[2]);
let mut v11 = splat(IV[3]);
let mut v12 = splat(IV[4]);
let mut v13 = splat(IV[5]);
let mut v14 = splat(IV[6]);
let mut v15 = splat(IV[7]);
macro_rules! round {
(
$m0:literal,
$m1:literal,
$m2:literal,
$m3:literal,
$m4:literal,
$m5:literal,
$m6:literal,
$m7:literal,
$m8:literal,
$m9:literal,
$m10:literal,
$m11:literal,
$m12:literal,
$m13:literal,
$m14:literal,
$m15:literal
) => {{
let m0 = load(&block[$m0]);
let m1 = load(&block[$m1]);
let m2 = load(&block[$m2]);
let m3 = load(&block[$m3]);
let m4 = load(&block[$m4]);
let m5 = load(&block[$m5]);
let m6 = load(&block[$m6]);
let m7 = load(&block[$m7]);
let m8 = load(&block[$m8]);
let m9 = load(&block[$m9]);
let m10 = load(&block[$m10]);
let m11 = load(&block[$m11]);
let m12 = load(&block[$m12]);
let m13 = load(&block[$m13]);
let m14 = load(&block[$m14]);
let m15 = load(&block[$m15]);
v0 = add(add(v0, v4), m0);
v1 = add(add(v1, v5), m2);
v2 = add(add(v2, v6), m4);
v3 = add(add(v3, v7), m6);
v12 = rotr16(xor(v12, v0));
v13 = rotr16(xor(v13, v1));
v14 = rotr16(xor(v14, v2));
v15 = rotr16(xor(v15, v3));
v8 = add(v8, v12);
v9 = add(v9, v13);
v10 = add(v10, v14);
v11 = add(v11, v15);
v4 = rotr12(xor(v4, v8));
v5 = rotr12(xor(v5, v9));
v6 = rotr12(xor(v6, v10));
v7 = rotr12(xor(v7, v11));
v0 = add(add(v0, v4), m1);
v1 = add(add(v1, v5), m3);
v2 = add(add(v2, v6), m5);
v3 = add(add(v3, v7), m7);
v12 = rotr8(xor(v12, v0));
v13 = rotr8(xor(v13, v1));
v14 = rotr8(xor(v14, v2));
v15 = rotr8(xor(v15, v3));
v8 = add(v8, v12);
v9 = add(v9, v13);
v10 = add(v10, v14);
v11 = add(v11, v15);
v4 = rotr7(xor(v4, v8));
v5 = rotr7(xor(v5, v9));
v6 = rotr7(xor(v6, v10));
v7 = rotr7(xor(v7, v11));
v0 = add(add(v0, v5), m8);
v1 = add(add(v1, v6), m10);
v2 = add(add(v2, v7), m12);
v3 = add(add(v3, v4), m14);
v15 = rotr16(xor(v15, v0));
v12 = rotr16(xor(v12, v1));
v13 = rotr16(xor(v13, v2));
v14 = rotr16(xor(v14, v3));
v10 = add(v10, v15);
v11 = add(v11, v12);
v8 = add(v8, v13);
v9 = add(v9, v14);
v5 = rotr12(xor(v5, v10));
v6 = rotr12(xor(v6, v11));
v7 = rotr12(xor(v7, v8));
v4 = rotr12(xor(v4, v9));
v0 = add(add(v0, v5), m9);
v1 = add(add(v1, v6), m11);
v2 = add(add(v2, v7), m13);
v3 = add(add(v3, v4), m15);
v15 = rotr8(xor(v15, v0));
v12 = rotr8(xor(v12, v1));
v13 = rotr8(xor(v13, v2));
v14 = rotr8(xor(v14, v3));
v10 = add(v10, v15);
v11 = add(v11, v12);
v8 = add(v8, v13);
v9 = add(v9, v14);
v5 = rotr7(xor(v5, v10));
v6 = rotr7(xor(v6, v11));
v7 = rotr7(xor(v7, v8));
v4 = rotr7(xor(v4, v9));
}};
}
round!(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
round!(2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8);
round!(3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1);
round!(10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6);
round!(12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4);
round!(9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7);
round!(11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13);
[
store(xor(v0, v8)),
store(xor(v1, v9)),
store(xor(v2, v10)),
store(xor(v3, v11)),
store(xor(v4, v12)),
store(xor(v5, v13)),
store(xor(v6, v14)),
store(xor(v7, v15)),
]
}
}
#[cfg(test)]
mod tests {
use super::*;
fn next_u32(state: &mut u64) -> u32 {
*state ^= *state << 13;
*state ^= *state >> 7;
*state ^= *state << 17;
*state as u32
}
fn random_packed<const W: usize>(state: &mut u64) -> ([[u32; W]; 8], [[u32; W]; 16]) {
let cv = array::from_fn(|_| array::from_fn(|_| next_u32(state)));
let block = array::from_fn(|_| array::from_fn(|_| next_u32(state)));
(cv, block)
}
#[test]
fn sub_batches_reassemble_lanes_in_order() {
let mut state = 0x1234_5678_9abc_def0u64;
let (cv, block) = random_packed::<PACKED_LANES>(&mut state);
let expected = compress_packed::<PACKED_LANES>(cv, block);
assert_eq!(compress_via_sub_batches::<4>(&cv, &block, compress_packed::<4>), expected);
assert_eq!(compress_via_sub_batches::<8>(&cv, &block, compress_packed::<8>), expected);
assert_eq!(compress_via_sub_batches::<16>(&cv, &block, compress_packed::<16>), expected);
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
mod wasm32_simd128_backend {
use super::*;
const PARAMETER_WORDS: [u32; 4] = [IV[4], IV[5], IV[6], IV[7]];
fn assert_matches_scalar(cv: [[u32; 4]; 8], block: [[u32; 4]; 16]) {
let actual = super::super::wasm32_simd128::compress_packed_4(cv, block);
for lane in 0..4 {
let cv_lane = array::from_fn(|word| cv[word][lane]);
let block_lane = array::from_fn(|word| block[word][lane]);
let expected =
compress_raw_with_parameter_words(cv_lane, block_lane, PARAMETER_WORDS);
let actual_lane = array::from_fn(|word| actual[word][lane]);
assert_eq!(actual_lane, expected, "lane {lane} diverged");
}
}
#[test]
fn wasm32_simd128_packed_backend_matches_scalar_for_adversarial_and_random_inputs() {
const VALUES: [u32; 8] = [
0,
u32::MAX,
0x0000_0001,
0x8000_0000,
0xaaaa_aaaa,
0x5555_5555,
0x0123_4567,
0x89ab_cdef,
];
let cv = array::from_fn(|word| {
array::from_fn(|lane| VALUES[(word * 3 + lane) % VALUES.len()])
});
let block = array::from_fn(|word| {
array::from_fn(|lane| VALUES[(word * 5 + lane * 3) % VALUES.len()])
});
assert_matches_scalar(cv, block);
let mut state = 0x9e37_79b9_7f4a_7c15u64;
for _ in 0..512 {
let (cv, block) = random_packed::<4>(&mut state);
assert_matches_scalar(cv, block);
}
}
#[test]
fn wasm32_simd128_packed_backend_preserves_logical_lane_order_across_sub_batches() {
let mut state = 0x0123_4567_89ab_cdefu64;
let (cv, block) = random_packed::<PACKED_LANES>(&mut state);
let expected = compress_packed::<PACKED_LANES>(cv, block);
let actual = compress_via_sub_batches::<4>(
&cv,
&block,
super::super::wasm32_simd128::compress_packed_4,
);
assert_eq!(actual, expected);
}
}
#[cfg(all(target_arch = "x86_64", feature = "std"))]
mod x86_64_backends {
use super::*;
const PARAMETER_WORDS: [u32; 4] = [IV[4], IV[5], IV[6], IV[7]];
fn assert_packed_backend_matches_reference<const W: usize>(
backend: impl Fn([[u32; W]; 8], [[u32; W]; 16]) -> [[u32; W]; 8],
) {
let mut state = 0x9e37_79b9_7f4a_7c15u64;
for _ in 0..512 {
let (cv, block) = random_packed::<W>(&mut state);
let out = backend(cv, block);
for lane in 0..W {
let cv_lane: [u32; 8] = array::from_fn(|word| cv[word][lane]);
let block_lane: [u32; 16] = array::from_fn(|word| block[word][lane]);
let expected =
compress_raw_with_parameter_words(cv_lane, block_lane, PARAMETER_WORDS);
let actual: [u32; 8] = array::from_fn(|word| out[word][lane]);
assert_eq!(actual, expected, "lane {lane} diverged");
}
}
}
fn assert_row_variant_matches_reference(
raw: impl Fn(&[u32; 8], &[u32; 16]) -> [u32; 8],
xof: impl Fn(&[u32; 8], &[u32; 16]) -> [u32; 16],
) {
let mut state = 0x0123_4567_89ab_cdefu64;
for _ in 0..2048 {
let cv: [u32; 8] = array::from_fn(|_| next_u32(&mut state));
let block: [u32; 16] = array::from_fn(|_| next_u32(&mut state));
assert_eq!(
raw(&cv, &block),
compress_raw_with_parameter_words(cv, block, PARAMETER_WORDS)
);
assert_eq!(
xof(&cv, &block),
compress_raw_xof_with_parameter_words(cv, block, PARAMETER_WORDS)
);
}
}
#[test]
fn sse2_packed_backend_matches_reference() {
assert_packed_backend_matches_reference::<4>(|cv, block| {
unsafe { x86_64_sse2::compress_packed_4(cv, block) }
});
}
#[test]
fn avx2_packed_backend_matches_reference() {
if !cpu::has_avx2() {
std::eprintln!("skipped: the running CPU lacks AVX2");
return;
}
assert_packed_backend_matches_reference::<8>(|cv, block| {
unsafe { x86_64_avx2::compress_packed_8(cv, block) }
});
}
#[test]
fn avx512_packed_backend_matches_reference() {
if !cpu::has_avx512f() {
std::eprintln!("skipped: the running CPU lacks AVX-512F");
return;
}
assert_packed_backend_matches_reference::<16>(|cv, block| {
unsafe { x86_64_avx512::compress_packed_16(cv, block) }
});
}
#[test]
fn sse2_row_variant_matches_reference() {
assert_row_variant_matches_reference(row_x86::compress_raw, row_x86::compress_raw_xof);
}
#[test]
fn sse41_row_variant_matches_reference() {
if !cpu::has_sse41() {
std::eprintln!("skipped: the running CPU lacks SSE4.1");
return;
}
assert_row_variant_matches_reference(
|cv, block| {
unsafe { row_x86::compress_raw_sse41(cv, block) }
},
|cv, block| {
unsafe { row_x86::compress_raw_xof_sse41(cv, block) }
},
);
}
#[test]
fn avx512vl_row_variant_matches_reference() {
if !cpu::has_avx512vl() {
std::eprintln!("skipped: the running CPU lacks AVX-512VL");
return;
}
assert_row_variant_matches_reference(
|cv, block| {
unsafe { row_x86::compress_raw_avx512vl(cv, block) }
},
|cv, block| {
unsafe { row_x86::compress_raw_xof_avx512vl(cv, block) }
},
);
}
}
}