miden-crypto 0.34.0

Miden Cryptographic primitives
Documentation
//! Row-wise (diagonalized) single-block BLAKE3 compression for x86_64.
//!
//! Adapted from the `blake3` crate's `rust_sse2::compress_pre` (MIT/Apache-2.0/CC0). It applies
//! Eidos's fixed parameter-word tail (`v[12..16] = IV[4..8]`, without counter, block length, or
//! flags) to a `[u32; 16]` message block using four diagonalized 128-bit rows.
//!
//! SSE2 is part of the x86_64 architectural baseline. The SSE4.1 variant uses native blends for
//! message permutation; its rotations keep the shift-based form used by upstream Rust BLAKE3.
//!
//! The AVX-512F/VL-targeted variant uses the same algorithm with access to XMM16-31, reducing
//! register spills without changing the logical width. Both variants are emitted from the same
//! macro body. `avx512f` is AVX-512VL's prerequisite.

use core::arch::x86_64::*;

use super::IV;

#[inline(always)]
unsafe fn loadu(src: *const u32) -> __m128i {
    unsafe { _mm_loadu_si128(src.cast()) }
}

#[inline(always)]
unsafe fn storeu(src: __m128i, dest: *mut u32) {
    unsafe { _mm_storeu_si128(dest.cast(), src) }
}

#[inline(always)]
unsafe fn add(a: __m128i, b: __m128i) -> __m128i {
    unsafe { _mm_add_epi32(a, b) }
}

#[inline(always)]
unsafe fn xor(a: __m128i, b: __m128i) -> __m128i {
    unsafe { _mm_xor_si128(a, b) }
}

#[inline(always)]
unsafe fn set4(a: u32, b: u32, c: u32, d: u32) -> __m128i {
    unsafe { _mm_setr_epi32(a as i32, b as i32, c as i32, d as i32) }
}

#[inline(always)]
unsafe fn rot16(a: __m128i) -> __m128i {
    unsafe { _mm_or_si128(_mm_srli_epi32(a, 16), _mm_slli_epi32(a, 32 - 16)) }
}

#[inline(always)]
unsafe fn rot12(a: __m128i) -> __m128i {
    unsafe { _mm_or_si128(_mm_srli_epi32(a, 12), _mm_slli_epi32(a, 32 - 12)) }
}

#[inline(always)]
unsafe fn rot8(a: __m128i) -> __m128i {
    unsafe { _mm_or_si128(_mm_srli_epi32(a, 8), _mm_slli_epi32(a, 32 - 8)) }
}

#[inline(always)]
unsafe fn rot7(a: __m128i) -> __m128i {
    unsafe { _mm_or_si128(_mm_srli_epi32(a, 7), _mm_slli_epi32(a, 32 - 7)) }
}

#[inline(always)]
unsafe fn g1(
    row0: &mut __m128i,
    row1: &mut __m128i,
    row2: &mut __m128i,
    row3: &mut __m128i,
    m: __m128i,
) {
    unsafe {
        *row0 = add(add(*row0, m), *row1);
        *row3 = xor(*row3, *row0);
        *row3 = rot16(*row3);
        *row2 = add(*row2, *row3);
        *row1 = xor(*row1, *row2);
        *row1 = rot12(*row1);
    }
}

#[inline(always)]
unsafe fn g2(
    row0: &mut __m128i,
    row1: &mut __m128i,
    row2: &mut __m128i,
    row3: &mut __m128i,
    m: __m128i,
) {
    unsafe {
        *row0 = add(add(*row0, m), *row1);
        *row3 = xor(*row3, *row0);
        *row3 = rot8(*row3);
        *row2 = add(*row2, *row3);
        *row1 = xor(*row1, *row2);
        *row1 = rot7(*row1);
    }
}

macro_rules! mm_shuffle {
    ($z:expr, $y:expr, $x:expr, $w:expr) => {
        ($z << 6) | ($y << 4) | ($x << 2) | $w
    };
}

