use super::limb::{
black_box_l, ct_add_l_l, ct_mul_add_l_l_l_c, ct_mul_l_l, LimbChoice, LimbType, LIMB_BITS,
};
#[cfg(test)]
use super::limbs_buffer::MpMutUIntSlice;
use super::limbs_buffer::{ct_mp_nlimbs, MpMutUInt, MpUIntCommon};
pub fn ct_mul_trunc_cond_mp_mp<T0: MpMutUInt, T1: MpUIntCommon>(
op0: &mut T0,
op0_in_len: usize,
op1: &T1,
cond: LimbChoice,
) {
debug_assert!(op0_in_len <= op0.len());
if op1.is_empty() {
let cond_mask = cond.select(0, !0);
for j in 0..op0.nlimbs() {
op0.store_l(j, op0.load_l(j) & !cond_mask);
}
return;
}
let op1_nlimbs = op1.nlimbs();
let result_high_mask = op0.partial_high_mask();
let op0_nlimbs = op0.nlimbs();
op0.clear_bytes_above(op0_in_len);
let op0_in_nlimbs = ct_mp_nlimbs(op0_in_len);
let mut j = op0_in_nlimbs;
while j > 0 {
j -= 1;
let op0_val = op0.load_l(j);
op0.store_l(j, 0);
let mut carry = 0;
let result_nlimbs = op0_nlimbs - j;
let mut cond_unit = cond.select(1, 0);
for k in 0..op1_nlimbs.min(result_nlimbs) {
let op1_val = cond.select(0, op1.load_l(k)) | cond_unit;
cond_unit = 0;
let mut result_val = op0.load_l(j + k);
(carry, result_val) = ct_mul_add_l_l_l_c(result_val, op0_val, op1_val, carry);
if k != result_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(j + k, result_val);
} else {
op0.store_l(j + k, result_val & result_high_mask);
}
}
for k in op1_nlimbs..result_nlimbs {
let mut result_val = op0.load_l(j + k);
(carry, result_val) = ct_add_l_l(result_val, carry);
if k != result_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(j + k, result_val);
} else {
op0.store_l(j + k, result_val & result_high_mask);
}
}
}
}
#[cfg(test)]
fn test_ct_mul_trunc_cond_mp_mp<T0: MpMutUIntSlice, T1: MpMutUIntSlice>() {
use super::limb::LIMB_BYTES;
let mut op0 = tst_mk_mp_backing_vec!(T0, 5 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
op1.store_l(0, !0);
op1.store_l(1, !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 0);
assert_eq!(op0.load_l(3), 0);
assert_eq!(op0.load_l(4), 0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0);
assert_eq!(op0.load_l(2), !1);
assert_eq!(op0.load_l(3), !0);
assert_eq!(op0.load_l(4), 0);
let mut op0 = tst_mk_mp_backing_vec!(T0, 3 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
op1.store_l(0, !0);
op1.store_l(1, !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0);
assert_eq!(op0.load_l(2), !1);
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
op1.store_l(0, !0);
op1.store_l(1, !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0);
if !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
return;
}
let mut op0 = tst_mk_mp_backing_vec!(T0, 4 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
op1.store_l(0, !0);
op1.store_l(1, !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 0);
assert_eq!(op0.load_l(3), 0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0);
assert_eq!(op0.load_l(2), !1);
assert_eq!(op0.load_l(3), !0 >> 8);
let mut op0 = tst_mk_mp_backing_vec!(T0, 3 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
op1.store_l(0, !0);
op1.store_l(1, !0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 0);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0);
assert_eq!(op0.load_l(2), (!0 >> 8) ^ 1);
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let mut op1 = tst_mk_mp_backing_vec!(T1, 2 * LIMB_BYTES);
let mut op1 = T1::from_slice(&mut op1).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0 >> 2 * 8);
op1.store_l(0, !0);
op1.store_l(1, !0 >> 2 * 8);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES - 1, &op1, LimbChoice::from(0));
assert_eq!(op0.load_l(0), !0);
assert_eq!(op0.load_l(1), !0 >> 2 * 8);
ct_mul_trunc_cond_mp_mp(&mut op0, 2 * LIMB_BYTES - 1, &op1, LimbChoice::from(1));
assert_eq!(op0.load_l(0), 1);
assert_eq!(op0.load_l(1), 0xfe << 8 * (LIMB_BYTES - 2));
}
#[test]
fn test_ct_mul_trunc_cond_be_be() {
use super::limbs_buffer::MpMutBigEndianUIntByteSlice;
test_ct_mul_trunc_cond_mp_mp::<MpMutBigEndianUIntByteSlice, MpMutBigEndianUIntByteSlice>()
}
#[test]
fn test_ct_mul_trunc_cond_le_le() {
use super::limbs_buffer::MpMutLittleEndianUIntByteSlice;
test_ct_mul_trunc_cond_mp_mp::<MpMutLittleEndianUIntByteSlice, MpMutLittleEndianUIntByteSlice>()
}
#[test]
fn test_ct_mul_trunc_cond_ne_ne() {
use super::limbs_buffer::MpMutNativeEndianUIntLimbsSlice;
test_ct_mul_trunc_cond_mp_mp::<MpMutNativeEndianUIntLimbsSlice, MpMutNativeEndianUIntLimbsSlice>(
)
}
pub fn ct_mul_trunc_mp_mp<T0: MpMutUInt, T1: MpUIntCommon>(
op0: &mut T0,
op0_in_len: usize,
op1: &T1,
) {
ct_mul_trunc_cond_mp_mp(op0, op0_in_len, op1, LimbChoice::from(1))
}
pub fn ct_square_trunc_mp<T0: MpMutUInt>(op0: &mut T0, op0_in_len: usize) {
debug_assert!(op0_in_len <= op0.len());
let result_high_mask = op0.partial_high_mask();
let op0_nlimbs = op0.nlimbs();
op0.clear_bytes_above(op0_in_len);
let op0_in_nlimbs = ct_mp_nlimbs(op0_in_len);
let mut j = op0_in_nlimbs;
while j > 0 {
j -= 1;
let op0_val = op0.load_l(j);
op0.store_l(j, 0);
let mut last_prod_high: LimbType = 0;
let mut carry = 0;
let result_nlimbs = op0_nlimbs - j;
for k in 0..j.min(result_nlimbs) {
let op1_val = op0.load_l(k);
let prod = ct_mul_l_l(op0_val, op1_val);
let mut result_val = op0.load_l(j + k);
let carry0 = black_box_l(last_prod_high >> (LIMB_BITS - 1));
last_prod_high = last_prod_high.wrapping_mul(2);
debug_assert!(last_prod_high <= !3);
debug_assert!(carry <= 2 || last_prod_high <= !5);
last_prod_high += carry;
let carry1;
(carry1, result_val) = ct_add_l_l(result_val, last_prod_high);
last_prod_high = prod.high();
let carry2 = black_box_l(prod.low() >> (LIMB_BITS - 1));
let prod_low = prod.low().wrapping_mul(2);
let carry3;
(carry3, result_val) = ct_add_l_l(result_val, prod_low);
carry = carry0 + carry1 + carry2 + carry3;
debug_assert!(last_prod_high <= !1);
debug_assert!(carry <= 2 || last_prod_high <= !2);
if k != result_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(j + k, result_val);
} else {
op0.store_l(j + k, result_val & result_high_mask);
}
}
if j >= result_nlimbs {
continue;
}
let prod = ct_mul_l_l(op0_val, op0_val);
let mut result_val = op0.load_l(2 * j);
let carry0 = black_box_l(last_prod_high >> (LIMB_BITS - 1));
last_prod_high = last_prod_high.wrapping_mul(2);
debug_assert!(last_prod_high <= !3);
debug_assert!(carry <= 2 || last_prod_high <= !5);
last_prod_high += carry;
let carry1;
(carry1, result_val) = ct_add_l_l(result_val, last_prod_high);
last_prod_high = prod.high();
let carry2;
(carry2, result_val) = ct_add_l_l(result_val, prod.low());
carry = carry0 + carry1 + carry2;
if j != result_nlimbs - 1 {
op0.store_l_full(2 * j, result_val);
} else {
op0.store_l(2 * j, result_val & result_high_mask);
}
for k in j + 1..result_nlimbs {
let mut result_val = op0.load_l(j + k);
let carry0;
(carry0, result_val) = ct_add_l_l(result_val, last_prod_high);
last_prod_high = 0;
let carry1;
(carry1, result_val) = ct_add_l_l(result_val, carry);
carry = carry0 + carry1;
if k != result_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(j + k, result_val);
} else {
op0.store_l(j + k, result_val & result_high_mask);
}
}
}
}
#[cfg(test)]
fn test_ct_square_trunc_mp<T0: MpMutUIntSlice>() {
extern crate alloc;
use super::limb::LIMB_BYTES;
use alloc::vec::Vec;
fn square_by_mul<T0: MpMutUIntSlice>(
op0: &[T0::BackingSliceElementType],
op0_in_len: usize,
) -> Vec<T0::BackingSliceElementType> {
let mut _result = Vec::from(op0);
let mut result = T0::from_slice(&mut _result).unwrap();
let mut op0 = Vec::from(op0);
let mut op0 = T0::from_slice(&mut op0).unwrap();
let op0 = op0.shrink_to(op0_in_len);
ct_mul_trunc_mp_mp(&mut result, op0_in_len, &op0);
drop(result);
_result
}
let mut _op0 = tst_mk_mp_backing_vec!(T0, 5 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES);
drop(op0);
assert_eq!(_op0, expected);
let mut _op0 = tst_mk_mp_backing_vec!(T0, 3 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES);
drop(op0);
assert_eq!(_op0, expected);
let mut _op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES);
drop(op0);
assert_eq!(_op0, expected);
if !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
return;
}
let mut _op0 = tst_mk_mp_backing_vec!(T0, 4 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES);
drop(op0);
assert_eq!(_op0, expected);
let mut _op0 = tst_mk_mp_backing_vec!(T0, 3 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES);
drop(op0);
assert_eq!(_op0, expected);
let mut _op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0 >> 2 * 8);
drop(op0);
let expected = square_by_mul::<T0>(&_op0, 2 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut _op0).unwrap();
ct_square_trunc_mp(&mut op0, 2 * LIMB_BYTES - 1);
drop(op0);
assert_eq!(_op0, expected);
}
#[test]
fn test_ct_square_trunc_be() {
use super::limbs_buffer::MpMutBigEndianUIntByteSlice;
test_ct_square_trunc_mp::<MpMutBigEndianUIntByteSlice>()
}
#[test]
fn test_ct_square_trunc_le() {
use super::limbs_buffer::MpMutLittleEndianUIntByteSlice;
test_ct_square_trunc_mp::<MpMutLittleEndianUIntByteSlice>()
}
#[test]
fn test_ct_square_trunc_ne() {
use super::limbs_buffer::MpMutNativeEndianUIntLimbsSlice;
test_ct_square_trunc_mp::<MpMutNativeEndianUIntLimbsSlice>()
}
pub fn ct_mul_trunc_mp_l<T0: MpMutUInt>(op0: &mut T0, op0_in_len: usize, op1: LimbType) {
debug_assert!(op0_in_len <= op0.len());
let result_high_mask = op0.partial_high_mask();
let op0_nlimbs = op0.nlimbs();
op0.clear_bytes_above(op0_in_len);
let op0_in_nlimbs = ct_mp_nlimbs(op0_in_len);
if op0_in_len == 0 {
return;
}
let mut carry = 0;
for j in 0..op0_in_nlimbs {
let op0_val = op0.load_l(j);
let result_val;
(carry, result_val) = ct_mul_add_l_l_l_c(0, op0_val, op1, carry);
if j != op0_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(j, result_val);
} else {
op0.store_l(j, result_val & result_high_mask);
}
}
if op0_in_nlimbs != op0_nlimbs {
if op0_in_nlimbs != op0_nlimbs - 1 || !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
op0.store_l_full(op0_in_nlimbs, carry);
} else {
op0.store_l(op0_in_nlimbs, carry & result_high_mask);
}
}
}
#[cfg(test)]
fn test_ct_mul_trunc_mp_l<T0: MpMutUIntSlice>() {
use super::limb::LIMB_BYTES;
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
let op1 = 0;
ct_mul_trunc_mp_l(&mut op0, 2 * LIMB_BYTES, op1);
assert_eq!(op0.load_l(0), 0);
assert_eq!(op0.load_l(1), 0);
let mut op0 = tst_mk_mp_backing_vec!(T0, 3 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
let op1 = 2;
ct_mul_trunc_mp_l(&mut op0, 2 * LIMB_BYTES, op1);
assert_eq!(op0.load_l(0), !1);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 1);
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES);
let mut op0 = T0::from_slice(&mut op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
let op1 = 2;
ct_mul_trunc_mp_l(&mut op0, 2 * LIMB_BYTES, op1);
assert_eq!(op0.load_l(0), !1);
assert_eq!(op0.load_l(1), !0);
if !T0::SUPPORTS_UNALIGNED_BUFFER_LENGTHS {
return;
}
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES + 1);
let mut op0 = T0::from_slice(&mut op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0);
let op1 = 2;
ct_mul_trunc_mp_l(&mut op0, 2 * LIMB_BYTES, op1);
assert_eq!(op0.load_l(0), !1);
assert_eq!(op0.load_l(1), !0);
assert_eq!(op0.load_l(2), 1);
let mut op0 = tst_mk_mp_backing_vec!(T0, 2 * LIMB_BYTES - 1);
let mut op0 = T0::from_slice(&mut op0).unwrap();
op0.store_l(0, !0);
op0.store_l(1, !0 >> 8);
let op1 = 2;
ct_mul_trunc_mp_l(&mut op0, 2 * LIMB_BYTES - 1, op1);
assert_eq!(op0.load_l(0), !1);
assert_eq!(op0.load_l(1), !0 >> 8);
}
#[test]
fn test_ct_mul_trunc_be_l() {
use super::limbs_buffer::MpMutBigEndianUIntByteSlice;
test_ct_mul_trunc_mp_l::<MpMutBigEndianUIntByteSlice>()
}
#[test]
fn test_ct_mul_trunc_le_l() {
use super::limbs_buffer::MpMutLittleEndianUIntByteSlice;
test_ct_mul_trunc_mp_l::<MpMutLittleEndianUIntByteSlice>()
}
#[test]
fn test_ct_mul_trunc_ne_l() {
use super::limbs_buffer::MpMutNativeEndianUIntLimbsSlice;
test_ct_mul_trunc_mp_l::<MpMutNativeEndianUIntLimbsSlice>()
}