Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 71 additions & 19 deletions src/target/r1cs/trans.rs
Original file line number Diff line number Diff line change
Expand Up @@ -644,6 +644,58 @@ impl<'cfg> ToR1cs<'cfg> {
acc
}

/// Subtract equally sized, least-significant-bit-first bit vectors modulo their width.
fn bv_sub_bits(&mut self, xs: &[TermLc], ys: &[TermLc]) -> Vec<TermLc> {
assert_eq!(xs.len(), ys.len());
let mut borrow = self.zero.clone();
let mut difference = Vec::with_capacity(xs.len());
for (x, y) in xs.iter().zip(ys) {
difference.push(self.nary_xor(vec![x.clone(), y.clone(), borrow.clone()].into_iter()));

// The next borrow is majority(!x, y, borrow).
let not_x = self.bool_not(x);
let y_or_borrow = self.nary_or(vec![y.clone(), borrow.clone()].into_iter());
let from_x = self.nary_and(vec![not_x, y_or_borrow].into_iter());
let from_borrow = self.nary_and(vec![y.clone(), borrow].into_iter());
borrow = self.nary_or(vec![from_x, from_borrow].into_iter());
}
difference
}

/// Unsigned restoring division over bit wires. The quotient and remainder are uniquely
/// determined without relying on a product that can wrap in the scalar field.
fn bv_udivrem_bits(
&mut self,
dividend: Vec<TermLc>,
divisor: Vec<TermLc>,
) -> (Vec<TermLc>, Vec<TermLc>) {
assert_eq!(dividend.len(), divisor.len());
let width = dividend.len();
let mut quotient = vec![self.zero.clone(); width];
let mut remainder = vec![self.zero.clone(); width + 1];
let mut extended_divisor = divisor;
extended_divisor.push(self.zero.clone());

for i in (0..width).rev() {
let mut shifted = Vec::with_capacity(width + 1);
shifted.push(dividend[i].clone());
shifted.extend(remainder.into_iter().take(width));

let subtract =
self.bv_bitwise_greater(shifted.clone(), extended_divisor.clone(), false);
let difference = self.bv_sub_bits(&shifted, &extended_divisor);
remainder = difference
.into_iter()
.zip(shifted)
.map(|(difference, original)| self.ite(subtract.clone(), difference, &original))
.collect();
quotient[i] = subtract;
}

remainder.truncate(width);
(quotient, remainder)
}

/// Shift `x` left by `2^(2^y)`, if bit-valued `c` is true.
fn const_pow_shift_bv_lit(&mut self, x: &TermLc, y: usize, c: TermLc) -> TermLc {
let two_to_the_y = 1usize.checked_shl(y as u32).unwrap();
Expand Down Expand Up @@ -844,25 +896,9 @@ impl<'cfg> ToR1cs<'cfg> {
self.set_bv_bits(bv, bits);
}
BvBinOp::Udiv | BvBinOp::Urem => {
let a_bv_term = term![Op::PfToBv(n); a.0.clone()];
let b_bv_term = term![Op::PfToBv(n); b.0.clone()];
let q_term = term![Op::new_ubv_to_pf(self.field.clone()); term![BV_UDIV; a_bv_term.clone(), b_bv_term.clone()]];
let r_term = term![Op::new_ubv_to_pf(self.field.clone()); term![BV_UREM; a_bv_term, b_bv_term]];
let q = self.fresh_wit("div_q", q_term);
let r = self.fresh_wit("div_r", r_term);
let qb = self.bitify("div_q", &q, n, false);
let rb = self.bitify("div_r", &r, n, false);
self.constraint(q.1.clone(), b.1.clone(), (a - &r).1);
// b == 0 -> q == M // b != 0 or q == M
// b != 0 -> r < b // b == 0 or r < b
// so, since we don't care about b == 0,
// q == M or r < b
// not(q != M and r >= b)
let r_ge_b = self.bv_greater(r, b, n, false);
let max = self.r1cs.modulus.new_v((Integer::from(1) << n) - 1);
let q_eq_max = self.is_zero(q - &max);
let q_ne_max = self.bool_not(&q_eq_max);
self.constraint(r_ge_b.1, q_ne_max.1, self.r1cs.zero());
let dividend = self.get_bv_bits(&bv.cs()[0]);
let divisor = self.get_bv_bits(&bv.cs()[1]);
let (qb, rb) = self.bv_udivrem_bits(dividend, divisor);
let bits = match o {
BvBinOp::Udiv => qb,
BvBinOp::Urem => rb,
Expand Down Expand Up @@ -1363,6 +1399,22 @@ pub mod test {
]);
}

#[test]
fn div128_test() {
init();
let high: Integer = Integer::from(1) << 127;
const_test(term![
Op::Eq;
term![Op::BvBinOp(BvBinOp::Udiv); bv_lit(1, 128), bv_lit(high.clone(), 128)],
bv_lit(0, 128)
]);
const_test(term![
Op::Eq;
term![Op::BvBinOp(BvBinOp::Urem); bv_lit(1, 128), bv_lit(high, 128)],
bv_lit(1, 128)
]);
}

#[test]
fn sh_test() {
init();
Expand Down