diff --git a/Cargo.lock b/Cargo.lock index 995b774..02379fe 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -676,6 +676,7 @@ dependencies = [ "divan", "flock-core", "num-traits", + "pastey", "rand 0.10.2", "rand_core 0.10.1", "rand_pcg 0.10.2", @@ -1821,6 +1822,7 @@ version = "0.1.0" dependencies = [ "field", "spongefish", + "thiserror", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 73469b3..831ed0e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,6 +36,7 @@ divan = "0.1.21" flock-core = { git = "https://github.com/succinctlabs/flock.git", rev = "879072249e52b8b9054bf0c6a034cec20f8f6fc7" } num-bigint = "0.4" num-traits = "0.2" +pastey = "0.2.3" proptest = "1.11.0" rand = "0.10" rand_chacha = "0.10" diff --git a/crates/circuit/src/constraints.rs b/crates/circuit/src/constraints.rs index c3463f6..bea93a0 100644 --- a/crates/circuit/src/constraints.rs +++ b/crates/circuit/src/constraints.rs @@ -29,6 +29,11 @@ impl SparseRow { pub fn entries(&self) -> &[(usize, C)] { &self.entries } + + /// The entries, consuming the row. + pub fn into_entries(self) -> Vec<(usize, C)> { + self.entries + } } /// A row-major sparse matrix. @@ -84,6 +89,11 @@ impl SparseMatrix { &self.rows } + /// The rows, consuming the matrix. + pub fn into_rows(self) -> Vec> { + self.rows + } + /// Number of rows. pub fn row_count(&self) -> usize { self.rows.len() @@ -96,7 +106,7 @@ impl SparseMatrix { } impl SparseMatrix { - fn map_values_with(self, map: M) -> SparseMatrix + fn map_values(self, map: M) -> SparseMatrix where D: Send + Sync, M: Fn(C) -> D + Send + Sync, @@ -116,6 +126,28 @@ impl SparseMatrix { columns: self.columns, } } + + /// Maps every coefficient, keeping this matrix. + pub fn map_values_ref(&self, map: M) -> SparseMatrix + where + D: Send + Sync, + M: Fn(&C) -> D + Send + Sync, + { + SparseMatrix { + rows: self + .rows + .par_iter() + .map(|row| SparseRow { + entries: row + .entries + .iter() + .map(|(column, coefficient)| (*column, map(coefficient))) + .collect(), + }) + .collect(), + columns: self.columns, + } + } } /// Why a sparse row representation is malformed. @@ -279,9 +311,9 @@ impl ConstraintMatrices { { ConstraintMatrices { m: self.m, - a: self.a.map_values_with(&map), - b: self.b.map_values_with(&map), - c: self.c.map_values_with(&map), + a: self.a.map_values(&map), + b: self.b.map_values(&map), + c: self.c.map_values(&map), } } } diff --git a/crates/circuit/src/matrix_products.rs b/crates/circuit/src/matrix_products.rs index e536ea0..93ed24c 100644 --- a/crates/circuit/src/matrix_products.rs +++ b/crates/circuit/src/matrix_products.rs @@ -9,6 +9,7 @@ use crate::witgen::Z as Integer; use crate::{BitWidth, IntoWords}; +use common::BitzConstraintRing; use num_traits::{One, Zero}; use rayon::prelude::*; use std::cmp::Ordering; @@ -242,6 +243,7 @@ fn add_mod_words( /// A dense vector of canonical runtime-field elements. #[derive(Clone, Debug, Eq, PartialEq)] pub struct ModularVector { + /// Vector of canonical field element in little-endian limbs form, in row order. values: Vec<[u64; PRIME_LIMBS]>, } @@ -256,7 +258,7 @@ impl ModularVector { self.values.is_empty() } - /// Dense canonical field elements in row order. + /// Slice of canonical field element in little-endian limbs form, in row order. pub fn values(&self) -> &[[u64; PRIME_LIMBS]] { &self.values } @@ -267,12 +269,13 @@ impl ModularVector { } } -impl From<&ModularVector<2>> for Vec { +/// Converts two little-endian limbs, treating them as one `u128`. +impl> From<&ModularVector<2>> for Vec { fn from(values: &ModularVector<2>) -> Self { values .values() .iter() - .map(|&[low, high]| field::FqDefault::from_limbs(low, high)) + .map(|&[low, high]| T::from(u128::from(low) | (u128::from(high) << 64))) .collect() } } @@ -354,6 +357,23 @@ impl From<&num_bigint::BigInt> for StoredInteger { } } +impl From<&StoredInteger> for num_bigint::BigInt { + fn from(value: &StoredInteger) -> Self { + let bytes = value + .words() + .iter() + .flat_map(|word| word.to_le_bytes()) + .collect::>(); + Self::from_signed_bytes_le(&bytes) + } +} + +impl From<&i128> for StoredInteger { + fn from(value: &i128) -> Self { + Self::from(&Integer::<2>::from(*value)) + } +} + impl From<&Integer> for StoredInteger { /// Stores a gadget-local fixed integer without changing its value. fn from(value: &Integer) -> Self { @@ -401,6 +421,22 @@ impl IntegerProducts { self.c_mw.push(StoredInteger::from(&c)); } + /// Whether `A(Mw) * B(Mw) = C(Mw)` holds row by row over the integers + /// (hence modulo every prime). + pub fn is_satisfied(&self) -> bool + where + R: BitzConstraintRing + for<'a> From<&'a StoredInteger>, + { + self.a_mw.len() == self.b_mw.len() + && self.a_mw.len() == self.c_mw.len() + && self + .a_mw + .par_iter() + .zip(&self.b_mw) + .zip(&self.c_mw) + .all(|((a, b), c)| R::from(a) * R::from(b) == R::from(c)) + } + /// Reduces every materialized element modulo `modulus`. /// /// Vectors with at least 32,768 entries use Rayon. Smaller vectors remain @@ -446,12 +482,7 @@ mod tests { type R = num_bigint::BigInt; fn stored_ring(value: &StoredInteger) -> R { - let bytes = value - .words() - .iter() - .flat_map(|word| word.to_le_bytes()) - .collect::>(); - R::from_signed_bytes_le(&bytes) + R::from(value) } fn direct_row(row: &SparseRow, integer_witness: &PackedWitness) -> R { @@ -540,6 +571,51 @@ mod tests { } } + #[test] + fn products_are_satisfied_over_the_integers_only_when_exact() { + let mut products = IntegerProducts::default(); + assert!(products.is_satisfied::()); + let z = |value: i128| Integer::<2>::from(value); + products.push(z(-3), z(5), z(-15)); + products.push( + z(i128::from(i64::MIN)), + z(i128::from(i64::MIN)), + z(1 << 126), + ); + assert!(products.is_satisfied::()); + products.push(z(2), z(2), z(5)); + assert!(!products.is_satisfied::()); + products.c_mw.pop(); + assert!( + !products.is_satisfied::(), + "a missing row is not satisfied" + ); + } + + #[test] + fn every_source_stores_an_integer_alike() { + for value in [ + 0_i128, + 1, + -1, + 7, + -7, + i128::from(i64::MAX), + i128::from(i64::MIN), + 1 << 100, + -(1 << 100), + ] { + let from_bigint = StoredInteger::from(&R::from(value)); + assert_eq!(StoredInteger::from(&value), from_bigint, "{value}"); + assert_eq!( + StoredInteger::from(&Integer::<2>::from(value)), + from_bigint, + "{value}" + ); + assert_eq!(stored_ring(&from_bigint), R::from(value)); + } + } + #[test] fn arbitrary_bigints_round_trip_through_stored_integers() { let boundary = R::one() << 128_usize; diff --git a/crates/circuit/src/witgen.rs b/crates/circuit/src/witgen.rs index 2f04d25..759cdfd 100644 --- a/crates/circuit/src/witgen.rs +++ b/crates/circuit/src/witgen.rs @@ -612,8 +612,8 @@ impl ProductWitgen { /// Consumes the runner into `w`, `M * w`, and exact matrix products. pub fn into_parts(self) -> (PackedWitness, PackedWitness, IntegerProducts) { - let (witness, integer_witness) = self.witgen.into_witnesses(); - (witness, integer_witness, self.products) + let (bool_witness, integer_witness) = self.witgen.into_witnesses(); + (bool_witness, integer_witness, self.products) } } diff --git a/crates/common/src/lib.rs b/crates/common/src/lib.rs index 7486d7d..726f513 100644 --- a/crates/common/src/lib.rs +++ b/crates/common/src/lib.rs @@ -17,7 +17,9 @@ pub use fold::{ Fold, FoldError, column_images, fold_column, fold_columns, reconstruct, row_images, }; pub use opening::OpeningQuery; -pub use params::{BitZParams, ParamsError, VirtualParams, VirtualParamsError}; +pub use params::{ + BitZParams, MIN_PRIME_BITS, ParamsError, VirtualParams, VirtualParamsError, prime_bits, +}; pub use shape::{Shape, ShapeError}; pub use table::{BitTable, TableError, TransposeError, TransposedBitTable}; pub use virtual_map::{ @@ -89,6 +91,7 @@ define_blanket_trait! { #[cfg(test)] mod tests { use super::*; + use field::dynamic::DynField; use field::{F128, FqDefault}; #[test] @@ -105,8 +108,10 @@ mod tests { fn assert_impl_field() {} assert_impl_field::(); assert_impl_field::(); + assert_impl_field::(); fn assert_impl_claim_field() {} assert_impl_claim_field::(); + assert_impl_claim_field::(); } } diff --git a/crates/common/src/params.rs b/crates/common/src/params.rs index 4eca260..8a2bb6f 100644 --- a/crates/common/src/params.rs +++ b/crates/common/src/params.rs @@ -2,7 +2,7 @@ use crate::{BitTable, BitzClaimField, Shape, TableError, VirtualMap}; use field::{ - F128, + F128, FqDefault, MAX_MODULUS_BITS, gf128::{MULT_ORDER, is_generator}, }; use num_traits::{CheckedMul, ToBytes}; @@ -18,6 +18,35 @@ pub enum ParamsError { /// The generator's order is not the full group, so a fold is not the only /// exponent producing its image. GeneratorOrderNotFull, + /// The fold gate leaves the fingerprint prime narrower than + /// [`MIN_PRIME_BITS`]: the shape is too tall. + PrimeTooNarrow, +} + +/// The narrowest fingerprint prime a proof is run over: the width of the +/// default modulus. The PIOP's soundness error is of order `1/q`, so no shape +/// may push the prime below what the fixed modulus gave. +pub const MIN_PRIME_BITS: u32 = FqDefault::META.bits; + +/// Width of the fingerprint prime for `shape`: the largest `bits` such that +/// every prime in `[2^(bits-1), 2^bits)` passes the fold gate of +/// [`BitZParams::new`], capped by the field's [`MAX_MODULUS_BITS`]. +/// +/// The prime is drawn from this interval by both sides, so the width must +/// be fixed by the shape. +/// +/// Wider is sounder, so the interval is the widest the gate admits. The gate +/// is the paper's `q < (|K| - 1) / k_1`, and with `k_1 = 2^log_rows` that is +/// exactly `128 - log_rows` bits. +pub fn prime_bits(shape: &Shape) -> Result { + // `rows * q < ord(g)` admits `q <= (ord(g) - 1) / rows`. + let max_modulus = (MULT_ORDER - 1) / shape.rows() as u128; + // `2^bits - 1 <= max_modulus < 2^(bits + 1) - 1`. + let bits = (max_modulus + 1).ilog2().min(MAX_MODULUS_BITS); + if bits < MIN_PRIME_BITS { + return Err(ParamsError::PrimeTooNarrow); + } + Ok(bits) } /// The shape, the modulus and the generator: what a proof is fixed against. @@ -41,17 +70,19 @@ impl BitZParams { // `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. + // The bound is the paper's `Q < (|K| - 1) / k_1`, under which Round 1 + // of 4.3. "The BitZ IOP for the core LinBitsRings relation" is + // skipped: `k_1 Q < ord(g)`. With `k_1` a power of two it admits every + // `(128 - log_rows)`-bit prime, so a prime of the width [`prime_bits`] + // derives always passes. // // `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) + let rows = F::Integer::from(shape.rows() as u64); + if !rows + .checked_mul(&F::modulus()) + .is_some_and(|reach| reach < f128_order) { return Err(ParamsError::FoldBoundExceeded); } @@ -300,12 +331,11 @@ mod tests { #[test] fn rejects_a_shape_the_modulus_is_too_large_for() { - // `t = 13` is the widest row count this prime admits; 14 is not. The - // `+1` is what separates the two: `k_1 (Q - 1)` alone would still fit - // at `t = 14`. - assert!(params_at(Shape::new(13, 22).unwrap()).is_ok()); + // `t = 14` is the widest row count this prime admits: `2^14 Q114` is + // `2^128 - 11 * 2^14`, under `ord(g)`; `2^15 Q114` is not. + assert!(params_at(Shape::new(14, 21).unwrap()).is_ok()); assert_eq!( - params_at(Shape::new(14, 21).unwrap()).err(), + params_at(Shape::new(15, 20).unwrap()).err(), Some(ParamsError::FoldBoundExceeded) ); } @@ -318,6 +348,59 @@ mod tests { ); } + /// The gate of [`BitZParams::new`] at the top of a `bits`-bit interval. + fn gate_admits(bits: u32, shape: &Shape) -> bool { + let top = (1u128 << bits) - 1; + top.checked_mul(shape.rows() as u128) + .is_some_and(|reach| reach < MULT_ORDER) + } + + /// One shape per admissible row count. + fn shapes_by_rows() -> impl Iterator { + use crate::shape::{MAX_LOG_BITS, MIN_LOG_BITS, PACK_BITS}; + (PACK_BITS as usize..=MAX_LOG_BITS) + .map(|log_rows| Shape::new(log_rows, MIN_LOG_BITS.saturating_sub(log_rows)).unwrap()) + } + + #[test] + fn the_prime_width_is_the_widest_interval_the_gate_admits() { + for shape in shapes_by_rows() { + match prime_bits(&shape) { + Ok(bits) => { + assert!(bits >= MIN_PRIME_BITS, "{shape:?}"); + assert!(gate_admits(bits, &shape), "{shape:?}"); + assert!( + bits == MAX_MODULUS_BITS || !gate_admits(bits + 1, &shape), + "{shape:?}" + ); + } + Err(error) => { + assert_eq!(error, ParamsError::PrimeTooNarrow, "{shape:?}"); + assert!(!gate_admits(MIN_PRIME_BITS, &shape), "{shape:?}"); + } + } + } + } + + #[test] + fn the_prime_width_at_the_row_counts_in_use() { + for (log_rows, bits) in [(7, 121), (14, 114), (21, 107), (28, 100)] { + let shape = Shape::new(log_rows, 22usize.saturating_sub(log_rows)).unwrap(); + assert_eq!(prime_bits(&shape), Ok(bits), "{log_rows}"); + } + let shape = Shape::new(29, 0).unwrap(); + assert_eq!(prime_bits(&shape), Err(ParamsError::PrimeTooNarrow)); + } + + /// `Q114` is admitted by exactly the shapes whose width reaches it. + #[test] + fn the_prime_width_agrees_with_the_gate() { + for shape in shapes_by_rows() { + let reaches = prime_bits(&shape).is_ok_and(|bits| bits >= F::META.bits); + assert_eq!(params_at(shape).is_ok(), reaches, "{shape:?}"); + } + } + #[test] fn the_fold_bound_is_the_widest_column_sum() { let params = params_at(shape()).unwrap(); diff --git a/crates/field/Cargo.toml b/crates/field/Cargo.toml index 13807f9..d053101 100644 --- a/crates/field/Cargo.toml +++ b/crates/field/Cargo.toml @@ -12,6 +12,7 @@ crypto-primitives = { workspace = true } crypto-primitives-proc-macros = { workspace = true } rand = { workspace = true, optional = true } spongefish = { workspace = true, optional = true } +pastey = { workspace = true } [features] rand = ["dep:rand"] diff --git a/crates/field/src/codec.rs b/crates/field/src/codec.rs index ab30161..816762e 100644 --- a/crates/field/src/codec.rs +++ b/crates/field/src/codec.rs @@ -9,14 +9,16 @@ //! rejects anything at or above `Q`, so each element has exactly one wire //! form. No `Decoding`: `Fq` is never sampled here, and sampling it would //! need a wider squeeze to kill the modular bias. +//! - `DynField`: as `Fq`, against the installed modulus. //! //! `NargSerialize` comes from spongefish's blanket impl over `Encoding`. -use crypto_primitives::LiftElement; +use crypto_primitives::{BaseField, LiftElement}; use spongefish::{ ByteArray, Decoding, Encoding, NargDeserialize, VerificationError, VerificationResult, }; +use crate::dynamic::DynField; use crate::{F128, Fq}; impl Encoding<[u8]> for F128 { @@ -45,17 +47,34 @@ impl Encoding<[u8]> for Fq { } } +/// Deserialized `u128` from 16 little-endian bytes, checking that it's below given modulus. +fn deserialized_reduced(buf: &mut &[u8], modulus: u128) -> VerificationResult { + // Stage the cursor: the contract requires `buf` untouched on failure, + // and the range check can still fail after the read succeeds. + let mut rest = *buf; + let value = u128::from_le_bytes(<[u8; 16]>::deserialize_from_narg(&mut rest)?); + if value >= modulus { + return Err(VerificationError); + } + *buf = rest; + Ok(value) +} + impl NargDeserialize for Fq { fn deserialize_from_narg(buf: &mut &[u8]) -> VerificationResult { - // Stage the cursor: the contract requires `buf` untouched on failure, - // and the range check can still fail after the read succeeds. - let mut rest = *buf; - let value = u128::from_le_bytes(<[u8; 16]>::deserialize_from_narg(&mut rest)?); - if value >= Q { - return Err(VerificationError); - } - *buf = rest; - Ok(Self::from(value)) + deserialized_reduced(buf, Q).map(Self::from) + } +} + +impl Encoding<[u8]> for DynField { + fn encode(&self) -> impl AsRef<[u8]> { + self.lift().to_le_bytes() + } +} + +impl NargDeserialize for DynField { + fn deserialize_from_narg(buf: &mut &[u8]) -> VerificationResult { + deserialized_reduced(buf, DynField::modulus()).map(Self::from) } } @@ -65,6 +84,7 @@ mod tests { use spongefish::NargSerialize; use super::*; + use crate::dynamic::test_support::with_modulus; use crate::{FqDefault, Q100}; fn f128_cases() -> [F128; 4] { @@ -142,4 +162,39 @@ mod tests { assert_eq!(buf, narg, "cursor must not move on failure"); } } + + #[test] + fn dyn_field_wire_form_is_fq_wire_form() { + with_modulus(Q100, || { + for v in [0u128, 1, 12345, Q100 - 1] { + let a = DynField::from(v); + assert_eq!(a.encode().as_ref(), FqDefault::from(v).encode().as_ref()); + assert_eq!(a.encode().as_ref(), v.to_le_bytes()); + + let mut narg = Vec::new(); + a.serialize_into_narg(&mut narg); + narg.extend_from_slice(b"tail"); + + let mut buf = narg.as_slice(); + assert_eq!(DynField::deserialize_from_narg(&mut buf).unwrap(), a); + assert_eq!(buf, b"tail"); + } + }); + } + + #[test] + fn dyn_field_narg_rejects_non_canonical() { + with_modulus(Q100, || { + for v in [Q100, Q100 + 1, u128::MAX] { + let narg = v.to_le_bytes(); + let mut buf = narg.as_slice(); + assert!(DynField::deserialize_from_narg(&mut buf).is_err()); + assert_eq!(buf, narg, "cursor must not move on failure"); + } + let narg = [0u8; 15]; + let mut buf = narg.as_slice(); + assert!(DynField::deserialize_from_narg(&mut buf).is_err()); + assert_eq!(buf, narg); + }); + } } diff --git a/crates/field/src/dynamic.rs b/crates/field/src/dynamic.rs new file mode 100644 index 0000000..18d25c9 --- /dev/null +++ b/crates/field/src/dynamic.rs @@ -0,0 +1,602 @@ +//! Field with modulus not known at compile-time. +//! +//! We only expect to have one of those at any given time, so the modulus is shared globally. + +use crate::{FieldWithDynamicModulus, helpers}; +use crypto_primitives::{BaseField, LiftElement, WithAssociatedInteger}; +use crypto_primitives_proc_macros::InfallibleCheckedOp; +use num_traits::{ + Bounded, CheckedAdd, CheckedDiv, CheckedMul, CheckedNeg, CheckedSub, ConstOne, ConstZero, Inv, + One, Pow, Zero, +}; +use pastey::paste; +use std::fmt::Display; +use std::iter::{Product, Sum}; +use std::ops::{Add, AddAssign, Div, DivAssign, Mul, MulAssign, Neg, Sub, SubAssign}; + +/// The installed [`FieldMetadata`], one atomic per word: the config shared by +/// every [`DynField`]. +/// +/// Written only by [`DynField::set_modulus`], read by every operation. Its +/// contract rules out a writer concurrent with a reader, so the words need +/// neither a lock nor an ordering among them: `Relaxed` loads are plain +/// loads, and every thread reads the same cache line without writing it. +mod global { + use crate::helpers::FieldMetadata; + use std::sync::atomic::{AtomicU32, AtomicU64, Ordering::Relaxed}; + + static MODULUS_LO: AtomicU64 = AtomicU64::new(0); + static MODULUS_HI: AtomicU64 = AtomicU64::new(0); + static MU_LO: AtomicU64 = AtomicU64::new(0); + static MU_HI: AtomicU64 = AtomicU64::new(0); + static BITS: AtomicU32 = AtomicU32::new(0); + + #[inline(always)] + fn load_u128(lo: &AtomicU64, hi: &AtomicU64) -> u128 { + u128::from(lo.load(Relaxed)) | (u128::from(hi.load(Relaxed)) << 64) + } + + fn store_u128(lo: &AtomicU64, hi: &AtomicU64, value: u128) { + lo.store(value as u64, Relaxed); + hi.store((value >> 64) as u64, Relaxed); + } + + #[inline(always)] + pub(super) fn modulus() -> u128 { + load_u128(&MODULUS_LO, &MODULUS_HI) + } + + #[inline(always)] + pub(super) fn load() -> FieldMetadata { + FieldMetadata { + modulus: modulus(), + bits: BITS.load(Relaxed), + mu: load_u128(&MU_LO, &MU_HI), + } + } + + pub(super) fn store(cfg: &FieldMetadata) { + store_u128(&MODULUS_LO, &MODULUS_HI, cfg.modulus); + store_u128(&MU_LO, &MU_HI, cfg.mu); + BITS.store(cfg.bits, Relaxed); + } +} + +/// An element of `Z/qZ` for the installed `q`, held reduced. +/// +/// The modulus is process-wide, so values made under different moduli are +/// the same Rust type; keeping them apart is the caller's job. +#[derive(Debug, Copy, Clone, Default, PartialEq, Eq, Hash, InfallibleCheckedOp)] +#[infallible_checked_unary_op((CheckedNeg, neg))] +#[infallible_checked_binary_op((CheckedAdd, add), (CheckedSub, sub), (CheckedMul, mul))] +#[repr(transparent)] +pub struct DynField { + reduced_value: u128, +} + +impl DynField { + /// The installed modulus and its reduction constants. + #[inline(always)] + pub fn config() -> helpers::FieldMetadata { + let cfg = global::load(); + debug_assert!(cfg.modulus != 0, "Field modulus has not been set yet!"); + cfg + } +} + +impl FieldWithDynamicModulus for DynField { + /// Set modulus globally. + /// + /// The modulus must be an odd prime below `2^126`, the bound of the + /// Barrett reduction shared with [`Fq`](crate::Fq); anything else panics. + unsafe fn set_modulus(modulus: u128) { + assert!(modulus >= 3, "modulus must be at least 3"); + assert_eq!(modulus % 2, 1, "modulus must be odd"); + assert!( + modulus < 1u128 << helpers::MAX_MODULUS_BITS, + "modulus must be below 2^126" + ); + // `DynField` claims to be a prime field, so it enforces that itself + // rather than trusting whoever drew the modulus. + assert!(helpers::is_prime(modulus), "modulus must be prime"); + global::store(&helpers::FieldMetadata::new(modulus)); + } +} + +// +// Core traits +// + +impl Display for DynField { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{} (mod {})", self.reduced_value, DynField::modulus()) + } +} + +// +// Zero and One traits +// + +impl Zero for DynField { + #[inline(always)] + fn zero() -> Self { + Self::ZERO + } + + #[inline(always)] + fn is_zero(&self) -> bool { + self.reduced_value == 0 + } +} + +impl One for DynField { + #[inline(always)] + fn one() -> Self { + Self::ONE + } +} + +impl ConstZero for DynField { + const ZERO: Self = DynField { reduced_value: 0 }; +} + +impl ConstOne for DynField { + const ONE: Self = DynField { reduced_value: 1 }; +} + +// +// Basic arithmetic operations +// + +impl Neg for DynField { + type Output = Self; + + #[inline(always)] + fn neg(self) -> Self::Output { + if self.is_zero() { + self + } else { + DynField { + reduced_value: Self::modulus() - self.reduced_value, + } + } + } +} + +macro_rules! impl_basic_op_forward_to_assign { + ($trait:ident, $method:ident, $assign_method:ident) => { + paste! { + impl $trait for DynField { + type Output = DynField; + + #[inline(always)] + fn $method(mut self, rhs: DynField) -> Self::Output { + [<$trait Assign>]::$assign_method(&mut self, rhs); + self + } + } + + impl $trait<&Self> for DynField { + type Output = DynField; + + #[inline(always)] + fn $method(mut self, rhs: &DynField) -> Self::Output { + [<$trait Assign>]::$assign_method(&mut self, rhs); + self + } + } + + impl $trait for &DynField { + type Output = DynField; + + #[inline(always)] + fn $method(self, rhs: DynField) -> Self::Output { + (*self).$method(rhs) + } + } + + impl $trait for &DynField { + type Output = DynField; + + #[inline(always)] + fn $method(self, rhs: &DynField) -> Self::Output { + (*self).$method(*rhs) + } + } + } + }; +} + +impl_basic_op_forward_to_assign!(Add, add, add_assign); +impl_basic_op_forward_to_assign!(Sub, sub, sub_assign); +impl_basic_op_forward_to_assign!(Mul, mul, mul_assign); +impl_basic_op_forward_to_assign!(Div, div, div_assign); + +// Required by `crypto_primitives::Field`, not implemented yet. + +impl Pow for DynField { + type Output = Self; + + fn pow(self, _exp: u32) -> Self::Output { + unimplemented!("exponentiation is not implemented yet") + } +} + +impl Pow for DynField { + type Output = Self; + + fn pow(self, _exp: u128) -> Self::Output { + unimplemented!("exponentiation is not implemented yet") + } +} + +impl Pow<&u128> for DynField { + type Output = Self; + + fn pow(self, _exp: &u128) -> Self::Output { + unimplemented!("exponentiation is not implemented yet") + } +} + +impl Inv for DynField { + type Output = Option; + + fn inv(self) -> Self::Output { + unimplemented!("inversion is not implemented yet") + } +} + +// +// Checked arithmetic operations +// (Note: Field operations do not overflow) +// + +impl CheckedDiv for DynField { + #[allow(clippy::arithmetic_side_effects)] // False alert + fn checked_div(&self, rhs: &Self) -> Option { + Some(self * Inv::inv(*rhs)?) + } +} + +// +// Arithmetic assign operations +// + +macro_rules! impl_op_assign_boilerplate { + ($trait:ident, $method:ident) => { + impl<'a> $trait<&'a DynField> for DynField { + #[inline(always)] + fn $method(&mut self, rhs: &'a DynField) { + self.$method(*rhs); + } + } + }; +} + +impl_op_assign_boilerplate!(AddAssign, add_assign); +impl_op_assign_boilerplate!(SubAssign, sub_assign); +impl_op_assign_boilerplate!(MulAssign, mul_assign); +impl_op_assign_boilerplate!(DivAssign, div_assign); + +impl AddAssign for DynField { + #[inline(always)] + fn add_assign(&mut self, rhs: Self) { + let modulus = Self::modulus(); + // SAFETY: Both operands are below `modulus < 2^126`, so the sum cannot wrap. + let sum = unsafe { self.reduced_value.unchecked_add(rhs.reduced_value) }; + self.reduced_value = if sum >= modulus { sum - modulus } else { sum }; + } +} + +impl SubAssign for DynField { + #[inline(always)] + fn sub_assign(&mut self, rhs: Self) { + self.reduced_value = if self.reduced_value >= rhs.reduced_value { + self.reduced_value - rhs.reduced_value + } else { + self.reduced_value + Self::modulus() - rhs.reduced_value + }; + } +} + +impl MulAssign for DynField { + #[inline(always)] + fn mul_assign(&mut self, rhs: Self) { + self.reduced_value = Self::config().mul(self.reduced_value, rhs.reduced_value); + } +} + +impl DivAssign for DynField { + fn div_assign(&mut self, _rhs: Self) { + unimplemented!("division is not implemented yet") + } +} + +// +// Aggregate operations +// + +impl Sum for DynField { + fn sum>(iter: I) -> Self { + iter.fold(Self::ZERO, |acc, x| acc + x) + } +} + +impl<'a> Sum<&'a Self> for DynField { + #[allow(clippy::arithmetic_side_effects)] // False alert + fn sum>(iter: I) -> Self { + iter.fold(Self::ZERO, |acc, x| acc + x) + } +} + +impl Product for DynField { + #[allow(clippy::arithmetic_side_effects)] // False alert + fn product>(iter: I) -> Self { + iter.fold(Self::ONE, |acc, x| acc * x) + } +} + +impl<'a> Product<&'a Self> for DynField { + #[allow(clippy::arithmetic_side_effects)] // False alert + fn product>(iter: I) -> Self { + iter.fold(Self::ONE, |acc, x| acc * x) + } +} + +// +// Conversions +// + +impl From for DynField { + fn from(value: bool) -> Self { + if value { Self::ONE } else { Self::ZERO } + } +} + +impl From<&u64> for DynField { + fn from(value: &u64) -> Self { + Self::from(*value) + } +} + +/// Reduces its input, so any `u64` is accepted. +impl From for DynField { + fn from(value: u64) -> Self { + Self::from(u128::from(value)) + } +} + +impl From<&u128> for DynField { + fn from(value: &u128) -> Self { + Self::from(*value) + } +} + +/// Reduces its input, so any `u128` is accepted. +impl From for DynField { + fn from(value: u128) -> Self { + DynField { + reduced_value: value % Self::modulus(), + } + } +} + +// +// crypto-primitives +// + +impl Bounded for DynField { + #[inline(always)] + fn min_value() -> Self { + Self::ZERO + } + + #[inline(always)] + fn max_value() -> Self { + DynField { + reduced_value: DynField::modulus() - 1, + } + } +} + +impl BaseField for DynField { + #[inline(always)] + fn modulus() -> Self::Integer { + let modulus = global::modulus(); + debug_assert!(modulus != 0, "Field modulus has not been set yet!"); + modulus + } + + fn modulus_minus_one_div_two() -> Self::Integer { + (Self::modulus() - 1) / 2 + } +} + +impl WithAssociatedInteger for DynField { + type Integer = u128; +} + +impl LiftElement for DynField { + fn lift(&self) -> u128 { + self.reduced_value + } +} + +// +// Other +// + +#[allow(dead_code)] // Cannot be gated for #[cfg(test)] or it won't be accessible to other crates +pub mod test_support { + use super::*; + use std::sync::{Mutex, PoisonError}; + + /// Runs `test_code` under `modulus`. The modulus is process-wide and the test + /// harness is multi-threaded, so every test that needs one holds this + /// lock for its whole run. + pub fn with_modulus(modulus: u128, test_code: impl FnOnce() -> T) -> T { + static LOCK: Mutex<()> = Mutex::new(()); + let _guard = LOCK.lock().unwrap_or_else(PoisonError::into_inner); + // SAFETY: the lock keeps every other test's values and operations out. + unsafe { DynField::set_modulus(modulus) }; + test_code() + } +} + +#[cfg(test)] +mod tests { + use super::test_support::with_modulus; + use super::*; + use crate::{Fq, FqDefault, Q100}; + use crypto_primitives::{ConstField, WithExtensionDegree}; + use rand_core::{Rng, SeedableRng}; + use rand_pcg::Pcg64; + + /// A prime small enough to check every pair of operands. + const SMALL: u128 = 251; + /// A 114-bit prime. + const WIDE: u128 = (1 << 114) - 11; + + fn u128_of(rng: &mut Pcg64) -> u128 { + (rng.next_u64() as u128) << 64 | rng.next_u64() as u128 + } + + /// Every pair of `extremes`, then `count` random pairs below `q`. + fn operand_pairs<'a>( + extremes: &'a [u128], + q: u128, + count: usize, + rng: &'a mut Pcg64, + ) -> impl Iterator + 'a { + extremes + .iter() + .flat_map(move |&a| extremes.iter().map(move |&b| (a, b))) + .chain((0..count).map(move |_| (u128_of(rng) % q, u128_of(rng) % q))) + } + + #[test] + fn ensure_traits() { + fn assert_impl() {} + assert_impl::(); + } + + #[test] + #[should_panic(expected = "modulus must be odd")] + fn set_modulus_rejects_even() { + // SAFETY: rejected before anything is stored. + unsafe { DynField::set_modulus(Q100 + 1) }; + } + + #[test] + #[should_panic(expected = "modulus must be prime")] + fn set_modulus_rejects_composite() { + // SAFETY: rejected before anything is stored. + unsafe { DynField::set_modulus(((1u128 << 54) - 33) * ((1u128 << 53) - 111)) }; + } + + #[test] + #[should_panic(expected = "modulus must be below 2^126")] + fn set_modulus_rejects_the_barrett_bound() { + // SAFETY: rejected before anything is stored. + unsafe { DynField::set_modulus((1 << 127) - 1) }; + } + + #[test] + fn config_reports_the_installed_modulus() { + with_modulus(Q100, || { + let cfg = DynField::config(); + assert_eq!(cfg.modulus, Q100); + assert_eq!(cfg.bits, FqDefault::META.bits); + assert_eq!(DynField::modulus(), Q100); + assert_eq!(DynField::modulus_minus_one_div_two(), (Q100 - 1) / 2); + assert_eq!(DynField::min_value(), DynField::ZERO); + assert_eq!(DynField::max_value().lift(), Q100 - 1); + assert_eq!(DynField::extension_degree(), 1); + assert_eq!(DynField::ONE.lift(), 1); + assert!(DynField::ZERO.is_zero()); + }); + } + + /// Every ring operation against `Fq`, whose Barrett is checked against + /// long division. + fn agrees_with_fq(rng: &mut Pcg64) { + with_modulus(Q, || { + for (a, b) in operand_pairs(&[0, 1, Q - 1, Q / 2], Q, 512, rng) { + let (x, y) = (DynField::from(a), DynField::from(b)); + let (fx, fy) = (Fq::::from(a), Fq::::from(b)); + assert_eq!((x + y).lift(), (fx + fy).lift(), "{a} + {b}"); + assert_eq!((x - y).lift(), (fx - fy).lift(), "{a} - {b}"); + assert_eq!((x * y).lift(), (fx * fy).lift(), "{a} * {b}"); + assert_eq!((-x).lift(), (-fx).lift(), "-{a}"); + } + }); + } + + #[test] + fn matches_the_compile_time_field() { + let mut rng = Pcg64::seed_from_u64(501); + agrees_with_fq::(&mut rng); + agrees_with_fq::(&mut rng); + agrees_with_fq::(&mut rng); + } + + /// The largest prime below the Barrett bound, where the quotient estimate + /// has the least slack, against the add-and-double multiply, which shares + /// nothing with Barrett. + #[test] + fn multiplies_at_the_widest_modulus() { + let mut q = (1u128 << 126) - 1; + while !helpers::is_prime(q) { + q -= 2; + } + let mut rng = Pcg64::seed_from_u64(502); + with_modulus(q, || { + for (a, b) in operand_pairs(&[0, 1, q - 1, q - 2, q / 2], q, 512, &mut rng) { + let got = (DynField::from(a) * DynField::from(b)).lift(); + assert_eq!(got, helpers::mul_mod(a, b, q), "{a} * {b}"); + } + }); + } + + #[test] + fn exhaustive_over_a_small_modulus() { + with_modulus(SMALL, || { + for a in 0..SMALL { + for b in 0..SMALL { + let (x, y) = (DynField::from(a), DynField::from(b)); + assert_eq!((x * y).lift(), a * b % SMALL, "{a} * {b}"); + assert_eq!((x + y).lift(), (a + b) % SMALL, "{a} + {b}"); + assert_eq!((x - y).lift(), (a + SMALL - b) % SMALL, "{a} - {b}"); + } + } + }); + } + + #[test] + fn from_reduces() { + with_modulus(Q100, || { + assert_eq!(DynField::from(Q100).lift(), 0); + assert_eq!(DynField::from(Q100 + 1).lift(), 1); + assert_eq!(DynField::from(u128::MAX).lift(), u128::MAX % Q100); + assert_eq!(DynField::from(&u128::MAX), DynField::from(u128::MAX)); + assert_eq!(DynField::from(u64::MAX).lift(), u128::from(u64::MAX)); + assert_eq!(DynField::from(&u64::MAX), DynField::from(u64::MAX)); + assert_eq!(DynField::from(true), DynField::ONE); + assert_eq!(DynField::from(false), DynField::ZERO); + }); + } + + #[test] + fn values_follow_the_installed_modulus() { + with_modulus(SMALL, || { + let seven = DynField::from(7u64); + assert_eq!((seven * seven * seven).lift(), 343 % SMALL); + // SAFETY: nothing made under `SMALL` is used past here; the lock + // is held. + unsafe { DynField::set_modulus(Q100) }; + let seven = DynField::from(7u64); + assert_eq!((seven * seven * seven).lift(), 343); + assert_eq!(DynField::from(SMALL).lift(), SMALL); + }); + } +} diff --git a/crates/field/src/fq.rs b/crates/field/src/fq.rs index ca67702..7799fd4 100644 --- a/crates/field/src/fq.rs +++ b/crates/field/src/fq.rs @@ -1,5 +1,6 @@ //! Arithmetic modulo a prime fixed at compile time. +use crate::helpers::{self, FieldMetadata}; use crypto_primitives::{ConstBaseField, LiftElement, WithAssociatedInteger}; use crypto_primitives_proc_macros::InfallibleCheckedOp; use num_traits::{ @@ -35,165 +36,12 @@ pub type FqDefault = Fq; #[infallible_checked_binary_op((CheckedAdd, add), (CheckedSub, sub), (CheckedMul, mul))] pub struct Fq(u128); -/// Full 128x128 -> 256-bit product as `(low, high)`. -const fn mul_wide(a: u128, b: u128) -> (u128, u128) { - let (a_lo, a_hi) = (a as u64 as u128, a >> 64); - let (b_lo, b_hi) = (b as u64 as u128, b >> 64); - - let ll = a_lo * b_lo; - let hh = a_hi * b_hi; - let (mid, mid_carry) = (a_lo * b_hi).overflowing_add(a_hi * b_lo); - - let (lo, lo_carry) = ll.overflowing_add(mid << 64); - let hi = hh + (mid >> 64) + ((mid_carry as u128) << 64) + lo_carry as u128; - (lo, hi) -} - -/// `(lo, hi) >> n` for `n < 128`, keeping the low 128 bits. Every use here -/// shifts far enough that nothing above them survives. -const fn shr_wide(lo: u128, hi: u128, n: u32) -> u128 { - if n == 0 { - lo - } else { - (lo >> n) | (hi << (128 - n)) - } -} - -/// `floor(2^(2k) / q)` by binary long division, `k` being `q`'s bit length. -/// -/// Barrett's precomputed reciprocal. The numerator has a single set bit, so -/// each step shifts the remainder up and subtracts `q` when it fits. -const fn barrett_mu(q: u128, k: u32) -> u128 { - let mut rem = 0u128; - let mut quo = 0u128; - let mut i = 2 * k; - loop { - rem = (rem << 1) | (i == 2 * k) as u128; - let fits = rem >= q; - if fits { - rem -= q; - } - quo = (quo << 1) | fits as u128; - if i == 0 { - return quo; - } - i -= 1; - } -} - -/// The first thirteen primes: the standard deterministic Miller-Rabin base set -/// below `2^81.4`. -const PRIMALITY_BASES: [u128; 13] = [2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41]; - -/// `(a + b) mod q`, for `a, b < q < 2^126`. The sum stays below `2^127`. -const fn add_mod(a: u128, b: u128, q: u128) -> u128 { - let sum = a + b; - if sum >= q { sum - q } else { sum } -} - -/// `(a * b) mod q` by doubling, avoiding the 256-bit product a u128 cannot -/// hold. Barrett is not an option here: `MU` depends on `BITS`, which is the -/// constant this feeds. -const fn mul_mod(a: u128, b: u128, q: u128) -> u128 { - let mut result = 0u128; - let mut addend = a % q; - let mut remaining = b; - while remaining != 0 { - if remaining & 1 == 1 { - result = add_mod(result, addend, q); - } - addend = add_mod(addend, addend, q); - remaining >>= 1; - } - result -} - -/// `(base ^ exponent) mod q`, by square-and-multiply. -const fn pow_mod(base: u128, exponent: u128, q: u128) -> u128 { - let mut result = 1u128 % q; - let mut square = base % q; - let mut remaining = exponent; - while remaining != 0 { - if remaining & 1 == 1 { - result = mul_mod(result, square, q); - } - square = mul_mod(square, square, q); - remaining >>= 1; - } - result -} - -/// Whether `candidate` is prime. -/// -/// `const` because [`Fq`] asserts it on its own modulus, which makes a -/// composite `Q` a build failure rather than a type that quietly is not a -/// field. The loops are written out for the same reason: iterator combinators -/// are not available in a const context. -/// -/// # What this does not decide -/// -/// Miller-Rabin against a fixed base set is **proven** only for -/// `n < 3_317_044_064_679_887_385_961_981`, about `2^81.4`. Moduli above that -/// are not decided by any theorem here, and because the bases are public, -/// someone choosing `Q` could in principle construct a composite that passes -/// all of them. Since `Q` is a compile-time constant, doing so means editing -/// the source rather than forging a proof. -const fn is_prime(candidate: u128) -> bool { - if candidate < 2 { - return false; - } - - let mut index = 0; - while index < PRIMALITY_BASES.len() { - let base = PRIMALITY_BASES[index]; - if candidate == base { - return true; - } - if candidate.is_multiple_of(base) { - return false; - } - index += 1; - } - - // `candidate - 1 = odd * 2^shift`. - let shift = (candidate - 1).trailing_zeros(); - let odd = (candidate - 1) >> shift; - - let mut index = 0; - while index < PRIMALITY_BASES.len() { - let mut witness = pow_mod(PRIMALITY_BASES[index], odd, candidate); - if witness != 1 && witness != candidate - 1 { - let mut round = 1; - loop { - if round >= shift { - return false; - } - witness = mul_mod(witness, witness, candidate); - if witness == candidate - 1 { - break; - } - round += 1; - } - } - index += 1; - } - true -} - impl Fq { - /// Bit length of the modulus, derived so it cannot disagree with `Q`. - /// - /// Its asserts are the only check on `Q`, and an associated constant is - /// evaluated where it is used, so an operation that does not need the - /// value reads it anyway rather than accept a modulus out of range. - pub const BITS: u32 = { + pub const META: FieldMetadata = { assert!(Q >= 3, "modulus must be at least 3"); - assert!(Q % 2 == 1, "modulus must be odd"); - assert!(Q < 1u128 << 126, "modulus must be below 2^126"); - // `Fq` claims to be a prime field, so it enforces that itself rather - // than trusting whoever names the constant. - assert!(is_prime(Q), "modulus must be prime"); - 128 - Q.leading_zeros() + // `Fq` claims to be a prime field, so it enforces that itself + assert!(helpers::is_prime(Q), "modulus must be prime"); + FieldMetadata::new(Q) }; /// Constructs an element from two little-endian `u64` limbs and reduces it @@ -201,34 +49,6 @@ impl Fq { pub fn from_limbs(low: u64, high: u64) -> Self { Self::from(u128::from(low) | (u128::from(high) << 64)) } - - const MU: u128 = barrett_mu(Q, Self::BITS); - - /// Barrett reduction, Handbook of Applied Cryptography Algorithm 14.42, - /// after Barrett, CRYPTO '86, LNCS 263:311-323: estimate the quotient - /// from the top bits of `x` and the precomputed reciprocal, subtract, - /// then correct. The estimate is never more than two too small, so two - /// conditional subtractions suffice. - /// - /// That algorithm is stated for a radix above 3 and takes the difference - /// modulo `b^(k+1)`, which is wide enough to hold `3Q`. Here the radix is - /// 2, where it is not, so the difference is taken modulo `2^128` instead. - /// Hence the bound on `Q`. - const fn reduce_wide(lo: u128, hi: u128) -> Self { - let k = Self::BITS; - let q1 = shr_wide(lo, hi, k - 1); - let (q2_lo, q2_hi) = mul_wide(q1, Self::MU); - let q3 = shr_wide(q2_lo, q2_hi, k + 1); - - let mut r = lo.wrapping_sub(mul_wide(q3, Q).0); - if r >= Q { - r -= Q; - } - if r >= Q { - r -= Q; - } - Self(r) - } } impl Display for Fq { @@ -241,7 +61,7 @@ impl Display for Fq { impl Distribution> for StandardUniform { fn sample(&self, rng: &mut R) -> Fq { // Force validation of the const-generic modulus before using it below. - let _ = Fq::::BITS; + let _ = Fq::::META; // A u128 range is not generally an exact multiple of Q. Reject its // incomplete final interval before reducing so every residue has the @@ -285,7 +105,7 @@ impl ConstOne for Fq { /// Reduces its input, so any `u128` is accepted. impl From for Fq { fn from(value: u128) -> Self { - let _ = Self::BITS; + let _ = Self::META; Self(value % Q) } } @@ -312,7 +132,7 @@ impl From for Fq { impl Neg for Fq { type Output = Self; fn neg(self) -> Self { - let _ = Self::BITS; + let _ = Self::META; Self(if self.0 == 0 { 0 } else { Q - self.0 }) } } @@ -320,7 +140,7 @@ impl Neg for Fq { impl Add for Fq { type Output = Self; fn add(self, rhs: Self) -> Self { - let _ = Self::BITS; + let _ = Self::META; // Both operands are below `Q < 2^126`, so the sum cannot wrap. let s = self.0 + rhs.0; Self(if s >= Q { s - Q } else { s }) @@ -330,7 +150,7 @@ impl Add for Fq { impl Sub for Fq { type Output = Self; fn sub(self, rhs: Self) -> Self { - let _ = Self::BITS; + let _ = Self::META; Self(if self.0 >= rhs.0 { self.0 - rhs.0 } else { @@ -341,9 +161,9 @@ impl Sub for Fq { impl Mul for Fq { type Output = Self; - fn mul(self, rhs: Self) -> Self { - let (lo, hi) = mul_wide(self.0, rhs.0); - Self::reduce_wide(lo, hi) + fn mul(mut self, rhs: Self) -> Self { + self.0 = Self::META.mul(self.0, rhs.0); + self } } @@ -507,14 +327,14 @@ impl Bounded for Fq { /// The largest residue, `Q - 1`. fn max_value() -> Self { - let _ = Self::BITS; + let _ = Self::META; Self(Q - 1) } } impl ConstBaseField for Fq { const MODULUS: Self::Integer = { - let _ = Self::BITS; + let _ = Self::META; Q }; const MODULUS_MINUS_ONE_DIV_TWO: Self::Integer = (Q - 1) / 2; @@ -523,6 +343,7 @@ impl ConstBaseField for Fq { #[cfg(test)] mod tests { use super::*; + use crate::helpers::tests::mulmod_reference; use crypto_primitives::{BaseField, WithExtensionDegree}; use rand_core::{Rng, SeedableRng}; use rand_pcg::Pcg64; @@ -530,38 +351,6 @@ mod tests { /// A prime small enough to check every pair of operands. const SMALL: u128 = 251; - #[test] - fn primality_agrees_with_trial_division() { - for candidate in 0u128..2_000 { - let expected = candidate >= 2 - && (2..candidate) - .take_while(|d| d * d <= candidate) - .all(|d| !candidate.is_multiple_of(d)); - assert_eq!(is_prime(candidate), expected, "{candidate}"); - } - } - - #[test] - fn primality_rejects_what_a_fermat_test_would_admit() { - // Carmichael numbers pass every Fermat test; Miller-Rabin does not. - for carmichael in [561u128, 41_041, 825_265] { - assert!(!is_prime(carmichael), "{carmichael}"); - } - } - - #[test] - fn primality_holds_at_the_widths_the_moduli_use() { - assert!(is_prime((1 << 100) - 15)); - assert!(is_prime((1 << 108) - 59)); - assert!(is_prime((1 << 114) - 11)); - // A semiprime with no factor small enough for the trial-division pass. - assert!(!is_prime(((1u128 << 54) - 33) * ((1u128 << 53) - 111))); - // Every odd value between the largest prime below 2^114 and 2^114. - for offset in (1..11).step_by(2) { - assert!(!is_prime((1 << 114) - offset), "2^114 - {offset}"); - } - } - #[test] fn ensure_traits() { fn assert_impl() {} @@ -570,7 +359,7 @@ mod tests { #[test] fn base_field_metadata() { - assert_eq!(FqDefault::MODULUS, Q100); + assert_eq!(FqDefault::META.modulus, Q100); assert_eq!(FqDefault::MODULUS_MINUS_ONE_DIV_TWO, (Q100 - 1) / 2); assert_eq!(FqDefault::modulus(), Q100); assert_eq!(FqDefault::min_value(), FqDefault::ZERO); @@ -582,41 +371,6 @@ mod tests { (rng.next_u64() as u128) << 64 | rng.next_u64() as u128 } - /// `(hi:lo) mod q` by binary long division — the independent reference - /// Barrett is checked against. Structurally unlike Barrett, which - /// estimates a quotient and corrects. - fn mod_reference(lo: u128, hi: u128, q: u128) -> u128 { - let mut rem = 0u128; - for i in (0..256).rev() { - let bit = if i >= 128 { - (hi >> (i - 128)) & 1 - } else { - (lo >> i) & 1 - }; - rem = (rem << 1) | bit; - if rem >= q { - rem -= q; - } - } - rem - } - - fn mulmod_reference(a: u128, b: u128, q: u128) -> u128 { - let (lo, hi) = mul_wide(a, b); - mod_reference(lo, hi, q) - } - - #[test] - fn mul_wide_matches_u128_on_small_inputs() { - let mut rng = Pcg64::seed_from_u64(401); - for _ in 0..512 { - let (a, b) = (rng.next_u64() as u128, rng.next_u64() as u128); - assert_eq!(mul_wide(a, b), (a * b, 0)); - } - assert_eq!(mul_wide(u128::MAX, u128::MAX), (1, u128::MAX - 1)); - assert_eq!(mul_wide(1u128 << 127, 2), (0, 1)); - } - #[test] fn barrett_matches_long_division() { let mut rng = Pcg64::seed_from_u64(402); @@ -679,11 +433,11 @@ mod tests { #[test] fn bits_matches_the_modulus() { - assert_eq!(Fq::::BITS, 100); - assert_eq!(Fq::::BITS, 8); + assert_eq!(Fq::::META.bits, 100); + assert_eq!(Fq::::META.bits, 8); // The defining bracket, `2^(BITS-1) <= Q < 2^BITS`. - const { assert!(1u128 << (Fq::::BITS - 1) <= Q100) }; - const { assert!(Q100 < 1u128 << Fq::::BITS) }; + const { assert!(1u128 << (Fq::::META.bits - 1) <= Q100) }; + const { assert!(Q100 < 1u128 << Fq::::META.bits) }; } #[test] diff --git a/crates/field/src/helpers.rs b/crates/field/src/helpers.rs new file mode 100644 index 0000000..882f3c5 --- /dev/null +++ b/crates/field/src/helpers.rs @@ -0,0 +1,351 @@ +//! Various helpers shared across multiple field implementations. + +/// The first thirteen primes: the standard deterministic Miller-Rabin base set +/// below `2^81.4`. +const PRIMALITY_BASES: [u128; 13] = [2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41]; + +/// Every modulus is below `2^MAX_MODULUS_BITS`: the bound of +/// [`barrett_reduce`], which every prime field here reduces by. +pub const MAX_MODULUS_BITS: u32 = 126; + +/// `Some` when trial division by [`PRIMALITY_BASES`] decides +/// `candidate`. +pub const fn sieve(candidate: u128) -> Option { + if candidate < 2 { + return Some(false); + } + let mut index = 0; + while index < PRIMALITY_BASES.len() { + let prime = PRIMALITY_BASES[index]; + if candidate == prime { + return Some(true); + } + if candidate.is_multiple_of(prime) { + return Some(false); + } + index += 1; + } + None +} + +/// Whether `candidate`, below `2^MAX_MODULUS_BITS`, is prime. +/// +/// Note that **this is not supposed to be used for testing adversarial primes**! +/// (See [`transcript::challenge::prime::sample`]) +pub const fn is_prime(candidate: u128) -> bool { + if let Some(decided) = sieve(candidate) { + return decided; + } + let modulus = FieldMetadata::new(candidate); + let mut index = 0; + while index < PRIMALITY_BASES.len() { + if !modulus.strong_round(PRIMALITY_BASES[index]) { + return false; + } + index += 1; + } + true +} + +/// `(a + b) mod q`, for `a, b < q < 2^127`. The sum cannot wrap. +const fn add_mod(a: u128, b: u128, q: u128) -> u128 { + let sum = a + b; + if sum >= q { sum - q } else { sum } +} + +/// `(a * b) mod q` by doubling, for any `q` below `2^127`: slow, but +/// independent of Barrett, so it is the reference [`FieldMetadata`] is checked +/// against. +pub const fn mul_mod(a: u128, b: u128, q: u128) -> u128 { + let mut result = 0u128; + let mut addend = a % q; + let mut remaining = b; + while remaining != 0 { + if remaining & 1 == 1 { + result = add_mod(result, addend, q); + } + addend = add_mod(addend, addend, q); + remaining >>= 1; + } + result +} + +/// Full 128x128 -> 256-bit product as `(low, high)`. +pub const fn mul_wide(a: u128, b: u128) -> (u128, u128) { + let (a_lo, a_hi) = (a as u64 as u128, a >> 64); + let (b_lo, b_hi) = (b as u64 as u128, b >> 64); + + let ll = a_lo * b_lo; + let hh = a_hi * b_hi; + let (mid, mid_carry) = (a_lo * b_hi).overflowing_add(a_hi * b_lo); + + let (lo, lo_carry) = ll.overflowing_add(mid << 64); + let hi = hh + (mid >> 64) + ((mid_carry as u128) << 64) + lo_carry as u128; + (lo, hi) +} + +/// `(lo, hi) >> n` for `n < 128`, keeping the low 128 bits. Every use here +/// shifts far enough that nothing above them survives. +pub const fn shr_wide(lo: u128, hi: u128, n: u32) -> u128 { + if n == 0 { + lo + } else { + (lo >> n) | (hi << (128 - n)) + } +} + +/// `floor(2^(2k) / q)` by binary long division, `k` being `q`'s bit length. +/// +/// Barrett's precomputed reciprocal. The numerator has a single set bit, so +/// each step shifts the remainder up and subtracts `q` when it fits. +pub const fn barrett_mu(q: u128, k: u32) -> u128 { + let mut rem = 0u128; + let mut quo = 0u128; + let mut i = 2 * k; + loop { + rem = (rem << 1) | (i == 2 * k) as u128; + let fits = rem >= q; + if fits { + rem -= q; + } + quo = (quo << 1) | fits as u128; + if i == 0 { + return quo; + } + i -= 1; + } +} + +/// `(lo, hi) mod q` for `q` of `k` bits and `mu` its reciprocal from +/// [`barrett_mu`]: Barrett reduction, Handbook of Applied Cryptography +/// Algorithm 14.42, after Barrett, CRYPTO '86, LNCS 263:311-323. Estimate the +/// quotient from the top bits of the input and the precomputed reciprocal, +/// subtract, then correct. The estimate is never more than two too small, so +/// two conditional subtractions suffice. +/// +/// That algorithm is stated for a radix above 3 and takes the difference +/// modulo `b^(k+1)`, which is wide enough to hold `3q`. Here the radix is 2, +/// where it is not, so the difference is taken modulo `2^128` instead. Hence +/// the bound `q < 2^126` on every modulus that comes through here. +#[inline] +pub const fn barrett_reduce(lo: u128, hi: u128, q: u128, mu: u128, k: u32) -> u128 { + let q1 = shr_wide(lo, hi, k - 1); + let (q2_lo, q2_hi) = mul_wide(q1, mu); + let q3 = shr_wide(q2_lo, q2_hi, k + 1); + + let mut r = lo.wrapping_sub(mul_wide(q3, q).0); + if r >= q { + r -= q; + } + if r >= q { + r -= q; + } + r +} + +/// A modulus in `[2, 2^MAX_MODULUS_BITS)` with its Barrett reciprocal: the +/// arithmetic of a modulus that is a value rather than a type, in `const` and +/// at runtime alike. What [`DynField`](crate::dynamic::DynField) installs. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FieldMetadata { + pub modulus: u128, + /// Bit length of the modulus. + pub bits: u32, + /// `floor(2^(2 * bits) / modulus)`. + pub mu: u128, +} + +impl FieldMetadata { + pub const fn new(modulus: u128) -> Self { + assert!( + modulus >= 2 && modulus < 1 << MAX_MODULUS_BITS, + "modulus out of range" + ); + assert!(modulus == 2 || modulus % 2 == 1, "modulus > 2 must be odd"); + let bits = u128::BITS - modulus.leading_zeros(); + Self { + modulus, + bits, + mu: barrett_mu(modulus, bits), + } + } + + /// `(a * b) mod q` for `a, b < q`. + #[inline] + pub const fn mul(&self, a: u128, b: u128) -> u128 { + let (lo, hi) = mul_wide(a, b); + barrett_reduce(lo, hi, self.modulus, self.mu, self.bits) + } + + /// `(base ^ exponent) mod q` for `base < q`, by square-and-multiply. + pub const fn pow(&self, base: u128, exponent: u128) -> u128 { + let mut result = 1 % self.modulus; + let mut bit = u128::BITS - exponent.leading_zeros(); + while bit > 0 { + bit -= 1; + result = self.mul(result, result); + if (exponent >> bit) & 1 == 1 { + result = self.mul(result, base); + } + } + result + } + + /// One strong probable-prime round of `q`, odd and at least 3, at `base` + /// in `[1, q - 1]`: with `q - 1 = odd * 2^shift`, `base^odd` is `±1` or + /// one of its next `shift - 1` squarings is `-1`. A prime passes every + /// base; a composite passes at most a quarter of them. + pub const fn strong_round(&self, base: u128) -> bool { + let shift = (self.modulus - 1).trailing_zeros(); + let odd = (self.modulus - 1) >> shift; + let mut witness = self.pow(base, odd); + if witness == 1 || witness == self.modulus - 1 { + return true; + } + let mut round = 1; + while round < shift { + witness = self.mul(witness, witness); + if witness == self.modulus - 1 { + return true; + } + round += 1; + } + false + } +} + +#[cfg(test)] +pub mod tests { + use super::*; + use rand_core::{Rng, SeedableRng}; + use rand_pcg::Pcg64; + + #[test] + fn primality_agrees_with_trial_division() { + for candidate in 0u128..2_000 { + let expected = candidate >= 2 + && (2..candidate) + .take_while(|d| d * d <= candidate) + .all(|d| !candidate.is_multiple_of(d)); + assert_eq!(is_prime(candidate), expected, "{candidate}"); + } + } + + #[test] + fn primality_rejects_what_a_fermat_test_would_admit() { + // Carmichael numbers pass every Fermat test; Miller-Rabin does not. + for carmichael in [561u128, 41_041, 825_265] { + assert!(!is_prime(carmichael), "{carmichael}"); + } + } + + #[test] + fn primality_holds_at_the_widths_the_moduli_use() { + assert!(is_prime((1 << 100) - 15)); + assert!(is_prime((1 << 108) - 59)); + assert!(is_prime((1 << 114) - 11)); + // A semiprime with no factor small enough for the trial-division pass. + assert!(!is_prime(((1u128 << 54) - 33) * ((1u128 << 53) - 111))); + // Every odd value between the largest prime below 2^114 and 2^114. + for offset in (1..11).step_by(2) { + assert!(!is_prime((1 << 114) - offset), "2^114 - {offset}"); + } + } + + #[test] + fn the_sieve_decides_only_what_trial_division_can() { + assert_eq!(sieve(0), Some(false)); + assert_eq!(sieve(1), Some(false)); + assert_eq!(sieve(2), Some(true)); + assert_eq!(sieve(41), Some(true)); + assert_eq!(sieve(43), None); + assert_eq!(sieve(1 << 100), Some(false)); + assert_eq!(sieve((1 << 100) - 15), None); + } + + #[test] + fn a_strong_round_catches_what_its_base_witnesses() { + // `2047 = 23 * 89` is a strong pseudoprime to base 2 and nothing else + // small; `1_373_653` is one to bases 2 and 3. + assert!(FieldMetadata::new(2047).strong_round(2)); + assert!(!FieldMetadata::new(2047).strong_round(3)); + assert!(FieldMetadata::new(1_373_653).strong_round(2)); + assert!(FieldMetadata::new(1_373_653).strong_round(3)); + assert!(!FieldMetadata::new(1_373_653).strong_round(5)); + // A prime passes every base. + let prime = FieldMetadata::new((1 << 61) - 1); + assert!((1..64).all(|base| prime.strong_round(base))); + } + + /// `(hi:lo) mod q` by binary long division — the independent reference + /// Barrett is checked against. Structurally unlike Barrett, which + /// estimates a quotient and corrects. + fn mod_reference(lo: u128, hi: u128, q: u128) -> u128 { + let mut rem = 0u128; + for i in (0..256).rev() { + let bit = if i >= 128 { + (hi >> (i - 128)) & 1 + } else { + (lo >> i) & 1 + }; + rem = (rem << 1) | bit; + if rem >= q { + rem -= q; + } + } + rem + } + + pub fn mulmod_reference(a: u128, b: u128, q: u128) -> u128 { + let (lo, hi) = mul_wide(a, b); + mod_reference(lo, hi, q) + } + + fn u128_of(rng: &mut Pcg64) -> u128 { + (rng.next_u64() as u128) << 64 | rng.next_u64() as u128 + } + + #[test] + fn the_modulus_multiplies_and_exponentiates_like_the_references() { + let mut rng = Pcg64::seed_from_u64(405); + for q in [ + 2u128, + 3, + 59, + 251, + (1 << 100) - 15, + (1 << 114) - 11, + (1 << 126) - 1, + ] { + let modulus = FieldMetadata::new(q); + for _ in 0..64 { + let (a, b) = (u128_of(&mut rng) % q, u128_of(&mut rng) % q); + assert_eq!( + modulus.mul(a, b), + mulmod_reference(a, b, q), + "{a} * {b} mod {q}" + ); + assert_eq!(modulus.mul(a, b), mul_mod(a, b, q), "{a} * {b} mod {q}"); + } + for a in [0, 1, q - 1] { + assert_eq!(modulus.mul(a, a), mul_mod(a, a, q), "{a}^2 mod {q}"); + let mut power = 1 % q; + for exponent in 0..20u128 { + assert_eq!(modulus.pow(a, exponent), power, "{a}^{exponent} mod {q}"); + power = mul_mod(power, a, q); + } + } + } + } + + #[test] + fn mul_wide_matches_u128_on_small_inputs() { + let mut rng = Pcg64::seed_from_u64(401); + for _ in 0..512 { + let (a, b) = (rng.next_u64() as u128, rng.next_u64() as u128); + assert_eq!(mul_wide(a, b), (a * b, 0)); + } + assert_eq!(mul_wide(u128::MAX, u128::MAX), (1, u128::MAX - 1)); + assert_eq!(mul_wide(1u128 << 127, 2), (0, 1)); + } +} diff --git a/crates/field/src/lib.rs b/crates/field/src/lib.rs index 5db31f2..45b5c32 100644 --- a/crates/field/src/lib.rs +++ b/crates/field/src/lib.rs @@ -2,8 +2,27 @@ #[cfg(feature = "spongefish")] mod codec; +pub mod dynamic; pub mod fq; pub mod gf128; +pub mod helpers; +pub use dynamic::DynField; pub use fq::{Fq, FqDefault, Q100}; pub use gf128::{F128, FixedBasePow, Wide256}; +pub use helpers::MAX_MODULUS_BITS; + +/// A prime field whose modulus is installed at run time, so a proof can run +/// over a prime its transcript draws. +pub trait FieldWithDynamicModulus { + /// Installs `modulus` process-wide. + /// + /// # Safety + /// + /// No instances of this field value may be alive, in any thread, or they may + /// become invalid. + /// + /// Despite that, no low-level Rust guarantees are violated and thus no true + /// Undefined Behavior™ could happen under any circumstances. + unsafe fn set_modulus(modulus: u128); +} diff --git a/crates/spartan/src/lib.rs b/crates/spartan/src/lib.rs index 70a4311..d3358e0 100644 --- a/crates/spartan/src/lib.rs +++ b/crates/spartan/src/lib.rs @@ -7,7 +7,8 @@ pub mod piop; pub mod sumcheck; pub use matrix::{ - PreparedConstraintMatrices, SpartanMatrixError, build_assignment_mle, build_product_mles, + PreparedConstraintMatrices, PreparedIntegerMatrices, SpartanMatrixError, build_assignment_mle, + build_product_mles, }; pub use piop::{ SpartanError, SpartanPiopProof, prove_spartan_piop, verify_spartan_proof, diff --git a/crates/spartan/src/matrix.rs b/crates/spartan/src/matrix.rs index 83e8bea..5d07089 100644 --- a/crates/spartan/src/matrix.rs +++ b/crates/spartan/src/matrix.rs @@ -1,14 +1,15 @@ //! R1CS matrix preparation and the sparse kernels used by Spartan. -use circuit::constraints::{ConstraintMatrices, SparseMatrix}; -use circuit::matrix_products::{IntegerProducts, ModularVector, RuntimeModulus}; +use circuit::constraints::{ConstraintMatrices, SparseBoolMatrix, SparseMatrix}; +use circuit::matrix_products::{IntegerProducts, ModularVector, RuntimeModulus, StoredInteger}; use circuit::witgen::PackedWitness; use circuit::{BitWidth, IntoWords}; use common::{BitzClaimField, BitzField}; +use num_traits::ToPrimitive; use poly::DenseMultilinearExtension; use rayon::prelude::*; use sha2::{Digest, Sha256}; -use transcript::Encoding; +use std::sync::Arc; use crate::sumcheck::R1csProductMles; @@ -26,22 +27,99 @@ pub enum SpartanMatrixError { InvalidMleOperation, } +/// Constraint matrices over their coefficient ring, prepared once for every +/// proof: shape validated, Boolean domains sized, nonzeros laid out row by +/// row and chunked by column, statement digested. None of that depends on a +/// modulus, so [`Self::project`] shares it and only reduces the coefficients. +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct PreparedIntegerMatrices { + m: SparseBoolMatrix, + layout: Arc, + /// The coefficients of `A`, `B` and `C`, each in its layout's order. + values: [Vec; 3], + /// `integer_matrix_digest`: over the integers, so it exists before a + /// modulus does. + digest: [u8; 32], +} + /// Immutable field-valued constraint matrices prepared for repeated Spartan -/// proofs and verification. +/// proofs and verification: the coefficients of `A`, `B` and `C` over the +/// field, their shared layout and the canonical statement digest. /// -/// Shape validation, Boolean-domain sizing, canonical statement hashing and -/// the column chunking are performed once during construction rather than -/// inside the prover or verifier. +/// Built by [`PreparedIntegerMatrices::project`] under a drawn modulus, or by +/// [`Self::new`] from matrices already over the field. The digest is stamped +/// by whichever built it: `project` carries the `integer_matrix_digest` its +/// source computed at setup, hashing nothing itself, while `new` hashes the +/// field coefficients it is given (`constraint_matrix_digest`). #[derive(Clone, Debug, Eq, PartialEq)] -pub struct PreparedConstraintMatrices { - matrices: ConstraintMatrices, - /// Nonzeros of `a`, `b` and `c` grouped by column chunk. - column_chunks: [ColumnChunkIndex; 3], +pub struct PreparedConstraintMatrices { + layout: Arc, + /// The coefficients of `A`, `B` and `C`, each in its layout's order. + values: [Vec; 3], digest: [u8; 32], +} + +/// Everything about `A`, `B` and `C` but their coefficients: positions in +/// compressed-row form, the column chunking and the Boolean domain sizes. +/// Modulus-free, so one layout serves the integer matrices and every +/// projection of them. +#[derive(Clone, Debug, Eq, PartialEq)] +struct Layout { + positions: [Positions; 3], + column_chunks: [ColumnChunkIndex; 3], + /// R1CS rows and assignment entries before padding. + row_count: usize, + column_count: usize, num_row_vars: usize, num_column_vars: usize, } +/// Where one matrix's nonzeros are: row `i` owns +/// `columns[row_starts[i]..row_starts[i + 1]]` and the same range of the +/// matrix's values. +#[derive(Clone, Debug, Eq, PartialEq)] +struct Positions { + row_starts: Vec, + columns: Vec, +} + +impl Positions { + /// Lays `matrix` out row by row, returning its coefficients in that order. + fn new(matrix: SparseMatrix) -> Result<(Self, Vec), SpartanMatrixError> { + let rows = matrix.into_rows(); + let nonzeros = rows.iter().map(|row| row.entries().len()).sum(); + let mut row_starts = Vec::with_capacity(rows.len() + 1); + let mut columns = Vec::with_capacity(nonzeros); + let mut values = Vec::with_capacity(nonzeros); + row_starts.push(0); + for row in rows { + for (column, coefficient) in row.into_entries() { + let column = + u32::try_from(column).map_err(|_| SpartanMatrixError::DomainTooLarge)?; + columns.push(column); + values.push(coefficient); + } + row_starts.push(columns.len()); + } + Ok(( + Self { + row_starts, + columns, + }, + values, + )) + } + + fn row_count(&self) -> usize { + self.row_starts.len() - 1 + } + + /// The positions of row `row`. + fn row(&self, row: usize) -> std::ops::Range { + self.row_starts[row]..self.row_starts[row + 1] + } +} + /// `bind_and_batch` splits the column domain into chunks of `2^16` columns and /// accumulates each chunk on its own thread. /// A chunk's slice of the dense table fits in cache. @@ -50,8 +128,8 @@ const BIND_CHUNK_COLUMN_VARS: usize = 16; /// The nonzeros of one sparse matrix grouped by column chunk. /// /// Chunk `c` owns columns `[c * chunk_len, (c + 1) * chunk_len)` and lists -/// the entries of every row that fall into it, in row order. The thread -/// accumulating chunk `c` reads only these entries and writes only these +/// the positions of every row that fall into it, in row order. The thread +/// accumulating chunk `c` reads only these positions and writes only these /// columns, so threads never share a column. #[derive(Clone, Debug, Eq, PartialEq)] struct ColumnChunkIndex { @@ -59,7 +137,7 @@ struct ColumnChunkIndex { spans: Vec>, } -/// The entries `[start, end)` of row `row`. +/// The positions `[start, end)` of row `row`. #[derive(Clone, Copy, Debug, Eq, PartialEq)] struct RowSpan { row: usize, @@ -68,21 +146,21 @@ struct RowSpan { } impl ColumnChunkIndex { - /// Groups `matrix` into `chunk_count` chunks of `chunk_len` columns. - fn new(matrix: &SparseMatrix, chunk_len: usize, chunk_count: usize) -> Self { + /// Groups `positions` into `chunk_count` chunks of `chunk_len` columns. + fn new(positions: &Positions, chunk_len: usize, chunk_count: usize) -> Self { debug_assert!(chunk_len.is_power_of_two()); - debug_assert!(chunk_len * chunk_count >= matrix.column_count()); let mut spans = vec![Vec::new(); chunk_count]; - for (row, entries) in matrix.rows().iter().enumerate() { - let entries = entries.entries(); - let mut start = 0; - while start < entries.len() { + for row in 0..positions.row_count() { + let range = positions.row(row); + let mut start = range.start; + while start < range.end { // Columns increase within a row, so a chunk's entries are a // contiguous run. - let chunk = entries[start].0 / chunk_len; + let chunk = positions.columns[start] as usize / chunk_len; let end = start - + entries[start..].partition_point(|(column, _)| column / chunk_len == chunk); + + positions.columns[start..range.end] + .partition_point(|&column| column as usize / chunk_len == chunk); spans[chunk].push(RowSpan { row, start, end }); start = end; } @@ -92,57 +170,138 @@ impl ColumnChunkIndex { } } -impl PreparedConstraintMatrices { - pub fn new(matrices: ConstraintMatrices) -> Result { +impl Layout { + /// Validates the shape, lays `a`, `b` and `c` out and chunks them, and + /// returns `m` and the coefficients in layout order. + #[allow(clippy::type_complexity)] + fn new( + matrices: ConstraintMatrices, + ) -> Result<(SparseBoolMatrix, Arc, [Vec; 3]), SpartanMatrixError> { let (num_row_vars, num_column_vars) = r1cs_num_vars(&matrices)?; + let row_count = matrices.a.row_count(); + let column_count = matrices.a.column_count(); + let ConstraintMatrices { m, a, b, c } = matrices; + let (a, a_values) = Positions::new(a)?; + let (b, b_values) = Positions::new(b)?; + let (c, c_values) = Positions::new(c)?; + let positions = [a, b, c]; let num_columns = 1_usize << num_column_vars; let chunk_len = num_columns.min(1_usize << BIND_CHUNK_COLUMN_VARS); let chunk_count = num_columns / chunk_len; + let column_chunks = [&positions[0], &positions[1], &positions[2]] + .map(|positions| ColumnChunkIndex::new(positions, chunk_len, chunk_count)); - let column_chunks = [&matrices.a, &matrices.b, &matrices.c] - .map(|matrix| ColumnChunkIndex::new(matrix, chunk_len, chunk_count)); - let digest = constraint_matrix_digest(&matrices)?; - - Ok(Self { - matrices, + let layout = Self { + positions, column_chunks, - digest, + row_count, + column_count, num_row_vars, num_column_vars, + }; + Ok((m, Arc::new(layout), [a_values, b_values, c_values])) + } +} + +impl PreparedIntegerMatrices +where + R: Send + Sync + ToPrimitive, + for<'a> StoredInteger: From<&'a R>, +{ + pub fn new(matrices: ConstraintMatrices) -> Result { + let (m, layout, values) = Layout::new(matrices)?; + let digest = integer_matrix_digest(&m, &layout, [&values[0], &values[1], &values[2]])?; + Ok(Self { + m, + layout, + values, + digest, + }) + } + + /// The Boolean witness matrix. + pub fn m(&self) -> &SparseBoolMatrix { + &self.m + } + + /// Number of R1CS rows before padding. + pub fn row_count(&self) -> usize { + self.layout.row_count + } + + /// Number of assignment entries before padding, the constant included. + pub fn column_count(&self) -> usize { + self.layout.column_count + } + + pub const fn digest(&self) -> &[u8; 32] { + &self.digest + } + + /// The matrices with `map` applied to every coefficient of `A`, `B` and + /// `C`, each into one contiguous table: the per-proof work. The layout + /// and the digest are shared. + pub fn project(&self, map: M) -> PreparedConstraintMatrices + where + F: BitzField, + M: Fn(&R) -> F + Sync, + { + let values = [&self.values[0], &self.values[1], &self.values[2]] + .map(|values| values.par_iter().map(&map).collect()); + PreparedConstraintMatrices { + layout: Arc::clone(&self.layout), + values, + digest: self.digest, + } + } +} + +impl PreparedConstraintMatrices { + /// Prepares matrices already over the field; the digest is + /// `constraint_matrix_digest`. + pub fn new(matrices: ConstraintMatrices) -> Result { + let (m, layout, values) = Layout::new(matrices)?; + let digest = constraint_matrix_digest(&m, &layout, [&values[0], &values[1], &values[2]])?; + Ok(Self { + layout, + values, + digest, }) } /// Returns human-readable info about R1CS matrices and their nonzero entries. pub fn short_debug_info(&self) -> String { - let nonzeros: usize = [&self.matrices.a, &self.matrices.b, &self.matrices.c] - .iter() - .flat_map(|matrix| matrix.rows()) - .map(|row| row.entries().len()) - .sum(); + let nonzeros: usize = self.values.iter().map(Vec::len).sum(); format!( "{} r1cs rows -> 2^{}, {} h entries -> 2^{}, {nonzeros} nonzeros", - self.matrices.a.row_count(), + self.row_count(), self.num_row_vars(), - self.matrices.a.column_count(), + self.column_count(), self.num_column_vars(), ) } - pub fn matrices(&self) -> &ConstraintMatrices { - &self.matrices + /// Number of R1CS rows before padding. + pub fn row_count(&self) -> usize { + self.layout.row_count + } + + /// Number of assignment entries before padding, the constant included. + pub fn column_count(&self) -> usize { + self.layout.column_count } pub const fn digest(&self) -> &[u8; 32] { &self.digest } - pub const fn num_row_vars(&self) -> usize { - self.num_row_vars + pub fn num_row_vars(&self) -> usize { + self.layout.num_row_vars } - pub const fn num_column_vars(&self) -> usize { - self.num_column_vars + pub fn num_column_vars(&self) -> usize { + self.layout.num_column_vars } } @@ -248,14 +407,40 @@ impl PreparedConstraintMatrices { row_point: &[F], rho: F, ) -> Result, SpartanMatrixError> { - bind_and_batch_with_num_vars( - &self.matrices, - &self.column_chunks, - row_point, - rho, - self.num_row_vars, - self.num_column_vars, - ) + let layout = &*self.layout; + if row_point.len() != layout.num_row_vars { + return Err(SpartanMatrixError::InvalidRowPointLength { + expected: layout.num_row_vars, + actual: row_point.len(), + }); + } + + let row_weights = poly::eq_table(row_point); + let batched = self.batched(rho); + let chunk_len = layout.column_chunks[0].chunk_len; + + // Every task owns one column chunk of the table and touches only the + // 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 << layout.num_column_vars).collect(); + evaluations + .par_chunks_mut(chunk_len) + .enumerate() + .for_each(|(chunk, output)| { + let base = chunk * chunk_len; + for (positions, index, values, batch_scale) in batched { + for span in &index.spans[chunk] { + let row_scale = row_weights[span.row] * batch_scale; + for k in span.start..span.end { + output[positions.columns[k] as usize - base] += row_scale * values[k]; + } + } + } + }); + + DenseMultilinearExtension::from_evaluations(layout.num_column_vars, evaluations) + .map_err(|_| SpartanMatrixError::InvalidMleOperation) } /// Directly evaluates @@ -271,127 +456,68 @@ impl PreparedConstraintMatrices { rho: F, column_point: &[F], ) -> Result { - evaluate_batched_with_num_vars( - &self.matrices, - &self.column_chunks, - row_point, - rho, - column_point, - self.num_row_vars, - self.num_column_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> { - if row_point.len() != num_row_vars { - return Err(SpartanMatrixError::InvalidRowPointLength { - expected: num_row_vars, - actual: row_point.len(), - }); - } + let layout = &*self.layout; + if row_point.len() != layout.num_row_vars { + return Err(SpartanMatrixError::InvalidRowPointLength { + expected: layout.num_row_vars, + actual: row_point.len(), + }); + } + if column_point.len() != layout.num_column_vars { + return Err(SpartanMatrixError::InvalidColumnPointLength { + expected: layout.num_column_vars, + actual: column_point.len(), + }); + } - let row_weights = poly::eq_table(row_point); - let batched = [ - (&matrices.a, &column_chunks[0], F::one()), - (&matrices.b, &column_chunks[1], rho), - (&matrices.c, &column_chunks[2], rho * rho), - ]; - let chunk_len = column_chunks[0].chunk_len; - - // Every task owns one column chunk of the table and touches only the - // 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(); - evaluations - .par_chunks_mut(chunk_len) - .enumerate() - .for_each(|(chunk, output)| { - let base = chunk * chunk_len; - for (matrix, index, batch_scale) in batched { - for span in &index.spans[chunk] { - let row_scale = row_weights[span.row] * batch_scale; - let entries = &matrix.rows()[span.row].entries()[span.start..span.end]; - for &(column, coefficient) in entries { - output[column - base] += row_scale * coefficient; + let row_weights = poly::eq_table(row_point); + let batched = self.batched(rho); + + // All tasks share one cache-sized `low` table and a per-chunk factor. + // Each task reduces one column chunk of nonzeros to a single field element. + let chunk_len = layout.column_chunks[0].chunk_len; + let (low_point, high_point) = column_point.split_at(chunk_len.ilog2() as usize); + let low_weights = poly::eq_table(low_point); + let high_weights = poly::eq_table(high_point); + debug_assert_eq!(high_weights.len(), layout.column_chunks[0].spans.len()); + + let evaluation = high_weights + .par_iter() + .enumerate() + .map(|(chunk, &chunk_weight)| { + let base = chunk * chunk_len; + let mut chunk_sum = F::zero(); + for (positions, index, values, batch_scale) in batched { + let mut matrix_sum = F::zero(); + for span in &index.spans[chunk] { + let span_sum = (span.start..span.end).fold(F::zero(), |sum, k| { + sum + low_weights[positions.columns[k] as usize - base] * values[k] + }); + matrix_sum += row_weights[span.row] * span_sum; } + chunk_sum += batch_scale * matrix_sum; } - } - }); - - DenseMultilinearExtension::from_evaluations(num_column_vars, evaluations) - .map_err(|_| SpartanMatrixError::InvalidMleOperation) -} + chunk_weight * chunk_sum + }) + .reduce(|| F::zero(), |left, right| left + right); -fn evaluate_batched_with_num_vars( - matrices: &ConstraintMatrices, - column_chunks: &[ColumnChunkIndex; 3], - row_point: &[F], - rho: F, - column_point: &[F], - num_row_vars: usize, - num_column_vars: usize, -) -> Result { - if row_point.len() != num_row_vars { - return Err(SpartanMatrixError::InvalidRowPointLength { - expected: num_row_vars, - actual: row_point.len(), - }); - } - if column_point.len() != num_column_vars { - return Err(SpartanMatrixError::InvalidColumnPointLength { - expected: num_column_vars, - actual: column_point.len(), - }); + Ok(evaluation) } - let row_weights = poly::eq_table(row_point); - let batched = [ - (&matrices.a, &column_chunks[0], F::one()), - (&matrices.b, &column_chunks[1], rho), - (&matrices.c, &column_chunks[2], rho * rho), - ]; - - // All tasks share one cache-sized `low` table and a per-chunk factor. - // Each task reduces one column chunk of nonzeros to a single field element. - let chunk_len = column_chunks[0].chunk_len; - let (low_point, high_point) = column_point.split_at(chunk_len.ilog2() as usize); - let low_weights = poly::eq_table(low_point); - let high_weights = poly::eq_table(high_point); - debug_assert_eq!(high_weights.len(), column_chunks[0].spans.len()); - - let evaluation = high_weights - .par_iter() - .enumerate() - .map(|(chunk, &chunk_weight)| { - let base = chunk * chunk_len; - let mut chunk_sum = F::zero(); - for (matrix, index, batch_scale) in batched { - 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 - }); - matrix_sum += row_weights[span.row] * span_sum; - } - chunk_sum += batch_scale * matrix_sum; - } - chunk_weight * chunk_sum + /// `A`, `B` and `C`, each with its positions, chunking and batching + /// scale: `1`, `rho` and `rho^2`. + fn batched(&self, rho: F) -> [(&Positions, &ColumnChunkIndex, &[F], F); 3] { + let layout = &*self.layout; + let scales = [F::one(), rho, rho * rho]; + std::array::from_fn(|i| { + ( + &layout.positions[i], + &layout.column_chunks[i], + self.values[i].as_slice(), + scales[i], + ) }) - .reduce(|| F::zero(), |left, right| left + right); - - Ok(evaluation) + } } pub(crate) fn r1cs_num_vars( @@ -407,61 +533,124 @@ pub(crate) fn r1cs_num_vars( } /// Canonically commits the complete public matrix statement before any -/// Fiat--Shamir challenge is sampled. -/// -/// 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( - matrices: &ConstraintMatrices, +/// Fiat--Shamir challenge is sampled: `domain`, the positions of `M`, then +/// `A`, `B` and `C` row by row with every coefficient written by `encode`. +fn matrix_digest( + domain: &[u8], + m: &SparseBoolMatrix, + layout: &Layout, + values: [&[C]; 3], + encode: impl Fn(&mut Sha256, &C) -> Result<(), SpartanMatrixError>, ) -> Result<[u8; 32], SpartanMatrixError> { - matrices - .validate_shape() - .map_err(|_| SpartanMatrixError::InvalidR1csShape)?; - let mut hash = Sha256::new(); - hash.update(b"bitz/spartan/constraint-matrices/v1"); + hash.update(domain); hash.update(b"M"); - hash_usize(&mut hash, matrices.m.row_count())?; - hash_usize(&mut hash, matrices.m.column_count())?; - for row in matrices.m.rows() { + hash_usize(&mut hash, m.row_count())?; + hash_usize(&mut hash, m.column_count())?; + for row in m.rows() { hash_usize(&mut hash, row.positions().len())?; for &column in row.positions() { hash_usize(&mut hash, column)?; } } - for (label, matrix) in [ - (b"A", &matrices.a), - (b"B", &matrices.b), - (b"C", &matrices.c), - ] { + for ((label, positions), values) in [b"A", b"B", b"C"] + .into_iter() + .zip(&layout.positions) + .zip(values) + { hash.update(label); - hash_sparse_matrix(&mut hash, matrix)?; + hash_usize(&mut hash, layout.row_count)?; + hash_usize(&mut hash, layout.column_count)?; + for row in 0..layout.row_count { + let range = positions.row(row); + hash_usize(&mut hash, range.len())?; + for k in range { + hash_usize(&mut hash, positions.columns[k] as usize)?; + encode(&mut hash, &values[k])?; + } + } } Ok(hash.finalize().into()) } -fn hash_sparse_matrix( - hash: &mut Sha256, - matrix: &SparseMatrix, -) -> Result<(), SpartanMatrixError> -where - F: Encoding<[u8]>, -{ - hash_usize(hash, matrix.row_count())?; - hash_usize(hash, matrix.column_count())?; - for row in matrix.rows() { - hash_usize(hash, row.entries().len())?; - for (column, coefficient) in row.entries() { - hash_usize(hash, *column)?; - let encoding = coefficient.encode(); +/// The statement digest over coefficients already in the field, encoded as +/// the transcript encodes them. +/// +/// 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. +fn constraint_matrix_digest( + m: &SparseBoolMatrix, + layout: &Layout, + values: [&[F]; 3], +) -> Result<[u8; 32], SpartanMatrixError> { + matrix_digest( + b"bitz/spartan/constraint-matrices/v1", + m, + layout, + values, + |hash, coefficient| { + let encoding = transcript::Encoding::encode(coefficient); let bytes = encoding.as_ref(); hash_usize(hash, bytes.len())?; hash.update(bytes); - } + Ok(()) + }, + ) +} + +/// The statement digest over integer coefficients, each in its normalised +/// two's-complement words, so it is fixed before any modulus is drawn and +/// shared by every projection. Its own domain: it never equals the digest of +/// the projected matrices. A protocol that projects must bind the modulus in +/// its transcript before Spartan reads this digest, as the e2e does by +/// absorbing its parameters right after the draw. +fn integer_matrix_digest( + m: &SparseBoolMatrix, + layout: &Layout, + values: [&[R]; 3], +) -> Result<[u8; 32], SpartanMatrixError> +where + R: ToPrimitive, + for<'a> StoredInteger: From<&'a R>, +{ + matrix_digest( + b"bitz/spartan/integer-constraint-matrices/v1", + m, + layout, + values, + |hash, coefficient| { + // Coefficients are small in practice; spare them the heap. + if let Some(value) = coefficient.to_i128() { + let (words, len) = canonical_words(value); + hash_words(hash, &words[..len]) + } else { + hash_words(hash, StoredInteger::from(coefficient).words()) + } + }, + ) +} + +/// The words [`StoredInteger`] holds for `value`, and how many: little-endian +/// two's complement with a redundant sign word dropped, none for zero. +fn canonical_words(value: i128) -> ([u64; 2], usize) { + let (low, high) = (value as u64, (value >> 64) as u64); + let sign_extended = (high == 0 && low >> 63 == 0) || (high == u64::MAX && low >> 63 == 1); + let len = match (value == 0, sign_extended) { + (true, _) => 0, + (false, true) => 1, + (false, false) => 2, + }; + ([low, high], len) +} + +fn hash_words(hash: &mut Sha256, words: &[u64]) -> Result<(), SpartanMatrixError> { + hash_usize(hash, words.len())?; + for word in words { + hash.update(word.to_le_bytes()); } Ok(()) } @@ -483,12 +672,14 @@ pub(crate) fn padded_num_vars(logical_len: usize) -> Result SparseMatrix { + fn random_sparse_matrix( + rng: &mut Pcg64, + mut coefficient: impl FnMut(&mut Pcg64) -> C, + ) -> SparseMatrix { let rows = (0..ROWS) .map(|_| { let mut columns: Vec = (0..ENTRIES_PER_ROW) @@ -506,19 +700,35 @@ mod tests { columns.dedup(); columns .into_iter() - .map(|column| (column, F::from(u128::from(rng.random::())))) + .map(|column| (column, coefficient(rng))) .collect() }) .collect(); SparseMatrix::try_from_rows(COLUMNS, rows).unwrap() } - 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); - let c = random_sparse_matrix(rng); - PreparedConstraintMatrices::new(ConstraintMatrices { m, a, b, c }).unwrap() + fn empty_m() -> SparseBoolMatrix { + SparseBoolMatrix::try_from_rows(1, vec![Vec::new(); COLUMNS]).unwrap() + } + + fn random_matrices(rng: &mut Pcg64) -> ConstraintMatrices { + let mut field = |rng: &mut Pcg64| F::from(u128::from(rng.random::())); + ConstraintMatrices { + m: empty_m(), + a: random_sparse_matrix(rng, &mut field), + b: random_sparse_matrix(rng, &mut field), + c: random_sparse_matrix(rng, &mut field), + } + } + + fn random_integer_matrices(rng: &mut Pcg64) -> ConstraintMatrices { + let mut integer = |rng: &mut Pcg64| i128::from(rng.random::()); + ConstraintMatrices { + m: empty_m(), + a: random_sparse_matrix(rng, &mut integer), + b: random_sparse_matrix(rng, &mut integer), + c: random_sparse_matrix(rng, &mut integer), + } } fn random_point(rng: &mut Pcg64, len: usize) -> Vec { @@ -552,15 +762,43 @@ mod tests { evaluation } + #[test] + fn the_layout_keeps_every_nonzero_in_row_order() { + let mut rng = Pcg64::seed_from_u64(5); + let matrices = random_matrices(&mut rng); + let prepared = PreparedConstraintMatrices::new(matrices.clone()).unwrap(); + assert_eq!(prepared.row_count(), ROWS); + assert_eq!(prepared.column_count(), COLUMNS); + + for ((matrix, positions), values) in [&matrices.a, &matrices.b, &matrices.c] + .into_iter() + .zip(&prepared.layout.positions) + .zip(&prepared.values) + { + assert_eq!(positions.row_count(), matrix.row_count()); + let mut k = 0; + for (row, entries) in matrix.rows().iter().enumerate() { + assert_eq!(positions.row(row), k..k + entries.entries().len()); + for &(column, coefficient) in entries.entries() { + assert_eq!(positions.columns[k] as usize, column); + assert_eq!(values[k], coefficient); + k += 1; + } + } + assert_eq!(k, values.len()); + } + } + #[test] fn column_chunks_partition_every_row() { let mut rng = Pcg64::seed_from_u64(7); - let prepared = random_prepared_matrices(&mut rng); - let matrices = prepared.matrices(); + let prepared = PreparedConstraintMatrices::new(random_matrices(&mut rng)).unwrap(); - for (matrix, index) in [&matrices.a, &matrices.b, &matrices.c] - .into_iter() - .zip(&prepared.column_chunks) + for (positions, index) in prepared + .layout + .positions + .iter() + .zip(&prepared.layout.column_chunks) { let mut spans: Vec<_> = index .spans @@ -571,17 +809,18 @@ mod tests { spans.sort_by_key(|(_, span)| (span.row, span.start)); let mut spans = spans.into_iter().peekable(); - for (row, entries) in matrix.rows().iter().enumerate() { - let mut next_start = 0; + for row in 0..positions.row_count() { + let range = positions.row(row); + let mut next_start = range.start; while let Some((chunk, span)) = spans.next_if(|(_, span)| span.row == row) { assert_eq!(span.start, next_start); assert!(span.end > span.start); - for &(column, _) in &entries.entries()[span.start..span.end] { - assert_eq!(column / index.chunk_len, chunk); + for &column in &positions.columns[span.start..span.end] { + assert_eq!(column as usize / index.chunk_len, chunk); } next_start = span.end; } - assert_eq!(next_start, entries.entries().len()); + assert_eq!(next_start, range.end); } assert!(spans.next().is_none()); } @@ -590,12 +829,13 @@ mod tests { #[test] fn evaluate_batched_matches_reference_and_bound_table() { let mut rng = Pcg64::seed_from_u64(11); - let prepared = random_prepared_matrices(&mut rng); + let matrices = random_matrices(&mut rng); + let prepared = PreparedConstraintMatrices::new(matrices.clone()).unwrap(); let row_point = random_point(&mut rng, prepared.num_row_vars()); let column_point = random_point(&mut rng, prepared.num_column_vars()); let rho = F::from(u128::from(rng.random::())); - let expected = reference_evaluation(prepared.matrices(), &row_point, rho, &column_point); + let expected = reference_evaluation(&matrices, &row_point, rho, &column_point); let evaluation = prepared .evaluate_batched(&row_point, rho, &column_point) .unwrap(); @@ -604,4 +844,84 @@ mod tests { let bound = prepared.bind_and_batch(&row_point, rho).unwrap(); assert_eq!(bound.evaluate(&column_point).unwrap(), expected); } + + #[test] + fn canonical_words_are_the_stored_integers() { + for value in [ + 0_i128, + 1, + -1, + i128::from(i64::MAX), + i128::from(i64::MIN), + i128::from(i64::MAX) + 1, + i128::from(i64::MIN) - 1, + i128::from(u64::MAX), + -i128::from(u64::MAX), + i128::MAX, + i128::MIN, + ] { + let (words, len) = super::canonical_words(value); + assert_eq!( + &words[..len], + StoredInteger::from(&value).words(), + "{value}" + ); + } + } + + /// Projecting integer matrices prepared once gives what preparing the + /// projected matrices gives, the layout shared rather than rebuilt; only + /// the digest differs, each over its own coefficients and domain. + #[test] + fn projection_matches_direct_preparation_up_to_the_digest_domain() { + let mut rng = Pcg64::seed_from_u64(13); + let project = |coefficient: &i128| { + let magnitude = F::from(coefficient.unsigned_abs()); + if *coefficient < 0 { + -magnitude + } else { + magnitude + } + }; + + let integer_matrices = random_integer_matrices(&mut rng); + let prepared = PreparedIntegerMatrices::new(integer_matrices.clone()).unwrap(); + let projected = prepared.project(project); + let direct = PreparedConstraintMatrices::new( + integer_matrices.clone().map_coefficients(|c| project(&c)), + ) + .unwrap(); + assert_eq!(projected.values, direct.values); + assert_eq!(projected.layout, direct.layout); + assert_ne!( + projected.digest(), + direct.digest(), + "the integer digest is its own domain" + ); + assert_eq!(projected.digest(), prepared.digest()); + assert!(Arc::ptr_eq( + &projected.layout, + &prepared.project(project).layout + )); + assert_eq!(prepared.row_count(), projected.row_count()); + assert_eq!(prepared.column_count(), projected.column_count()); + + // The integer digest is the integer matrices': the same under every + // projection, different as soon as a coefficient or a position is. + type Wide = field::Fq<{ (1 << 114) - 11 }>; + let wide = |coefficient: &i128| Wide::from(coefficient.unsigned_abs()); + assert_eq!(prepared.project(wide).digest(), prepared.digest()); + let mut negated = integer_matrices.clone(); + negated.a = negated.a.map_values_ref(|c| -c); + assert_ne!( + PreparedIntegerMatrices::new(negated).unwrap().digest(), + prepared.digest() + ); + let mut moved = integer_matrices; + moved.m = SparseBoolMatrix::try_from_rows(2, vec![vec![1]; COLUMNS]).unwrap(); + assert_ne!( + PreparedIntegerMatrices::new(moved).unwrap().digest(), + prepared.digest() + ); + } } diff --git a/crates/spartan/src/piop.rs b/crates/spartan/src/piop.rs index 332a16a..5a8271d 100644 --- a/crates/spartan/src/piop.rs +++ b/crates/spartan/src/piop.rs @@ -5,7 +5,7 @@ use common::BitzField; use poly::{ DenseMultilinearExtension, MleClaimError, ScaledMleEvaluationClaim, make_equality_factors, }; -use transcript::{ProverState, VerifierState}; +use transcript::{ProverState, SqueezableTranscript, VerifierState}; use crate::matrix::{PreparedConstraintMatrices, SpartanMatrixError, build_assignment_mle}; use crate::sumcheck::{ @@ -160,7 +160,7 @@ pub fn verify_spartan_with_mle_claim( mle_claim: &ScaledMleEvaluationClaim, assignment: &PackedWitness, ) -> Result<(), SpartanError> { - let assignment = build_assignment_mle::(assignment, matrices.matrices().a.column_count())?; + let assignment = build_assignment_mle::(assignment, matrices.column_count())?; let expected_claim = verify_spartan_proof(transcript, matrices, proof)?; if mle_claim != &expected_claim { return Err(SpartanError::InvalidMleClaim); diff --git a/crates/spartan/src/sumcheck.rs b/crates/spartan/src/sumcheck.rs index c2d9b6c..ed91277 100644 --- a/crates/spartan/src/sumcheck.rs +++ b/crates/spartan/src/sumcheck.rs @@ -32,7 +32,7 @@ use crypto_primitives::Semiring; use poly::DenseMultilinearExtension; use rayon::prelude::*; use std::array; -use transcript::{ProverState, VerifierState}; +use transcript::{ProverState, SqueezableTranscript, VerifierState}; /// Failures produced while reducing or checking a sumcheck claim. #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/crates/transcript/Cargo.toml b/crates/transcript/Cargo.toml index ab3f237..cf92566 100644 --- a/crates/transcript/Cargo.toml +++ b/crates/transcript/Cargo.toml @@ -8,3 +8,4 @@ license.workspace = true [dependencies] field = { workspace = true, features = ["spongefish"] } spongefish = { workspace = true } +thiserror = { workspace = true } diff --git a/crates/transcript/src/challenge.rs b/crates/transcript/src/challenge.rs index 2341383..c024213 100644 --- a/crates/transcript/src/challenge.rs +++ b/crates/transcript/src/challenge.rs @@ -1,6 +1,9 @@ //! Typed Fiat–Shamir challenge sampling. +pub mod prime; + use crate::{ProverState, VerifierState}; +use field::dynamic::DynField; use field::{F128, Fq}; /// A type that knows how to construct itself from transcript squeezes. @@ -19,37 +22,63 @@ pub trait TranscriptChallenge: Sized { fn from_squeezes(next_u128: impl FnMut() -> u128) -> Self; } -impl ProverState { +pub trait SqueezableTranscript { /// Samples a typed Fiat–Shamir challenge from this transcript. - pub fn squeeze(&mut self) -> T { + fn squeeze(&mut self) -> T; + + /// Squeezes a `bits`-bit prime from this transcript. + fn squeeze_prime(&mut self, bits: u32) -> u128; +} + +impl SqueezableTranscript for ProverState { + fn squeeze(&mut self) -> T { T::from_squeezes(|| self.verifier_message::()) } + + fn squeeze_prime(&mut self, bits: u32) -> u128 { + prime::sample(|| self.verifier_message::(), bits) + } } -impl VerifierState<'_> { - /// Samples a typed Fiat–Shamir challenge from this transcript. - pub fn squeeze(&mut self) -> T { +impl SqueezableTranscript for VerifierState<'_> { + fn squeeze(&mut self) -> T { T::from_squeezes(|| self.verifier_message::()) } + + fn squeeze_prime(&mut self, bits: u32) -> u128 { + prime::sample(|| self.verifier_message::(), bits) + } +} + +/// Sample `u128` whose residue modulo `modulus` is uniform, from successive +/// squeezes. +/// +/// The u128 range is not generally an exact multiple of the modulus. Reject +/// its incomplete final interval before reducing so every residue has the +/// same number of preimages. +fn unbiased_u128(modulus: u128, mut next_u128: impl FnMut() -> u128) -> u128 { + let rejection_remainder = (u128::MAX % modulus + 1) % modulus; + let max_accepted = u128::MAX - rejection_remainder; + + loop { + let candidate = next_u128(); + if candidate <= max_accepted { + return candidate; + } + } } impl TranscriptChallenge for Fq { - fn from_squeezes(mut next_u128: impl FnMut() -> u128) -> Self { + fn from_squeezes(next_u128: impl FnMut() -> u128) -> Self { // Validate the const-generic modulus before using it below. - let _ = Self::BITS; - - // The u128 range is not generally an exact multiple of Q. Reject its - // incomplete final interval before reducing so every residue has the - // same number of preimages. - let rejection_remainder = (u128::MAX % Q + 1) % Q; - let max_accepted = u128::MAX - rejection_remainder; + let _ = Self::META; + Self::from(unbiased_u128(Q, next_u128)) + } +} - loop { - let candidate = next_u128(); - if candidate <= max_accepted { - return Self::from(candidate); - } - } +impl TranscriptChallenge for DynField { + fn from_squeezes(next_u128: impl FnMut() -> u128) -> Self { + Self::from(unbiased_u128(DynField::config().modulus, next_u128)) } } @@ -62,10 +91,10 @@ impl TranscriptChallenge for F128 { #[cfg(test)] mod tests { - use field::{F128, Q100}; - - use super::TranscriptChallenge; + use super::*; use crate::{build_prover, build_verifier}; + use field::Q100; + use spongefish::Encoding; const SESSION: &[u8] = b"transcript/typed-challenge/test"; const INSTANCE: &[u8] = b"fq-rejection-sampling"; @@ -91,6 +120,24 @@ mod tests { assert_eq!(challenge, F::from(Q100 - 1)); } + /// Under the same modulus, the same stream gives the same element. + #[test] + fn dyn_field_squeezes_as_fq() { + field::dynamic::test_support::with_modulus(Q100, || { + let mut typed = build_prover(SESSION, INSTANCE); + let dynamic = typed.squeeze::(); + let mut fixed = build_prover(SESSION, INSTANCE); + let fq = fixed.squeeze::(); + assert_eq!(dynamic.encode().as_ref(), fq.encode().as_ref()); + + let rejection_remainder = (u128::MAX % Q100 + 1) % Q100; + let max_accepted = u128::MAX - rejection_remainder; + let mut candidates = [max_accepted + 1, max_accepted].into_iter(); + let challenge = DynField::from_squeezes(|| candidates.next().unwrap()); + assert_eq!(challenge, DynField::from(Q100 - 1)); + }); + } + #[test] fn prover_and_verifier_squeeze_in_lockstep() { let mut prover = build_prover(SESSION, INSTANCE); @@ -107,6 +154,21 @@ mod tests { verifier.check_eof().unwrap(); } + #[test] + fn prover_and_verifier_squeeze_the_same_prime() { + let mut prover = build_prover(SESSION, INSTANCE); + let prime = prover.squeeze_prime(F::META.bits); + let after = prover.squeeze::(); + let proof = prover.finish(); + assert_eq!(u128::BITS - prime.leading_zeros(), F::META.bits); + assert!(field::helpers::is_prime(prime)); + + let mut verifier = build_verifier(SESSION, INSTANCE, &proof); + assert_eq!(verifier.squeeze_prime(F::META.bits), prime); + assert_eq!(verifier.squeeze::(), after); + verifier.check_eof().unwrap(); + } + #[test] fn typed_f128_squeeze_matches_direct_decoding() { let mut typed = build_prover(SESSION, INSTANCE); diff --git a/crates/transcript/src/challenge/prime.rs b/crates/transcript/src/challenge/prime.rs new file mode 100644 index 0000000..58b3ad6 --- /dev/null +++ b/crates/transcript/src/challenge/prime.rs @@ -0,0 +1,120 @@ +//! Public prime sampling from the challenge stream. +//! +//! Candidates uniform among the odd integers of the requested width, sieved +//! by the small primes, then Miller–Rabin with bases drawn from the same +//! stream, so a prover grinding the transcript faces fresh bases at every +//! candidate. +//! +//! Arithmetic is variable-time. + +use field::MAX_MODULUS_BITS; +use field::helpers::{self, FieldMetadata}; + +/// Miller–Rabin bases per candidate. +/// +/// A composite passes a base with probability at most `1/4`, so it survives +/// with probability at most `2^-160`, and a search of up to `2^16` candidates +/// accepts a composite with probability below `2^-144`: the 128 bits and its +/// 16 bits of slack. +/// +/// A longer search does not happen: at least one odd integer in 44 is prime +/// below `2^126`, so `2^16` composites in a row have probability below `2^-2000`. +const MILLER_RABIN_ROUNDS: u32 = 80; + +/// A prime of exactly `bits` bits, `2 <= bits <= MAX_MODULUS_BITS`, from +/// successive `u128` draws, the same on both sides: the fingerprint prime. +/// +/// Each candidate is one draw, uniform among the odd integers in +/// `[2^(bits-1), 2^bits)`; each base is uniform in `[2, candidate - 2]` by +/// rejection. The search stops at the first candidate that passes. +/// +/// Should be drawn once per proof after the commitment and the statement +/// are absorbed and before the PIOP. +pub fn sample(mut next_u128: impl FnMut() -> u128, bits: u32) -> u128 { + assert!( + (2..=MAX_MODULUS_BITS).contains(&bits), + "prime width out of range" + ); + let bottom = 1u128 << (bits - 1); + loop { + let candidate = bottom | (next_u128() & (bottom - 1)) | 1; + let is_prime = helpers::sieve(candidate).unwrap_or_else(|| { + let modulus = FieldMetadata::new(candidate); + (0..MILLER_RABIN_ROUNDS) + .all(|_| modulus.strong_round(2 + sample_below(&mut next_u128, candidate - 3))) + }); + if is_prime { + return candidate; + } + } +} + +/// Uniform in `[0, bound)` for `bound >= 2`: draws masked to the bit length +/// of `bound - 1`, accepted below `bound`, so at most half are rejected. +fn sample_below(next_u128: &mut impl FnMut() -> u128, bound: u128) -> u128 { + debug_assert!(bound >= 2, "nothing to sample"); + let mask = u128::MAX >> (bound - 1).leading_zeros(); + loop { + let candidate = next_u128() & mask; + if candidate < bound { + return candidate; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// SplitMix64, two outputs per draw. + fn stream(mut state: u64) -> impl FnMut() -> u128 { + let mut next_u64 = move || { + state = state.wrapping_add(0x9e37_79b9_7f4a_7c15); + let mut z = state; + z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9); + z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb); + z ^ (z >> 31) + }; + move || u128::from(next_u64()) | (u128::from(next_u64()) << 64) + } + + /// `is_prime` is exact below `2^81.4` and a strong probable-prime test + /// above, which is all a check on the sampler's own test needs. + #[test] + fn sampled_primes_are_prime_and_of_the_requested_width() { + for bits in [2, 3, 8, 40, 64, 100, 113, MAX_MODULUS_BITS] { + for seed in 0..8 { + let prime = sample(stream(seed | u64::from(bits) << 8), bits); + assert_eq!( + u128::BITS - prime.leading_zeros(), + bits, + "{bits} bits, seed {seed}" + ); + assert!( + helpers::is_prime(prime), + "{bits} bits, seed {seed}: {prime} is composite" + ); + } + } + } + + #[test] + fn the_prime_is_a_function_of_the_stream() { + assert_eq!(sample(stream(7), 100), sample(stream(7), 100)); + assert_ne!(sample(stream(7), 100), sample(stream(8), 100)); + } + + #[test] + fn sampling_below_masks_then_rejects() { + // Bound 5 masks to three bits: 7 and 5 are rejected, `11 & 7 = 3` is not. + let mut draws = [7u128, 5, 8 | 3].into_iter(); + assert_eq!(sample_below(&mut || draws.next().unwrap(), 5), 3); + assert!(draws.next().is_none()); + } + + #[test] + #[should_panic(expected = "prime width out of range")] + fn rejects_a_width_the_field_cannot_hold() { + sample(stream(0), MAX_MODULUS_BITS + 1); + } +} diff --git a/crates/transcript/src/lib.rs b/crates/transcript/src/lib.rs index eb9b981..a4663ba 100644 --- a/crates/transcript/src/lib.rs +++ b/crates/transcript/src/lib.rs @@ -33,7 +33,7 @@ mod proof; mod prover; mod verifier; -pub use challenge::TranscriptChallenge; +pub use challenge::{SqueezableTranscript, TranscriptChallenge}; pub use domain::{PROTOCOL_LABEL, build_prover, build_verifier}; pub use proof::Proof; pub use prover::ProverState; diff --git a/crates/verifier/src/verify.rs b/crates/verifier/src/verify.rs index 7b47d9f..1376265 100644 --- a/crates/verifier/src/verify.rs +++ b/crates/verifier/src/verify.rs @@ -143,8 +143,8 @@ impl BitZVerifier { claim: &LinearClaim, transcript: &mut VerifierState<'_>, ) -> Result { - // Step 2 is absent: the field modulus is fixed, and BitZParams::new - // checks its fold bound. + // Step 2 is skipped: BitZParams::new checks the claim field bound + // `q < (|K| - 1) / k_1`, so a fold cannot wrap in the exponent. // Step 3: read the folds, range-check them, reconstruct against mu. let fold = self diff --git a/tooling/cli/benches/circuits.rs b/tooling/cli/benches/circuits.rs index 822d3b0..dfb0985 100644 --- a/tooling/cli/benches/circuits.rs +++ b/tooling/cli/benches/circuits.rs @@ -1,15 +1,15 @@ //! End-to-end and per-stage benchmarks; instance generation is outside timing. use bitz_cli::{ - ProjectBigIntToFq, benchmark, + ProjectBigIntToField, benchmark, circuits::{BuiltinCircuit, CircuitInstance}, end_to_end::CircuitProofSystem, }; use divan::Bencher; type R = num_bigint::BigInt; -type F = field::FqDefault; -type Proj = ProjectBigIntToFq; +type F = field::DynField; +type Proj = ProjectBigIntToField; fn main() { divan::main(); @@ -21,10 +21,10 @@ 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::<_, F>::new::(statement).unwrap(), + CircuitProofSystem::<_, F, _, Proj>::new(statement).unwrap(), inputs, ) } @@ -43,7 +43,7 @@ fn circuit_setup(bencher: Bencher, circuit: BuiltinCircuit) { bencher .with_inputs(|| CircuitInstance::random(circuit, None, None).unwrap()) .bench_local_values(|statement| { - CircuitProofSystem::<_, F>::new::(statement).unwrap() + CircuitProofSystem::<_, F, _, Proj>::new(statement).unwrap() }); } diff --git a/tooling/cli/benches/sha256_spartan.rs b/tooling/cli/benches/sha256_spartan.rs index 0066d51..47736ed 100644 --- a/tooling/cli/benches/sha256_spartan.rs +++ b/tooling/cli/benches/sha256_spartan.rs @@ -6,8 +6,8 @@ use std::sync::{Arc, Mutex}; -use bitz_cli::{ProjectBigIntToFq, ProjectConstraint}; -use circuit::constraints::ConstraintGenerator; +use bitz_cli::{ProjectBigIntToField, ProjectConstraint}; +use circuit::constraints::{ConstraintGenerator, ConstraintMatrices}; use circuit::sha256::sha256_block_aligned_circuit; use circuit::witgen::ProductWitgen; use divan::{AllocProfiler, Bencher, black_box}; @@ -30,7 +30,7 @@ const BLOCKS: &[usize] = &[608]; type R = num_bigint::BigInt; type F = field::FqDefault; -type Proj = ProjectBigIntToFq; +type Proj = ProjectBigIntToField; /// An R1CS instance with a witness satisfying it. #[derive(Debug, Clone)] @@ -85,7 +85,8 @@ fn build(blocks: usize) -> R1csInstanceWitness { // Lower to Q100 and pad to the Boolean domains. let projection = >::prepare(); - let matrices = integer_matrices.map_coefficients(|c| projection.project(&c)); + let matrices: ConstraintMatrices = + 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(); diff --git a/tooling/cli/src/benchmark.rs b/tooling/cli/src/benchmark.rs index 8511532..0ac921d 100644 --- a/tooling/cli/src/benchmark.rs +++ b/tooling/cli/src/benchmark.rs @@ -2,9 +2,10 @@ use crate::ProjectConstraint; use crate::end_to_end::{CircuitProofSystem, CircuitStatement, CircuitStats, Error}; -use circuit::matrix_products::ModularVector; +use circuit::matrix_products::StoredInteger; use circuit::{BitWidth, IntoWords}; use common::{BitzClaimField, BitzConstraintRing}; +use field::FieldWithDynamicModulus; use std::{ fmt, time::{Duration, Instant}, @@ -25,14 +26,14 @@ pub struct Timings { pub fn run(statement: S, inputs: &[bool]) -> Result where S: CircuitStatement, - F: BitzClaimField, + F: BitzClaimField + FieldWithDynamicModulus, F::Integer: BitWidth + IntoWords, - Vec: for<'a> From<&'a ModularVector<2>>, - R: BitzConstraintRing, + R: BitzConstraintRing + for<'a> From<&'a StoredInteger>, + for<'a> StoredInteger: From<&'a R>, Proj: ProjectConstraint, { let started = Instant::now(); - let prepared = CircuitProofSystem::::new::(statement)?; + let prepared = CircuitProofSystem::<_, _, _, Proj>::new(statement)?; let setup = started.elapsed(); let circuit = prepared.stats(); tracing::info!( @@ -41,6 +42,7 @@ where assignment_bits = circuit.assignment_bits, committed_bits = circuit.committed_bits, padded_committed_bits = circuit.padded_committed_bits, + prime_bits = circuit.prime_bits, "Circuit prepared", ); let started = Instant::now(); @@ -69,12 +71,13 @@ impl fmt::Display for Timings { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!( f, - "opening_path={:?} constraints={} assignment_bits={} committed_bits={} padded_committed_bits={} setup_ms={:.3} witness_ms={:.3} commit_ms={:.3} prove_ms={:.3} total_prove_ms={:.3} verify_ms={:.3}", + "opening_path={:?} constraints={} assignment_bits={} committed_bits={} padded_committed_bits={} prime_bits={} setup_ms={:.3} witness_ms={:.3} commit_ms={:.3} prove_ms={:.3} total_prove_ms={:.3} verify_ms={:.3}", self.circuit.opening_path, self.circuit.constraints, self.circuit.assignment_bits, self.circuit.committed_bits, self.circuit.padded_committed_bits, + self.circuit.prime_bits, self.setup.as_secs_f64() * 1000., self.witness.as_secs_f64() * 1000., self.commit.as_secs_f64() * 1000., diff --git a/tooling/cli/src/cmd/circuit_e2e.rs b/tooling/cli/src/cmd/circuit_e2e.rs index 31ff5a5..4f5184c 100644 --- a/tooling/cli/src/cmd/circuit_e2e.rs +++ b/tooling/cli/src/cmd/circuit_e2e.rs @@ -8,10 +8,10 @@ use { }, }; -type F = field::FqDefault; -type Proj = bitz_cli::ProjectBigIntToFq; +type F = field::DynField; +type Proj = bitz_cli::ProjectBigIntToField; -/// Prove generated circuit constraints over Q100 and verify the proof. +/// Prove generated circuit constraints over a transcript-drawn prime and verify the proof. #[derive(FromArgs, PartialEq, Debug)] #[argh(subcommand, name = "circuit-e2e")] pub struct Args { @@ -48,7 +48,7 @@ impl Command for Args { "circuit_e2e", circuit = %self.circuit, threads = rayon::current_num_threads(), - field = "Q100", + field = "random-prime", pcs = "Fast", hash = "Blake3", ).entered(); @@ -58,7 +58,7 @@ impl Command for Args { let inputs = statement.inputs.clone(); 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!("circuit={} threads={} field=random-prime pcs=Fast hash=Blake3 relation=r1cs constraints_verified=true", self.circuit, rayon::current_num_threads()); println!("{timings}"); Ok(()) }) diff --git a/tooling/cli/src/end_to_end.rs b/tooling/cli/src/end_to_end.rs index 6ddf1f9..39cc22d 100644 --- a/tooling/cli/src/end_to_end.rs +++ b/tooling/cli/src/end_to_end.rs @@ -1,9 +1,14 @@ -//! Circuit constraints over a prime field, 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. +//! +//! The field's modulus is sampled at a runtime once the commitment and the +//! statement are absorbed. +use circuit::matrix_products::StoredInteger; use circuit::{ BitWidth, Circuit, IntoWords, constraints::{ConstraintGenerator, SparseBoolMatrix}, - matrix_products::ModularVector, + matrix_products::IntegerProducts, matrix_transpose::{MTransposeGenerator, MaterializedMTranspose}, witgen::{PackedWitness, ProductWitgen}, }; @@ -12,27 +17,37 @@ use common::{ VirtualMap, VirtualStatement, shape::{MIN_LOG_BITS, PACK_BITS}, }; -use field::{F128, gf128::smallest_generator}; +use field::{F128, FieldWithDynamicModulus, gf128::smallest_generator}; use num_traits::{ConstOne, ConstZero}; use pcs::{CommitScheme, HashKind, LigeritoProfile, Pcs, ProverData, StatementBinding}; -use poly::{DenseMultilinearExtension, ScaledMleEvaluationClaim}; +use poly::ScaledMleEvaluationClaim; use prover::{BitZProver, VirtualWitness}; -use transcript::{ProverState, PublicTranscript, build_prover, build_verifier}; +use std::marker::PhantomData; +use std::sync::{Mutex, MutexGuard, PoisonError}; +use transcript::{ + ProverState, PublicTranscript, SqueezableTranscript, build_prover, build_verifier, +}; use verifier::BitZVerifier; use crate::ProjectConstraint; use spartan::{ - PreparedConstraintMatrices, R1csProductMles, SpartanPiopProof, build_assignment_mle, - build_product_mles, prove_spartan_piop, verify_spartan_proof, + PreparedConstraintMatrices, PreparedIntegerMatrices, SpartanMatrixError, SpartanPiopProof, + build_assignment_mle, build_product_mles, prove_spartan_piop, verify_spartan_proof, }; -const SESSION: &[u8] = b"bitz/circuit-e2e/v1"; +const SESSION: &[u8] = b"bitz/circuit-e2e/v2"; const WINDOW: u32 = 8; +/// Held for the whole of a proof or a verification: the modulus they install +/// is process-wide. +static DYNAMIC_MODULUS_LOCK: Mutex<()> = Mutex::new(()); + +type ModulusLock<'a> = MutexGuard<'a, ()>; + /// 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 -/// the modulus of the field the proof system is built over. +/// the prime the proof draws. pub trait CircuitStatement { fn domain(&self) -> &'static [u8]; fn public_bytes(&self) -> Vec; @@ -47,7 +62,7 @@ pub enum Error { #[error("invalid proof configuration: {0}")] Configuration(&'static str), #[error("matrix preparation failed: {0:?}")] - Matrix(spartan::SpartanMatrixError), + Matrix(SpartanMatrixError), #[error("circuit witness does not satisfy the constraints")] Unsatisfied, #[error("Spartan failed: {0:?}")] @@ -81,26 +96,36 @@ pub struct CircuitStats { /// Meaningful committed witness bits before PCS zero padding. pub committed_bits: usize, pub padded_committed_bits: usize, + /// Width of the fingerprint prime each proof draws. + pub prime_bits: u32, } -/// A circuit prepared for proving over the prime field `F`. +/// A circuit prepared for proving over the prime field `F`, whose modulus +/// each proof draws; `Proj` projects the constraints from `R` onto it. #[derive(Debug)] -pub struct CircuitProofSystem { +pub struct CircuitProofSystem { statement: S, opening_path: OpeningPath, - matrices: PreparedConstraintMatrices, + /// Prepared once over the integers; projected onto `F` under the drawn prime. + matrices: PreparedIntegerMatrices, map: MaterializedMTranspose, - params: BitZParams, + claim_shape: Shape, + generator: F128, + /// Width of the fingerprint prime, fixed by the claim shape. + prime_bits: u32, committed_shape: Shape, pcs: Pcs, + _phantom: PhantomData<(F, Proj)>, } +/// A witness kept exact: it takes field values once the prime is drawn. #[derive(Debug)] -pub struct Witness { - committed: Vec, - assignment_bits: Vec, - assignment: DenseMultilinearExtension, - products: R1csProductMles, +pub struct Witness { + committed_bits: Vec, + /// Only exist on a virtual opening path + virtual_bits: Option>, + assignment: PackedWitness, + products: IntegerProducts, } /// Commitment data and the transcript that sampled its OOD claim. @@ -109,6 +134,7 @@ pub struct CommittedWitness { transcript: ProverState, } +/// The residues are under the prime the transcript yields for `root`. #[derive(Clone, Debug)] pub struct Proof { pub root: Root, @@ -116,43 +142,36 @@ pub struct Proof { pub opening: transcript::Proof, } -impl CircuitProofSystem +impl CircuitProofSystem where S: CircuitStatement, - F: BitzClaimField, + F: BitzClaimField + FieldWithDynamicModulus, F::Integer: BitWidth + IntoWords, - Vec: for<'a> From<&'a ModularVector<2>>, // Needed for `build_product_mles` + R: BitzConstraintRing + for<'a> From<&'a StoredInteger>, + for<'a> StoredInteger: From<&'a R>, + Proj: ProjectConstraint, { #[tracing::instrument(name = "setup", skip_all)] - pub fn new(statement: S) -> Result - where - R: BitzConstraintRing, - Proj: ProjectConstraint, - { + pub fn new(statement: S) -> Result { 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| projection.project(&c)), - ) - .map_err(Error::Matrix)?; + let matrices = + PreparedIntegerMatrices::new(constraints.into_matrices()).map_err(Error::Matrix)?; let mut generator = MTransposeGenerator::new(statement.input_bits()); let inputs = generator.take_inputs(); statement.synthesize(&mut generator, &inputs)?; let map = generator.finish(); - if map.h_len() != matrices.matrices().a.column_count() { + if map.h_len() != matrices.column_count() { return Err(Error::Configuration("map and assignment dimensions differ")); } - if map.f_len() != matrices.matrices().m.column_count() { + if map.f_len() != matrices.m().column_count() { return Err(Error::Configuration("map and witness dimensions differ")); } let claim_shape = shape_for(map.h_len())?; - let opening_path = if is_identity(&matrices.matrices().m) { + let opening_path = if is_identity(matrices.m()) { OpeningPath::Direct } else { OpeningPath::Virtual @@ -161,8 +180,8 @@ where OpeningPath::Direct => claim_shape, OpeningPath::Virtual => shape_for(map.f_len() - 1)?, }; - let params = BitZParams::new(claim_shape, smallest_generator()) - .map_err(|_| Error::Configuration("inadmissible BitZ parameters"))?; + let prime_bits = common::prime_bits(&claim_shape) + .map_err(|_| Error::Configuration("shape too tall for the fingerprint prime"))?; let pcs = Pcs::new(&committed_shape, LigeritoProfile::Fast, HashKind::Blake3) .map_err(|_| Error::Configuration("unsupported PCS shape"))?; Ok(Self { @@ -170,27 +189,31 @@ where opening_path, matrices, map, - params, + claim_shape, + generator: smallest_generator(), + prime_bits, committed_shape, pcs, + _phantom: PhantomData, }) } pub fn stats(&self) -> CircuitStats { CircuitStats { opening_path: self.opening_path, - constraints: self.matrices.matrices().a.row_count(), + constraints: self.matrices.row_count(), assignment_bits: self.map.h_len(), committed_bits: match self.opening_path { OpeningPath::Direct => self.map.h_len(), OpeningPath::Virtual => self.map.f_len() - 1, }, padded_committed_bits: self.pcs.bit_len(), + prime_bits: self.prime_bits, } } #[tracing::instrument(name = "witness", skip_all)] - pub fn witness(&self, inputs: &[bool]) -> Result, Error> { + pub fn witness(&self, inputs: &[bool]) -> Result { if inputs.len() != self.statement.input_bits() { return Err(Error::Input("wrong witness input length")); } @@ -200,27 +223,22 @@ where if f.bit_len() + 1 != self.map.f_len() || h.bit_len() != self.map.h_len() { return Err(Error::Input("circuit replay changed witness dimensions")); } - let products = build_product_mles(&products, self.matrices.matrices().a.row_count()) - .map_err(Error::Matrix)?; - if products - .az - .iter() - .zip(products.bz.iter()) - .zip(products.cz.iter()) - .any(|((&a, &b), &c)| a * b != c) - { + if products.a_mw.len() != self.matrices.row_count() { + return Err(Error::Input("circuit replay changed constraint count")); + } + // Over the integers, so modulo whichever prime is drawn. + if !products.is_satisfied::() { return Err(Error::Unsatisfied); } - let assignment = build_assignment_mle(&h, self.map.h_len()).map_err(Error::Matrix)?; - let assignment_bits = pack(&h, *self.params.shape()); - let (committed, assignment_bits) = match self.opening_path { - OpeningPath::Direct => (assignment_bits, Vec::new()), - OpeningPath::Virtual => (pack(&f, self.committed_shape), assignment_bits), + let assignment_bits = pack(&h, self.claim_shape); + let (committed_bits, virtual_bits) = match self.opening_path { + OpeningPath::Direct => (assignment_bits, None), + OpeningPath::Virtual => (pack(&f, self.committed_shape), Some(assignment_bits)), }; Ok(Witness { - committed, - assignment_bits, - assignment, + committed_bits, + virtual_bits, + assignment: h, products, }) } @@ -228,22 +246,22 @@ where /// 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 - .commit(&witness.committed, &mut transcript) + .commit(&witness.committed_bits, &mut transcript) .map_err(Error::Commit)?; Ok(CommittedWitness { data, transcript }) } - /// Continues the commitment transcript through Spartan and the BitZ opening. + /// Continues the commitment transcript through the prime draw, 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, Error> { + pub fn prove(&self, witness: Witness, commitment: CommittedWitness) -> Result, Error> { + let prime_lock = DYNAMIC_MODULUS_LOCK + .lock() + .unwrap_or_else(PoisonError::into_inner); let CommittedWitness { data, mut transcript, @@ -254,37 +272,46 @@ where self.pcs .prove_lin( &data, - witness.committed.clone(), + witness.committed_bits.clone(), &self.constant_query(), StatementBinding::Bind, &mut transcript, ) .map_err(Error::ConstantProve)?; } - let (spartan, terminal) = prove_spartan_piop( - &mut transcript, - &self.matrices, - &witness.products, - &witness.assignment, - ) - .map_err(Error::Spartan)?; - let claim = opening_claim(&self.params, &terminal)?; - let prover = BitZProver::new(self.params, WINDOW); + let prime = transcript.squeeze_prime(self.prime_bits); + let (params, matrices) = self.under_prime(prime, &prime_lock, &mut transcript)?; + let products = + build_product_mles(&witness.products, matrices.row_count()).map_err(Error::Matrix)?; + let assignment = + build_assignment_mle(&witness.assignment, self.map.h_len()).map_err(Error::Matrix)?; + let (spartan, terminal) = + prove_spartan_piop(&mut transcript, &matrices, &products, &assignment) + .map_err(Error::Spartan)?; + let claim = opening_claim(¶ms, &terminal)?; + let prover = BitZProver::new(params, WINDOW); match self.opening_path { - OpeningPath::Direct => { - prover.prove(&claim, &self.pcs, &data, witness.committed, &mut transcript) - } + OpeningPath::Direct => prover.prove( + &claim, + &self.pcs, + &data, + witness.committed_bits, + &mut transcript, + ), OpeningPath::Virtual => { let statement = - VirtualStatement::new(self.params, self.committed_shape, &self.map, &claim) + VirtualStatement::new(params, self.committed_shape, &self.map, &claim) .map_err(|_| Error::Configuration("invalid virtual statement"))?; prover.prove_virtual( &statement, &self.pcs, &data, VirtualWitness { - committed_bits: witness.committed, - virtual_bits: &witness.assignment_bits, + committed_bits: witness.committed_bits, + virtual_bits: witness + .virtual_bits + .as_ref() + .expect("virtual_bits must exist on virtual path"), }, &mut transcript, ) @@ -300,6 +327,9 @@ where #[tracing::instrument(name = "verify", skip_all)] pub fn verify(&self, proof: &Proof) -> Result<(), Error> { + let prime_lock = DYNAMIC_MODULUS_LOCK + .lock() + .unwrap_or_else(PoisonError::into_inner); let mut transcript = build_verifier(SESSION, self.statement.domain(), &proof.opening); let commitment = self .pcs @@ -316,17 +346,19 @@ where ) .map_err(Error::ConstantVerify)?; } - let terminal = verify_spartan_proof(&mut transcript, &self.matrices, &proof.spartan) + let prime = transcript.squeeze_prime(self.prime_bits); + let (params, matrices) = self.under_prime(prime, &prime_lock, &mut transcript)?; + let terminal = verify_spartan_proof(&mut transcript, &matrices, &proof.spartan) .map_err(Error::Spartan)?; - let claim = opening_claim(&self.params, &terminal)?; - let verifier = BitZVerifier::new(self.params, WINDOW); + let claim = opening_claim(¶ms, &terminal)?; + let verifier = BitZVerifier::new(params, WINDOW); match self.opening_path { OpeningPath::Direct => { verifier.verify_with_commitment(&claim, &self.pcs, &commitment, transcript) } OpeningPath::Virtual => { let statement = - VirtualStatement::new(self.params, self.committed_shape, &self.map, &claim) + VirtualStatement::new(params, self.committed_shape, &self.map, &claim) .map_err(|_| Error::Configuration("invalid virtual statement"))?; verifier.verify_virtual_with_commitment( &statement, @@ -339,6 +371,26 @@ where .map_err(Error::Verify) } + /// Installs the sampled prime and builds the parameters and the matrices + /// under it. The parameters' frame carries the prime into the transcript. + fn under_prime( + &self, + prime: u128, + _prime_lock: &ModulusLock<'_>, + transcript: &mut impl PublicTranscript, + ) -> Result<(BitZParams, PreparedConstraintMatrices), Error> { + tracing::info!(bits = self.prime_bits, prime = %prime, "Fingerprint prime drawn"); + // SAFETY: No value of `F` exists yet, and `DYNAMIC_MODULUS_LOCK` is held, so + // no operation is in flight elsewhere. + unsafe { F::set_modulus(prime) }; + let params = BitZParams::new(self.claim_shape, self.generator) + .map_err(|_| Error::Configuration("inadmissible BitZ parameters"))?; + transcript.public_message(¶ms); + let projection = Proj::prepare(); + let matrices = self.matrices.project(|c| projection.project(c)); + Ok((params, matrices)) + } + // Direct commitment includes h[0]; unlike the virtual map, it does not // supply that coordinate as a fixed one. Opening at zero enforces h[0] = 1. fn constant_query(&self) -> OpeningQuery { @@ -348,12 +400,15 @@ where } } + /// Everything public and fixed before the prime is drawn. fn bind(&self, transcript: &mut impl PublicTranscript, root: Root) { let public = self.statement.public_bytes(); transcript.public_message(&(public.len() as u64)); transcript.public_message(public.as_slice()); transcript.public_message(&root.0); - transcript.public_message(&self.params); + transcript.public_message(&(self.claim_shape.log_rows() as u64)); + transcript.public_message(&(self.claim_shape.log_columns() as u64)); + transcript.public_message(&self.generator); transcript.public_message(&self.pcs); transcript.public_message(&self.map.digest()); transcript.public_message(b"bitz/circuit-opening-path/v1"); @@ -414,10 +469,13 @@ fn opening_claim( #[cfg(test)] mod tests { use super::*; - use crate::ProjectBigIntToFq; + use crate::ProjectBigIntToField; + use poly::DenseMultilinearExtension; - type F = field::FqDefault; - type Proj = ProjectBigIntToFq; + type F = field::DynField; + type R = num_bigint::BigInt; + type Proj = ProjectBigIntToField; + type System = CircuitProofSystem; #[test] fn only_exact_identity_maps_select_direct_opening() { @@ -458,7 +516,7 @@ mod tests { #[test] fn commitment_sends_ood_before_proving() { - let system = CircuitProofSystem::<_, F>::new::<_, Proj>(IdentityBit).unwrap(); + let system = System::new(IdentityBit).unwrap(); let witness = system.witness(&[true]).unwrap(); let committed = system.commit(&witness).unwrap(); let proof = committed.transcript.finish(); @@ -474,7 +532,7 @@ mod tests { #[test] fn direct_opening_requires_constant_one_on_both_sides() { - let mut system = CircuitProofSystem::<_, F>::new::<_, Proj>(IdentityBit).unwrap(); + let mut system = System::new(IdentityBit).unwrap(); let witness = system.witness(&[true]).unwrap(); let data = system.commit(&witness).unwrap(); let mut proof = system.prove(witness, data).unwrap(); @@ -484,7 +542,7 @@ mod tests { system.opening_path = OpeningPath::Direct; let mut bad_witness = system.witness(&[true]).unwrap(); - bad_witness.committed.fill(F128::ZERO); + bad_witness.committed_bits.fill(F128::ZERO); let bad_data = system.commit(&bad_witness).unwrap(); assert!(matches!( system.prove(bad_witness, bad_data), @@ -518,8 +576,31 @@ mod tests { )); } + /// The prime is the transcript's: the same statement and commitment + /// yield it again, another statement or commitment yields another. + #[test] + fn the_prime_follows_the_transcript() { + let system = System::new(IdentityBit).unwrap(); + let witness = system.witness(&[true]).unwrap(); + let root = system.commit(&witness).unwrap().data.root(); + let draw = |domain: &'static [u8], root: Root| { + let mut transcript = build_prover(SESSION, domain); + system.bind(&mut transcript, root); + transcript.squeeze_prime(system.prime_bits) + }; + let prime = draw(system.statement.domain(), root); + assert_eq!(u128::BITS - prime.leading_zeros(), system.prime_bits); + assert!(field::helpers::is_prime(prime)); + assert_eq!(draw(system.statement.domain(), root), prime); + assert_ne!(draw(b"another/statement", root), prime); + assert_ne!(draw(system.statement.domain(), Root([0; 32])), prime); + } + + /// `opening_claim` is generic, so a fixed field keeps this test off the + /// installed modulus. #[test] fn scaled_claim_conversion_preserves_values_and_zero_scale() { + type F = field::FqDefault; let params = BitZParams::::new(Shape::new(7, 15).unwrap(), smallest_generator()).unwrap(); let assignment = diff --git a/tooling/cli/src/lib.rs b/tooling/cli/src/lib.rs index 56990ee..9aae2a9 100644 --- a/tooling/cli/src/lib.rs +++ b/tooling/cli/src/lib.rs @@ -4,7 +4,7 @@ pub mod benchmark; pub mod circuits; pub mod end_to_end; -use field::Fq; +use common::BitzClaimField; use num_bigint::BigInt; use num_traits::{Signed, ToPrimitive}; @@ -17,27 +17,94 @@ pub trait ProjectConstraint: Send + Sync { fn project(&self, constraint: &R) -> F; } -/// Projects [`BigInt`] constraints onto [`Fq`] by reducing it canonically modulo `Q`. -pub struct ProjectBigIntToFq { - modulus: BigInt, +/// Projects [`BigInt`] constraints onto a prime field by reducing them +/// canonically modulo its modulus, read when prepared. +#[derive(Debug, Clone)] +pub struct ProjectBigIntToField { + modulus: u128, + wide_modulus: BigInt, } -impl ProjectConstraint> for ProjectBigIntToFq { +impl ProjectConstraint for ProjectBigIntToField { fn prepare() -> Self { + let modulus = F::modulus().to_u128().expect("the modulus fits a u128"); Self { - modulus: BigInt::from(Q), + modulus, + wide_modulus: BigInt::from(modulus), } } - fn project(&self, constraint: &BigInt) -> Fq { - let mut reduced = constraint % &self.modulus; + fn project(&self, constraint: &BigInt) -> F { + // Coefficients are small in practice; reduce them without the heap. + if let Some(value) = constraint.to_i128() { + let magnitude = value.unsigned_abs(); + let magnitude = if magnitude < self.modulus { + magnitude + } else { + magnitude % self.modulus + }; + return F::from(if value < 0 && magnitude != 0 { + self.modulus - magnitude + } else { + magnitude + }); + } + let mut reduced = constraint % &self.wide_modulus; if reduced.is_negative() { - reduced += &self.modulus; + reduced += &self.wide_modulus; } - Fq::from( + F::from( reduced .to_u128() .expect("a canonical residue always fits a u128"), ) } } + +#[cfg(test)] +mod tests { + use super::*; + use field::Q100; + use num_traits::{One, Pow}; + + type F = field::FqDefault; + + /// Every coefficient lands on its canonical residue, whichever path + /// reduces it. + #[test] + fn projection_is_canonical_reduction() { + let projection = >::prepare(); + let reference = |value: &BigInt| { + let modulus = BigInt::from(Q100); + let mut residue = value % &modulus; + if residue.is_negative() { + residue += &modulus; + } + F::from(residue.to_u128().unwrap()) + }; + let q = Q100 as i128; + let mut values: Vec = [ + 0, + 1, + -1, + 7, + -7, + q, + -q, + q + 1, + -q - 1, + 2 * q, + i128::MAX, + i128::MIN, + ] + .into_iter() + .map(BigInt::from) + .collect(); + let wide = BigInt::from(3).pow(130u32) + BigInt::one(); + values.extend([wide.clone(), -wide.clone(), &wide * &wide, -(&wide * &wide)]); + for value in &values { + let projected: F = projection.project(value); + assert_eq!(projected, reference(value), "{value}"); + } + } +} diff --git a/tooling/cli/tests/circuits.rs b/tooling/cli/tests/circuits.rs index 063c568..0569baf 100644 --- a/tooling/cli/tests/circuits.rs +++ b/tooling/cli/tests/circuits.rs @@ -1,5 +1,5 @@ use bitz_cli::{ - ProjectBigIntToFq, + ProjectBigIntToField, circuits::{BuiltinCircuit, CircuitInstance}, end_to_end::{CircuitProofSystem, CircuitStatement, Error}, }; @@ -10,11 +10,14 @@ use circuit::{ use num_traits::{Signed, ToPrimitive}; type R = num_bigint::BigInt; -type F = field::FqDefault; -type Proj = ProjectBigIntToFq; +type F = field::DynField; +type Proj = ProjectBigIntToField; +/// Every residual is below the narrowest prime a proof may draw, so no +/// Boolean assignment satisfies a row modulo the prime without satisfying it +/// over the integers. #[test] -fn sha_constraint_residuals_cannot_wrap_modulo_q100() { +fn sha_constraint_residuals_cannot_wrap_modulo_the_prime() { for circuit in BuiltinCircuit::ALL { let statement = CircuitInstance::random(circuit, None, None).unwrap(); let mut generator = ConstraintGenerator::::new(statement.input_bits()); @@ -45,7 +48,10 @@ fn sha_constraint_residuals_cannot_wrap_modulo_q100() { .unwrap() .checked_add(bound(c)) .unwrap(); - assert!(residual_bound < field::Q100, "{circuit}: residual may wrap"); + assert!( + residual_bound < 1u128 << (common::MIN_PRIME_BITS - 1), + "{circuit}: residual may wrap" + ); } } } @@ -55,7 +61,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::<_, F>::new::(statement).unwrap(); + let system = CircuitProofSystem::<_, F, R, Proj>::new(statement).unwrap(); let witness = system.witness(&inputs).unwrap(); let data = system.commit(&witness).unwrap(); let proof = system.prove(witness, data).unwrap(); @@ -78,17 +84,17 @@ fn sha_compression_matches_abc_and_binds_public_values() { inputs: inputs.clone(), output: bits(&ABC_DIGEST), }; - let system = CircuitProofSystem::<_, F>::new::(statement.clone()).unwrap(); + let system = CircuitProofSystem::<_, F, R, Proj>::new(statement.clone()).unwrap(); let witness = system.witness(&inputs).unwrap(); let data = system.commit(&witness).unwrap(); let proof = system.prove(witness, data).unwrap(); - CircuitProofSystem::<_, F>::new::(statement.clone()) + CircuitProofSystem::<_, F, R, Proj>::new(statement.clone()) .unwrap() .verify(&proof) .unwrap(); let mut changed = statement.clone(); changed.output[0] ^= true; - let wrong_output = CircuitProofSystem::<_, F>::new::(changed).unwrap(); + let wrong_output = CircuitProofSystem::<_, F, R, Proj>::new(changed).unwrap(); assert!(wrong_output.verify(&proof).is_err()); assert!(matches!( wrong_output.witness(&inputs), @@ -97,7 +103,7 @@ fn sha_compression_matches_abc_and_binds_public_values() { let mut changed = statement; changed.inputs[0] ^= true; assert!( - CircuitProofSystem::<_, F>::new::(changed) + CircuitProofSystem::<_, F, R, Proj>::new(changed) .unwrap() .verify(&proof) .is_err() @@ -114,7 +120,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::<_, F>::new::(statement).unwrap(); + let system = CircuitProofSystem::<_, F, R, Proj>::new(statement).unwrap(); system.witness(&inputs).unwrap(); } for (circuit, count) in [ @@ -131,5 +137,5 @@ fn variable_lengths_custom_state_and_invalid_dimensions() { let mut malformed = CircuitInstance::random(BuiltinCircuit::Sha256Compression, None, None).unwrap(); malformed.inputs.pop(); - assert!(CircuitProofSystem::<_, F>::new::(malformed).is_err()); + assert!(CircuitProofSystem::<_, F, R, Proj>::new(malformed).is_err()); } diff --git a/tooling/cli/tests/end_to_end.rs b/tooling/cli/tests/end_to_end.rs index f9d387c..48236bb 100644 --- a/tooling/cli/tests/end_to_end.rs +++ b/tooling/cli/tests/end_to_end.rs @@ -1,14 +1,14 @@ -use bitz_cli::ProjectBigIntToFq; +use bitz_cli::ProjectBigIntToField; 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; +type F = field::DynField; +type Proj = ProjectBigIntToField; fn rejects_changed_or_missing_ood( - system: &CircuitProofSystem, + system: &CircuitProofSystem, proof: &Proof, ) { // These Fast-profile fixtures have zero initial grinding bits, so the first @@ -51,13 +51,13 @@ impl CircuitStatement for PublicBit { #[test] fn generic_driver_accepts_a_non_sha_circuit() { - let prepared = CircuitProofSystem::<_, F>::new::(PublicBit).unwrap(); + let prepared = CircuitProofSystem::<_, F, R, Proj>::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::<_, F>::new::(PublicBit) + CircuitProofSystem::<_, F, R, Proj>::new(PublicBit) .unwrap() .verify(&proof) .unwrap(); @@ -123,14 +123,14 @@ impl CircuitStatement for PublicXor { #[test] fn nonidentity_map_uses_virtual_opening_and_checks_xor_relation() { - let system = CircuitProofSystem::<_, F>::new::(PublicXor).unwrap(); + let system = CircuitProofSystem::<_, F, R, Proj>::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::<_, F>::new::(PublicXor) + CircuitProofSystem::<_, F, R, Proj>::new(PublicXor) .unwrap() .verify(&proof) .unwrap(); diff --git a/tooling/cli/tests/sha256_spartan.rs b/tooling/cli/tests/sha256_spartan.rs index 82b29d1..2256273 100644 --- a/tooling/cli/tests/sha256_spartan.rs +++ b/tooling/cli/tests/sha256_spartan.rs @@ -1,4 +1,5 @@ -use bitz_cli::{ProjectBigIntToFq, ProjectConstraint}; +use bitz_cli::{ProjectBigIntToField, ProjectConstraint}; +use circuit::constraints::ConstraintMatrices; use circuit::{ constraints::ConstraintGenerator, sha256::{ @@ -16,7 +17,7 @@ use transcript::{build_prover, build_verifier}; type R = num_bigint::BigInt; type F = field::FqDefault; -type Proj = ProjectBigIntToFq; +type Proj = ProjectBigIntToField; const SESSION: &[u8] = b"spartan/piop/sha256-compression/v1"; const INSTANCE: &[u8] = b"abc-single-compression"; @@ -55,7 +56,8 @@ fn sha256_compression_verifies_through_spartan_piop() { assert_assignment_matches(&recomputed_assignment, &recorded_assignment); let projection = >::prepare(); - let matrices = integer_matrices.map_coefficients(|c| projection.project(&c)); + let matrices: ConstraintMatrices = + 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();