diff --git a/Cargo.lock b/Cargo.lock index d5eb1780..995b774b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -384,10 +384,12 @@ version = "0.1.0" dependencies = [ "crypto-primitives", "field", + "num-bigint 0.4.8", "num-traits", "poly", "rayon", "spongefish", + "transcript", ] [[package]] diff --git a/crates/circuit/Cargo.toml b/crates/circuit/Cargo.toml index b6057bd6..2c16bfae 100644 --- a/crates/circuit/Cargo.toml +++ b/crates/circuit/Cargo.toml @@ -8,7 +8,7 @@ license.workspace = true [dependencies] blake3.workspace = true common.workspace = true -field = { path = "../field" } +field.workspace = true num-bigint.workspace = true poly.workspace = true num-traits.workspace = true diff --git a/crates/circuit/benches/support/p256.rs b/crates/circuit/benches/support/p256.rs index 3216b382..4e38cf0a 100644 --- a/crates/circuit/benches/support/p256.rs +++ b/crates/circuit/benches/support/p256.rs @@ -1,29 +1,32 @@ use circuit::p256::VERIFY_DIGEST_INPUT_BITS; -use num_bigint::{BigInt, BigUint, Sign}; +use num_bigint::Sign; use num_traits::{One, Zero}; pub const AUX_INPUT_BITS: usize = 6 * 256; -fn from_hex(value: &[u8]) -> BigUint { - BigUint::parse_bytes(value, 16).unwrap() +type S = num_bigint::BigUint; +type R = num_bigint::BigInt; + +fn from_hex(value: &[u8]) -> S { + S::parse_bytes(value, 16).unwrap() } -fn scalar_modulus() -> BigUint { +fn scalar_modulus() -> S { from_hex(b"ffffffff00000000ffffffffffffffffbce6faada7179e84f3b9cac2fc632551") } -fn inverse(value: &BigUint, modulus: &BigUint) -> BigUint { - let mut t = BigInt::zero(); - let mut new_t = BigInt::one(); - let mut r = BigInt::from(modulus.clone()); - let mut new_r = BigInt::from(value.clone()); +fn inverse(value: &S, modulus: &S) -> S { + let mut t = R::zero(); + let mut new_t = R::one(); + let mut r = R::from(modulus.clone()); + let mut new_r = R::from(value.clone()); while !new_r.is_zero() { let quotient = &r / &new_r; (t, new_t) = (new_t.clone(), t - "ient * new_t); (r, new_r) = (new_r.clone(), r - quotient * new_r); } - assert_eq!(r, BigInt::one()); - let modulus = BigInt::from(modulus.clone()); + assert_eq!(r, R::one()); + let modulus = R::from(modulus.clone()); let mut t = t % &modulus; if t.sign() == Sign::Minus { t += modulus; @@ -32,7 +35,7 @@ fn inverse(value: &BigUint, modulus: &BigUint) -> BigUint { } /// P-256 inputs for d = k = 1: Q = G, r = G.x, and s = z + r*d. -pub fn valid_aux_input(digest: &BigUint) -> Box<[bool; AUX_INPUT_BITS]> { +pub fn valid_aux_input(digest: &S) -> Box<[bool; AUX_INPUT_BITS]> { let modulus = scalar_modulus(); let gx = from_hex(b"6b17d1f2e12c4247f8bce6e563a440f277037d812deb33a0f4a13945d898c296"); let gy = from_hex(b"4fe342e2fe1a7f9b8ee7eb4a7c0f9e162bce33576b315ececbb6406837bf51f5"); @@ -53,7 +56,7 @@ pub fn valid_aux_input(digest: &BigUint) -> Box<[bool; AUX_INPUT_BITS]> { /// A small valid ECDSA instance with z = 1. pub fn valid_input() -> Box<[bool; VERIFY_DIGEST_INPUT_BITS]> { - let digest = BigUint::one(); + let digest = S::one(); let aux = valid_aux_input(&digest); let bits: Box<[bool]> = (0..VERIFY_DIGEST_INPUT_BITS) .map(|index| { diff --git a/crates/circuit/src/constraints.rs b/crates/circuit/src/constraints.rs index 732b6baa..c3463f6c 100644 --- a/crates/circuit/src/constraints.rs +++ b/crates/circuit/src/constraints.rs @@ -4,11 +4,11 @@ //! witness, prefixed by a constant one, to the integer witness. Its first row //! is the implicit integer constant one. `A`, `B`, and `C` then encode the //! rank-1 constraints `(A z) * (B z) = C z` over that integer witness. Every -//! integer coefficient is an arbitrary-precision signed [`BitzRing`]. +//! integer coefficient is an arbitrary-precision signed [`BitzConstraintRing`]. use crate::witgen::PackedWitness; use crate::{BoolWitness, Circuit, HintResult, PackedBits, ScalarBits, WitnessContext}; -use common::{BitzRing, BitzSemiring}; +use common::{BitzConstraintRing, BitzSemiring}; use num_traits::Zero; use rayon::prelude::*; use std::array; @@ -232,7 +232,7 @@ pub enum ConstraintMatrixShapeError { AssignmentLengthMismatch { m_rows: usize, r1cs_columns: usize }, } -impl ConstraintMatrices { +impl ConstraintMatrices { /// Checks that A, B, and C share a shape and consume the assignment /// produced by M. pub fn validate_shape(&self) -> Result<(), ConstraintMatrixShapeError> { @@ -284,7 +284,9 @@ impl ConstraintMatrices { c: self.c.map_values_with(&map), } } +} +impl ConstraintMatrices { /// Applies `M` to a packed Boolean witness. /// /// The returned vector starts with the implicit constant one and is the @@ -498,7 +500,7 @@ impl AddAssign for LinearCombination { } } -impl Neg for LinearCombination { +impl Neg for LinearCombination { type Output = Self; fn neg(mut self) -> Self::Output { @@ -510,7 +512,7 @@ impl Neg for LinearCombination { } } -impl Sub for LinearCombination { +impl Sub for LinearCombination { type Output = Self; fn sub(self, rhs: Self) -> Self::Output { @@ -518,7 +520,7 @@ impl Sub for LinearCombination { } } -impl SubAssign for LinearCombination { +impl SubAssign for LinearCombination { fn sub_assign(&mut self, rhs: Self) { *self += -rhs; } @@ -715,7 +717,7 @@ fn bool_sparse_row(value: BoolLinearCombination) -> SparseBoolRow { } } -impl Circuit for ConstraintGenerator { +impl Circuit for ConstraintGenerator { type Bool = BoolLinearCombination; type Coefficient = R; type Z = LinearCombination; diff --git a/crates/circuit/src/lib.rs b/crates/circuit/src/lib.rs index 59024b4d..2ccc5038 100644 --- a/crates/circuit/src/lib.rs +++ b/crates/circuit/src/lib.rs @@ -538,21 +538,40 @@ pub trait Circuit { ) -> Self::Z; } -pub trait Bits { +pub trait BitWidth { /// Determines the fewest bits necessary to express this value - fn bits(&self) -> u64; + fn bit_width(&self) -> u64; } pub trait IntoWords { + /// Convert a value into u64 words, lowest limb first fn into_words(self) -> [u64; LIMBS]; } -impl Bits for num_bigint::BigUint { - fn bits(&self) -> u64 { +impl BitWidth for u128 { + fn bit_width(&self) -> u64 { + u128::bit_width(*self) as u64 + } +} + +impl BitWidth for num_bigint::BigUint { + fn bit_width(&self) -> u64 { num_bigint::BigUint::bits(self) } } +impl IntoWords for u128 { + fn into_words(self) -> [u64; LIMBS] { + const { + assert!(LIMBS >= 2); + } + let mut result = [0; LIMBS]; + result[0] = self as u64; + result[1] = (self >> 64) as u64; + result + } +} + impl IntoWords for num_bigint::BigUint { fn into_words(self) -> [u64; LIMBS] { let mut digits = self.iter_u64_digits(); diff --git a/crates/circuit/src/matrix_products.rs b/crates/circuit/src/matrix_products.rs index 30e2ca46..e536ea03 100644 --- a/crates/circuit/src/matrix_products.rs +++ b/crates/circuit/src/matrix_products.rs @@ -5,9 +5,10 @@ //! and `C(Mw)` vectors during witness generation. This module subsequently //! reduces those vectors modulo a runtime modulus. Large batches use Rayon; //! small batches stay sequential to avoid scheduling overhead. +// TODO(alex): Should this be generalized over [`BitzClaimField`] as well? use crate::witgen::Z as Integer; -use crate::{Bits, IntoWords}; +use crate::{BitWidth, IntoWords}; use num_traits::{One, Zero}; use rayon::prelude::*; use std::cmp::Ordering; @@ -21,14 +22,16 @@ pub struct RuntimeModulus { impl RuntimeModulus { /// Validates and stores a runtime modulus. - pub fn new(modulus: S) -> Result { + pub fn new( + modulus: S, + ) -> Result { if PRIME_LIMBS == 0 { return Err("a runtime field needs at least one limb"); } if modulus <= S::one() { return Err("the modulus must be greater than one"); } - if modulus.bits() > (PRIME_LIMBS as u64) * 64 { + if modulus.bit_width() > (PRIME_LIMBS as u64) * 64 { return Err("the modulus does not fit the selected limb count"); } Ok(Self { @@ -264,6 +267,16 @@ impl ModularVector { } } +impl From<&ModularVector<2>> for Vec { + fn from(values: &ModularVector<2>) -> Self { + values + .values() + .iter() + .map(|&[low, high]| field::FqDefault::from_limbs(low, high)) + .collect() + } +} + /// Dense runtime-field `A(Mw)`, `B(Mw)`, and `C(Mw)` vectors. #[derive(Clone, Debug, Eq, PartialEq)] pub struct MatrixProducts { diff --git a/crates/circuit/src/matrix_sparse.rs b/crates/circuit/src/matrix_sparse.rs index 680d6e93..fa29a982 100644 --- a/crates/circuit/src/matrix_sparse.rs +++ b/crates/circuit/src/matrix_sparse.rs @@ -8,7 +8,7 @@ use crate::constraints::{ConstraintMatrices, SparseMatrix}; use crate::matrix_products::{RuntimeModulus, StoredInteger}; use crate::matrix_wengert::{add_mod_words, montgomery_mul_2, neg_mod_words}; -use common::BitzRing; +use common::BitzConstraintRing; use crypto_bigint::modular::{FixedMontyForm, FixedMontyParams}; use crypto_bigint::{Odd, U128}; use rayon::prelude::*; @@ -74,7 +74,7 @@ impl MaterializedAbc { /// Transposes and stores the integer matrices without choosing a modulus. pub fn from_matrices(matrices: &ConstraintMatrices) -> Self where - R: BitzRing, + R: BitzConstraintRing, StoredInteger: for<'a> From<&'a R>, { let row_count = matrices.a.row_count(); diff --git a/crates/circuit/src/matrix_transpose.rs b/crates/circuit/src/matrix_transpose.rs index 986fb4e3..54edd64c 100644 --- a/crates/circuit/src/matrix_transpose.rs +++ b/crates/circuit/src/matrix_transpose.rs @@ -514,11 +514,10 @@ mod tests { use crate::sha256::{COMPRESSION_HINT_BITS, COMPRESSION_INPUT_BITS, compression_circuit}; use crate::witgen::Witgen; use crate::{BoolRepresentation, BoolWitness, Circuit}; - use num_bigint::BigInt; use super::*; - type R = BigInt; + type R = num_bigint::BigInt; fn example_circuit(circuit: &mut CS, inputs: &[CS::Bool; 3]) { let xy = circuit.xor(inputs[0].clone(), inputs[1].clone()); diff --git a/crates/circuit/src/projection_tests.rs b/crates/circuit/src/projection_tests.rs index c0db7dc5..b8373c1a 100644 --- a/crates/circuit/src/projection_tests.rs +++ b/crates/circuit/src/projection_tests.rs @@ -1,7 +1,7 @@ use std::array; use field::F128; -use num_bigint::{BigInt, BigUint, Sign}; +use num_bigint::Sign; use num_traits::{One, Signed, Zero}; use crate::constraints::BoolLinearCombination; @@ -20,26 +20,29 @@ use crate::stats::Dummy; use crate::witgen::{PackedWitness, ProductWitgen}; use crate::{Circuit, HintResult, PackedBits, ScalarBits, WitnessContext}; -fn from_hex(value: &[u8]) -> BigUint { - BigUint::parse_bytes(value, 16).unwrap() +type S = num_bigint::BigUint; +type R = num_bigint::BigInt; + +fn from_hex(value: &[u8]) -> S { + S::parse_bytes(value, 16).unwrap() } -fn scalar_modulus() -> BigUint { +fn scalar_modulus() -> S { from_hex(b"ffffffff00000000ffffffffffffffffbce6faada7179e84f3b9cac2fc632551") } -fn inverse(value: &BigUint, modulus: &BigUint) -> BigUint { - let mut t = BigInt::zero(); - let mut new_t = BigInt::one(); - let mut r = BigInt::from(modulus.clone()); - let mut new_r = BigInt::from(value.clone()); +fn inverse(value: &S, modulus: &S) -> S { + let mut t = R::zero(); + let mut new_t = R::one(); + let mut r = R::from(modulus.clone()); + let mut new_r = R::from(value.clone()); while !new_r.is_zero() { let quotient = &r / &new_r; (t, new_t) = (new_t.clone(), t - "ient * new_t); (r, new_r) = (new_r.clone(), r - quotient * new_r); } - assert_eq!(r, BigInt::one()); - let modulus = BigInt::from(modulus.clone()); + assert_eq!(r, R::one()); + let modulus = R::from(modulus.clone()); let mut t = t % &modulus; if t.sign() == Sign::Minus { t += modulus; @@ -47,7 +50,7 @@ fn inverse(value: &BigUint, modulus: &BigUint) -> BigUint { t.to_biguint().unwrap() } -fn signature_aux(digest: &BigUint) -> [BigUint; 6] { +fn signature_aux(digest: &S) -> [S; 6] { let modulus = scalar_modulus(); let gx = from_hex(b"6b17d1f2e12c4247f8bce6e563a440f277037d812deb33a0f4a13945d898c296"); let gy = from_hex(b"4fe342e2fe1a7f9b8ee7eb4a7c0f9e162bce33576b315ececbb6406837bf51f5"); @@ -63,7 +66,7 @@ fn signature_aux(digest: &BigUint) -> [BigUint; 6] { } fn p256_input() -> Box<[bool; VERIFY_DIGEST_INPUT_BITS]> { - let digest = BigUint::one(); + let digest = S::one(); let aux = signature_aux(&digest); let bits: Box<[bool]> = (0..VERIFY_DIGEST_INPUT_BITS) .map(|index| { @@ -102,17 +105,17 @@ fn dummy_inputs() -> Box<[Dummy; N]> { .unwrap_or_else(|_| unreachable!("dummy input length is fixed")) } -fn stored_bigint(value: &StoredInteger) -> BigInt { +fn stored_bigint(value: &StoredInteger) -> R { let bytes = value .words() .iter() .flat_map(|word| word.to_le_bytes()) .collect::>(); - BigInt::from_signed_bytes_le(&bytes) + R::from_signed_bytes_le(&bytes) } -fn reduced_words(value: &BigInt, modulus: &BigUint) -> [u64; 2] { - let modulus = BigInt::from(modulus.clone()); +fn reduced_words(value: &R, modulus: &S) -> [u64; 2] { + let modulus = R::from(modulus.clone()); let mut value = value % &modulus; if value.is_negative() { value += modulus; @@ -134,12 +137,12 @@ struct DirectAbcProjector<'a> { integer_witness: &'a PackedWitness, exact: &'a IntegerProducts, reduced: MatrixProducts<2>, - modulus: BigUint, + modulus: S, } impl<'a> DirectAbcProjector<'a> { fn new(label: &'a str, witgen: &'a ProductWitgen) -> Self { - let modulus = (BigUint::one() << 128_usize) - BigUint::from(159_u64); + let modulus = (S::one() << 128_usize) - S::from(159_u64); let runtime_modulus = RuntimeModulus::<2>::new(modulus.clone()).unwrap(); Self { label, @@ -165,15 +168,15 @@ impl<'a> DirectAbcProjector<'a> { impl Circuit for DirectAbcProjector<'_> { type Bool = Dummy; - type Coefficient = BigInt; - type Z = BigInt; + type Coefficient = R; + type Z = R; - fn coefficient_from_le_words(words: &[u64]) -> BigInt { + fn coefficient_from_le_words(words: &[u64]) -> R { let bytes = words .iter() .flat_map(|word| word.to_le_bytes()) .collect::>(); - BigInt::from_bytes_le(Sign::Plus, &bytes) + R::from_bytes_le(Sign::Plus, &bytes) } fn xor(&mut self, _: Dummy, _: Dummy) -> Dummy { @@ -185,7 +188,7 @@ impl Circuit for DirectAbcProjector<'_> { _: H, ) -> ScalarBits where - H: Fn(&dyn WitnessContext) -> HintResult> + H: Fn(&dyn WitnessContext) -> HintResult> + Send + Sync + 'static, @@ -194,18 +197,13 @@ impl Circuit for DirectAbcProjector<'_> { ScalarBits([Dummy; N]) } - fn bitz(&mut self, _: Dummy) -> BigInt { + fn bitz(&mut self, _: Dummy) -> R { let witness = self.next_integer_witness; self.next_integer_witness += 1; - BigInt::from(self.integer_witness.bit(witness + 1)) + R::from(self.integer_witness.bit(witness + 1)) } - fn assert_r1c( - &mut self, - expected_a: BigInt, - expected_b: BigInt, - expected_c: BigInt, - ) { + fn assert_r1c(&mut self, expected_a: R, expected_b: R, expected_c: R) { let row = self.row; self.row += 1; assert_eq!( @@ -246,10 +244,7 @@ impl Circuit for DirectAbcProjector<'_> { ); } - fn sign_extend_z( - &mut self, - value: BigInt, - ) -> BigInt { + fn sign_extend_z(&mut self, value: R) -> R { assert!(TO_LIMBS >= FROM_LIMBS); value } diff --git a/crates/common/Cargo.toml b/crates/common/Cargo.toml index b3feed87..bffc3e9b 100644 --- a/crates/common/Cargo.toml +++ b/crates/common/Cargo.toml @@ -16,3 +16,7 @@ spongefish = { workspace = true } poly = { workspace = true } rayon = { workspace = true, optional = true } num-traits = { workspace = true } +transcript = { workspace = true } + +[dev-dependencies] +num-bigint = { workspace = true } diff --git a/crates/common/src/claim.rs b/crates/common/src/claim.rs index e52b6c15..8863f48b 100644 --- a/crates/common/src/claim.rs +++ b/crates/common/src/claim.rs @@ -1,10 +1,8 @@ //! The linear claim BitZ is asked to discharge. -use crypto_primitives::LiftElement; -use field::Fq; use spongefish::Encoding; -use crate::{BitZParams, Shape}; +use crate::{BitZParams, BitzClaimField, Shape}; /// A Merkle root over the committed codeword. /// @@ -24,8 +22,9 @@ pub enum ClaimError { /// The caller's `x_core`: the weights and the value they are claimed to give. /// -/// BitZ uses `LinearClaim>`; opening queries use `LinearClaim`. -/// The following BitZ requirements apply to the `Fq` input claim. +/// BitZ uses `LinearClaim` over a [`BitzClaimField`]; opening queries use +/// `LinearClaim`. The following BitZ requirements apply to the +/// prime-field input claim. /// BitZ verifies nothing upstream of this. The caller runs its own PIOP, and /// establishes that its claim holds, that `q` is prime, and that the /// coefficient factors as `v = v^(1) (x) v^(2)`. A claim whose coefficient @@ -55,16 +54,16 @@ impl> Encoding<[u8]> for LinearClaim { } } -impl LinearClaim> { +impl LinearClaim { /// Checks the weights against `config` and returns the claim. /// /// `row_weights` is `v^(1)`, one element per row; `column_weights` is /// `v^(2)`, one per column; `target` is the claimed value `mu`. pub fn new( - params: &BitZParams, - row_weights: Vec>, - column_weights: Vec>, - target: Fq, + params: &BitZParams, + row_weights: Vec, + column_weights: Vec, + target: F, ) -> Result { Self::from_shape(params.shape(), row_weights, column_weights, target) } @@ -75,7 +74,7 @@ impl LinearClaim> { /// The claim lives in `F_q` while the exponent is an integer, and taking /// the representative is what bounds it: each is below `q`, so a fold over /// `k_1` rows lands in `[0, k_1(q-1)]`. - pub fn row_exponents(&self) -> Vec { + pub fn row_exponents(&self) -> Vec { self.row_weights .iter() .map(|weight| weight.lift()) @@ -132,25 +131,26 @@ mod tests { /// The largest prime below `2^114`, the top of the sampling range. const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; /// `m = 22`: 128 rows per column, 32768 columns. - fn params() -> BitZParams { + fn params() -> BitZParams { BitZParams::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap() } - fn claim(row_weights: Vec>) -> Result>, ClaimError> { + fn claim(row_weights: Vec) -> Result, ClaimError> { let params = params(); LinearClaim::new( ¶ms, row_weights, - vec![Fq::ONE; params.shape().columns()], - Fq::from(0u128), + vec![F::ONE; params.shape().columns()], + F::from(0u128), ) } - fn weights() -> Vec> { + fn weights() -> Vec { (0..params().shape().rows()) - .map(|row| Fq::from(row as u128)) + .map(|row| F::from(row as u128)) .collect() } @@ -174,7 +174,7 @@ mod tests { } check([1u64, 2, 3, 4].map(field::F128::from)); - check([1u128, 2, 3, 4].map(Fq::::from)); + check([1u128, 2, 3, 4].map(F::from)); } #[test] @@ -209,13 +209,13 @@ mod tests { #[test] fn rejects_weight_vectors_that_do_not_fit_the_shape() { assert_eq!( - claim(vec![Fq::ONE]).err(), + claim(vec![F::ONE]).err(), Some(ClaimError::RowWeightCountMismatch) ); let params = params(); assert_eq!( - LinearClaim::new(¶ms, weights(), vec![Fq::ONE], Fq::from(0u128)).err(), + LinearClaim::new(¶ms, weights(), vec![F::ONE], F::from(0u128)).err(), Some(ClaimError::ColumnWeightCountMismatch) ); } @@ -223,11 +223,11 @@ mod tests { #[test] fn a_row_weight_lifts_to_its_canonical_representative() { // The type makes an out-of-range weight unrepresentable, so there is - // no range check to test: `Fq::from` reduces on the way in. + // no range check to test: `F::from` reduces on the way in. let mut weights = weights(); - weights[3] = Fq::from(Q114 - 1); - weights[4] = Fq::from(Q114); - weights[5] = Fq::from(Q114 + 6); + weights[3] = F::from(Q114 - 1); + weights[4] = F::from(Q114); + weights[5] = F::from(Q114 + 6); let exponents = claim(weights).unwrap().row_exponents(); assert_eq!(exponents[3], Q114 - 1); diff --git a/crates/common/src/fold.rs b/crates/common/src/fold.rs index 50753969..102f14a3 100644 --- a/crates/common/src/fold.rs +++ b/crates/common/src/fold.rs @@ -1,11 +1,12 @@ //! The column fold, and the round state both sides hold once it closes. -use field::{F128, FixedBasePow, Fq}; +use crypto_primitives::Semiring; +use field::{F128, FixedBasePow}; use poly::DenseMultilinearExtension; #[cfg(feature = "parallel")] use rayon::prelude::*; -use crate::{BitTable, LinearClaim, Shape, table::PACKED_BITS}; +use crate::{BitTable, BitzClaimField, BitzSemiring, LinearClaim, Shape, table::PACKED_BITS}; /// A round whose parts do not describe the shape they belong to. #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -26,7 +27,7 @@ pub enum FoldError { /// `exponents` must be the claim's own — [`LinearClaim::row_exponents`]. /// That is what makes the sum safe: each is below `q` and there are `k_1` of /// them, so it is at most `k_1 (q - 1)`, which admissibility put below `|K|`. -pub fn fold_column(table: &BitTable<'_>, exponents: &[u128], column: usize) -> u128 { +pub fn fold_column(table: &BitTable<'_>, exponents: &[S], column: usize) -> S { table .column(column) .iter() @@ -36,15 +37,15 @@ pub fn fold_column(table: &BitTable<'_>, exponents: &[u128], column: usize) -> u [(0, element.lo), (64, element.hi)] .into_iter() .map(|(half, mut remaining)| { - let mut total = 0u128; + let mut total = S::zero(); while remaining != 0 { - total += exponents[base + half + remaining.trailing_zeros() as usize]; + total += &exponents[base + half + remaining.trailing_zeros() as usize]; // Clears the lowest set bit. remaining &= remaining - 1; } total }) - .sum::() + .sum::() }) .sum() } @@ -53,7 +54,7 @@ pub fn fold_column(table: &BitTable<'_>, exponents: &[u128], column: usize) -> u /// /// Columns are independent and read disjoint slices of the witness, so the /// only sharing is the read-only `exponents`. -pub fn fold_columns(table: &BitTable<'_>, exponents: &[u128]) -> Vec { +pub fn fold_columns(table: &BitTable<'_>, exponents: &[S]) -> Vec { let columns = 0..table.shape().columns(); #[cfg(feature = "parallel")] { @@ -76,14 +77,17 @@ pub fn fold_columns(table: &BitTable<'_>, exponents: &[u128]) -> Vec { /// at most `2^21` -- and each is a windowed exponentiation over the whole /// 128-bit range, so this is what the round costs once the fold itself runs in /// parallel. -pub fn column_images(comb: &FixedBasePow, folds: &[u128]) -> Vec { +pub fn column_images(comb: &FixedBasePow, folds: &[S]) -> Vec { #[cfg(feature = "parallel")] { - folds.par_iter().map(|&fold| comb.pow(fold)).collect() + folds + .par_iter() + .map(|fold| comb.pow(fold.clone())) + .collect() } #[cfg(not(feature = "parallel"))] { - folds.iter().map(|&fold| comb.pow(fold)).collect() + folds.iter().map(|fold| comb.pow(fold.clone())).collect() } } @@ -94,10 +98,13 @@ pub fn column_images(comb: &FixedBasePow, folds: &[u128]) -> Vec { /// /// Takes the exponents rather than the claim: the fold has already lifted /// them, and lifting is a pass over `k_1` weights. -pub fn row_images(comb: &FixedBasePow, exponents: &[u128]) -> Vec { +pub fn row_images(comb: &FixedBasePow, exponents: &[S]) -> Vec +where + S: BitzSemiring, +{ exponents .iter() - .map(|&exponent| comb.pow(exponent)) + .map(|exponent| comb.pow(exponent.clone())) .collect() } @@ -110,10 +117,10 @@ pub fn row_images(comb: &FixedBasePow, exponents: &[u128]) -> Vec { /// The length is checked rather than zipped away. A short `folds` would /// otherwise sum over a prefix and return a value that is right for no /// instance -- and for the common `y = 0` it would look correct. -pub fn reconstruct( - claim: &LinearClaim>, - folds: &[u128], -) -> Result, FoldError> { +pub fn reconstruct( + claim: &LinearClaim, + folds: &[F::Integer], +) -> Result { if folds.len() != claim.column_weights().len() { return Err(FoldError::ColumnCountMismatch); } @@ -121,7 +128,7 @@ pub fn reconstruct( .column_weights() .iter() .zip(folds) - .map(|(&weight, &fold)| weight * Fq::from(fold)) + .map(|(weight, fold)| F::from(fold) * weight) .sum()) } @@ -131,9 +138,9 @@ pub fn reconstruct( /// claim, or from the transcript, so the two sides must arrive at /// identical values. #[derive(Debug, Clone, PartialEq, Eq)] -pub struct Fold { +pub struct Fold { /// The column folds `eta_j`, as integers. - pub folds: Vec, + pub folds: Vec, /// `g^{eta_j}` over `j in {0,1}^s`, derived on both sides rather than /// transmitted. The grand product reads the table; `e0` is its extension /// at `zeta`. @@ -147,7 +154,7 @@ pub struct Fold { pub e0: F128, } -impl Fold { +impl Fold { /// Derives `e0` and takes ownership of the round. /// /// Every length is checked against `shape`, including `row_images`, which @@ -155,7 +162,7 @@ impl Fold { /// consumed it. pub fn new( shape: &Shape, - folds: Vec, + folds: Vec, images: Vec, row_images: Vec, zeta: Vec, @@ -193,6 +200,8 @@ mod tests { use crate::{BitZParams, Shape}; const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; + /// Comb window: `FixedBasePow` always covers the full 128-bit exponent /// range, and `win` trades table size against multiplies per call. const WINDOW: u32 = 8; @@ -202,7 +211,7 @@ mod tests { Shape::new(7, 15).unwrap() } - fn params() -> BitZParams { + fn params() -> BitZParams { BitZParams::new(shape(), smallest_generator()).unwrap() } @@ -210,8 +219,8 @@ mod tests { FixedBasePow::new(smallest_generator(), WINDOW) } - fn claim(row_weights: Vec>, column_weights: Vec>) -> LinearClaim> { - LinearClaim::new(¶ms(), row_weights, column_weights, Fq::from(0u128)).unwrap() + fn claim(row_weights: Vec, column_weights: Vec) -> LinearClaim { + LinearClaim::new(¶ms(), row_weights, column_weights, F::from(0u128)).unwrap() } fn witness(shape: &Shape, bits: &[(usize, usize)]) -> Vec { @@ -233,8 +242,8 @@ mod tests { fn a_fold_adds_the_weights_of_the_set_rows() { let shape = shape(); let weights: Vec = (0..shape.rows()).map(|row| (row as u128) * 1_000).collect(); - let field_weights: Vec> = weights.iter().map(|&w| Fq::from(w)).collect(); - let claim = claim(field_weights, vec![Fq::ONE; shape.columns()]); + let field_weights: Vec = weights.iter().map(|&w| F::from(w)).collect(); + let claim = claim(field_weights, vec![F::ONE; shape.columns()]); let packed = witness(&shape, &[(1, 0), (5, 0), (127, 0), (64, 4)]); let table = BitTable::new(shape, &packed).unwrap(); @@ -252,8 +261,8 @@ mod tests { let shape = shape(); // Every weight at `q - 1` and every bit of the column set. let claim = claim( - vec![Fq::from(Q114 - 1); shape.rows()], - vec![Fq::ONE; shape.columns()], + vec![F::from(Q114 - 1); shape.rows()], + vec![F::ONE; shape.columns()], ); let all: Vec<(usize, usize)> = (0..shape.rows()).map(|row| (row, 0)).collect(); let packed = witness(&shape, &all); @@ -268,10 +277,10 @@ mod tests { #[test] fn the_fold_agrees_with_the_bit_by_bit_definition() { let shape = shape(); - let weights: Vec> = (0..shape.rows()) - .map(|row| Fq::from((row as u128 + 1) * (Q114 / 137))) + let weights: Vec = (0..shape.rows()) + .map(|row| F::from((row as u128 + 1) * (Q114 / 137))) .collect(); - let claim = claim(weights, vec![Fq::ONE; shape.columns()]); + let claim = claim(weights, vec![F::ONE; shape.columns()]); let bits: Vec<(usize, usize)> = (0..shape.rows()) .filter(|row| row % 3 == 0 || row % 7 == 1) @@ -290,22 +299,22 @@ mod tests { #[test] fn the_reconstruction_reduces_a_fold_that_runs_past_the_modulus() { let shape = shape(); - let claim = claim(vec![Fq::ONE; shape.rows()], vec![Fq::ONE; shape.columns()]); + let claim = claim(vec![F::ONE; shape.rows()], vec![F::ONE; shape.columns()]); // One column folding to exactly `q` contributes nothing. let mut folds = vec![0u128; shape.columns()]; folds[0] = Q114; - assert_eq!(reconstruct(&claim, &folds), Ok(Fq::from(0u128))); + assert_eq!(reconstruct(&claim, &folds), Ok(F::from(0u128))); folds[0] = Q114 + 5; - assert_eq!(reconstruct(&claim, &folds), Ok(Fq::from(5u128))); + assert_eq!(reconstruct(&claim, &folds), Ok(F::from(5u128))); } #[test] fn the_row_images_are_the_generator_raised_to_each_weight() { let shape = shape(); - let weights: Vec> = (0..shape.rows()).map(|row| Fq::from(row as u128)).collect(); - let claim = claim(weights, vec![Fq::ONE; shape.columns()]); + let weights: Vec = (0..shape.rows()).map(|row| F::from(row as u128)).collect(); + let claim = claim(weights, vec![F::ONE; shape.columns()]); let comb = comb(); let images = row_images(&comb, &claim.row_exponents()); @@ -318,7 +327,7 @@ mod tests { #[test] fn the_reconstruction_refuses_a_fold_vector_that_is_not_one_per_column() { let shape = shape(); - let claim = claim(vec![Fq::ONE; shape.rows()], vec![Fq::ONE; shape.columns()]); + let claim = claim(vec![F::ONE; shape.rows()], vec![F::ONE; shape.columns()]); // The dangerous case: a short vector would sum over a prefix, and at // the common `y = 0` an empty one would look correct. @@ -422,7 +431,7 @@ mod tests { .collect(); let round = Fold::new( &shape, - vec![0; shape.columns()], + vec![0_u64; shape.columns()], images.clone(), vec![F128::ONE; shape.rows()], zeta.clone(), diff --git a/crates/common/src/lib.rs b/crates/common/src/lib.rs index 1897a00b..7486d7d5 100644 --- a/crates/common/src/lib.rs +++ b/crates/common/src/lib.rs @@ -24,14 +24,89 @@ pub use virtual_map::{ TransposedWeights, VirtualMap, VirtualMapError, VirtualStatement, VirtualStatementError, }; -use crypto_primitives::Semiring; -use std::ops::Neg; +use crypto_primitives::{BaseField, Field, Semiring, WithAssociatedInteger}; +use num_traits::{Bounded, FromBytes, ToBytes, ToPrimitive}; +use spongefish::{Encoding, NargDeserialize}; +use std::ops::{BitAnd, Neg, ShrAssign}; -pub trait BitzSemiring: Semiring + From {} +/// Define an empty trait with the given supertraits, and make a blanket +/// implementation for it. +#[macro_export] +macro_rules! define_blanket_trait { + ($(#[$attr:meta])* $vis:vis trait $trait_name:ident: $($bound:tt)+) => { + $(#[$attr])* + $vis trait $trait_name: $($bound)+ {} -impl BitzSemiring for T where T: Semiring + From {} + impl $trait_name for T where T: $($bound)* {} + }; +} -// Since BigInt does not support CheckedNeg and CheckedRem, we can't use Ring here -pub trait BitzRing: BitzSemiring + Neg {} +define_blanket_trait! { + pub trait BitzSemiring: + Semiring + + BitAnd + + ShrAssign + + From + + ToPrimitive +} -impl BitzRing for T where T: BitzSemiring + Neg {} +define_blanket_trait! { + // Since BigInt does not support CheckedNeg and CheckedRem, we can't use Ring here + + /// The ring in which BitZ constraints live. + pub trait BitzConstraintRing: BitzSemiring + Neg +} + +define_blanket_trait! { + /// Any BitZ field type (base or GF128) + pub trait BitzField: + Field + + Copy + + Encoding<[u8]> + + transcript::TranscriptChallenge +} + +define_blanket_trait! { + /// The prime field of a BitZ claim. Its representatives are the fold + /// exponents, so the modulus must fit the `u128` exponent of the `F128` + /// group; `From` takes a fold back into the field. + pub trait BitzClaimField: + BitzField + + BaseField + + WithAssociatedInteger< + Integer: + BitzSemiring + + Bounded + + From + + From + + FromBytes + Encoding<[u8]> + NargDeserialize> + + ToBytes + > + + From> + + From +} + +#[cfg(test)] +mod tests { + use super::*; + use field::{F128, FqDefault}; + + #[test] + fn ensure_traits() { + fn assert_impl_semiring() {} + assert_impl_semiring::(); + assert_impl_semiring::(); + assert_impl_semiring::(); + + fn assert_impl_ring() {} + assert_impl_ring::(); + assert_impl_ring::(); + + fn assert_impl_field() {} + assert_impl_field::(); + assert_impl_field::(); + + fn assert_impl_claim_field() {} + assert_impl_claim_field::(); + } +} diff --git a/crates/common/src/params.rs b/crates/common/src/params.rs index e8ca41e1..4eca2607 100644 --- a/crates/common/src/params.rs +++ b/crates/common/src/params.rs @@ -1,15 +1,19 @@ //! The protocol parameters: everything both sides fix before a claim exists. -use field::{F128, gf128::is_generator}; +use crate::{BitTable, BitzClaimField, Shape, TableError, VirtualMap}; +use field::{ + F128, + gf128::{MULT_ORDER, is_generator}, +}; +use num_traits::{CheckedMul, ToBytes}; use spongefish::Encoding; - -use crate::{BitTable, Shape, TableError, VirtualMap}; +use std::marker::PhantomData; /// A parameter set one of the pre-claim gates rejects. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ParamsError { - /// `(k_1 + 1)(Q - 1)` reaches `ord(g)`, so two folds could collide in the - /// exponent. + /// `(k_1 + 1)(Q - 1)` reaches `ord(g)`, so a sent fold and the honest + /// exponent could collide in the exponent. FoldBoundExceeded, /// The generator's order is not the full group, so a fold is not the only /// exponent producing its image. @@ -21,42 +25,45 @@ pub enum ParamsError { /// The two roles derive their own setups from this, so neither can be built /// against parameters the other did not see. /// -/// `Q` is a const parameter, not a field: the weights are `Fq`, whose -/// modulus lives in the type. +/// The modulus is not a value here, it's accessible as `F::modulus()`. #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct BitZParams { +pub struct BitZParams { shape: Shape, generator: F128, + _phantom: PhantomData, } -impl BitZParams { +impl BitZParams { /// Runs the gates that need only the parameters. pub fn new(shape: Shape, generator: F128) -> Result { - // `Fq` asserts Q is an odd prime below 2^126 on its own behalf, so - // no modulus gate is needed here. - - // `ord(g) > (k_1 + 1)(Q - 1)`, the paper's requisite. A fold is an - // integer at most `k_1 (Q - 1)` while the value it is compared against - // is at most `Q - 1`, so the two differ by at most the sum; the - // exponent is only ever seen modulo `ord(g)`, and a gap that never - // reaches the group order cannot close. + // The field asserts its modulus is an odd prime below 2^126 on its own + // behalf, so no modulus gate is needed here. + + // `g^{} = g^{eta_j}` implies integer equality only if both + // sides, each in `[0, k_1 (Q - 1)]`, differ by less than `ord(g)`. + // The paper asks for `Q < (|K| - 1) / k_1`; the PoC's slightly stricter + // `(k_1 + 1)(Q - 1) < ord(g)` is kept, and it implies the paper's. // - // Overflow is itself a rejection: past `2^128 - 1` there is no room - // left. `ord(g)` is `u128::MAX`, the full order the generator gate - // below establishes. - let Some(gap) = (Q - 1).checked_mul(shape.rows() as u128 + 1) else { - return Err(ParamsError::FoldBoundExceeded); - }; - // A `u128` cannot exceed `ord(g)`, so equalling it is the only way - // left to reach it. - if gap == u128::MAX { + // `ord(g) = 2^128 - 1` is a property of `F128`, established by the + // generator gate below. A product that overflows exceeds it too. + let f128_order = F::Integer::from(MULT_ORDER); + let max_f = F::max_value().lift(); + let rows_plus_one = F::Integer::from(shape.rows() as u64 + 1); + if !max_f + .checked_mul(&rows_plus_one) + .is_some_and(|gap| gap < f128_order) + { return Err(ParamsError::FoldBoundExceeded); } if !is_generator(generator) { return Err(ParamsError::GeneratorOrderNotFull); } - Ok(Self { shape, generator }) + Ok(Self { + shape, + generator, + _phantom: PhantomData, + }) } /// Views a packed witness through the configured shape. @@ -78,25 +85,24 @@ impl BitZParams { /// The largest fold the verifier may accept, `k_1 (Q - 1)`. /// /// [`Self::new`]'s gate puts it below `ord(g)`. - pub fn fold_bound(&self) -> u128 { - (self.shape.rows() as u128) * (Q - 1) + pub fn fold_bound(&self) -> F::Integer { + let rows = u64::try_from(self.shape.rows()).expect("Too many rows"); + let rows = F::Integer::from(rows); + let max_f = F::max_value().lift(); + rows.checked_mul(&max_f).expect("Multiplication overflow") } } -/// Every field is fixed width, so distinct parameter sets cannot encode alike. -impl Encoding<[u8]> for BitZParams { +/// Every field is fixed width, the modulus at `F::Integer`'s, so distinct +/// parameter sets cannot encode alike. +impl Encoding<[u8]> for BitZParams { fn encode(&self) -> impl AsRef<[u8]> { - let mut frame = [0u8; 48]; - let mut at = 0; - let mut put = |bytes: &[u8]| { - frame[at..at + bytes.len()].copy_from_slice(bytes); - at += bytes.len(); - }; - - put(&(self.shape.log_rows() as u64).to_le_bytes()); - put(&(self.shape.log_columns() as u64).to_le_bytes()); - put(&Q.to_le_bytes()); - put(&self.generator.to_bytes()); + let modulus = F::modulus().to_le_bytes(); + let mut frame = Vec::with_capacity(32 + modulus.as_ref().len()); + frame.extend_from_slice(&(self.shape.log_rows() as u64).to_le_bytes()); + frame.extend_from_slice(&(self.shape.log_columns() as u64).to_le_bytes()); + frame.extend_from_slice(modulus.as_ref()); + frame.extend_from_slice(&self.generator.to_bytes()); frame } } @@ -120,19 +126,19 @@ pub enum VirtualParamsError { /// exponent, and what the claim's weights are counted against. The committed /// shape belongs to `f` alone and reaches only the table and the opening. #[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct VirtualParams { - claim: BitZParams, +pub struct VirtualParams { + claim: BitZParams, committed: Shape, } -impl VirtualParams { +impl VirtualParams { /// Checks both shapes against the map before either is used. /// /// Zero padding is what makes the inequalities rather than equalities: both /// vectors are padded up to their shape so that they have multilinear /// extensions, and padding contributes nothing. pub fn new( - claim: BitZParams, + claim: BitZParams, committed: Shape, map: &impl VirtualMap, ) -> Result { @@ -151,7 +157,7 @@ impl VirtualParams { } /// The parameters the fold and the reduction read, shaped to `h`. - pub fn claim(&self) -> &BitZParams { + pub fn claim(&self) -> &BitZParams { &self.claim } @@ -169,12 +175,11 @@ impl VirtualParams { /// Distinct from a plain [`BitZParams`] frame by length, so a proof of one /// cannot replay as a proof of the other. -impl Encoding<[u8]> for VirtualParams { +impl Encoding<[u8]> for VirtualParams { fn encode(&self) -> impl AsRef<[u8]> { - let mut frame = [0u8; 64]; - frame[..48].copy_from_slice(self.claim.encode().as_ref()); - frame[48..56].copy_from_slice(&(self.committed.log_rows() as u64).to_le_bytes()); - frame[56..].copy_from_slice(&(self.committed.log_columns() as u64).to_le_bytes()); + let mut frame = self.claim.encode().as_ref().to_vec(); + frame.extend_from_slice(&(self.committed.log_rows() as u64).to_le_bytes()); + frame.extend_from_slice(&(self.committed.log_columns() as u64).to_le_bytes()); frame } } @@ -187,8 +192,9 @@ mod tests { /// The largest prime below `2^114`, the top of the sampling range. const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; - fn params_at(shape: Shape) -> Result, ParamsError> { + fn params_at(shape: Shape) -> Result, ParamsError> { BitZParams::new(shape, smallest_generator()) } @@ -227,7 +233,7 @@ mod tests { fn virtual_params_for( h_len: usize, f_len: usize, - ) -> Result, VirtualParamsError> { + ) -> Result, VirtualParamsError> { VirtualParams::new( params_at(shape()).unwrap(), shape(), @@ -283,8 +289,7 @@ mod tests { /// through. #[test] fn the_table_is_shaped_by_the_committed_bits() { - let claim = - BitZParams::::new(Shape::new(8, 15).unwrap(), smallest_generator()).unwrap(); + let claim = BitZParams::::new(Shape::new(8, 15).unwrap(), smallest_generator()).unwrap(); let committed = shape(); let params = VirtualParams::new(claim, committed, &Dimensions { h_len: 4, f_len: 4 }).unwrap(); @@ -308,7 +313,7 @@ mod tests { #[test] fn rejects_a_generator_of_partial_order() { assert_eq!( - BitZParams::::new(shape(), F128::ONE).err(), + BitZParams::::new(shape(), F128::ONE).err(), Some(ParamsError::GeneratorOrderNotFull) ); } diff --git a/crates/common/src/virtual_map.rs b/crates/common/src/virtual_map.rs index a59e2d03..0908d3c2 100644 --- a/crates/common/src/virtual_map.rs +++ b/crates/common/src/virtual_map.rs @@ -11,13 +11,14 @@ //! //! `M` stays with the caller. This crate needs only `M^T v` and a digest. -use field::{F128, Fq}; +use field::F128; use num_traits::ConstZero; #[cfg(feature = "parallel")] use rayon::prelude::*; use crate::{ - BitZParams, ClaimError, LinearClaim, OpeningQuery, Shape, VirtualParams, VirtualParamsError, + BitZParams, BitzClaimField, ClaimError, LinearClaim, OpeningQuery, Shape, VirtualParams, + VirtualParamsError, }; /// A claim on virtual bits `h = M (1 || f)` with checked dimensions. @@ -25,10 +26,10 @@ use crate::{ /// Both roles bind the virtual domain, commitment root, shapes, modulus, /// generator, map digest, and input claim to the transcript before folding. #[derive(Debug)] -pub struct VirtualStatement<'a, const Q: u128, M: VirtualMap> { - params: VirtualParams, +pub struct VirtualStatement<'a, F, M: VirtualMap> { + params: VirtualParams, map: &'a M, - claim: &'a LinearClaim>, + claim: &'a LinearClaim, } /// A map or claim with mismatched dimensions. @@ -115,16 +116,16 @@ impl TransposedWeights { } } -impl<'a, const Q: u128, M: VirtualMap> VirtualStatement<'a, Q, M> { +impl<'a, F: BitzClaimField, M: VirtualMap> VirtualStatement<'a, F, M> { /// Checks the map and claim against both witness shapes. /// /// The statement retains the checked inputs. The map must keep the same /// linear transformation while the statement borrows it. pub fn new( - claim_params: BitZParams, + claim_params: BitZParams, committed_shape: Shape, map: &'a M, - claim: &'a LinearClaim>, + claim: &'a LinearClaim, ) -> Result { let params = VirtualParams::new(claim_params, committed_shape, map) .map_err(VirtualStatementError::Parameters)?; @@ -141,7 +142,7 @@ impl<'a, const Q: u128, M: VirtualMap> VirtualStatement<'a, Q, M> { Ok(Self { params, map, claim }) } - pub fn params(&self) -> &VirtualParams { + pub fn params(&self) -> &VirtualParams { &self.params } @@ -149,7 +150,7 @@ impl<'a, const Q: u128, M: VirtualMap> VirtualStatement<'a, Q, M> { self.map } - pub fn claim(&self) -> &LinearClaim> { + pub fn claim(&self) -> &LinearClaim { self.claim } @@ -228,6 +229,9 @@ mod tests { use super::*; use num_traits::{ConstOne, ConstZero}; + const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; + /// Dense `M`, the reference the sparse implementations are checked against. struct DenseMap { /// Row `i` is the `1 ‖ f` indicator of the bits `h_i` sums. @@ -345,19 +349,17 @@ mod tests { assert_eq!(map.transpose(&weights), map.transpose(&padded)); } - const Q: u128 = (1 << 114) - 11; - - fn params() -> BitZParams { + fn params() -> BitZParams { let shape = Shape::new(7, 15).unwrap(); BitZParams::new(shape, field::gf128::smallest_generator()).unwrap() } - fn input_claim(params: &BitZParams) -> LinearClaim> { + fn input_claim(params: &BitZParams) -> LinearClaim { LinearClaim::new( params, - vec![Fq::ONE; params.shape().rows()], - vec![Fq::ONE; params.shape().columns()], - Fq::from(0u128), + vec![F::ONE; params.shape().rows()], + vec![F::ONE; params.shape().columns()], + F::from(0u128), ) .unwrap() } diff --git a/crates/field/src/fq.rs b/crates/field/src/fq.rs index ae62acc2..ca677026 100644 --- a/crates/field/src/fq.rs +++ b/crates/field/src/fq.rs @@ -28,7 +28,9 @@ pub type FqDefault = Fq; /// Barrett leaves a value below `3Q`, which has to fit a `u128`. That admits /// anything up to `floor((2^128 - 1)/3)`, just over `2^126.41`; the bound is /// rounded down to a power of two. -#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash, InfallibleCheckedOp)] +#[derive( + Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash, InfallibleCheckedOp, +)] #[infallible_checked_unary_op((CheckedNeg, neg))] #[infallible_checked_binary_op((CheckedAdd, add), (CheckedSub, sub), (CheckedMul, mul))] pub struct Fq(u128); diff --git a/crates/field/src/gf128/pow.rs b/crates/field/src/gf128/pow.rs index d90e9197..98d689d9 100644 --- a/crates/field/src/gf128/pow.rs +++ b/crates/field/src/gf128/pow.rs @@ -1,9 +1,10 @@ //! Exponentiation, inversion, and the primitive-element test. -use std::fmt::{Debug, Formatter, Result as FmtResult}; - use super::{F128, kernel}; -use num_traits::{ConstOne, Inv, Pow, Zero}; +use crypto_primitives::Semiring; +use num_traits::{ConstOne, Inv, Pow, ToPrimitive, Zero}; +use std::fmt::{Debug, Formatter, Result as FmtResult}; +use std::ops::{BitAnd, ShrAssign}; /// The order of the multiplicative group, `2^128 - 1`. /// @@ -175,14 +176,19 @@ impl FixedBasePow { Self { table, win } } - pub fn pow(&self, exp: u128) -> F128 { - let mask = (1u128 << self.win) - 1; + pub fn pow(&self, exp: S) -> F128 + where + S: Semiring + BitAnd + ToPrimitive + ShrAssign, + { + let two = S::one() + S::one(); + let mask = two.pow(self.win) - S::one(); let mut acc = F128::ONE; let mut e = exp; let mut i = 0; - while e != 0 { - let d = (e & mask) as usize; - if d != 0 { + while !e.is_zero() { + let d = e.clone() & mask.clone(); + if !d.is_zero() { + let d = d.to_usize().expect("Value is too large"); acc *= self.table[i][d]; } e >>= self.win; diff --git a/crates/poly/src/eq.rs b/crates/poly/src/eq.rs index d28db403..4dbb024e 100644 --- a/crates/poly/src/eq.rs +++ b/crates/poly/src/eq.rs @@ -1,10 +1,9 @@ -use crypto_primitives::ConstField; -#[cfg(feature = "parallel")] -use rayon::prelude::*; - use crate::mle::{DenseMleError, DenseMultilinearExtension}; #[cfg(feature = "parallel")] -use crate::parallel::workload_size; +use crate::parallel; +use crypto_primitives::Field; +#[cfg(feature = "parallel")] +use rayon::prelude::*; /// Evaluates the multilinear equality polynomial. /// @@ -20,14 +19,14 @@ use crate::parallel::workload_size; /// saving one multiplication per coordinate compared with the defining /// expression. Seeding the accumulator with the first coordinate factor gives /// `2n - 1` multiplications for nonempty points. -pub fn eq_eval(left: &[F], right: &[F]) -> F { +pub fn eq_eval(left: &[F], right: &[F]) -> F { assert_eq!( left.len(), right.len(), "equality points must have the same length" ); - let one = F::ONE; + let one = F::one(); let mut coordinates = left.iter().copied().zip(right.iter().copied()); let Some((first_left, first_right)) = coordinates.next() else { return one; @@ -59,13 +58,13 @@ pub fn eq_eval(left: &[F], right: &[F]) -> F { /// `b_i = (j >> i) & 1`, so variable `i` corresponds to bit `i` of /// the index (little-endian order). /// -/// For `n = 0`, the table is `[F::ONE]`, corresponding to the +/// For `n = 0`, the table is `[F::one()]`, corresponding to the /// empty product. -pub fn eq_table(r: &[F]) -> Vec { +pub fn eq_table(r: &[F]) -> Vec { let n = 1 << r.len(); // Allocate the final output once. - let mut table = vec![F::ZERO; n]; - table[0] = F::ONE; + let mut table = vec![F::zero(); n]; + table[0] = F::one(); for (i, &r_i) in r.iter().enumerate() { let half = 1usize << i; @@ -85,7 +84,7 @@ pub fn eq_table(r: &[F]) -> Vec { // Per-level parallel doubling adapted from Flock: // https://github.com/succinctlabs/flock/blob/85fc0e7cc002e7ca4dffdff805ba89976e9a5293/crates/flock-core/src/pcs/ring_switch.rs#L276-L320 #[cfg(feature = "parallel")] - if half >= workload_size::() { + if half >= parallel::workload_size::() { zero_children .par_iter_mut() .zip(one_children.par_iter_mut()) @@ -105,7 +104,7 @@ pub fn eq_table(r: &[F]) -> Vec { /// The first MLE covers the low, earlier coordinates and the second covers the /// remaining high coordinates. Their tensor product is `eq(·, point)` under /// this crate's little-endian variable order. -pub fn make_equality_factors( +pub fn make_equality_factors( point: &[F], ) -> Result<(DenseMultilinearExtension, DenseMultilinearExtension), DenseMleError> { let split = point.len() / 2; @@ -117,16 +116,17 @@ pub fn make_equality_factors( #[cfg(test)] pub mod tests { - use super::{eq_eval, eq_table, make_equality_factors}; - use crypto_primitives::ConstField; - use field::{F128, FqDefault}; - use num_traits::{ConstOne, ConstZero}; + use super::*; + use field::F128; + use num_traits::{One, Zero}; use proptest::prelude::*; - fn direct_table_entry(r: &[F], index: usize) -> F { - r.iter().enumerate().fold(F::ONE, |acc, (bit, &r_i)| { + type F = field::FqDefault; + + fn direct_table_entry(r: &[F], index: usize) -> F { + r.iter().enumerate().fold(F::one(), |acc, (bit, &r_i)| { let factor = if (index >> bit) & 1 == 0 { - F::ONE - r_i + F::one() - r_i } else { r_i }; @@ -136,17 +136,13 @@ pub mod tests { #[test] fn empty_eq_tables_contain_one() { - assert_eq!(eq_table::(&[]), vec![F128::ONE]); - assert_eq!(eq_table::(&[]), vec![FqDefault::ONE]); + assert_eq!(eq_table::(&[]), vec![F128::one()]); + assert_eq!(eq_table::(&[]), vec![F::one()]); } #[test] fn equality_factors_split_low_coordinates_first() { - let point = [ - FqDefault::from(2u128), - FqDefault::from(3u128), - FqDefault::from(5u128), - ]; + let point = [F::from(2u128), F::from(3u128), F::from(5u128)]; let (low, high) = make_equality_factors(&point).unwrap(); assert_eq!(low.num_vars(), 1); @@ -166,9 +162,9 @@ pub mod tests { let (low, high) = make_equality_factors::(&[]).unwrap(); assert_eq!(low.num_vars(), 0); - assert_eq!(low.iter().copied().collect::>(), vec![F128::ONE]); + assert_eq!(low.iter().copied().collect::>(), vec![F128::one()]); assert_eq!(high.num_vars(), 0); - assert_eq!(high.iter().copied().collect::>(), vec![F128::ONE]); + assert_eq!(high.iter().copied().collect::>(), vec![F128::one()]); } #[test] @@ -178,21 +174,21 @@ pub mod tests { assert_eq!( eq_table(&[r_0, r_1]), vec![ - (F128::ONE - r_0) * (F128::ONE - r_1), - r_0 * (F128::ONE - r_1), - (F128::ONE - r_0) * r_1, + (F128::one() - r_0) * (F128::one() - r_1), + r_0 * (F128::one() - r_1), + (F128::one() - r_0) * r_1, r_0 * r_1, ] ); - let r_0 = FqDefault::from(2u128); - let r_1 = FqDefault::from(7u128); + let r_0 = F::from(2u128); + let r_1 = F::from(7u128); assert_eq!( eq_table(&[r_0, r_1]), vec![ - (FqDefault::ONE - r_0) * (FqDefault::ONE - r_1), - r_0 * (FqDefault::ONE - r_1), - (FqDefault::ONE - r_0) * r_1, + (F::one() - r_0) * (F::one() - r_1), + r_0 * (F::one() - r_1), + (F::one() - r_0) * r_1, r_0 * r_1, ] ); @@ -205,16 +201,16 @@ pub mod tests { .map(|bit| F128::from((selected >> bit) & 1 == 1)) .collect(); let r_fq: Vec<_> = (0..4) - .map(|bit| FqDefault::from(((selected >> bit) & 1) as u128)) + .map(|bit| F::from(((selected >> bit) & 1) as u128)) .collect(); for (index, &weight) in eq_table(&r_f128).iter().enumerate() { assert_eq!( weight, if index == selected { - F128::ONE + F128::one() } else { - F128::ZERO + F128::zero() } ); } @@ -222,9 +218,9 @@ pub mod tests { assert_eq!( weight, if index == selected { - FqDefault::ONE + F::one() } else { - FqDefault::ZERO + F::zero() } ); } @@ -232,14 +228,14 @@ pub mod tests { #[test] fn empty_vectors_give_one() { - assert_eq!(eq_eval::(&[], &[]), F128::ONE); - assert_eq!(eq_eval::(&[], &[]), FqDefault::ONE); + assert_eq!(eq_eval::(&[], &[]), F128::one()); + assert_eq!(eq_eval::(&[], &[]), F::one()); } #[test] fn one_coordinate_boolean_truth_table() { - let zero = F128::ZERO; - let one = F128::ONE; + let zero = F128::zero(); + let one = F128::one(); assert_eq!(eq_eval(&[zero], &[zero]), one); assert_eq!(eq_eval(&[zero], &[one]), zero); @@ -249,8 +245,8 @@ pub mod tests { #[test] fn boolean_vectors_are_equality_indicators() { - let zero = F128::ZERO; - let one = F128::ONE; + let zero = F128::zero(); + let one = F128::one(); let x = [zero, one, one, zero]; let equal = [zero, one, one, zero]; @@ -271,7 +267,7 @@ pub mod tests { #[test] #[should_panic(expected = "equality points must have the same length")] fn different_widths_panic() { - eq_eval(&[F128::ZERO], &[F128::ZERO, F128::ONE]); + eq_eval(&[F128::zero()], &[F128::zero(), F128::one()]); } #[test] @@ -282,7 +278,7 @@ pub mod tests { F128::from(19u128), F128::from(31u128), ]; - let mut sum = F128::ZERO; + let mut sum = F128::zero(); for index in 0..1usize << r.len() { let point: Vec<_> = (0..r.len()) @@ -291,13 +287,13 @@ pub mod tests { sum += eq_eval(&point, &r); } - assert_eq!(sum, F128::ONE); + assert_eq!(sum, F128::one()); } #[cfg(feature = "parallel")] #[test] fn parallel_eq_table_matches_serial_around_threshold() { - let threshold = crate::parallel::workload_size::(); + let threshold = parallel::workload_size::(); assert!(threshold.is_power_of_two()); let threshold_log = threshold.trailing_zeros() as usize; @@ -333,7 +329,7 @@ pub mod tests { prop_assert_eq!(weight, direct_table_entry(&r_f128, index)); } - let r_fq: Vec<_> = raw.iter().copied().map(FqDefault::from).collect(); + let r_fq: Vec<_> = raw.iter().copied().map(F::from).collect(); let table_fq = eq_table(&r_fq); prop_assert_eq!(table_fq.len(), 1usize << raw.len()); for (index, &weight) in table_fq.iter().enumerate() { @@ -348,14 +344,14 @@ pub mod tests { let r_f128: Vec<_> = raw.iter().copied().map(F128::from).collect(); let sum_f128 = eq_table(&r_f128) .into_iter() - .fold(F128::ZERO, |sum, weight| sum + weight); - prop_assert_eq!(sum_f128, F128::ONE); + .fold(F128::zero(), |sum, weight| sum + weight); + prop_assert_eq!(sum_f128, F128::one()); - let r_fq: Vec<_> = raw.iter().copied().map(FqDefault::from).collect(); + let r_fq: Vec<_> = raw.iter().copied().map(F::from).collect(); let sum_fq = eq_table(&r_fq) .into_iter() - .fold(FqDefault::ZERO, |sum, weight| sum + weight); - prop_assert_eq!(sum_fq, FqDefault::ONE); + .fold(F::zero(), |sum, weight| sum + weight); + prop_assert_eq!(sum_fq, F::one()); } #[test] @@ -373,12 +369,12 @@ pub mod tests { .collect(); let expected = x.iter().zip(&y).fold( - F128::ONE, + F128::one(), |acc, (&x_i, &y_i)| { acc * ( x_i * y_i - + (F128::ONE - x_i) - * (F128::ONE - y_i) + + (F128::one() - x_i) + * (F128::one() - y_i) ) }, ); diff --git a/crates/poly/src/mle.rs b/crates/poly/src/mle.rs index ba1d331e..476ba2a0 100644 --- a/crates/poly/src/mle.rs +++ b/crates/poly/src/mle.rs @@ -281,10 +281,12 @@ mod tests { DenseMleError, DenseMultilinearExtension, MleClaimError, ScaledMleEvaluationClaim, }; use crate::eq::eq_table; - use field::{F128, FqDefault}; + use field::F128; use num_traits::{ConstOne, ConstZero}; use proptest::prelude::*; + type F = field::FqDefault; + fn fold_layer(evaluations: &[F128], challenge: F128) -> Vec { evaluations .chunks_exact(2) @@ -309,14 +311,12 @@ mod tests { #[test] fn scaled_mle_claim_verifies_and_exposes_its_terms() { - let polynomial = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(2u128), FqDefault::from(5u128)], - ) - .unwrap(); - let point = FqDefault::from(3u128); - let scale = FqDefault::from(7u128); - let value = FqDefault::from(77u128); + let polynomial = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(2u128), F::from(5u128)]) + .unwrap(); + let point = F::from(3u128); + let scale = F::from(7u128); + let value = F::from(77u128); let claim = ScaledMleEvaluationClaim::new(vec![point].into_boxed_slice(), scale, value); assert_eq!(claim.point(), &[point]); @@ -327,15 +327,13 @@ mod tests { #[test] fn scaled_mle_claim_rejects_an_incorrect_value() { - let polynomial = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(2u128), FqDefault::from(5u128)], - ) - .unwrap(); + let polynomial = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(2u128), F::from(5u128)]) + .unwrap(); let claim = ScaledMleEvaluationClaim::new( - vec![FqDefault::from(3u128)].into_boxed_slice(), - FqDefault::from(7u128), - FqDefault::from(78u128), + vec![F::from(3u128)].into_boxed_slice(), + F::from(7u128), + F::from(78u128), ); assert_eq!( @@ -346,15 +344,13 @@ mod tests { #[test] fn scaled_mle_claim_rejects_a_point_with_the_wrong_width() { - let polynomial = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(2u128), FqDefault::from(5u128)], - ) - .unwrap(); + let polynomial = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(2u128), F::from(5u128)]) + .unwrap(); let claim = ScaledMleEvaluationClaim::new( - vec![FqDefault::from(3u128), FqDefault::from(11u128)].into_boxed_slice(), - FqDefault::from(7u128), - FqDefault::from(77u128), + vec![F::from(3u128), F::from(11u128)].into_boxed_slice(), + F::from(7u128), + F::from(77u128), ); assert_eq!( diff --git a/crates/prover/benches/prover.rs b/crates/prover/benches/prover.rs index 6c70a9f4..1264a101 100644 --- a/crates/prover/benches/prover.rs +++ b/crates/prover/benches/prover.rs @@ -6,7 +6,7 @@ use std::hint::black_box; use common::{BitTable, BitZParams, Fold, Shape}; use divan::Bencher; -use field::F128; +use field::{F128, Fq}; use num_traits::ConstOne; fn main() { @@ -20,7 +20,7 @@ fn random_table(shape: Shape) -> BitTable<'static> { let packed: Box> = Box::new((0..n).map(|_| F128::from(rand::random::())).collect()); let packed: &'static _ = packed.leak(); - let params: BitZParams = + let params: BitZParams> = BitZParams::new(shape, field::gf128::smallest_generator()).unwrap(); params.table(packed).unwrap() diff --git a/crates/prover/examples/profile.rs b/crates/prover/examples/profile.rs index b2e97402..f9af864a 100644 --- a/crates/prover/examples/profile.rs +++ b/crates/prover/examples/profile.rs @@ -13,6 +13,8 @@ use num_traits::ConstOne; use transcript::ProverState; const Q114: u128 = (1 << 114) - 11; +type F = field::Fq; + const ITERATIONS: usize = 30; fn random_table(shape: Shape) -> BitTable<'static> { @@ -20,8 +22,7 @@ fn random_table(shape: Shape) -> BitTable<'static> { let packed: Box> = Box::new((0..n).map(|_| F128::from(rand::random::())).collect()); let packed: &'static _ = packed.leak(); - let params: BitZParams = - BitZParams::new(shape, field::gf128::smallest_generator()).unwrap(); + let params: BitZParams = BitZParams::new(shape, field::gf128::smallest_generator()).unwrap(); params.table(packed).unwrap() } @@ -49,7 +50,7 @@ fn main() { } #[inline(never)] -fn gkr_wrapper(mut transcript: ProverState, fold: &Fold, table: BitTable<'_>) { +fn gkr_wrapper(mut transcript: ProverState, fold: &Fold, table: BitTable<'_>) { black_box( prover::gkr_reduce(&mut transcript, black_box(fold), black_box(&table)) .expect("profiling fold matches the table shape"), diff --git a/crates/prover/src/fold.rs b/crates/prover/src/fold.rs index 8c272fe5..56b58523 100644 --- a/crates/prover/src/fold.rs +++ b/crates/prover/src/fold.rs @@ -1,8 +1,10 @@ //! The fold round: send the column folds, then take the challenge. -use common::{BitTable, Fold, FoldError, LinearClaim, column_images, fold_columns, row_images}; - use crate::BitZProver; +use common::{ + BitTable, BitzClaimField, Fold, FoldError, LinearClaim, column_images, fold_columns, row_images, +}; +use num_traits::ToBytes; use transcript::ProverState; /// A fold the prover cannot produce. @@ -14,7 +16,7 @@ pub enum SendError { Fold(FoldError), } -impl BitZProver { +impl BitZProver { /// Runs the fold round. /// /// Only the folds `eta_j` are sent. Their images `g^{eta_j}` are what the @@ -34,10 +36,10 @@ impl BitZProver { #[tracing::instrument(name = "Fold columns", skip_all)] pub fn send_fold( &self, - claim: &LinearClaim>, + claim: &LinearClaim, table: &BitTable<'_>, transcript: &mut ProverState, - ) -> Result { + ) -> Result, SendError> { // The weights are sized by the configured row count while the bits are // read at the table's. Disagreement is a panic, a silently wrong fold, or // a desynchronised transcript depending on which way it goes. @@ -51,7 +53,7 @@ impl BitZProver { let folds = fold_columns(table, &exponents); for fold in &folds { - transcript.prover_message(&fold.to_le_bytes()); + transcript.prover_message(&fold.to_le_bytes().as_ref()); } let images = column_images(self.comb(), &folds); diff --git a/crates/prover/src/prove.rs b/crates/prover/src/prove.rs index 49fc2789..e5cd4cd9 100644 --- a/crates/prover/src/prove.rs +++ b/crates/prover/src/prove.rs @@ -1,10 +1,10 @@ //! `ProveBitZ`. use common::{ - BitTable, ClaimError, LinearClaim, OpeningQuery, TableError, VirtualMap, VirtualMapError, - VirtualStatement, + BitTable, BitzClaimField, ClaimError, LinearClaim, OpeningQuery, TableError, VirtualMap, + VirtualMapError, VirtualStatement, }; -use field::{F128, Fq}; +use field::F128; use pcs::{CommitScheme, Pcs, ProveError as OpeningProveError, ProverData, StatementBinding}; use transcript::ProverState; @@ -39,7 +39,7 @@ pub struct VirtualWitness<'a> { pub virtual_bits: &'a [F128], } -impl BitZProver { +impl BitZProver { /// Proves the caller's linear claim about the committed bits. /// /// Call `Pcs::commit` on this transcript, then pass its retained @@ -52,7 +52,7 @@ impl BitZProver { #[tracing::instrument(name = "Prove BitZ", skip_all)] pub fn prove( &self, - claim: &LinearClaim>, + claim: &LinearClaim, pcs: &Pcs, data: &ProverData, packed: Vec, @@ -92,7 +92,7 @@ impl BitZProver { #[tracing::instrument(name = "Prove virtual BitZ", skip_all)] pub fn prove_virtual( &self, - statement: &VirtualStatement<'_, Q, impl VirtualMap>, + statement: &VirtualStatement<'_, F, impl VirtualMap>, pcs: &Pcs, data: &ProverData, witness: VirtualWitness<'_>, @@ -148,11 +148,12 @@ impl BitZProver { /// The caller binds the statement before this call and opens the returned claim. pub(crate) fn fold_and_reduce( &self, - claim: &LinearClaim>, + claim: &LinearClaim, table: &BitTable<'_>, transcript: &mut ProverState, ) -> Result { - // Step 2 is absent: Q is fixed, and BitZParams::new checks its fold bound. + // Step 2 is absent: the field modulus is fixed, and BitZParams::new + // checks its fold bound. // Step 3: fold each column into an integer exponent. let fold = self diff --git a/crates/prover/src/reduce.rs b/crates/prover/src/reduce.rs index 4364cf26..95245b78 100644 --- a/crates/prover/src/reduce.rs +++ b/crates/prover/src/reduce.rs @@ -14,7 +14,7 @@ use transcript::ProverState; #[inline(never)] #[tracing::instrument(name = "Build grand-product circuit", level = "debug", skip_all)] -fn init_circuit(table: &BitTable, fold: &Fold) -> GrandProductCircuit { +fn init_circuit(table: &BitTable, fold: &Fold) -> GrandProductCircuit { let columns = table.shape().columns(); let dim = columns * table.shape().rows(); let mut leafs = F128::zeroed_vec(dim); @@ -52,9 +52,9 @@ fn init_circuit(table: &BitTable, fold: &Fold) -> GrandProductCircuit { /// Reduces the grand-product circuit to a factored claim on the committed bits. #[tracing::instrument(name = "Reduce grand products", skip_all)] -pub fn gkr_reduce( +pub fn gkr_reduce( transcript: &mut ProverState, - fold: &Fold, + fold: &Fold, table: &BitTable, ) -> Result { let circuit = init_circuit(table, fold); @@ -91,13 +91,16 @@ mod order_check_ai_test { use field::gf128::smallest_generator; use num_traits::ConstZero; - const Q: u128 = (1 << 114) - 11; + const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; + + type F2 = field::FqDefault; fn shape() -> Shape { Shape::new(7, 15).unwrap() } - fn params() -> BitZParams { + fn params() -> BitZParams { BitZParams::new(shape(), smallest_generator()).unwrap() } @@ -150,12 +153,10 @@ mod order_check_ai_test { #[test] fn factored_claim_matches_for_narrow_tables() { - const Q100: u128 = (1 << 100) - 15; - // Cover one column and both sides of the 128-column transpose boundary. for log_columns in [0, 6, 7] { let shape = Shape::new(22 - log_columns, log_columns).unwrap(); - let params = BitZParams::::new(shape, smallest_generator()).unwrap(); + let params = BitZParams::::new(shape, smallest_generator()).unwrap(); let packed = packed_witness(&shape, |column, row| { let bits = (row as u64).wrapping_mul(0x9E3779B97F4A7C15) ^ (column as u64).wrapping_mul(0xD1B54A32D192ED03); @@ -170,7 +171,7 @@ mod order_check_ai_test { .collect(); let fold = Fold::new( &shape, - vec![0; shape.columns()], + vec![0_u64; shape.columns()], vec![F128::ONE; shape.columns()], row_images, zeta, diff --git a/crates/prover/src/setup.rs b/crates/prover/src/setup.rs index 19fa1898..c7855c78 100644 --- a/crates/prover/src/setup.rs +++ b/crates/prover/src/setup.rs @@ -1,6 +1,6 @@ //! The prover's derived setup. -use common::BitZParams; +use common::{BitZParams, BitzClaimField}; use field::FixedBasePow; /// The parameters, with the comb table derived from them. @@ -8,12 +8,12 @@ use field::FixedBasePow; /// Deriving the comb here is what removes the pairing check the two used to /// need: there is no second generator for the comb to disagree with. #[derive(Debug)] -pub struct BitZProver { - params: BitZParams, +pub struct BitZProver { + params: BitZParams, comb: FixedBasePow, } -impl BitZProver { +impl BitZProver { /// Derives the comb over the parameters' generator. /// /// `window` trades the comb's size against multiplies per exponentiation. @@ -21,12 +21,12 @@ impl BitZProver { /// # Panics /// /// [`FixedBasePow::new`] requires `window` in `1..=16`. - pub fn new(params: BitZParams, window: u32) -> Self { + pub fn new(params: BitZParams, window: u32) -> Self { let comb = FixedBasePow::new(params.generator(), window); Self { params, comb } } - pub fn params(&self) -> &BitZParams { + pub fn params(&self) -> &BitZParams { &self.params } @@ -42,12 +42,13 @@ mod tests { use field::gf128::smallest_generator; const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; #[test] fn the_comb_is_built_on_the_generator_the_parameters_name() { // `pow(1)` reads the base straight out of the comb. let params = - BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); + BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); let setup = BitZProver::new(params, 8); assert_eq!(setup.comb().pow(1), params.generator()); diff --git a/crates/spartan/Cargo.toml b/crates/spartan/Cargo.toml index cd68ad22..32036c27 100644 --- a/crates/spartan/Cargo.toml +++ b/crates/spartan/Cargo.toml @@ -24,7 +24,3 @@ divan = { workspace = true } field = { path = "../field", features = ["rand"] } rand = "0.9" rand_pcg = "0.9" - -[[bench]] -name = "sha256" -harness = false diff --git a/crates/spartan/src/lib.rs b/crates/spartan/src/lib.rs index 7ef48aca..70a43113 100644 --- a/crates/spartan/src/lib.rs +++ b/crates/spartan/src/lib.rs @@ -7,14 +7,12 @@ pub mod piop; pub mod sumcheck; pub use matrix::{ - PreparedConstraintMatrices, SpartanMatrixError, bigint_to_fq, build_assignment_mle, - build_product_mles, + PreparedConstraintMatrices, SpartanMatrixError, build_assignment_mle, build_product_mles, }; pub use piop::{ SpartanError, SpartanPiopProof, prove_spartan_piop, verify_spartan_proof, verify_spartan_with_mle_claim, }; - pub use sumcheck::{ InnerSumcheckOutput, OuterSumcheckOutput, OuterSumcheckProof, OuterSumcheckVerifierOutput, R1csProductMles, SumcheckError, SumcheckProof, SumcheckProverOutput, prove_inner_sumcheck, diff --git a/crates/spartan/src/matrix.rs b/crates/spartan/src/matrix.rs index 1ed7aeea..83e8bea0 100644 --- a/crates/spartan/src/matrix.rs +++ b/crates/spartan/src/matrix.rs @@ -3,21 +3,15 @@ use circuit::constraints::{ConstraintMatrices, SparseMatrix}; use circuit::matrix_products::{IntegerProducts, ModularVector, RuntimeModulus}; use circuit::witgen::PackedWitness; -use common::BitzRing; -use crypto_primitives::ConstField; -use field::{FqDefault, Q100}; -use num_bigint::{BigInt, BigUint}; -use num_traits::{Signed, ToPrimitive}; +use circuit::{BitWidth, IntoWords}; +use common::{BitzClaimField, BitzField}; use poly::DenseMultilinearExtension; use rayon::prelude::*; use sha2::{Digest, Sha256}; -use std::sync::LazyLock; use transcript::Encoding; use crate::sumcheck::R1csProductMles; -static FQ_DEFAULT_MODULUS: LazyLock = LazyLock::new(|| BigInt::from(Q100)); - /// Failures while preparing or evaluating Spartan's R1CS matrices. #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum SpartanMatrixError { @@ -98,10 +92,7 @@ impl ColumnChunkIndex { } } -impl PreparedConstraintMatrices -where - F: ConstField + Copy + Encoding<[u8]>, -{ +impl PreparedConstraintMatrices { pub fn new(matrices: ConstraintMatrices) -> Result { let (num_row_vars, num_column_vars) = r1cs_num_vars(&matrices)?; @@ -155,26 +146,19 @@ where } } -/// Reduces a signed integer canonically modulo Q100. -pub fn bigint_to_fq(value: &BigInt) -> FqDefault { - let modulus = &*FQ_DEFAULT_MODULUS; - let mut reduced = value % modulus; - if reduced.is_negative() { - reduced += modulus; - } - FqDefault::from( - reduced - .to_u128() - .expect("a canonical Q100 residue always fits a u128"), - ) -} +const PRIME_LIMBS: usize = 2; -/// Reduces exact `Ah`, `Bh`, and `Ch` values modulo Q100 and pads their row +/// Reduces exact `Ah`, `Bh`, and `Ch` values modulo `F::modulus` and pads their row /// tables with trailing zeros to the next power of two. -pub fn build_product_mles( +pub fn build_product_mles( products: &IntegerProducts, expected_rows: usize, -) -> Result, SpartanMatrixError> { +) -> Result, SpartanMatrixError> +where + F: BitzClaimField, + F::Integer: BitWidth + IntoWords, + Vec: for<'a> From<&'a ModularVector>, +{ for actual in [ products.a_mw.len(), products.b_mw.len(), @@ -188,7 +172,9 @@ pub fn build_product_mles( } } - let modulus = RuntimeModulus::<2>::new(BigUint::from(Q100)) + // We cannot define RuntimeModulus in stable Rust, so + // use Vec: for<'a> From<&'a ModularVector> as a workaround + let modulus = RuntimeModulus::::new(F::modulus()) .map_err(|_| SpartanMatrixError::InvalidModulus)?; let reduced = products.reduce_parallel(&modulus); let num_vars = padded_num_vars(expected_rows)?; @@ -200,16 +186,30 @@ pub fn build_product_mles( }) } +fn modular_vector_mle( + values: &ModularVector, + num_vars: usize, +) -> Result, SpartanMatrixError> +where + F: BitzClaimField, + Vec: for<'a> From<&'a ModularVector>, +{ + let zero = F::zero(); + let padded_len = 1usize << num_vars; + let mut evaluations: Vec = values.into(); + evaluations.resize(padded_len, zero); + + DenseMultilinearExtension::from_evaluations(num_vars, evaluations) + .map_err(|_| SpartanMatrixError::InvalidMleOperation) +} + /// Converts the packed assignment `h = M(1 || f)` into the selected field and /// pads it with trailing zeros to the next power-of-two column domain. The /// first bit must be the R1CS constant one. -pub fn build_assignment_mle( +pub fn build_assignment_mle( assignment: &PackedWitness, expected_columns: usize, -) -> Result, SpartanMatrixError> -where - F: ConstField + Copy, -{ +) -> Result, SpartanMatrixError> { if assignment.bit_len() != expected_columns { return Err(SpartanMatrixError::InvalidAssignmentLength { expected: expected_columns, @@ -225,21 +225,18 @@ where let mut evaluations = Vec::with_capacity(padded_len); evaluations.extend((0..assignment.bit_len()).map(|index| { if assignment.bit(index) { - F::ONE + F::one() } else { - F::ZERO + F::zero() } })); - evaluations.resize(padded_len, F::ZERO); + evaluations.resize(padded_len, F::zero()); DenseMultilinearExtension::from_evaluations(num_vars, evaluations) .map_err(|_| SpartanMatrixError::InvalidMleOperation) } -impl PreparedConstraintMatrices -where - F: ConstField + Copy, -{ +impl PreparedConstraintMatrices { /// Constructs /// /// `D(j) = sum_i eq(i,r_x) (A[i,j] + rho B[i,j] + rho^2 C[i,j])` @@ -286,17 +283,14 @@ where } } -fn bind_and_batch_with_num_vars( +fn bind_and_batch_with_num_vars( matrices: &ConstraintMatrices, column_chunks: &[ColumnChunkIndex; 3], row_point: &[F], rho: F, num_row_vars: usize, num_column_vars: usize, -) -> Result, SpartanMatrixError> -where - F: ConstField + Copy, -{ +) -> Result, SpartanMatrixError> { if row_point.len() != num_row_vars { return Err(SpartanMatrixError::InvalidRowPointLength { expected: num_row_vars, @@ -306,7 +300,7 @@ where let row_weights = poly::eq_table(row_point); let batched = [ - (&matrices.a, &column_chunks[0], F::ONE), + (&matrices.a, &column_chunks[0], F::one()), (&matrices.b, &column_chunks[1], rho), (&matrices.c, &column_chunks[2], rho * rho), ]; @@ -316,7 +310,7 @@ where // nonzeros landing in it: no two tasks write the same column, and each // task's writes stay within a cache-sized slice. let mut evaluations: Vec = - rayon::iter::repeat_n(F::ZERO, 1usize << num_column_vars).collect(); + rayon::iter::repeat_n(F::zero(), 1usize << num_column_vars).collect(); evaluations .par_chunks_mut(chunk_len) .enumerate() @@ -337,7 +331,7 @@ where .map_err(|_| SpartanMatrixError::InvalidMleOperation) } -fn evaluate_batched_with_num_vars( +fn evaluate_batched_with_num_vars( matrices: &ConstraintMatrices, column_chunks: &[ColumnChunkIndex; 3], row_point: &[F], @@ -345,10 +339,7 @@ fn evaluate_batched_with_num_vars( column_point: &[F], num_row_vars: usize, num_column_vars: usize, -) -> Result -where - F: ConstField + Copy, -{ +) -> Result { if row_point.len() != num_row_vars { return Err(SpartanMatrixError::InvalidRowPointLength { expected: num_row_vars, @@ -364,7 +355,7 @@ where let row_weights = poly::eq_table(row_point); let batched = [ - (&matrices.a, &column_chunks[0], F::ONE), + (&matrices.a, &column_chunks[0], F::one()), (&matrices.b, &column_chunks[1], rho), (&matrices.c, &column_chunks[2], rho * rho), ]; @@ -382,26 +373,28 @@ where .enumerate() .map(|(chunk, &chunk_weight)| { let base = chunk * chunk_len; - let mut chunk_sum = F::ZERO; + let mut chunk_sum = F::zero(); for (matrix, index, batch_scale) in batched { - let mut matrix_sum = F::ZERO; + let mut matrix_sum = F::zero(); for span in &index.spans[chunk] { let entries = &matrix.rows()[span.row].entries()[span.start..span.end]; - let span_sum = entries.iter().fold(F::ZERO, |sum, &(column, coefficient)| { - sum + low_weights[column - base] * coefficient - }); + let span_sum = entries + .iter() + .fold(F::zero(), |sum, &(column, coefficient)| { + sum + low_weights[column - base] * coefficient + }); matrix_sum += row_weights[span.row] * span_sum; } chunk_sum += batch_scale * matrix_sum; } chunk_weight * chunk_sum }) - .reduce(|| F::ZERO, |left, right| left + right); + .reduce(|| F::zero(), |left, right| left + right); Ok(evaluation) } -pub(crate) fn r1cs_num_vars( +pub(crate) fn r1cs_num_vars( matrices: &ConstraintMatrices, ) -> Result<(usize, usize), SpartanMatrixError> { matrices @@ -419,12 +412,9 @@ pub(crate) fn r1cs_num_vars( /// The digest domain is intentionally field-neutral. A protocol that supports /// more than one field must bind the field choice in its transcript session or /// instance; canonical coefficient encodings need not identify their field. -pub(crate) fn constraint_matrix_digest( +pub(crate) fn constraint_matrix_digest( matrices: &ConstraintMatrices, -) -> Result<[u8; 32], SpartanMatrixError> -where - F: ConstField + Copy + Encoding<[u8]>, -{ +) -> Result<[u8; 32], SpartanMatrixError> { matrices .validate_shape() .map_err(|_| SpartanMatrixError::InvalidR1csShape)?; @@ -490,30 +480,14 @@ pub(crate) fn padded_num_vars(logical_len: usize) -> Result, - num_vars: usize, -) -> Result, SpartanMatrixError> { - let zero = FqDefault::from(0u128); - let padded_len = 1usize << num_vars; - let mut evaluations: Vec<_> = values - .values() - .iter() - .map(|&[low, high]| FqDefault::from_limbs(low, high)) - .collect(); - evaluations.resize(padded_len, zero); - - DenseMultilinearExtension::from_evaluations(num_vars, evaluations) - .map_err(|_| SpartanMatrixError::InvalidMleOperation) -} - #[cfg(test)] mod tests { use circuit::constraints::{ConstraintMatrices, SparseBoolMatrix, SparseMatrix}; - use field::FqDefault; use rand::{Rng, SeedableRng}; use rand_pcg::Pcg64; + type F = field::FqDefault; + use super::{BIND_CHUNK_COLUMN_VARS, PreparedConstraintMatrices}; /// Three chunks of columns plus one chunk of padding, so rows straddle @@ -522,7 +496,7 @@ mod tests { const ROWS: usize = 37; const ENTRIES_PER_ROW: usize = 24; - fn random_sparse_matrix(rng: &mut Pcg64) -> SparseMatrix { + fn random_sparse_matrix(rng: &mut Pcg64) -> SparseMatrix { let rows = (0..ROWS) .map(|_| { let mut columns: Vec = (0..ENTRIES_PER_ROW) @@ -532,14 +506,14 @@ mod tests { columns.dedup(); columns .into_iter() - .map(|column| (column, FqDefault::from(u128::from(rng.random::())))) + .map(|column| (column, F::from(u128::from(rng.random::())))) .collect() }) .collect(); SparseMatrix::try_from_rows(COLUMNS, rows).unwrap() } - fn random_prepared_matrices(rng: &mut Pcg64) -> PreparedConstraintMatrices { + fn random_prepared_matrices(rng: &mut Pcg64) -> PreparedConstraintMatrices { let m = SparseBoolMatrix::try_from_rows(1, vec![Vec::new(); COLUMNS]).unwrap(); let a = random_sparse_matrix(rng); let b = random_sparse_matrix(rng); @@ -547,24 +521,24 @@ mod tests { PreparedConstraintMatrices::new(ConstraintMatrices { m, a, b, c }).unwrap() } - fn random_point(rng: &mut Pcg64, len: usize) -> Vec { + fn random_point(rng: &mut Pcg64, len: usize) -> Vec { (0..len) - .map(|_| FqDefault::from(u128::from(rng.random::()))) + .map(|_| F::from(u128::from(rng.random::()))) .collect() } /// `D(r_y)` by the direct triple loop over every nonzero. fn reference_evaluation( - matrices: &ConstraintMatrices, - row_point: &[FqDefault], - rho: FqDefault, - column_point: &[FqDefault], - ) -> FqDefault { + matrices: &ConstraintMatrices, + row_point: &[F], + rho: F, + column_point: &[F], + ) -> F { let row_weights = poly::eq_table(row_point); let column_weights = poly::eq_table(column_point); - let mut evaluation = FqDefault::from(0u128); + let mut evaluation = F::from(0u128); for (matrix, batch_scale) in [ - (&matrices.a, FqDefault::from(1u128)), + (&matrices.a, F::from(1u128)), (&matrices.b, rho), (&matrices.c, rho * rho), ] { @@ -619,7 +593,7 @@ mod tests { let prepared = random_prepared_matrices(&mut rng); let row_point = random_point(&mut rng, prepared.num_row_vars()); let column_point = random_point(&mut rng, prepared.num_column_vars()); - let rho = FqDefault::from(u128::from(rng.random::())); + let rho = F::from(u128::from(rng.random::())); let expected = reference_evaluation(prepared.matrices(), &row_point, rho, &column_point); let evaluation = prepared diff --git a/crates/spartan/src/piop.rs b/crates/spartan/src/piop.rs index 4a486417..332a16a0 100644 --- a/crates/spartan/src/piop.rs +++ b/crates/spartan/src/piop.rs @@ -1,11 +1,11 @@ //! Composition of Spartan's outer and inner sumchecks. use circuit::witgen::PackedWitness; -use crypto_primitives::ConstField; +use common::BitzField; use poly::{ DenseMultilinearExtension, MleClaimError, ScaledMleEvaluationClaim, make_equality_factors, }; -use transcript::{Encoding, ProverState, TranscriptChallenge, VerifierState}; +use transcript::{ProverState, VerifierState}; use crate::matrix::{PreparedConstraintMatrices, SpartanMatrixError, build_assignment_mle}; use crate::sumcheck::{ @@ -64,15 +64,12 @@ impl From for SpartanError { /// session or instance must bind that choice so proofs from different fields /// occupy distinct Fiat--Shamir domains. #[tracing::instrument(name = "Prove Spartan", skip_all)] -pub fn prove_spartan_piop( +pub fn prove_spartan_piop( transcript: &mut ProverState, matrices: &PreparedConstraintMatrices, products: &R1csProductMles, assignment: &DenseMultilinearExtension, -) -> Result<(SpartanPiopProof, ScaledMleEvaluationClaim), SpartanError> -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +) -> Result<(SpartanPiopProof, ScaledMleEvaluationClaim), SpartanError> { let num_row_vars = matrices.num_row_vars(); let num_column_vars = matrices.num_column_vars(); if products.az.num_vars() != num_row_vars @@ -91,7 +88,7 @@ where .collect::>(); let equality_factors = make_equality_factors(&tau).map_err(|_| SpartanMatrixError::InvalidMleOperation)?; - let outer = prove_outer_sumcheck(transcript, F::ZERO, equality_factors, products)?; + let outer = prove_outer_sumcheck(transcript, F::zero(), equality_factors, products)?; // The outer prover absorbed these evaluations before returning. let rho = transcript.squeeze::(); @@ -117,21 +114,18 @@ where /// Verifies both sumchecks and returns their terminal scaled assignment claim /// `D(r_y) * h(r_y) = final_claim`. #[tracing::instrument(name = "Verify Spartan", skip_all)] -pub fn verify_spartan_proof( +pub fn verify_spartan_proof( transcript: &mut VerifierState<'_>, matrices: &PreparedConstraintMatrices, proof: &SpartanPiopProof, -) -> Result, SpartanError> -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +) -> Result, SpartanError> { let num_row_vars = matrices.num_row_vars(); let num_column_vars = matrices.num_column_vars(); transcript.public_message(matrices.digest()); let tau = (0..num_row_vars) .map(|_| transcript.squeeze::()) .collect::>(); - let outer = proof.outer.verify(transcript, F::ZERO, &tau)?; + let outer = proof.outer.verify(transcript, F::zero(), &tau)?; // The outer verifier absorbed the product evaluations before returning. let rho = transcript.squeeze::(); @@ -159,16 +153,13 @@ where /// Boolean witness `f`. The eventual virtual opening protocol must enforce /// that map from a commitment to `f`; this witness-aware path checks the /// resulting assignment directly. -pub fn verify_spartan_with_mle_claim( +pub fn verify_spartan_with_mle_claim( transcript: &mut VerifierState<'_>, matrices: &PreparedConstraintMatrices, proof: &SpartanPiopProof, mle_claim: &ScaledMleEvaluationClaim, assignment: &PackedWitness, -) -> Result<(), SpartanError> -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +) -> Result<(), SpartanError> { let assignment = build_assignment_mle::(assignment, matrices.matrices().a.column_count())?; let expected_claim = verify_spartan_proof(transcript, matrices, proof)?; if mle_claim != &expected_claim { @@ -185,8 +176,8 @@ mod tests { constraints::{ConstraintMatrices, SparseBoolMatrix, SparseMatrix}, witgen::PackedWitness, }; - use crypto_primitives::ConstField; - use field::{F128, FqDefault}; + use crypto_primitives::Field; + use field::F128; use poly::DenseMultilinearExtension; use rand::{Rng, SeedableRng}; use rand_pcg::Pcg64; @@ -202,27 +193,29 @@ mod tests { const ROWS: usize = 5; const COLUMNS: usize = 7; + type F = field::FqDefault; + #[test] fn random_satisfying_r1cs_reduces_through_both_sumchecks() { - check_random_satisfying_r1cs::(FQ_SESSION); + check_random_satisfying_r1cs::(FQ_SESSION); check_random_satisfying_r1cs::(F128_SESSION); } #[test] fn verifier_rejects_proof_with_unsatisfied_witness() { - check_verifier_rejects_proof_with_unsatisfied_witness::(FQ_SESSION); + check_verifier_rejects_proof_with_unsatisfied_witness::(FQ_SESSION); check_verifier_rejects_proof_with_unsatisfied_witness::(F128_SESSION); } fn check_verifier_rejects_proof_with_unsatisfied_witness(session: &[u8]) where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, + F: Field + Copy + Encoding<[u8]> + TranscriptChallenge, { const SIZE: usize = 2; const UNSATISFIED_INSTANCE: &[u8] = b"unsatisfied-r1cs"; - let one = F::ONE; - let zero = F::ZERO; + let one = F::one(); + let zero = F::zero(); let satisfying_assignment = [one, zero]; let unsatisfied_assignment = [one, one]; let a = SparseMatrix::try_from_rows(SIZE, vec![vec![(1, one)], vec![(1, one)]]).unwrap(); @@ -267,7 +260,7 @@ mod tests { fn check_random_satisfying_r1cs(session: &[u8]) where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, + F: Field + Copy + Encoding<[u8]> + TranscriptChallenge, { let mut rng = Pcg64::seed_from_u64(0x5a17_c0de); let mut assignment_bits = Vec::with_capacity(COLUMNS); @@ -276,7 +269,7 @@ mod tests { let assignment_values: Vec<_> = assignment_bits .iter() .copied() - .map(|bit| if bit { F::ONE } else { F::ZERO }) + .map(|bit| if bit { F::one() } else { F::zero() }) .collect(); let a = random_sparse_matrix::(&mut rng); @@ -337,7 +330,7 @@ mod tests { complete_verifier.check_eof().unwrap(); } - fn random_sparse_matrix(rng: &mut Pcg64) -> SparseMatrix { + fn random_sparse_matrix(rng: &mut Pcg64) -> SparseMatrix { let rows = (0..ROWS) .map(|_| { // A nonzero constant-column entry keeps each row evaluation @@ -356,25 +349,25 @@ mod tests { SparseMatrix::try_from_rows(COLUMNS, rows).unwrap() } - fn multiply(matrix: &SparseMatrix, assignment: &[F]) -> Vec { + fn multiply(matrix: &SparseMatrix, assignment: &[F]) -> Vec { matrix .rows() .iter() .map(|row| { row.entries() .iter() - .fold(F::ZERO, |sum, &(column, coefficient)| { + .fold(F::zero(), |sum, &(column, coefficient)| { sum + coefficient * assignment[column] }) }) .collect() } - fn padded_mle(values: Vec) -> DenseMultilinearExtension { + fn padded_mle(values: Vec) -> DenseMultilinearExtension { let padded_len = values.len().max(1).next_power_of_two(); let num_vars = padded_len.ilog2() as usize; let mut evaluations = values; - evaluations.resize(padded_len, F::ZERO); + evaluations.resize(padded_len, F::zero()); DenseMultilinearExtension::from_evaluations(num_vars, evaluations).unwrap() } } diff --git a/crates/spartan/src/sumcheck.rs b/crates/spartan/src/sumcheck.rs index 1704ddde..c2d9b6c6 100644 --- a/crates/spartan/src/sumcheck.rs +++ b/crates/spartan/src/sumcheck.rs @@ -27,10 +27,12 @@ //! //! [*More Optimizations to Sum-Check Proving*]: https://eprint.iacr.org/2024/1210.pdf -use crypto_primitives::ConstField; +use common::BitzField; +use crypto_primitives::Semiring; use poly::DenseMultilinearExtension; use rayon::prelude::*; -use transcript::{Encoding, ProverState, TranscriptChallenge, VerifierState}; +use std::array; +use transcript::{ProverState, VerifierState}; /// Failures produced while reducing or checking a sumcheck claim. #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -53,10 +55,7 @@ pub struct SumcheckProof { pub round_polynomials: Vec<[F; COEFFS]>, } -impl SumcheckProof -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +impl SumcheckProof { /// Verifies the round reductions and returns `(r, final_claim)`. /// /// The caller supplies the expected number of rounds from the statement. @@ -79,7 +78,7 @@ where }); } - let zero = F::ZERO; + let zero = F::zero(); let mut current_claim = initial_claim; let mut eval_points = Vec::with_capacity(expected_rounds); @@ -215,10 +214,7 @@ pub struct InnerSumcheckOutput { pub witness_evaluation: F, } -impl OuterSumcheckProof -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +impl OuterSumcheckProof { /// Verifies the outer reduction and its terminal R1CS identity. #[tracing::instrument(name = "Verify outer sumcheck", skip_all)] pub fn verify( @@ -253,15 +249,12 @@ where /// session or instance must bind that choice so proofs from different fields /// occupy distinct Fiat–Shamir domains. #[tracing::instrument(name = "Prove outer sumcheck", skip_all)] -pub fn prove_outer_sumcheck( +pub fn prove_outer_sumcheck( transcript: &mut ProverState, initial_claim: F, (eq_low, eq_high): (DenseMultilinearExtension, DenseMultilinearExtension), products: &R1csProductMles, -) -> Result, SumcheckError> -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +) -> Result, SumcheckError> { let num_vars = products.az.num_vars(); if products.bz.num_vars() != num_vars || products.cz.num_vars() != num_vars { return Err(SumcheckError::InvalidProductDimensions); @@ -275,7 +268,7 @@ where return Err(SumcheckError::InvalidEqualityDimensions); } - let zero = F::ZERO; + let zero = F::zero(); let mut eq_low: Vec<_> = eq_low.into_iter().collect(); let mut eq_high: Vec<_> = eq_high.into_iter().collect(); // It's possible to avoid cloning here but its impact is negligible @@ -446,21 +439,18 @@ impl R1csProductTableBuffers { /// session or instance must bind that choice so proofs from different fields /// occupy distinct Fiat–Shamir domains. #[tracing::instrument(name = "Prove inner sumcheck", skip_all)] -pub fn prove_inner_sumcheck( +pub fn prove_inner_sumcheck( transcript: &mut ProverState, initial_claim: F, batched_matrix_mle: DenseMultilinearExtension, witness_mle: &DenseMultilinearExtension, -) -> Result, SumcheckError> -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ +) -> Result, SumcheckError> { let num_vars = batched_matrix_mle.num_vars(); if witness_mle.num_vars() != num_vars { return Err(SumcheckError::InvalidProductDimensions); } - let zero = F::ZERO; + let zero = F::zero(); let mut batched_matrix = batched_matrix_mle.evaluations; let mut current_claim = initial_claim; let mut eval_points = Vec::with_capacity(num_vars); @@ -562,13 +552,10 @@ where } #[inline] -fn compute_inner_pair_coefficients_without_linear( +fn compute_inner_pair_coefficients_without_linear( batched_matrix: [F; 2], witness: [F; 2], -) -> [F; 2] -where - F: ConstField + Copy, -{ +) -> [F; 2] { let [matrix_zero, matrix_one] = batched_matrix; let [witness_zero, witness_one] = witness; @@ -578,14 +565,14 @@ where ] } -fn sum_inner_round_coefficients_without_linear(batched_matrix: &[F], witness: &[F]) -> [F; 2] -where - F: ConstField + Copy, -{ +fn sum_inner_round_coefficients_without_linear( + batched_matrix: &[F], + witness: &[F], +) -> [F; 2] { debug_assert_eq!(batched_matrix.len(), witness.len()); debug_assert!(batched_matrix.len() >= 2); - let zero = F::ZERO; + let zero = F::zero(); let pair_count = batched_matrix.len() / 2; if should_parallelize(pair_count) { @@ -622,16 +609,13 @@ where } #[inline] -fn fold_inner_chunk( +fn fold_inner_chunk( batched_matrix: &[F], witness: &[F], batched_matrix_output: &mut [F], witness_output: &mut [F], challenge: F, -) -> [F; 2] -where - F: ConstField + Copy, -{ +) -> [F; 2] { debug_assert_eq!(batched_matrix.len(), 4); debug_assert_eq!(witness.len(), 4); debug_assert_eq!(batched_matrix_output.len(), 2); @@ -654,22 +638,19 @@ where /// Binds the current variable in both tables and simultaneously prepares the /// next round's `[c0, c2]`. The challenge has already been sampled, so this /// does not move any work across the Fiat-Shamir boundary. -fn fold_and_compute_next_inner_round_coefficients_without_linear( +fn fold_and_compute_next_inner_round_coefficients_without_linear( batched_matrix: &[F], witness: &[F], batched_matrix_output: &mut [F], witness_output: &mut [F], challenge: F, -) -> [F; 2] -where - F: ConstField + Copy, -{ +) -> [F; 2] { debug_assert_eq!(batched_matrix.len(), witness.len()); debug_assert!(batched_matrix.len() >= 4); debug_assert_eq!(batched_matrix_output.len(), batched_matrix.len() / 2); debug_assert_eq!(witness_output.len(), witness.len() / 2); - let zero = F::ZERO; + let zero = F::zero(); let chunk_count = batched_matrix.len() / 4; if should_parallelize(chunk_count) { @@ -707,29 +688,26 @@ where } #[inline] -fn add_coefficients(left: [F; COEFFS], right: [F; COEFFS]) -> [F; COEFFS] -where - F: ConstField + Copy, -{ - std::array::from_fn(|index| left[index] + right[index]) +fn add_coefficients( + left: [S; COEFFS], + right: [S; COEFFS], +) -> [S; COEFFS] { + array::from_fn(|index| left[index].clone() + &right[index]) } -fn sum_coefficients( +fn sum_coefficients( len: usize, - contribution: impl Fn(usize) -> [F; COEFFS] + Sync, -) -> [F; COEFFS] -where - F: ConstField + Copy, -{ - let zero = F::ZERO; + contribution: impl Fn(usize) -> [S; COEFFS] + Sync, +) -> [S; COEFFS] { + let make_zero_arr = || array::repeat::<_, COEFFS>(S::zero()); if should_parallelize(len) { (0..len) .into_par_iter() .map(&contribution) - .reduce(|| [zero; COEFFS], add_coefficients::) + .reduce(make_zero_arr, add_coefficients::) } else { - (0..len).fold([zero; COEFFS], |sum, index| { + (0..len).fold(make_zero_arr(), |sum, index| { add_coefficients(sum, contribution(index)) }) } @@ -743,10 +721,7 @@ fn should_parallelize(work_items: usize) -> bool { } #[inline] -fn cubic_contribution(eq: [F; 2], az: [F; 2], bz: [F; 2], cz: [F; 2]) -> [F; 3] -where - F: ConstField + Copy, -{ +fn cubic_contribution(eq: [F; 2], az: [F; 2], bz: [F; 2], cz: [F; 2]) -> [F; 3] { let [eq_zero, eq_one] = eq; let [az_zero, az_one] = az; let [bz_zero, bz_one] = bz; @@ -787,12 +762,12 @@ fn reconstruct_round_coefficients [F; COEFFS] where - F: ConstField + Copy, + F: BitzField, { assert!(INPUT_COEFFS >= 1); assert_eq!(COEFFS, INPUT_COEFFS + 1); - let zero = F::ZERO; + let zero = F::zero(); let mut coefficients = [zero; COEFFS]; coefficients[0] = coefficients_without_linear[0]; coefficients[2..].copy_from_slice(&coefficients_without_linear[1..]); @@ -806,11 +781,11 @@ where } #[inline] -fn evaluate_polynomial(coefficients: &[F; COEFFS], point: F) -> F -where - F: ConstField + Copy, -{ - let zero = F::ZERO; +fn evaluate_polynomial( + coefficients: &[F; COEFFS], + point: F, +) -> F { + let zero = F::zero(); coefficients .iter() .rev() @@ -826,7 +801,7 @@ where /// absorbs the completed polynomial, samples `r_i`, and updates /// `current_claim` to `g_i(r_i)`. fn recover_full_round_polynomial_and_sample_next_challenge< - F, + F: BitzField, const INPUT_COEFFS: usize, const COEFFS: usize, >( @@ -835,11 +810,8 @@ fn recover_full_round_polynomial_and_sample_next_challenge< coefficients_without_linear: [F; INPUT_COEFFS], round_polynomials: &mut Vec<[F; COEFFS]>, eval_points: &mut Vec, -) -> F -where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, -{ - let zero = F::ZERO; +) -> F { + let zero = F::zero(); let coefficients = reconstruct_round_coefficients(*current_claim, coefficients_without_linear); let at_one = coefficients .iter() @@ -867,10 +839,7 @@ struct EqualityPairs<'a, F> { low_mask: usize, } -impl<'a, F> EqualityPairs<'a, F> -where - F: ConstField + Copy, -{ +impl<'a, F: BitzField> EqualityPairs<'a, F> { fn new(low: &'a [F], high: &'a [F]) -> Self { debug_assert!(low.len().is_power_of_two()); debug_assert!(high.len().is_power_of_two()); @@ -902,13 +871,10 @@ where } /// Computes `[c0, c2, c3]` from adjacent pairs in the current product tables. -fn compute_coefficients_without_linear( +fn compute_coefficients_without_linear( products: &R1csProductTableBuffers, equality_pairs: EqualityPairs<'_, F>, -) -> [F; 3] -where - F: ConstField + Copy, -{ +) -> [F; 3] { let pair_count = products.len() / 2; sum_coefficients(pair_count, |pair| { @@ -923,19 +889,13 @@ where } #[inline] -fn interpolate_pair(pair: [F; 2], challenge: F) -> F -where - F: ConstField + Copy, -{ +fn interpolate_pair(pair: [F; 2], challenge: F) -> F { let [zero, one] = pair; zero + challenge * (one - zero) } #[inline] -fn fold_two_pairs(values: &[F], challenge: F) -> [F; 2] -where - F: ConstField + Copy, -{ +fn fold_two_pairs(values: &[F], challenge: F) -> [F; 2] { debug_assert_eq!(values.len(), 4); [ interpolate_pair([values[0], values[1]], challenge), @@ -944,7 +904,7 @@ where } #[inline] -fn fold_product_chunk( +fn fold_product_chunk( az: &[F], bz: &[F], cz: &[F], @@ -952,10 +912,7 @@ fn fold_product_chunk( bz_output: &mut [F], cz_output: &mut [F], challenge: F, -) -> [[F; 2]; 3] -where - F: ConstField + Copy, -{ +) -> [[F; 2]; 3] { debug_assert_eq!(az_output.len(), 2); debug_assert_eq!(bz_output.len(), 2); debug_assert_eq!(cz_output.len(), 2); @@ -972,10 +929,7 @@ where } /// Folds one evaluation table into an already initialized destination. -fn fold_table(input: &[F], output: &mut [F], challenge: F) -where - F: ConstField + Copy, -{ +fn fold_table(input: &[F], output: &mut [F], challenge: F) { debug_assert_eq!(input.len(), 2 * output.len()); let fold = |(pair, value): (&[F], &mut F)| { @@ -992,13 +946,11 @@ where } /// Folds all three product tables into reusable scratch storage. -fn fold_product_tables( +fn fold_product_tables( input: &R1csProductTableBuffers, output: &mut R1csProductTableBuffers, challenge: F, -) where - F: ConstField + Copy, -{ +) { debug_assert_eq!(input.len(), 2 * output.len()); fold_table(&input.az, &mut output.az, challenge); @@ -1008,18 +960,15 @@ fn fold_product_tables( /// Folds all three product tables and accumulates the next round polynomial /// from the freshly folded pairs. -fn fold_products_and_compute_next( +fn fold_products_and_compute_next( input: &R1csProductTableBuffers, output: &mut R1csProductTableBuffers, challenge: F, equality_pairs: EqualityPairs<'_, F>, -) -> [F; 3] -where - F: ConstField + Copy, -{ +) -> [F; 3] { debug_assert_eq!(input.len(), 2 * output.len()); - let zero = F::ZERO; + let zero = F::zero(); let accumulate = |sum: [F; 3], chunk: usize, az: &[F], @@ -1076,15 +1025,13 @@ where } /// Folds the high equality factor and all product tables together. -fn fold_products_and_eq( +fn fold_products_and_eq( products: &R1csProductTableBuffers, product_output: &mut R1csProductTableBuffers, eq: &[F], eq_output: &mut [F], challenge: F, -) where - F: ConstField + Copy, -{ +) { debug_assert_eq!(products.len(), eq.len()); debug_assert_eq!(products.len(), 2 * product_output.len()); debug_assert_eq!(eq.len(), 2 * eq_output.len()); @@ -1095,7 +1042,7 @@ fn fold_products_and_eq( #[cfg(test)] mod tests { - use field::{F128, FqDefault}; + use field::F128; use rand::{Rng, SeedableRng}; use rand_pcg::Pcg64; use transcript::{build_prover, build_verifier}; @@ -1107,8 +1054,10 @@ mod tests { const INNER_SESSION: &[u8] = b"spartan/inner-sumcheck/test"; const F128_INNER_SESSION: &[u8] = b"spartan/inner-sumcheck/f128/test"; - fn fq(value: u128) -> FqDefault { - FqDefault::from(value) + type F = field::FqDefault; + + fn fq(value: u128) -> F { + F::from(value) } #[test] @@ -1125,15 +1074,15 @@ mod tests { .iter() .map(|round| { prover.public_message(round); - prover.squeeze::() + prover.squeeze::() }) .collect(); - let next_prover_challenge = prover.squeeze::(); + let next_prover_challenge = prover.squeeze::(); let transcript_proof = prover.finish(); let mut verifier = build_verifier(SESSION, instance, &transcript_proof); let (verifier_points, final_claim) = sumcheck.verify(&mut verifier, fq(20), 2).unwrap(); - let next_verifier_challenge = verifier.squeeze::(); + let next_verifier_challenge = verifier.squeeze::(); let expected_final_claim = second_round .iter() @@ -1166,7 +1115,7 @@ mod tests { #[test] fn zero_round_sumcheck_preserves_the_initial_claim() { - let sumcheck_proof = SumcheckProof:: { + let sumcheck_proof = SumcheckProof:: { round_polynomials: vec![], }; let transcript_proof = transcript::Proof::default(); @@ -1181,7 +1130,7 @@ mod tests { #[test] fn sumcheck_rejects_zero_coefficient_rounds() { - let sumcheck = SumcheckProof:: { + let sumcheck = SumcheckProof:: { round_polynomials: vec![[]], }; let transcript_proof = transcript::Proof::default(); @@ -1201,20 +1150,14 @@ mod tests { } /// Builds random product MLEs and equality factors from `poly::eq_table`. - fn build_outer_sumcheck_inputs(num_vars: usize) -> OuterSumcheckTestInputs - where - F: ConstField + Copy, - { + fn build_outer_sumcheck_inputs(num_vars: usize) -> OuterSumcheckTestInputs { build_outer_sumcheck_inputs_with_split(num_vars, num_vars / 2) } - fn build_outer_sumcheck_inputs_with_split( + fn build_outer_sumcheck_inputs_with_split( num_vars: usize, split: usize, - ) -> OuterSumcheckTestInputs - where - F: ConstField + Copy, - { + ) -> OuterSumcheckTestInputs { assert!(num_vars < usize::BITS as usize); assert!(split <= num_vars); @@ -1257,10 +1200,7 @@ mod tests { } /// Runs only the outer prover and outer verifier on prebuilt inputs. - fn check_outer_sumcheck(session: &[u8], inputs: OuterSumcheckTestInputs) - where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, - { + fn check_outer_sumcheck(session: &[u8], inputs: OuterSumcheckTestInputs) { let OuterSumcheckTestInputs { tau, eq_factors, @@ -1272,13 +1212,13 @@ mod tests { let mut prover = build_prover(session, &instance); let prover_output = - prove_outer_sumcheck(&mut prover, F::ZERO, eq_factors, &products).unwrap(); + prove_outer_sumcheck(&mut prover, F::zero(), eq_factors, &products).unwrap(); let proof = prover.finish(); let mut verifier = build_verifier(session, &instance, &proof); let verifier_output = prover_output .proof - .verify(&mut verifier, F::ZERO, &tau) + .verify(&mut verifier, F::zero(), &tau) .unwrap(); verifier.check_eof().unwrap(); @@ -1322,10 +1262,7 @@ mod tests { ); } - fn check_outer_sumcheck_input_construction() - where - F: ConstField + Copy, - { + fn check_outer_sumcheck_input_construction() { for num_vars in [0, 1, 3, 10] { let inputs = build_outer_sumcheck_inputs::(num_vars); let table_len = 1usize << num_vars; @@ -1356,7 +1293,7 @@ mod tests { #[test] fn outer_sumcheck_inputs_have_pointwise_products_and_factored_eq() { - check_outer_sumcheck_input_construction::(); + check_outer_sumcheck_input_construction::(); check_outer_sumcheck_input_construction::(); } @@ -1365,7 +1302,7 @@ mod tests { for split in 0..=5 { check_outer_sumcheck( SESSION, - build_outer_sumcheck_inputs_with_split::(5, split), + build_outer_sumcheck_inputs_with_split::(5, split), ); check_outer_sumcheck( F128_SESSION, @@ -1395,10 +1332,10 @@ mod tests { ), Err(SumcheckError::InvalidProductDimensions) ); - let challenge_after_product_error = invalid_product_prover.squeeze::(); + let challenge_after_product_error = invalid_product_prover.squeeze::(); let mut invalid_equality_prover = build_prover(SESSION, instance); - let inputs = build_outer_sumcheck_inputs::(1); + let inputs = build_outer_sumcheck_inputs::(1); assert_eq!( prove_outer_sumcheck( &mut invalid_equality_prover, @@ -1411,10 +1348,10 @@ mod tests { ), Err(SumcheckError::InvalidEqualityDimensions) ); - let challenge_after_equality_error = invalid_equality_prover.squeeze::(); + let challenge_after_equality_error = invalid_equality_prover.squeeze::(); let mut clean_prover = build_prover(SESSION, instance); - let clean_challenge = clean_prover.squeeze::(); + let clean_challenge = clean_prover.squeeze::(); assert_eq!(challenge_after_product_error, clean_challenge); assert_eq!(challenge_after_equality_error, clean_challenge); } @@ -1466,57 +1403,44 @@ mod tests { #[test] fn outer_sumcheck_zero_vars() { - check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(0)); + check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(0)); check_outer_sumcheck(F128_SESSION, build_outer_sumcheck_inputs::(0)); } #[test] fn outer_sumcheck_one_var() { - check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(1)); + check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(1)); check_outer_sumcheck(F128_SESSION, build_outer_sumcheck_inputs::(1)); } #[test] fn outer_sumcheck_three_vars() { - check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(3)); + check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(3)); check_outer_sumcheck(F128_SESSION, build_outer_sumcheck_inputs::(3)); } #[test] fn outer_sumcheck_ten_vars() { - check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(10)); + check_outer_sumcheck(SESSION, build_outer_sumcheck_inputs::(10)); check_outer_sumcheck(F128_SESSION, build_outer_sumcheck_inputs::(10)); } #[test] fn inner_sumcheck_one_variable_has_expected_quadratic() { - let batched_matrix = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(2u128), FqDefault::from(5u128)], - ) - .unwrap(); - let witness = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(3u128), FqDefault::from(7u128)], - ) - .unwrap(); + let batched_matrix = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(2u128), F::from(5u128)]) + .unwrap(); + let witness = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(3u128), F::from(7u128)]) + .unwrap(); let mut prover = build_prover(INNER_SESSION, b"one-variable"); - let output = prove_inner_sumcheck( - &mut prover, - FqDefault::from(41u128), - batched_matrix, - &witness, - ) - .unwrap(); + let output = + prove_inner_sumcheck(&mut prover, F::from(41u128), batched_matrix, &witness).unwrap(); assert_eq!( output.sumcheck.proof.round_polynomials, - vec![[ - FqDefault::from(6u128), - FqDefault::from(17u128), - FqDefault::from(12u128), - ]] + vec![[F::from(6u128), F::from(17u128), F::from(12u128),]] ); } @@ -1524,37 +1448,22 @@ mod tests { fn inner_sumcheck_binds_lowest_index_variable_first() { let batched_matrix = DenseMultilinearExtension::from_evaluations( 2, - [2u128, 5, 11, 17] - .into_iter() - .map(FqDefault::from) - .collect(), + [2u128, 5, 11, 17].into_iter().map(F::from).collect(), ) .unwrap(); let witness = DenseMultilinearExtension::from_evaluations( 2, - [3u128, 7, 13, 19] - .into_iter() - .map(FqDefault::from) - .collect(), + [3u128, 7, 13, 19].into_iter().map(F::from).collect(), ) .unwrap(); let mut prover = build_prover(INNER_SESSION, b"lowest-variable-first"); - let output = prove_inner_sumcheck( - &mut prover, - FqDefault::from(507u128), - batched_matrix, - &witness, - ) - .unwrap(); + let output = + prove_inner_sumcheck(&mut prover, F::from(507u128), batched_matrix, &witness).unwrap(); assert_eq!( output.sumcheck.proof.round_polynomials[0], - [ - FqDefault::from(149u128), - FqDefault::from(161u128), - FqDefault::from(48u128), - ] + [F::from(149u128), F::from(161u128), F::from(48u128),] ); } @@ -1569,7 +1478,7 @@ mod tests { for num_column_vars in [0, 1, 3, 12, 13] { check_inner_sumcheck( INNER_SESSION, - build_inner_sumcheck_inputs::(2, num_column_vars), + build_inner_sumcheck_inputs::(2, num_column_vars), ); check_inner_sumcheck( F128_INNER_SESSION, @@ -1578,20 +1487,17 @@ mod tests { } } - fn build_inner_sumcheck_inputs( + fn build_inner_sumcheck_inputs( num_row_vars: usize, num_column_vars: usize, - ) -> InnerSumcheckTestInputs - where - F: ConstField + Copy, - { + ) -> InnerSumcheckTestInputs { assert!(num_row_vars < usize::BITS as usize); assert!(num_column_vars < usize::BITS as usize); let num_rows = 1usize << num_row_vars; let num_columns = 1usize << num_column_vars; let matrix_len = num_rows.checked_mul(num_columns).unwrap(); - let zero = F::ZERO; + let zero = F::zero(); let mut rng = Pcg64::seed_from_u64( 0x1a2b_3c4d ^ ((num_row_vars as u64) << 32) ^ num_column_vars as u64, ); @@ -1653,10 +1559,7 @@ mod tests { } } - fn check_inner_sumcheck(session: &[u8], inputs: InnerSumcheckTestInputs) - where - F: ConstField + Copy + Encoding<[u8]> + TranscriptChallenge, - { + fn check_inner_sumcheck(session: &[u8], inputs: InnerSumcheckTestInputs) { let InnerSumcheckTestInputs { initial_claim, batched_matrix_mle, @@ -1705,9 +1608,9 @@ mod tests { #[test] fn inner_sumcheck_supports_zero_variables() { - let batched_matrix = DenseMultilinearExtension::zero_vars(FqDefault::from(5u128)); - let witness = DenseMultilinearExtension::zero_vars(FqDefault::from(7u128)); - let initial_claim = FqDefault::from(35u128); + let batched_matrix = DenseMultilinearExtension::zero_vars(F::from(5u128)); + let witness = DenseMultilinearExtension::zero_vars(F::from(7u128)); + let initial_claim = F::from(35u128); let mut prover = build_prover(INNER_SESSION, b"zero-variables"); let output = @@ -1716,49 +1619,36 @@ mod tests { assert!(output.sumcheck.proof.round_polynomials.is_empty()); assert!(output.sumcheck.eval_points.is_empty()); assert_eq!(output.sumcheck.final_claim, initial_claim); - assert_eq!(output.batched_matrix_evaluation, FqDefault::from(5u128)); - assert_eq!(output.witness_evaluation, FqDefault::from(7u128)); + assert_eq!(output.batched_matrix_evaluation, F::from(5u128)); + assert_eq!(output.witness_evaluation, F::from(7u128)); let mut control = build_prover(INNER_SESSION, b"zero-variables"); - assert_eq!( - prover.squeeze::(), - control.squeeze::() - ); + assert_eq!(prover.squeeze::(), control.squeeze::()); } #[test] fn inner_sumcheck_rejects_mismatched_dimensions() { - let batched_matrix = DenseMultilinearExtension::from_evaluations( - 1, - vec![FqDefault::from(1u128), FqDefault::from(2u128)], - ) - .unwrap(); + let batched_matrix = + DenseMultilinearExtension::from_evaluations(1, vec![F::from(1u128), F::from(2u128)]) + .unwrap(); let witness = DenseMultilinearExtension::from_evaluations( 2, vec![ - FqDefault::from(1u128), - FqDefault::from(2u128), - FqDefault::from(3u128), - FqDefault::from(4u128), + F::from(1u128), + F::from(2u128), + F::from(3u128), + F::from(4u128), ], ) .unwrap(); let mut prover = build_prover(INNER_SESSION, b"mismatched-dimensions"); assert_eq!( - prove_inner_sumcheck( - &mut prover, - FqDefault::from(0u128), - batched_matrix, - &witness - ), + prove_inner_sumcheck(&mut prover, F::from(0u128), batched_matrix, &witness), Err(SumcheckError::InvalidProductDimensions) ); let mut control = build_prover(INNER_SESSION, b"mismatched-dimensions"); - assert_eq!( - prover.squeeze::(), - control.squeeze::() - ); + assert_eq!(prover.squeeze::(), control.squeeze::()); } } diff --git a/crates/tests/examples/dump_bitz.rs b/crates/tests/examples/dump_bitz.rs index e873a5ed..e0666bff 100644 --- a/crates/tests/examples/dump_bitz.rs +++ b/crates/tests/examples/dump_bitz.rs @@ -12,6 +12,8 @@ use crypto_primitives::LiftElement; use support::{hex, write_binary, write_witness}; use tests::{Instance, Q, verifier_transcript}; +type F = field::FqDefault; + fn main() -> Result<(), Box> { const USAGE: &str = "usage: dump_bitz "; let args: Vec = std::env::args().skip(1).collect(); @@ -26,7 +28,7 @@ fn main() -> Result<(), Box> { Shape::for_log_bits(log_bits).map_err(|error| format!("invalid shape: {error:?}"))?; std::fs::create_dir_all(out)?; let (t, s) = (shape.log_rows(), shape.log_columns()); - let mut instance = Instance::honest(shape, seed); + let mut instance = Instance::::honest(shape, seed); let mut transcript = instance.transcript.take().unwrap(); let started = Instant::now(); diff --git a/crates/tests/examples/verify_bitz.rs b/crates/tests/examples/verify_bitz.rs index de6f74f1..2705aeeb 100644 --- a/crates/tests/examples/verify_bitz.rs +++ b/crates/tests/examples/verify_bitz.rs @@ -6,6 +6,8 @@ use common::Shape; use tests::{Instance, verifier_transcript}; use transcript::Proof; +type F = field::FqDefault; + fn main() -> Result<(), Box> { const USAGE: &str = "usage: verify_bitz "; let args: Vec = std::env::args().skip(1).collect(); @@ -20,7 +22,7 @@ fn main() -> Result<(), Box> { narg_string: std::fs::read(narg_file)?, hints: std::fs::read(hints_file)?, }; - let instance = Instance::honest(shape, seed); + let instance = Instance::::honest(shape, seed); let started = std::time::Instant::now(); let result = instance.verifier.verify( &instance.claim, diff --git a/crates/tests/src/lib.rs b/crates/tests/src/lib.rs index 4da32491..4a43969e 100644 --- a/crates/tests/src/lib.rs +++ b/crates/tests/src/lib.rs @@ -9,9 +9,8 @@ //! The fixtures live here rather than under `tests/` so they compile once //! rather than once per test binary. -use common::{BitTable, BitZParams, LinearClaim, Root, Shape}; -use crypto_primitives::LiftElement; -use field::{F128, Fq, gf128::smallest_generator}; +use common::{BitTable, BitZParams, BitzClaimField, LinearClaim, Root, Shape}; +use field::{F128, gf128::smallest_generator}; use pcs::{HashKind, LigeritoProfile, Pcs, ProverData}; use rand_chacha::ChaCha8Rng; use rand_core::{Rng, SeedableRng}; @@ -36,13 +35,13 @@ pub fn packed_witness(shape: Shape, rng: &mut impl Rng) -> Vec { /// A claim that actually holds, with the witness it is about, before any /// commitment. What the fold round needs and nothing the opening does. #[derive(Debug, Clone)] -pub struct HonestClaim { - pub params: BitZParams, - pub claim: LinearClaim>, +pub struct HonestClaim { + pub params: BitZParams, + pub claim: LinearClaim, pub packed: Vec, } -impl HonestClaim { +impl HonestClaim { /// Builds a random witness and the target its own fold produces, so the /// claim is true by construction rather than by asserting the code agrees /// with itself. @@ -51,25 +50,25 @@ impl HonestClaim { /// with `eta_j` read bit by bit — not from the reconstruction the verifier /// runs. pub fn new(shape: Shape, rng: &mut impl Rng) -> Self { - let params = BitZParams::::new(shape, smallest_generator()).unwrap(); + let params = BitZParams::::new(shape, smallest_generator()).unwrap(); let packed = packed_witness(shape, rng); - let row_weights: Vec> = (0..shape.rows()) - .map(|_| Fq::from(sample_below_q(rng))) + let row_weights: Vec = (0..shape.rows()) + .map(|_| F::from(sample_below_q(rng))) .collect(); - let column_weights: Vec> = (0..shape.columns()) - .map(|_| Fq::from(sample_below_q(rng))) + let column_weights: Vec = (0..shape.columns()) + .map(|_| F::from(sample_below_q(rng))) .collect(); let table = params.table(&packed).unwrap(); - let exponents: Vec = row_weights.iter().map(|weight| weight.lift()).collect(); - let target: Fq = (0..shape.columns()) + let exponents: Vec = row_weights.iter().map(|weight| weight.lift()).collect(); + let target: F = (0..shape.columns()) .map(|column| { - let fold: u128 = (0..shape.rows()) + let fold: F::Integer = (0..shape.rows()) .filter(|&row| table.bit(column, row)) - .map(|row| exponents[row]) + .map(|row| &exponents[row]) .sum(); - column_weights[column] * Fq::from(fold) + F::from(fold) * column_weights[column] }) .sum(); @@ -88,11 +87,11 @@ impl HonestClaim { } /// An instance whose claim actually holds, committed under a real scheme. -pub struct Instance { - pub params: BitZParams, - pub prover: prover::BitZProver, - pub verifier: verifier::BitZVerifier, - pub claim: LinearClaim>, +pub struct Instance { + pub params: BitZParams, + pub prover: prover::BitZProver, + pub verifier: verifier::BitZVerifier, + pub claim: LinearClaim, pub pcs: Pcs, pub com: Root, pub data: ProverData, @@ -100,7 +99,7 @@ pub struct Instance { pub packed: Vec, } -impl Instance { +impl Instance { /// [`HonestClaim::new`] under the `Fast` profile, with a setup per role. pub fn honest(shape: Shape, seed: u64) -> Self { let mut rng = ChaCha8Rng::seed_from_u64(seed); @@ -132,7 +131,7 @@ impl Instance { } /// The same instance under a different claimed value. - pub fn with_target(&self, target: Fq) -> LinearClaim> { + pub fn with_target(&self, target: F) -> LinearClaim { LinearClaim::new( &self.params, self.claim.row_weights().to_vec(), diff --git a/crates/tests/tests/fold.rs b/crates/tests/tests/fold.rs index 2d2e9c15..858d45d4 100644 --- a/crates/tests/tests/fold.rs +++ b/crates/tests/tests/fold.rs @@ -1,17 +1,18 @@ //! The fold round, prover against verifier. use common::{BitZParams, FoldError, LinearClaim}; -use field::{F128, Fq, gf128::smallest_generator}; -use num_traits::{ConstOne, ConstZero}; +use crypto_primitives::{BaseField, WithAssociatedInteger}; +use field::{F128, gf128::smallest_generator}; +use num_traits::{Bounded, ConstOne, ConstZero, ToBytes}; use prover::{BitZProver, SendError}; -use tests::{ - Instance, Q, WINDOW, narrow_shape, prover_transcript, verifier_transcript, wide_shape, -}; +use tests::{Instance, WINDOW, narrow_shape, prover_transcript, verifier_transcript, wide_shape}; use transcript::Proof; use verifier::{BitZVerifier, ReceiveError}; +type F = field::FqDefault; + /// Runs an honest prover and returns the round it produced with its proof. -fn prove(instance: &Instance) -> (common::Fold, Proof) { +fn prove(instance: &Instance) -> (common::Fold<::Integer>, Proof) { let mut transcript = prover_transcript(); let round = instance .prover @@ -22,10 +23,10 @@ fn prove(instance: &Instance) -> (common::Fold, Proof) { /// A proof carrying `folds` and nothing else, which is the whole wire format /// of this round. -fn forge(folds: &[u128]) -> Proof { +fn forge>>(folds: &[I]) -> Proof { let mut transcript = prover_transcript(); - for &fold in folds { - transcript.prover_message(&fold.to_le_bytes()); + for fold in folds { + transcript.prover_message(&fold.to_le_bytes().as_ref()); } transcript.finish() } @@ -33,7 +34,7 @@ fn forge(folds: &[u128]) -> Proof { #[test] fn the_two_sides_agree_on_every_shape_the_profile_admits() { for shape in [narrow_shape(), wide_shape()] { - let instance = Instance::honest(shape, 7); + let instance = Instance::::honest(shape, 7); let (sent, proof) = prove(&instance); let mut transcript = verifier_transcript(&proof); @@ -53,7 +54,7 @@ fn the_two_sides_agree_on_every_shape_the_profile_admits() { fn the_proof_carries_only_the_folds() { // Sending the images too would double this round's proof size: at // k_2 = 2^21 the second copy is 32 MiB. - let instance = Instance::honest(narrow_shape(), 2); + let instance = Instance::::honest(narrow_shape(), 2); let (_, proof) = prove(&instance); assert_eq!( @@ -68,7 +69,7 @@ fn the_challenge_depends_on_the_folds() { // The folds are written and absorbed by the same call, so the challenge // must move when they do. That is what binds the images the grand product // consumes, since `u -> g^u` is injective over the admitted range. - let instance = Instance::honest(narrow_shape(), 3); + let instance = Instance::::honest(narrow_shape(), 3); let (round, _) = prove(&instance); let echoed = forge(&round.folds); @@ -96,16 +97,17 @@ fn a_fold_at_the_bound_is_accepted_and_one_past_it_is_not() { let shape = narrow_shape(); let packed = vec![F128::new(u64::MAX, u64::MAX); (1 << shape.log_bits()) / 128]; - let fold = (shape.rows() as u128) * (Q - 1); - let params = BitZParams::::new(shape, smallest_generator()).unwrap(); + let modulus = F::modulus(); + let fold = (shape.rows() as u128) * (modulus - 1); + let params = BitZParams::::new(shape, smallest_generator()).unwrap(); let prover = BitZProver::new(params, WINDOW); let verifier = BitZVerifier::new(params, WINDOW); let table = params.table(&packed).unwrap(); let claim = LinearClaim::new( ¶ms, - vec![Fq::from(Q - 1); shape.rows()], - vec![Fq::ONE; shape.columns()], - Fq::from(fold) * Fq::from(shape.columns() as u128), + vec![F::max_value(); shape.rows()], + vec![F::ONE; shape.columns()], + F::from(fold) * F::from(shape.columns() as u128), ) .unwrap(); @@ -129,7 +131,7 @@ fn a_fold_at_the_bound_is_accepted_and_one_past_it_is_not() { fn the_range_check_fires_before_the_reconstruction() { // A proof violating both. The order is fixed, so the earlier obligation is // the one that must be reported. - let instance = Instance::honest(narrow_shape(), 6); + let instance = Instance::::honest(narrow_shape(), 6); let shape = instance.params.shape(); let over = vec![instance.params.fold_bound() + 1; shape.columns()]; @@ -151,11 +153,11 @@ fn the_range_check_fires_before_the_reconstruction() { #[test] fn folds_that_do_not_reconstruct_the_target_are_rejected() { - let instance = Instance::honest(narrow_shape(), 5); + let instance = Instance::::honest(narrow_shape(), 5); let (_, proof) = prove(&instance); // The proof is honest; the claim it is replayed against is not. - let retargeted = instance.with_target(instance.claim.target() + Fq::ONE); + let retargeted = instance.with_target(instance.claim.target() + F::ONE); let mut transcript = verifier_transcript(&proof); assert_eq!( instance.verifier.receive_fold(&retargeted, &mut transcript), @@ -165,8 +167,8 @@ fn folds_that_do_not_reconstruct_the_target_are_rejected() { #[test] fn a_witness_of_a_different_shape_is_refused_before_anything_is_written() { - let instance = Instance::honest(narrow_shape(), 8); - let other = Instance::honest(wide_shape(), 8); + let instance = Instance::::honest(narrow_shape(), 8); + let other = Instance::::honest(wide_shape(), 8); let mut transcript = prover_transcript(); assert_eq!( @@ -180,7 +182,7 @@ fn a_witness_of_a_different_shape_is_refused_before_anything_is_written() { #[test] fn a_truncated_proof_is_refused_rather_than_read_past() { - let instance = Instance::honest(narrow_shape(), 10); + let instance = Instance::::honest(narrow_shape(), 10); let (_, proof) = prove(&instance); let mut short = proof.clone(); @@ -198,14 +200,14 @@ fn a_truncated_proof_is_refused_rather_than_read_past() { fn an_all_zero_witness_folds_to_zero_and_still_round_trips() { let shape = narrow_shape(); let packed = vec![F128::ZERO; (1 << shape.log_bits()) / 128]; - let params = BitZParams::::new(shape, smallest_generator()).unwrap(); + let params = BitZParams::::new(shape, smallest_generator()).unwrap(); let prover = BitZProver::new(params, WINDOW); let verifier = BitZVerifier::new(params, WINDOW); let claim = LinearClaim::new( ¶ms, - vec![Fq::from(Q - 1); shape.rows()], - vec![Fq::from(3u128); shape.columns()], - Fq::from(0u128), + vec![F::max_value(); shape.rows()], + vec![F::from(3u128); shape.columns()], + F::from(0u128), ) .unwrap(); let table = params.table(&packed).unwrap(); @@ -231,7 +233,7 @@ fn a_round_is_refused_when_its_parts_do_not_match_the_shape() { assert_eq!( common::Fold::new( &shape, - vec![0; shape.columns()], + vec![0_u64; shape.columns()], vec![F128::ONE; shape.columns()], Vec::new(), vec![F128::ONE; shape.log_columns()], diff --git a/crates/tests/tests/host.rs b/crates/tests/tests/host.rs index 39ed8616..7977c4a1 100644 --- a/crates/tests/tests/host.rs +++ b/crates/tests/tests/host.rs @@ -8,8 +8,10 @@ use host::wire_proof; use tests::{Instance, narrow_shape, verifier_transcript, wide_shape}; +type F = field::FqDefault; + /// Runs an honest prover and hands back what a caller would ship. -fn shipped(instance: &mut Instance) -> Vec { +fn shipped(instance: &mut Instance) -> Vec { let mut transcript = instance.transcript.take().unwrap(); instance .prover @@ -28,7 +30,7 @@ fn shipped(instance: &mut Instance) -> Vec { #[test] fn a_proof_survives_the_round_trip_through_bytes() { for shape in [narrow_shape(), wide_shape()] { - let mut instance = Instance::honest(shape, 41); + let mut instance = Instance::::honest(shape, 41); let proof_bytes = shipped(&mut instance); let proof = wire_proof::decode(&proof_bytes).expect("its own encoding"); @@ -58,7 +60,7 @@ fn a_tampered_fold_is_left_for_the_verifier_to_catch() { // The container frames but does not authenticate: a tampered fold decodes // cleanly and the sponge refuses it on replay. let shape = wide_shape(); - let mut instance = Instance::honest(shape, 43); + let mut instance = Instance::::honest(shape, 43); let proof_bytes = shipped(&mut instance); let mut tampered = proof_bytes.clone(); diff --git a/crates/tests/tests/large.rs b/crates/tests/tests/large.rs index 0f94b66c..81a97654 100644 --- a/crates/tests/tests/large.rs +++ b/crates/tests/tests/large.rs @@ -8,7 +8,6 @@ //! This test is designed to test what's possible in under a minute. use common::LinearClaim; -use field::Fq; use num_traits::ConstOne; use prover::BitZProver; use rand_chacha::ChaCha8Rng; @@ -16,10 +15,12 @@ use rand_core::SeedableRng; use tests::{HonestClaim, WINDOW, large_shape, prover_transcript, verifier_transcript}; use verifier::{BitZVerifier, ReceiveError}; +type F = field::FqDefault; + #[test] fn the_fold_round_trips_on_the_large_shape() { let shape = large_shape(); - let honest = HonestClaim::new(shape, &mut ChaCha8Rng::seed_from_u64(31)); + let honest = HonestClaim::::new(shape, &mut ChaCha8Rng::seed_from_u64(31)); let prover = BitZProver::new(honest.params, WINDOW); let verifier = BitZVerifier::new(honest.params, WINDOW); @@ -56,7 +57,7 @@ fn the_fold_round_trips_on_the_large_shape() { &honest.params, honest.claim.row_weights().to_vec(), honest.claim.column_weights().to_vec(), - honest.claim.target() + Fq::ONE, + honest.claim.target() + F::ONE, ) .unwrap(); assert_eq!( diff --git a/crates/tests/tests/prove.rs b/crates/tests/tests/prove.rs index 07cc90b5..7c592a73 100644 --- a/crates/tests/tests/prove.rs +++ b/crates/tests/tests/prove.rs @@ -1,7 +1,7 @@ //! The top-level prove and verify, through the real opening. use common::{Root, TableError}; -use field::{F128, Fq}; +use field::F128; use num_traits::{ConstOne, ConstZero}; use pcs::{HashKind, LigeritoProfile, Pcs, VerifyError as PcsVerifyError}; use prover::ProveError; @@ -9,7 +9,9 @@ use tests::{Instance, narrow_shape, verifier_transcript, wide_shape}; use transcript::Proof; use verifier::{ReceiveError, VerifyError}; -fn prove(instance: &mut Instance) -> Proof { +type F = field::FqDefault; + +fn prove(instance: &mut Instance) -> Proof { let mut transcript = instance.transcript.take().unwrap(); instance .prover @@ -27,7 +29,7 @@ fn prove(instance: &mut Instance) -> Proof { #[test] fn an_honest_proof_verifies_on_both_floor_shapes() { for shape in [narrow_shape(), wide_shape()] { - let mut instance = Instance::honest(shape, 31); + let mut instance = Instance::::honest(shape, 31); let proof = prove(&mut instance); instance @@ -44,7 +46,7 @@ fn an_honest_proof_verifies_on_both_floor_shapes() { #[test] fn a_proof_replayed_under_a_different_commitment_is_refused() { - let mut instance = Instance::honest(narrow_shape(), 32); + let mut instance = Instance::::honest(narrow_shape(), 32); let proof = prove(&mut instance); // Binding a different root changes the fold batching point, so GKR rejects. @@ -61,13 +63,13 @@ fn a_proof_replayed_under_a_different_commitment_is_refused() { #[test] fn the_statement_is_bound_before_the_first_challenge() { - let mut instance = Instance::honest(narrow_shape(), 33); + let mut instance = Instance::::honest(narrow_shape(), 33); let proof = prove(&mut instance); // Same folds, same commitment, a claim that differs only in its claimed // value. The fold's own reconstruction rejects it, which is the check the // binding backs up rather than replaces. - let retargeted = instance.with_target(instance.claim.target() + Fq::ONE); + let retargeted = instance.with_target(instance.claim.target() + F::ONE); assert_eq!( instance.verifier.verify( &retargeted, @@ -81,7 +83,7 @@ fn the_statement_is_bound_before_the_first_challenge() { #[test] fn a_proof_with_trailing_bytes_is_refused() { - let mut instance = Instance::honest(narrow_shape(), 34); + let mut instance = Instance::::honest(narrow_shape(), 34); let mut proof = prove(&mut instance); proof.hints.push(0); @@ -102,8 +104,8 @@ fn an_opening_against_another_commitment_is_refused() { // the verifier is given, the folds are over the witness the claim describes, // and the GKR claim is true of that witness. Only the codeword and // the Merkle tree the opening reads belong to a different commitment. - let proved = Instance::honest(narrow_shape(), 35); - let mut committed = Instance::honest(narrow_shape(), 36); + let proved = Instance::::honest(narrow_shape(), 35); + let mut committed = Instance::::honest(narrow_shape(), 36); let mut transcript = committed.transcript.take().unwrap(); proved @@ -133,7 +135,7 @@ fn an_opening_against_another_commitment_is_refused() { fn a_tampered_opening_proof_is_refused() { // The opening rides the hint channel, which the sponge never sees, so // nothing upstream of the opening notices this. The opening itself must. - let mut instance = Instance::honest(narrow_shape(), 37); + let mut instance = Instance::::honest(narrow_shape(), 37); let mut proof = prove(&mut instance); let middle = proof.hints.len() / 2; proof.hints[middle] ^= 0xff; @@ -153,7 +155,7 @@ fn a_tampered_opening_proof_is_refused() { fn a_proof_verified_under_a_different_profile_is_refused() { // OOD binds PCS parameters before the first fold challenge, so a different // profile changes the fold transcript and GKR rejects. - let mut instance = Instance::honest(narrow_shape(), 38); + let mut instance = Instance::::honest(narrow_shape(), 38); let slim = Pcs::new( instance.params.shape(), LigeritoProfile::Slim, @@ -175,7 +177,7 @@ fn a_proof_verified_under_a_different_profile_is_refused() { #[test] fn a_witness_of_the_wrong_length_is_refused_without_changing_the_commitment_transcript() { - let mut instance = Instance::honest(narrow_shape(), 39); + let mut instance = Instance::::honest(narrow_shape(), 39); let mut transcript = instance.transcript.take().unwrap(); assert_eq!( diff --git a/crates/tests/tests/virtual_prove.rs b/crates/tests/tests/virtual_prove.rs index 452637ae..de117eeb 100644 --- a/crates/tests/tests/virtual_prove.rs +++ b/crates/tests/tests/virtual_prove.rs @@ -9,7 +9,7 @@ use common::{ BitZParams, LinearClaim, Root, Shape, TableError, TransposedWeights, VirtualMap, VirtualMapError, VirtualStatement, }; -use field::{F128, Fq, gf128::smallest_generator}; +use field::{F128, gf128::smallest_generator}; use num_traits::{ConstOne, ConstZero}; use pcs::{HashKind, LigeritoProfile, Pcs, ProverData}; use prover::{BitZProver, ProveError, VirtualWitness}; @@ -17,6 +17,8 @@ use tests::{Q, WINDOW, prover_transcript, verifier_transcript}; use transcript::{Proof, ProverState}; use verifier::{BitZVerifier, VerifyError}; +type F = field::Fq; + /// `h[0] = 1`, `h[1] = f[0]`, `h[128] = f[1]`, `h[129] = f[0] XOR f[1]`. /// All other virtual bits are zero. struct Map(u8); @@ -46,9 +48,9 @@ impl VirtualMap for Map { } struct Instance { - params: BitZParams, + params: BitZParams, committed_shape: Shape, - claim: LinearClaim>, + claim: LinearClaim, committed_bits: Vec, virtual_bits: Vec, pcs: Pcs, @@ -61,12 +63,12 @@ impl Instance { fn new() -> Self { let claim_shape = Shape::new(7, 15).unwrap(); let committed_shape = Shape::new(8, 14).unwrap(); - let params = BitZParams::::new(claim_shape, smallest_generator()).unwrap(); + let params = BitZParams::::new(claim_shape, smallest_generator()).unwrap(); let claim = LinearClaim::new( ¶ms, - vec![Fq::ONE; claim_shape.rows()], - vec![Fq::ONE; claim_shape.columns()], - Fq::from(3u128), + vec![F::ONE; claim_shape.rows()], + vec![F::ONE; claim_shape.columns()], + F::from(3u128), ) .unwrap(); let mut committed_bits = vec![F128::ZERO; 1 << committed_shape.log_packed_len()]; @@ -91,7 +93,7 @@ impl Instance { } } - fn statement(&self) -> VirtualStatement<'_, Q, Map> { + fn statement(&self) -> VirtualStatement<'_, F, Map> { VirtualStatement::new(self.params, self.committed_shape, &Map(7), &self.claim).unwrap() } @@ -114,7 +116,7 @@ impl Instance { fn verify( &self, - statement: &VirtualStatement<'_, Q, Map>, + statement: &VirtualStatement<'_, F, Map>, root: Root, proof: &Proof, ) -> Result<(), VerifyError> { @@ -163,7 +165,7 @@ fn changed_virtual_statements_are_rejected() { // Changing a weight on a zero virtual row preserves the integer target. // Rejection must therefore depend on the statement, not a false claim. let mut rows = instance.claim.row_weights().to_vec(); - rows[2] += Fq::ONE; + rows[2] += F::ONE; let changed_claim = LinearClaim::new( &instance.params, rows, @@ -340,7 +342,7 @@ fn sha256_virtual_inner_product_opens_the_committed_bits() { let claim_shape = Shape::new(7, 15).unwrap(); let committed_shape = Shape::new(8, 14).unwrap(); - let params = BitZParams::::new(claim_shape, smallest_generator()).unwrap(); + let params = BitZParams::::new(claim_shape, smallest_generator()).unwrap(); let pack = |witness: &PackedWitness, shape: Shape| { assert!(witness.bit_len() <= 1 << shape.log_bits()); let mut packed: Vec<_> = witness @@ -354,10 +356,10 @@ fn sha256_virtual_inner_product_opens_the_committed_bits() { let committed_bits = pack(&f, committed_shape); let virtual_bits = pack(&h, claim_shape); let rows: Vec<_> = (0..claim_shape.rows()) - .map(|row| Fq::::from((row + 1) as u128)) + .map(|row| F::from((row + 1) as u128)) .collect(); let columns: Vec<_> = (0..claim_shape.columns()) - .map(|column| Fq::::from((column + 1) as u128)) + .map(|column| F::from((column + 1) as u128)) .collect(); let target = (0..h.bit_len()) .filter(|&index| h.bit(index)) diff --git a/crates/transcript/src/challenge.rs b/crates/transcript/src/challenge.rs index e1920097..23413838 100644 --- a/crates/transcript/src/challenge.rs +++ b/crates/transcript/src/challenge.rs @@ -1,8 +1,7 @@ //! Typed Fiat–Shamir challenge sampling. -use field::{F128, Fq}; - use crate::{ProverState, VerifierState}; +use field::{F128, Fq}; /// A type that knows how to construct itself from transcript squeezes. /// @@ -63,7 +62,7 @@ impl TranscriptChallenge for F128 { #[cfg(test)] mod tests { - use field::{F128, FqDefault, Q100}; + use field::{F128, Q100}; use super::TranscriptChallenge; use crate::{build_prover, build_verifier}; @@ -71,6 +70,8 @@ mod tests { const SESSION: &[u8] = b"transcript/typed-challenge/test"; const INSTANCE: &[u8] = b"fq-rejection-sampling"; + type F = field::FqDefault; + #[test] fn fq_rejects_the_incomplete_final_interval() { let rejection_remainder = (u128::MAX % Q100 + 1) % Q100; @@ -81,24 +82,24 @@ mod tests { assert_eq!(max_accepted % Q100, Q100 - 1); assert_eq!((max_accepted + 1) % Q100, 0); - let challenge = FqDefault::from_squeezes(|| { + let challenge = F::from_squeezes(|| { squeezes += 1; candidates.next().unwrap() }); assert_eq!(squeezes, 3); - assert_eq!(challenge, FqDefault::from(Q100 - 1)); + assert_eq!(challenge, F::from(Q100 - 1)); } #[test] fn prover_and_verifier_squeeze_in_lockstep() { let mut prover = build_prover(SESSION, INSTANCE); - let prover_fq = prover.squeeze::(); + let prover_fq = prover.squeeze::(); let prover_f128 = prover.squeeze::(); let proof = prover.finish(); let mut verifier = build_verifier(SESSION, INSTANCE, &proof); - let verifier_fq = verifier.squeeze::(); + let verifier_fq = verifier.squeeze::(); let verifier_f128 = verifier.squeeze::(); assert_eq!(verifier_fq, prover_fq); diff --git a/crates/transcript/tests/golden.rs b/crates/transcript/tests/golden.rs index df697bc1..ebbeafc5 100644 --- a/crates/transcript/tests/golden.rs +++ b/crates/transcript/tests/golden.rs @@ -1,7 +1,7 @@ //! Pins the challenge sequence and wire bytes for one fixed transcript. //! A mismatch means the framing changed and every existing proof is invalid. -use field::{F128, FqDefault}; +use field::F128; use transcript::{Proof, build_prover, build_verifier}; const SESSION: &[u8] = b"golden-session"; @@ -16,12 +16,14 @@ const HINTS: &str = "aaaaaaaaaa"; const C1: &str = "1d1bcda36aa1541c3752f7b23fedef81"; const C2: &str = "e80729d4aa76e8e4d28dd085143eb3d5"; +type F = field::FqDefault; + fn prove() -> (Proof, F128, F128) { let mut prover = build_prover(SESSION, INSTANCE); prover.prover_message(&MSG_F128); let c1: F128 = prover.verifier_message(); prover.hint(&HINT); - prover.prover_message(&FqDefault::from(MSG_FQ)); + prover.prover_message(&F::from(MSG_FQ)); let c2: F128 = prover.verifier_message(); (prover.finish(), c1, c2) } @@ -43,10 +45,7 @@ fn verifier_replays_the_golden_transcript() { assert_eq!(verifier.prover_message::().unwrap(), MSG_F128); assert_eq!(verifier.verifier_message::(), c1); assert_eq!(verifier.hint::<[u8; 5]>().unwrap(), HINT); - assert_eq!( - verifier.prover_message::().unwrap(), - FqDefault::from(MSG_FQ) - ); + assert_eq!(verifier.prover_message::().unwrap(), F::from(MSG_FQ)); assert_eq!(verifier.verifier_message::(), c2); verifier.check_eof().unwrap(); } diff --git a/crates/transcript/tests/hints.rs b/crates/transcript/tests/hints.rs index 081ff1cd..ccd61494 100644 --- a/crates/transcript/tests/hints.rs +++ b/crates/transcript/tests/hints.rs @@ -1,7 +1,7 @@ //! The hint channel: round-trips beside the narg string, never touches the //! sponge. -use field::{F128, FqDefault}; +use field::F128; use transcript::{Proof, build_prover, build_verifier}; const SESSION: &[u8] = b"hints-session"; @@ -12,12 +12,14 @@ const HINT_BYTES: [u8; 5] = [0xAA; 5]; const HINT_F128: F128 = F128::new(7, 8); const HINT_U32: u32 = 42; +type F = field::FqDefault; + fn prove() -> (Proof, F128, F128) { let mut prover = build_prover(SESSION, INSTANCE); prover.prover_message(&MSG_1); let c1: F128 = prover.verifier_message(); prover.hint(&HINT_BYTES); - prover.prover_message(&FqDefault::from(12345u128)); + prover.prover_message(&F::from(12345u128)); prover.hint(&HINT_F128); let c2: F128 = prover.verifier_message(); prover.hint(&HINT_U32); @@ -32,10 +34,7 @@ fn mixed_messages_and_hints_round_trip() { assert_eq!(verifier.prover_message::().unwrap(), MSG_1); assert_eq!(verifier.verifier_message::(), c1); assert_eq!(verifier.hint::<[u8; 5]>().unwrap(), HINT_BYTES); - assert_eq!( - verifier.prover_message::().unwrap(), - FqDefault::from(12345u128) - ); + assert_eq!(verifier.prover_message::().unwrap(), F::from(12345u128)); assert_eq!(verifier.hint::().unwrap(), HINT_F128); assert_eq!(verifier.verifier_message::(), c2); assert_eq!(verifier.hint::().unwrap(), HINT_U32); @@ -51,7 +50,7 @@ fn tampered_hint_leaves_challenges_unchanged() { assert_eq!(verifier.prover_message::().unwrap(), MSG_1); assert_eq!(verifier.verifier_message::(), c1); assert_ne!(verifier.hint::<[u8; 5]>().unwrap(), HINT_BYTES); - verifier.prover_message::().unwrap(); + verifier.prover_message::().unwrap(); verifier.hint::().unwrap(); assert_eq!(verifier.verifier_message::(), c2); verifier.hint::().unwrap(); @@ -67,7 +66,7 @@ fn truncated_hints_fail_the_read() { verifier.prover_message::().unwrap(); verifier.verifier_message::(); verifier.hint::<[u8; 5]>().unwrap(); - verifier.prover_message::().unwrap(); + verifier.prover_message::().unwrap(); verifier.hint::().unwrap(); verifier.verifier_message::(); assert!(verifier.hint::().is_err()); @@ -82,7 +81,7 @@ fn trailing_hint_bytes_fail_eof() { verifier.prover_message::().unwrap(); verifier.verifier_message::(); verifier.hint::<[u8; 5]>().unwrap(); - verifier.prover_message::().unwrap(); + verifier.prover_message::().unwrap(); verifier.hint::().unwrap(); verifier.verifier_message::(); verifier.hint::().unwrap(); diff --git a/crates/verifier/src/fold.rs b/crates/verifier/src/fold.rs index 4c23e0b9..1af0fbd4 100644 --- a/crates/verifier/src/fold.rs +++ b/crates/verifier/src/fold.rs @@ -1,9 +1,11 @@ //! The fold round: read the column folds, check them, then take the //! challenge. -use common::{Fold, FoldError, LinearClaim, column_images, reconstruct, row_images}; - use crate::BitZVerifier; +use common::{ + BitzClaimField, Fold, FoldError, LinearClaim, column_images, reconstruct, row_images, +}; +use num_traits::FromBytes; use transcript::VerifierState; /// A fold the verifier rejects. @@ -19,7 +21,7 @@ pub enum ReceiveError { Fold(FoldError), } -impl BitZVerifier { +impl BitZVerifier { /// Reads the fold round and checks it. /// /// The proof carries only the folds; their images are derived here rather than @@ -36,21 +38,21 @@ impl BitZVerifier { #[tracing::instrument(name = "Verify column folds", skip_all)] pub fn receive_fold( &self, - claim: &LinearClaim>, + claim: &LinearClaim, transcript: &mut VerifierState<'_>, - ) -> Result { + ) -> Result, ReceiveError> { let shape = self.params().shape(); let folds = (0..shape.columns()) .map(|_| { transcript - .prover_message::<[u8; 16]>() - .map(u128::from_le_bytes) + .prover_message::<::Bytes>() + .map(|bs| F::Integer::from_le_bytes(&bs)) .map_err(|_| ReceiveError::MalformedProof) }) .collect::, _>>()?; - if folds.iter().any(|&fold| fold > self.fold_bound()) { + if folds.iter().any(|fold| fold > self.fold_bound()) { return Err(ReceiveError::FoldOutOfRange); } if reconstruct(claim, &folds).map_err(ReceiveError::Fold)? != claim.target() { diff --git a/crates/verifier/src/reduce.rs b/crates/verifier/src/reduce.rs index eec8a599..fcb78c42 100644 --- a/crates/verifier/src/reduce.rs +++ b/crates/verifier/src/reduce.rs @@ -21,9 +21,9 @@ pub enum ReduceError { } #[tracing::instrument(name = "Verify grand-product reduction", skip_all)] -pub(crate) fn gkr_reduce( +pub(crate) fn gkr_reduce( transcript: &mut VerifierState, - fold: &Fold, + fold: &Fold, shape: &Shape, ) -> Result { // Each layer halves the row count, leaving one product per column. @@ -60,13 +60,14 @@ mod round_trip_ai_test { use super::*; - const Q: u128 = (1 << 114) - 11; + const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; fn shape() -> Shape { Shape::new(7, 15).unwrap() } - fn params() -> BitZParams { + fn params() -> BitZParams { BitZParams::new(shape(), smallest_generator()).unwrap() } diff --git a/crates/verifier/src/setup.rs b/crates/verifier/src/setup.rs index 492dbcb3..414a60b2 100644 --- a/crates/verifier/src/setup.rs +++ b/crates/verifier/src/setup.rs @@ -1,6 +1,6 @@ //! The verifier's derived setup. -use common::BitZParams; +use common::{BitZParams, BitzClaimField}; use field::FixedBasePow; /// The parameters, with the comb table and the fold bound derived from them. @@ -9,19 +9,19 @@ use field::FixedBasePow; /// verifier is its only reader: the prover never range-checks a fold it /// produced itself. #[derive(Debug)] -pub struct BitZVerifier { - params: BitZParams, +pub struct BitZVerifier { + params: BitZParams, comb: FixedBasePow, - fold_bound: u128, + fold_bound: F::Integer, } -impl BitZVerifier { +impl BitZVerifier { /// Derives the comb over the parameters' generator, and the fold bound. /// /// # Panics /// /// [`FixedBasePow::new`] requires `window` in `1..=16`. - pub fn new(params: BitZParams, window: u32) -> Self { + pub fn new(params: BitZParams, window: u32) -> Self { let comb = FixedBasePow::new(params.generator(), window); let fold_bound = params.fold_bound(); Self { @@ -31,7 +31,7 @@ impl BitZVerifier { } } - pub fn params(&self) -> &BitZParams { + pub fn params(&self) -> &BitZParams { &self.params } @@ -40,8 +40,8 @@ impl BitZVerifier { } /// The largest fold this verifier accepts, `k_1 (Q - 1)`. - pub fn fold_bound(&self) -> u128 { - self.fold_bound + pub fn fold_bound(&self) -> &F::Integer { + &self.fold_bound } } @@ -52,12 +52,13 @@ mod tests { use field::gf128::smallest_generator; const Q114: u128 = (1 << 114) - 11; + type F = field::Fq; #[test] fn the_comb_is_built_on_the_generator_the_parameters_name() { // `pow(1)` reads the base straight out of the comb. let params = - BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); + BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); let setup = BitZVerifier::new(params, 8); assert_eq!(setup.comb().pow(1), params.generator()); diff --git a/crates/verifier/src/verify.rs b/crates/verifier/src/verify.rs index 0510f1bb..7b47d9f0 100644 --- a/crates/verifier/src/verify.rs +++ b/crates/verifier/src/verify.rs @@ -1,7 +1,8 @@ //! `VerifyBitZ`. -use common::{LinearClaim, OpeningQuery, Root, VirtualMap, VirtualMapError, VirtualStatement}; -use field::Fq; +use common::{ + BitzClaimField, LinearClaim, OpeningQuery, Root, VirtualMap, VirtualMapError, VirtualStatement, +}; use pcs::{CommitScheme, Commitment, Pcs, StatementBinding, VerifyError as OpeningVerifyError}; use transcript::VerifierState; @@ -24,11 +25,11 @@ pub enum VerifyError { TrailingData, } -impl BitZVerifier { +impl BitZVerifier { /// Receives the commitment's OOD claim and verifies the virtual BitZ proof. pub fn verify_virtual( &self, - statement: &VirtualStatement<'_, Q, impl VirtualMap>, + statement: &VirtualStatement<'_, F, impl VirtualMap>, pcs: &Pcs, root: Root, mut transcript: VerifierState<'_>, @@ -47,7 +48,7 @@ impl BitZVerifier { /// Receives the commitment's OOD claim and verifies the BitZ proof. pub fn verify( &self, - claim: &LinearClaim>, + claim: &LinearClaim, pcs: &Pcs, root: Root, mut transcript: VerifierState<'_>, @@ -57,7 +58,6 @@ impl BitZVerifier { .map_err(VerifyError::Opening)?; self.verify_with_commitment(claim, pcs, &commitment, transcript) } - /// Verifies a claim on `h = M (1 || f)` against the commitment to `f`. /// /// Build the setup from `statement.params().claim()` and match the commitment's @@ -72,7 +72,7 @@ impl BitZVerifier { #[tracing::instrument(name = "Verify virtual BitZ", skip_all)] pub fn verify_virtual_with_commitment( &self, - statement: &VirtualStatement<'_, Q, impl VirtualMap>, + statement: &VirtualStatement<'_, F, impl VirtualMap>, pcs: &Pcs, commitment: &Commitment, mut transcript: VerifierState<'_>, @@ -107,7 +107,7 @@ impl BitZVerifier { #[tracing::instrument(name = "Verify BitZ", skip_all)] pub fn verify_with_commitment( &self, - claim: &LinearClaim>, + claim: &LinearClaim, pcs: &Pcs, commitment: &Commitment, mut transcript: VerifierState<'_>, @@ -140,10 +140,11 @@ impl BitZVerifier { /// The caller binds the statement before this call and verifies the opening afterward. pub(crate) fn fold_and_reduce( &self, - claim: &LinearClaim>, + claim: &LinearClaim, transcript: &mut VerifierState<'_>, ) -> Result { - // Step 2 is absent: Q is fixed, and BitZParams::new checks its fold bound. + // Step 2 is absent: the field modulus is fixed, and BitZParams::new + // checks its fold bound. // Step 3: read the folds, range-check them, reconstruct against mu. let fold = self diff --git a/tooling/cli/Cargo.toml b/tooling/cli/Cargo.toml index a1090a2c..f8fcce66 100644 --- a/tooling/cli/Cargo.toml +++ b/tooling/cli/Cargo.toml @@ -12,6 +12,7 @@ argh.workspace = true circuit.workspace = true common.workspace = true field = { workspace = true, features = ["spongefish"] } +num-bigint.workspace = true num-traits.workspace = true rand.workspace = true rayon.workspace = true @@ -28,8 +29,11 @@ verifier.workspace = true [dev-dependencies] divan.workspace = true -num-bigint.workspace = true [[bench]] name = "circuits" harness = false + +[[bench]] +name = "sha256_spartan" +harness = false diff --git a/tooling/cli/benches/circuits.rs b/tooling/cli/benches/circuits.rs index 13fec1c1..822d3b08 100644 --- a/tooling/cli/benches/circuits.rs +++ b/tooling/cli/benches/circuits.rs @@ -1,12 +1,16 @@ //! End-to-end and per-stage benchmarks; instance generation is outside timing. use bitz_cli::{ - benchmark, + ProjectBigIntToFq, benchmark, circuits::{BuiltinCircuit, CircuitInstance}, end_to_end::CircuitProofSystem, }; use divan::Bencher; +type R = num_bigint::BigInt; +type F = field::FqDefault; +type Proj = ProjectBigIntToFq; + fn main() { divan::main(); } @@ -17,23 +21,30 @@ fn instance(circuit: BuiltinCircuit) -> (CircuitInstance, Vec) { (statement, inputs) } -fn setup(circuit: BuiltinCircuit) -> (CircuitProofSystem, Vec) { +fn setup(circuit: BuiltinCircuit) -> (CircuitProofSystem, Vec) { let (statement, inputs) = instance(circuit); - (CircuitProofSystem::new(statement).unwrap(), inputs) + ( + CircuitProofSystem::<_, F>::new::(statement).unwrap(), + inputs, + ) } #[divan::bench(args = BuiltinCircuit::ALL)] fn end_to_end(bencher: Bencher, circuit: BuiltinCircuit) { bencher .with_inputs(|| instance(circuit)) - .bench_local_values(|(statement, inputs)| benchmark::run(statement, &inputs).unwrap()); + .bench_local_values(|(statement, inputs)| { + benchmark::run::<_, F, R, Proj>(statement, &inputs).unwrap() + }); } #[divan::bench(args = BuiltinCircuit::ALL)] fn circuit_setup(bencher: Bencher, circuit: BuiltinCircuit) { bencher .with_inputs(|| CircuitInstance::random(circuit, None, None).unwrap()) - .bench_local_values(|statement| CircuitProofSystem::new(statement).unwrap()); + .bench_local_values(|statement| { + CircuitProofSystem::<_, F>::new::(statement).unwrap() + }); } #[divan::bench(args = BuiltinCircuit::ALL)] diff --git a/crates/spartan/benches/sha256.rs b/tooling/cli/benches/sha256_spartan.rs similarity index 86% rename from crates/spartan/benches/sha256.rs rename to tooling/cli/benches/sha256_spartan.rs index 40f56539..0066d51b 100644 --- a/crates/spartan/benches/sha256.rs +++ b/tooling/cli/benches/sha256_spartan.rs @@ -6,15 +6,15 @@ use std::sync::{Arc, Mutex}; +use bitz_cli::{ProjectBigIntToFq, ProjectConstraint}; use circuit::constraints::ConstraintGenerator; use circuit::sha256::sha256_block_aligned_circuit; use circuit::witgen::ProductWitgen; use divan::{AllocProfiler, Bencher, black_box}; -use field::FqDefault; use poly::{DenseMultilinearExtension, ScaledMleEvaluationClaim}; use spartan::{ - PreparedConstraintMatrices, R1csProductMles, SpartanPiopProof, bigint_to_fq, - build_assignment_mle, build_product_mles, prove_spartan_piop, verify_spartan_proof, + PreparedConstraintMatrices, R1csProductMles, SpartanPiopProof, build_assignment_mle, + build_product_mles, prove_spartan_piop, verify_spartan_proof, }; use transcript::{Proof, build_prover, build_verifier}; @@ -28,13 +28,17 @@ const INSTANCE: &[u8] = b"bench"; /// 2^22-bit committed shape. const BLOCKS: &[usize] = &[608]; +type R = num_bigint::BigInt; +type F = field::FqDefault; +type Proj = ProjectBigIntToFq; + /// An R1CS instance with a witness satisfying it. #[derive(Debug, Clone)] struct R1csInstanceWitness { - instance: PreparedConstraintMatrices, - witness: DenseMultilinearExtension, + instance: PreparedConstraintMatrices, + witness: DenseMultilinearExtension, /// Holds precomputed `Az`, `Bz` and `Cz` tables. - products: R1csProductMles, + products: R1csProductMles, } static R1CS_INSTANCE_WITNESSES: Mutex)>> = @@ -80,10 +84,11 @@ fn build(blocks: usize) -> R1csInstanceWitness { let (_witness, assignment_bits, exact_products) = witgen.into_parts(); // Lower to Q100 and pad to the Boolean domains. - let matrices = integer_matrices.map_coefficients(|c| bigint_to_fq(&c)); + let projection = >::prepare(); + let matrices = integer_matrices.map_coefficients(|c| projection.project(&c)); let products = build_product_mles(&exact_products, matrices.a.row_count()).unwrap(); let assignment = - build_assignment_mle::(&assignment_bits, matrices.a.column_count()).unwrap(); + build_assignment_mle::(&assignment_bits, matrices.a.column_count()).unwrap(); let matrices = PreparedConstraintMatrices::new(matrices).unwrap(); eprintln!("SHA-256 {blocks} blocks: {}", matrices.short_debug_info()); @@ -97,11 +102,7 @@ fn build(blocks: usize) -> R1csInstanceWitness { fn spartan_prove( r1cs: &R1csInstanceWitness, -) -> ( - Proof, - SpartanPiopProof, - ScaledMleEvaluationClaim, -) { +) -> (Proof, SpartanPiopProof, ScaledMleEvaluationClaim) { let mut prover = build_prover(SESSION, INSTANCE); let (piop, claim) = prove_spartan_piop(&mut prover, &r1cs.instance, &r1cs.products, &r1cs.witness).unwrap(); @@ -110,10 +111,10 @@ fn spartan_prove( } fn spartan_verify( - r1cs_instance: &PreparedConstraintMatrices, + r1cs_instance: &PreparedConstraintMatrices, proof: &Proof, - piop: &SpartanPiopProof, -) -> ScaledMleEvaluationClaim { + piop: &SpartanPiopProof, +) -> ScaledMleEvaluationClaim { let mut verifier = build_verifier(SESSION, INSTANCE, proof); let claim = verify_spartan_proof(&mut verifier, r1cs_instance, piop).unwrap(); verifier.check_eof().unwrap(); diff --git a/tooling/cli/src/benchmark.rs b/tooling/cli/src/benchmark.rs index cbd8f6a9..85115321 100644 --- a/tooling/cli/src/benchmark.rs +++ b/tooling/cli/src/benchmark.rs @@ -1,6 +1,10 @@ //! Shared execution and timing for circuit proof benchmarks. +use crate::ProjectConstraint; use crate::end_to_end::{CircuitProofSystem, CircuitStatement, CircuitStats, Error}; +use circuit::matrix_products::ModularVector; +use circuit::{BitWidth, IntoWords}; +use common::{BitzClaimField, BitzConstraintRing}; use std::{ fmt, time::{Duration, Instant}, @@ -18,9 +22,17 @@ pub struct Timings { /// Runs setup, witness generation, commitment, proving, and verification. /// Callers generate inputs and initialize worker threads before calling this. -pub fn run(statement: S, inputs: &[bool]) -> Result { +pub fn run(statement: S, inputs: &[bool]) -> Result +where + S: CircuitStatement, + F: BitzClaimField, + F::Integer: BitWidth + IntoWords, + Vec: for<'a> From<&'a ModularVector<2>>, + R: BitzConstraintRing, + Proj: ProjectConstraint, +{ let started = Instant::now(); - let prepared = CircuitProofSystem::new(statement)?; + let prepared = CircuitProofSystem::::new::(statement)?; let setup = started.elapsed(); let circuit = prepared.stats(); tracing::info!( diff --git a/tooling/cli/src/cmd/circuit_e2e.rs b/tooling/cli/src/cmd/circuit_e2e.rs index c56ad04b..31ff5a57 100644 --- a/tooling/cli/src/cmd/circuit_e2e.rs +++ b/tooling/cli/src/cmd/circuit_e2e.rs @@ -8,6 +8,9 @@ use { }, }; +type F = field::FqDefault; +type Proj = bitz_cli::ProjectBigIntToFq; + /// Prove generated circuit constraints over Q100 and verify the proof. #[derive(FromArgs, PartialEq, Debug)] #[argh(subcommand, name = "circuit-e2e")] @@ -53,7 +56,7 @@ impl Command for Args { CircuitInstance::random(self.circuit, self.num_blocks, self.initial_state) })?; let inputs = statement.inputs.clone(); - let timings = benchmark::run(statement, &inputs)?; + let timings = benchmark::run::<_, F, _, Proj>(statement, &inputs)?; tracing::info!("Proof verified successfully"); println!("circuit={} threads={} field=Q100 pcs=Fast hash=Blake3 relation=Q100-r1cs constraints_verified=true", self.circuit, rayon::current_num_threads()); println!("{timings}"); diff --git a/tooling/cli/src/end_to_end.rs b/tooling/cli/src/end_to_end.rs index 9e4c9d96..6ddf1f9c 100644 --- a/tooling/cli/src/end_to_end.rs +++ b/tooling/cli/src/end_to_end.rs @@ -1,16 +1,18 @@ -//! Circuit constraints over Q100, reduced by Spartan and opened through direct or virtual BitZ. +//! Circuit constraints over a prime field, reduced by Spartan and opened through direct or virtual BitZ. use circuit::{ - Circuit, + BitWidth, Circuit, IntoWords, constraints::{ConstraintGenerator, SparseBoolMatrix}, + matrix_products::ModularVector, matrix_transpose::{MTransposeGenerator, MaterializedMTranspose}, witgen::{PackedWitness, ProductWitgen}, }; use common::{ - BitZParams, LinearClaim, OpeningQuery, Root, Shape, VirtualMap, VirtualStatement, + BitZParams, BitzClaimField, BitzConstraintRing, LinearClaim, OpeningQuery, Root, Shape, + VirtualMap, VirtualStatement, shape::{MIN_LOG_BITS, PACK_BITS}, }; -use field::{F128, FqDefault, Q100, gf128::smallest_generator}; +use field::{F128, gf128::smallest_generator}; use num_traits::{ConstOne, ConstZero}; use pcs::{CommitScheme, HashKind, LigeritoProfile, Pcs, ProverData, StatementBinding}; use poly::{DenseMultilinearExtension, ScaledMleEvaluationClaim}; @@ -18,9 +20,10 @@ use prover::{BitZProver, VirtualWitness}; use transcript::{ProverState, PublicTranscript, build_prover, build_verifier}; use verifier::BitZVerifier; +use crate::ProjectConstraint; use spartan::{ - PreparedConstraintMatrices, R1csProductMles, SpartanPiopProof, bigint_to_fq, - build_assignment_mle, build_product_mles, prove_spartan_piop, verify_spartan_proof, + PreparedConstraintMatrices, R1csProductMles, SpartanPiopProof, build_assignment_mle, + build_product_mles, prove_spartan_piop, verify_spartan_proof, }; const SESSION: &[u8] = b"bitz/circuit-e2e/v1"; @@ -28,7 +31,8 @@ const WINDOW: u32 = 8; /// A trusted, deterministic circuit and its public inputs. Implementations must /// emit identical operations for symbolic and concrete backends and constrain -/// every public input/output. The proved constraints are interpreted modulo Q100. +/// every public input/output. The proved constraints are interpreted modulo +/// the modulus of the field the proof system is built over. pub trait CircuitStatement { fn domain(&self) -> &'static [u8]; fn public_bytes(&self) -> Vec; @@ -79,23 +83,24 @@ pub struct CircuitStats { pub padded_committed_bits: usize, } +/// A circuit prepared for proving over the prime field `F`. #[derive(Debug)] -pub struct CircuitProofSystem { +pub struct CircuitProofSystem { statement: S, opening_path: OpeningPath, - matrices: PreparedConstraintMatrices, + matrices: PreparedConstraintMatrices, map: MaterializedMTranspose, - params: BitZParams, + params: BitZParams, committed_shape: Shape, pcs: Pcs, } #[derive(Debug)] -pub struct Witness { +pub struct Witness { committed: Vec, assignment_bits: Vec, - assignment: DenseMultilinearExtension, - products: R1csProductMles, + assignment: DenseMultilinearExtension, + products: R1csProductMles, } /// Commitment data and the transcript that sampled its OOD claim. @@ -105,24 +110,35 @@ pub struct CommittedWitness { } #[derive(Clone, Debug)] -pub struct Proof { +pub struct Proof { pub root: Root, - pub spartan: SpartanPiopProof, + pub spartan: SpartanPiopProof, pub opening: transcript::Proof, } -impl CircuitProofSystem { +impl CircuitProofSystem +where + S: CircuitStatement, + F: BitzClaimField, + F::Integer: BitWidth + IntoWords, + Vec: for<'a> From<&'a ModularVector<2>>, // Needed for `build_product_mles` +{ #[tracing::instrument(name = "setup", skip_all)] - pub fn new(statement: S) -> Result { - let mut constraints = ConstraintGenerator::new(statement.input_bits()); + pub fn new(statement: S) -> Result + where + R: BitzConstraintRing, + Proj: ProjectConstraint, + { + let mut constraints = ConstraintGenerator::::new(statement.input_bits()); let inputs: Vec<_> = (0..statement.input_bits()) .map(|i| constraints.input(i)) .collect(); statement.synthesize(&mut constraints, &inputs)?; + let projection = Proj::prepare(); let matrices = PreparedConstraintMatrices::new( constraints .into_matrices() - .map_coefficients(|c| bigint_to_fq(&c)), + .map_coefficients(|c| projection.project(&c)), ) .map_err(Error::Matrix)?; let mut generator = MTransposeGenerator::new(statement.input_bits()); @@ -174,7 +190,7 @@ impl CircuitProofSystem { } #[tracing::instrument(name = "witness", skip_all)] - pub fn witness(&self, inputs: &[bool]) -> Result { + pub fn witness(&self, inputs: &[bool]) -> Result, Error> { if inputs.len() != self.statement.input_bits() { return Err(Error::Input("wrong witness input length")); } @@ -212,7 +228,7 @@ impl CircuitProofSystem { /// Commits and sends the initial OOD evaluation before any PIOP challenge. /// The returned state retains both PCS data and the transcript for proving. #[tracing::instrument(name = "commit", skip_all)] - pub fn commit(&self, witness: &Witness) -> Result { + pub fn commit(&self, witness: &Witness) -> Result { let mut transcript = build_prover(SESSION, self.statement.domain()); let (_, data) = self .pcs @@ -223,7 +239,11 @@ impl CircuitProofSystem { /// Continues the commitment transcript through Spartan and the BitZ opening. #[tracing::instrument(name = "prove", skip_all, fields(opening_path = ?self.opening_path))] - pub fn prove(&self, witness: Witness, commitment: CommittedWitness) -> Result { + pub fn prove( + &self, + witness: Witness, + commitment: CommittedWitness, + ) -> Result, Error> { let CommittedWitness { data, mut transcript, @@ -279,7 +299,7 @@ impl CircuitProofSystem { } #[tracing::instrument(name = "verify", skip_all)] - pub fn verify(&self, proof: &Proof) -> Result<(), Error> { + pub fn verify(&self, proof: &Proof) -> Result<(), Error> { let mut transcript = build_verifier(SESSION, self.statement.domain(), &proof.opening); let commitment = self .pcs @@ -370,10 +390,10 @@ fn pack(witness: &PackedWitness, shape: Shape) -> Vec { packed } -fn opening_claim( - params: &BitZParams, - terminal: &ScaledMleEvaluationClaim, -) -> Result, Error> { +fn opening_claim( + params: &BitZParams, + terminal: &ScaledMleEvaluationClaim, +) -> Result, Error> { let shape = params.shape(); if terminal.point().len() > shape.log_bits() { return Err(Error::Configuration("Spartan point exceeds virtual shape")); @@ -381,7 +401,7 @@ fn opening_claim( // Zero high coordinates select the original assignment inside its zero padding. // Put the scale in one factor, avoiding division even when the scale is zero. let mut point = terminal.point().to_vec(); - point.resize(shape.log_bits(), FqDefault::ZERO); + point.resize(shape.log_bits(), F::zero()); let rows = poly::eq_table(&point[..shape.log_rows()]) .into_iter() .map(|weight| terminal.scale() * weight) @@ -394,6 +414,10 @@ fn opening_claim( #[cfg(test)] mod tests { use super::*; + use crate::ProjectBigIntToFq; + + type F = field::FqDefault; + type Proj = ProjectBigIntToFq; #[test] fn only_exact_identity_maps_select_direct_opening() { @@ -434,7 +458,7 @@ mod tests { #[test] fn commitment_sends_ood_before_proving() { - let system = CircuitProofSystem::new(IdentityBit).unwrap(); + let system = CircuitProofSystem::<_, F>::new::<_, Proj>(IdentityBit).unwrap(); let witness = system.witness(&[true]).unwrap(); let committed = system.commit(&witness).unwrap(); let proof = committed.transcript.finish(); @@ -450,7 +474,7 @@ mod tests { #[test] fn direct_opening_requires_constant_one_on_both_sides() { - let mut system = CircuitProofSystem::new(IdentityBit).unwrap(); + let mut system = CircuitProofSystem::<_, F>::new::<_, Proj>(IdentityBit).unwrap(); let witness = system.witness(&[true]).unwrap(); let data = system.commit(&witness).unwrap(); let mut proof = system.prove(witness, data).unwrap(); @@ -496,37 +520,28 @@ mod tests { #[test] fn scaled_claim_conversion_preserves_values_and_zero_scale() { - let params = BitZParams::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); - let assignment = DenseMultilinearExtension::from_evaluations( - 2, - [0u128, 1, 1, 0].map(FqDefault::from).to_vec(), - ) - .unwrap(); - let point = vec![FqDefault::from(3u128), FqDefault::from(5u128)]; + let params = + BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); + let assignment = + DenseMultilinearExtension::from_evaluations(2, [0u128, 1, 1, 0].map(F::from).to_vec()) + .unwrap(); + let point = vec![F::from(3u128), F::from(5u128)]; let evaluation = assignment.evaluate(&point).unwrap(); - for scale in [FqDefault::ZERO, FqDefault::from(7u128)] { + for scale in [F::ZERO, F::from(7u128)] { let terminal = ScaledMleEvaluationClaim::new( point.clone().into_boxed_slice(), scale, scale * evaluation, ); let claim = opening_claim(¶ms, &terminal).unwrap(); - let value: FqDefault = assignment + let value: F = assignment .iter() .enumerate() .map(|(i, bit)| *bit * claim.row_weights()[i] * claim.column_weights()[0]) .sum(); assert_eq!(value, claim.target()); - assert!( - claim.row_weights()[4..] - .iter() - .all(|w| *w == FqDefault::ZERO) - ); - assert!( - claim.column_weights()[1..] - .iter() - .all(|w| *w == FqDefault::ZERO) - ); + assert!(claim.row_weights()[4..].iter().all(|w| *w == F::ZERO)); + assert!(claim.column_weights()[1..].iter().all(|w| *w == F::ZERO)); } } } diff --git a/tooling/cli/src/lib.rs b/tooling/cli/src/lib.rs index 747e6843..56990eec 100644 --- a/tooling/cli/src/lib.rs +++ b/tooling/cli/src/lib.rs @@ -3,3 +3,41 @@ pub mod benchmark; pub mod circuits; pub mod end_to_end; + +use field::Fq; +use num_bigint::BigInt; +use num_traits::{Signed, ToPrimitive}; + +/// Trait for projecting a constraint onto a field. +/// Needs to be prepared right before the projection, as field might be +/// reconfigured. +pub trait ProjectConstraint: Send + Sync { + fn prepare() -> Self; + + fn project(&self, constraint: &R) -> F; +} + +/// Projects [`BigInt`] constraints onto [`Fq`] by reducing it canonically modulo `Q`. +pub struct ProjectBigIntToFq { + modulus: BigInt, +} + +impl ProjectConstraint> for ProjectBigIntToFq { + fn prepare() -> Self { + Self { + modulus: BigInt::from(Q), + } + } + + fn project(&self, constraint: &BigInt) -> Fq { + let mut reduced = constraint % &self.modulus; + if reduced.is_negative() { + reduced += &self.modulus; + } + Fq::from( + reduced + .to_u128() + .expect("a canonical residue always fits a u128"), + ) + } +} diff --git a/tooling/cli/tests/circuits.rs b/tooling/cli/tests/circuits.rs index bb0f5153..063c5686 100644 --- a/tooling/cli/tests/circuits.rs +++ b/tooling/cli/tests/circuits.rs @@ -1,4 +1,5 @@ use bitz_cli::{ + ProjectBigIntToFq, circuits::{BuiltinCircuit, CircuitInstance}, end_to_end::{CircuitProofSystem, CircuitStatement, Error}, }; @@ -6,14 +7,17 @@ use circuit::{ constraints::ConstraintGenerator, sha256::{ABC_BLOCK, ABC_DIGEST, INITIAL_STATE}, }; -use num_bigint::BigInt; use num_traits::{Signed, ToPrimitive}; +type R = num_bigint::BigInt; +type F = field::FqDefault; +type Proj = ProjectBigIntToFq; + #[test] fn sha_constraint_residuals_cannot_wrap_modulo_q100() { for circuit in BuiltinCircuit::ALL { let statement = CircuitInstance::random(circuit, None, None).unwrap(); - let mut generator = ConstraintGenerator::::new(statement.input_bits()); + let mut generator = ConstraintGenerator::::new(statement.input_bits()); let inputs: Vec<_> = (0..statement.input_bits()) .map(|i| generator.input(i)) .collect(); @@ -51,7 +55,7 @@ fn supported_sha_circuits_prove_and_verify() { for circuit in BuiltinCircuit::ALL { let statement = CircuitInstance::random(circuit, None, None).unwrap(); let inputs = statement.inputs.clone(); - let system = CircuitProofSystem::new(statement).unwrap(); + let system = CircuitProofSystem::<_, F>::new::(statement).unwrap(); let witness = system.witness(&inputs).unwrap(); let data = system.commit(&witness).unwrap(); let proof = system.prove(witness, data).unwrap(); @@ -74,17 +78,17 @@ fn sha_compression_matches_abc_and_binds_public_values() { inputs: inputs.clone(), output: bits(&ABC_DIGEST), }; - let system = CircuitProofSystem::new(statement.clone()).unwrap(); + let system = CircuitProofSystem::<_, F>::new::(statement.clone()).unwrap(); let witness = system.witness(&inputs).unwrap(); let data = system.commit(&witness).unwrap(); let proof = system.prove(witness, data).unwrap(); - CircuitProofSystem::new(statement.clone()) + CircuitProofSystem::<_, F>::new::(statement.clone()) .unwrap() .verify(&proof) .unwrap(); let mut changed = statement.clone(); changed.output[0] ^= true; - let wrong_output = CircuitProofSystem::new(changed).unwrap(); + let wrong_output = CircuitProofSystem::<_, F>::new::(changed).unwrap(); assert!(wrong_output.verify(&proof).is_err()); assert!(matches!( wrong_output.witness(&inputs), @@ -93,7 +97,7 @@ fn sha_compression_matches_abc_and_binds_public_values() { let mut changed = statement; changed.inputs[0] ^= true; assert!( - CircuitProofSystem::new(changed) + CircuitProofSystem::<_, F>::new::(changed) .unwrap() .verify(&proof) .is_err() @@ -110,7 +114,7 @@ fn variable_lengths_custom_state_and_invalid_dimensions() { ] { let statement = CircuitInstance::random(circuit, count, state).unwrap(); let inputs = statement.inputs.clone(); - let system = CircuitProofSystem::new(statement).unwrap(); + let system = CircuitProofSystem::<_, F>::new::(statement).unwrap(); system.witness(&inputs).unwrap(); } for (circuit, count) in [ @@ -127,5 +131,5 @@ fn variable_lengths_custom_state_and_invalid_dimensions() { let mut malformed = CircuitInstance::random(BuiltinCircuit::Sha256Compression, None, None).unwrap(); malformed.inputs.pop(); - assert!(CircuitProofSystem::new(malformed).is_err()); + assert!(CircuitProofSystem::<_, F>::new::(malformed).is_err()); } diff --git a/tooling/cli/tests/end_to_end.rs b/tooling/cli/tests/end_to_end.rs index 857f7b80..f9d387ce 100644 --- a/tooling/cli/tests/end_to_end.rs +++ b/tooling/cli/tests/end_to_end.rs @@ -1,10 +1,15 @@ +use bitz_cli::ProjectBigIntToFq; use bitz_cli::end_to_end::{CircuitProofSystem, CircuitStatement, Error, OpeningPath, Proof}; use circuit::Circuit; use pcs::VerifyError; +type R = num_bigint::BigInt; +type F = field::FqDefault; +type Proj = ProjectBigIntToFq; + fn rejects_changed_or_missing_ood( - system: &CircuitProofSystem, - proof: &Proof, + system: &CircuitProofSystem, + proof: &Proof, ) { // These Fast-profile fixtures have zero initial grinding bits, so the first // 16 transcript bytes encode the OOD evaluation. @@ -46,13 +51,13 @@ impl CircuitStatement for PublicBit { #[test] fn generic_driver_accepts_a_non_sha_circuit() { - let prepared = CircuitProofSystem::new(PublicBit).unwrap(); + let prepared = CircuitProofSystem::<_, F>::new::(PublicBit).unwrap(); assert_eq!(prepared.stats().opening_path, OpeningPath::Direct); assert_eq!(prepared.stats().committed_bits, 2); let witness = prepared.witness(&[true]).unwrap(); let data = prepared.commit(&witness).unwrap(); let proof = prepared.prove(witness, data).unwrap(); - CircuitProofSystem::new(PublicBit) + CircuitProofSystem::<_, F>::new::(PublicBit) .unwrap() .verify(&proof) .unwrap(); @@ -62,7 +67,7 @@ fn generic_driver_accepts_a_non_sha_circuit() { changed.root.0[0] ^= 1; assert!(prepared.verify(&changed).is_err()); let mut changed = proof.clone(); - changed.spartan.inner.round_polynomials[0][0] += field::FqDefault::from(1u128); + changed.spartan.inner.round_polynomials[0][0] += F::from(1u128); assert!(prepared.verify(&changed).is_err()); for hints in [false, true] { let mut changed = proof.clone(); @@ -83,12 +88,12 @@ fn generic_driver_accepts_a_non_sha_circuit() { #[test] fn benchmark_runs_a_generic_circuit_and_propagates_failure() { - let timings = bitz_cli::benchmark::run(PublicBit, &[true]).unwrap(); + let timings = bitz_cli::benchmark::run::<_, F, R, Proj>(PublicBit, &[true]).unwrap(); let output = timings.to_string(); assert!(output.contains("total_prove_ms=")); assert!(output.contains("verify_ms=")); assert!(matches!( - bitz_cli::benchmark::run(PublicBit, &[false]), + bitz_cli::benchmark::run::<_, F, R, Proj>(PublicBit, &[false]), Err(Error::Unsatisfied) )); } @@ -118,14 +123,14 @@ impl CircuitStatement for PublicXor { #[test] fn nonidentity_map_uses_virtual_opening_and_checks_xor_relation() { - let system = CircuitProofSystem::new(PublicXor).unwrap(); + let system = CircuitProofSystem::<_, F>::new::(PublicXor).unwrap(); assert_eq!(system.stats().opening_path, OpeningPath::Virtual); assert_eq!(system.stats().assignment_bits, 3); assert_eq!(system.stats().committed_bits, 2); let witness = system.witness(&[true, false]).unwrap(); let data = system.commit(&witness).unwrap(); let proof = system.prove(witness, data).unwrap(); - CircuitProofSystem::new(PublicXor) + CircuitProofSystem::<_, F>::new::(PublicXor) .unwrap() .verify(&proof) .unwrap(); @@ -140,6 +145,6 @@ fn nonidentity_map_uses_virtual_opening_and_checks_xor_relation() { let mut changed = proof; changed.opening.narg_string.push(0); assert!(system.verify(&changed).is_err()); - let timings = bitz_cli::benchmark::run(PublicXor, &[true, false]).unwrap(); + let timings = bitz_cli::benchmark::run::<_, F, R, Proj>(PublicXor, &[true, false]).unwrap(); assert_eq!(timings.circuit.opening_path, OpeningPath::Virtual); } diff --git a/crates/spartan/tests/sha256_piop.rs b/tooling/cli/tests/sha256_spartan.rs similarity index 88% rename from crates/spartan/tests/sha256_piop.rs rename to tooling/cli/tests/sha256_spartan.rs index 172e40a1..82b29d17 100644 --- a/crates/spartan/tests/sha256_piop.rs +++ b/tooling/cli/tests/sha256_spartan.rs @@ -1,5 +1,4 @@ -use std::array; - +use bitz_cli::{ProjectBigIntToFq, ProjectConstraint}; use circuit::{ constraints::ConstraintGenerator, sha256::{ @@ -8,13 +7,17 @@ use circuit::{ }, witgen::{PackedWitness, ProductWitgen}, }; -use num_bigint::BigInt; use spartan::{ - PreparedConstraintMatrices, bigint_to_fq, build_assignment_mle, build_product_mles, - prove_spartan_piop, verify_spartan_with_mle_claim, + PreparedConstraintMatrices, build_assignment_mle, build_product_mles, prove_spartan_piop, + verify_spartan_with_mle_claim, }; +use std::array; use transcript::{build_prover, build_verifier}; +type R = num_bigint::BigInt; +type F = field::FqDefault; +type Proj = ProjectBigIntToFq; + const SESSION: &[u8] = b"spartan/piop/sha256-compression/v1"; const INSTANCE: &[u8] = b"abc-single-compression"; @@ -51,7 +54,8 @@ fn sha256_compression_verifies_through_spartan_piop() { let recomputed_assignment = integer_matrices.integer_witness(&boolean_witness).unwrap(); assert_assignment_matches(&recomputed_assignment, &recorded_assignment); - let matrices = integer_matrices.map_coefficients(|coefficient| bigint_to_fq(&coefficient)); + let projection = >::prepare(); + let matrices = integer_matrices.map_coefficients(|c| projection.project(&c)); let products = build_product_mles(&exact_products, matrices.a.row_count()).unwrap(); let assignment = build_assignment_mle(&recorded_assignment, matrices.a.column_count()).unwrap(); let matrices = PreparedConstraintMatrices::new(matrices).unwrap(); @@ -103,12 +107,12 @@ fn words_from_le_bits(bits: &[bool; 256]) -> [u32; 8] { }) } -fn assert_assignment_matches(computed: &[BigInt], recorded: &PackedWitness) { +fn assert_assignment_matches(computed: &[R], recorded: &PackedWitness) { assert_eq!(computed.len(), recorded.bit_len()); for (index, expected) in computed.iter().enumerate() { assert_eq!( expected, - &BigInt::from(recorded.bit(index)), + &R::from(recorded.bit(index)), "assignment differs at index {index}", ); }