use super::{HeaplessBigInt, zero};
use crate::MachineWord;
use const_num_traits::{Bounded, Personality, PersonalityTag};
use core::marker::PhantomData;
use core::ops::{Shl, ShlAssign, Shr, ShrAssign};
fn shl_wp<T: MachineWord, const CAP: usize, P: Personality>(
value: HeaplessBigInt<T, CAP, P>,
bits: usize,
) -> HeaplessBigInt<T, CAP, P> {
let word_bits = core::mem::size_of::<T>() * 8;
let word_shift = bits / word_bits;
let bit_shift = bits % word_bits;
let out_len = value.len as usize;
let mut limbs = [zero::<T>(); CAP];
let mut i = 0;
while i < out_len {
let dst_lo = i + word_shift;
if dst_lo < out_len {
limbs[dst_lo] |= value.limbs[i] << bit_shift;
if bit_shift > 0 {
let dst_hi = dst_lo + 1;
if dst_hi < out_len {
limbs[dst_hi] |= value.limbs[i] >> (word_bits - bit_shift);
}
}
}
i += 1;
}
HeaplessBigInt {
limbs,
len: value.len,
_p: PhantomData,
}
}
fn shr_wp<T: MachineWord, const CAP: usize, P: Personality>(
value: HeaplessBigInt<T, CAP, P>,
bits: usize,
) -> HeaplessBigInt<T, CAP, P> {
let word_bits = core::mem::size_of::<T>() * 8;
let word_shift = bits / word_bits;
let bit_shift = bits % word_bits;
let n = value.len as usize;
let mut limbs = [zero::<T>(); CAP];
let mut i = 0;
while i < n {
let src_lo = i + word_shift;
let lo = if src_lo < n {
value.limbs[src_lo] >> bit_shift
} else {
zero::<T>()
};
let hi = if bit_shift > 0 && src_lo + 1 < n {
value.limbs[src_lo + 1] << (word_bits - bit_shift)
} else {
zero::<T>()
};
limbs[i] = lo | hi;
i += 1;
}
HeaplessBigInt {
limbs,
len: value.len,
_p: PhantomData,
}
}
#[inline]
fn ct_mask<T: MachineWord>(choice_bit: u8) -> T {
let bit = core::hint::black_box(choice_bit & 1);
let bit_t = <T as core::convert::From<u8>>::from(bit);
<T as core::ops::Mul>::mul(bit_t, <T as Bounded>::max_value())
}
pub(crate) fn const_shl_ct<T: MachineWord, const CAP: usize, P: Personality>(
value: HeaplessBigInt<T, CAP, P>,
bits: usize,
) -> HeaplessBigInt<T, CAP, P> {
let n = value.len as usize;
let layers = core::mem::size_of::<usize>() * 8;
let mut target = value;
let mut k = 0;
while k < layers {
let shifted = shl_wp(target, 1usize << k);
let mask = ct_mask::<T>(((bits >> k) & 1) as u8);
let mut i = 0;
while i < n {
let diff = target.limbs[i] ^ shifted.limbs[i];
target.limbs[i] ^= mask & diff;
i += 1;
}
k += 1;
}
target
}
pub(crate) fn const_shr_ct<T: MachineWord, const CAP: usize, P: Personality>(
value: HeaplessBigInt<T, CAP, P>,
bits: usize,
) -> HeaplessBigInt<T, CAP, P> {
let n = value.len as usize;
let layers = core::mem::size_of::<usize>() * 8;
let mut target = value;
let mut k = 0;
while k < layers {
let shifted = shr_wp(target, 1usize << k);
let mask = ct_mask::<T>(((bits >> k) & 1) as u8);
let mut i = 0;
while i < n {
let diff = target.limbs[i] ^ shifted.limbs[i];
target.limbs[i] ^= mask & diff;
i += 1;
}
k += 1;
}
target
}
pub(crate) fn ct_shl<T: MachineWord, const CAP: usize, P: Personality>(
value: HeaplessBigInt<T, CAP, P>,
amount: u32,
) -> HeaplessBigInt<T, CAP, P> {
const_shl_ct(value, amount as usize)
}
impl<T: MachineWord, const CAP: usize, P: Personality> Shl<u32> for HeaplessBigInt<T, CAP, P> {
type Output = Self;
fn shl(self, bits: u32) -> Self::Output {
match P::TAG {
PersonalityTag::Ct => const_shl_ct(self, bits as usize),
PersonalityTag::Nct => {
let value_bits = self.len as u32 * (core::mem::size_of::<T>() as u32 * 8);
if bits >= value_bits {
Self::new_zero_with_len(self.len())
} else {
self << (bits as usize)
}
}
}
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> Shr<u32> for HeaplessBigInt<T, CAP, P> {
type Output = Self;
fn shr(self, bits: u32) -> Self::Output {
match P::TAG {
PersonalityTag::Ct => const_shr_ct(self, bits as usize),
PersonalityTag::Nct => {
let value_bits = self.len as u32 * (core::mem::size_of::<T>() as u32 * 8);
if bits >= value_bits {
Self::new_zero_with_len(0)
} else {
self >> (bits as usize)
}
}
}
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> Shl<usize> for HeaplessBigInt<T, CAP, P> {
type Output = Self;
fn shl(self, bits: usize) -> Self::Output {
match P::TAG {
PersonalityTag::Nct => shl_wp(self, bits),
PersonalityTag::Ct => const_shl_ct(self, bits),
}
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShlAssign<usize>
for HeaplessBigInt<T, CAP, P>
{
fn shl_assign(&mut self, bits: usize) {
*self = *self << bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShrAssign<usize>
for HeaplessBigInt<T, CAP, P>
{
fn shr_assign(&mut self, bits: usize) {
*self = *self >> bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> Shr<usize> for HeaplessBigInt<T, CAP, P> {
type Output = Self;
fn shr(self, bits: usize) -> Self::Output {
match P::TAG {
PersonalityTag::Ct => const_shr_ct(self, bits),
PersonalityTag::Nct => {
let word_bits = core::mem::size_of::<T>() * 8;
let word_shift = bits / word_bits;
let bit_shift = bits % word_bits;
let mut limbs = [zero::<T>(); CAP];
if word_shift >= self.len as usize {
return Self {
limbs,
len: 0,
_p: PhantomData,
};
}
let out_len = self.len as usize - word_shift;
let mut i = 0;
while i < out_len {
let src_lo = i + word_shift;
let lo = self.limbs[src_lo] >> bit_shift;
let hi = if bit_shift > 0 && src_lo + 1 < self.len as usize {
self.limbs[src_lo + 1] << (word_bits - bit_shift)
} else {
zero::<T>()
};
limbs[i] = lo | hi;
i += 1;
}
Self {
limbs,
len: out_len as u16,
_p: PhantomData,
}
}
}
}
}
macro_rules! shift_operand_forms {
($imp:ident, $method:ident, $op:tt, $scalar:ty) => {
impl<T: MachineWord, const CAP: usize, P: Personality> $imp<&$scalar>
for HeaplessBigInt<T, CAP, P>
{
type Output = HeaplessBigInt<T, CAP, P>;
fn $method(self, bits: &$scalar) -> Self::Output {
self $op *bits
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> $imp<$scalar>
for &HeaplessBigInt<T, CAP, P>
{
type Output = HeaplessBigInt<T, CAP, P>;
fn $method(self, bits: $scalar) -> Self::Output {
*self $op bits
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> $imp<&$scalar>
for &HeaplessBigInt<T, CAP, P>
{
type Output = HeaplessBigInt<T, CAP, P>;
fn $method(self, bits: &$scalar) -> Self::Output {
*self $op *bits
}
}
};
}
shift_operand_forms!(Shl, shl, <<, usize);
shift_operand_forms!(Shl, shl, <<, u32);
shift_operand_forms!(Shr, shr, >>, usize);
shift_operand_forms!(Shr, shr, >>, u32);
impl<T: MachineWord, const CAP: usize, P: Personality> ShlAssign<u32>
for HeaplessBigInt<T, CAP, P>
{
fn shl_assign(&mut self, bits: u32) {
*self = *self << bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShrAssign<u32>
for HeaplessBigInt<T, CAP, P>
{
fn shr_assign(&mut self, bits: u32) {
*self = *self >> bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShlAssign<&usize>
for HeaplessBigInt<T, CAP, P>
{
fn shl_assign(&mut self, bits: &usize) {
*self = *self << *bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShrAssign<&usize>
for HeaplessBigInt<T, CAP, P>
{
fn shr_assign(&mut self, bits: &usize) {
*self = *self >> *bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShlAssign<&u32>
for HeaplessBigInt<T, CAP, P>
{
fn shl_assign(&mut self, bits: &u32) {
*self = *self << *bits;
}
}
impl<T: MachineWord, const CAP: usize, P: Personality> ShrAssign<&u32>
for HeaplessBigInt<T, CAP, P>
{
fn shr_assign(&mut self, bits: &u32) {
*self = *self >> *bits;
}
}
#[cfg(test)]
mod ct_shl_tests {
use super::{HeaplessBigInt, ct_shl};
use const_num_traits::{Ct, Nct};
type HC = HeaplessBigInt<u8, 4, Ct>; type HN = HeaplessBigInt<u8, 4, Nct>;
#[test]
fn ct_shl_matches_plain_shift_all_amounts() {
for &raw in &[1u32, 0x1234_5678, 0xFFFF_FFFF, 0x8000_0000] {
let v = HC::from(raw);
for amount in 0u32..=40 {
assert_eq!(
ct_shl(v, amount),
v << (amount as usize),
"ct_shl({raw:#x}, {amount})"
);
}
}
}
#[test]
fn ct_shifts_match_nct_reference() {
let cases = [
[1u8, 0, 0, 0],
[0x78, 0x56, 0x34, 0x12],
[0xFF, 0xFF, 0xFF, 0xFF],
[0, 0, 0, 0x80],
];
for a in cases {
for amount in 0usize..=40 {
assert_eq!(
(HC::from_limbs(a, 4) << amount).all_limbs(),
(HN::from_limbs(a, 4) << amount).all_limbs(),
"shl {a:?} << {amount}"
);
assert_eq!(
(HC::from_limbs(a, 4) >> amount).all_limbs(),
(HN::from_limbs(a, 4) >> amount).all_limbs(),
"shr {a:?} >> {amount}"
);
}
}
}
}