use snarkvm_fields::{FieldParameters, PrimeField};
use snarkvm_r1cs::{errors::SynthesisError, Assignment, ConstraintSystem, LinearCombination};
use snarkvm_utilities::biginteger::{BigInteger, BigInteger256};
use crate::{
bits::boolean::{AllocatedBit, Boolean},
integers::uint::UInt,
traits::{
alloc::AllocGadget,
eq::{ConditionalEqGadget, EqGadget},
integers::{Integer, Pow},
select::CondSelectGadget,
},
uint_impl_common,
UnsignedIntegerError,
};
uint_impl_common!(UInt128, u128, 128);
impl UInt for UInt128 {
fn negate(&self) -> Self {
Self {
bits: self.bits.clone(),
negated: true,
value: self.value,
}
}
fn rotr(&self, by: usize) -> Self {
let by = by % 128;
let new_bits = self
.bits
.iter()
.skip(by)
.chain(self.bits.iter())
.take(128)
.cloned()
.collect();
Self {
bits: new_bits,
negated: false,
value: self.value.map(|v| v.rotate_right(by as u32) as u128),
}
}
fn addmany<F: PrimeField, CS: ConstraintSystem<F>>(mut cs: CS, operands: &[Self]) -> Result<Self, SynthesisError> {
assert!(F::Parameters::MODULUS_BITS <= 253);
assert!(operands.len() >= 2);
let mut max_value = BigInteger256::from_u128(u128::max_value());
max_value.muln(operands.len() as u32);
let mut big_result_value = Some(BigInteger256::default());
let mut lc = LinearCombination::zero();
let mut all_constants = true;
for op in operands {
match op.value {
Some(val) => {
if op.negated {
big_result_value
.as_mut()
.map(|v| v.sub_noborrow(&BigInteger256::from_u128(val)));
} else {
big_result_value
.as_mut()
.map(|v| v.add_nocarry(&BigInteger256::from_u128(val)));
}
}
None => {
big_result_value = None;
}
}
let mut coeff = F::one();
for bit in &op.bits {
match *bit {
Boolean::Is(ref bit) => {
all_constants = false;
if op.negated {
lc = lc - (coeff, bit.get_variable());
} else {
lc += (coeff, bit.get_variable());
}
}
Boolean::Not(ref bit) => {
all_constants = false;
if op.negated {
lc = lc - (coeff, CS::one()) + (coeff, bit.get_variable());
} else {
lc = lc + (coeff, CS::one()) - (coeff, bit.get_variable());
}
}
Boolean::Constant(bit) => {
if bit {
if op.negated {
lc = lc - (coeff, CS::one());
} else {
lc += (coeff, CS::one());
}
}
}
}
coeff.double_in_place();
}
}
let modular_value = big_result_value.map(|v| v.to_u128());
if all_constants {
if let Some(val) = modular_value {
return Ok(Self::constant(val));
}
}
let mut result_bits = vec![];
let mut coeff = F::one();
let mut i = 0;
while !max_value.is_zero() {
let b = AllocatedBit::alloc(cs.ns(|| format!("result bit_gadget {}", i)), || {
big_result_value.map(|v| v.get_bit(i)).get()
})?;
lc = lc - (coeff, b.get_variable());
if result_bits.len() < 128 {
result_bits.push(b.into());
}
max_value.div2();
i += 1;
coeff.double_in_place();
}
cs.enforce(|| "modular addition", |lc| lc, |lc| lc, |_| lc);
Ok(Self {
bits: result_bits,
negated: false,
value: modular_value,
})
}
fn mul<F: PrimeField, CS: ConstraintSystem<F>>(
&self,
mut cs: CS,
other: &Self,
) -> Result<Self, UnsignedIntegerError> {
let is_constant = Boolean::constant(Self::result_is_constant(self, other));
let constant_result = Self::constant(0u128);
let allocated_result = Self::alloc(&mut cs.ns(|| "allocated_0u128"), || Ok(0u128))?;
let zero_result = Self::conditionally_select(
&mut cs.ns(|| "constant_or_allocated"),
&is_constant,
&constant_result,
&allocated_result,
)?;
let mut left_shift = self.clone();
let partial_products = other
.bits
.iter()
.enumerate()
.map(|(i, bit)| {
let current_left_shift = left_shift.clone();
left_shift = Self::addmany(&mut cs.ns(|| format!("shift_left_{}", i)), &[
left_shift.clone(),
left_shift.clone(),
])
.unwrap();
Self::conditionally_select(
&mut cs.ns(|| format!("calculate_product_{}", i)),
bit,
¤t_left_shift,
&zero_result,
)
.unwrap()
})
.collect::<Vec<Self>>();
Self::addmany(&mut cs.ns(|| "partial_products"), &partial_products)
.map_err(UnsignedIntegerError::SynthesisError)
}
}
impl<F: PrimeField> Pow<F> for UInt128 {
type ErrorType = UnsignedIntegerError;
fn pow<CS: ConstraintSystem<F>>(&self, mut cs: CS, other: &Self) -> Result<Self, Self::ErrorType> {
let is_constant = Boolean::constant(Self::result_is_constant(self, other));
let constant_result = Self::constant(1u128);
let allocated_result = Self::alloc(&mut cs.ns(|| "allocated_1u128"), || Ok(1u128))?;
let mut result = Self::conditionally_select(
&mut cs.ns(|| "constant_or_allocated"),
&is_constant,
&constant_result,
&allocated_result,
)?;
for (i, bit) in other.bits.iter().rev().enumerate() {
let found_one = Boolean::Constant(result.eq(&Self::constant(1u128)));
let cond1 = Boolean::and(cs.ns(|| format!("found_one_{}", i)), &bit.not(), &found_one)?;
let square = result.mul(cs.ns(|| format!("square_{}", i)), &result).unwrap();
result = Self::conditionally_select(
&mut cs.ns(|| format!("result_or_sqaure_{}", i)),
&cond1,
&result,
&square,
)?;
let mul_by_self = result.mul(cs.ns(|| format!("multiply_by_self_{}", i)), self).unwrap();
result = Self::conditionally_select(
&mut cs.ns(|| format!("mul_by_self_or_result_{}", i)),
bit,
&mul_by_self,
&result,
)?;
}
Ok(result)
}
}