macro_rules! shuffle2 {
    ($a:expr, $b:expr, $c:expr) => {
        _mm_castps_si128(_mm_shuffle_ps(_mm_castsi128_ps($a), _mm_castsi128_ps($b), $c))
    };
}

// Leaving row1 unrotated avoids an extra shuffle; the message loads below account for this layout.
#[inline(always)]
unsafe fn diagonalize(row0: &mut __m128i, row2: &mut __m128i, row3: &mut __m128i) {
    unsafe {
        *row0 = _mm_shuffle_epi32(*row0, mm_shuffle!(2, 1, 0, 3));
        *row3 = _mm_shuffle_epi32(*row3, mm_shuffle!(1, 0, 3, 2));
        *row2 = _mm_shuffle_epi32(*row2, mm_shuffle!(0, 3, 2, 1));
    }
}

#[inline(always)]
unsafe fn undiagonalize(row0: &mut __m128i, row2: &mut __m128i, row3: &mut __m128i) {
    unsafe {
        *row0 = _mm_shuffle_epi32(*row0, mm_shuffle!(0, 3, 2, 1));
        *row3 = _mm_shuffle_epi32(*row3, mm_shuffle!(1, 0, 3, 2));
        *row2 = _mm_shuffle_epi32(*row2, mm_shuffle!(2, 1, 0, 3));
    }
}

#[inline(always)]
#[cfg(any(feature = "std", not(target_feature = "sse4.1"), target_feature = "avx512vl"))]
unsafe fn blend_epi16<const IMM8: i32>(a: __m128i, b: __m128i) -> __m128i {
    unsafe {
        let bits = _mm_set_epi16(0x80, 0x40, 0x20, 0x10, 0x08, 0x04, 0x02, 0x01);
        let mut mask = _mm_set1_epi16(IMM8 as i16);
        mask = _mm_and_si128(mask, bits);
        mask = _mm_cmpeq_epi16(mask, bits);
        _mm_or_si128(_mm_and_si128(mask, b), _mm_andnot_si128(mask, a))
    }
}

#[inline]
#[cfg(any(
    feature = "std",
    all(target_feature = "sse4.1", not(target_feature = "avx512vl"))
))]
#[target_feature(enable = "sse4.1")]
unsafe fn blend_sse41<const IMM8: i32>(a: __m128i, b: __m128i) -> __m128i {
    _mm_blend_epi16::<IMM8>(a, b)
}

