From 032a74a9d5649ae1f8467ae08aac7190df9da380 Mon Sep 17 00:00:00 2001 From: NullWitnessZK <312565654+NullWitnessZK@users.noreply.github.com> Date: Mon, 3 Aug 2026 23:17:33 +0800 Subject: [PATCH] fix(r1cs): constrain wide unsigned division bitwise --- src/target/r1cs/trans.rs | 90 +++++++++++++++++++++++++++++++--------- 1 file changed, 71 insertions(+), 19 deletions(-) diff --git a/src/target/r1cs/trans.rs b/src/target/r1cs/trans.rs index e5b6defd..a4cf9ed7 100644 --- a/src/target/r1cs/trans.rs +++ b/src/target/r1cs/trans.rs @@ -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 { + 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, + divisor: Vec, + ) -> (Vec, Vec) { + 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(); @@ -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, @@ -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();