#![allow(clippy::indexing_slicing)] #![allow(unsafe_code)]
use core::arch::global_asm;
#[cfg(any(feature = "diag", all(test, target_os = "linux")))]
use super::PolyVec;
#[cfg(all(test, target_os = "linux"))]
use super::SAMPLE_NTT_ACC_CHUNK_COEFFS;
use super::{GAMMAS_MONT, Poly};
#[cfg(target_os = "macos")]
global_asm!(include_str!("../asm/rscrypto_mlkem_basemul_aarch64_apple_darwin.s"));
#[cfg(target_os = "linux")]
global_asm!(include_str!("../asm/rscrypto_mlkem_rej_uniform_aarch64_linux.s"));
#[cfg(target_os = "linux")]
global_asm!(include_str!("../asm/rscrypto_mlkem_basemul_aarch64_linux.s"));
#[cfg(target_os = "macos")]
unsafe extern "C" {
fn rscrypto_mlkem_basemul_accumulate_aarch64_apple_darwin(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k2_aarch64_apple_darwin(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k3_aarch64_apple_darwin(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k4_aarch64_apple_darwin(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
}
#[cfg(target_os = "linux")]
unsafe extern "C" {
fn rscrypto_mlkem_rej_uniform_block_aarch64_linux(out: *mut u16, input: *const u8) -> usize;
fn rscrypto_mlkem_rej_uniform_block_bounded_aarch64_linux(out: *mut u16, input: *const u8, cap: usize) -> usize;
fn rscrypto_mlkem_rej_uniform_triple_block_aarch64_linux(
out0: *mut u16,
input0: *const u8,
out1: *mut u16,
input1: *const u8,
out2: *mut u16,
input2: *const u8,
) -> u64;
fn rscrypto_mlkem_rej_uniform_triple_block_bounded_aarch64_linux(
out0: *mut u16,
input0: *const u8,
out1: *mut u16,
input1: *const u8,
out2: *mut u16,
input2: *const u8,
caps: *const usize,
) -> u64;
#[cfg(any(test, feature = "diag"))]
fn rscrypto_mlkem_rej_uniform_3blocks_aarch64_linux(out: *mut u16, input: *const u8) -> usize;
#[cfg(any(test, feature = "diag"))]
fn rscrypto_mlkem_basemul_accumulate_aarch64_linux(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
#[cfg(test)]
fn rscrypto_mlkem_basemul_accumulate_chunk_aarch64_linux(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k2_aarch64_linux(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k3_aarch64_linux(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
fn rscrypto_mlkem_basemul_accumulate_k4_aarch64_linux(
acc: *mut u16,
a: *const u16,
b: *const u16,
gammas_mont: *const i16,
);
}
#[cfg(target_os = "linux")]
#[inline(always)]
fn unpack_triple_counts(packed: u64) -> [usize; 3] {
[
(packed & 0xffff) as usize,
((packed >> 16) & 0xffff) as usize,
((packed >> 32) & 0xffff) as usize,
]
}
#[cfg(target_os = "linux")]
#[inline]
pub(super) unsafe fn sample_ntt_rej_uniform_block_asm(out: *mut u16, input: *const u8) -> usize {
unsafe { rscrypto_mlkem_rej_uniform_block_aarch64_linux(out, input) }
}
#[cfg(target_os = "linux")]
#[inline]
pub(super) unsafe fn sample_ntt_rej_uniform_block_bounded_asm(out: *mut u16, input: *const u8, cap: usize) -> usize {
unsafe { rscrypto_mlkem_rej_uniform_block_bounded_aarch64_linux(out, input, cap) }
}
#[cfg(target_os = "linux")]
#[inline]
pub(super) unsafe fn sample_ntt_rej_uniform_triple_block_asm(
out0: *mut u16,
input0: *const u8,
out1: *mut u16,
input1: *const u8,
out2: *mut u16,
input2: *const u8,
) -> [usize; 3] {
let packed =
unsafe { rscrypto_mlkem_rej_uniform_triple_block_aarch64_linux(out0, input0, out1, input1, out2, input2) };
unpack_triple_counts(packed)
}
#[cfg(target_os = "linux")]
#[inline]
pub(super) unsafe fn sample_ntt_rej_uniform_triple_block_bounded_asm(
out0: *mut u16,
input0: *const u8,
out1: *mut u16,
input1: *const u8,
out2: *mut u16,
input2: *const u8,
caps: [usize; 3],
) -> [usize; 3] {
let packed = unsafe {
rscrypto_mlkem_rej_uniform_triple_block_bounded_aarch64_linux(
out0,
input0,
out1,
input1,
out2,
input2,
caps.as_ptr(),
)
};
unpack_triple_counts(packed)
}
#[cfg(all(any(test, feature = "diag"), target_os = "linux"))]
#[inline]
pub(super) unsafe fn sample_ntt_rej_uniform_3blocks_asm(out: *mut u16, input: *const u8) -> usize {
unsafe { rscrypto_mlkem_rej_uniform_3blocks_aarch64_linux(out, input) }
}
#[inline]
#[cfg(any(test, feature = "diag", target_os = "macos"))]
pub(super) unsafe fn basemul_accumulate_asm(acc: &mut Poly, a: &Poly, b: &Poly) {
#[cfg(target_os = "macos")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_aarch64_apple_darwin(
acc.as_mut_ptr(),
a.as_ptr(),
b.as_ptr(),
GAMMAS_MONT.as_ptr(),
);
}
}
#[cfg(target_os = "linux")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_aarch64_linux(acc.as_mut_ptr(), a.as_ptr(), b.as_ptr(), GAMMAS_MONT.as_ptr());
}
}
}
#[cfg(all(test, target_os = "linux"))]
pub(super) unsafe fn test_basemul_accumulate_asm(acc: &mut Poly, a: &Poly, b: &Poly) {
unsafe {
basemul_accumulate_asm(acc, a, b);
}
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_asm_digest(seed: u16) -> u16 {
let a = super::diag_poly(seed);
let b = super::diag_poly(seed.wrapping_add(1));
let acc = super::diag_poly(seed.wrapping_add(2));
unsafe { diag_basemul_accumulate_asm_input_digest(a, b, acc) }
}
#[cfg(any(target_os = "macos", target_os = "linux"))]
#[inline]
pub(super) unsafe fn basemul_accumulate_k2_asm_ptr(acc: &mut Poly, a: *const u16, b: *const u16) {
#[cfg(target_os = "macos")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k2_aarch64_apple_darwin(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
#[cfg(target_os = "linux")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k2_aarch64_linux(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
}
#[cfg(any(target_os = "macos", target_os = "linux"))]
#[inline]
pub(super) unsafe fn basemul_accumulate_k3_asm_ptr(acc: &mut Poly, a: *const u16, b: *const u16) {
#[cfg(target_os = "macos")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k3_aarch64_apple_darwin(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
#[cfg(target_os = "linux")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k3_aarch64_linux(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
}
#[cfg(any(target_os = "macos", target_os = "linux"))]
#[inline]
pub(super) unsafe fn basemul_accumulate_k4_asm_ptr(acc: &mut Poly, a: *const u16, b: *const u16) {
#[cfg(target_os = "macos")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k4_aarch64_apple_darwin(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
#[cfg(target_os = "linux")]
{
unsafe {
rscrypto_mlkem_basemul_accumulate_k4_aarch64_linux(acc.as_mut_ptr(), a, b, GAMMAS_MONT.as_ptr());
}
}
}
#[inline]
#[cfg(all(test, target_os = "linux"))]
unsafe fn basemul_accumulate_k2_asm(acc: &mut Poly, a: &PolyVec<2>, b: &PolyVec<2>) {
unsafe {
basemul_accumulate_k2_asm_ptr(acc, a.as_ptr().cast::<u16>(), b.as_ptr().cast::<u16>());
}
}
#[inline]
#[cfg(any(feature = "diag", all(test, target_os = "linux")))]
unsafe fn basemul_accumulate_k3_asm(acc: &mut Poly, a: &PolyVec<3>, b: &PolyVec<3>) {
unsafe {
basemul_accumulate_k3_asm_ptr(acc, a.as_ptr().cast::<u16>(), b.as_ptr().cast::<u16>());
}
}
#[inline]
#[cfg(any(feature = "diag", all(test, target_os = "linux")))]
unsafe fn basemul_accumulate_k4_asm(acc: &mut Poly, a: &PolyVec<4>, b: &PolyVec<4>) {
unsafe {
basemul_accumulate_k4_asm_ptr(acc, a.as_ptr().cast::<u16>(), b.as_ptr().cast::<u16>());
}
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_asm_input_digest(a: Poly, b: Poly, mut acc: Poly) -> u16 {
unsafe {
basemul_accumulate_asm(&mut acc, &a, &b);
}
let digest = super::diag_fold_poly(&acc);
super::zeroize_poly(&mut acc);
digest
}
#[cfg(all(test, target_os = "linux"))]
pub(super) unsafe fn test_basemul_accumulate_k2_asm(acc: &mut Poly, a: &PolyVec<2>, b: &PolyVec<2>) {
unsafe {
basemul_accumulate_k2_asm(acc, a, b);
}
}
#[cfg(all(test, target_os = "linux"))]
pub(super) unsafe fn test_basemul_accumulate_k3_asm(acc: &mut Poly, a: &PolyVec<3>, b: &PolyVec<3>) {
unsafe {
basemul_accumulate_k3_asm(acc, a, b);
}
}
#[cfg(all(test, target_os = "linux"))]
pub(super) unsafe fn test_basemul_accumulate_k4_asm(acc: &mut Poly, a: &PolyVec<4>, b: &PolyVec<4>) {
unsafe {
basemul_accumulate_k4_asm(acc, a, b);
}
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_k3_asm_digest(seed: u16) -> u16 {
let a = [
super::diag_poly(seed),
super::diag_poly(seed.wrapping_add(1)),
super::diag_poly(seed.wrapping_add(2)),
];
let b = [
super::diag_poly(seed.wrapping_add(3)),
super::diag_poly(seed.wrapping_add(4)),
super::diag_poly(seed.wrapping_add(5)),
];
let acc = super::diag_poly(seed.wrapping_add(6));
unsafe { diag_basemul_accumulate_k3_asm_input_digest(a, b, acc) }
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_k4_asm_digest(seed: u16) -> u16 {
let a = [
super::diag_poly(seed),
super::diag_poly(seed.wrapping_add(1)),
super::diag_poly(seed.wrapping_add(2)),
super::diag_poly(seed.wrapping_add(3)),
];
let b = [
super::diag_poly(seed.wrapping_add(4)),
super::diag_poly(seed.wrapping_add(5)),
super::diag_poly(seed.wrapping_add(6)),
super::diag_poly(seed.wrapping_add(7)),
];
let acc = super::diag_poly(seed.wrapping_add(8));
unsafe { diag_basemul_accumulate_k4_asm_input_digest(a, b, acc) }
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_k3_asm_input_digest(
mut a: PolyVec<3>,
mut b: PolyVec<3>,
mut acc: Poly,
) -> u16 {
unsafe {
basemul_accumulate_k3_asm(&mut acc, &a, &b);
}
let digest = super::diag_fold_poly(&acc);
super::zeroize_polyvec(&mut a);
super::zeroize_polyvec(&mut b);
super::zeroize_poly(&mut acc);
digest
}
#[cfg(feature = "diag")]
pub(super) unsafe fn diag_basemul_accumulate_k4_asm_input_digest(
mut a: PolyVec<4>,
mut b: PolyVec<4>,
mut acc: Poly,
) -> u16 {
unsafe {
basemul_accumulate_k4_asm(&mut acc, &a, &b);
}
let digest = super::diag_fold_poly(&acc);
super::zeroize_polyvec(&mut a);
super::zeroize_polyvec(&mut b);
super::zeroize_poly(&mut acc);
digest
}
#[cfg(all(test, target_os = "linux"))]
pub(super) unsafe fn test_basemul_accumulate_chunk_asm(
acc: &mut Poly,
a: &[u16; SAMPLE_NTT_ACC_CHUNK_COEFFS],
b: &Poly,
coeff_offset: usize,
) {
debug_assert_eq!(coeff_offset % SAMPLE_NTT_ACC_CHUNK_COEFFS, 0);
debug_assert!(coeff_offset.strict_add(SAMPLE_NTT_ACC_CHUNK_COEFFS) <= acc.len());
let gamma_offset = coeff_offset / 2;
unsafe {
rscrypto_mlkem_basemul_accumulate_chunk_aarch64_linux(
acc.as_mut_ptr().add(coeff_offset),
a.as_ptr(),
b.as_ptr().add(coeff_offset),
GAMMAS_MONT.as_ptr().add(gamma_offset),
);
}
}