/// Row-wise diagonalized permutation.
///
/// Returns `[row0, row1, row2, row3] = [v[0..4], v[4..8], v[8..12], v[12..16]]` of the standard
/// BLAKE3 permuted state after all seven rounds, with Eidos's fixed parameter-word tail
/// (`v[12..16] = IV[4..8]`, matching `super::permuted_state_with_parameter_words`).
///
/// The macro emits baseline SSE2, SSE4.1, and AVX-512F/VL variants from the same body.
macro_rules! define_compress_pre {
    ($(#[$attr:meta])* $name:ident, $blend:ident) => {
        $(#[$attr])*
        #[inline]
        unsafe fn $name(cv: &[u32; 8], block: &[u32; 16]) -> [__m128i; 4] {
            unsafe {
                let row0 = &mut loadu(cv.as_ptr());
                let row1 = &mut loadu(cv.as_ptr().add(4));
                let row2 = &mut set4(IV[0], IV[1], IV[2], IV[3]);
                let row3 = &mut set4(IV[4], IV[5], IV[6], IV[7]);

                let mut m0 = loadu(block.as_ptr());
                let mut m1 = loadu(block.as_ptr().add(4));
                let mut m2 = loadu(block.as_ptr().add(8));
                let mut m3 = loadu(block.as_ptr().add(12));

                let mut t0;
                let mut t1;
                let mut t2;
                let mut t3;
                let mut tt;

                // Round 1 permutes the message words from the original input order into the
                // groups that get mixed in parallel.
                t0 = shuffle2!(m0, m1, mm_shuffle!(2, 0, 2, 0));
                g1(row0, row1, row2, row3, t0);
                t1 = shuffle2!(m0, m1, mm_shuffle!(3, 1, 3, 1));
                g2(row0, row1, row2, row3, t1);
                diagonalize(row0, row2, row3);
                t2 = shuffle2!(m2, m3, mm_shuffle!(2, 0, 2, 0));
                t2 = _mm_shuffle_epi32(t2, mm_shuffle!(2, 1, 0, 3));
                g1(row0, row1, row2, row3, t2);
                t3 = shuffle2!(m2, m3, mm_shuffle!(3, 1, 3, 1));
                t3 = _mm_shuffle_epi32(t3, mm_shuffle!(2, 1, 0, 3));
                g2(row0, row1, row2, row3, t3);
                undiagonalize(row0, row2, row3);
                m0 = t0;
                m1 = t1;
                m2 = t2;
                m3 = t3;

                // Rounds 2-7 apply a fixed permutation to the message words produced by the
                // round before, so the same shuffle sequence repeats.
                for _ in 0..6 {
                    t0 = shuffle2!(m0, m1, mm_shuffle!(3, 1, 1, 2));
                    t0 = _mm_shuffle_epi32(t0, mm_shuffle!(0, 3, 2, 1));
                    g1(row0, row1, row2, row3, t0);
                    t1 = shuffle2!(m2, m3, mm_shuffle!(3, 3, 2, 2));
                    tt = _mm_shuffle_epi32(m0, mm_shuffle!(0, 0, 3, 3));
                    t1 = $blend::<0xcc>(tt, t1);
                    g2(row0, row1, row2, row3, t1);
                    diagonalize(row0, row2, row3);
                    t2 = _mm_unpacklo_epi64(m3, m1);
                    tt = $blend::<0xc0>(t2, m2);
                    t2 = _mm_shuffle_epi32(tt, mm_shuffle!(1, 3, 2, 0));
                    g1(row0, row1, row2, row3, t2);
                    t3 = _mm_unpackhi_epi32(m1, m3);
                    tt = _mm_unpacklo_epi32(m2, t3);
                    t3 = _mm_shuffle_epi32(tt, mm_shuffle!(0, 1, 3, 2));
                    g2(row0, row1, row2, row3, t3);
                    undiagonalize(row0, row2, row3);
                    m0 = t0;
                    m1 = t1;
                    m2 = t2;
                    m3 = t3;
                }

                [*row0, *row1, *row2, *row3]
            }
        }
    };
}

define_compress_pre!(
    #[cfg(any(
        feature = "std",
        not(any(target_feature = "avx512vl", target_feature = "sse4.1"))
    ))]
    compress_pre,
    blend_epi16
);
define_compress_pre!(
    #[cfg(any(
        feature = "std",
        all(target_feature = "sse4.1", not(target_feature = "avx512vl"))
    ))]
    #[target_feature(enable = "sse4.1")]
    compress_pre_sse41,
    blend_sse41
);
define_compress_pre!(
    #[cfg(any(feature = "std", target_feature = "avx512vl"))]
    #[target_feature(enable = "avx512f,avx512vl")]
    compress_pre_avx512vl,
    blend_epi16
);

/// Returns the raw eight-word CV fold with Eidos compression's fixed parameter words:
/// `out[i] = v[i] ^ v[i + 8]`.
///
/// Generated twice (see module docs): a plain SSE2 variant needing no runtime feature check, and
/// an `avx512f,avx512vl`-attributed variant for a wider register file.
macro_rules! define_compress_raw {
    ($(#[$attr:meta])* $name:ident, $compress_pre:ident) => {
        $(#[$attr])*
        #[inline]
        pub(super) unsafe fn $name(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 8] {
            unsafe {
                let [row0, row1, row2, row3] = $compress_pre(cv, block);
                let mut out = [0u32; 8];
                storeu(xor(row0, row2), out.as_mut_ptr());
                storeu(xor(row1, row3), out.as_mut_ptr().add(4));
                out
            }
        }
    };
}

define_compress_raw!(
    #[cfg(any(
        feature = "std",
        not(any(target_feature = "avx512vl", target_feature = "sse4.1"))
    ))]
    compress_raw_impl,
    compress_pre
);
define_compress_raw!(
    #[cfg(any(
        feature = "std",
        all(target_feature = "sse4.1", not(target_feature = "avx512vl"))
    ))]
    #[target_feature(enable = "sse4.1")]
    compress_raw_sse41,
    compress_pre_sse41
);
define_compress_raw!(
    #[cfg(any(feature = "std", target_feature = "avx512vl"))]
    #[target_feature(enable = "avx512f,avx512vl")]
    compress_raw_avx512vl,
    compress_pre_avx512vl
);

/// Returns the raw sixteen-word XOF fold with Eidos compression's fixed parameter words:
/// `out[i] = v[i] ^ v[i + 8]` for `i < 8`, `out[i] = v[i] ^ cv[i - 8]` for `i >= 8`.
///
/// Generated twice (see module docs): a plain SSE2 variant needing no runtime feature check, and
/// an `avx512f,avx512vl`-attributed variant for a wider register file.
macro_rules! define_compress_raw_xof {
    ($(#[$attr:meta])* $name:ident, $compress_pre:ident) => {
        $(#[$attr])*
        #[inline]
        pub(super) unsafe fn $name(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 16] {
            unsafe {
                let [row0, row1, row2, row3] = $compress_pre(cv, block);
                let cv_row0 = loadu(cv.as_ptr());
                let cv_row1 = loadu(cv.as_ptr().add(4));
                let mut out = [0u32; 16];
                storeu(xor(row0, row2), out.as_mut_ptr());
                storeu(xor(row1, row3), out.as_mut_ptr().add(4));
                storeu(xor(row2, cv_row0), out.as_mut_ptr().add(8));
                storeu(xor(row3, cv_row1), out.as_mut_ptr().add(12));
                out
            }
        }
    };
}

define_compress_raw_xof!(
    #[cfg(any(
        feature = "std",
        not(any(target_feature = "avx512vl", target_feature = "sse4.1"))
    ))]
    compress_raw_xof_impl,
    compress_pre
);
define_compress_raw_xof!(
    #[cfg(any(
        feature = "std",
        all(target_feature = "sse4.1", not(target_feature = "avx512vl"))
    ))]
    #[target_feature(enable = "sse4.1")]
    compress_raw_xof_sse41,
    compress_pre_sse41
);
define_compress_raw_xof!(
    #[cfg(any(feature = "std", target_feature = "avx512vl"))]
    #[target_feature(enable = "avx512f,avx512vl")]
    compress_raw_xof_avx512vl,
    compress_pre_avx512vl
);

/// Returns the raw eight-word CV fold with Eidos compression's fixed parameter words. Needs no
/// runtime feature check: SSE2 is part of the x86_64 architectural baseline.
#[cfg(any(
    feature = "std",
    not(any(target_feature = "avx512vl", target_feature = "sse4.1"))
))]
#[inline]
pub(super) fn compress_raw(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 8] {
    // SAFETY: SSE2 is part of the x86_64 architectural baseline.
    unsafe { compress_raw_impl(cv, block) }
}

/// Returns the raw sixteen-word XOF fold with Eidos compression's fixed parameter words. Needs no
/// runtime feature check: SSE2 is part of the x86_64 architectural baseline.
#[cfg(any(
    feature = "std",
    not(any(target_feature = "avx512vl", target_feature = "sse4.1"))
))]
#[inline]
pub(super) fn compress_raw_xof(cv: &[u32; 8], block: &[u32; 16]) -> [u32; 16] {
    // SAFETY: SSE2 is part of the x86_64 architectural baseline.
    unsafe { compress_raw_xof_impl(cv, block) }
}