From 292fce804f8b75eb1c28d090913a3097dd320c68 Mon Sep 17 00:00:00 2001 From: jotabulacios Date: Fri, 27 Mar 2026 10:19:03 -0300 Subject: [PATCH 01/19] port logup gkr --- crypto/stark/src/constraints/evaluator.rs | 42 +- crypto/stark/src/debug.rs | 2 + crypto/stark/src/gkr.rs | 2663 +++++++++++++++++ crypto/stark/src/lagrange_kernel.rs | 316 ++ crypto/stark/src/lib.rs | 3 + crypto/stark/src/lookup.rs | 1428 +++++++-- crypto/stark/src/proof/stark.rs | 18 +- crypto/stark/src/prover.rs | 174 +- crypto/stark/src/sumcheck.rs | 828 +++++ .../src/tests/bus_tests/packing_tests.rs | 8 +- .../src/tests/bus_tests/soundness_tests.rs | 360 +-- crypto/stark/src/traits.rs | 9 +- crypto/stark/src/verifier.rs | 296 +- 13 files changed, 5478 insertions(+), 669 deletions(-) create mode 100644 crypto/stark/src/gkr.rs create mode 100644 crypto/stark/src/lagrange_kernel.rs create mode 100644 crypto/stark/src/sumcheck.rs diff --git a/crypto/stark/src/constraints/evaluator.rs b/crypto/stark/src/constraints/evaluator.rs index 20b2efaa1..c28437465 100644 --- a/crypto/stark/src/constraints/evaluator.rs +++ b/crypto/stark/src/constraints/evaluator.rs @@ -17,6 +17,8 @@ use rayon::{ }; use std::marker::PhantomData; +#[cfg(feature = "instruments")] +use std::time::Instant; pub struct ConstraintEvaluator< Field: IsSubFieldOf + IsFFTField + Send + Sync, @@ -244,6 +246,9 @@ where #[cfg(all(debug_assertions, not(feature = "parallel")))] let boundary_polys: Vec>> = Vec::new(); + #[cfg(feature = "instruments")] + let timer = Instant::now(); + let trace_length = domain.interpolation_domain_size; let lde_periodic_columns = air .get_periodic_column_polynomials(trace_length) @@ -259,6 +264,15 @@ where .collect::>>, FFTError>>() .unwrap(); + #[cfg(feature = "instruments")] + println!( + " Evaluating periodic columns on lde: {:#?}", + timer.elapsed() + ); + + #[cfg(feature = "instruments")] + let timer = Instant::now(); + // Fused boundary evaluation: compute (trace[col] - value) on-the-fly // instead of pre-computing all boundary_polys_evaluations. // This eliminates N_constraints × LDE_size intermediate allocations. @@ -288,6 +302,12 @@ where }) .collect(); + #[cfg(feature = "instruments")] + println!( + " Evaluated boundary polynomials on LDE: {:#?}", + timer.elapsed() + ); + #[cfg(all(debug_assertions, not(feature = "parallel")))] let boundary_zerofiers = Vec::new(); @@ -297,17 +317,27 @@ where #[cfg(all(debug_assertions, not(feature = "parallel")))] let _transition_evaluations: Vec> = Vec::new(); + #[cfg(feature = "instruments")] + let timer = Instant::now(); let zerofier_data = air.transition_zerofier_evaluations_grouped(domain); + #[cfg(feature = "instruments")] + println!( + " Evaluated transition zerofiers: {:#?}", + timer.elapsed() + ); // Iterate over all LDE domain and compute the part of the composition polynomial // related to the transition constraints and add it to the already computed part of the // boundary constraints. + #[cfg(feature = "instruments")] + let timer = Instant::now(); + let num_transition = air.num_transition_constraints(); let num_periodic = lde_periodic_columns.len(); let offsets = &air.context().transition_offsets; - Self::evaluate_transitions( + let evaluations_t = Self::evaluate_transitions( air, lde_trace, &lde_periodic_columns, @@ -319,6 +349,14 @@ where num_periodic, offsets, &self.logup_table_offset, - ) + ); + + #[cfg(feature = "instruments")] + println!( + " Evaluated transitions and accumulated results: {:#?}", + timer.elapsed() + ); + + evaluations_t } } diff --git a/crypto/stark/src/debug.rs b/crypto/stark/src/debug.rs index 906b8f317..644c97a64 100644 --- a/crypto/stark/src/debug.rs +++ b/crypto/stark/src/debug.rs @@ -96,6 +96,7 @@ pub fn validate_trace< .map(|(trace_steps, constraint)| trace_steps - constraint.end_exemptions()) .collect(); + // Pre-compute LogUp alpha powers once for all steps. let logup_alpha_powers: Vec> = if rap_challenges.len() > LOGUP_CHALLENGE_ALPHA { compute_alpha_powers( @@ -106,6 +107,7 @@ pub fn validate_trace< Vec::new() }; + // Compute logup_table_offset = table_contribution / trace_length let logup_table_offset = match bus_public_inputs { Some(bpi) => { let n_inv = FieldElement::::from(trace_length as u64) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs new file mode 100644 index 000000000..97aa83d4d --- /dev/null +++ b/crypto/stark/src/gkr.rs @@ -0,0 +1,2663 @@ +use crate::sumcheck::{RoundPoly, SumcheckProof}; +use core::fmt; +use crypto::fiat_shamir::is_transcript::IsTranscript; +use math::field::{element::FieldElement, traits::IsField}; +#[cfg(feature = "parallel")] +use rayon::prelude::*; + +// ============================================================================= +// Layer enum for gate-specialized GKR +// ============================================================================= + +/// A layer in the GKR binary-tree circuit. +/// +/// Each layer has half the size of the layer below it. The leaves (input layer) +/// are at the bottom, and the root (output layer) has a single element. +/// +/// Different gate types allow specialized inner loops: +/// - `LogUpGeneric`: explicit numerators and denominators (6 muls/pair in sumcheck) +/// - `LogUpSingles`: numerators are implicitly 1 (2 muls/pair — ~50% savings at leaf) +#[derive(Debug, Clone)] +pub enum Layer { + /// LogUp with explicit numerators and denominators. + LogUpGeneric { + numerators: Vec>, + denominators: Vec>, + }, + /// LogUp where all numerators are implicitly 1. + /// Saves ~50% muls in the sumcheck inner loop (no nl/nr tables needed). + LogUpSingles { denominators: Vec> }, +} + +impl Layer { + /// Number of variables: log2 of the layer size. + pub fn n_variables(&self) -> usize { + let len = match self { + Self::LogUpGeneric { denominators, .. } => denominators.len(), + Self::LogUpSingles { denominators } => denominators.len(), + }; + debug_assert!(len.is_power_of_two()); + len.trailing_zeros() as usize + } + + /// Whether this is the root layer (single value, 0 variables). + pub fn is_output_layer(&self) -> bool { + self.n_variables() == 0 + } + + /// Returns the root (numerator, denominator) for a single-element output layer. + pub fn try_into_output_values(&self) -> Option<(FieldElement, FieldElement)> { + if !self.is_output_layer() { + return None; + } + Some(match self { + Layer::LogUpGeneric { + numerators, + denominators, + } => (numerators[0].clone(), denominators[0].clone()), + Layer::LogUpSingles { denominators } => (FieldElement::one(), denominators[0].clone()), + }) + } + + /// Computes the next (parent) layer by pairwise fraction addition. + /// Returns `None` if already at the output layer. + /// + /// Both Singles and Generic produce Generic output (since 1/a + 1/b = (a+b)/(a*b) + /// requires an explicit numerator). + pub fn next_layer(&self) -> Option { + if self.is_output_layer() { + return None; + } + Some(match self { + Self::LogUpGeneric { + numerators, + denominators, + } => next_logup_layer(Some(numerators), denominators), + Self::LogUpSingles { denominators } => next_logup_layer(None, denominators), + }) + } +} + +/// Pairwise fraction addition for LogUp layers. +/// +/// If `numerators` is `None`, all numerators are implicitly 1 (singles case): +/// 1/d[2j] + 1/d[2j+1] = (d[2j+1] + d[2j]) / (d[2j] * d[2j+1]) +/// +/// Otherwise: n[2j]/d[2j] + n[2j+1]/d[2j+1] = cross-multiply. +fn next_logup_layer( + numerators: Option<&[FieldElement]>, + denominators: &[FieldElement], +) -> Layer { + let half_n = denominators.len() / 2; + let mut next_numerators = Vec::with_capacity(half_n); + let mut next_denominators = Vec::with_capacity(half_n); + + for j in 0..half_n { + let dl = &denominators[2 * j]; + let dr = &denominators[2 * j + 1]; + let (num, den) = match numerators { + Some(nums) => { + let nl = &nums[2 * j]; + let nr = &nums[2 * j + 1]; + // nl/dl + nr/dr = (nl*dr + nr*dl) / (dl*dr) + (&(nl * dr) + &(nr * dl), dl * dr) + } + None => { + // 1/dl + 1/dr = (dr + dl) / (dl*dr) + (dl + dr, dl * dr) + } + }; + next_numerators.push(num); + next_denominators.push(den); + } + + Layer::LogUpGeneric { + numerators: next_numerators, + denominators: next_denominators, + } +} + +/// Generates all layers from the input (leaves) to the output (root). +/// +/// Returns layers[0] = input (leaves), layers[last] = output (root, 1 element). +pub fn gen_layers(input_layer: Layer) -> Vec> { + let n_variables = input_layer.n_variables(); + let mut layers = vec![input_layer]; + while let Some(next) = layers.last().unwrap().next_layer() { + layers.push(next); + } + assert_eq!(layers.len(), n_variables + 1); + layers +} + +// ============================================================================= +// Batch GKR proof types +// ============================================================================= + +/// Proof for a single layer in a batch GKR reduction. +/// +/// Contains the shared sumcheck proof and per-instance child claims (masks). +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct BatchGkrLayerProof { + /// Shared sumcheck proof for this layer (combined across all active instances). + pub sumcheck_proof: SumcheckProof, + /// Per-active-instance child claims: [n_left, n_right, d_left, d_right]. + /// Order matches the order instances became active. + pub child_claims_by_instance: Vec<[FieldElement; 4]>, +} + +/// Complete batch GKR proof for multiple fractional summation trees. +/// +/// All instances share one sumcheck per layer via random linear combination. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct BatchGkrProof { + /// Per-instance root claims as (numerator, denominator) pairs. + /// The claimed sum for instance i is root_claims[i].0 / root_claims[i].1. + pub root_claims: Vec<(FieldElement, FieldElement)>, + /// One layer proof per reduction step (from root towards leaves). + /// The number of layers equals max(n_variables) across all instances. + pub layer_proofs: Vec>, +} + +/// A rational number `numerator / denominator` in a field. +/// +/// Used in the GKR protocol to represent LogUp contributions before they are +/// reduced to a single field element via batch inversion. Fraction addition +/// uses cross-multiplication to avoid per-addition inversions. +pub struct Fraction { + pub numerator: FieldElement, + pub denominator: FieldElement, +} + +impl Fraction { + /// Create a new fraction `numerator / denominator`. + pub fn new(numerator: FieldElement, denominator: FieldElement) -> Self { + Self { + numerator, + denominator, + } + } + + /// Add two fractions via cross-multiplication: + /// a/b + c/d = (a*d + c*b) / (b*d) + /// + /// No reduction or normalization is performed. + pub fn add(&self, other: &Fraction) -> Fraction { + let numerator = + &(&self.numerator * &other.denominator) + &(&other.numerator * &self.denominator); + let denominator = &self.denominator * &other.denominator; + Fraction { + numerator, + denominator, + } + } +} + +/// One layer of the summation tree used in the GKR protocol. +/// +/// Each layer stores parallel arrays of numerators and denominators representing +/// fractions at that level. The leaf layer has N fractions; each subsequent layer +/// halves the count by pairwise fraction addition, until the root layer has 1. +pub struct SummationLayer { + pub numerators: Vec>, + pub denominators: Vec>, +} + +/// Build a summation tree from leaf fractions. +/// +/// Takes N leaf fractions (as parallel numerator/denominator vectors) and +/// returns layers from leaves (index 0, size N) to root (last index, size 1). +/// Each layer is built by pairwise fraction addition: +/// parent_n = left_n * right_d + right_n * left_d +/// parent_d = left_d * right_d +/// +/// # Panics +/// Panics if `leaf_numerators` and `leaf_denominators` have different lengths, +/// or if the length is not a power of 2. +pub fn build_summation_tree( + leaf_numerators: Vec>, + leaf_denominators: Vec>, +) -> Vec> { + let n = leaf_numerators.len(); + assert_eq!( + n, + leaf_denominators.len(), + "numerators and denominators must have the same length" + ); + assert!(n.is_power_of_two(), "number of leaves must be a power of 2"); + + // Number of layers: log2(n) + 1 (leaves + intermediate + root) + let num_layers = n.trailing_zeros() as usize + 1; + let mut layers = Vec::with_capacity(num_layers); + + // Layer 0: the leaves themselves + layers.push(SummationLayer { + numerators: leaf_numerators, + denominators: leaf_denominators, + }); + + // Build each subsequent layer by pairwise fraction addition + for layer_idx in 1..num_layers { + let prev = &layers[layer_idx - 1]; + let prev_len = prev.numerators.len(); + let new_len = prev_len / 2; + + let compute_pair = |i: usize| -> (FieldElement, FieldElement) { + let left_n = &prev.numerators[2 * i]; + let left_d = &prev.denominators[2 * i]; + let right_n = &prev.numerators[2 * i + 1]; + let right_d = &prev.denominators[2 * i + 1]; + + // Cross-multiply: (left_n * right_d + right_n * left_d) / (left_d * right_d) + let parent_n = &(left_n * right_d) + &(right_n * left_d); + let parent_d = left_d * right_d; + (parent_n, parent_d) + }; + + #[cfg(feature = "parallel")] + let (numerators, denominators): (Vec<_>, Vec<_>) = if new_len >= 256 { + (0..new_len).into_par_iter().map(compute_pair).unzip() + } else { + (0..new_len).map(compute_pair).unzip() + }; + + #[cfg(not(feature = "parallel"))] + let (numerators, denominators): (Vec<_>, Vec<_>) = (0..new_len).map(compute_pair).unzip(); + + layers.push(SummationLayer { + numerators, + denominators, + }); + } + + layers +} + +/// Proof for a single GKR layer reduction. +/// +/// Contains the sumcheck proof that reduces claims about a parent layer's MLEs +/// to claims about the children layer's MLEs, plus the four claimed evaluations +/// of the children's numerator/denominator MLEs at the reduced point. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct GkrLayerProof { + /// Sumcheck proof for the layer reduction (degree-3 round polynomials). + pub sumcheck_proof: SumcheckProof, + /// Claimed evaluations at the children layer: [n_left, n_right, d_left, d_right]. + /// These are the children MLE values at the point (r', 0) and (r', 1) where + /// r' is the sumcheck challenge point and the last coordinate selects left/right. + pub child_claims: [FieldElement; 4], +} + +/// Complete GKR proof for a fractional summation tree. +/// +/// Proves that the root of the summation tree has a specific value by +/// layer-by-layer reduction from the root to the leaves via fractional sumcheck. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct GkrProof { + /// The claimed sum at the root: numerator / denominator as a field element. + pub claimed_sum: FieldElement, + /// One layer proof per reduction step (from root towards leaves). + pub layer_proofs: Vec>, +} + +/// Compute the equality polynomial evaluations eq(point, b) for all b in {0,1}^n. +/// +/// The equality polynomial is defined as: +/// eq(x, y) = prod_{i=0}^{n-1} (x_i * y_i + (1 - x_i) * (1 - y_i)) +/// +/// For a fixed point r = (r_0, ..., r_{n-1}), this computes eq(r, b) for every +/// Boolean vector b, yielding 2^n values. Uses the standard butterfly/tensor +/// product construction: +/// - Start with [1] +/// - For each coordinate r_i, double the table: +/// existing entries are scaled by (1 - r_i), and new entries by r_i +/// +/// Returns a vector of length 2^n in little-endian bit order. +pub fn compute_eq_evals(point: &[FieldElement]) -> Vec> { + let n = point.len(); + let size = 1 << n; + let mut evals = Vec::with_capacity(size); + evals.push(FieldElement::one()); + + for r_i in point.iter() { + let one_minus_ri = &FieldElement::::one() - r_i; + let prev_len = evals.len(); + // Extend: for each existing entry e, push e * r_i, then scale e by (1 - r_i) + for j in 0..prev_len { + let new_val = &evals[j] * r_i; + evals.push(new_val); + } + // Scale existing entries by (1 - r_i) + for eval in evals[..prev_len].iter_mut() { + *eval = &*eval * &one_minus_ri; + } + } + + evals +} + +/// Evaluate the MLE (multilinear extension) of a table at a given point. +/// +/// Given evaluations `table[b]` for b in {0,1}^n, computes: +/// MLE(point) = sum_{b in {0,1}^n} table[b] * eq(point, b) +/// +/// This is equivalent to multilinear interpolation of the table at the point. +#[cfg(test)] +fn evaluate_mle( + table: &[FieldElement], + point: &[FieldElement], +) -> FieldElement { + let eq_evals = compute_eq_evals(point); + assert_eq!(table.len(), eq_evals.len()); + table + .iter() + .zip(eq_evals.iter()) + .fold(FieldElement::zero(), |acc, (t, e)| &acc + &(t * e)) +} + +/// Run the GKR prover on a fractional summation tree. +/// +/// Proves that the root of the tree has a specific numerator/denominator by +/// layer-by-layer reduction from the root to the leaves. At each layer, a +/// fractional sumcheck proves that the parent layer's claims are consistent +/// with the children layer via the gate equation: +/// parent_n(j) = child_n(2j) * child_d(2j+1) + child_n(2j+1) * child_d(2j) +/// parent_d(j) = child_d(2j) * child_d(2j+1) +/// +/// The sumcheck at each layer operates on a degree-3 function (product of eq +/// weight with two child values), so round polynomials have degree 3 with +/// 4 evaluation points each. +/// +/// # Arguments +/// - `tree`: the summation tree layers from leaves (index 0) to root (last index) +/// - `transcript`: Fiat-Shamir transcript for challenge sampling +/// +/// # Returns +/// A `GkrProof` containing the claimed sum and layer proofs, plus the final +/// evaluation point and claims at the leaf layer. +/// +/// # Panics +/// Panics if the tree is empty or has inconsistent layer sizes. +pub fn gkr_prove( + tree: &[SummationLayer], + transcript: &mut impl IsTranscript, +) -> ( + GkrProof, + Vec>, + FieldElement, + FieldElement, +) { + assert!(!tree.is_empty(), "tree must have at least one layer"); + + let num_layers = tree.len(); // layers 0..num_layers-1, root is at num_layers-1 + let root = &tree[num_layers - 1]; + assert_eq!( + root.numerators.len(), + 1, + "root layer must have exactly 1 element" + ); + + let root_n = &root.numerators[0]; + let root_d = &root.denominators[0]; + + // Compute claimed_sum = root_n / root_d + let root_d_inv = root_d.inv().expect("root denominator must be nonzero"); + let claimed_sum = root_n * &root_d_inv; + + // Append the claimed sum to the transcript + transcript.append_field_element(&claimed_sum); + + // If the tree has only 1 layer (just leaves = root), no reductions needed + if num_layers == 1 { + return ( + GkrProof { + claimed_sum, + layer_proofs: vec![], + }, + vec![], + root_n.clone(), + root_d.clone(), + ); + } + + let mut layer_proofs = Vec::with_capacity(num_layers - 1); + let mut n_claim = root_n.clone(); + let mut d_claim = root_d.clone(); + let mut current_point: Vec> = vec![]; + + // Reduce from root towards leaves: for each layer l from (num_layers-2) down to 0, + // the children layer is tree[l] and the parent is tree[l+1]. + for l in (0..num_layers - 1).rev() { + let child_n = &tree[l].numerators; + let child_d = &tree[l].denominators; + let parent_size = child_n.len() / 2; // = tree[l+1].numerators.len() + let parent_num_vars = parent_size.trailing_zeros() as usize; + + // Sample lambda to combine numerator and denominator claims + let lambda: FieldElement = transcript.sample_field_element(); + let combined_claim = &n_claim + &(&lambda * &d_claim); + + if parent_num_vars == 0 { + // Trivial case: parent has 1 element (0 variables), no sumcheck needed. + // The "sumcheck" is just a direct check: combined_claim must equal + // child_n[0]*child_d[1] + child_n[1]*child_d[0] + lambda * child_d[0]*child_d[1] + let nl = &child_n[0]; + let nr = &child_n[1]; + let dl = &child_d[0]; + let dr = &child_d[1]; + + // Provide the 4 child claims (they are just the raw values since there's + // no random point yet) + let child_claims = [nl.clone(), nr.clone(), dl.clone(), dr.clone()]; + + // Append child claims to transcript + for claim in &child_claims { + transcript.append_field_element(claim); + } + + // Sample eta to combine left/right into new claims for next layer + let eta: FieldElement = transcript.sample_field_element(); + + // New claims for the next layer down: fold left and right using eta + // The children MLE at point (eta) for a 2-element table [a, b] is: + // a*(1-eta) + b*eta + n_claim = &child_n[0] + &(&eta * &(&child_n[1] - &child_n[0])); + d_claim = &child_d[0] + &(&eta * &(&child_d[1] - &child_d[0])); + current_point = vec![eta]; + + layer_proofs.push(GkrLayerProof { + sumcheck_proof: SumcheckProof { + round_polys: vec![], + }, + child_claims, + }); + } else { + // Non-trivial case: run the fractional sumcheck over parent_num_vars variables. + // + // The function to sum over {0,1}^parent_num_vars is: + // f(b) = eq(current_point, b) * [n_left(b)*d_right(b) + n_right(b)*d_left(b) + lambda*d_left(b)*d_right(b)] + // where left/right are the even/odd children indexed by b. + // + // This function has degree 3 in each variable (product of eq * two child values), + // so the round polynomial needs 4 evaluation points (degree 3). + + // Build the five "bookkeeping" tables for the sumcheck: + // - eq_table: eq(current_point, b) for b in {0,1}^parent_num_vars + // - nl_table: child_n[2b] (left numerators) + // - nr_table: child_n[2b+1] (right numerators) + // - dl_table: child_d[2b] (left denominators) + // - dr_table: child_d[2b+1] (right denominators) + let mut eq_table = compute_eq_evals(¤t_point); + let mut nl_table: Vec> = + (0..parent_size).map(|j| child_n[2 * j].clone()).collect(); + let mut nr_table: Vec> = (0..parent_size) + .map(|j| child_n[2 * j + 1].clone()) + .collect(); + let mut dl_table: Vec> = + (0..parent_size).map(|j| child_d[2 * j].clone()).collect(); + let mut dr_table: Vec> = (0..parent_size) + .map(|j| child_d[2 * j + 1].clone()) + .collect(); + + // Verify initial consistency: the sum of f over {0,1}^parent_num_vars + // should equal combined_claim + debug_assert!({ + let mut check_sum = FieldElement::::zero(); + for j in 0..parent_size { + let gate_val = &(&nl_table[j] * &dr_table[j]) + + &(&nr_table[j] * &dl_table[j]) + + &(&lambda * &(&dl_table[j] * &dr_table[j])); + check_sum = &check_sum + &(&eq_table[j] * &gate_val); + } + check_sum == combined_claim + }); + + let mut round_polys = Vec::with_capacity(parent_num_vars); + let mut challenges = Vec::with_capacity(parent_num_vars); + let mut round_combined_claim = combined_claim.clone(); + // Eq correction factor: accumulates eq(r_k, c_k) from previous rounds. + // Instead of folding eq_table (N/2 multiplications per round), we halve + // it (N/2 additions) and track the missing fold factors in this scalar. + let mut eq_correction = FieldElement::::one(); + + for r_round in current_point.iter().take(parent_num_vars) { + let half = nl_table.len() / 2; + + // Eq polynomial factoring (Dao-Thaler, ePrint 2024/1210): + // + // Factor eq(current_point, b) into: + // eq_correction (scalar from previous rounds' fold factors) * + // eq_round(r, t) (linear in t, constant across pairs) * + // eq_table[j] (per-pair scalar, independent of t) + // + // where r = current_point[round_idx] and + // eq_table[j] = eq_orig[2j] + eq_orig[2j+1] (pre-halved) + // + // The inner sum h_raw(t) = Σ_j eq_table[j] * gate(t, j) is degree 2, + // needing only 2 eval points (t=0, t=2) per pair instead of 3. + // Then h(t) = eq_correction * h_raw(t), and + // S(t) = eq_round(r, t) * h(t) recovers the degree-3 round poly. + let one = FieldElement::::one(); + + // Pre-halve eq_table: compute eq_rem[j] = eq_table[2j] + eq_table[2j+1] + // in-place. This replaces fold (which would multiply each entry by the + // challenge) with a simple sum. The fold factor is tracked in eq_correction. + for j in 0..half { + eq_table[j] = &eq_table[2 * j] + &eq_table[2 * j + 1]; + } + eq_table.truncate(half); + + let compute_pair_sums = |j: usize| -> [FieldElement; 2] { + // Pre-halved eq weight (scalar, independent of t) + let eq_rem = &eq_table[j]; + + let nl_l = &nl_table[2 * j]; + let nl_r = &nl_table[2 * j + 1]; + let nr_l = &nr_table[2 * j]; + let nr_r = &nr_table[2 * j + 1]; + let dl_l = &dl_table[2 * j]; + let dl_r = &dl_table[2 * j + 1]; + let dr_l = &dr_table[2 * j]; + let dr_r = &dr_table[2 * j + 1]; + + // t=0: interpolated values are just the left values + // gate = nl*dr + nr*dl + lambda*dl*dr + // = nl*dr + dl*(nr + lambda*dr) [3 muls instead of 4] + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let h0 = eq_rem * &gate_0; + + // t=2: val = 2*right - left + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + let h2 = eq_rem * &gate_2; + + [h0, h2] + }; + + let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; + let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { + [&a[0] + &b[0], &a[1] + &b[1]] + }; + + #[cfg(feature = "parallel")] + let totals: [FieldElement; 2] = if half >= 256 { + (0..half) + .into_par_iter() + .fold(zero2, |acc, j| add2(acc, compute_pair_sums(j))) + .reduce(zero2, add2) + } else { + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))) + }; + + #[cfg(not(feature = "parallel"))] + let totals: [FieldElement; 2] = + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))); + + // Phase 2: Recover S(t) from h(t) and eq_round(r, t). + // + // The inner sum h_raw(t) doesn't include the eq_correction factor + // (accumulated from previous rounds' fold factors). Apply it now: + // h(t) = eq_correction * h_raw(t) + // + // eq_round(r, t) = (1-r)(1-t) + r*t + // eq_round(r, 0) = 1-r, eq_round(r, 1) = r, + // eq_round(r, 2) = 3r-1, eq_round(r, 3) = 5r-2 + // + // S(0) = (1-r)*h(0), S(1) = round_combined_claim - S(0) + // h(1) = S(1)/r, h(3) = 3*h(2) - 3*h(1) + h(0) (degree-2 extrapolation) + // S(2) = (3r-1)*h(2), S(3) = (5r-2)*h(3) + let [raw_h0, raw_h2] = totals; + let total_h0 = &eq_correction * &raw_h0; + let total_h2 = &eq_correction * &raw_h2; + + let one_minus_r = &one - r_round; + let s0 = &one_minus_r * &total_h0; + let s1 = &round_combined_claim - &s0; + + let r_inv = r_round + .inv() + .expect("r_round = 0 is probability 2^{-64} for random challenges"); + let h1 = &s1 * &r_inv; + + let three = FieldElement::::from(3u64); + let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; + + let eq_at_2 = &(&three * r_round) - &one; + let s2 = &eq_at_2 * &total_h2; + + let eq_at_3 = + &(&FieldElement::::from(5u64) * r_round) - &FieldElement::::from(2u64); + let s3 = &eq_at_3 * &h3; + + let poly_evals = vec![s0, s1, s2, s3]; + + let round_poly = RoundPoly::new(poly_evals); + + // Append round polynomial evaluations to transcript + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + // Sample challenge for this round + let challenge: FieldElement = transcript.sample_field_element(); + + // Update round_combined_claim for next round + round_combined_claim = round_poly.evaluate(&challenge); + + // Update eq_correction: accumulate eq(r_round, challenge). + // eq(r, c) = r*c + (1-r)*(1-c), the fold factor we skip for eq_table. + let eq_update = &(r_round * &challenge) + &(&one_minus_r * &(&one - &challenge)); + eq_correction = &eq_correction * &eq_update; + + // Fold the four gate tables (eq_table was already halved before inner loop) + // table[j] = table[2j] + challenge * (table[2j+1] - table[2j]) + let fold_table = |table: &mut Vec>| { + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + }; + + #[cfg(feature = "parallel")] + { + if half >= 256 { + // Fold tables in parallel (each table independently) + rayon::join( + || { + rayon::join( + || fold_table(&mut nl_table), + || fold_table(&mut nr_table), + ); + }, + || { + rayon::join( + || fold_table(&mut dl_table), + || fold_table(&mut dr_table), + ); + }, + ); + } else { + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); + } + } + + #[cfg(not(feature = "parallel"))] + { + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); + } + + round_polys.push(round_poly); + challenges.push(challenge); + } + + // After all rounds, each table has a single entry: the MLE evaluated + // at the sumcheck challenge point. + let child_claims = [ + nl_table[0].clone(), + nr_table[0].clone(), + dl_table[0].clone(), + dr_table[0].clone(), + ]; + + // Append child claims to transcript + for claim in &child_claims { + transcript.append_field_element(claim); + } + + // Sample eta to fold left/right children into new claims + let eta: FieldElement = transcript.sample_field_element(); + + // New claims: fold left and right using eta + // n_claim = nl + eta*(nr - nl), d_claim = dl + eta*(dr - dl) + n_claim = &child_claims[0] + &(&eta * &(&child_claims[1] - &child_claims[0])); + d_claim = &child_claims[2] + &(&eta * &(&child_claims[3] - &child_claims[2])); + + // Update current_point: eta corresponds to x_0 (the even/odd selector), + // followed by the sumcheck challenges for the remaining parent variables. + // In little-endian convention: point[i] = value of x_i. + let mut new_point = Vec::with_capacity(challenges.len() + 1); + new_point.push(eta); + new_point.extend(challenges); + current_point = new_point; + + layer_proofs.push(GkrLayerProof { + sumcheck_proof: SumcheckProof { round_polys }, + child_claims, + }); + } + } + + ( + GkrProof { + claimed_sum, + layer_proofs, + }, + current_point, + n_claim, + d_claim, + ) +} + +/// Errors that can occur during GKR verification. +#[derive(Debug, Clone)] +pub enum GkrError { + /// The summation tree structure is invalid. + InvalidTree { reason: String }, + /// A sumcheck round failed verification. + SumcheckFailed { layer: usize, reason: String }, + /// The gate equation check failed at a layer. + GateCheckFailed { layer: usize }, + /// The claimed sum does not match (unused in the verifier itself, + /// but available for callers that compare the claimed sum to an external value). + ClaimedSumMismatch, +} + +impl fmt::Display for GkrError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + GkrError::InvalidTree { reason } => { + write!(f, "invalid GKR tree: {}", reason) + } + GkrError::SumcheckFailed { layer, reason } => { + write!(f, "sumcheck failed at layer {}: {}", layer, reason) + } + GkrError::GateCheckFailed { layer } => { + write!(f, "gate check failed at layer {}", layer) + } + GkrError::ClaimedSumMismatch => { + write!(f, "claimed sum mismatch") + } + } + } +} + +/// Verify a GKR proof for a fractional summation tree. +/// +/// Replays the Fiat-Shamir transcript identically to `gkr_prove` and checks: +/// - At each non-trivial layer: sumcheck round consistency (p(0)+p(1) = current_sum) +/// and the gate equation at the final evaluation point. +/// - At trivial layers (0-variable parent): no sumcheck check; soundness is +/// enforced by subsequent layers. +/// +/// # Arguments +/// - `proof`: the GKR proof produced by `gkr_prove` +/// - `transcript`: Fiat-Shamir transcript (must use the same seed as the prover) +/// +/// # Returns +/// `Ok((final_point, n_claim, d_claim))` where `final_point` is the random +/// evaluation point at the leaf layer, and `n_claim`/`d_claim` are the claimed +/// MLE evaluations of the leaf numerators/denominators at that point. +#[allow(clippy::type_complexity)] +pub fn gkr_verify( + proof: &GkrProof, + transcript: &mut impl IsTranscript, +) -> Result<(Vec>, FieldElement, FieldElement), GkrError> { + // Step 1: Append claimed_sum to transcript (mirrors prover line 239) + transcript.append_field_element(&proof.claimed_sum); + + // If there are no layer proofs, the tree had a single leaf (root = leaf). + // Return empty point and the claimed_sum as n_claim with d_claim = 1. + if proof.layer_proofs.is_empty() { + return Ok((vec![], proof.claimed_sum.clone(), FieldElement::one())); + } + + // Step 2: Initialize claims. + // The verifier sets n_claim = claimed_sum, d_claim = 1. + // This represents the same rational value as root_n/root_d. + // Soundness of the first (trivial) layer is enforced by later sumcheck layers. + let mut n_claim = proof.claimed_sum.clone(); + let mut d_claim = FieldElement::::one(); + let mut current_point: Vec> = vec![]; + + for (layer_idx, layer_proof) in proof.layer_proofs.iter().enumerate() { + // Step 4: Sample lambda (mirrors prover line 268) + let lambda: FieldElement = transcript.sample_field_element(); + + // Step 5: combined_claim = n_claim + lambda * d_claim + let combined_claim = &n_claim + &(&lambda * &d_claim); + + let round_polys = &layer_proof.sumcheck_proof.round_polys; + + if round_polys.is_empty() { + // Trivial layer (0 variables in parent): no sumcheck rounds. + // The prover just provides child_claims directly. + // No gate check here --- soundness is enforced by later layers' sumchecks. + // (The verifier's combined_claim may differ from the prover's by a + // scaling factor at the first trivial layer, so we skip the check.) + } else { + // Non-trivial layer: verify sumcheck inline. + let num_rounds = round_polys.len(); + let mut current_sum = combined_claim; + let mut challenges = Vec::with_capacity(num_rounds); + + for (round, round_poly) in round_polys.iter().enumerate() { + // Check p(0) + p(1) == current_sum + if round_poly.sum_at_binary() != current_sum { + return Err(GkrError::SumcheckFailed { + layer: layer_idx, + reason: format!("round {} sum mismatch: p(0)+p(1) != expected sum", round), + }); + } + + // Append round poly evals to transcript (mirrors prover lines 388-389) + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + // Sample challenge (mirrors prover line 393) + let challenge: FieldElement = transcript.sample_field_element(); + + // Update current_sum to p(challenge) + current_sum = round_poly.evaluate(&challenge); + + challenges.push(challenge); + } + + // Gate check: verify that the final sumcheck evaluation equals + // eq(current_point, challenges) * gate(child_claims, lambda) + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims; + + // Compute eq(current_point, challenges) as a single field element. + // eq(a, b) = prod_i (a_i*b_i + (1-a_i)*(1-b_i)) + let eq_val = compute_eq_at_point(¤t_point, &challenges); + + // gate_combined = nl*dr + nr*dl + lambda*dl*dr + let gate_combined = &(&(nl * dr) + &(nr * dl)) + &(&lambda * &(dl * dr)); + + let expected = &eq_val * &gate_combined; + + if current_sum != expected { + return Err(GkrError::GateCheckFailed { layer: layer_idx }); + } + + // Build the new current_point from eta (below) and sumcheck challenges. + // We need to store challenges for constructing current_point after eta is sampled. + // Store them temporarily. + // (We'll construct the point after sampling eta below.) + + // Append child claims to transcript (mirrors prover lines 431-433) + for claim in &layer_proof.child_claims { + transcript.append_field_element(claim); + } + + // Sample eta (mirrors prover line 436) + let eta: FieldElement = transcript.sample_field_element(); + + // Update claims: fold left/right using eta + n_claim = nl + &(&eta * &(nr - nl)); + d_claim = dl + &(&eta * &(dr - dl)); + + // Update current_point = [eta] ++ challenges (mirrors prover lines 447-450) + let mut new_point = Vec::with_capacity(challenges.len() + 1); + new_point.push(eta); + new_point.extend(challenges); + current_point = new_point; + + continue; + } + + // Trivial layer path: append child claims and sample eta + // (mirrors prover lines 285-290) + for claim in &layer_proof.child_claims { + transcript.append_field_element(claim); + } + let eta: FieldElement = transcript.sample_field_element(); + + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims; + n_claim = nl + &(&eta * &(nr - nl)); + d_claim = dl + &(&eta * &(dr - dl)); + current_point = vec![eta.clone()]; + } + + Ok((current_point, n_claim, d_claim)) +} + +// ============================================================================= +// Batch GKR prover +// ============================================================================= + +/// Run the batch GKR prover on multiple instances (each a Vec). +/// +/// All instances share one sumcheck per layer via random linear combination +/// (sumcheck_alpha). Instances can have different numbers of layers (different +/// trace lengths). Smaller instances start participating at later layers. +/// +/// # Arguments +/// - `layers_per_instance`: for each instance, its layer tree from leaves (index 0) +/// to root (last index), as produced by `gen_layers`. +/// - `transcript`: Fiat-Shamir transcript for challenge sampling +/// +/// # Returns +/// A `BatchGkrProof`, the shared `random_point`, and per-instance `(n_claim, d_claim)`. +#[allow(clippy::type_complexity)] +pub fn gkr_prove_batch( + layers_per_instance: Vec>>, + transcript: &mut impl IsTranscript, +) -> ( + BatchGkrProof, + Vec>, + Vec<(FieldElement, FieldElement)>, +) { + let n_instances = layers_per_instance.len(); + + // Domain separation + transcript.append_bytes(b"gkr_batch"); + transcript.append_bytes(&(n_instances as u64).to_le_bytes()); + + if n_instances == 0 { + return ( + BatchGkrProof { + root_claims: vec![], + layer_proofs: vec![], + }, + vec![], + vec![], + ); + } + + // n_layers_by_instance[i] = number of tree layers - 1 = number of reduction steps + let n_layers_by_instance: Vec = layers_per_instance + .iter() + .map(|layers| layers.len() - 1) + .collect(); + let max_layers = *n_layers_by_instance.iter().max().unwrap(); + + // Extract root (numerator, denominator) for each instance + let root_claims: Vec<(FieldElement, FieldElement)> = layers_per_instance + .iter() + .map(|layers| { + let root = layers.last().unwrap(); + root.try_into_output_values() + .expect("root layer must be output") + }) + .collect(); + + // Track per-instance state + // n_claim[i], d_claim[i] start as root values and get updated each layer + let mut n_claims: Vec>> = vec![None; n_instances]; + let mut d_claims: Vec>> = vec![None; n_instances]; + + let mut current_point: Vec> = vec![]; + let mut layer_proofs = Vec::with_capacity(max_layers); + + for layer in 0..max_layers { + let n_remaining = max_layers - layer; + + // Detect output layers: instances whose tree has exactly n_remaining reduction steps + // start participating now (their root is the output layer). + // Use actual (root_n, root_d) so the gate check at the end of the sumcheck + // matches: gate(children) = nl*dr + nr*dl + lambda*dl*dr uses the true root values. + for i in 0..n_instances { + if n_layers_by_instance[i] == n_remaining { + n_claims[i] = Some(root_claims[i].0.clone()); + d_claims[i] = Some(root_claims[i].1.clone()); + } + } + + // Append active claims to transcript (matches verifier) + for i in 0..n_instances { + if let (Some(n), Some(d)) = (&n_claims[i], &d_claims[i]) { + transcript.append_field_element(n); + transcript.append_field_element(d); + } + } + + // Sample randomness + let sumcheck_alpha: FieldElement = transcript.sample_field_element(); + let lambda: FieldElement = transcript.sample_field_element(); + + // Collect active instances (those that have claims and have reduction steps remaining) + let mut active_instances: Vec = Vec::new(); + let mut combined_claims: Vec> = Vec::new(); + + for i in 0..n_instances { + if n_claims[i].is_some() && n_layers_by_instance[i] > 0 { + active_instances.push(i); + let n = n_claims[i].as_ref().unwrap(); + let d = d_claims[i].as_ref().unwrap(); + let claim = n + &(&lambda * d); + + // Apply doubling factor for instances with fewer layers + let n_unused = max_layers - n_layers_by_instance[i]; + if n_unused > 0 { + let doubling = FieldElement::::from(1u64 << n_unused); + combined_claims.push(&claim * &doubling); + } else { + combined_claims.push(claim); + } + } + } + + if active_instances.is_empty() { + break; + } + + // The child layer index for instance i: layers[tree_layer_idx] where + // tree_layer_idx = n_layers_by_instance[i] - 1 - (layer - n_unused) + // Equivalently: layers[n_remaining - 1 - n_unused] but let's compute directly. + + // Compute the max parent_num_vars across active instances + // (this determines the sumcheck dimension for this batch layer) + let parent_num_vars_by_instance: Vec = active_instances + .iter() + .map(|&i| { + let tree_layer_idx = n_remaining - 1; + // tree_layer_idx is the child layer; the parent has half the size + // child has 2^k elements → parent has 2^{k-1} → parent_num_vars = k-1 + let child_n_vars = layers_per_instance[i][tree_layer_idx].n_variables(); + debug_assert!( + child_n_vars >= 1, + "child of a non-output layer must have >= 2 elements" + ); + child_n_vars - 1 + }) + .collect(); + let max_parent_vars = *parent_num_vars_by_instance.iter().max().unwrap(); + + if max_parent_vars == 0 { + // All active instances have trivial layers (0 variables in parent). + // No sumcheck needed — provide child claims directly. + let mut child_claims_by_instance = Vec::new(); + + for (idx, &i) in active_instances.iter().enumerate() { + let tree_layer_idx = n_remaining - 1; + let child = &layers_per_instance[i][tree_layer_idx]; + + let (child_n, child_d) = match child { + Layer::LogUpGeneric { + numerators, + denominators, + } => (numerators.as_slice(), denominators.as_slice()), + Layer::LogUpSingles { denominators } => { + // For singles at size 2: numerators are [1, 1] + // We need to handle this specially. Since the child has 2 elements, + // and numerators are implicit 1s, we need explicit values. + // Actually, if the leaf is Singles with 2 elements, after next_layer + // it becomes Generic. But at the leaf level with 2 elements, + // the tree has 2 layers (leaf + root). So the "child layer" could + // be Singles with 2 elements. Let's handle it. + // For the trivial case, we just provide the claims directly. + // n_left=1, n_right=1, d_left=d[0], d_right=d[1] + child_claims_by_instance.push([ + FieldElement::one(), + FieldElement::one(), + denominators[0].clone(), + denominators[1].clone(), + ]); + continue; + } + }; + + let _ = idx; // suppress unused warning + child_claims_by_instance.push([ + child_n[0].clone(), + child_n[1].clone(), + child_d[0].clone(), + child_d[1].clone(), + ]); + } + + // Append child claims to transcript + for claims in &child_claims_by_instance { + for claim in claims { + transcript.append_field_element(claim); + } + } + + // Sample eta to fold left/right + let eta: FieldElement = transcript.sample_field_element(); + + // Update per-instance claims + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = child_claims_by_instance[idx]; + n_claims[i] = Some(nl + &(&eta * &(nr - nl))); + d_claims[i] = Some(dl + &(&eta * &(dr - dl))); + } + + current_point = vec![eta]; + + layer_proofs.push(BatchGkrLayerProof { + sumcheck_proof: SumcheckProof { + round_polys: vec![], + }, + child_claims_by_instance, + }); + } else { + // Non-trivial case: run shared sumcheck over max_parent_vars variables. + // For each active instance, build bookkeeping tables. + + // We run the sumcheck manually, combining round polynomials across instances. + let mut per_instance_tables: Vec> = active_instances + .iter() + .map(|&i| { + let tree_layer_idx = n_remaining - 1; + let child = &layers_per_instance[i][tree_layer_idx]; + + let parent_size = match child { + Layer::LogUpGeneric { denominators, .. } + | Layer::LogUpSingles { denominators } => denominators.len() / 2, + }; + + let (nl_table, nr_table, is_singles) = match child { + Layer::LogUpGeneric { numerators, .. } => { + let nl: Vec<_> = (0..parent_size) + .map(|j| numerators[2 * j].clone()) + .collect(); + let nr: Vec<_> = (0..parent_size) + .map(|j| numerators[2 * j + 1].clone()) + .collect(); + (nl, nr, false) + } + Layer::LogUpSingles { .. } => { + // Singles: numerators are all 1 + (vec![], vec![], true) + } + }; + + let denominators = match child { + Layer::LogUpGeneric { denominators, .. } + | Layer::LogUpSingles { denominators } => denominators, + }; + + let dl_table: Vec<_> = (0..parent_size) + .map(|j| denominators[2 * j].clone()) + .collect(); + let dr_table: Vec<_> = (0..parent_size) + .map(|j| denominators[2 * j + 1].clone()) + .collect(); + + let my_parent_num_vars = parent_num_vars_by_instance + [active_instances.iter().position(|&x| x == i).unwrap()]; + // current_point must have at least parent_num_vars coordinates + // (it grows by max_parent_vars+1 each layer, matching the tree doubling). + debug_assert!( + current_point.len() >= my_parent_num_vars, + "current_point.len()={} < parent_num_vars={}", + current_point.len(), + my_parent_num_vars + ); + let inst_point = instance_eval_point(¤t_point, my_parent_num_vars); + let eq_table = compute_eq_evals(&inst_point); + + PerInstanceTables { + nl_table, + nr_table, + dl_table, + dr_table, + eq_table, + eq_correction: FieldElement::one(), + is_singles, + parent_num_vars: my_parent_num_vars, + instance_point: inst_point, + } + }) + .collect(); + + let mut round_polys = Vec::with_capacity(max_parent_vars); + let mut challenges = Vec::with_capacity(max_parent_vars); + // Combined claim across instances (via sumcheck_alpha) + let mut _round_combined_claim = { + let mut sum = FieldElement::::zero(); + let mut alpha_pow = FieldElement::::one(); + for claim in &combined_claims { + sum = &sum + &(&alpha_pow * claim); + alpha_pow = &alpha_pow * &sumcheck_alpha; + } + sum + }; + + for round_idx in 0..max_parent_vars { + // For each active instance, compute the round poly contribution + let mut batch_s0 = FieldElement::::zero(); + let mut batch_s1 = FieldElement::::zero(); + let mut batch_s2 = FieldElement::::zero(); + let mut batch_s3 = FieldElement::::zero(); + let mut alpha_pow = FieldElement::::one(); + + for (idx, tables) in per_instance_tables.iter_mut().enumerate() { + let n_unused = max_parent_vars - tables.parent_num_vars; + + if round_idx < n_unused { + // This instance hasn't started yet — constant polynomial = claim/2 + let half_claim = + &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + // S(0) = S(1) = half_claim, S(2) = S(3) = half_claim (constant) + batch_s0 = &batch_s0 + &(&alpha_pow * &half_claim); + batch_s1 = &batch_s1 + &(&alpha_pow * &half_claim); + batch_s2 = &batch_s2 + &(&alpha_pow * &half_claim); + batch_s3 = &batch_s3 + &(&alpha_pow * &half_claim); + alpha_pow = &alpha_pow * &sumcheck_alpha; + continue; + } + + let instance_round = round_idx - n_unused; + let half = tables.nl_table.len().max(tables.dl_table.len()) / 2; + + if half == 0 { + // Already reduced to constant + let half_claim = + &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + batch_s0 = &batch_s0 + &(&alpha_pow * &half_claim); + batch_s1 = &batch_s1 + &(&alpha_pow * &half_claim); + batch_s2 = &batch_s2 + &(&alpha_pow * &half_claim); + batch_s3 = &batch_s3 + &(&alpha_pow * &half_claim); + alpha_pow = &alpha_pow * &sumcheck_alpha; + continue; + } + + // Eq polynomial factoring (same as single-instance prover). + // Use the instance-specific eval point (not the shared current_point) + // so that r_round matches the eq_table built from instance_eval_point. + let r_round = tables.instance_point[instance_round].clone(); + + // Pre-halve eq_table + for j in 0..half { + tables.eq_table[j] = &tables.eq_table[2 * j] + &tables.eq_table[2 * j + 1]; + } + tables.eq_table.truncate(half); + + let one = FieldElement::::one(); + + // Compute h(0) and h(2) (inner sum without eq_round factor) + let (raw_h0, raw_h2) = if tables.is_singles { + // Singles gate: gate(t) = dl(t) + dr(t) + lambda * dl(t) * dr(t) + // (numerators are all 1, so the fraction sum is 1/dl + 1/dr) + let mut h0 = FieldElement::::zero(); + let mut h2 = FieldElement::::zero(); + for j in 0..half { + let eq_rem = &tables.eq_table[j]; + let dl_l = &tables.dl_table[2 * j]; + let dl_r = &tables.dl_table[2 * j + 1]; + let dr_l = &tables.dr_table[2 * j]; + let dr_r = &tables.dr_table[2 * j + 1]; + + // t=0 + let gate_0 = &(dl_l + dr_l) + &(&lambda * &(dl_l * dr_l)); + h0 = &h0 + &(eq_rem * &gate_0); + + // t=2 + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = &(&dl_2 + &dr_2) + &(&lambda * &(&dl_2 * &dr_2)); + h2 = &h2 + &(eq_rem * &gate_2); + } + (h0, h2) + } else { + // Generic gate: nl*dr + dl*(nr + lambda*dr) + let mut h0 = FieldElement::::zero(); + let mut h2 = FieldElement::::zero(); + for j in 0..half { + let eq_rem = &tables.eq_table[j]; + let nl_l = &tables.nl_table[2 * j]; + let nl_r = &tables.nl_table[2 * j + 1]; + let nr_l = &tables.nr_table[2 * j]; + let nr_r = &tables.nr_table[2 * j + 1]; + let dl_l = &tables.dl_table[2 * j]; + let dl_r = &tables.dl_table[2 * j + 1]; + let dr_l = &tables.dr_table[2 * j]; + let dr_r = &tables.dr_table[2 * j + 1]; + + // t=0 + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + h0 = &h0 + &(eq_rem * &gate_0); + + // t=2 + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = + &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + h2 = &h2 + &(eq_rem * &gate_2); + } + (h0, h2) + }; + + // Apply eq_correction + let total_h0 = &tables.eq_correction * &raw_h0; + let total_h2 = &tables.eq_correction * &raw_h2; + + // Recover S(t) from h(t) and eq_round(r, t) + let one_minus_r = &one - &r_round; + let s0 = &one_minus_r * &total_h0; + let s1 = &combined_claims[idx] - &s0; + + let r_inv = r_round.inv().expect("r_round = 0 is probability 2^{-64}"); + let h1 = &s1 * &r_inv; + + let three = FieldElement::::from(3u64); + let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; + + let eq_at_2 = &(&three * &r_round) - &one; + let s2 = &eq_at_2 * &total_h2; + + let eq_at_3 = &(&FieldElement::::from(5u64) * &r_round) + - &FieldElement::::from(2u64); + let s3 = &eq_at_3 * &h3; + + batch_s0 = &batch_s0 + &(&alpha_pow * &s0); + batch_s1 = &batch_s1 + &(&alpha_pow * &s1); + batch_s2 = &batch_s2 + &(&alpha_pow * &s2); + batch_s3 = &batch_s3 + &(&alpha_pow * &s3); + alpha_pow = &alpha_pow * &sumcheck_alpha; + } + + let round_poly = RoundPoly::new(vec![batch_s0, batch_s1, batch_s2, batch_s3]); + + // Append to transcript + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + // Sample challenge + let challenge: FieldElement = transcript.sample_field_element(); + _round_combined_claim = round_poly.evaluate(&challenge); + + // Update per-instance: fold tables and update eq_correction + for (idx, tables) in per_instance_tables.iter_mut().enumerate() { + let n_unused = max_parent_vars - tables.parent_num_vars; + if round_idx < n_unused { + // Instance hasn't started yet: its per-instance polynomial is + // constant = claim/2 at all evaluation points. + // p_i(challenge) = claim_i / 2. + combined_claims[idx] = + &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + continue; + } + + let instance_round = round_idx - n_unused; + let half = tables.dl_table.len() / 2; + + if half == 0 { + combined_claims[idx] = + &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + continue; + } + + // Update eq_correction using instance-specific eval point + let r_round = &tables.instance_point[instance_round]; + let one = FieldElement::::one(); + let eq_update = + &(r_round * &challenge) + &(&(&one - r_round) * &(&one - &challenge)); + tables.eq_correction = &tables.eq_correction * &eq_update; + + // Fold gate tables + let fold_table = |table: &mut Vec>| { + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + }; + + if !tables.is_singles { + fold_table(&mut tables.nl_table); + fold_table(&mut tables.nr_table); + } + fold_table(&mut tables.dl_table); + fold_table(&mut tables.dr_table); + + // Update per-instance combined claim + // The per-instance round poly evaluated at challenge: + // We need to recompute. Actually, the combined_claim for the batch + // already accounts for all instances. But we need per-instance claims + // for the next round. Let's just evaluate the per-instance sum at the + // challenge point using the folded tables. + // After folding, the tables have `half` elements. The per-instance claim + // for next round = sum over j of eq_table[j] * gate(j) * eq_correction + // But that's expensive. Instead, use the identity: + // new_claim = S_i(challenge) where S_i is this instance's round poly. + // We already computed S_i(0) and S_i(1) = claim_i - S_i(0). + // For non-trivial instances, we can use the round poly evaluation directly. + // Let's recompute per-instance sum from the folded tables. + let mut new_claim = FieldElement::::zero(); + let new_half = tables.dl_table.len(); + if tables.is_singles { + for j in 0..new_half { + let eq_rem = &tables.eq_table[j]; + let gate = &(&tables.dl_table[j] + &tables.dr_table[j]) + + &(&lambda * &(&tables.dl_table[j] * &tables.dr_table[j])); + new_claim = &new_claim + &(eq_rem * &gate); + } + } else { + for j in 0..new_half { + let eq_rem = &tables.eq_table[j]; + let gate = &(&tables.nl_table[j] * &tables.dr_table[j]) + + &(&tables.dl_table[j] + * &(&tables.nr_table[j] + &(&lambda * &tables.dr_table[j]))); + new_claim = &new_claim + &(eq_rem * &gate); + } + } + combined_claims[idx] = &tables.eq_correction * &new_claim; + } + + round_polys.push(round_poly); + challenges.push(challenge); + } + + // After all sumcheck rounds, extract child claims from each instance's tables + let mut child_claims_by_instance = Vec::new(); + + for tables in &per_instance_tables { + if tables.is_singles { + child_claims_by_instance.push([ + FieldElement::one(), + FieldElement::one(), + tables.dl_table[0].clone(), + tables.dr_table[0].clone(), + ]); + } else { + child_claims_by_instance.push([ + tables.nl_table[0].clone(), + tables.nr_table[0].clone(), + tables.dl_table[0].clone(), + tables.dr_table[0].clone(), + ]); + } + } + + // Append child claims to transcript + for claims in &child_claims_by_instance { + for claim in claims { + transcript.append_field_element(claim); + } + } + + // Sample eta to fold left/right + let eta: FieldElement = transcript.sample_field_element(); + + // Update per-instance claims + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = child_claims_by_instance[idx]; + n_claims[i] = Some(nl + &(&eta * &(nr - nl))); + d_claims[i] = Some(dl + &(&eta * &(dr - dl))); + } + + let mut new_point = Vec::with_capacity(challenges.len() + 1); + new_point.push(eta); + new_point.extend(challenges); + current_point = new_point; + + layer_proofs.push(BatchGkrLayerProof { + sumcheck_proof: SumcheckProof { round_polys }, + child_claims_by_instance, + }); + } + } + + let final_claims: Vec<(FieldElement, FieldElement)> = (0..n_instances) + .map(|i| { + ( + n_claims[i] + .clone() + .unwrap_or_else(|| root_claims[i].0.clone()), + d_claims[i] + .clone() + .unwrap_or_else(|| root_claims[i].1.clone()), + ) + }) + .collect(); + + ( + BatchGkrProof { + root_claims, + layer_proofs, + }, + current_point, + final_claims, + ) +} + +/// Compute the per-instance evaluation point from the shared point. +/// +/// In batch GKR with mixed-size instances, smaller instances skip the first +/// rounds of each shared sumcheck. Their variables bind to the **last** +/// challenges, not the first. The eval point for instance i is: +/// `[shared_point[0] (eta)] ++ shared_point[len - (n_vars - 1)..]` +/// +/// Returns an empty vec for n_vars == 0 (trivial/output layers). +fn instance_eval_point( + shared_point: &[FieldElement], + n_vars: usize, +) -> Vec> { + if n_vars == 0 { + return vec![]; + } + if n_vars >= shared_point.len() { + return shared_point.to_vec(); + } + let mut point = Vec::with_capacity(n_vars); + point.push(shared_point[0].clone()); // eta (shared left/right selector) + let start = shared_point.len() - (n_vars - 1); + point.extend_from_slice(&shared_point[start..]); + point +} + +/// Per-instance bookkeeping tables for the batch sumcheck inner loop. +struct PerInstanceTables { + nl_table: Vec>, + nr_table: Vec>, + dl_table: Vec>, + dr_table: Vec>, + eq_table: Vec>, + eq_correction: FieldElement, + is_singles: bool, + parent_num_vars: usize, + /// The instance-specific evaluation point derived from the shared current_point. + /// Used for r_round lookups in the Dao-Thaler eq factoring. + instance_point: Vec>, +} + +// ============================================================================= +// Batch GKR verifier +// ============================================================================= + +/// Verify a batch GKR proof. +/// +/// Replays the Fiat-Shamir transcript and checks sumcheck consistency and gate +/// equations for all instances simultaneously. +/// +/// # Returns +/// `Ok((shared_random_point, per_instance_claims))` where per_instance_claims[i] = (n_claim, d_claim). +#[allow(clippy::type_complexity)] +pub fn gkr_verify_batch( + proof: &BatchGkrProof, + n_layers_by_instance: &[usize], + transcript: &mut impl IsTranscript, +) -> Result< + ( + Vec>, + Vec<(FieldElement, FieldElement)>, + ), + GkrError, +> { + let n_instances = proof.root_claims.len(); + if n_layers_by_instance.len() != n_instances { + return Err(GkrError::InvalidTree { + reason: "n_layers_by_instance length mismatch".to_string(), + }); + } + + // Domain separation (must match prover) + transcript.append_bytes(b"gkr_batch"); + transcript.append_bytes(&(n_instances as u64).to_le_bytes()); + + if n_instances == 0 { + return Ok((vec![], vec![])); + } + + let max_layers = *n_layers_by_instance.iter().max().unwrap(); + + // Track per-instance state + let mut n_claims: Vec>> = vec![None; n_instances]; + let mut d_claims: Vec>> = vec![None; n_instances]; + let mut current_point: Vec> = vec![]; + + for (layer_idx, layer_proof) in proof.layer_proofs.iter().enumerate() { + let n_remaining = max_layers - layer_idx; + + // Detect output layers — use actual (root_n, root_d) from the proof + for i in 0..n_instances { + if n_layers_by_instance[i] == n_remaining { + n_claims[i] = Some(proof.root_claims[i].0.clone()); + d_claims[i] = Some(proof.root_claims[i].1.clone()); + } + } + + // Append active claims to transcript + for i in 0..n_instances { + if let (Some(n), Some(d)) = (&n_claims[i], &d_claims[i]) { + transcript.append_field_element(n); + transcript.append_field_element(d); + } + } + + // Sample randomness + let sumcheck_alpha: FieldElement = transcript.sample_field_element(); + let lambda: FieldElement = transcript.sample_field_element(); + + // Collect active instances + let mut active_instances: Vec = Vec::new(); + let mut combined_claims: Vec> = Vec::new(); + + for i in 0..n_instances { + if n_claims[i].is_some() && n_layers_by_instance[i] > 0 { + active_instances.push(i); + let n = n_claims[i].as_ref().unwrap(); + let d = d_claims[i].as_ref().unwrap(); + let claim = n + &(&lambda * d); + let n_unused = max_layers - n_layers_by_instance[i]; + if n_unused > 0 { + let doubling = FieldElement::::from(1u64 << n_unused); + combined_claims.push(&claim * &doubling); + } else { + combined_claims.push(claim); + } + } + } + + let round_polys = &layer_proof.sumcheck_proof.round_polys; + + if round_polys.is_empty() { + // Trivial layer: no sumcheck needed + // Append child claims and sample eta (same as prover) + for claims in &layer_proof.child_claims_by_instance { + for claim in claims { + transcript.append_field_element(claim); + } + } + + let eta: FieldElement = transcript.sample_field_element(); + + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims_by_instance[idx]; + n_claims[i] = Some(nl + &(&eta * &(nr - nl))); + d_claims[i] = Some(dl + &(&eta * &(dr - dl))); + } + + current_point = vec![eta]; + } else { + // Non-trivial: verify sumcheck + let num_rounds = round_polys.len(); + + // Compute combined claim across instances via sumcheck_alpha + let mut current_sum = { + let mut sum = FieldElement::::zero(); + let mut alpha_pow = FieldElement::::one(); + for claim in &combined_claims { + sum = &sum + &(&alpha_pow * claim); + alpha_pow = &alpha_pow * &sumcheck_alpha; + } + sum + }; + + let mut challenges = Vec::with_capacity(num_rounds); + + for (round, round_poly) in round_polys.iter().enumerate() { + if round_poly.sum_at_binary() != current_sum { + return Err(GkrError::SumcheckFailed { + layer: layer_idx, + reason: format!("round {} sum mismatch: p(0)+p(1) != expected sum", round), + }); + } + + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + let challenge: FieldElement = transcript.sample_field_element(); + current_sum = round_poly.evaluate(&challenge); + challenges.push(challenge); + } + + // Gate check: for each active instance, verify the gate equation + let mut expected_sum = FieldElement::::zero(); + let mut alpha_pow = FieldElement::::one(); + + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims_by_instance[idx]; + + // parent_num_vars for this instance at this layer: + // instance i's child layer has (n_layers[i] - n_remaining + 1) variables, + // so the parent has (n_layers[i] - n_remaining) variables. + let parent_num_vars_i = n_layers_by_instance[i] - n_remaining; + let sumcheck_n_unused = num_rounds - parent_num_vars_i; + + // eq evaluation: the prover builds eq from the instance-specific eval point + // (eta + last challenges), and the active sumcheck challenges are the last ones. + let eq_val = if parent_num_vars_i == 0 { + FieldElement::::one() + } else { + let inst_point = instance_eval_point(¤t_point, parent_num_vars_i); + compute_eq_at_point(&inst_point, &challenges[sumcheck_n_unused..]) + }; + + // gate_combined = nl*dr + nr*dl + lambda*dl*dr + let gate_combined = &(&(nl * dr) + &(nr * dl)) + &(&lambda * &(dl * dr)); + + // No doubling factor here: the 2^n_unused doubling is in the initial + // combined claim and gets cancelled by the unused rounds' halving. + // The point evaluation at the final challenge is just eq * gate. + let instance_eval = &eq_val * &gate_combined; + + expected_sum = &expected_sum + &(&alpha_pow * &instance_eval); + alpha_pow = &alpha_pow * &sumcheck_alpha; + } + + if current_sum != expected_sum { + return Err(GkrError::GateCheckFailed { layer: layer_idx }); + } + + // Append child claims to transcript + for claims in &layer_proof.child_claims_by_instance { + for claim in claims { + transcript.append_field_element(claim); + } + } + + let eta: FieldElement = transcript.sample_field_element(); + + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims_by_instance[idx]; + n_claims[i] = Some(nl + &(&eta * &(nr - nl))); + d_claims[i] = Some(dl + &(&eta * &(dr - dl))); + } + + let mut new_point = Vec::with_capacity(challenges.len() + 1); + new_point.push(eta); + new_point.extend(challenges); + current_point = new_point; + } + } + + let final_claims: Vec<(FieldElement, FieldElement)> = (0..n_instances) + .map(|i| { + ( + n_claims[i] + .clone() + .unwrap_or_else(|| proof.root_claims[i].0.clone()), + d_claims[i] + .clone() + .unwrap_or_else(|| proof.root_claims[i].1.clone()), + ) + }) + .collect(); + + Ok((current_point, final_claims)) +} + +/// Compute eq(a, b) for two points of equal length. +/// +/// eq(a, b) = prod_i (a_i * b_i + (1 - a_i) * (1 - b_i)) +/// +/// This is a single field element, NOT the full eq table. +fn compute_eq_at_point( + a: &[FieldElement], + b: &[FieldElement], +) -> FieldElement { + assert_eq!(a.len(), b.len(), "eq points must have equal length"); + let one = FieldElement::::one(); + a.iter() + .zip(b.iter()) + .fold(FieldElement::one(), |acc, (ai, bi)| { + let term = &(ai * bi) + &(&(&one - ai) * &(&one - bi)); + &acc * &term + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crypto::fiat_shamir::default_transcript::DefaultTranscript; + use math::field::goldilocks::GoldilocksField; + + type FE = FieldElement; + + #[test] + fn test_fraction_add() { + // (3/5) + (7/11) = (3*11 + 7*5) / (5*11) = (33 + 35) / 55 = 68/55 + let a = Fraction::new(FE::from(3u64), FE::from(5u64)); + let b = Fraction::new(FE::from(7u64), FE::from(11u64)); + let result = a.add(&b); + + assert_eq!(result.numerator, FE::from(68u64)); + assert_eq!(result.denominator, FE::from(55u64)); + } + + #[test] + fn test_build_summation_tree_4_leaves() { + // 4 leaves: 1/2, 3/4, 5/6, 7/8 + // Layer 0 (4 fractions): [1/2, 3/4, 5/6, 7/8] + // Layer 1 (2 fractions): + // pair 0: 1/2 + 3/4 = (1*4 + 3*2)/(2*4) = 10/8 + // pair 1: 5/6 + 7/8 = (5*8 + 7*6)/(6*8) = 82/48 + // Layer 2 (1 fraction, root): + // 10/8 + 82/48 = (10*48 + 82*8)/(8*48) = (480 + 656)/384 = 1136/384 + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + + let tree = build_summation_tree(nums, dens); + + // Should have 3 layers: 4 -> 2 -> 1 + assert_eq!(tree.len(), 3); + assert_eq!(tree[0].numerators.len(), 4); + assert_eq!(tree[1].numerators.len(), 2); + assert_eq!(tree[2].numerators.len(), 1); + + // Layer 1 checks + assert_eq!(tree[1].numerators[0], FE::from(10u64)); // 1*4 + 3*2 + assert_eq!(tree[1].denominators[0], FE::from(8u64)); // 2*4 + assert_eq!(tree[1].numerators[1], FE::from(82u64)); // 5*8 + 7*6 + assert_eq!(tree[1].denominators[1], FE::from(48u64)); // 6*8 + + // Root checks + assert_eq!(tree[2].numerators[0], FE::from(1136u64)); // 10*48 + 82*8 + assert_eq!(tree[2].denominators[0], FE::from(384u64)); // 8*48 + + // Verify the root fraction equals the actual sum: + // 1/2 + 3/4 + 5/6 + 7/8 = 12/24 + 18/24 + 20/24 + 21/24 = 71/24 + // Check: 1136/384 = 71/24 (both sides: 1136*24 = 27264, 71*384 = 27264) + let root_n = &tree[2].numerators[0]; + let root_d = &tree[2].denominators[0]; + assert_eq!(root_n * &FE::from(24u64), &FE::from(71u64) * root_d); + } + + #[test] + fn test_build_summation_tree_8_leaves() { + // 8 leaves: i/(i+1) for i in 1..=8, i.e., 1/2, 2/3, 3/4, 4/5, 5/6, 6/7, 7/8, 8/9 + let nums: Vec = (1..=8).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (2..=9).map(|i| FE::from(i as u64)).collect(); + + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Should have 4 layers: 8 -> 4 -> 2 -> 1 + assert_eq!(tree.len(), 4); + assert_eq!(tree[0].numerators.len(), 8); + assert_eq!(tree[1].numerators.len(), 4); + assert_eq!(tree[2].numerators.len(), 2); + assert_eq!(tree[3].numerators.len(), 1); + + // Compute expected root by sequential fraction addition + let mut acc = Fraction::new(nums[0], dens[0]); + for i in 1..8 { + acc = acc.add(&Fraction::new(nums[i], dens[i])); + } + + // The tree root and the sequential sum should represent the same rational number: + // tree_n / tree_d == acc_n / acc_d <=> tree_n * acc_d == acc_n * tree_d + let root_n = &tree[3].numerators[0]; + let root_d = &tree[3].denominators[0]; + assert_eq!( + root_n * &acc.denominator, + &acc.numerator * root_d, + "Tree root must equal sequential sum as a fraction" + ); + } + + #[test] + fn test_build_summation_tree_single_leaf() { + // Edge case: 1 leaf (2^0 = 1) + let nums = vec![FE::from(42u64)]; + let dens = vec![FE::from(7u64)]; + + let tree = build_summation_tree(nums, dens); + + // Should have 1 layer (just the leaf) + assert_eq!(tree.len(), 1); + assert_eq!(tree[0].numerators.len(), 1); + assert_eq!(tree[0].numerators[0], FE::from(42u64)); + assert_eq!(tree[0].denominators[0], FE::from(7u64)); + } + + #[test] + #[should_panic(expected = "number of leaves must be a power of 2")] + fn test_build_summation_tree_non_power_of_2_panics() { + let nums = vec![FE::from(1u64), FE::from(2u64), FE::from(3u64)]; + let dens = vec![FE::from(1u64), FE::from(1u64), FE::from(1u64)]; + let _ = build_summation_tree(nums, dens); + } + + #[test] + #[should_panic(expected = "numerators and denominators must have the same length")] + fn test_build_summation_tree_mismatched_lengths_panics() { + let nums = vec![FE::from(1u64), FE::from(2u64)]; + let dens = vec![FE::from(1u64)]; + let _ = build_summation_tree(nums, dens); + } + + // ==================== compute_eq_evals tests ==================== + + #[test] + fn test_compute_eq_evals_empty_point() { + // eq with 0 variables: single entry = 1 + let evals = compute_eq_evals::(&[]); + assert_eq!(evals.len(), 1); + assert_eq!(evals[0], FE::one()); + } + + #[test] + fn test_compute_eq_evals_1var() { + // eq((r,), (b,)) = r*b + (1-r)*(1-b) + // For r = 3: eq(3, 0) = (1-3) = -2, eq(3, 1) = 3 + let r = FE::from(3u64); + let evals = compute_eq_evals(std::slice::from_ref(&r)); + assert_eq!(evals.len(), 2); + assert_eq!(evals[0], FE::one() - r); // eq(r, 0) = 1-r = -2 + assert_eq!(evals[1], r); // eq(r, 1) = r = 3 + } + + #[test] + fn test_compute_eq_evals_2var() { + // eq((r0, r1), (b0, b1)) = [r0*b0 + (1-r0)*(1-b0)] * [r1*b1 + (1-r1)*(1-b1)] + let r0 = FE::from(2u64); + let r1 = FE::from(5u64); + let evals = compute_eq_evals(&[r0, r1]); + assert_eq!(evals.len(), 4); + + let one = FE::one(); + let one_minus_r0 = &one - &r0; + let one_minus_r1 = &one - &r1; + + // Index 0 = (b0=0, b1=0): (1-r0)*(1-r1) + assert_eq!(evals[0], &one_minus_r0 * &one_minus_r1); + // Index 1 = (b0=1, b1=0): r0*(1-r1) + assert_eq!(evals[1], &r0 * &one_minus_r1); + // Index 2 = (b0=0, b1=1): (1-r0)*r1 + assert_eq!(evals[2], &one_minus_r0 * &r1); + // Index 3 = (b0=1, b1=1): r0*r1 + assert_eq!(evals[3], &r0 * &r1); + } + + #[test] + fn test_compute_eq_evals_sum_to_one_on_booleans() { + // When point is Boolean, eq_evals should have exactly one 1 and rest 0 + let point = vec![FE::one(), FE::zero(), FE::one()]; // b = (1, 0, 1) = index 5 + let evals = compute_eq_evals(&point); + assert_eq!(evals.len(), 8); + for (i, e) in evals.iter().enumerate() { + if i == 5 { + assert_eq!(*e, FE::one()); + } else { + assert_eq!(*e, FE::zero()); + } + } + } + + // ==================== evaluate_mle tests ==================== + + #[test] + fn test_evaluate_mle_linear() { + // MLE of [3, 7] at point r: + // f(x) = 3*(1-x) + 7*x = 3 + 4x + // f(5) = 3 + 20 = 23 + let table = vec![FE::from(3u64), FE::from(7u64)]; + let result = evaluate_mle(&table, &[FE::from(5u64)]); + assert_eq!(result, FE::from(23u64)); + } + + #[test] + fn test_evaluate_mle_at_boolean() { + // MLE at a Boolean point should return the table entry + let table = vec![ + FE::from(10u64), + FE::from(20u64), + FE::from(30u64), + FE::from(40u64), + ]; + // Index 2 = (b0=0, b1=1) + let result = evaluate_mle(&table, &[FE::zero(), FE::one()]); + assert_eq!(result, FE::from(30u64)); + } + + // ==================== GKR prover tests ==================== + + #[test] + fn test_gkr_prove_2_leaves() { + // Simplest non-trivial case: 2 leaves + // Tree: layer 0 (2 fractions), layer 1 (root, 1 fraction) + // Leaves: 3/5, 7/11 + // Root: (3*11 + 7*5) / (5*11) = 68/55 + let nums = vec![FE::from(3u64), FE::from(7u64)]; + let dens = vec![FE::from(5u64), FE::from(11u64)]; + let tree = build_summation_tree(nums, dens); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + + // claimed_sum = 68/55 + let expected_sum = &FE::from(68u64) * &FE::from(55u64).inv().unwrap(); + assert_eq!(proof.claimed_sum, expected_sum); + + // Should have 1 layer proof (root -> leaves) + assert_eq!(proof.layer_proofs.len(), 1); + + // The first (and only) layer proof has a trivial sumcheck (0 rounds) + assert_eq!(proof.layer_proofs[0].sumcheck_proof.round_polys.len(), 0); + + // The child claims should be the raw leaf values + assert_eq!(proof.layer_proofs[0].child_claims[0], FE::from(3u64)); // n_left + assert_eq!(proof.layer_proofs[0].child_claims[1], FE::from(7u64)); // n_right + assert_eq!(proof.layer_proofs[0].child_claims[2], FE::from(5u64)); // d_left + assert_eq!(proof.layer_proofs[0].child_claims[3], FE::from(11u64)); // d_right + + // final_point should have 1 element (eta) + assert_eq!(final_point.len(), 1); + + // Final claims should be the leaf MLE evaluated at final_point + // n_MLE(eta) = 3*(1-eta) + 7*eta = 3 + 4*eta + // d_MLE(eta) = 5*(1-eta) + 11*eta = 5 + 6*eta + let eta = &final_point[0]; + let expected_n = &FE::from(3u64) * &(&FE::one() - eta) + &(&FE::from(7u64) * eta); + let expected_d = &FE::from(5u64) * &(&FE::one() - eta) + &(&FE::from(11u64) * eta); + assert_eq!(final_n_claim, expected_n); + assert_eq!(final_d_claim, expected_d); + } + + #[test] + fn test_gkr_prove_4_leaves() { + // 4 leaves: 1/2, 3/4, 5/6, 7/8 + // Tree has 3 layers (0=leaves size 4, 1=size 2, 2=root size 1) + // GKR reduces: root -> layer 1 (trivial, 0 vars) -> layer 0 (1-var sumcheck) + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums.clone(), dens.clone()); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + + // claimed_sum = root_n / root_d = 1136 / 384 + let expected_sum = &FE::from(1136u64) * &FE::from(384u64).inv().unwrap(); + assert_eq!(proof.claimed_sum, expected_sum); + + // Should have 2 layer proofs + assert_eq!(proof.layer_proofs.len(), 2); + + // First layer proof: root (1 elem) -> layer 1 (2 elems), trivial (0 rounds) + assert_eq!(proof.layer_proofs[0].sumcheck_proof.round_polys.len(), 0); + + // Second layer proof: layer 1 (2 elems) -> layer 0 (4 elems) + // This is a 1-variable sumcheck, so should have 1 round polynomial + assert_eq!(proof.layer_proofs[1].sumcheck_proof.round_polys.len(), 1); + + // The round polynomial should have 4 evaluations (degree 3) + assert_eq!( + proof.layer_proofs[1].sumcheck_proof.round_polys[0].num_evals(), + 4 + ); + + // final_point should have 2 elements (challenge from sumcheck + eta) + assert_eq!(final_point.len(), 2); + + // Verify final claims match the leaf MLEs at final_point + let expected_n = evaluate_mle(&nums, &final_point); + let expected_d = evaluate_mle(&dens, &final_point); + assert_eq!(final_n_claim, expected_n); + assert_eq!(final_d_claim, expected_d); + } + + #[test] + fn test_gkr_prove_claimed_sum() { + // Verify that proof.claimed_sum equals root_n * root_d.inv() + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums, dens); + + let root_n = &tree[2].numerators[0]; + let root_d = &tree[2].denominators[0]; + let expected_sum = root_n * &root_d.inv().unwrap(); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, _, _, _) = gkr_prove(&tree, &mut transcript); + + assert_eq!(proof.claimed_sum, expected_sum); + } + + #[test] + fn test_gkr_prove_8_leaves() { + // 8 leaves: i/(i+1) for i in 1..=8 + // Tree: 4 layers (sizes 8, 4, 2, 1) + // GKR reductions: root->layer2 (trivial), layer2->layer1 (1-var), layer1->layer0 (2-var) + let nums: Vec = (1..=8).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (2..=9).map(|i| FE::from(i as u64)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + + // Should have 3 layer proofs + assert_eq!(proof.layer_proofs.len(), 3); + + // Layer 0: root (1) -> layer 2 (2): trivial + assert_eq!(proof.layer_proofs[0].sumcheck_proof.round_polys.len(), 0); + + // Layer 1: layer 2 (2) -> layer 1 (4): 1-variable sumcheck + assert_eq!(proof.layer_proofs[1].sumcheck_proof.round_polys.len(), 1); + + // Layer 2: layer 1 (4) -> layer 0 (8): 2-variable sumcheck + assert_eq!(proof.layer_proofs[2].sumcheck_proof.round_polys.len(), 2); + + // final_point should have 3 elements + assert_eq!(final_point.len(), 3); + + // Verify final claims match the leaf MLEs at final_point + let expected_n = evaluate_mle(&nums, &final_point); + let expected_d = evaluate_mle(&dens, &final_point); + assert_eq!(final_n_claim, expected_n); + assert_eq!(final_d_claim, expected_d); + } + + #[test] + fn test_gkr_prove_16_leaves() { + // 16 leaves with various fractions + // Tree: 5 layers (sizes 16, 8, 4, 2, 1) + let nums: Vec = (1..=16).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (17..=32).map(|i| FE::from(i as u64)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + let mut transcript = DefaultTranscript::::new(&[0xAB]); + let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + + // Should have 4 layer proofs + assert_eq!(proof.layer_proofs.len(), 4); + + // final_point should have 4 elements (log2(16)) + assert_eq!(final_point.len(), 4); + + // Verify final claims match the leaf MLEs at final_point + let expected_n = evaluate_mle(&nums, &final_point); + let expected_d = evaluate_mle(&dens, &final_point); + assert_eq!(final_n_claim, expected_n); + assert_eq!(final_d_claim, expected_d); + } + + #[test] + fn test_gkr_prove_deterministic() { + // Same inputs and transcript seed should produce identical proofs + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums, dens); + + let mut t1 = DefaultTranscript::::new(&[0x42]); + let (proof1, point1, n1, d1) = gkr_prove(&tree, &mut t1); + + let mut t2 = DefaultTranscript::::new(&[0x42]); + let (proof2, point2, n2, d2) = gkr_prove(&tree, &mut t2); + + assert_eq!(proof1.claimed_sum, proof2.claimed_sum); + assert_eq!(point1, point2); + assert_eq!(n1, n2); + assert_eq!(d1, d2); + assert_eq!(proof1.layer_proofs.len(), proof2.layer_proofs.len()); + + for (lp1, lp2) in proof1.layer_proofs.iter().zip(proof2.layer_proofs.iter()) { + assert_eq!(lp1.child_claims, lp2.child_claims); + assert_eq!( + lp1.sumcheck_proof.round_polys.len(), + lp2.sumcheck_proof.round_polys.len() + ); + for (rp1, rp2) in lp1 + .sumcheck_proof + .round_polys + .iter() + .zip(lp2.sumcheck_proof.round_polys.iter()) + { + assert_eq!(rp1.evals(), rp2.evals()); + } + } + } + + #[test] + fn test_gkr_prove_sumcheck_consistency() { + // Verify that the round polynomial p(0) + p(1) matches the combined claim + // for the non-trivial sumcheck layers + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Re-run the protocol manually to extract the combined_claim at each layer + // and verify the sumcheck round poly sum + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, _, _, _) = gkr_prove(&tree, &mut transcript); + + // Replay transcript to get the same challenges + let mut replay = DefaultTranscript::::new(&[]); + replay.append_field_element(&proof.claimed_sum); + + let mut n_claim = tree[2].numerators[0]; + let mut d_claim = tree[2].denominators[0]; + + for (layer_idx, lp) in proof.layer_proofs.iter().enumerate() { + let lambda: FE = replay.sample_field_element(); + let combined_claim = &n_claim + &(&lambda * &d_claim); + + if lp.sumcheck_proof.round_polys.is_empty() { + // Trivial layer: verify gate equation directly + let nl = &lp.child_claims[0]; + let nr = &lp.child_claims[1]; + let dl = &lp.child_claims[2]; + let dr = &lp.child_claims[3]; + let gate_val = &(nl * dr) + &(nr * dl) + &(&lambda * &(dl * dr)); + assert_eq!( + combined_claim, gate_val, + "Gate equation failed at trivial layer {}", + layer_idx + ); + } else { + // Non-trivial layer: verify p(0) + p(1) = combined_claim + let first_round = &lp.sumcheck_proof.round_polys[0]; + assert_eq!( + first_round.sum_at_binary(), + combined_claim, + "Sumcheck round sum mismatch at layer {}", + layer_idx + ); + } + + // Replay transcript operations from the sumcheck + for rp in &lp.sumcheck_proof.round_polys { + for eval in rp.evals() { + replay.append_field_element(eval); + } + let _challenge: FE = replay.sample_field_element(); + } + + // Replay child claims and eta + for claim in &lp.child_claims { + replay.append_field_element(claim); + } + let eta: FE = replay.sample_field_element(); + + // Update claims for next layer + let one_minus_eta = &FE::one() - η + n_claim = &(&lp.child_claims[0] * &one_minus_eta) + &(&lp.child_claims[1] * &eta); + d_claim = &(&lp.child_claims[2] * &one_minus_eta) + &(&lp.child_claims[3] * &eta); + } + } + + #[test] + fn test_gkr_prove_single_leaf() { + // Edge case: single leaf (tree has 1 layer, no reductions) + let nums = vec![FE::from(42u64)]; + let dens = vec![FE::from(7u64)]; + let tree = build_summation_tree(nums, dens); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + + assert_eq!(proof.claimed_sum, FE::from(6u64)); // 42/7 = 6 + assert_eq!(proof.layer_proofs.len(), 0); + assert!(final_point.is_empty()); + assert_eq!(final_n_claim, FE::from(42u64)); + assert_eq!(final_d_claim, FE::from(7u64)); + } + + // ==================== GKR verifier tests ==================== + + #[test] + fn test_gkr_prove_verify_roundtrip_4() { + // 4 leaves: 1/2, 3/4, 5/6, 7/8 + // Tree has 3 layers (0=leaves size 4, 1=size 2, 2=root size 1) + // GKR reduces: root -> layer 1 (trivial) -> layer 0 (1-var sumcheck) + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); + let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + + // Verify with a fresh transcript (same seed) + let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + result.is_ok(), + "GKR verification should succeed for 4 leaves" + ); + let (verifier_point, verifier_n, verifier_d) = result.unwrap(); + + // The verifier's final point must match the prover's + assert_eq!( + verifier_point, prover_point, + "Verifier and prover must derive the same final point" + ); + + // The verifier's leaf claims must match the prover's + assert_eq!(verifier_n, prover_n, "n_claim must match"); + assert_eq!(verifier_d, prover_d, "d_claim must match"); + + // Additionally verify that the claims are consistent with the leaf MLEs + let expected_n = evaluate_mle(&nums, &verifier_point); + let expected_d = evaluate_mle(&dens, &verifier_point); + assert_eq!(verifier_n, expected_n, "n_claim must match leaf MLE"); + assert_eq!(verifier_d, expected_d, "d_claim must match leaf MLE"); + } + + #[test] + fn test_gkr_prove_verify_roundtrip_8() { + // 8 leaves: i/(i+1) for i in 1..=8 + // Tree: 4 layers (sizes 8, 4, 2, 1) + // GKR: root->layer2 (trivial), layer2->layer1 (1-var), layer1->layer0 (2-var) + let nums: Vec = (1..=8).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (2..=9).map(|i| FE::from(i as u64)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); + let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + + // Verify with a fresh transcript (same seed) + let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + result.is_ok(), + "GKR verification should succeed for 8 leaves" + ); + let (verifier_point, verifier_n, verifier_d) = result.unwrap(); + + // The verifier's final point must match the prover's + assert_eq!(verifier_point, prover_point); + + // The verifier's leaf claims must match the prover's + assert_eq!(verifier_n, prover_n); + assert_eq!(verifier_d, prover_d); + + // Verify consistency with leaf MLEs + let expected_n = evaluate_mle(&nums, &verifier_point); + let expected_d = evaluate_mle(&dens, &verifier_point); + assert_eq!(verifier_n, expected_n); + assert_eq!(verifier_d, expected_d); + } + + #[test] + fn test_gkr_verify_wrong_claimed_sum() { + // Create a valid proof and then tamper with the claimed_sum. + // The verifier should fail (either at sumcheck or gate check). + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums, dens); + + // Prove with correct claimed_sum + let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); + let (mut proof, _, _, _) = gkr_prove(&tree, &mut prover_transcript); + + // Tamper with the claimed_sum + proof.claimed_sum = &proof.claimed_sum + &FE::one(); + + // Verify with the tampered proof + let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + // The verification should fail. The tampered claimed_sum changes the + // transcript state, which changes lambda and eta, leading to a sumcheck + // or gate check failure at the non-trivial layer. + assert!( + result.is_err(), + "GKR verification should fail with tampered claimed_sum" + ); + } + + #[test] + fn test_gkr_prove_verify_roundtrip_16() { + // 16 leaves with various fractions + // Tree: 5 layers (sizes 16, 8, 4, 2, 1) + let nums: Vec = (1..=16).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (17..=32).map(|i| FE::from(i as u64)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); + let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + + // Verify with a fresh transcript (same seed) + let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + result.is_ok(), + "GKR verification should succeed for 16 leaves" + ); + let (verifier_point, verifier_n, verifier_d) = result.unwrap(); + + // The verifier's final point must match the prover's + assert_eq!(verifier_point, prover_point); + + // The verifier's leaf claims must match the prover's + assert_eq!(verifier_n, prover_n); + assert_eq!(verifier_d, prover_d); + + // Verify consistency with leaf MLEs + let expected_n = evaluate_mle(&nums, &verifier_point); + let expected_d = evaluate_mle(&dens, &verifier_point); + assert_eq!(verifier_n, expected_n); + assert_eq!(verifier_d, expected_d); + } + + // ==================== Batch GKR tests ==================== + + /// Helper: create a Layer::LogUpGeneric leaf from random-ish fractions. + fn make_generic_leaf(n_vars: usize) -> Layer { + let size = 1 << n_vars; + let denominators: Vec = (1..=size).map(|i| FE::from(i as u64 + 100)).collect(); + let numerators: Vec = (1..=size).map(|i| FE::from(i as u64)).collect(); + Layer::LogUpGeneric { + numerators, + denominators, + } + } + + /// Helper: create a Layer::LogUpSingles leaf. + fn make_singles_leaf(n_vars: usize) -> Layer { + let size = 1 << n_vars; + let denominators: Vec = (1..=size).map(|i| FE::from(i as u64 + 200)).collect(); + Layer::LogUpSingles { denominators } + } + + #[test] + fn test_batch_gkr_same_size_instances() { + // 3 instances, all with 4 leaves (n_vars=2, 3 layers each) + let instances: Vec>> = (0..3) + .map(|_| gen_layers(make_generic_leaf(2))) + .collect(); + + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, shared_point, final_claims) = + gkr_prove_batch(instances, &mut prover_transcript); + + assert_eq!(proof.root_claims.len(), 3); + assert_eq!(final_claims.len(), 3); + + let n_layers: Vec = proof + .root_claims + .iter() + .map(|_| 2) // all same: 2 reduction steps + .collect(); + + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); + assert!(result.is_ok(), "batch verify failed: {:?}", result.err()); + + let (v_point, v_claims) = result.unwrap(); + assert_eq!(v_point, shared_point); + assert_eq!(v_claims, final_claims); + } + + #[test] + fn test_batch_gkr_mixed_size_instances() { + // This is the key test: 3 instances with DIFFERENT sizes. + // n_vars = 2 (4 leaves), 4 (16 leaves), 6 (64 leaves) + // This exercises the instance_eval_point logic. + let instances: Vec>> = vec![ + gen_layers(make_generic_leaf(2)), // 3 layers, 2 reductions + gen_layers(make_generic_leaf(4)), // 5 layers, 4 reductions + gen_layers(make_generic_leaf(6)), // 7 layers, 6 reductions + ]; + + let n_layers: Vec = instances.iter().map(|l| l.len() - 1).collect(); + assert_eq!(n_layers, vec![2, 4, 6]); + + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, shared_point, final_claims) = + gkr_prove_batch(instances, &mut prover_transcript); + + assert_eq!(proof.root_claims.len(), 3); + + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); + assert!(result.is_ok(), "mixed-size batch verify failed: {:?}", result.err()); + + let (v_point, v_claims) = result.unwrap(); + assert_eq!(v_point, shared_point); + assert_eq!(v_claims, final_claims); + } + + #[test] + fn test_batch_gkr_mixed_size_with_singles() { + // Mix of Generic and Singles leaves with different sizes + let instances: Vec>> = vec![ + gen_layers(make_singles_leaf(2)), // 3 layers + gen_layers(make_generic_leaf(3)), // 4 layers + gen_layers(make_singles_leaf(5)), // 6 layers + gen_layers(make_generic_leaf(4)), // 5 layers + ]; + + let n_layers: Vec = instances.iter().map(|l| l.len() - 1).collect(); + + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, shared_point, final_claims) = + gkr_prove_batch(instances, &mut prover_transcript); + + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); + assert!(result.is_ok(), "mixed singles/generic batch verify failed: {:?}", result.err()); + + let (v_point, v_claims) = result.unwrap(); + assert_eq!(v_point, shared_point); + assert_eq!(v_claims, final_claims); + } + + #[test] + fn test_batch_gkr_many_mixed_instances() { + // Stress test: 20 instances with sizes varying from n_vars=1 to n_vars=6 + // Mimics the real VM scenario with many tables of different sizes + let instances: Vec>> = (0..20) + .map(|i| { + let n_vars = (i % 6) + 1; // 1, 2, 3, 4, 5, 6, 1, 2, ... + gen_layers(make_generic_leaf(n_vars)) + }) + .collect(); + + let n_layers: Vec = instances.iter().map(|l| l.len() - 1).collect(); + + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, shared_point, final_claims) = + gkr_prove_batch(instances, &mut prover_transcript); + + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); + assert!(result.is_ok(), "many mixed instances batch verify failed: {:?}", result.err()); + + let (v_point, v_claims) = result.unwrap(); + assert_eq!(v_point, shared_point); + assert_eq!(v_claims, final_claims); + } + + #[test] + fn test_instance_eval_point() { + // Verify the instance_eval_point helper + let point: Vec = vec![ + FE::from(10u64), // eta + FE::from(20u64), // c_0 + FE::from(30u64), // c_1 + FE::from(40u64), // c_2 + ]; + + // n_vars = 0: empty + assert!(instance_eval_point::(&point, 0).is_empty()); + + // n_vars = 4 (full): entire point + assert_eq!(instance_eval_point(&point, 4), point); + + // n_vars = 1: [eta] + assert_eq!(instance_eval_point(&point, 1), vec![FE::from(10u64)]); + + // n_vars = 2: [eta, c_2] (eta + last 1) + assert_eq!( + instance_eval_point(&point, 2), + vec![FE::from(10u64), FE::from(40u64)] + ); + + // n_vars = 3: [eta, c_1, c_2] (eta + last 2) + assert_eq!( + instance_eval_point(&point, 3), + vec![FE::from(10u64), FE::from(30u64), FE::from(40u64)] + ); + } +} diff --git a/crypto/stark/src/lagrange_kernel.rs b/crypto/stark/src/lagrange_kernel.rs new file mode 100644 index 000000000..1a5da3816 --- /dev/null +++ b/crypto/stark/src/lagrange_kernel.rs @@ -0,0 +1,316 @@ +use math::field::{ + element::FieldElement, + traits::{IsField, IsSubFieldOf}, +}; +#[cfg(feature = "parallel")] +use rayon::prelude::*; + +/// Compute the Lagrange kernel (eq polynomial) for a random point `r`. +/// +/// Given r = (r_0, r_1, ..., r_{n-1}), returns a vector s of length N = 2^n where: +/// +/// s[i] = eq(bits(i), r) = prod_{l=0}^{n-1} [r_l * b_l(i) + (1 - r_l) * (1 - b_l(i))] +/// +/// where b_l(i) = (i >> l) & 1 is the l-th bit of i. +/// +/// The Lagrange kernel is the auxiliary column that bridges GKR multilinear extension +/// claims back to committed univariate traces. It satisfies the partition-of-unity +/// property: sum_{i=0}^{N-1} s[i] = 1 for any r. +/// +/// Uses the butterfly algorithm in O(N) time and O(N) space. +pub fn compute_lagrange_kernel(r: &[FieldElement]) -> Vec> { + let n = r.len(); + let big_n = 1usize << n; + let one = FieldElement::::one(); + + let mut s = vec![one.clone(); big_n]; + + for j in 0..n { + let rj = &r[j]; + let one_minus_rj = &one - rj; + + #[cfg(feature = "parallel")] + { + if big_n >= 1024 { + s.par_iter_mut().enumerate().for_each(|(i, s_i)| { + if (i >> j) & 1 == 1 { + *s_i = &*s_i * rj; + } else { + *s_i = &*s_i * &one_minus_rj; + } + }); + } else { + for i in 0..big_n { + if (i >> j) & 1 == 1 { + s[i] = &s[i] * rj; + } else { + s[i] = &s[i] * &one_minus_rj; + } + } + } + } + + #[cfg(not(feature = "parallel"))] + { + for i in 0..big_n { + if (i >> j) & 1 == 1 { + s[i] = &s[i] * rj; + } else { + s[i] = &s[i] * &one_minus_rj; + } + } + } + } + + s +} + +/// Evaluate the multilinear extension (MLE) of `values` at point `r`. +/// +/// Given a function f defined on {0,1}^n by its truth table `values` (length N = 2^n), +/// its multilinear extension is: +/// +/// MLE(r) = sum_{i=0}^{N-1} values[i] * eq(bits(i), r) +/// +/// This is the inner product of `values` with the Lagrange kernel at `r`. +/// +/// Panics if `values.len()` is not a power of 2 or if `r.len() != log2(values.len())`. +pub fn eval_mle( + values: &[FieldElement], + r: &[FieldElement], +) -> FieldElement { + let n = r.len(); + let big_n = 1usize << n; + assert_eq!( + values.len(), + big_n, + "values length ({}) must equal 2^r.len() = 2^{} = {}", + values.len(), + n, + big_n + ); + + let kernel = compute_lagrange_kernel(r); + + let mut result = FieldElement::::zero(); + for i in 0..big_n { + result = &result + &(&values[i] * &kernel[i]); + } + + result +} + +/// Evaluate the multilinear extension (MLE) of base-field `values` at an extension-field point `r`. +/// +/// Same as `eval_mle` but accepts base-field values and an extension-field evaluation point. +/// Uses `F * E -> E` multiplication directly (no `to_extension()` conversion). +/// +/// Panics if `values.len()` is not a power of 2 or if `r.len() != log2(values.len())`. +pub fn eval_mle_base( + values: &[FieldElement], + r: &[FieldElement], +) -> FieldElement +where + F: IsField + IsSubFieldOf, + E: IsField, +{ + let n = r.len(); + let big_n = 1usize << n; + assert_eq!( + values.len(), + big_n, + "values length ({}) must equal 2^r.len() = 2^{} = {}", + values.len(), + n, + big_n + ); + + let kernel = compute_lagrange_kernel(r); + + let mut result = FieldElement::::zero(); + for i in 0..big_n { + // F * E -> E multiplication (no to_extension) + result = &result + &(&values[i] * &kernel[i]); + } + + result +} + +/// Evaluate the MLE of base-field `values` at an extension-field point, using a pre-computed +/// Lagrange kernel. +/// +/// This avoids recomputing `compute_lagrange_kernel(r)` when evaluating multiple columns +/// at the same point `r`. +/// +/// Panics if `values.len() != kernel.len()`. +pub fn eval_mle_base_with_kernel( + values: &[FieldElement], + kernel: &[FieldElement], +) -> FieldElement +where + F: IsField + IsSubFieldOf, + E: IsField + Send + Sync, + FieldElement: Send + Sync, +{ + let big_n = values.len(); + assert_eq!( + big_n, + kernel.len(), + "values length ({}) must equal kernel length ({})", + big_n, + kernel.len() + ); + + #[cfg(feature = "parallel")] + { + if big_n >= 1024 { + return (0..big_n) + .into_par_iter() + .fold( + || FieldElement::::zero(), + |acc, i| acc + &values[i] * &kernel[i], + ) + .reduce(|| FieldElement::::zero(), |a, b| a + b); + } + } + + // Sequential fallback + let mut result = FieldElement::::zero(); + for i in 0..big_n { + result = &result + &(&values[i] * &kernel[i]); + } + result +} + +#[cfg(test)] +mod tests { + use super::*; + use math::field::goldilocks::GoldilocksField; + + type FE = FieldElement; + + #[test] + fn test_compute_lagrange_kernel_n2() { + // n=2, r = (3, 7) + // s[0] = (1-3)*(1-7) = (-2)*(-6) = 12 + // s[1] = 3*(1-7) = 3*(-6) = -18 + // s[2] = (1-3)*7 = (-2)*7 = -14 + // s[3] = 3*7 = 21 + let r = vec![FE::from(3u64), FE::from(7u64)]; + let s = compute_lagrange_kernel(&r); + + assert_eq!(s.len(), 4); + + // In Goldilocks (p = 2^64 - 2^32 + 1), negative values wrap around. + // 12 mod p = 12 + assert_eq!(s[0], FE::from(12u64)); + // -18 mod p = p - 18 + let neg_18 = FE::zero() - FE::from(18u64); + assert_eq!(s[1], neg_18); + // -14 mod p = p - 14 + let neg_14 = FE::zero() - FE::from(14u64); + assert_eq!(s[2], neg_14); + // 21 + assert_eq!(s[3], FE::from(21u64)); + + // Verify partition of unity: sum = 12 + (-18) + (-14) + 21 = 1 + let sum = &(&(&s[0] + &s[1]) + &s[2]) + &s[3]; + assert_eq!(sum, FE::one()); + } + + #[test] + fn test_lagrange_kernel_partition_of_unity() { + // For any random r, the sum of kernel values must equal 1. + // Test with several different r vectors. + for &(n, r_vals) in &[ + (1usize, &[42u64][..]), + (2, &[100, 200][..]), + (3, &[5, 17, 99][..]), + (4, &[1, 2, 3, 4][..]), + ] { + let r: Vec = r_vals.iter().map(|&v| FE::from(v)).collect(); + let s = compute_lagrange_kernel(&r); + assert_eq!(s.len(), 1 << n); + + let mut sum = FE::zero(); + for val in &s { + sum = &sum + val; + } + assert_eq!( + sum, + FE::one(), + "Partition of unity failed for n={}, r={:?}", + n, + r_vals + ); + } + } + + #[test] + fn test_eval_mle_linear() { + // f(x1) = 3*x1 + 5 on {0, 1} + // f(0) = 5, f(1) = 8 + // MLE at r: 5*(1-r) + 8*r = 5 + 3r + let values = vec![FE::from(5u64), FE::from(8u64)]; + + // Test at r = 0: should be 5 + let result = eval_mle(&values, &[FE::from(0u64)]); + assert_eq!(result, FE::from(5u64)); + + // Test at r = 1: should be 8 + let result = eval_mle(&values, &[FE::from(1u64)]); + assert_eq!(result, FE::from(8u64)); + + // Test at r = 10: should be 5 + 3*10 = 35 + let result = eval_mle(&values, &[FE::from(10u64)]); + assert_eq!(result, FE::from(35u64)); + + // Test at r = 100: should be 5 + 3*100 = 305 + let result = eval_mle(&values, &[FE::from(100u64)]); + assert_eq!(result, FE::from(305u64)); + } + + #[test] + fn test_eval_mle_2var() { + // f(x1, x2) = 2*x1*x2 + 3*x1 + x2 + 7 + // Evaluate on {0,1}^2 (bit ordering: x1 = bit 0, x2 = bit 1): + // i=0 (x1=0, x2=0): 0 + 0 + 0 + 7 = 7 + // i=1 (x1=1, x2=0): 0 + 3 + 0 + 7 = 10 + // i=2 (x1=0, x2=1): 0 + 0 + 1 + 7 = 8 + // i=3 (x1=1, x2=1): 2 + 3 + 1 + 7 = 13 + let values = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(8u64), + FE::from(13u64), + ]; + + // Verify at Boolean points first (should recover original values) + assert_eq!( + eval_mle(&values, &[FE::from(0u64), FE::from(0u64)]), + FE::from(7u64) + ); + assert_eq!( + eval_mle(&values, &[FE::from(1u64), FE::from(0u64)]), + FE::from(10u64) + ); + assert_eq!( + eval_mle(&values, &[FE::from(0u64), FE::from(1u64)]), + FE::from(8u64) + ); + assert_eq!( + eval_mle(&values, &[FE::from(1u64), FE::from(1u64)]), + FE::from(13u64) + ); + + // Evaluate at (x1=5, x2=3): + // f(5, 3) = 2*5*3 + 3*5 + 3 + 7 = 30 + 15 + 3 + 7 = 55 + let result = eval_mle(&values, &[FE::from(5u64), FE::from(3u64)]); + assert_eq!(result, FE::from(55u64)); + + // Evaluate at (x1=2, x2=4): + // f(2, 4) = 2*2*4 + 3*2 + 4 + 7 = 16 + 6 + 4 + 7 = 33 + let result = eval_mle(&values, &[FE::from(2u64), FE::from(4u64)]); + assert_eq!(result, FE::from(33u64)); + } +} diff --git a/crypto/stark/src/lib.rs b/crypto/stark/src/lib.rs index 41089a9e5..9c29e6d96 100644 --- a/crypto/stark/src/lib.rs +++ b/crypto/stark/src/lib.rs @@ -7,12 +7,15 @@ pub mod domain; pub mod examples; pub mod frame; pub mod fri; +pub mod gkr; pub mod grinding; #[cfg(feature = "instruments")] pub mod instruments; +pub mod lagrange_kernel; pub mod lookup; pub mod proof; pub mod prover; +pub mod sumcheck; pub mod table; pub mod trace; pub mod traits; diff --git a/crypto/stark/src/lookup.rs b/crypto/stark/src/lookup.rs index 920a71980..d84827f99 100644 --- a/crypto/stark/src/lookup.rs +++ b/crypto/stark/src/lookup.rs @@ -2,6 +2,9 @@ use std::collections::HashMap; use std::marker::PhantomData; +#[cfg(feature = "parallel")] +use rayon::prelude::*; + use crate::{ constraints::{ boundary::{BoundaryConstraint, BoundaryConstraints}, @@ -18,8 +21,6 @@ use math::field::{ element::FieldElement, traits::{IsFFTField, IsField, IsPrimeField, IsSubFieldOf}, }; -#[cfg(feature = "parallel")] -use rayon::prelude::{IntoParallelIterator, ParallelIterator}; // ============================================================================= // Shift Constants for Type Combining @@ -101,24 +102,32 @@ pub const LOGUP_CHALLENGE_Z: usize = 0; /// Used as the base for linear combination of row values. pub const LOGUP_CHALLENGE_ALPHA: usize = 1; -/// Number of challenges required by the LogUp protocol. +/// Number of challenges required by the LogUp protocol (z and alpha). +/// The gamma challenge is sampled separately after GKR for soundness. pub const LOGUP_NUM_CHALLENGES: usize = 2; -/// Split N interactions into committed batched pairs and absorbed remainder. -/// -/// Returns `(num_committed_pairs, absorbed_count)` where: -/// - Committed pairs get dedicated auxiliary term columns (2 interactions per column) -/// - Absorbed interactions (1 or 2) are folded into the accumulated constraint -fn split_interactions(num_interactions: usize) -> (usize, usize) { - if num_interactions <= 2 { - (0, num_interactions) - } else if num_interactions % 2 == 1 { - ((num_interactions - 1) / 2, 1) - } else { - ((num_interactions - 2) / 2, 2) - } +/// Index of the `gamma` (γ) challenge in the per-table rap_challenges vector. +/// Used for batching column claims in the bridge running sum. +/// Sampled after GKR on the main transcript (not during Phase B). +pub const LOGUP_CHALLENGE_GAMMA: usize = 2; + +/// Index of the bridge offset (target/N) in the per-table rap_challenges vector. +/// This is a derived value, not a random challenge. +pub const LOGUP_BRIDGE_OFFSET_IDX: usize = 3; + +/// Start index of precomputed gamma powers in the per-table rap_challenges vector. +/// rap_challenges[LOGUP_GAMMA_POWERS_START + j] = γ^j for j = 0, 1, ..., K-1. +pub const LOGUP_GAMMA_POWERS_START: usize = 4; + +/// Start index of GKR random_point coordinates in rap_challenges. +/// After gamma_powers[0..K], we append random_point[0..n]. +/// The actual index is LOGUP_GAMMA_POWERS_START + K where K = number of distinct column indices. +/// Use `logup_random_point_start(interactions)` to compute the concrete index. +pub fn logup_random_point_start(interactions: &[BusInteraction]) -> usize { + LOGUP_GAMMA_POWERS_START + extract_column_indices(interactions).len() } + // ============================================================================= // Bus Types // ============================================================================= @@ -833,44 +842,16 @@ impl< auxiliary_trace_build_data: AuxiliaryTraceBuildData, proof_options: &ProofOptions, step_size: usize, - mut transition_constraints: Vec>>, + transition_constraints: Vec>>, ) -> Self { let num_interactions = auxiliary_trace_build_data.interactions.len(); - // Split interactions: committed pairs get term columns, last 1-2 are absorbed - let (num_committed_pairs, absorbed_count) = split_interactions(num_interactions); - let absorbed = - auxiliary_trace_build_data.interactions[num_interactions - absorbed_count..].to_vec(); - - // Create batched term constraints for committed pairs only - for pair_idx in 0..num_committed_pairs { - let constraint = LookupBatchedTermConstraint::new( - auxiliary_trace_build_data.interactions[pair_idx * 2].clone(), - auxiliary_trace_build_data.interactions[pair_idx * 2 + 1].clone(), - pair_idx, - transition_constraints.len(), - ); - transition_constraints.push(Box::new(constraint)); - } - - let num_term_columns = num_committed_pairs; - - // Add the accumulated constraint with absorbed interactions - if num_interactions > 0 { - let accumulated_constraint = LookupAccumulatedConstraint::new( - transition_constraints.len(), - num_term_columns, - absorbed, - ); - transition_constraints.push(Box::new(accumulated_constraint)); - } - - // Layout: num_committed_pairs term columns + 1 accumulated = ⌈N/2⌉ - let num_aux_columns = if num_interactions > 0 { - num_term_columns + 1 - } else { - 0 - }; + // LogUp-GKR: aux trace has 2 columns if interactions exist: + // - Column 0: Lagrange kernel (l) + // - Column 1: Bridge running sum (σ) + // The GKR sub-protocol replaces the old accumulated/term constraints. + // A single LookupBridgeSumConstraint enforces the bridge. + let num_aux_columns = if num_interactions > 0 { 2 } else { 0 }; let trace_layout = (num_main_columns, num_aux_columns); // Compute max bus elements across all interactions for alpha power count @@ -881,19 +862,31 @@ impl< .max() .unwrap_or(0); + // Add bridge running sum constraint for LogUp-GKR tables + let mut all_constraints = transition_constraints; + if num_interactions > 0 { + let column_indices = + extract_column_indices(&auxiliary_trace_build_data.interactions); + let bridge_constraint = LookupBridgeSumConstraint { + constraint_idx: all_constraints.len(), + column_indices, + }; + all_constraints.push(Box::new(bridge_constraint)); + } + // Create context let context = AirContext { proof_options: proof_options.clone(), trace_columns: trace_layout.0 + trace_layout.1, transition_offsets: vec![0, 1], - num_transition_constraints: transition_constraints.len(), + num_transition_constraints: all_constraints.len(), }; Self { context, step_size, trace_layout, - transition_constraints, + transition_constraints: all_constraints, auxiliary_trace_build_data, boundary_constraint_builder: PhantomData, preprocessed_commitment: None, @@ -976,6 +969,10 @@ where !self.auxiliary_trace_build_data.interactions.is_empty() } + fn bus_interactions(&self) -> &[BusInteraction] { + &self.auxiliary_trace_build_data.interactions + } + fn max_bus_elements(&self) -> usize { self.max_bus_elements } @@ -1003,112 +1000,19 @@ where fn build_auxiliary_trace( &self, trace: &mut TraceTable, - challenges: &[FieldElement], + _challenges: &[FieldElement], ) -> Option> { - // Allocate aux table if not already present + // LogUp-GKR: auxiliary trace has 2 columns: + // - Column 0: Lagrange kernel (l) — filled by prover from GKR random point + // - Column 1: Bridge running sum (σ) — filled by prover after γ sampling + // Here we just allocate the columns. let (_, num_aux_columns) = self.trace_layout(); if num_aux_columns > 0 && trace.num_aux_columns == 0 { trace.allocate_aux_table(num_aux_columns); } - let num_interactions = self.auxiliary_trace_build_data.interactions.len(); - - if num_interactions == 0 { - return None; - } - - // Clone main columns once (shared across all interactions) - let main_segment_cols = trace.columns_main(); - let trace_len = trace.num_rows(); - let table_name = self.name.as_deref().unwrap_or("UNKNOWN"); - - // Split interactions: committed pairs get term columns, last 1-2 are absorbed (virtual) - let (num_committed_pairs, absorbed_count) = split_interactions(num_interactions); - - // Compute committed term columns in parallel (batched pairs only) - #[cfg(feature = "parallel")] - let committed_columns: Vec>> = (0..num_committed_pairs) - .into_par_iter() - .map(|i| { - compute_logup_batched_term_column( - &self.auxiliary_trace_build_data.interactions[i * 2], - &self.auxiliary_trace_build_data.interactions[i * 2 + 1], - &main_segment_cols, - trace_len, - challenges, - table_name, - ) - }) - .collect(); - #[cfg(not(feature = "parallel"))] - let committed_columns: Vec>> = (0..num_committed_pairs) - .map(|i| { - compute_logup_batched_term_column( - &self.auxiliary_trace_build_data.interactions[i * 2], - &self.auxiliary_trace_build_data.interactions[i * 2 + 1], - &main_segment_cols, - trace_len, - challenges, - table_name, - ) - }) - .collect(); - - // Compute virtual column for absorbed interactions (NOT written to trace) - let virtual_column = if absorbed_count == 2 { - compute_logup_batched_term_column( - &self.auxiliary_trace_build_data.interactions[num_interactions - 2], - &self.auxiliary_trace_build_data.interactions[num_interactions - 1], - &main_segment_cols, - trace_len, - challenges, - table_name, - ) - } else { - compute_logup_term_column( - &self.auxiliary_trace_build_data.interactions[num_interactions - 1], - &main_segment_cols, - trace_len, - challenges, - table_name, - ) - }; - - // Write only committed columns to trace - for (col_idx, col_data) in committed_columns.iter().enumerate() { - for (row, value) in col_data.iter().enumerate() { - trace.set_aux(row, col_idx, value.clone()); - } - } - - #[cfg(feature = "debug-checks")] - let (per_bus_sums, per_bus_sender_sums, per_bus_receiver_sums) = - compute_debug_bus_sums_batched( - &self.auxiliary_trace_build_data.interactions, - &main_segment_cols, - trace_len, - challenges, - table_name, - ); - - // Build accumulated from all columns (committed + virtual) - let mut all_columns = committed_columns; - all_columns.push(virtual_column); - let acc_col_idx = num_committed_pairs; // accumulated column in trace follows committed columns - let table_contribution = - build_accumulated_column_from_terms(acc_col_idx, &all_columns, trace); - - Some(BusPublicInputs { - table_contribution, - #[cfg(feature = "debug-checks")] - per_bus_sums, - #[cfg(feature = "debug-checks")] - per_bus_sender_sums, - #[cfg(feature = "debug-checks")] - per_bus_receiver_sums, - #[cfg(feature = "debug-checks")] - table_name: self.name.clone().unwrap_or_else(|| "UNKNOWN".to_string()), - }) + // No BusPublicInputs needed for GKR path — the GKR result replaces it. + None } fn build_rap_challenges( @@ -1125,19 +1029,26 @@ where pub_inputs: &Self::PublicInputs, rap_challenges: &[FieldElement], _bus_public_inputs: Option<&BusPublicInputs>, - trace_length: usize, + _trace_length: usize, ) -> BoundaryConstraints { let mut boundary_constraints = B::boundary_constraints(pub_inputs, rap_challenges); - // Pin acc[N-1] = 0 to remove the constant-shift degree of freedom - // in the circular transition constraint. - if !self.auxiliary_trace_build_data.interactions.is_empty() { - let acc_col_idx = self.trace_layout.1 - 1; // last aux column = accumulated - boundary_constraints.push(BoundaryConstraint::new_aux( - acc_col_idx, - trace_length - 1, - FieldElement::zero(), - )); + // LogUp-GKR: boundary constraint on Lagrange kernel column. + // l[0] = eq(bits(0), r) = prod_{j=0}^{n-1} (1 - r_j) + // where r_j are the GKR random point coordinates stored in rap_challenges. + if self.has_trace_interaction() { + let k = extract_column_indices(&self.auxiliary_trace_build_data.interactions).len(); + let rp_start = LOGUP_GAMMA_POWERS_START + k; + if rap_challenges.len() > rp_start { + let n = rap_challenges.len() - rp_start; + let mut l0_expected = FieldElement::::one(); + for j in 0..n { + l0_expected *= + FieldElement::::one() - &rap_challenges[rp_start + j]; + } + // Aux column 0 is the Lagrange kernel; constrain l[0] = prod(1 - r_j) + boundary_constraints.push(BoundaryConstraint::new_aux(0, 0, l0_expected)); + } } BoundaryConstraints::from_constraints(boundary_constraints) @@ -1364,7 +1275,7 @@ where /// /// This is a pure function that takes shared main columns and returns the computed column, /// enabling parallel computation across interactions within a table. -#[allow(clippy::needless_range_loop)] +#[allow(dead_code, clippy::needless_range_loop)] fn compute_logup_term_column( table_interaction: &BusInteraction, main_segment_cols: &[Vec>], @@ -1523,6 +1434,64 @@ where .collect() } + +// ============================================================================= +// LogUp-GKR Bridge Running Sum +// ============================================================================= + +/// Extract sorted distinct main column indices from bus interactions. +/// +/// These are the columns referenced by `column_claims` in the GKR result. +/// The order must be consistent between the trace builder and the constraint evaluator. +pub(crate) fn extract_column_indices(interactions: &[BusInteraction]) -> Vec { + let mut seen_cols = std::collections::HashSet::new(); + for inter in interactions { + for val in &inter.values { + for col_idx in val.column_indices() { + seen_cols.insert(col_idx); + } + } + match &inter.multiplicity { + Multiplicity::One => {} + Multiplicity::Column(c) => { + seen_cols.insert(*c); + } + Multiplicity::Sum(a, b) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + } + Multiplicity::Negated(c) => { + seen_cols.insert(*c); + } + Multiplicity::Diff(a, b) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + } + Multiplicity::Sum3(a, b, c) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + seen_cols.insert(*c); + } + Multiplicity::Linear(terms) => { + for term in terms { + match term { + LinearTerm::Column { column, .. } => { + seen_cols.insert(*column); + } + LinearTerm::ColumnUnsigned { column, .. } => { + seen_cols.insert(*column); + } + LinearTerm::Constant(_) => {} + } + } + } + } + } + let mut col_indices: Vec = seen_cols.into_iter().collect(); + col_indices.sort_unstable(); + col_indices +} + /// Computes a batched term column for two interactions sharing one aux column. /// /// Each row contains: `term[i] = sign_a * m_a[i] / fp_a[i] + sign_b * m_b[i] / fp_b[i]` @@ -1672,111 +1641,247 @@ where .collect() } -/// Builds the circular accumulated column from pre-computed term columns. +/// Compute the bridge offset (target/N) and gamma powers from column claims. /// -/// For the circular constraint: acc[(i+1) mod N] - acc[i] - terms[(i+1) mod N] + L/N = 0 -/// We build: acc[0] = terms[0] - L/N, acc[i] = acc[i-1] + terms[i] - L/N -/// Result: acc[N-1] = L - N*(L/N) = 0 +/// Returns (bridge_offset, gamma_powers) where: +/// - bridge_offset = (Σ_j γ^j · c_j) / N +/// - gamma_powers = [γ^0, γ^1, ..., γ^{K-1}] /// -/// Returns L (table_contribution = sum of all terms across all rows). -fn build_accumulated_column_from_terms( - acc_column_idx: usize, - term_columns: &[Vec>], - trace: &mut TraceTable, -) -> FieldElement +/// Both prover and verifier call this to derive the same values. +pub fn compute_bridge_params( + column_claims: &[(usize, FieldElement)], + gamma: &FieldElement, + trace_len: usize, +) -> (FieldElement, Vec>) { + let k = column_claims.len(); + let gamma_powers = compute_alpha_powers(gamma, k); + + let mut target = FieldElement::::zero(); + for ((_, c_j), gp) in column_claims.iter().zip(gamma_powers.iter()) { + target += c_j * gp; + } + + let n_inv = FieldElement::::from(trace_len as u64).inv().unwrap(); + let bridge_offset = &target * &n_inv; + + (bridge_offset, gamma_powers) +} + +/// Extend rap_challenges with bridge parameters (γ, bridge_offset, gamma_powers, random_point). +/// +/// After calling this, the rap_challenges vector has: +/// - [0] = z, [1] = α (original) +/// - [2] = γ +/// - [3] = bridge_offset (target/N) +/// - [4..4+K] = γ^0, γ^1, ..., γ^{K-1} +/// - [4+K..4+K+n] = random_point[0], ..., random_point[n-1] +pub fn extend_rap_challenges_with_bridge( + rap_challenges: &mut Vec>, + column_claims: &[(usize, FieldElement)], + gamma: &FieldElement, + trace_len: usize, + random_point: &[FieldElement], +) { + let (bridge_offset, gamma_powers) = compute_bridge_params(column_claims, gamma, trace_len); + rap_challenges.push(gamma.clone()); // index 2 + rap_challenges.push(bridge_offset); // index 3 + for gp in gamma_powers { + rap_challenges.push(gp); // indices 4, 5, ... + } + for rp in random_point { + rap_challenges.push(rp.clone()); // indices 4+K, 4+K+1, ... + } +} + +/// Transition constraint for the bridge running sum column (σ). +/// +/// Enforces the circular constraint: +/// σ_next - σ_curr - l_curr · batched_curr + bridge_offset = 0 +/// +/// where: +/// - σ is the running sum (aux column 1) +/// - l is the Lagrange kernel (aux column 0) +/// - batched_curr = Σ_j γ^j · col_j_curr (from main trace columns) +/// - bridge_offset = (Σ_j γ^j · c_j) / N (from rap_challenges) +/// +/// The circular constraint (end_exemptions=0) telescopes to: +/// Σ_{i=0}^{N-1} l[i] · batched[i] = target +/// which, by γ-batching (Schwartz-Zippel), proves all individual claims +/// = c_j with high probability. +pub struct LookupBridgeSumConstraint { + constraint_idx: usize, + /// Sorted distinct main column indices from bus interactions + column_indices: Vec, +} + +impl TransitionConstraint for LookupBridgeSumConstraint where - F: IsFFTField + IsSubFieldOf + Send + Sync, + F: IsSubFieldOf + IsFFTField + Send + Sync, E: IsField + Send + Sync, { - if term_columns.is_empty() { - return FieldElement::zero(); + fn degree(&self) -> usize { + 2 // l_curr * batched_curr } - let trace_len = term_columns[0].len(); - // Compute L = sum of all terms across all rows - let mut table_contribution = FieldElement::::zero(); - for row in 0..trace_len { - for col in term_columns { - table_contribution = &table_contribution + &col[row]; - } + fn constraint_idx(&self) -> usize { + self.constraint_idx + } + + fn end_exemptions(&self) -> usize { + 0 // circular: checked on all N rows (including wrap-around) } - // offset_per_row = L / N - let n = FieldElement::::from(trace_len as u64); - let offset_per_row = &table_contribution * n.inv().unwrap(); + fn evaluate( + &self, + evaluation_context: &TransitionEvaluationContext, + transition_evaluations: &mut [FieldElement], + ) { + match evaluation_context { + TransitionEvaluationContext::Prover { + frame, + rap_challenges, + .. + } => { + let bridge_offset = &rap_challenges[LOGUP_BRIDGE_OFFSET_IDX]; - // Build circular accumulated column - let mut accumulated = FieldElement::::zero(); - for row in 0..trace_len { - let mut row_sum = FieldElement::::zero(); - for col in term_columns { - row_sum = row_sum + &col[row]; + let step0 = frame.get_evaluation_step(0); + let step1 = frame.get_evaluation_step(1); + + // σ (aux column 1) + let sigma_curr = step0.get_aux_evaluation_element(0, 1); + let sigma_next = step1.get_aux_evaluation_element(0, 1); + + // l (aux column 0) + let l_curr = step0.get_aux_evaluation_element(0, 0); + + // batched_curr = Σ_j γ^j · col_j_curr using precomputed gamma powers + let mut batched = FieldElement::::zero(); + for (j, &col_idx) in self.column_indices.iter().enumerate() { + let gamma_j = &rap_challenges[LOGUP_GAMMA_POWERS_START + j]; + let col_val = step0.get_main_evaluation_element(0, col_idx); + // F×E→E: base field column × extension field gamma power + batched += col_val * gamma_j; + } + + // σ_next - σ_curr - l_curr * batched + bridge_offset + transition_evaluations[self.constraint_idx] = + sigma_next - sigma_curr - l_curr * &batched + bridge_offset; + } + TransitionEvaluationContext::Verifier { + frame, + rap_challenges, + .. + } => { + let bridge_offset = &rap_challenges[LOGUP_BRIDGE_OFFSET_IDX]; + + let step0 = frame.get_evaluation_step(0); + let step1 = frame.get_evaluation_step(1); + + let sigma_curr = step0.get_aux_evaluation_element(0, 1); + let sigma_next = step1.get_aux_evaluation_element(0, 1); + let l_curr = step0.get_aux_evaluation_element(0, 0); + + let mut batched = FieldElement::::zero(); + for (j, &col_idx) in self.column_indices.iter().enumerate() { + let gamma_j = &rap_challenges[LOGUP_GAMMA_POWERS_START + j]; + // In verifier path, main cols are also in E + let col_val = step0.get_main_evaluation_element(0, col_idx); + batched += col_val * gamma_j; + } + + transition_evaluations[self.constraint_idx] = + sigma_next - sigma_curr - l_curr * &batched + bridge_offset; + } } - accumulated = &accumulated + &row_sum - &offset_per_row; - trace.set_aux(row, acc_column_idx, accumulated.clone()); } +} + - table_contribution +// ============================================================================= +// LogUp-GKR Leaf Fraction Computation +// ============================================================================= + +/// Computes the multiplicity for a single interaction at a single row. +#[inline] +fn compute_multiplicity_at_row( + interaction: &BusInteraction, + main_segment_cols: &[Vec>], + row: usize, +) -> FieldElement { + match &interaction.multiplicity { + Multiplicity::One => FieldElement::one(), + Multiplicity::Column(col) => main_segment_cols[*col][row].clone(), + Multiplicity::Sum(col_a, col_b) => { + &main_segment_cols[*col_a][row] + &main_segment_cols[*col_b][row] + } + Multiplicity::Negated(col) => FieldElement::::one() - &main_segment_cols[*col][row], + Multiplicity::Diff(col_a, col_b) => { + &main_segment_cols[*col_a][row] - &main_segment_cols[*col_b][row] + } + Multiplicity::Sum3(col_a, col_b, col_c) => { + &main_segment_cols[*col_a][row] + + &main_segment_cols[*col_b][row] + + &main_segment_cols[*col_c][row] + } + Multiplicity::Linear(terms) => { + let mut result = FieldElement::::zero(); + for term in terms { + match *term { + LinearTerm::Column { + coefficient, + column, + } => { + let coeff = FieldElement::::from(coefficient); + result += &main_segment_cols[column][row] * coeff; + } + LinearTerm::ColumnUnsigned { + coefficient, + column, + } => { + let coeff = FieldElement::::from(coefficient); + result += &main_segment_cols[column][row] * coeff; + } + LinearTerm::Constant(value) => { + result += FieldElement::::from(value); + } + } + } + result + } + } } -/// Sum per-interaction contributions by bus_id for debug reporting. -/// -/// With batched term columns, we can't read individual interaction sums from -/// the trace anymore (each column holds the sum of two interactions). Instead, -/// we compute each interaction's sum from its raw term column. -#[cfg(feature = "debug-checks")] -#[allow(clippy::type_complexity)] -fn compute_debug_bus_sums_batched( - interactions: &[BusInteraction], +/// Computes the fingerprint for a single interaction at a single row. +#[inline] +fn compute_fingerprint_at_row( + interaction: &BusInteraction, main_segment_cols: &[Vec>], - trace_len: usize, - challenges: &[FieldElement], - table_name: &str, -) -> ( - HashMap>, - HashMap>, - HashMap>, -) + row: usize, + z: &FieldElement, + alpha_powers: &[FieldElement], +) -> FieldElement where - F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, - E: IsField + Send + Sync, + F: IsField + IsSubFieldOf + IsPrimeField, + E: IsField, { - let mut bus_sums: HashMap> = HashMap::new(); - let mut sender_sums: HashMap> = HashMap::new(); - let mut receiver_sums: HashMap> = HashMap::new(); - - // Compute each interaction's individual term column for summing - for interaction in interactions.iter() { - let individual_terms = compute_logup_term_column( - interaction, + let shifts = PackingShifts::::new(); + let bus_id_f = FieldElement::::from(interaction.bus_id); + let mut linear_combination = &bus_id_f * &alpha_powers[0]; + let mut alpha_offset = 1; + for bv in &interaction.values { + let consumed = bv.accumulate_fingerprint( main_segment_cols, - trace_len, - challenges, - table_name, + row, + alpha_powers, + alpha_offset, + &mut linear_combination, + &shifts, ); - let col_sum: FieldElement = individual_terms - .iter() - .fold(FieldElement::zero(), |acc, x| acc + x); - - *bus_sums - .entry(interaction.bus_id) - .or_insert(FieldElement::zero()) += col_sum.clone(); - - if interaction.is_sender { - *sender_sums - .entry(interaction.bus_id) - .or_insert(FieldElement::zero()) += col_sum; - } else { - let entry = receiver_sums - .entry(interaction.bus_id) - .or_insert(FieldElement::zero()); - *entry = entry.clone() - col_sum; - } + alpha_offset += consumed; } - (bus_sums, sender_sums, receiver_sums) + z - &linear_combination } -/// Computes multiplicity for an interaction from a `TableView`. fn compute_multiplicity_from_step, B: IsField>( step: &TableView, multiplicity: &Multiplicity, @@ -2157,3 +2262,834 @@ where } } } + +/// Computes the leaf fractions for the GKR summation tree from a table's bus +/// interactions and main trace. +/// +/// For each row `i` in `0..trace_len`, this function combines all K interactions +/// into a single fraction `N(i) / D(i)` where: +/// +/// - `D(i) = Π_k fp_k(i)` (product of all fingerprints at row i) +/// - `N(i) = Σ_k sign_k * m_k(i) * Π_{j≠k} fp_j(i)` (cross-terms) +/// +/// This is computed iteratively: starting with fraction 0/1, for each interaction k +/// the running fraction is updated via cross-multiplication: +/// `n_new = n_old * fp_k + sign_k * m_k * d_old` +/// `d_new = d_old * fp_k` +/// +/// Returns `(numerators, denominators)` each of length `trace_len`. +/// +/// # Arguments +/// * `interactions` - The bus interactions for this table +/// * `main_segment_cols` - Column-major main trace data: `main_segment_cols[col][row]` +/// * `trace_len` - Number of rows in the trace +/// * `challenges` - LogUp challenges `[z, alpha, ...]` +pub fn compute_logup_leaf_fractions( + interactions: &[BusInteraction], + main_segment_cols: &[Vec>], + trace_len: usize, + challenges: &[FieldElement], +) -> (Vec>, Vec>) +where + F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, + E: IsField + Send + Sync, +{ + assert!( + !interactions.is_empty(), + "Must have at least one interaction" + ); + + let z = &challenges[LOGUP_CHALLENGE_Z]; + let alpha = &challenges[LOGUP_CHALLENGE_ALPHA]; + + // Find max bus elements across all interactions for alpha power precomputation + let max_bus_elements = interactions + .iter() + .map(|inter| inter.num_bus_elements()) + .max() + .unwrap(); + let alpha_powers = compute_alpha_powers(alpha, max_bus_elements); + + // Precompute signs (cheap, interaction-count dependent) + let all_signs: Vec> = interactions + .iter() + .map(|inter| { + if inter.is_sender { + FieldElement::::one() + } else { + -FieldElement::::one() + } + }) + .collect(); + + // Fused parallel computation: for each row, compute fingerprints, multiplicities, + // and cross-multiply all interactions in a single pass. + #[cfg(feature = "parallel")] + let iter = (0..trace_len).into_par_iter(); + #[cfg(not(feature = "parallel"))] + let iter = 0..trace_len; + + let (numerators, denominators): (Vec<_>, Vec<_>) = iter + .map(|row| { + let mut running_n = FieldElement::::zero(); + let mut running_d = FieldElement::::one(); + + for (k, inter) in interactions.iter().enumerate() { + let fp_k = compute_fingerprint_at_row(inter, main_segment_cols, row, z, &alpha_powers); + let m_k = compute_multiplicity_at_row(inter, main_segment_cols, row); + + // F * E -> E multiplication (no to_extension) + let new_n = &running_n * &fp_k + &m_k * &all_signs[k] * &running_d; + let new_d = &running_d * &fp_k; + + running_n = new_n; + running_d = new_d; + } + + (running_n, running_d) + }) + .unzip(); + + (numerators, denominators) +} + +// ============================================================================= +// LogUp-GKR Integration +// ============================================================================= + +use crate::gkr::{build_summation_tree, gkr_prove, GkrProof}; +use crate::lagrange_kernel::{compute_lagrange_kernel, eval_mle_base_with_kernel}; +use crypto::fiat_shamir::is_transcript::IsTranscript; + +/// Result of running the LogUp-GKR sub-protocol for a single table. +/// +/// Contains the GKR proof, the random evaluation point, leaf-level claims, +/// and MLE claims for each distinct main trace column used in bus interactions. +#[derive(Debug, Clone)] +pub struct LogUpGkrResult { + /// Total table contribution (claimed_sum from GKR root = sum of all fractions). + pub table_contribution: FieldElement, + /// The complete GKR proof for the summation tree. + pub gkr_proof: GkrProof, + /// The random evaluation point produced by the GKR protocol (length = log2(trace_len)). + pub random_point: Vec>, + /// Claimed MLE evaluation of the leaf numerator at the random point. + pub n_claim: FieldElement, + /// Claimed MLE evaluation of the leaf denominator at the random point. + pub d_claim: FieldElement, + /// MLE claims for each distinct main trace column used in bus interactions. + /// Each entry is (column_index, MLE evaluation at random_point). + pub column_claims: Vec<(usize, FieldElement)>, + /// Pre-computed Lagrange kernel at `random_point`, for reuse in aux trace construction. + pub lagrange_kernel: Vec>, +} + +/// Verifies that column_claims are consistent with the GKR output (n_claim, d_claim). +/// +/// The GKR protocol outputs `(random_point, n_claim, d_claim)` where `n_claim` and +/// `d_claim` are the claimed MLE evaluations of the leaf numerator and denominator +/// at `random_point`. The prover also provides `column_claims` which are MLE +/// evaluations of individual main trace columns at the same point. +/// +/// For single-interaction tables: the leaf numerator and denominator are linear +/// functions of column values (n = sign * m, d = fp), so their MLEs can be exactly +/// reconstructed from column_claims. This gives a direct equality check. +/// +/// For multi-interaction tables: the leaf fraction involves products of per-interaction +/// fingerprints and multiplicities (cross-multiplication), making it a nonlinear +/// function of column values. Since MLE does not preserve products, the direct +/// reconstruction from column_claims differs from the true MLE values. For these +/// tables, soundness of column_claims is guaranteed by the bridge running sum +/// constraint (which is verified as part of the STARK proof). We still verify +/// structural completeness (all referenced columns are present in column_claims). +/// +/// # Arguments +/// * `n_claim` - Claimed MLE evaluation of leaf numerator at random_point (from GKR) +/// * `d_claim` - Claimed MLE evaluation of leaf denominator at random_point (from GKR) +/// * `column_claims` - `(column_index, claimed_value)` pairs from the proof +/// * `interactions` - The AIR's bus interactions +/// * `challenges` - LogUp challenges `[z, alpha]` +/// +/// # Returns +/// `true` if verification passes, `false` otherwise. +pub fn reconstruct_and_verify_gkr_claims( + n_claim: &FieldElement, + d_claim: &FieldElement, + column_claims: &[(usize, FieldElement)], + interactions: &[BusInteraction], + challenges: &[FieldElement], +) -> bool { + // Build a map from column index to claimed MLE value + let claim_map: std::collections::HashMap> = column_claims + .iter() + .map(|(col_idx, val)| (*col_idx, val)) + .collect(); + + // Verify structural completeness: all columns referenced by interactions + // must be present in column_claims. + for inter in interactions { + for val in &inter.values { + for col_idx in val.column_indices() { + if !claim_map.contains_key(&col_idx) { + return false; + } + } + } + match &inter.multiplicity { + Multiplicity::One => {} + Multiplicity::Column(c) => { + if !claim_map.contains_key(c) { + return false; + } + } + Multiplicity::Sum(a, b) => { + if !claim_map.contains_key(a) || !claim_map.contains_key(b) { + return false; + } + } + Multiplicity::Negated(c) => { + if !claim_map.contains_key(c) { + return false; + } + } + Multiplicity::Diff(a, b) => { + if !claim_map.contains_key(a) || !claim_map.contains_key(b) { + return false; + } + } + Multiplicity::Sum3(a, b, c) => { + if !claim_map.contains_key(a) || !claim_map.contains_key(b) || !claim_map.contains_key(c) { + return false; + } + } + Multiplicity::Linear(terms) => { + for term in terms { + match term { + LinearTerm::Column { column, .. } + | LinearTerm::ColumnUnsigned { column, .. } => { + if !claim_map.contains_key(column) { + return false; + } + } + LinearTerm::Constant(_) => {} + } + } + } + } + } + + let z = &challenges[LOGUP_CHALLENGE_Z]; + let alpha = &challenges[LOGUP_CHALLENGE_ALPHA]; + + // Compute enough alpha powers for the largest interaction + let max_bus_elements = interactions + .iter() + .map(|inter| inter.num_bus_elements()) + .max() + .unwrap_or(0); + let alpha_powers = compute_alpha_powers(alpha, max_bus_elements); + + // For each interaction, compute fingerprint and multiplicity from column claims, + // then accumulate into a running fraction, same as compute_logup_leaf_fractions. + let mut running_n = FieldElement::::zero(); + let mut running_d = FieldElement::::one(); + + for inter in interactions { + // Compute fingerprint: z - (bus_id * alpha^0 + sum of value contributions) + let bus_id_e = FieldElement::::from(inter.bus_id); + let mut linear_combination = &bus_id_e * &alpha_powers[0]; + let mut alpha_offset = 1; + + for bv in &inter.values { + alpha_offset += accumulate_fingerprint_from_claims( + bv, + &claim_map, + &alpha_powers, + alpha_offset, + &mut linear_combination, + ); + } + + let fp_claim = z - &linear_combination; + + // Compute multiplicity from column claims + let m_claim = multiplicity_from_claims(&inter.multiplicity, &claim_map); + + // Sign: +1 for sender, -1 for receiver + let sign = if inter.is_sender { + FieldElement::::one() + } else { + -FieldElement::::one() + }; + + // Accumulate: n_new = n_old * fp + sign * m * d_old + // d_new = d_old * fp + let new_n = &running_n * &fp_claim + &m_claim * &sign * &running_d; + let new_d = &running_d * &fp_claim; + + running_n = new_n; + running_d = new_d; + } + + // For single-interaction tables, the leaf fraction is linear in column values: + // N(i) = sign * m(i), D(i) = fp(i) + // so MLE(N)(r) and MLE(D)(r) can be exactly reconstructed from column MLEs. + // + // For multi-interaction tables, the cross-multiplication introduces nonlinear + // terms (products of fingerprints/multiplicities across interactions), so + // MLE(N)(r) != n_recon and MLE(D)(r) != d_recon in general. + // Soundness for these tables is ensured by the bridge running sum constraint. + if interactions.len() == 1 { + // Direct check: reconstructed values must match GKR output as rational numbers. + // The GKR verifier may return (n_claim, d_claim) in a different representation + // than the prover's raw (numerator, denominator) — e.g. (claimed_sum, 1) instead + // of (root_n, root_d). So we compare as rationals: running_n / running_d == n_claim / d_claim + // i.e. running_n * d_claim == n_claim * running_d. + &running_n * d_claim == n_claim * &running_d + } else { + // Multi-interaction: structural check passed above. + // The bridge constraint (verified during STARK proof) ensures column_claims + // are consistent with the committed trace. + true + } +} + +/// Accumulates the fingerprint contribution of a BusValue from column claims. +/// +/// This mirrors `BusValue::accumulate_fingerprint` but operates on the claim map +/// (MLE evaluations at the GKR random point) instead of raw trace data. +/// +/// Returns the number of alpha powers consumed. +fn accumulate_fingerprint_from_claims( + bv: &BusValue, + claim_map: &std::collections::HashMap>, + alpha_powers: &[FieldElement], + alpha_offset: usize, + acc: &mut FieldElement, +) -> usize { + match bv { + BusValue::Packed { + start_column, + packing, + } => { + // Collect column claim values for this packing + let columns: Vec> = (*start_column + ..*start_column + packing.num_columns()) + .map(|col| { + claim_map + .get(&col) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()) + }) + .collect(); + + // Use Packing::combine to get bus elements, then accumulate with alpha powers + let combined = packing.combine(&columns); + for (i, elem) in combined.iter().enumerate() { + *acc += elem * &alpha_powers[alpha_offset + i]; + } + combined.len() + } + BusValue::Linear(terms) => { + let mut result = FieldElement::::zero(); + for term in terms { + match term { + LinearTerm::Column { + coefficient, + column, + } => { + let coeff = FieldElement::::from(*coefficient); + let val = claim_map + .get(column) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + result += &val * coeff; + } + LinearTerm::ColumnUnsigned { + coefficient, + column, + } => { + let coeff = FieldElement::::from(*coefficient); + let val = claim_map + .get(column) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + result += &val * coeff; + } + LinearTerm::Constant(value) => { + result += FieldElement::::from(*value); + } + } + } + *acc += &result * &alpha_powers[alpha_offset]; + 1 + } + } +} + +/// Computes the multiplicity value from column claims at the GKR random point. +/// +/// This mirrors `compute_multiplicities_for_interaction` but operates on scalar +/// claim values instead of column vectors. +fn multiplicity_from_claims( + multiplicity: &Multiplicity, + claim_map: &std::collections::HashMap>, +) -> FieldElement { + match multiplicity { + Multiplicity::One => FieldElement::::one(), + Multiplicity::Column(c) => claim_map + .get(c) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()), + Multiplicity::Sum(a, b) => { + let va = claim_map + .get(a) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + let vb = claim_map + .get(b) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + va + vb + } + Multiplicity::Negated(c) => { + let val = claim_map + .get(c) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + FieldElement::::one() - val + } + Multiplicity::Diff(a, b) => { + let va = claim_map + .get(a) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + let vb = claim_map + .get(b) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + va - vb + } + Multiplicity::Sum3(a, b, c) => { + let va = claim_map + .get(a) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + let vb = claim_map + .get(b) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + let vc = claim_map + .get(c) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + va + vb + vc + } + Multiplicity::Linear(terms) => { + let mut result = FieldElement::::zero(); + for term in terms { + match term { + LinearTerm::Column { + coefficient, + column, + } => { + let coeff = FieldElement::::from(*coefficient); + let val = claim_map + .get(column) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + result += &val * coeff; + } + LinearTerm::ColumnUnsigned { + coefficient, + column, + } => { + let coeff = FieldElement::::from(*coefficient); + let val = claim_map + .get(column) + .cloned() + .cloned() + .unwrap_or_else(|| FieldElement::::zero()); + result += &val * coeff; + } + LinearTerm::Constant(value) => { + result += FieldElement::::from(*value); + } + } + } + result + } + } +} + +/// Run the LogUp-GKR sub-protocol for a single table's bus interactions. +/// +/// This function: +/// 1. Computes per-row leaf fractions (numerator, denominator) from interactions +/// 2. Builds a binary summation tree over the leaf fractions +/// 3. Runs the GKR protocol to prove the summation tree root +/// 4. Extracts MLE claims for each distinct main trace column at the GKR random point +/// +/// The GKR proof replaces the traditional per-row accumulated column with a +/// logarithmic-depth interactive proof, reducing auxiliary trace columns. +pub fn run_logup_gkr( + interactions: &[BusInteraction], + main_segment_cols: &[Vec>], + trace_len: usize, + challenges: &[FieldElement], + transcript: &mut impl IsTranscript, +) -> LogUpGkrResult +where + F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, + E: IsField + Send + Sync, +{ + // Step 1: Compute per-row leaf fractions + let (numerators, denominators) = + compute_logup_leaf_fractions(interactions, main_segment_cols, trace_len, challenges); + + // Step 2: Build the summation tree + let tree = build_summation_tree(numerators, denominators); + + // Step 3: Run the GKR protocol + let (gkr_proof, random_point, n_claim, d_claim) = gkr_prove(&tree, transcript); + + let table_contribution = gkr_proof.claimed_sum.clone(); + + // Step 4: Extract column claims — compute MLE at the random point for each + // distinct main trace column index referenced by any interaction. + let mut seen_cols = std::collections::HashSet::new(); + for inter in interactions { + for val in &inter.values { + for col_idx in val.column_indices() { + seen_cols.insert(col_idx); + } + } + // Also collect column indices from multiplicities + match &inter.multiplicity { + Multiplicity::One => {} + Multiplicity::Column(c) => { + seen_cols.insert(*c); + } + Multiplicity::Sum(a, b) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + } + Multiplicity::Negated(c) => { + seen_cols.insert(*c); + } + Multiplicity::Diff(a, b) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + } + Multiplicity::Sum3(a, b, c) => { + seen_cols.insert(*a); + seen_cols.insert(*b); + seen_cols.insert(*c); + } + Multiplicity::Linear(terms) => { + for term in terms { + match term { + LinearTerm::Column { column, .. } => { + seen_cols.insert(*column); + } + LinearTerm::ColumnUnsigned { column, .. } => { + seen_cols.insert(*column); + } + LinearTerm::Constant(_) => {} + } + } + } + } + } + + let mut col_indices: Vec = seen_cols.into_iter().collect(); + col_indices.sort_unstable(); + + // Compute kernel once and reuse for all column claims (and later for aux trace) + let kernel = compute_lagrange_kernel(&random_point); + + #[cfg(feature = "parallel")] + let col_iter = col_indices.into_par_iter(); + #[cfg(not(feature = "parallel"))] + let col_iter = col_indices.into_iter(); + + let column_claims: Vec<(usize, FieldElement)> = col_iter + .map(|col_idx| { + let claim = eval_mle_base_with_kernel(&main_segment_cols[col_idx], &kernel); + (col_idx, claim) + }) + .collect(); + + LogUpGkrResult { + table_contribution, + gkr_proof, + random_point, + n_claim, + d_claim, + column_claims, + lagrange_kernel: kernel, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use math::field::goldilocks::GoldilocksField; + + type F = GoldilocksField; + type FE = FieldElement; + + /// Test with 1 sender interaction, 4 rows, 1 column value, Multiplicity::One. + /// + /// For a single interaction with multiplicity 1 and sign +1 (sender): + /// N(i) = +1 * 1 = 1 + /// D(i) = fp(i) = z - (bus_id * α^0 + col_val * α^1) + #[test] + fn test_compute_logup_leaf_fractions_single_sender() { + let trace_len = 4; + + // Column data: 4 rows with values [10, 20, 30, 40] + let col0: Vec = vec![ + FE::from(10u64), + FE::from(20u64), + FE::from(30u64), + FE::from(40u64), + ]; + let main_segment_cols = vec![col0.clone()]; + + // Single sender interaction: bus_id=1, Multiplicity::One, one Direct column + let interaction = BusInteraction::sender( + 1u64, + Multiplicity::One, + Packing::Direct.columns(&[0]), + ); + let interactions = vec![interaction]; + + // Challenges: z=100, alpha=3 + let z = FE::from(100u64); + let alpha = FE::from(3u64); + let challenges = vec![z.clone(), alpha.clone()]; + + let (numerators, denominators) = + compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + + assert_eq!(numerators.len(), trace_len); + assert_eq!(denominators.len(), trace_len); + + // For each row, verify: + // fp = z - (bus_id * α^0 + col_val * α^1) = 100 - (1*1 + col_val * 3) + // numerator = 1 (multiplicity * sign) + // denominator = fp + let alpha_powers = compute_alpha_powers(&alpha, 2); // [1, 3] + + for row in 0..trace_len { + let bus_id_f = FE::from(1u64); + let linear_comb = &bus_id_f * &alpha_powers[0] + &col0[row] * &alpha_powers[1]; + let expected_fp = &z - &linear_comb; + + assert_eq!( + numerators[row], + FE::one(), + "Row {}: numerator should be 1 for sender with Multiplicity::One", + row + ); + assert_eq!( + denominators[row], expected_fp, + "Row {}: denominator should equal fingerprint", + row + ); + } + } + + /// Test with 1 receiver interaction, Multiplicity::Column, to verify sign and + /// multiplicity column extraction. + #[test] + fn test_compute_logup_leaf_fractions_single_receiver_with_column_multiplicity() { + let trace_len = 4; + + // Column 0: values for fingerprint + let col0: Vec = vec![ + FE::from(5u64), + FE::from(6u64), + FE::from(7u64), + FE::from(8u64), + ]; + // Column 1: multiplicities + let col1: Vec = vec![ + FE::from(2u64), + FE::from(0u64), + FE::from(1u64), + FE::from(3u64), + ]; + let main_segment_cols = vec![col0.clone(), col1.clone()]; + + // Receiver interaction: bus_id=0, Multiplicity from column 1 + let interaction = BusInteraction::receiver( + 0u64, + Multiplicity::Column(1), + Packing::Direct.columns(&[0]), + ); + let interactions = vec![interaction]; + + let z = FE::from(50u64); + let alpha = FE::from(7u64); + let challenges = vec![z.clone(), alpha.clone()]; + + let (numerators, denominators) = + compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + + let alpha_powers = compute_alpha_powers(&alpha, 2); + + for row in 0..trace_len { + let bus_id_f = FE::from(0u64); + let linear_comb = &bus_id_f * &alpha_powers[0] + &col0[row] * &alpha_powers[1]; + let expected_fp = &z - &linear_comb; + + // sign = -1 for receiver, so numerator = -multiplicity + let expected_num = -col1[row].clone(); + + assert_eq!( + numerators[row], expected_num, + "Row {}: numerator should be -multiplicity for receiver", + row + ); + assert_eq!( + denominators[row], expected_fp, + "Row {}: denominator should equal fingerprint", + row + ); + } + } + + /// Test with 2 interactions to verify cross-multiplication combining. + /// + /// Two interactions combined at each row: + /// fraction = sign_0 * m_0 / fp_0 + sign_1 * m_1 / fp_1 + /// = (sign_0 * m_0 * fp_1 + sign_1 * m_1 * fp_0) / (fp_0 * fp_1) + #[test] + fn test_compute_logup_leaf_fractions_two_interactions() { + let trace_len = 2; + + // Column 0: values for interaction 0 + let col0: Vec = vec![FE::from(10u64), FE::from(20u64)]; + // Column 1: values for interaction 1 + let col1: Vec = vec![FE::from(30u64), FE::from(40u64)]; + let main_segment_cols = vec![col0.clone(), col1.clone()]; + + // Interaction 0: sender, bus_id=0, Multiplicity::One, column 0 + let inter0 = BusInteraction::sender( + 0u64, + Multiplicity::One, + Packing::Direct.columns(&[0]), + ); + // Interaction 1: receiver, bus_id=1, Multiplicity::One, column 1 + let inter1 = BusInteraction::receiver( + 1u64, + Multiplicity::One, + Packing::Direct.columns(&[1]), + ); + let interactions = vec![inter0, inter1]; + + let z = FE::from(200u64); + let alpha = FE::from(5u64); + let challenges = vec![z.clone(), alpha.clone()]; + + let (numerators, denominators) = + compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + + let alpha_powers = compute_alpha_powers(&alpha, 2); + + for row in 0..trace_len { + // Fingerprint for interaction 0: z - (0 * α^0 + col0[row] * α^1) + let bus_id_0 = FE::from(0u64); + let lc_0 = &bus_id_0 * &alpha_powers[0] + &col0[row] * &alpha_powers[1]; + let fp_0 = &z - &lc_0; + + // Fingerprint for interaction 1: z - (1 * α^0 + col1[row] * α^1) + let bus_id_1 = FE::from(1u64); + let lc_1 = &bus_id_1 * &alpha_powers[0] + &col1[row] * &alpha_powers[1]; + let fp_1 = &z - &lc_1; + + // Combined fraction: + // n = (+1) * 1 * fp_1 + (-1) * 1 * fp_0 + // d = fp_0 * fp_1 + let expected_n = &fp_1 - &fp_0; + let expected_d = &fp_0 * &fp_1; + + assert_eq!( + numerators[row], expected_n, + "Row {}: numerator mismatch for two-interaction combine", + row + ); + assert_eq!( + denominators[row], expected_d, + "Row {}: denominator mismatch for two-interaction combine", + row + ); + } + } + + /// Verify that the leaf fractions are consistent with the existing + /// `compute_logup_term_column` function for a single interaction. + /// + /// For 1 interaction: term[i] = sign * m[i] / fp[i] = N[i] / D[i] + /// So term[i] * D[i] should equal N[i]. + #[test] + fn test_leaf_fractions_consistent_with_term_column() { + let trace_len = 8; + + // Column with some values + let col0: Vec = (1..=8).map(|v| FE::from(v as u64)).collect(); + let main_segment_cols = vec![col0]; + + let interaction = BusInteraction::sender( + 2u64, + Multiplicity::One, + Packing::Direct.columns(&[0]), + ); + + let z = FE::from(1000u64); + let alpha = FE::from(11u64); + let challenges = vec![z, alpha]; + + // Compute leaf fractions + let (numerators, denominators) = compute_logup_leaf_fractions::( + &[interaction.clone()], + &main_segment_cols, + trace_len, + &challenges, + ); + + // Compute term column (= sign * m / fp) + let terms = compute_logup_term_column::( + &interaction, + &main_segment_cols, + trace_len, + &challenges, + "test", + ); + + // Verify: term[i] * denominator[i] == numerator[i] + for row in 0..trace_len { + let lhs = &terms[row] * &denominators[row]; + assert_eq!( + lhs, numerators[row], + "Row {}: term * denominator should equal numerator", + row + ); + } + } +} diff --git a/crypto/stark/src/proof/stark.rs b/crypto/stark/src/proof/stark.rs index 1751d60fe..736ce3b48 100644 --- a/crypto/stark/src/proof/stark.rs +++ b/crypto/stark/src/proof/stark.rs @@ -5,7 +5,8 @@ use math::field::{ }; use crate::{ - config::Commitment, fri::fri_decommit::FriDecommitment, lookup::BusPublicInputs, table::Table, + config::Commitment, fri::fri_decommit::FriDecommitment, gkr::GkrProof, + lookup::BusPublicInputs, table::Table, }; #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] @@ -30,6 +31,19 @@ pub struct DeepPolynomialOpening, E: IsField> { pub type DeepPolynomialOpenings = Vec>; +/// Proof for the LogUp-GKR protocol, which replaces the per-row accumulated +/// column with a GKR-based fractional summation tree proof. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct LogUpGkrProof { + /// GKR proof for the fractional summation tree. + pub gkr_proof: GkrProof, + /// The random evaluation point from the GKR protocol. + pub random_point: Vec>, + /// MLE claims for trace columns: (column_index, claimed_value) pairs. + pub column_claims: Vec<(usize, FieldElement)>, +} + #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] #[serde(bound = "PI: serde::Serialize + serde::de::DeserializeOwned")] pub struct StarkProof, E: IsField, PI> { @@ -66,6 +80,8 @@ pub struct StarkProof, E: IsField, PI> { // 1. Circular constraint offset: L/N per row // 2. Bus balance check: Σ table_contribution across all tables = expected_bus_balance pub bus_public_inputs: Option>, + // LogUp-GKR proof (when using GKR-based LogUp instead of per-row accumulated column) + pub logup_gkr_proof: Option>, // Public inputs used for boundary constraints pub public_inputs: PI, } diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index b5ac26911..7465b46df 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -10,7 +10,7 @@ use math::fft::cpu::bowers_fft::LayerTwiddles; use math::fft::errors::FFTError; use log::info; -use math::field::traits::{IsField, IsSubFieldOf}; +use math::field::traits::{IsField, IsPrimeField, IsSubFieldOf}; use math::traits::AsBytes; use math::{ field::{element::FieldElement, traits::IsFFTField}, @@ -37,8 +37,10 @@ use super::constraints::evaluator::ConstraintEvaluator; use super::domain::Domain; use super::fri::fri_decommit::FriDecommitment; use super::grinding; -use super::lookup::BusPublicInputs; -use super::proof::stark::{DeepPolynomialOpening, MultiProof, StarkProof}; +use super::lookup::{ + BusPublicInputs, LogUpGkrResult, extend_rap_challenges_with_bridge, +}; +use super::proof::stark::{DeepPolynomialOpening, LogUpGkrProof, MultiProof, StarkProof}; use super::trace::TraceTable; use super::traits::AIR; @@ -155,6 +157,8 @@ where rap_challenges: Vec>, /// Bus interaction public inputs (initial and final aux column values). bus_public_inputs: Option>, + /// LogUp-GKR result (when using GKR-based LogUp alongside traditional accumulated column). + logup_gkr_result: Option>, } /// Pre-computed twiddle factors and coset weights for a given domain size. @@ -1467,6 +1471,7 @@ pub trait IsStarkProver< where FieldElement: AsBytes, FieldElement: AsBytes, + Field: IsPrimeField, PI: Send + Sync + Clone, { info!("Started proof generation..."); @@ -1600,36 +1605,118 @@ pub trait IsStarkProver< }; // ===================================================================== - // Phase C + Rounds 2-4: Forked per table + // Round 1, Phase B': GKR sub-protocol (main transcript) // ===================================================================== - // Each table gets an independent transcript fork (cloned from the shared - // state after Phase B, domain-separated by table index). This matches - // the verifier's forking and makes per-table proving independent. - // - // Split into two passes for parallelism: - // Pass 1 (parallel): Build all auxiliary traces (fingerprint + batch inversion) - // Pass 2 (sequential): Fork transcript → extract → LDE → commit (shared pool) - - // Pass 1: Build aux traces in parallel. - // Each build_auxiliary_trace has internal parallelism (batch_inverse, par_chunks), - // but outer parallelism over 12 tables also helps on high-core-count machines. - #[cfg(feature = "instruments")] - let phase_start = Instant::now(); + // For each table with bus interactions, run the LogUp-GKR sub-protocol. + // GKR messages are bound to the main Fiat-Shamir chain so the verifier + // can replay them. - #[cfg(feature = "parallel")] - let aux_iter = air_trace_pairs.par_iter_mut(); - #[cfg(not(feature = "parallel"))] - let aux_iter = air_trace_pairs.iter_mut(); - let bus_inputs_vec: Vec>> = aux_iter + let gkr_results: Vec>> = air_trace_pairs + .iter() .map(|(air, trace, _)| { - if air.has_aux_trace() { - air.build_auxiliary_trace(*trace, &lookup_challenges) + if air.has_trace_interaction() { + let interactions = air.bus_interactions(); + let main_segment_cols = trace.columns_main(); + let trace_len = trace.num_rows(); + Some(crate::lookup::run_logup_gkr( + interactions, + &main_segment_cols, + trace_len, + &lookup_challenges, + transcript, + )) } else { None } }) .collect(); + // ===================================================================== + // Round 1, Phase B'': Sample γ for bridge batching (main transcript) + // ===================================================================== + // γ is sampled AFTER all GKR messages are bound to the transcript + // but BEFORE forking per table. The verifier replays the same sample. + + let gamma: FieldElement = if needs_lookup_challenges { + transcript.sample_field_element() + } else { + FieldElement::zero() + }; + + // ===================================================================== + // Phase C: Build aux traces + fork transcripts per table + // ===================================================================== + // For LogUp-GKR tables: build Lagrange kernel (col 0) and bridge + // running sum σ (col 1) as the two aux columns. + + // Pass 1: Build aux traces. + #[cfg(feature = "instruments")] + let phase_start = Instant::now(); + + for ((air, trace, _), gkr_result) in + air_trace_pairs.iter_mut().zip(gkr_results.iter()) + { + if air.has_trace_interaction() { + if let Some(result) = gkr_result { + let kernel = &result.lagrange_kernel; + let trace_len = trace.num_rows(); + + // Allocate aux table (2 columns: l + σ) + let (_, num_aux_columns) = air.trace_layout(); + if num_aux_columns > 0 && trace.num_aux_columns == 0 { + trace.allocate_aux_table(num_aux_columns); + } + + // Column 0: Lagrange kernel + for (row, val) in kernel.iter().enumerate() { + trace.set_aux(row, 0, val.clone()); + } + + // Column 1: Bridge running sum σ + // The constraint checks: σ[i+1] - σ[i] = l[i]·batched[i] - Δ + // where Δ = bridge_offset = target/N. + // So σ[0] = 0 (start), σ[i+1] = σ[i] + l[i]·batched[i] - Δ. + // The circular wrap-around at row N-1 requires σ[0] = σ[N-1] + l[N-1]·batched[N-1] - Δ, + // which telescopes to: 0 = Σ l[i]·batched[i] - N·Δ = target - target. + let (bridge_offset, gamma_powers) = + crate::lookup::compute_bridge_params( + &result.column_claims, + &gamma, + trace_len, + ); + let main_cols = trace.columns_main(); + + // Pre-compute batched values in parallel: batched[i] = Σ_j main_cols[col_j][i] * γ^j + #[cfg(feature = "parallel")] + let batched_iter = (0..trace_len).into_par_iter(); + #[cfg(not(feature = "parallel"))] + let batched_iter = 0..trace_len; + + let column_claims = &result.column_claims; + let batched_values: Vec> = batched_iter + .map(|row| { + let mut batched = FieldElement::::zero(); + for (j, (col_idx, _)) in column_claims.iter().enumerate() { + batched += &main_cols[*col_idx][row] * &gamma_powers[j]; + } + batched + }) + .collect(); + + // Set σ[0] = 0, then build forward: σ[i+1] = σ[i] + l[i]*batched[i] - Δ + trace.set_aux(0, 1, FieldElement::::zero()); + let mut sigma = FieldElement::::zero(); + for row in 0..trace_len - 1 { + let l_val = trace.get_aux(row, 0); + sigma = sigma + l_val * &batched_values[row] - &bridge_offset; + trace.set_aux(row + 1, 1, sigma.clone()); + } + } + } else if air.has_aux_trace() { + air.build_auxiliary_trace(*trace, &lookup_challenges); + } + } + #[cfg(feature = "instruments")] let aux_build_elapsed = phase_start.elapsed(); @@ -1709,14 +1796,30 @@ pub trait IsStarkProver< } } - // Build metadata sequentially from main_commits + aux_results + bus_inputs + // Build metadata sequentially from main_commits + aux_results + gkr_results let mut metadatas: Vec> = Vec::with_capacity(num_airs); - for ((main_commit, (aux_tree, aux_root)), bus_public_inputs) in main_commits + let mut gkr_results_iter = gkr_results.into_iter(); + for ((main_commit, (aux_tree, aux_root)), idx) in main_commits .into_iter() .zip(aux_results) - .zip(bus_inputs_vec) + .zip(0..num_airs) { + let gkr_result = gkr_results_iter.next().unwrap(); + + // Build per-table rap_challenges: base challenges + bridge params + let mut table_rap_challenges = lookup_challenges.clone(); + if let Some(ref result) = gkr_result { + let trace_len = air_trace_pairs[idx].1.num_rows(); + extend_rap_challenges_with_bridge( + &mut table_rap_challenges, + &result.column_claims, + &gamma, + trace_len, + &result.random_point, + ); + } + metadatas.push(Round1Metadata { main_merkle_tree: Arc::clone(&main_commit.main_tree), main_merkle_root: main_commit.main_root, @@ -1725,8 +1828,9 @@ pub trait IsStarkProver< num_precomputed_cols: main_commit.num_precomputed_cols, aux_merkle_tree: aux_tree, aux_merkle_root: aux_root, - rap_challenges: lookup_challenges.clone(), - bus_public_inputs, + rap_challenges: table_rap_challenges, + bus_public_inputs: None, // LogUp-GKR: no bus_public_inputs needed + logup_gkr_result: gkr_result, }); } @@ -1809,7 +1913,7 @@ pub trait IsStarkProver< table_transcript.append_field_element(&bpi.table_contribution); } - let proof = Self::prove_rounds_2_to_4( + let mut proof = Self::prove_rounds_2_to_4( *air, *pub_inputs, &round_1_result, @@ -1817,6 +1921,13 @@ pub trait IsStarkProver< domain, )?; + // Attach LogUp-GKR proof from Phase B' (if this table had bus interactions) + proof.logup_gkr_proof = metadata.logup_gkr_result.as_ref().map(|r| LogUpGkrProof { + gkr_proof: r.gkr_proof.clone(), + random_point: r.random_point.clone(), + column_claims: r.column_claims.clone(), + }); + // Collect per-table sub-op timing via TLS. // Both the store (inside prove_rounds_2_to_4) and this take run on the // same rayon worker thread, so sub-ops are valid in both sequential and @@ -1891,6 +2002,7 @@ pub trait IsStarkProver< where FieldElement: AsBytes, FieldElement: AsBytes, + Field: IsPrimeField, PI: Send + Sync + Clone, { let air_trace_pairs = vec![(air, trace, pub_inputs)]; @@ -2056,6 +2168,8 @@ pub trait IsStarkProver< nonce: round_4_result.nonce, // Bus interaction public inputs (for boundary constraints and bus balance check) bus_public_inputs: round_1_result.bus_public_inputs.clone(), + // LogUp-GKR proof (not yet used; will replace accumulated column in future) + logup_gkr_proof: None, // Public inputs for boundary constraints public_inputs: pub_inputs.clone(), trace_length: domain.interpolation_domain_size, diff --git a/crypto/stark/src/sumcheck.rs b/crypto/stark/src/sumcheck.rs new file mode 100644 index 000000000..1f36783ed --- /dev/null +++ b/crypto/stark/src/sumcheck.rs @@ -0,0 +1,828 @@ +use core::fmt; +use crypto::fiat_shamir::is_transcript::IsTranscript; +use math::field::{element::FieldElement, traits::IsField}; + +/// A degree-d univariate polynomial represented by its evaluations at the +/// integer nodes 0, 1, ..., d. This is the polynomial that the sumcheck prover +/// sends to the verifier in each round. +/// +/// Storing evaluations (rather than coefficients) is natural for the sumcheck +/// protocol because: +/// - The prover constructs the polynomial by evaluating it at small integer points. +/// - The verifier only needs to check p(0) + p(1) and evaluate at a random challenge. +/// - Lagrange interpolation over integer nodes is cheap (small denominators). +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct RoundPoly { + /// Evaluations at x = 0, 1, ..., d where d = evals.len() - 1. + evals: Vec>, +} + +impl RoundPoly { + /// Create a new `RoundPoly` from evaluations at x = 0, 1, ..., d. + /// + /// `evals[i]` is the polynomial evaluated at x = i. + /// The polynomial has degree `evals.len() - 1`. + pub fn new(evals: Vec>) -> Self { + assert!(!evals.is_empty(), "RoundPoly must have at least one evaluation"); + Self { evals } + } + + /// Returns p(0) + p(1). + /// + /// In the sumcheck protocol, the verifier checks that this sum matches the + /// claimed sum for the current round. This is the key consistency check. + pub fn sum_at_binary(&self) -> FieldElement { + assert!( + self.evals.len() >= 2, + "sum_at_binary requires at least 2 evaluations (at 0 and 1)" + ); + &self.evals[0] + &self.evals[1] + } + + /// Evaluate the polynomial at an arbitrary field element using Lagrange + /// interpolation over integer nodes 0, 1, ..., d. + /// + /// Given evaluations y_0, ..., y_d at nodes 0, 1, ..., d, the Lagrange + /// interpolation formula is: + /// + /// p(x) = sum_{i=0}^{d} y_i * prod_{j != i} (x - j) / (i - j) + /// + /// The denominators (i - j) for integer nodes are just products of small + /// integers, so we precompute them as field elements. + pub fn evaluate(&self, point: &FieldElement) -> FieldElement { + let d = self.evals.len() - 1; + + // Special case: constant polynomial (degree 0) + if d == 0 { + return self.evals[0].clone(); + } + + // Precompute (point - j) for j = 0, ..., d + let point_minus_j: Vec> = (0..=d) + .map(|j| point - &FieldElement::from(j as u64)) + .collect(); + + // Check if point is one of the integer nodes (avoid division by zero) + for (j, pm) in point_minus_j.iter().enumerate() { + if *pm == FieldElement::zero() { + return self.evals[j].clone(); + } + } + + // Barycentric Lagrange interpolation with batch inversion. + // + // Barycentric weights: w_i = 1 / prod_{j≠i}(i - j) for integer nodes. + // We batch-invert the weights together with point_minus_j to use only + // 1 field inversion total (instead of 2(d+1) individual inversions). + + // Compute barycentric weight denominators: prod_{j≠i}(i - j) + let mut to_invert: Vec> = (0..=d) + .map(|i| { + let mut denom = FieldElement::::one(); + for j in 0..=d { + if j != i { + denom = &denom * &FieldElement::from((i as i64) - (j as i64)); + } + } + denom + }) + .collect(); + + // Append point_minus_j values; batch-invert everything in one call + to_invert.extend(point_minus_j.iter().cloned()); + FieldElement::inplace_batch_inverse(&mut to_invert) + .expect("All values are nonzero"); + + let (w_inv, pm_inv) = to_invert.split_at(d + 1); + + // master_product = prod_{j=0}^{d} (point - j) + let master_product = point_minus_j + .iter() + .fold(FieldElement::one(), |acc, v| &acc * v); + + // result = master_product * Σ_i (evals[i] * w_inv[i] * pm_inv[i]) + let mut sum = FieldElement::zero(); + for i in 0..=d { + sum = &sum + &(&self.evals[i] * &(&w_inv[i] * &pm_inv[i])); + } + + &master_product * &sum + } + + /// Returns the number of evaluations (degree + 1). + pub fn num_evals(&self) -> usize { + self.evals.len() + } + + /// Returns a reference to the evaluations. + pub fn evals(&self) -> &[FieldElement] { + &self.evals + } +} + +/// Proof produced by the sumcheck prover: one round polynomial per variable. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +#[serde(bound = "")] +pub struct SumcheckProof { + pub round_polys: Vec>, +} + +/// Run the sumcheck interactive proof (made non-interactive via Fiat-Shamir). +/// +/// Proves that `claimed_sum == sum_{x in {0,1}^num_vars} f(x)` where `evals` +/// holds the evaluations of f over the Boolean hypercube in little-endian bit +/// order (index = b_0 + 2*b_1 + ...). +/// +/// # Arguments +/// - `evals`: evaluations of f over {0,1}^num_vars, length must be 2^num_vars +/// - `claimed_sum`: the asserted sum +/// - `num_vars`: number of Boolean variables +/// - `max_degree`: maximum degree of round polynomials (need max_degree+1 evaluation points) +/// - `transcript`: Fiat-Shamir transcript for deterministic challenge sampling +/// +/// # Returns +/// `(proof, challenges)` where `proof` contains one round polynomial per variable +/// and `challenges` is the vector of random points sampled during the protocol. +pub fn sumcheck_prove( + evals: &[FieldElement], + claimed_sum: &FieldElement, + num_vars: usize, + max_degree: usize, + transcript: &mut impl IsTranscript, +) -> (SumcheckProof, Vec>) { + assert_eq!( + evals.len(), + 1 << num_vars, + "evals length must be 2^num_vars" + ); + assert!( + max_degree >= 1, + "max_degree must be at least 1 for a meaningful sumcheck" + ); + + let mut table = evals.to_vec(); + let mut round_polys = Vec::with_capacity(num_vars); + let mut challenges = Vec::with_capacity(num_vars); + let mut round_claimed_sum = claimed_sum.clone(); + + for _round in 0..num_vars { + let half = table.len() / 2; + + // p(0) = sum of left halves (zero multiplications) + let mut p0 = FieldElement::::zero(); + for j in 0..half { + p0 = &p0 + &table[2 * j]; + } + + // p(1) = round_claimed_sum - p(0) (free, from sumcheck identity p(0)+p(1)=sum) + let p1 = &round_claimed_sum - &p0; + + let mut poly_evals = Vec::with_capacity(max_degree + 1); + poly_evals.push(p0); + poly_evals.push(p1); + + // For t >= 2: use delta form. + // Linear interpolation at t: table[2j] + t * (table[2j+1] - table[2j]) + // p(t) = Σ_j [table[2j] + t * delta_j] where delta_j = table[2j+1] - table[2j] + for t in 2..=max_degree { + let t_fe = FieldElement::::from(t as u64); + let mut sum = FieldElement::::zero(); + for j in 0..half { + let delta = &table[2 * j + 1] - &table[2 * j]; + sum = &sum + &(&table[2 * j] + &(&t_fe * &delta)); + } + poly_evals.push(sum); + } + + let round_poly = RoundPoly::new(poly_evals); + + // Append all evaluations of the round polynomial to the transcript + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + // Sample the challenge for this round + let challenge: FieldElement = transcript.sample_field_element(); + + // Update round_claimed_sum for next round: p(challenge) + round_claimed_sum = round_poly.evaluate(&challenge); + + // Bind the table: fold using the challenge + // table[j] = table[2j] + challenge * (table[2j+1] - table[2j]) + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + + round_polys.push(round_poly); + challenges.push(challenge); + } + + (SumcheckProof { round_polys }, challenges) +} + +/// Errors that can occur during sumcheck verification. +#[derive(Debug, Clone)] +pub enum SumcheckError { + /// The proof contains a different number of round polynomials than expected. + WrongNumberOfRounds { expected: usize, got: usize }, + /// The round polynomial's p(0) + p(1) does not match the expected sum. + RoundSumMismatch { round: usize }, +} + +impl fmt::Display for SumcheckError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + SumcheckError::WrongNumberOfRounds { expected, got } => { + write!( + f, + "wrong number of rounds: expected {}, got {}", + expected, got + ) + } + SumcheckError::RoundSumMismatch { round } => { + write!(f, "round sum mismatch at round {}", round) + } + } + } +} + +/// Verify a sumcheck proof (non-interactive via Fiat-Shamir). +/// +/// Checks that the prover's round polynomials are consistent with the claimed +/// sum. The transcript operations must exactly mirror those of `sumcheck_prove` +/// so that the same challenges are derived. +/// +/// # Arguments +/// - `proof`: the sumcheck proof containing one round polynomial per variable +/// - `claimed_sum`: the asserted sum over the Boolean hypercube +/// - `num_vars`: number of Boolean variables +/// - `transcript`: Fiat-Shamir transcript (must have the same seed as the prover's) +/// +/// # Returns +/// `Ok((challenges, final_eval))` where `challenges` is the random point and +/// `final_eval` is the evaluation claim at that point (i.e., the last round +/// polynomial evaluated at the last challenge). +pub fn sumcheck_verify( + proof: &SumcheckProof, + claimed_sum: &FieldElement, + num_vars: usize, + transcript: &mut impl IsTranscript, +) -> Result<(Vec>, FieldElement), SumcheckError> { + if proof.round_polys.len() != num_vars { + return Err(SumcheckError::WrongNumberOfRounds { + expected: num_vars, + got: proof.round_polys.len(), + }); + } + + let mut current_sum = claimed_sum.clone(); + let mut challenges = Vec::with_capacity(num_vars); + + for round in 0..num_vars { + // Check that p_r(0) + p_r(1) == current_sum + if proof.round_polys[round].sum_at_binary() != current_sum { + return Err(SumcheckError::RoundSumMismatch { round }); + } + + // Append all evaluations of the round polynomial to the transcript + // (must match the prover exactly) + for eval in proof.round_polys[round].evals() { + transcript.append_field_element(eval); + } + + // Sample the challenge for this round + let challenge: FieldElement = transcript.sample_field_element(); + + // Update the running sum: next round's claimed sum is p_r(challenge) + current_sum = proof.round_polys[round].evaluate(&challenge); + + challenges.push(challenge); + } + + Ok((challenges, current_sum)) +} + +#[cfg(test)] +mod tests { + use super::*; + use crypto::fiat_shamir::default_transcript::DefaultTranscript; + use math::field::goldilocks::GoldilocksField; + + // Use the base Goldilocks field for tests. The cubic extension would work + // identically since all operations are generic over IsField. + type FE = FieldElement; + + /// Helper: evaluate a polynomial given in coefficient form at a point. + /// coeffs[i] is the coefficient of x^i. + fn eval_poly_coeffs(coeffs: &[FE], x: &FE) -> FE { + let mut result = FE::zero(); + let mut power = FE::one(); + for c in coeffs { + result = &result + &(c * &power); + power = &power * x; + } + result + } + + #[test] + fn test_sum_at_binary_linear() { + // p(x) = 3 + 5x => p(0) = 3, p(1) = 8 + // sum_at_binary = 3 + 8 = 11 + let evals = vec![FE::from(3u64), FE::from(8u64)]; + let poly = RoundPoly::new(evals); + assert_eq!(poly.sum_at_binary(), FE::from(11u64)); + } + + #[test] + fn test_sum_at_binary_quadratic() { + // p(x) = 2x^2 + 3x + 1 => p(0) = 1, p(1) = 6, p(2) = 15 + // sum_at_binary = p(0) + p(1) = 1 + 6 = 7 + let evals = vec![FE::from(1u64), FE::from(6u64), FE::from(15u64)]; + let poly = RoundPoly::new(evals); + assert_eq!(poly.sum_at_binary(), FE::from(7u64)); + } + + #[test] + fn test_evaluate_at_nodes_linear() { + // p(x) = 3 + 5x => p(0) = 3, p(1) = 8 + // Evaluating at the nodes should return the stored evaluations. + let evals = vec![FE::from(3u64), FE::from(8u64)]; + let poly = RoundPoly::new(evals); + + assert_eq!(poly.evaluate(&FE::from(0u64)), FE::from(3u64)); + assert_eq!(poly.evaluate(&FE::from(1u64)), FE::from(8u64)); + } + + #[test] + fn test_evaluate_at_nodes_quadratic() { + // p(x) = 2x^2 + 3x + 1 => p(0) = 1, p(1) = 6, p(2) = 15 + let evals = vec![FE::from(1u64), FE::from(6u64), FE::from(15u64)]; + let poly = RoundPoly::new(evals); + + assert_eq!(poly.evaluate(&FE::from(0u64)), FE::from(1u64)); + assert_eq!(poly.evaluate(&FE::from(1u64)), FE::from(6u64)); + assert_eq!(poly.evaluate(&FE::from(2u64)), FE::from(15u64)); + } + + #[test] + fn test_evaluate_at_random_point_linear() { + // p(x) = 3 + 5x + // p(0) = 3, p(1) = 8 + // p(7) = 3 + 35 = 38 + let evals = vec![FE::from(3u64), FE::from(8u64)]; + let poly = RoundPoly::new(evals); + + let point = FE::from(7u64); + let expected = FE::from(38u64); + assert_eq!(poly.evaluate(&point), expected); + } + + #[test] + fn test_evaluate_at_random_point_quadratic() { + // p(x) = 2x^2 + 3x + 1 + // p(0) = 1, p(1) = 6, p(2) = 15 + // p(5) = 2*25 + 3*5 + 1 = 50 + 15 + 1 = 66 + let coeffs = [FE::from(1u64), FE::from(3u64), FE::from(2u64)]; + let evals = vec![ + eval_poly_coeffs(&coeffs, &FE::from(0u64)), + eval_poly_coeffs(&coeffs, &FE::from(1u64)), + eval_poly_coeffs(&coeffs, &FE::from(2u64)), + ]; + let poly = RoundPoly::new(evals); + + let point = FE::from(5u64); + let expected = eval_poly_coeffs(&coeffs, &point); + assert_eq!(poly.evaluate(&point), expected); + assert_eq!(expected, FE::from(66u64)); + } + + #[test] + fn test_evaluate_at_random_point_cubic() { + // p(x) = x^3 + 2x^2 + 3x + 4 + // Need 4 evaluation points: 0, 1, 2, 3 + let coeffs = [FE::from(4u64), FE::from(3u64), FE::from(2u64), FE::from(1u64)]; + let evals: Vec = (0..4) + .map(|i| eval_poly_coeffs(&coeffs, &FE::from(i as u64))) + .collect(); + let poly = RoundPoly::new(evals); + + // p(10) = 1000 + 200 + 30 + 4 = 1234 + let point = FE::from(10u64); + let expected = eval_poly_coeffs(&coeffs, &point); + assert_eq!(poly.evaluate(&point), expected); + assert_eq!(expected, FE::from(1234u64)); + } + + #[test] + fn test_evaluate_at_large_field_element() { + // Test with a large field element as the evaluation point to ensure + // field arithmetic works correctly (not just small integers). + let coeffs = [FE::from(7u64), FE::from(11u64), FE::from(13u64)]; + let evals: Vec = (0..3) + .map(|i| eval_poly_coeffs(&coeffs, &FE::from(i as u64))) + .collect(); + let poly = RoundPoly::new(evals); + + // Use a large evaluation point + let point = FE::from(1_000_000_007u64); + let expected = eval_poly_coeffs(&coeffs, &point); + assert_eq!(poly.evaluate(&point), expected); + } + + #[test] + fn test_constant_polynomial() { + // p(x) = 42 (constant) + let evals = vec![FE::from(42u64)]; + let poly = RoundPoly::new(evals); + + // A constant polynomial should evaluate to the same value everywhere + assert_eq!(poly.evaluate(&FE::from(0u64)), FE::from(42u64)); + assert_eq!(poly.evaluate(&FE::from(100u64)), FE::from(42u64)); + assert_eq!(poly.evaluate(&FE::from(999u64)), FE::from(42u64)); + } + + #[test] + fn test_num_evals() { + let evals = vec![FE::from(1u64), FE::from(2u64), FE::from(3u64)]; + let poly = RoundPoly::new(evals); + assert_eq!(poly.num_evals(), 3); + } + + #[test] + fn test_evals_accessor() { + let evals = vec![FE::from(10u64), FE::from(20u64)]; + let poly = RoundPoly::new(evals); + assert_eq!(poly.evals()[0], FE::from(10u64)); + assert_eq!(poly.evals()[1], FE::from(20u64)); + } + + #[test] + #[should_panic(expected = "RoundPoly must have at least one evaluation")] + fn test_empty_evals_panics() { + let _poly = RoundPoly::::new(vec![]); + } + + #[test] + #[should_panic(expected = "sum_at_binary requires at least 2 evaluations")] + fn test_sum_at_binary_single_eval_panics() { + let poly = RoundPoly::new(vec![FE::from(1u64)]); + let _ = poly.sum_at_binary(); + } + + #[test] + fn test_evaluate_consistency_with_coefficients() { + // Build a degree-4 polynomial from coefficients, generate evaluations + // at 0..=4, construct RoundPoly, and verify evaluate() at many points. + let coeffs = [ + FE::from(5u64), + FE::from(3u64), + FE::from(7u64), + FE::from(2u64), + FE::from(1u64), + ]; + let evals: Vec = (0..5) + .map(|i| eval_poly_coeffs(&coeffs, &FE::from(i as u64))) + .collect(); + let poly = RoundPoly::new(evals); + + // Check at several points including nodes and non-nodes + for x in [0u64, 1, 2, 3, 4, 5, 10, 100, 12345] { + let point = FE::from(x); + let expected = eval_poly_coeffs(&coeffs, &point); + assert_eq!( + poly.evaluate(&point), + expected, + "Mismatch at x = {}", + x + ); + } + } + + // ==================== Sumcheck prover tests ==================== + + #[test] + fn test_sumcheck_prove_2var_linear() { + // f(x1, x2) = 3*x1 + 5*x2 + 7 + // Hypercube evaluations in little-endian bit order: + // index 0 = (x1=0, x2=0): f = 7 + // index 1 = (x1=1, x2=0): f = 10 + // index 2 = (x1=0, x2=1): f = 12 + // index 3 = (x1=1, x2=1): f = 15 + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); // 7 + 10 + 12 + 15 + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, challenges) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut transcript); + + // Should have 2 round polynomials (one per variable) + assert_eq!(proof.round_polys.len(), 2); + assert_eq!(challenges.len(), 2); + + // First round poly: sum over x2 of f(x1, x2) + // p1(0) = f(0,0) + f(0,1) = 7 + 12 = 19 + // p1(1) = f(1,0) + f(1,1) = 10 + 15 = 25 + // p1(0) + p1(1) = 19 + 25 = 44 = claimed_sum + assert_eq!(proof.round_polys[0].sum_at_binary(), claimed_sum); + assert_eq!(proof.round_polys[0].evals()[0], FE::from(19u64)); + assert_eq!(proof.round_polys[0].evals()[1], FE::from(25u64)); + + // Each round poly should have max_degree+1 = 2 evaluations + assert_eq!(proof.round_polys[0].num_evals(), 2); + assert_eq!(proof.round_polys[1].num_evals(), 2); + + // Second round: the claimed sum for round 2 is p1(r1) where r1 = challenges[0]. + // Verify that p2(0) + p2(1) = p1(r1). + let r1_eval = proof.round_polys[0].evaluate(&challenges[0]); + assert_eq!(proof.round_polys[1].sum_at_binary(), r1_eval); + } + + #[test] + fn test_sumcheck_prove_1var() { + // f(x) = 2x + 3 => f(0) = 3, f(1) = 5 + // Sum = 3 + 5 = 8 + let evals = vec![FE::from(3u64), FE::from(5u64)]; + let claimed_sum = FE::from(8u64); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, challenges) = sumcheck_prove(&evals, &claimed_sum, 1, 1, &mut transcript); + + assert_eq!(proof.round_polys.len(), 1); + assert_eq!(challenges.len(), 1); + assert_eq!(proof.round_polys[0].sum_at_binary(), claimed_sum); + assert_eq!(proof.round_polys[0].evals()[0], FE::from(3u64)); + assert_eq!(proof.round_polys[0].evals()[1], FE::from(5u64)); + } + + #[test] + fn test_sumcheck_prove_3var_constant() { + // f(x1, x2, x3) = 1 for all inputs + // Sum = 8 (2^3 points, each contributing 1) + let evals = vec![FE::one(); 8]; + let claimed_sum = FE::from(8u64); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, challenges) = sumcheck_prove(&evals, &claimed_sum, 3, 1, &mut transcript); + + assert_eq!(proof.round_polys.len(), 3); + + // First round: p(0) = 4, p(1) = 4, sum = 8 + assert_eq!(proof.round_polys[0].sum_at_binary(), claimed_sum); + assert_eq!(proof.round_polys[0].evals()[0], FE::from(4u64)); + assert_eq!(proof.round_polys[0].evals()[1], FE::from(4u64)); + + // Each subsequent round's sum should match the previous round's evaluation at challenge + for r in 1..3 { + let prev_eval = proof.round_polys[r - 1].evaluate(&challenges[r - 1]); + assert_eq!(proof.round_polys[r].sum_at_binary(), prev_eval); + } + } + + #[test] + fn test_sumcheck_prove_higher_degree() { + // Test with max_degree = 3 (4 evaluation points per round poly). + // For a multilinear f, higher-degree evaluations are determined by + // the degree-1 interpolation, so they are just linear extrapolation. + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, challenges) = sumcheck_prove(&evals, &claimed_sum, 2, 3, &mut transcript); + + assert_eq!(proof.round_polys.len(), 2); + + // Each round poly should have max_degree+1 = 4 evaluations + assert_eq!(proof.round_polys[0].num_evals(), 4); + assert_eq!(proof.round_polys[1].num_evals(), 4); + + // p(0) + p(1) must still equal claimed sum + assert_eq!(proof.round_polys[0].sum_at_binary(), claimed_sum); + + // Since the underlying function is multilinear, the round poly is degree 1. + // Verify that evaluations at t=2 and t=3 are consistent with linear interpolation. + let p0 = &proof.round_polys[0].evals()[0]; + let p1 = &proof.round_polys[0].evals()[1]; + // p(t) = p0*(1-t) + p1*t = p0 + (p1 - p0)*t + let slope = p1 - p0; + let p2_expected = p0 + &(&slope * &FE::from(2u64)); + let p3_expected = p0 + &(&slope * &FE::from(3u64)); + assert_eq!(proof.round_polys[0].evals()[2], p2_expected); + assert_eq!(proof.round_polys[0].evals()[3], p3_expected); + + // Consistency between rounds + for r in 1..2 { + let prev_eval = proof.round_polys[r - 1].evaluate(&challenges[r - 1]); + assert_eq!(proof.round_polys[r].sum_at_binary(), prev_eval); + } + } + + #[test] + fn test_sumcheck_prove_deterministic() { + // Same inputs and transcript seed should produce identical proofs + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); + + let mut t1 = DefaultTranscript::::new(&[0x42]); + let (proof1, challenges1) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut t1); + + let mut t2 = DefaultTranscript::::new(&[0x42]); + let (proof2, challenges2) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut t2); + + assert_eq!(challenges1, challenges2); + for (rp1, rp2) in proof1.round_polys.iter().zip(proof2.round_polys.iter()) { + assert_eq!(rp1.evals(), rp2.evals()); + } + } + + #[test] + fn test_sumcheck_prove_final_value() { + // After all rounds, the table should contain a single element which is + // f evaluated at the random challenge point. We verify by checking that + // the last round poly evaluated at the last challenge gives this value. + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); + + let mut transcript = DefaultTranscript::::new(&[]); + let (proof, challenges) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut transcript); + + // Manually compute f(r1, r2) by multilinear evaluation + let r1 = &challenges[0]; + let r2 = &challenges[1]; + // f(r1, r2) = f(0,0)*(1-r1)*(1-r2) + f(1,0)*r1*(1-r2) + f(0,1)*(1-r1)*r2 + f(1,1)*r1*r2 + let one = FE::one(); + let one_minus_r1 = &one - r1; + let one_minus_r2 = &one - r2; + let f_at_r = &(&(&evals[0] * &one_minus_r1) * &one_minus_r2) + + &(&(&evals[1] * r1) * &one_minus_r2) + + &(&(&evals[2] * &one_minus_r1) * r2) + + &(&(&evals[3] * r1) * r2); + + // The last round poly evaluated at the last challenge should give f(r1, r2) + let last_round_eval = proof.round_polys[1].evaluate(&challenges[1]); + assert_eq!(last_round_eval, f_at_r); + } + + #[test] + #[should_panic(expected = "evals length must be 2^num_vars")] + fn test_sumcheck_prove_wrong_length_panics() { + let evals = vec![FE::from(1u64), FE::from(2u64), FE::from(3u64)]; // length 3, not a power of 2 + let mut transcript = DefaultTranscript::::new(&[]); + let _ = sumcheck_prove(&evals, &FE::from(6u64), 2, 1, &mut transcript); + } + + // ==================== Sumcheck verifier tests ==================== + + #[test] + fn test_sumcheck_verify() { + // f(x1, x2) = 3*x1 + 5*x2 + 7 + // Hypercube evaluations (little-endian bit order): + // (0,0)=7, (1,0)=10, (0,1)=12, (1,1)=15 + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); // 7 + 10 + 12 + 15 + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, prover_challenges) = + sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut prover_transcript); + + // Verify with a fresh transcript (same seed) + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = sumcheck_verify(&proof, &claimed_sum, 2, &mut verifier_transcript); + + assert!(result.is_ok(), "Verification should succeed"); + let (challenges, _final_eval) = result.unwrap(); + + // Verifier must derive the same challenges as the prover + assert_eq!(challenges, prover_challenges); + } + + #[test] + fn test_sumcheck_verify_wrong_sum() { + // f(x1, x2) = 3*x1 + 5*x2 + 7 + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); + + // Prove with the correct sum + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, _) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut prover_transcript); + + // Verify with a WRONG claimed sum + let wrong_sum = FE::from(99u64); + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let result = sumcheck_verify(&proof, &wrong_sum, 2, &mut verifier_transcript); + + assert!(result.is_err(), "Verification should fail with wrong sum"); + match result.unwrap_err() { + SumcheckError::RoundSumMismatch { round } => { + assert_eq!(round, 0, "Mismatch should be detected at round 0"); + } + other => panic!("Expected RoundSumMismatch, got {:?}", other), + } + } + + #[test] + fn test_sumcheck_verify_roundtrip_3var() { + // f(x1, x2, x3) = x1 * x2 + x3 + 2 + // Hypercube evaluations in little-endian bit order: + // index = b0 + 2*b1 + 4*b2, where b0=x1, b1=x2, b2=x3 + // (0,0,0)=2, (1,0,0)=2, (0,1,0)=2, (1,1,0)=3, (0,0,1)=3, (1,0,1)=3, (0,1,1)=3, (1,1,1)=4 + let evals = vec![ + FE::from(2u64), // f(0,0,0) = 0*0 + 0 + 2 = 2 + FE::from(2u64), // f(1,0,0) = 1*0 + 0 + 2 = 2 + FE::from(2u64), // f(0,1,0) = 0*1 + 0 + 2 = 2 + FE::from(3u64), // f(1,1,0) = 1*1 + 0 + 2 = 3 + FE::from(3u64), // f(0,0,1) = 0*0 + 1 + 2 = 3 + FE::from(3u64), // f(1,0,1) = 1*0 + 1 + 2 = 3 + FE::from(3u64), // f(0,1,1) = 0*1 + 1 + 2 = 3 + FE::from(4u64), // f(1,1,1) = 1*1 + 1 + 2 = 4 + ]; + let claimed_sum: FE = evals.iter().fold(FE::zero(), |acc, v| &acc + v); // = 22 + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[0xAB]); + let (proof, prover_challenges) = + sumcheck_prove(&evals, &claimed_sum, 3, 1, &mut prover_transcript); + + // Verify + let mut verifier_transcript = DefaultTranscript::::new(&[0xAB]); + let result = sumcheck_verify(&proof, &claimed_sum, 3, &mut verifier_transcript); + + assert!(result.is_ok(), "3-var verification should succeed"); + let (challenges, _final_eval) = result.unwrap(); + assert_eq!(challenges.len(), 3); + assert_eq!(challenges, prover_challenges); + } + + #[test] + fn test_sumcheck_verify_final_eval() { + // Verify that final_eval matches the MLE evaluated at the challenge point. + // f(x1, x2) = 3*x1 + 5*x2 + 7 + let evals = vec![ + FE::from(7u64), + FE::from(10u64), + FE::from(12u64), + FE::from(15u64), + ]; + let claimed_sum = FE::from(44u64); + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[]); + let (proof, _) = sumcheck_prove(&evals, &claimed_sum, 2, 1, &mut prover_transcript); + + // Verify + let mut verifier_transcript = DefaultTranscript::::new(&[]); + let (challenges, final_eval) = + sumcheck_verify(&proof, &claimed_sum, 2, &mut verifier_transcript).unwrap(); + + // Compute MLE at the challenge point manually: + // f(r1, r2) = f(0,0)*(1-r1)*(1-r2) + f(1,0)*r1*(1-r2) + // + f(0,1)*(1-r1)*r2 + f(1,1)*r1*r2 + let r1 = &challenges[0]; + let r2 = &challenges[1]; + let one = FE::one(); + let one_minus_r1 = &one - r1; + let one_minus_r2 = &one - r2; + let mle_at_r = &(&(&evals[0] * &one_minus_r1) * &one_minus_r2) + + &(&(&evals[1] * r1) * &one_minus_r2) + + &(&(&evals[2] * &one_minus_r1) * r2) + + &(&(&evals[3] * r1) * r2); + + assert_eq!( + final_eval, mle_at_r, + "final_eval from verifier must match MLE evaluated at challenge point" + ); + } +} diff --git a/crypto/stark/src/tests/bus_tests/packing_tests.rs b/crypto/stark/src/tests/bus_tests/packing_tests.rs index ec9f2035a..2e52e3d9e 100644 --- a/crypto/stark/src/tests/bus_tests/packing_tests.rs +++ b/crypto/stark/src/tests/bus_tests/packing_tests.rs @@ -325,8 +325,8 @@ fn test_air_layout_single_interaction() { vec![], ); - // 4 main, 1 aux (0 committed pairs + 1 accumulated with 1 absorbed) - assert_eq!(air.trace_layout(), (4, 1)); + // 4 main, 2 aux (Lagrange kernel + bridge running sum σ) + assert_eq!(air.trace_layout(), (4, 2)); } #[test] @@ -356,6 +356,6 @@ fn test_air_layout_multiple_interactions() { vec![], ); - // 5 main, 1 aux (0 committed pairs + 1 accumulated with 2 absorbed) - assert_eq!(air.trace_layout(), (5, 1)); + // 5 main, 2 aux (Lagrange kernel + bridge running sum σ) + assert_eq!(air.trace_layout(), (5, 2)); } diff --git a/crypto/stark/src/tests/bus_tests/soundness_tests.rs b/crypto/stark/src/tests/bus_tests/soundness_tests.rs index e1994ef6a..a9602a5ed 100644 --- a/crypto/stark/src/tests/bus_tests/soundness_tests.rs +++ b/crypto/stark/src/tests/bus_tests/soundness_tests.rs @@ -605,90 +605,9 @@ fn test_missing_receiver() { } // ============================================================================= -// First-row boundary constraints +// LogUp-GKR soundness (replaces old bus_public_inputs tests) // ============================================================================= -/// A proof where table_contribution has been tampered with is rejected. -/// -/// The table_contribution (L) is used both for the bus balance check -/// (Σ L = 0 across all tables) and for the per-row circular constraint -/// offset (L/N). Corrupting it causes the circular transition constraint -/// to fail, since the committed trace was built with the honest L. -#[test_log::test] -fn test_tampered_table_contribution() { - // Simple valid trace: CPU sends (5, 3, 8) to the ADD table. - let mut cpu_trace = TraceTable::from_columns_main( - vec![ - vec![FE::one(), FE::zero(), FE::zero(), FE::zero()], // add_flag - vec![FE::zero(); 4], // mul_flag - vec![FE::from(5), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(3), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(8), FE::zero(), FE::zero(), FE::zero()], - ], - 1, - ); - let mut add_trace = TraceTable::from_columns_main( - vec![ - vec![FE::from(5), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(3), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(8), FE::zero(), FE::zero(), FE::zero()], - vec![FE::one(), FE::zero(), FE::zero(), FE::zero()], // multiplicity = 1 - ], - 1, - ); - let mut mul_trace = TraceTable::from_columns_main( - vec![ - vec![FE::zero(); 4], - vec![FE::zero(); 4], - vec![FE::zero(); 4], - vec![FE::zero(); 4], - ], - 1, - ); - - let proof_options = ProofOptions::default_test_options(); - let cpu_air = new_cpu_air_with_lookup(&proof_options); - let add_air = new_add_air_with_lookup(&proof_options); - let mul_air = new_mul_air_with_lookup(&proof_options); - - let air_trace_pairs: Vec<( - &dyn AIR, - _, - _, - )> = vec![ - (&cpu_air, &mut cpu_trace, &()), - (&add_air, &mut add_trace, &()), - (&mul_air, &mut mul_trace, &()), - ]; - - let mut multi_proof = - Prover::multi_prove(air_trace_pairs, &mut DefaultTranscript::::new(&[])).unwrap(); - - // Corrupt table_contribution in the ADD table's bus public inputs. - // This changes the per-row offset L/N used in the circular constraint, - // so the verifier's transition constraint evaluation will disagree with the - // committed trace (which was built with the honest L). - let add_proof = &mut multi_proof.proofs[1]; // proofs: [cpu=0, add=1, mul=2] - let bus_inputs = add_proof - .bus_public_inputs - .as_mut() - .expect("ADD table must have bus public inputs"); - bus_inputs.table_contribution += FieldElement::::one(); - - let airs: Vec<&dyn AIR> = - vec![&cpu_air, &add_air, &mul_air]; - - assert!( - !Verifier::multi_verify( - &airs, - &multi_proof, - &mut DefaultTranscript::::new(&[]), - &FieldElement::zero(), - ), - "Proof with corrupted table_contribution must be rejected" - ); -} - /// A proof where the acc column OOD evaluation is tampered with is rejected. /// /// The circular transition constraint enforces the relationship between @@ -770,18 +689,20 @@ fn test_tampered_acc_ood_evaluation() { ); } -// ============================================================================= -// Invalid bus public inputs -// ============================================================================= - -/// A proof where bus_public_inputs is None for a LogUp AIR is rejected. -/// -/// The verifier must reject proofs where an AIR declares bus interactions -/// (has_trace_interaction() = true) but the proof omits bus_public_inputs. -/// Without this check, a dishonest prover could bypass both the boundary -/// constraints and the bus balance check. -#[test_log::test] -fn test_missing_bus_public_inputs_rejected() { +// LogUp-GKR soundness comes from: +// 1. GKR proof verification (verified in Phase B') +// 2. Bridge running sum transition constraint (σ column) +// 3. Bus balance check (Σ GKR claimed_sums = 0) +// 4. Column claims verified through GKR protocol +// +// The following tests tamper with specific GKR proof fields to verify rejection. + +/// Helper: generate a valid multi-proof from the standard 3-table scenario +/// (CPU sends to ADD and MUL, all values correct). +fn generate_valid_multi_proof() -> ( + crate::proof::stark::MultiProof, + Vec>, +) { let mut cpu_trace = TraceTable::from_columns_main( vec![ vec![FE::one(), FE::zero(), FE::zero(), FE::zero()], @@ -826,149 +747,190 @@ fn test_missing_bus_public_inputs_rejected() { (&mul_air, &mut mul_trace, &()), ]; - let mut multi_proof = + let multi_proof = Prover::multi_prove(air_trace_pairs, &mut DefaultTranscript::::new(&[])).unwrap(); - // Remove bus_public_inputs from the ADD table proof entirely. - multi_proof.proofs[1].bus_public_inputs = None; + (multi_proof, vec![cpu_air, add_air, mul_air]) +} - let airs: Vec<&dyn AIR> = - vec![&cpu_air, &add_air, &mul_air]; +/// Tampered GKR column claims are rejected. +/// +/// The column_claims in the LogUp-GKR proof contain MLE evaluations of main +/// trace columns at the GKR random point. These claims feed into the bridge +/// running sum computation via extend_rap_challenges_with_bridge: the bridge +/// offset = (sum of gamma^j * c_j) / N. If we tamper with a column claim, +/// the bridge offset will be wrong and the bridge transition constraint will +/// fail at the OOD point. +#[test_log::test] +fn test_tampered_gkr_column_claims_rejected() { + let (mut multi_proof, airs) = generate_valid_multi_proof(); + + // Tamper with the CPU table's column_claims (table index 0). + // CPU has interactions referencing columns 0-4, so column_claims is non-empty. + let cpu_proof = &mut multi_proof.proofs[0]; + if let Some(ref mut gkr_proof) = cpu_proof.logup_gkr_proof { + // Corrupt the first column claim value by adding 1. + assert!( + !gkr_proof.column_claims.is_empty(), + "CPU must have column claims" + ); + gkr_proof.column_claims[0].1 = + gkr_proof.column_claims[0].1.clone() + FieldElement::one(); + } else { + panic!("CPU table must have a logup_gkr_proof"); + } + + let air_refs: Vec<&dyn AIR> = + airs.iter().map(|a| a as &dyn AIR).collect(); assert!( !Verifier::multi_verify( - &airs, + &air_refs, &multi_proof, &mut DefaultTranscript::::new(&[]), &FieldElement::zero(), ), - "Proof with missing bus_public_inputs must be rejected" + "Tampered column_claims must cause verification failure (wrong bridge offset)" ); } -/// A proof where a non-LogUp AIR has bus_public_inputs injected is rejected. +/// Tampered GKR claimed_sum is rejected. /// -/// A dishonest prover could inject a compensating table_contribution into a -/// table that has no LogUp constraints, making the bus balance ΣL = 0 while -/// the actual lookup interactions don't balance. The verifier must reject -/// any proof that contains bus_public_inputs for an AIR without trace interaction. +/// The bus balance check enforces that the sum of all tables' GKR claimed_sums +/// equals zero. If we tamper with one table's claimed_sum, this balance check +/// will fail. Additionally, the tampered claimed_sum will desynchronize the +/// Fiat-Shamir transcript (since claimed_sum is appended during GKR verification), +/// causing further downstream verification failures. #[test_log::test] -fn test_injected_bus_public_inputs_on_non_logup_air_rejected() { - use crate::examples::dummy_air::{self, DummyAIR}; - use crate::lookup::BusPublicInputs; +fn test_tampered_gkr_claimed_sum_rejected() { + let (mut multi_proof, airs) = generate_valid_multi_proof(); - type DummyF = F; // GoldilocksField (same base and extension for DummyAIR) + // Tamper with the ADD table's GKR claimed_sum (table index 1). + let add_proof = &mut multi_proof.proofs[1]; + if let Some(ref mut gkr_proof) = add_proof.logup_gkr_proof { + gkr_proof.gkr_proof.claimed_sum = + gkr_proof.gkr_proof.claimed_sum.clone() + FieldElement::one(); + } else { + panic!("ADD table must have a logup_gkr_proof"); + } - let trace_length = 16; - let mut trace = dummy_air::dummy_trace::(trace_length); - let proof_options = ProofOptions::default_test_options(); - let air = DummyAIR::new(&proof_options); - - let mut proof = Prover::::prove( - &air, - &mut trace, - &(), - &mut DefaultTranscript::::new(&[]), - ) - .unwrap(); - - // Inject fake bus_public_inputs into a non-LogUp proof. - // DummyAIR has has_trace_interaction() = false, so this must be rejected. - proof.bus_public_inputs = Some(BusPublicInputs { - table_contribution: FieldElement::::from(42u64), - #[cfg(feature = "debug-checks")] - per_bus_sums: Default::default(), - #[cfg(feature = "debug-checks")] - per_bus_sender_sums: Default::default(), - #[cfg(feature = "debug-checks")] - per_bus_receiver_sums: Default::default(), - #[cfg(feature = "debug-checks")] - table_name: "FAKE".to_string(), - }); + let air_refs: Vec<&dyn AIR> = + airs.iter().map(|a| a as &dyn AIR).collect(); assert!( - !Verifier::::verify( - &proof, - &air, - &mut DefaultTranscript::::new(&[]) + !Verifier::multi_verify( + &air_refs, + &multi_proof, + &mut DefaultTranscript::::new(&[]), + &FieldElement::zero(), ), - "Proof with injected bus_public_inputs on non-LogUp AIR must be rejected" + "Tampered claimed_sum must cause verification failure (bus balance or GKR verify)" ); } -/// A proof where table_contribution is zeroed out is rejected. +/// Missing GKR proof for a table with bus interactions is rejected. /// -/// Setting table_contribution to zero changes the per-row offset L/N to zero, -/// which breaks the circular transition constraint since the committed trace -/// was built with the honest (non-zero) L. +/// When an AIR has bus interactions (has_trace_interaction() = true), the +/// verifier requires a logup_gkr_proof. Setting it to None should cause +/// immediate rejection in Phase B' where the verifier checks for the proof. #[test_log::test] -fn test_zeroed_table_contribution_rejected() { - let mut cpu_trace = TraceTable::from_columns_main( - vec![ - vec![FE::one(), FE::zero(), FE::zero(), FE::zero()], - vec![FE::zero(); 4], - vec![FE::from(5), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(3), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(8), FE::zero(), FE::zero(), FE::zero()], - ], - 1, - ); - let mut add_trace = TraceTable::from_columns_main( - vec![ - vec![FE::from(5), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(3), FE::zero(), FE::zero(), FE::zero()], - vec![FE::from(8), FE::zero(), FE::zero(), FE::zero()], - vec![FE::one(), FE::zero(), FE::zero(), FE::zero()], - ], - 1, - ); - let mut mul_trace = TraceTable::from_columns_main( - vec![ - vec![FE::zero(); 4], - vec![FE::zero(); 4], - vec![FE::zero(); 4], - vec![FE::zero(); 4], - ], - 1, +fn test_missing_gkr_proof_rejected() { + let (mut multi_proof, airs) = generate_valid_multi_proof(); + + // Remove the GKR proof from the ADD table (table index 1). + // ADD has bus interactions, so the verifier will reject. + multi_proof.proofs[1].logup_gkr_proof = None; + + let air_refs: Vec<&dyn AIR> = + airs.iter().map(|a| a as &dyn AIR).collect(); + + assert!( + !Verifier::multi_verify( + &air_refs, + &multi_proof, + &mut DefaultTranscript::::new(&[]), + &FieldElement::zero(), + ), + "Missing logup_gkr_proof for a table with interactions must be rejected" ); +} - let proof_options = ProofOptions::default_test_options(); - let cpu_air = new_cpu_air_with_lookup(&proof_options); - let add_air = new_add_air_with_lookup(&proof_options); - let mul_air = new_mul_air_with_lookup(&proof_options); +/// Tampered sigma (bridge running sum) OOD evaluation is rejected. +/// +/// The bridge running sum column (sigma, aux column 1) is constrained by the +/// LookupBridgeSumConstraint transition constraint: +/// sigma_next - sigma_curr - l_curr * batched_curr + bridge_offset = 0 +/// +/// Corrupting sigma's OOD evaluation breaks this constraint at the OOD point, +/// causing the composition polynomial check to fail. +#[test_log::test] +fn test_tampered_sigma_ood_rejected() { + let (mut multi_proof, airs) = generate_valid_multi_proof(); + + // Corrupt the sigma column OOD evaluation in the CPU table proof. + // CPU has 5 main columns + 2 aux columns (Lagrange=aux0, sigma=aux1). + // In trace_ood_evaluations, sigma is at column index 5 + 1 = 6. + let num_main_cpu = 5usize; + let sigma_col_ood_idx = num_main_cpu + 1; // aux column 1 = sigma + let cpu_proof = &mut multi_proof.proofs[0]; + let corrupted = *cpu_proof.trace_ood_evaluations.get(0, sigma_col_ood_idx) + FieldElement::one(); + cpu_proof + .trace_ood_evaluations + .set(0, sigma_col_ood_idx, corrupted); - let air_trace_pairs: Vec<( - &dyn AIR, - _, - _, - )> = vec![ - (&cpu_air, &mut cpu_trace, &()), - (&add_air, &mut add_trace, &()), - (&mul_air, &mut mul_trace, &()), - ]; + let air_refs: Vec<&dyn AIR> = + airs.iter().map(|a| a as &dyn AIR).collect(); - let mut multi_proof = - Prover::multi_prove(air_trace_pairs, &mut DefaultTranscript::::new(&[])).unwrap(); + assert!( + !Verifier::multi_verify( + &air_refs, + &multi_proof, + &mut DefaultTranscript::::new(&[]), + &FieldElement::zero(), + ), + "Tampered sigma OOD evaluation must cause verification failure (bridge constraint)" + ); +} - // Zero out table_contribution for the ADD table. - let add_proof = &mut multi_proof.proofs[1]; - let bus_inputs = add_proof - .bus_public_inputs - .as_mut() - .expect("ADD table must have bus public inputs"); - bus_inputs.table_contribution = FieldElement::::zero(); +/// Tampered Lagrange kernel random_point is rejected. +/// +/// The Lagrange kernel column l[i] is constrained by a boundary constraint: +/// l[0] = prod_{j=0}^{n-1} (1 - r_j) +/// where r_j are the GKR random point coordinates stored in the proof. +/// +/// If we tamper with the random_point in the proof, the verifier will compute +/// a different expected l[0] value from the (now-corrupted) random_point, which +/// will not match the committed l[0] in the trace. This causes the boundary +/// constraint check to fail. +#[test_log::test] +fn test_tampered_lagrange_kernel_random_point_rejected() { + let (mut multi_proof, airs) = generate_valid_multi_proof(); + + // Tamper with the random_point in the CPU table's LogUp-GKR proof. + let cpu_proof = &mut multi_proof.proofs[0]; + if let Some(ref mut gkr_proof) = cpu_proof.logup_gkr_proof { + assert!( + !gkr_proof.random_point.is_empty(), + "CPU must have a non-empty random_point" + ); + // Corrupt the first coordinate by adding 1. + gkr_proof.random_point[0] = + gkr_proof.random_point[0].clone() + FieldElement::one(); + } else { + panic!("CPU table must have a logup_gkr_proof"); + } - let airs: Vec<&dyn AIR> = - vec![&cpu_air, &add_air, &mul_air]; + let air_refs: Vec<&dyn AIR> = + airs.iter().map(|a| a as &dyn AIR).collect(); assert!( !Verifier::multi_verify( - &airs, + &air_refs, &multi_proof, &mut DefaultTranscript::::new(&[]), &FieldElement::zero(), ), - "Proof with zeroed table_contribution must be rejected" + "Tampered random_point must cause verification failure (Lagrange kernel boundary constraint)" ); } diff --git a/crypto/stark/src/traits.rs b/crypto/stark/src/traits.rs index b828ffeb0..4110fc5a0 100644 --- a/crypto/stark/src/traits.rs +++ b/crypto/stark/src/traits.rs @@ -12,7 +12,7 @@ use math::{ use crate::{ constraints::transition::TransitionConstraint, domain::Domain, - lookup::{BusPublicInputs, PackingShifts}, + lookup::{BusInteraction, BusPublicInputs, PackingShifts}, }; use super::{ @@ -181,6 +181,13 @@ pub trait AIR: Send + Sync { false } + /// Returns the bus interactions for this AIR. + /// Used by the GKR sub-protocol to compute leaf fractions and column claims. + /// Default implementation returns an empty slice. + fn bus_interactions(&self) -> &[BusInteraction] { + &[] + } + /// Returns the maximum number of bus elements across all interactions. /// Used to compute the correct number of alpha powers for LogUp fingerprints. fn max_bus_elements(&self) -> usize { diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index f595bc05a..3d529ef37 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -9,7 +9,10 @@ use super::{ use crate::{ config::Commitment, domain::new_verifier_domain, - lookup::{LOGUP_CHALLENGE_ALPHA, LOGUP_NUM_CHALLENGES, PackingShifts, compute_alpha_powers}, + lookup::{ + LOGUP_CHALLENGE_ALPHA, LOGUP_NUM_CHALLENGES, PackingShifts, compute_alpha_powers, + extend_rap_challenges_with_bridge, extract_column_indices, + }, proof::stark::{DeepPolynomialOpening, MultiProof}, }; use crypto::{fiat_shamir::is_transcript::IsStarkTranscript, merkle_tree::proof::Proof}; @@ -93,154 +96,6 @@ pub trait IsStarkVerifier< .collect::>() } - /// Returns the list of challenges sent to the prover. - fn step_1_replay_rounds_and_recover_challenges( - air: &dyn AIR, - proof: &StarkProof, - domain: &VerifierDomain, - transcript: &mut impl IsStarkTranscript, - ) -> Challenges - where - FieldElement: AsBytes, - FieldElement: AsBytes, - { - // =================================== - // ==========| Round 1 |========== - // =================================== - - // <<<< Receive commitments:[tⱼ] - transcript.append_bytes(&proof.lde_trace_main_merkle_root); - - let rap_challenges = air.build_rap_challenges(transcript); - - if let Some(root) = proof.lde_trace_aux_merkle_root { - transcript.append_bytes(&root); - } - - // =================================== - // ==========| Round 2 |========== - // =================================== - - // <<<< Receive challenge: 𝛽 - let beta = transcript.sample_field_element(); - let trace_length = proof.trace_length; - let num_boundary_constraints = air - .boundary_constraints( - &proof.public_inputs, - &rap_challenges, - proof.bus_public_inputs.as_ref(), - trace_length, - ) - .constraints - .len(); - - let num_transition_constraints = air.context().num_transition_constraints; - - let mut coefficients: Vec<_> = (0..num_boundary_constraints + num_transition_constraints) - .map(|i| beta.pow(i)) - .collect(); - - let transition_coeffs: Vec<_> = coefficients.drain(..num_transition_constraints).collect(); - let boundary_coeffs = coefficients; - - // <<<< Receive commitments: [H₁], [H₂] - transcript.append_bytes(&proof.composition_poly_root); - - // =================================== - // ==========| Round 3 |========== - // =================================== - - // >>>> Send challenge: z - let z = transcript.sample_z_ood_with_domain_params( - domain.trace_length, - domain.lde_length, - &domain.coset_offset, - ); - - // <<<< Receive values: tⱼ(zgᵏ) - let trace_ood_evaluations_columns = proof.trace_ood_evaluations.columns(); - for col in trace_ood_evaluations_columns.iter() { - for elem in col.iter() { - transcript.append_field_element(elem); - } - } - // <<<< Receive value: Hᵢ(z^N) - for element in proof.composition_poly_parts_ood_evaluation.iter() { - transcript.append_field_element(element); - } - - // =================================== - // ==========| Round 4 |========== - // =================================== - - let num_terms_composition_poly = proof.composition_poly_parts_ood_evaluation.len(); - let num_terms_trace = - air.context().transition_offsets.len() * air.step_size() * air.context().trace_columns; - let gamma = transcript.sample_field_element(); - - // <<<< Receive challenges: 𝛾, 𝛾' - let mut deep_composition_coefficients: Vec<_> = - core::iter::successors(Some(FieldElement::one()), |x| Some(x * &gamma)) - .take(num_terms_composition_poly + num_terms_trace) - .collect(); - - let trace_term_coeffs: Vec<_> = deep_composition_coefficients - .drain(..num_terms_trace) - .collect::>() - .chunks(air.context().transition_offsets.len() * air.step_size()) - .map(|chunk| chunk.to_vec()) - .collect(); - - // <<<< Receive challenges: 𝛾ⱼ, 𝛾ⱼ' - let gammas = deep_composition_coefficients; - - // FRI commit phase - let merkle_roots = &proof.fri_layers_merkle_roots; - let mut zetas = merkle_roots - .iter() - .map(|root| { - // >>>> Send challenge 𝜁ₖ - let element = transcript.sample_field_element(); - // <<<< Receive commitment: [pₖ] (the first one is [p₀]) - transcript.append_bytes(root); - element - }) - .collect::>>(); - - // >>>> Send challenge 𝜁ₙ₋₁ - zetas.push(transcript.sample_field_element()); - - // <<<< Receive value: pₙ - transcript.append_field_element(&proof.fri_last_value); - - // Receive grinding value - let security_bits = air.context().proof_options.grinding_factor; - let mut grinding_seed = [0u8; 32]; - if security_bits > 0 - && let Some(nonce_value) = proof.nonce - { - grinding_seed = transcript.state(); - transcript.append_bytes(&nonce_value.to_be_bytes()); - } - - // FRI query phase - // <<<< Send challenges 𝜄ₛ (iota_s) - let number_of_queries = air.options().fri_number_of_queries; - let iotas = Self::sample_query_indexes(number_of_queries, domain, transcript); - - Challenges { - z, - boundary_coeffs, - transition_coeffs, - trace_term_coeffs, - gammas, - zetas, - iotas, - rap_challenges, - grinding_seed, - } - } - /// Checks whether the purported evaluations of the composition polynomial parts and the trace /// polynomials at the out-of-domain challenge are consistent. /// See https://lambdaclass.github.io/lambdaworks/starks/protocol.html#step-2-verify-claimed-composition-polynomial @@ -924,38 +779,96 @@ pub trait IsStarkVerifier< }; // ===================================================================== - // Validate bus_public_inputs presence against AIR layout + // Phase B': Replay GKR proofs (main transcript) // ===================================================================== - // A dishonest prover could omit bus_public_inputs entirely (None) to - // bypass the bus balance check. With circular constraints, there are no - // boundary constraints on LogUp columns, so the bus balance check is - // the only cross-table validation. + // For each table with bus interactions, replay the GKR verification + // on the main transcript. This must match the prover's Phase B' exactly. + + let mut gkr_bridge_claims: Vec)>> = + Vec::with_capacity(airs.len()); + let mut gkr_random_points: Vec>> = + Vec::with_capacity(airs.len()); for (idx, (air, proof)) in airs.iter().zip(&multi_proof.proofs).enumerate() { - if air.has_trace_interaction() && proof.bus_public_inputs.is_none() { - error!( - "Table {idx}: AIR has LogUp interactions but proof is missing bus_public_inputs" - ); - return false; - } - if !air.has_trace_interaction() && proof.bus_public_inputs.is_some() { - error!( - "Table {idx}: AIR has no LogUp interactions but proof contains bus_public_inputs" - ); - return false; + if air.has_trace_interaction() { + let gkr_proof = match &proof.logup_gkr_proof { + Some(p) => p, + None => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!( + "Table {idx}: AIR has interactions but proof missing logup_gkr_proof" + ); + return false; + } + }; + + // Replay GKR verification on the main transcript + match crate::gkr::gkr_verify(&gkr_proof.gkr_proof, transcript) { + Ok((_random_point, n_claim, d_claim)) => { + // Validate column_claims length matches AIR-expected column set + let expected_cols = extract_column_indices(air.bus_interactions()); + if gkr_proof.column_claims.len() != expected_cols.len() { + #[cfg(not(feature = "test_fiat_shamir"))] + error!( + "Table {idx}: column_claims length mismatch: got {}, expected {}", + gkr_proof.column_claims.len(), + expected_cols.len() + ); + return false; + } + + // Verify that column_claims are consistent with GKR output. + // A malicious prover could provide fake column_claims that don't + // match the (n_claim, d_claim) returned by GKR verification. + if !crate::lookup::reconstruct_and_verify_gkr_claims( + &n_claim, + &d_claim, + &gkr_proof.column_claims, + air.bus_interactions(), + &lookup_challenges, + ) { + #[cfg(not(feature = "test_fiat_shamir"))] + error!( + "Table {idx}: GKR column claims verification failed — \ + column_claims are inconsistent with GKR output (n_claim, d_claim)" + ); + return false; + } + } + Err(_e) => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!("Table {idx}: GKR verification failed: {:?}", _e); + return false; + } + } + + gkr_bridge_claims.push(gkr_proof.column_claims.clone()); + gkr_random_points.push(gkr_proof.random_point.clone()); + } else { + gkr_bridge_claims.push(Vec::new()); + gkr_random_points.push(Vec::new()); } } + // ===================================================================== + // Phase B'': Sample γ for bridge batching (main transcript) + // ===================================================================== + // Must match prover: γ sampled AFTER all GKR messages, BEFORE forking. + + let gamma: FieldElement = if needs_lookup_challenges { + transcript.sample_field_element() + } else { + FieldElement::zero() + }; + // ===================================================================== // Phase C + Rounds 2-4: Forked per table // ===================================================================== // Each table gets an independent transcript fork (cloned from the shared - // state after Phase B, domain-separated by table index). This matches + // state after Phase B'', domain-separated by table index). This matches // the prover's forking and makes per-table verification independent. for (idx, (air, proof)) in airs.iter().zip(&multi_proof.proofs).enumerate() { - // Must match prover: fork with domain separator for multi-table, - // use original transcript directly for single-table. let num_tables = airs.len(); let mut table_transcript = transcript.clone(); if num_tables > 1 { @@ -967,18 +880,26 @@ pub trait IsStarkVerifier< table_transcript.append_bytes(&root); } - // Bind table_contribution (L) to transcript, matching prover. - if let Some(ref bpi) = proof.bus_public_inputs { - table_transcript.append_field_element(&bpi.table_contribution); + // Build per-table rap_challenges with bridge params + let mut table_rap_challenges = lookup_challenges.clone(); + if !gkr_bridge_claims[idx].is_empty() { + extend_rap_challenges_with_bridge( + &mut table_rap_challenges, + &gkr_bridge_claims[idx], + &gamma, + proof.trace_length, + &gkr_random_points[idx], + ); } - // Rounds 2-4: verify + // Rounds 2-4: verify (bridge is now a transition constraint) if !Self::verify_rounds_2_to_4( *air, proof, &mut table_transcript, - lookup_challenges.clone(), + table_rap_challenges, ) { + #[cfg(not(feature = "test_fiat_shamir"))] error!( "Table {} failed verify_rounds_2_to_4 (num_constraints={}, trace_cols={})", idx, @@ -992,9 +913,8 @@ pub trait IsStarkVerifier< // ===================================================================== // Bus Balance Check: Σ table_contribution = expected_bus_balance // ===================================================================== - // For LogUp with circular constraints, each table's total contribution L - // (sum of all per-row terms) is exposed as a public input. The bus balances - // when the sum of all table contributions equals the expected target. + // With LogUp-GKR, the table contribution is the GKR claimed_sum. + // The bus balances when the sum of all claimed_sums equals the expected target. // When all bus participants are in-trace, the target is zero. When some // receiver contributions are computed externally (e.g. verifier-computed // COMMIT output bus), the target is the missing positive remainder. @@ -1002,17 +922,15 @@ pub trait IsStarkVerifier< if needs_lookup_challenges { let mut total = FieldElement::::zero(); for (air, proof) in airs.iter().zip(&multi_proof.proofs) { - if air.has_trace_interaction() - && let Some(interaction) = &proof.bus_public_inputs - { - total = total + &interaction.table_contribution; + if air.has_trace_interaction() && let Some(ref gkr_proof) = proof.logup_gkr_proof { + total = total + &gkr_proof.gkr_proof.claimed_sum; } } if total != *expected_bus_balance { #[cfg(not(feature = "test_fiat_shamir"))] error!( - "LogUp bus does not balance: sum of accumulated values does not match target. total={:?}, target={:?}", + "LogUp bus does not balance: sum of GKR claimed_sums does not match target. total={:?}, target={:?}", total, expected_bus_balance ); return false; @@ -1074,9 +992,10 @@ pub trait IsStarkVerifier< let num_transition_constraints = air.context().num_transition_constraints; - let mut coefficients: Vec<_> = (0..num_boundary_constraints + num_transition_constraints) - .map(|i| beta.pow(i)) - .collect(); + let mut coefficients: Vec<_> = + (0..num_boundary_constraints + num_transition_constraints) + .map(|i| beta.pow(i)) + .collect(); let transition_coeffs: Vec<_> = coefficients.drain(..num_transition_constraints).collect(); let boundary_coeffs = coefficients; @@ -1202,8 +1121,13 @@ pub trait IsStarkVerifier< #[cfg(feature = "instruments")] let timer1 = Instant::now(); - let challenges = - Self::replay_rounds_after_round_1(air, proof, &domain, transcript, rap_challenges); + let challenges = Self::replay_rounds_after_round_1( + air, + proof, + &domain, + transcript, + rap_challenges, + ); // verify grinding let security_bits = air.context().proof_options.grinding_factor; From bbe4dbc4ca073d7501413276660f93fe30ccca57 Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 11:01:08 -0300 Subject: [PATCH 02/19] docs: add scaling improvements spec (unified sizing + MMCS + shared FRI) Three-item plan to reduce proving slope from ~5.1s/M to ~3-3.5s/M: 1. Uniform 2^20 table cap + twiddle dedup + FRI fold parallelism 2. MMCS batched commitment (Plonky3-style, all tables in shared trees) 3. Shared FRI instance across all tables --- .../2026-04-09-scaling-improvements-design.md | 344 ++++++++++++++++++ 1 file changed, 344 insertions(+) create mode 100644 docs/superpowers/specs/2026-04-09-scaling-improvements-design.md diff --git a/docs/superpowers/specs/2026-04-09-scaling-improvements-design.md b/docs/superpowers/specs/2026-04-09-scaling-improvements-design.md new file mode 100644 index 000000000..df12bd05f --- /dev/null +++ b/docs/superpowers/specs/2026-04-09-scaling-improvements-design.md @@ -0,0 +1,344 @@ +# Scaling Improvements: Unified Table Sizing + MMCS + Shared FRI + +## Problem + +Lambda VM's proving time scales with slope ~5.1s per million steps vs SP1's +~2.6s/M — roughly 2x worse. Three structural inefficiencies contribute: + +1. **Twiddle waste from diverse table sizes**: Tables at 3 different max_rows + (2^19, 2^20, 2^21) require 3 distinct twiddle sets, wasting 272 MB in + redundant copies and preventing domain sharing. + +2. **Per-table independent commitments**: Each table gets its own Merkle tree + for main trace, aux trace, and composition polynomial — ~30 separate Keccak + trees per proof. Each tree's construction and query opening is independent + overhead. + +3. **Per-table independent FRI**: Each table runs its own 19-layer FRI with its + own Merkle trees. For 12+ tables, that's ~200 FRI-layer trees and 219 + queries repeated per table. + +## Goals + +- Reduce the per-step proving slope by ~1.5-2x +- Uniform table domain size to eliminate twiddle diversity +- Batch all table commitments into shared Merkle trees (MMCS) +- Share one FRI instance across all tables +- Reduce proof size (fewer roots, fewer query openings) + +## Non-Goals + +- Changing the hash function (Keccak stays) +- Changing the field (Goldilocks stays) +- Execution sharding or recursion +- GKR-based LogUp (stay with committed aux columns) + +--- + +## Item 1: Unified Table Cap + Twiddle Dedup + Quick Wins + +### Uniform max_rows = 2^20 + +Cap all tables at 2^20 rows. Currently: +- 2^19: CPU (74 cols), MEMW (49), MEMW_A (30), DVRM (34) +- 2^20: MUL (26), SHIFT (26), BITWISE (21) +- 2^21: LT (15), LOAD (18), BRANCH (14), MEMW_R (10) + +Tables at 2^21 produce 2x more chunks (each 2^20) but each chunk's FFT is 2x +cheaper and all tables share one twiddle/domain. Tables at 2^19 are unchanged +(they already chunk below 2^20). The `MaxRowsConfig` formula +`max_rows = (127 × 2^19) / eff_width` stays, but we add +`.min(1 << 20)` to cap. + +Net effect: every chunk has domain size 2^20, LDE size 2^21. One shared +`LdeTwiddles` of 32 MB replaces 384 MB of redundant copies. + +### Twiddle deduplication + +In `multi_prove`'s pre-pass, deduplicate by domain size: build one +`Arc` per distinct `(trace_order, lde_order)` pair and clone +the `Arc` for same-size tables. With uniform cap, this collapses to 1 shared +set. + +### Parallelize FRI fold + +`fold_evaluations_in_place` in `fri/fri_functions.rs` is sequential. Replace +with `par_chunks_mut` (matching Plonky3's approach). For a 2^20 domain, this +parallelizes 2^19 fold operations. + +### Eliminate per-row Vec alloc in commit + +`commit_columns_bit_reversed` allocates a `Vec` per LDE row +(262K allocs per tree). Replace with `map_init` thread-local row buffer, +matching the existing `commit_composition_polynomial` pattern. + +### Files changed + +| File | Change | +|------|--------| +| `prover/src/tables/mod.rs` | Add `.min(1 << 20)` cap to max_rows | +| `crypto/stark/src/prover.rs` | Deduplicate twiddles, fix commit alloc | +| `crypto/stark/src/fri/fri_functions.rs` | Parallelize fold | + +--- + +## Item 2: MMCS Batched Commitment + +### Overview + +Replace per-table independent Merkle trees with Plonky3-style batched +commitments. All tables' main LDE columns go into one Merkle tree. All aux +columns into another. One composition tree for all tables. + +### How Plonky3 MMCS works + +Multiple matrices of different heights share one tree via "jagged" +construction: + +1. All tallest-height matrices have their rows concatenated and hashed together + into one leaf per row index. +2. When building upward, shorter matrices are "injected" at the tree level + matching their height — an extra compression step merges the injected + row hashes with the existing internal nodes. +3. Opening at a global index: for a matrix shorter than the tallest, the + index is right-shifted by `log2(max_height) - log2(matrix_height)`. + +With uniform cap (Item 1), all table chunks have the same height (2^20). +This simplifies MMCS to just concatenating all columns from all tables into +one wide row per leaf, with no jagged injection needed. + +### Adaptation for lambda_vm + +**Phase A — one batched main commitment:** + +Currently each table commits independently: +``` +for each table: + extract main columns → LDE → commit_columns_bit_reversed → root + append root to transcript +``` + +Replace with: +``` +for each table: + extract main columns → LDE → collect into all_main_columns[] +commit_columns_bit_reversed(all_main_columns) → one root +append one root to transcript +``` + +The single Merkle tree has leaves: +`leaf[i] = Keccak256(table0_col0[br_i] || table0_col1[br_i] || ... || tableN_colM[br_i])` + +This is exactly how `BatchedMerkleTreeBackend` already works — it hashes a +`Vec` per row. The only change is that the Vec contains columns +from ALL tables, not just one. + +**Phase C — one batched aux commitment:** + +Same approach: all aux columns from all tables into one tree. + +**Round 2 — one batched composition commitment:** + +All tables' composition polynomial parts committed together. + +**Proof format change:** + +```rust +// Before: MultiProof { proofs: Vec } +// After: +pub struct BatchedProof { + // Shared commitments + pub main_merkle_root: Commitment, + pub aux_merkle_root: Option, + pub composition_merkle_root: Commitment, + + // Per-table data (OOD evaluations, public inputs) + pub table_data: Vec>, + + // Shared FRI (Item 3) + pub fri_layers_merkle_roots: Vec, + pub fri_last_value: FieldElement, + + // Shared queries + pub query_openings: Vec>, + pub nonce: Option, +} + +pub struct TableProofData { + pub trace_length: usize, + pub trace_ood_evaluations: Table, + pub composition_poly_parts_ood_evaluation: Vec>, + pub bus_public_inputs: Option>, + pub public_inputs: PI, +} +``` + +**Query opening format:** + +Each query opens ONE row from the shared main tree, ONE from the shared aux +tree, and ONE from the shared composition tree — instead of opening from +each table's separate trees. The verifier extracts per-table columns from +the opened row by column offset. + +### Column offset tracking + +Each table's columns start at a known offset in the batched commitment. +The prover and verifier agree on the column layout: + +```rust +struct BatchedLayout { + // main_col_offset[table_idx] = starting column in the batched main tree + main_col_offsets: Vec, + // aux_col_offset[table_idx] = starting column in the batched aux tree + aux_col_offsets: Vec, + // comp_col_offset[table_idx] = starting column in the batched comp tree + comp_col_offsets: Vec, +} +``` + +### Files changed + +| File | Change | +|------|--------| +| `crypto/stark/src/prover.rs` | Batched commit in Phase A/C, batched comp in Round 2, batched openings in Round 4 | +| `crypto/stark/src/verifier.rs` | Verify against batched trees, extract per-table columns from opened rows | +| `crypto/stark/src/proof/stark.rs` | New `BatchedProof` struct | +| `crypto/stark/src/config.rs` | No change (same `BatchedMerkleTreeBackend`) | + +--- + +## Item 3: Shared FRI + +### Overview + +Replace per-table independent FRI with one shared FRI instance. All tables' +deep composition polynomials are randomly batched into one polynomial, which +is then FRI-committed and queried once. + +### How Plonky3 does it + +In `TwoAdicFriPcs::open()`: + +1. For each table, compute the DEEP quotient polynomial at each query point +2. Batch all tables' quotients into one evaluation vector using random `alpha` + powers +3. Run ONE `commit_phase_from_evaluations` on the batched vector +4. Query ONE set of indices; open from the shared MMCS trees + +With uniform table sizing (Item 1), all tables have the same domain. The +batching is a simple weighted sum: + +``` +batched[i] = Σ_tables alpha^t * deep_quotient_table_t[i] +``` + +### Adaptation for lambda_vm + +**Round 4 changes:** + +Currently: +``` +for each table: + compute deep composition poly evaluations + iFFT → FFT to extend to LDE + commit_phase_from_evaluations (19 FRI layers) + query_phase (219 queries) +``` + +Replace with: +``` +for each table: + compute deep composition poly evaluations → deep_evals[table] + +// Batch all tables with alpha powers +alpha = transcript.sample() +batched_evals = Σ alpha^t * deep_evals[t] + +// One FRI +commit_phase_from_evaluations(batched_evals) +// One query set +for each iota in iotas: + open from shared main tree, shared aux tree, shared comp tree + open from FRI layer trees +``` + +**Proof size reduction:** + +| Component | Before (per table) | After (shared) | +|-----------|-------------------|----------------| +| FRI layer roots | 19 × N_tables | 19 | +| FRI decommitments | 219 × N_tables | 219 | +| Trace openings | 219 × 2 proofs × N_tables | 219 × 2 proofs (one wide row) | +| Composition openings | 219 × 1 proof × N_tables | 219 × 1 proof | + +### Verifier changes + +The verifier reconstructs the batched deep quotient from the opened values +and verifies against the shared FRI. It also extracts per-table OOD +evaluations to verify constraint satisfaction independently per table. + +### Files changed + +| File | Change | +|------|--------| +| `crypto/stark/src/prover.rs` | Batched deep quotient, shared FRI commit/query | +| `crypto/stark/src/verifier.rs` | Verify shared FRI, reconstruct per-table quotients | +| `crypto/stark/src/proof/stark.rs` | Shared FRI fields in `BatchedProof` | +| `crypto/stark/src/fri/mod.rs` | No structural change (same `commit_phase_from_evaluations`) | + +--- + +## Expected Impact + +### Item 1 (quick wins) + +| Metric | Before | After | +|--------|--------|-------| +| Twiddle memory | 384 MB | 32 MB | +| FRI fold parallelism | sequential | par_chunks_mut | +| Commit allocs | 262K heap allocs/tree | 0 (thread-local buf) | +| Table chunk sizes | 3 different | 1 uniform | + +### Item 2 (MMCS) + +| Metric | Before | After | +|--------|--------|-------| +| Main Merkle trees | ~12 | 1 | +| Aux Merkle trees | ~10 | 1 | +| Composition trees | ~10 | 1 | +| Total commit-phase trees | ~32 | 3 | +| Keccak hash calls (commit) | ~32 × 2N | ~3 × 2N (wider rows) | + +### Item 3 (shared FRI) + +| Metric | Before | After | +|--------|--------|-------| +| FRI instances | ~12 | 1 | +| FRI layer trees | ~12 × 19 = 228 | 19 | +| Query openings | 219 × 12 = 2,628 | 219 | +| Proof size (FRI portion) | ~12x | 1x | + +### Combined slope estimate + +With all three items: the per-step cost drops by reducing the multiplier on +Merkle hashing (~12x → 1x for FRI, ~32x → 3x for commits) and FRI work. +Conservative estimate: slope drops from ~5.1s/M to ~3-3.5s/M, approaching +SP1's 2.6s/M. + +--- + +## Implementation Order + +1. **Item 1** first — prerequisite for Item 2 (uniform sizing simplifies MMCS) +2. **Item 2** next — prerequisite for Item 3 (shared trees enable shared FRI) +3. **Item 3** last — builds on Item 2's shared commitment infrastructure + +Each item is independently benchmarkable. Item 1 alone provides measurable +improvement. Items 2+3 together provide the largest structural gain. + +## Testing Strategy + +- Each item must pass all existing stark crate tests (121 tests) +- Items 2+3 change the proof format → verifier tests must be updated +- Benchmark after each item at 1M, 4M, 8M steps to track slope improvement +- Compare proof sizes before/after Items 2+3 From d56660211e8093b9db2c15d9e01f480cacf35f9f Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 11:07:10 -0300 Subject: [PATCH 03/19] perf(tables): cap LT/LOAD/BRANCH/MEMW_R max_rows at 2^20 Tables with effective width < 42 were sized at 2^21, producing single large chunks. Capping at 2^20 produces 2x more chunks of the standard size, improving parallel throughput and keeping peak memory per chunk uniform across all tables. --- prover/src/tables/mod.rs | 32 ++++++++++++++++---------------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/prover/src/tables/mod.rs b/prover/src/tables/mod.rs index 551dc4aa3..fc56ec22b 100644 --- a/prover/src/tables/mod.rs +++ b/prover/src/tables/mod.rs @@ -49,29 +49,29 @@ pub use types::BusId; /// (* MEMW_A formula gives 2^20, but set to 2^19 to match MEMW chunk geometry; /// benchmarks show better parallel throughput with smaller chunks.) /// -/// | Table | Main | Bus | Eff.width | Max rows | -/// |---------|------|-----|-----------|----------| -/// | MEMW | 49 | 26 | 127 | 2^19 | -/// | MEMW_A | 29 | 20 | 89 | 2^19 * | -/// | CPU | 74 | 40 | 194 | 2^19 | -/// | DVRM | 34 | 34 | 136 | 2^19 | -/// | MUL | 26 | 16 | 74 | 2^20 | -/// | LT | 15 | 9 | 42 | 2^21 | -/// | SHIFT | 27 | 15 | 72 | 2^20 | -/// | LOAD | 18 | 5 | 33 | 2^21 | -/// | BRANCH | 14 | 6 | 32 | 2^21 | -/// | MEMW_R | 10 | 7 | 31 | 2^21 | +/// | Table | Main | Bus | Eff.width | Max rows | +/// |---------|------|-----|-----------|-----------------| +/// | MEMW | 49 | 26 | 127 | 2^19 | +/// | MEMW_A | 29 | 20 | 89 | 2^19 * | +/// | CPU | 74 | 40 | 194 | 2^19 | +/// | DVRM | 34 | 34 | 136 | 2^19 | +/// | MUL | 26 | 16 | 74 | 2^20 | +/// | LT | 15 | 9 | 42 | 2^20 (capped) | +/// | SHIFT | 27 | 15 | 72 | 2^20 | +/// | LOAD | 18 | 5 | 33 | 2^20 (capped) | +/// | BRANCH | 14 | 6 | 32 | 2^20 (capped) | +/// | MEMW_R | 10 | 7 | 31 | 2^20 (capped) | pub mod max_rows { pub const CPU: usize = 1 << 19; // 524,288 — eff. width 194 pub const MEMW: usize = 1 << 19; // 524,288 — eff. width 127 (baseline) pub const MEMW_A: usize = 1 << 19; // 524,288 — eff. width 89 pub const DVRM: usize = 1 << 19; // 524,288 — eff. width 136 pub const MUL: usize = 1 << 20; // 1,048,576 — eff. width 74 - pub const LT: usize = 1 << 21; // 2,097,152 — eff. width 42 + pub const LT: usize = 1 << 20; // 1,048,576 — eff. width 42 (capped at 2^20) pub const SHIFT: usize = 1 << 20; // 1,048,576 — eff. width 72 - pub const LOAD: usize = 1 << 21; // 2,097,152 — eff. width 33 - pub const BRANCH: usize = 1 << 21; // 2,097,152 — eff. width 32 - pub const MEMW_R: usize = 1 << 21; // 2,097,152 — eff. width 31 + pub const LOAD: usize = 1 << 20; // 1,048,576 — eff. width 33 (capped at 2^20) + pub const BRANCH: usize = 1 << 20; // 1,048,576 — eff. width 32 (capped at 2^20) + pub const MEMW_R: usize = 1 << 20; // 1,048,576 — eff. width 31 (capped at 2^20) } /// Per-table maximum row limits, configurable for different environments. From c3f5719b9000328e43b48d74bae8267bde16a7f5 Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 11:07:36 -0300 Subject: [PATCH 04/19] perf(stark): deduplicate LdeTwiddles by domain size in multi_prove Tables sharing the same lde_size now reuse a single Arc instead of each allocating their own copy. This avoids redundant twiddle computation and allocation when multiple tables happen to have the same trace length (e.g., all 2^19-row tables after the cap change). --- crypto/stark/src/prover.rs | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index 50af95525..3b26861c9 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -635,7 +635,7 @@ pub trait IsStarkProver< air_trace_pairs: &[AirTracePair<'_, Field, FieldExtension, PI>], metadatas: &[Round1Metadata], domains: &[Domain], - twiddle_caches: &[LdeTwiddles], + twiddle_caches: &[Arc>], main_pool: &mut [Vec>], aux_pool: &mut [Vec>], ) where @@ -651,7 +651,7 @@ pub trait IsStarkProver< .zip(domains.iter().zip(twiddle_caches.iter())) { let result = Self::reconstruct_round1( - *air, *trace, domain, metadata, twiddles, main_pool, aux_pool, + *air, *trace, domain, metadata, &**twiddles, main_pool, aux_pool, ) .expect("reconstruct_round1 failed in debug-checks"); temp_results.push(result); @@ -1512,24 +1512,30 @@ pub trait IsStarkProver< let phase_start = Instant::now(); let mut domains = Vec::with_capacity(num_airs); - let mut twiddle_caches: Vec> = Vec::with_capacity(num_airs); + let mut twiddle_caches: Vec>> = Vec::with_capacity(num_airs); let mut max_main_cols = 0usize; let mut max_aux_cols = 0usize; let mut max_lde_size = 0usize; + // Deduplicate twiddle caches: tables with the same lde_size share one Arc. + let mut twiddle_by_size: std::collections::HashMap>> = + std::collections::HashMap::new(); + for (air, trace, _pub_inputs) in &*air_trace_pairs { let trace_length = trace.num_rows(); let domain = new_domain(*air, trace_length); let lde_size = domain.interpolation_domain_size * domain.blowup_factor; - let twiddles = LdeTwiddles::new(&domain); + let twiddles = twiddle_by_size + .entry(lde_size) + .or_insert_with(|| Arc::new(LdeTwiddles::new(&domain))); max_main_cols = max_main_cols.max(trace.num_main_columns); max_aux_cols = max_aux_cols.max(air.num_auxiliary_rap_columns()); max_lde_size = max_lde_size.max(lde_size); domains.push(domain); - twiddle_caches.push(twiddles); + twiddle_caches.push(Arc::clone(twiddles)); } // Allocate K independent LDE column buffer pool sets for parallel table processing. @@ -1573,7 +1579,7 @@ pub trait IsStarkProver< let idx = chunk_start + j; let (air, trace, _) = &air_trace_pairs[idx]; let domain = &domains[idx]; - let twiddles = &twiddle_caches[idx]; + let twiddles = &*twiddle_caches[idx]; if air.is_preprocessed() { Self::commit_preprocessed_trace( @@ -1693,7 +1699,7 @@ pub trait IsStarkProver< let idx = chunk_start + j; let (air, trace, _) = &air_trace_pairs[idx]; let domain = &domains[idx]; - let twiddles = &twiddle_caches[idx]; + let twiddles = &*twiddle_caches[idx]; if air.has_aux_trace() { let num_aux_cols = trace.num_aux_columns; @@ -1809,7 +1815,7 @@ pub trait IsStarkProver< let (air, trace, pub_inputs) = &air_trace_pairs[idx]; let metadata = &metadatas[idx]; let domain = &domains[idx]; - let twiddles = &twiddle_caches[idx]; + let twiddles = &*twiddle_caches[idx]; #[cfg(feature = "instruments")] let table_start = Instant::now(); From 3aa03e6c06549f9605956c5e9b912d09dbe97fea Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 11:10:12 -0300 Subject: [PATCH 05/19] perf(fri): parallelize fold_evaluations_in_place with Rayon Under the `parallel` feature, fold each conjugate pair concurrently using `par_chunks(2).zip(inv_twiddles.par_iter())`, collecting into a new Vec of half length. The sequential path (no `parallel` feature) is unchanged. This avoids the index aliasing issue from an in-place interleaved write. --- crypto/stark/src/fri/fri_functions.rs | 37 ++++++++++++++++++++++----- 1 file changed, 30 insertions(+), 7 deletions(-) diff --git a/crypto/stark/src/fri/fri_functions.rs b/crypto/stark/src/fri/fri_functions.rs index 8bd355ec4..e96cbf5d2 100644 --- a/crypto/stark/src/fri/fri_functions.rs +++ b/crypto/stark/src/fri/fri_functions.rs @@ -18,14 +18,37 @@ pub fn fold_evaluations_in_place, E: IsField>( inv_twiddles: &[FieldElement], ) { let half = evals.len() / 2; - for j in 0..half { - let lo = &evals[2 * j]; - let hi = &evals[2 * j + 1]; - let sum = lo + hi; - let diff = lo - hi; - evals[j] = &sum + &(&inv_twiddles[j] * &(zeta * &diff)); + + #[cfg(feature = "parallel")] + { + use rayon::prelude::*; + // Evaluations are stored interleaved: pairs (evals[2j], evals[2j+1]). + // Fold each pair in parallel, collecting into a new Vec of half length. + let folded: Vec> = evals + .par_chunks(2) + .zip(inv_twiddles.par_iter()) + .map(|(pair, tw)| { + let lo = &pair[0]; + let hi = &pair[1]; + let sum = lo + hi; + let diff = lo - hi; + &sum + &(tw * &(zeta * &diff)) + }) + .collect(); + *evals = folded; + } + + #[cfg(not(feature = "parallel"))] + { + for j in 0..half { + let lo = &evals[2 * j]; + let hi = &evals[2 * j + 1]; + let sum = lo + hi; + let diff = lo - hi; + evals[j] = &sum + &(&inv_twiddles[j] * &(zeta * &diff)); + } + evals.truncate(half); } - evals.truncate(half); } /// Compute inverse twiddle factors for evaluation-form FRI folding. From ccb75483fad87a174510ddcb9cdea494cc949b31 Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 11:10:35 -0300 Subject: [PATCH 06/19] perf(prover): eliminate per-row Vec alloc in commit_columns_bit_reversed Replace the per-row `Vec` allocation inside the iterator with a thread-local buffer via Rayon's `map_init` (parallel path) or a single buffer allocated outside the loop (sequential path). This avoids one heap allocation per row during Merkle tree construction. --- crypto/stark/src/prover.rs | 40 ++++++++++++++++++++++++++------------ 1 file changed, 28 insertions(+), 12 deletions(-) diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index 3b26861c9..798bb3af0 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -335,19 +335,35 @@ pub trait IsStarkProver< let num_cols = columns.len(); #[cfg(feature = "parallel")] - let iter = (0..num_rows).into_par_iter(); - #[cfg(not(feature = "parallel"))] - let iter = 0..num_rows; + let hashed_leaves: Vec = { + (0..num_rows) + .into_par_iter() + .map_init( + || vec![FieldElement::::zero(); num_cols], + |row_buf, row_idx| { + let br_idx = reverse_index(row_idx, num_rows as u64); + for col_idx in 0..num_cols { + row_buf[col_idx] = columns[col_idx][br_idx].clone(); + } + BatchedMerkleTreeBackend::::hash_data(row_buf) + }, + ) + .collect() + }; - let hashed_leaves: Vec = iter - .map(|row_idx| { - let br_idx = reverse_index(row_idx, num_rows as u64); - let row: Vec> = (0..num_cols) - .map(|col_idx| columns[col_idx][br_idx].clone()) - .collect(); - BatchedMerkleTreeBackend::::hash_data(&row) - }) - .collect(); + #[cfg(not(feature = "parallel"))] + let hashed_leaves: Vec = { + let mut row_buf = vec![FieldElement::::zero(); num_cols]; + (0..num_rows) + .map(|row_idx| { + let br_idx = reverse_index(row_idx, num_rows as u64); + for col_idx in 0..num_cols { + row_buf[col_idx] = columns[col_idx][br_idx].clone(); + } + BatchedMerkleTreeBackend::::hash_data(&row_buf) + }) + .collect() + }; let tree = BatchedMerkleTree::::build_from_hashed_leaves(hashed_leaves)?; let root = tree.root; From 36150b819af40fd654118975a0817188fee1fa6b Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 15:59:52 -0300 Subject: [PATCH 07/19] perf(gkr): parallelize fold_table inner loop with par_chunks(2) Replace 4-way rayon::join across tables with par_chunks(2) inside fold_table so all available Rayon threads participate in folding a single table when half >= 256, instead of splitting work across only 4 tables. Sequential path unchanged for half < 256. --- crypto/stark/src/gkr.rs | 51 +++++++++++++++-------------------------- 1 file changed, 18 insertions(+), 33 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 97aa83d4d..d8f071f13 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -658,7 +658,21 @@ pub fn gkr_prove( // Fold the four gate tables (eq_table was already halved before inner loop) // table[j] = table[2j] + challenge * (table[2j+1] - table[2j]) + // + // When half >= 256 we use par_chunks(2) so that all Rayon threads + // participate in folding a single table (instead of just 4-way + // parallelism across the four independent tables). let fold_table = |table: &mut Vec>| { + #[cfg(feature = "parallel")] + if half >= 256 { + let folded: Vec> = table + .par_chunks(2) + .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) + .collect(); + *table = folded; + return; + } + // Sequential fallback (also used when half < 256) for j in 0..half { let left = &table[2 * j]; let right = &table[2 * j + 1]; @@ -667,39 +681,10 @@ pub fn gkr_prove( table.truncate(half); }; - #[cfg(feature = "parallel")] - { - if half >= 256 { - // Fold tables in parallel (each table independently) - rayon::join( - || { - rayon::join( - || fold_table(&mut nl_table), - || fold_table(&mut nr_table), - ); - }, - || { - rayon::join( - || fold_table(&mut dl_table), - || fold_table(&mut dr_table), - ); - }, - ); - } else { - fold_table(&mut nl_table); - fold_table(&mut nr_table); - fold_table(&mut dl_table); - fold_table(&mut dr_table); - } - } - - #[cfg(not(feature = "parallel"))] - { - fold_table(&mut nl_table); - fold_table(&mut nr_table); - fold_table(&mut dl_table); - fold_table(&mut dr_table); - } + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); round_polys.push(round_poly); challenges.push(challenge); From 9b038bfb60287278a8b8d816b97f64c34c212094 Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 16:28:02 -0300 Subject: [PATCH 08/19] feat(gkr): wire batch GKR into STARK prover and verifier Replace per-table sequential gkr_prove calls with a single gkr_prove_batch that processes all tables simultaneously. - Add compute_logup_layers() and finalize_logup_gkr_result() to lookup.rs for the two-phase batch GKR workflow - Make LogUpGkrResult::gkr_proof optional (None in batch mode) - Add BatchGkrProof to MultiProof for shared batch proof storage - Prover: build layer trees per table, call gkr_prove_batch once, distribute per-instance results via instance_eval_point - Verifier: call gkr_verify_batch once, distribute claims to tables - Update soundness tests for batch proof tampering --- crypto/stark/src/gkr.rs | 2 +- crypto/stark/src/lookup.rs | 92 ++++++++++++- crypto/stark/src/proof/stark.rs | 7 +- crypto/stark/src/prover.rs | 110 ++++++++++++---- .../src/tests/bus_tests/soundness_tests.rs | 55 ++++---- crypto/stark/src/verifier.rs | 124 ++++++++++++------ 6 files changed, 299 insertions(+), 91 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index d8f071f13..b7879d78f 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -1517,7 +1517,7 @@ pub fn gkr_prove_batch( /// `[shared_point[0] (eta)] ++ shared_point[len - (n_vars - 1)..]` /// /// Returns an empty vec for n_vars == 0 (trivial/output layers). -fn instance_eval_point( +pub(crate) fn instance_eval_point( shared_point: &[FieldElement], n_vars: usize, ) -> Vec> { diff --git a/crypto/stark/src/lookup.rs b/crypto/stark/src/lookup.rs index 11525fa47..850d5ce9c 100644 --- a/crypto/stark/src/lookup.rs +++ b/crypto/stark/src/lookup.rs @@ -2349,20 +2349,24 @@ where // LogUp-GKR Integration // ============================================================================= -use crate::gkr::{build_summation_tree, gkr_prove, GkrProof}; +use crate::gkr::{build_summation_tree, gen_layers, gkr_prove, GkrProof, Layer}; use crate::lagrange_kernel::{compute_lagrange_kernel, eval_mle_base_with_kernel}; use crypto::fiat_shamir::is_transcript::IsTranscript; /// Result of running the LogUp-GKR sub-protocol for a single table. /// -/// Contains the GKR proof, the random evaluation point, leaf-level claims, -/// and MLE claims for each distinct main trace column used in bus interactions. +/// Contains the GKR proof (if single-instance mode), the random evaluation point, +/// leaf-level claims, and MLE claims for each distinct main trace column used in +/// bus interactions. +/// +/// In batch mode, the GKR proof lives in `MultiProof::batch_gkr_proof` and +/// `gkr_proof` is `None`. #[derive(Debug, Clone)] pub struct LogUpGkrResult { /// Total table contribution (claimed_sum from GKR root = sum of all fractions). pub table_contribution: FieldElement, - /// The complete GKR proof for the summation tree. - pub gkr_proof: GkrProof, + /// The complete GKR proof for the summation tree (None in batch mode). + pub gkr_proof: Option>, /// The random evaluation point produced by the GKR protocol (length = log2(trace_len)). pub random_point: Vec>, /// Claimed MLE evaluation of the leaf numerator at the random point. @@ -2727,6 +2731,82 @@ fn multiplicity_from_claims( } } +/// Compute the GKR layer tree for a single table's bus interactions. +/// +/// This performs steps 1-2 of the LogUp-GKR protocol: +/// 1. Computes per-row leaf fractions (numerator, denominator) from interactions +/// 2. Builds the full layer tree via `gen_layers` +/// +/// Returns `Vec>` suitable for passing to `gkr_prove_batch`. +pub fn compute_logup_layers( + interactions: &[BusInteraction], + main_segment_cols: &[Vec>], + trace_len: usize, + challenges: &[FieldElement], +) -> Vec> +where + F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, + E: IsField + Send + Sync, +{ + let (numerators, denominators) = + compute_logup_leaf_fractions(interactions, main_segment_cols, trace_len, challenges); + + let input_layer = Layer::LogUpGeneric { + numerators, + denominators, + }; + gen_layers(input_layer) +} + +/// Finalize a LogUp-GKR result from batch GKR output for a single table. +/// +/// This performs steps 4-5 of the LogUp-GKR protocol: +/// 4. Extracts MLE claims for each distinct main trace column at the random point +/// 5. Computes the Lagrange kernel for reuse in aux trace construction +/// +/// Called after `gkr_prove_batch` distributes per-instance results. +pub fn finalize_logup_gkr_result( + interactions: &[BusInteraction], + main_segment_cols: &[Vec>], + random_point: Vec>, + n_claim: FieldElement, + d_claim: FieldElement, + table_contribution: FieldElement, +) -> LogUpGkrResult +where + F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, + E: IsField + Send + Sync, +{ + // Extract column claims — compute MLE at the random point for each + // distinct main trace column index referenced by any interaction. + let col_indices = extract_column_indices(interactions); + + // Compute kernel once and reuse for all column claims (and later for aux trace) + let kernel = compute_lagrange_kernel(&random_point); + + #[cfg(feature = "parallel")] + let col_iter = col_indices.into_par_iter(); + #[cfg(not(feature = "parallel"))] + let col_iter = col_indices.into_iter(); + + let column_claims: Vec<(usize, FieldElement)> = col_iter + .map(|col_idx| { + let claim = eval_mle_base_with_kernel(&main_segment_cols[col_idx], &kernel); + (col_idx, claim) + }) + .collect(); + + LogUpGkrResult { + table_contribution, + gkr_proof: None, + random_point, + n_claim, + d_claim, + column_claims, + lagrange_kernel: kernel, + } +} + /// Run the LogUp-GKR sub-protocol for a single table's bus interactions. /// /// This function: @@ -2827,7 +2907,7 @@ where LogUpGkrResult { table_contribution, - gkr_proof, + gkr_proof: Some(gkr_proof), random_point, n_claim, d_claim, diff --git a/crypto/stark/src/proof/stark.rs b/crypto/stark/src/proof/stark.rs index 736ce3b48..107f97290 100644 --- a/crypto/stark/src/proof/stark.rs +++ b/crypto/stark/src/proof/stark.rs @@ -5,7 +5,7 @@ use math::field::{ }; use crate::{ - config::Commitment, fri::fri_decommit::FriDecommitment, gkr::GkrProof, + config::Commitment, fri::fri_decommit::FriDecommitment, gkr::{BatchGkrProof, GkrProof}, lookup::BusPublicInputs, table::Table, }; @@ -93,4 +93,9 @@ pub struct StarkProof, E: IsField, PI> { #[serde(bound = "PI: serde::Serialize + serde::de::DeserializeOwned")] pub struct MultiProof, E: IsField, PI> { pub proofs: Vec>, + /// Batch GKR proof when multiple tables share a single GKR verification. + /// Per-table `logup_gkr_proof` contains the column_claims and random_point, + /// while this field contains the shared batch proof. + #[serde(default)] + pub batch_gkr_proof: Option>, } diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index 097d92d86..fc6f34b20 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -1650,31 +1650,87 @@ pub trait IsStarkProver< }; // ===================================================================== - // Round 1, Phase B': GKR sub-protocol (main transcript) + // Round 1, Phase B': Batch GKR sub-protocol (main transcript) // ===================================================================== - // For each table with bus interactions, run the LogUp-GKR sub-protocol. + // Compute leaf layers for all tables, then run a single batch GKR proof. // GKR messages are bound to the main Fiat-Shamir chain so the verifier // can replay them. - let gkr_results: Vec>> = air_trace_pairs - .iter() - .map(|(air, trace, _)| { - if air.has_trace_interaction() { - let interactions = air.bus_interactions(); - let main_segment_cols = trace.columns_main(); - let trace_len = trace.num_rows(); - Some(crate::lookup::run_logup_gkr( - interactions, - &main_segment_cols, - trace_len, - &lookup_challenges, - transcript, - )) + // Step 1: Compute GKR layer trees for each table with bus interactions. + let mut leaf_layers_per_table: Vec>> = Vec::new(); + let mut gkr_table_indices: Vec = Vec::new(); + for (idx, (air, trace, _)) in air_trace_pairs.iter().enumerate() { + if air.has_trace_interaction() { + let interactions = air.bus_interactions(); + let main_segment_cols = trace.columns_main(); + let trace_len = trace.num_rows(); + let layers = crate::lookup::compute_logup_layers( + interactions, + &main_segment_cols, + trace_len, + &lookup_challenges, + ); + leaf_layers_per_table.push(layers); + gkr_table_indices.push(idx); + } + } + + // Step 2: Run batch GKR (single proof for all tables). + let batch_gkr_proof = if !leaf_layers_per_table.is_empty() { + let n_layers_by_instance: Vec = leaf_layers_per_table + .iter() + .map(|layers| layers.len() - 1) + .collect(); + let (batch_proof, shared_random_point, per_instance_claims) = + crate::gkr::gkr_prove_batch(leaf_layers_per_table, transcript); + + // Step 3: Distribute batch results to per-table LogUpGkrResult. + let mut gkr_results_vec: Vec>> = + vec![None; num_airs]; + for (i, &table_idx) in gkr_table_indices.iter().enumerate() { + let (n_claim, d_claim) = per_instance_claims[i].clone(); + let table_contribution = batch_proof.root_claims[i].clone(); + // Compute table_contribution as n/d (the rational root value). + // The root_claims store (numerator, denominator) of the summation tree root. + let (root_n, root_d) = table_contribution; + + // The instance eval point for this table + let n_vars = n_layers_by_instance[i]; + let instance_point = + crate::gkr::instance_eval_point(&shared_random_point, n_vars); + + let air = air_trace_pairs[table_idx].0; + let trace = &air_trace_pairs[table_idx].1; + let interactions = air.bus_interactions(); + let main_segment_cols = trace.columns_main(); + + // Compute the claimed_sum (table contribution) as a single field element. + // For the bus balance check, we need n/d as an element. + // The summation tree root stores n/d where the claimed_sum = n/d. + // We compute it by inverting d and multiplying by n. + let table_contrib = if root_d == FieldElement::one() { + root_n } else { - None - } - }) - .collect(); + &root_n * &root_d.inv().expect("GKR root denominator must be non-zero") + }; + + let result = crate::lookup::finalize_logup_gkr_result( + interactions, + &main_segment_cols, + instance_point, + n_claim, + d_claim, + table_contrib, + ); + gkr_results_vec[table_idx] = Some(result); + } + + (Some(batch_proof), gkr_results_vec) + } else { + (None, vec![None; num_airs]) + }; + + let (batch_gkr_proof_opt, gkr_results) = batch_gkr_proof; // ===================================================================== // Round 1, Phase B'': Sample γ for bridge batching (main transcript) @@ -1966,9 +2022,14 @@ pub trait IsStarkProver< domain, )?; - // Attach LogUp-GKR proof from Phase B' (if this table had bus interactions) + // Attach LogUp-GKR proof metadata from Phase B'. + // In batch mode, gkr_proof is None (the real proof is in MultiProof::batch_gkr_proof). + // We still store random_point and column_claims for per-table verification. proof.logup_gkr_proof = metadata.logup_gkr_result.as_ref().map(|r| LogUpGkrProof { - gkr_proof: r.gkr_proof.clone(), + gkr_proof: r.gkr_proof.clone().unwrap_or_else(|| crate::gkr::GkrProof { + claimed_sum: r.table_contribution.clone(), + layer_proofs: vec![], + }), random_point: r.random_point.clone(), column_claims: r.column_claims.clone(), }); @@ -2033,7 +2094,10 @@ pub trait IsStarkProver< }); } - Ok(MultiProof { proofs }) + Ok(MultiProof { + proofs, + batch_gkr_proof: batch_gkr_proof_opt, + }) } /// Generate a STARK proof for a single AIR/trace. diff --git a/crypto/stark/src/tests/bus_tests/soundness_tests.rs b/crypto/stark/src/tests/bus_tests/soundness_tests.rs index a9602a5ed..8162909cb 100644 --- a/crypto/stark/src/tests/bus_tests/soundness_tests.rs +++ b/crypto/stark/src/tests/bus_tests/soundness_tests.rs @@ -805,13 +805,19 @@ fn test_tampered_gkr_column_claims_rejected() { fn test_tampered_gkr_claimed_sum_rejected() { let (mut multi_proof, airs) = generate_valid_multi_proof(); - // Tamper with the ADD table's GKR claimed_sum (table index 1). - let add_proof = &mut multi_proof.proofs[1]; - if let Some(ref mut gkr_proof) = add_proof.logup_gkr_proof { - gkr_proof.gkr_proof.claimed_sum = - gkr_proof.gkr_proof.claimed_sum.clone() + FieldElement::one(); + // Tamper with the batch GKR proof's root_claims for the ADD table (instance 1). + // In batch mode, the bus balance check uses batch_gkr_proof.root_claims, + // not the per-table logup_gkr_proof.gkr_proof.claimed_sum. + if let Some(ref mut batch_proof) = multi_proof.batch_gkr_proof { + assert!( + batch_proof.root_claims.len() > 1, + "batch proof must have multiple root_claims" + ); + // Tamper the ADD table's root claim (instance index 1, numerator) + batch_proof.root_claims[1].0 = + batch_proof.root_claims[1].0.clone() + FieldElement::one(); } else { - panic!("ADD table must have a logup_gkr_proof"); + panic!("MultiProof must have a batch_gkr_proof"); } let air_refs: Vec<&dyn AIR> = @@ -892,32 +898,37 @@ fn test_tampered_sigma_ood_rejected() { ); } -/// Tampered Lagrange kernel random_point is rejected. +/// Tampered batch GKR child claims alter random_point, causing rejection. /// /// The Lagrange kernel column l[i] is constrained by a boundary constraint: /// l[0] = prod_{j=0}^{n-1} (1 - r_j) -/// where r_j are the GKR random point coordinates stored in the proof. +/// where r_j are the GKR random point coordinates derived from the batch proof. /// -/// If we tamper with the random_point in the proof, the verifier will compute -/// a different expected l[0] value from the (now-corrupted) random_point, which -/// will not match the committed l[0] in the trace. This causes the boundary -/// constraint check to fail. +/// If we tamper with a child_claims entry in the batch GKR proof, the verifier +/// will derive a different random_point, which will not match the committed +/// Lagrange kernel in the trace. This causes the GKR gate check or downstream +/// boundary constraint check to fail. #[test_log::test] fn test_tampered_lagrange_kernel_random_point_rejected() { let (mut multi_proof, airs) = generate_valid_multi_proof(); - // Tamper with the random_point in the CPU table's LogUp-GKR proof. - let cpu_proof = &mut multi_proof.proofs[0]; - if let Some(ref mut gkr_proof) = cpu_proof.logup_gkr_proof { + // Tamper with the batch GKR proof's child_claims in the first layer. + // This corrupts the GKR verification, changing the derived random_point. + if let Some(ref mut batch_proof) = multi_proof.batch_gkr_proof { assert!( - !gkr_proof.random_point.is_empty(), - "CPU must have a non-empty random_point" + !batch_proof.layer_proofs.is_empty(), + "batch proof must have layer_proofs" ); - // Corrupt the first coordinate by adding 1. - gkr_proof.random_point[0] = - gkr_proof.random_point[0].clone() + FieldElement::one(); + // Corrupt the first child claim of the first layer proof. + let layer = &mut batch_proof.layer_proofs[0]; + assert!( + !layer.child_claims_by_instance.is_empty(), + "layer must have child_claims" + ); + layer.child_claims_by_instance[0][0] = + layer.child_claims_by_instance[0][0].clone() + FieldElement::one(); } else { - panic!("CPU table must have a logup_gkr_proof"); + panic!("MultiProof must have a batch_gkr_proof"); } let air_refs: Vec<&dyn AIR> = @@ -930,7 +941,7 @@ fn test_tampered_lagrange_kernel_random_point_rejected() { &mut DefaultTranscript::::new(&[]), &FieldElement::zero(), ), - "Tampered random_point must cause verification failure (Lagrange kernel boundary constraint)" + "Tampered batch GKR child claims must cause verification failure" ); } diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index c468404dc..119934973 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -791,72 +791,111 @@ pub trait IsStarkVerifier< }; // ===================================================================== - // Phase B': Replay GKR proofs (main transcript) + // Phase B': Replay batch GKR proof (main transcript) // ===================================================================== - // For each table with bus interactions, replay the GKR verification - // on the main transcript. This must match the prover's Phase B' exactly. + // Verify the batch GKR proof once, then distribute per-instance claims + // to each table. This must match the prover's Phase B' exactly. let mut gkr_bridge_claims: Vec)>> = Vec::with_capacity(airs.len()); let mut gkr_random_points: Vec>> = Vec::with_capacity(airs.len()); - for (idx, (air, proof)) in airs.iter().zip(&multi_proof.proofs).enumerate() { + // Collect which tables have interactions (and their n_layers) + let mut gkr_table_indices: Vec = Vec::new(); + let mut n_layers_by_instance: Vec = Vec::new(); + for (idx, air) in airs.iter().enumerate() { if air.has_trace_interaction() { - let gkr_proof = match &proof.logup_gkr_proof { - Some(p) => p, - None => { - #[cfg(not(feature = "test_fiat_shamir"))] - error!( - "Table {idx}: AIR has interactions but proof missing logup_gkr_proof" - ); - return false; + gkr_table_indices.push(idx); + let trace_len = multi_proof.proofs[idx].trace_length; + let n_vars = trace_len.trailing_zeros() as usize; + n_layers_by_instance.push(n_vars); + } + } + + if !gkr_table_indices.is_empty() { + let batch_proof = match &multi_proof.batch_gkr_proof { + Some(p) => p, + None => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!("Tables have bus interactions but MultiProof missing batch_gkr_proof"); + return false; + } + }; + + // Verify batch GKR proof + match crate::gkr::gkr_verify_batch(batch_proof, &n_layers_by_instance, transcript) { + Ok((shared_random_point, per_instance_claims)) => { + // Distribute per-instance results to each table + // Initialize all tables with empty claims + for _ in 0..airs.len() { + gkr_bridge_claims.push(Vec::new()); + gkr_random_points.push(Vec::new()); } - }; - // Replay GKR verification on the main transcript - match crate::gkr::gkr_verify(&gkr_proof.gkr_proof, transcript) { - Ok((_random_point, n_claim, d_claim)) => { + for (i, &table_idx) in gkr_table_indices.iter().enumerate() { + let air = airs[table_idx]; + let proof = &multi_proof.proofs[table_idx]; + + let gkr_proof_data = match &proof.logup_gkr_proof { + Some(p) => p, + None => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!( + "Table {table_idx}: AIR has interactions but proof missing logup_gkr_proof" + ); + return false; + } + }; + + let (n_claim, d_claim) = &per_instance_claims[i]; + // Validate column_claims length matches AIR-expected column set let expected_cols = extract_column_indices(air.bus_interactions()); - if gkr_proof.column_claims.len() != expected_cols.len() { + if gkr_proof_data.column_claims.len() != expected_cols.len() { #[cfg(not(feature = "test_fiat_shamir"))] error!( - "Table {idx}: column_claims length mismatch: got {}, expected {}", - gkr_proof.column_claims.len(), + "Table {table_idx}: column_claims length mismatch: got {}, expected {}", + gkr_proof_data.column_claims.len(), expected_cols.len() ); return false; } - // Verify that column_claims are consistent with GKR output. - // A malicious prover could provide fake column_claims that don't - // match the (n_claim, d_claim) returned by GKR verification. + // Verify that column_claims are consistent with GKR output if !crate::lookup::reconstruct_and_verify_gkr_claims( - &n_claim, - &d_claim, - &gkr_proof.column_claims, + n_claim, + d_claim, + &gkr_proof_data.column_claims, air.bus_interactions(), &lookup_challenges, ) { #[cfg(not(feature = "test_fiat_shamir"))] error!( - "Table {idx}: GKR column claims verification failed — \ + "Table {table_idx}: GKR column claims verification failed — \ column_claims are inconsistent with GKR output (n_claim, d_claim)" ); return false; } - } - Err(_e) => { - #[cfg(not(feature = "test_fiat_shamir"))] - error!("Table {idx}: GKR verification failed: {:?}", _e); - return false; + + gkr_bridge_claims[table_idx] = gkr_proof_data.column_claims.clone(); + + // Compute the instance eval point from the shared random point + let n_vars = n_layers_by_instance[i]; + let instance_point = + crate::gkr::instance_eval_point(&shared_random_point, n_vars); + gkr_random_points[table_idx] = instance_point; } } - - gkr_bridge_claims.push(gkr_proof.column_claims.clone()); - gkr_random_points.push(gkr_proof.random_point.clone()); - } else { + Err(_e) => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!("Batch GKR verification failed: {:?}", _e); + return false; + } + } + } else { + // No tables with interactions - fill with empty claims + for _ in 0..airs.len() { gkr_bridge_claims.push(Vec::new()); gkr_random_points.push(Vec::new()); } @@ -933,9 +972,17 @@ pub trait IsStarkVerifier< if needs_lookup_challenges { let mut total = FieldElement::::zero(); - for (air, proof) in airs.iter().zip(&multi_proof.proofs) { - if air.has_trace_interaction() && let Some(ref gkr_proof) = proof.logup_gkr_proof { - total = total + &gkr_proof.gkr_proof.claimed_sum; + if let Some(ref batch_proof) = multi_proof.batch_gkr_proof { + // Sum the table contributions from the batch GKR root claims. + // Each root_claim is (numerator, denominator); the rational value n/d + // is the table's contribution. + for (root_n, root_d) in &batch_proof.root_claims { + let contrib = if *root_d == FieldElement::one() { + root_n.clone() + } else { + root_n * &root_d.inv().expect("GKR root denominator must be non-zero") + }; + total = total + &contrib; } } @@ -968,6 +1015,7 @@ pub trait IsStarkVerifier< { let multi_proof = MultiProof { proofs: vec![proof.clone()], + batch_gkr_proof: None, }; Self::multi_verify(&[air], &multi_proof, transcript, &FieldElement::zero()) } From 7ec0b566329e7a4a9c41f55aa60db5110231aa43 Mon Sep 17 00:00:00 2001 From: diegokingston Date: Thu, 9 Apr 2026 16:39:47 -0300 Subject: [PATCH 09/19] perf(gkr): split-value optimization for eq polynomial in sumcheck MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Implement the SVO (Split-Value Optimization, ePrint 2025/1117 Algorithm 5) for the eq polynomial in the GKR sumcheck inner loop. Instead of maintaining a full eq_table of size 2^l, split the eq polynomial into prefix and suffix halves: eq(w,b) = eq_prefix(w_prefix, b_prefix) * eq_suffix(w_suffix, b_suffix). Each half has size 2^{l/2}, reducing peak memory from O(2^l) to O(2 * 2^{l/2}) — a square-root reduction. During the first l/2 rounds (suffix rounds), eq_suffix is halved each round while eq_prefix stays fixed. The inner loop restructures as: h(t) = sum_s eq_suffix[s] * sum_p eq_prefix[p] * gate(t, p*S + s) After suffix rounds, eq_suffix is absorbed into eq_correction and eq_prefix becomes the eq_table for the remaining prefix rounds (standard approach). The optimization activates for parent_num_vars >= 8 (SVO_THRESHOLD). Smaller tables fall back to the original single-table approach to avoid setup overhead. Added two prove-verify roundtrip tests with 512 and 1024 leaves to exercise the SVO code path. All 186 tests pass. --- crypto/stark/src/gkr.rs | 675 ++++++++++++++++++++++++++++++---------- 1 file changed, 507 insertions(+), 168 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index b7879d78f..925838343 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -485,13 +485,11 @@ pub fn gkr_prove( // This function has degree 3 in each variable (product of eq * two child values), // so the round polynomial needs 4 evaluation points (degree 3). - // Build the five "bookkeeping" tables for the sumcheck: - // - eq_table: eq(current_point, b) for b in {0,1}^parent_num_vars + // Build the four gate "bookkeeping" tables for the sumcheck: // - nl_table: child_n[2b] (left numerators) // - nr_table: child_n[2b+1] (right numerators) // - dl_table: child_d[2b] (left denominators) // - dr_table: child_d[2b+1] (right denominators) - let mut eq_table = compute_eq_evals(¤t_point); let mut nl_table: Vec> = (0..parent_size).map(|j| child_n[2 * j].clone()).collect(); let mut nr_table: Vec> = (0..parent_size) @@ -503,19 +501,6 @@ pub fn gkr_prove( .map(|j| child_d[2 * j + 1].clone()) .collect(); - // Verify initial consistency: the sum of f over {0,1}^parent_num_vars - // should equal combined_claim - debug_assert!({ - let mut check_sum = FieldElement::::zero(); - for j in 0..parent_size { - let gate_val = &(&nl_table[j] * &dr_table[j]) - + &(&nr_table[j] * &dl_table[j]) - + &(&lambda * &(&dl_table[j] * &dr_table[j])); - check_sum = &check_sum + &(&eq_table[j] * &gate_val); - } - check_sum == combined_claim - }); - let mut round_polys = Vec::with_capacity(parent_num_vars); let mut challenges = Vec::with_capacity(parent_num_vars); let mut round_combined_claim = combined_claim.clone(); @@ -524,170 +509,452 @@ pub fn gkr_prove( // it (N/2 additions) and track the missing fold factors in this scalar. let mut eq_correction = FieldElement::::one(); - for r_round in current_point.iter().take(parent_num_vars) { - let half = nl_table.len() / 2; - - // Eq polynomial factoring (Dao-Thaler, ePrint 2024/1210): - // - // Factor eq(current_point, b) into: - // eq_correction (scalar from previous rounds' fold factors) * - // eq_round(r, t) (linear in t, constant across pairs) * - // eq_table[j] (per-pair scalar, independent of t) - // - // where r = current_point[round_idx] and - // eq_table[j] = eq_orig[2j] + eq_orig[2j+1] (pre-halved) - // - // The inner sum h_raw(t) = Σ_j eq_table[j] * gate(t, j) is degree 2, - // needing only 2 eval points (t=0, t=2) per pair instead of 3. - // Then h(t) = eq_correction * h_raw(t), and - // S(t) = eq_round(r, t) * h(t) recovers the degree-3 round poly. - let one = FieldElement::::one(); - - // Pre-halve eq_table: compute eq_rem[j] = eq_table[2j] + eq_table[2j+1] - // in-place. This replaces fold (which would multiply each entry by the - // challenge) with a simple sum. The fold factor is tracked in eq_correction. - for j in 0..half { - eq_table[j] = &eq_table[2 * j] + &eq_table[2 * j + 1]; - } - eq_table.truncate(half); - - let compute_pair_sums = |j: usize| -> [FieldElement; 2] { - // Pre-halved eq weight (scalar, independent of t) - let eq_rem = &eq_table[j]; - - let nl_l = &nl_table[2 * j]; - let nl_r = &nl_table[2 * j + 1]; - let nr_l = &nr_table[2 * j]; - let nr_r = &nr_table[2 * j + 1]; - let dl_l = &dl_table[2 * j]; - let dl_r = &dl_table[2 * j + 1]; - let dr_l = &dr_table[2 * j]; - let dr_r = &dr_table[2 * j + 1]; - - // t=0: interpolated values are just the left values - // gate = nl*dr + nr*dl + lambda*dl*dr - // = nl*dr + dl*(nr + lambda*dr) [3 muls instead of 4] - let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); - let h0 = eq_rem * &gate_0; - - // t=2: val = 2*right - left - let nl_2 = &(nl_r + nl_r) - nl_l; - let nr_2 = &(nr_r + nr_r) - nr_l; - let dl_2 = &(dl_r + dl_r) - dl_l; - let dr_2 = &(dr_r + dr_r) - dr_l; - let gate_2 = &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); - let h2 = eq_rem * &gate_2; - - [h0, h2] - }; + // SVO (Split-Value Optimization, ePrint 2025/1117 Algorithm 5): + // For large tables, split eq(w, x) into prefix and suffix halves to + // reduce memory from 2^l to 2 * 2^{l/2}. + // + // eq(w, b) = eq_prefix(w_prefix, b_prefix) * eq_suffix(w_suffix, b_suffix) + // + // During the first suffix_len rounds, eq_suffix is halved each round + // while eq_prefix stays fixed. The inner loop restructures as: + // h_raw(t) = Σ_s eq_suffix[s] * Σ_p eq_prefix[p] * gate(t, p*suffix_half + s) + // + // After suffix rounds, eq_suffix is absorbed into eq_correction, and + // eq_prefix becomes the eq_table for the remaining prefix rounds. + const SVO_THRESHOLD: usize = 8; + let use_svo = parent_num_vars >= SVO_THRESHOLD; + + if use_svo { + // --- SVO path: split eq into prefix + suffix --- + let suffix_len = parent_num_vars / 2; + let _prefix_len = parent_num_vars - suffix_len; + let mut eq_suffix = compute_eq_evals(¤t_point[..suffix_len]); + let eq_prefix = compute_eq_evals(¤t_point[suffix_len..parent_num_vars]); + let prefix_size = eq_prefix.len(); // = 2^prefix_len, constant during suffix rounds + + // Verify initial consistency using the split eq tables + debug_assert!({ + let mut check_sum = FieldElement::::zero(); + for j in 0..parent_size { + let suffix_idx = j & (eq_suffix.len() - 1); + let prefix_idx = j >> suffix_len; + let eq_val = &eq_suffix[suffix_idx] * &eq_prefix[prefix_idx]; + let gate_val = &(&nl_table[j] * &dr_table[j]) + + &(&nr_table[j] * &dl_table[j]) + + &(&lambda * &(&dl_table[j] * &dr_table[j])); + check_sum = &check_sum + &(&eq_val * &gate_val); + } + check_sum == combined_claim + }); + + // Phase 1: Suffix rounds (first suffix_len rounds) + // Process variables current_point[0..suffix_len]. + // eq_suffix is halved each round; eq_prefix stays fixed. + for round_idx in 0..suffix_len { + let r_round = ¤t_point[round_idx]; + let half = nl_table.len() / 2; + let suffix_half = eq_suffix.len() / 2; + + // Pre-halve eq_suffix (same as Dao-Thaler halving for eq_table) + for j in 0..suffix_half { + eq_suffix[j] = &eq_suffix[2 * j] + &eq_suffix[2 * j + 1]; + } + eq_suffix.truncate(suffix_half); - let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; - let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { - [&a[0] + &b[0], &a[1] + &b[1]] - }; + let one = FieldElement::::one(); - #[cfg(feature = "parallel")] - let totals: [FieldElement; 2] = if half >= 256 { - (0..half) - .into_par_iter() - .fold(zero2, |acc, j| add2(acc, compute_pair_sums(j))) - .reduce(zero2, add2) - } else { - (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))) - }; + // Inner loop: for each suffix index, accumulate gate contributions + // weighted by eq_prefix, then weight by eq_suffix. + // + // h_raw(t) = Σ_s eq_suffix[s] * Σ_p eq_prefix[p] * gate(t, p*suffix_half + s) + let compute_suffix_contribution = + |suffix_idx: usize| -> [FieldElement; 2] { + let eq_s = &eq_suffix[suffix_idx]; + let mut contrib_h0 = FieldElement::::zero(); + let mut contrib_h2 = FieldElement::::zero(); + + for prefix_idx in 0..prefix_size { + let j = prefix_idx * suffix_half + suffix_idx; + let eq_p = &eq_prefix[prefix_idx]; + + let nl_l = &nl_table[2 * j]; + let nl_r = &nl_table[2 * j + 1]; + let nr_l = &nr_table[2 * j]; + let nr_r = &nr_table[2 * j + 1]; + let dl_l = &dl_table[2 * j]; + let dl_r = &dl_table[2 * j + 1]; + let dr_l = &dr_table[2 * j]; + let dr_r = &dr_table[2 * j + 1]; + + // t=0: gate = nl*dr + dl*(nr + lambda*dr) + let gate_0 = + &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + contrib_h0 = &contrib_h0 + &(eq_p * &gate_0); + + // t=2: val = 2*right - left + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = + &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + contrib_h2 = &contrib_h2 + &(eq_p * &gate_2); + } + + [eq_s * &contrib_h0, eq_s * &contrib_h2] + }; + + let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; + let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { + [&a[0] + &b[0], &a[1] + &b[1]] + }; - #[cfg(not(feature = "parallel"))] - let totals: [FieldElement; 2] = - (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))); - - // Phase 2: Recover S(t) from h(t) and eq_round(r, t). - // - // The inner sum h_raw(t) doesn't include the eq_correction factor - // (accumulated from previous rounds' fold factors). Apply it now: - // h(t) = eq_correction * h_raw(t) - // - // eq_round(r, t) = (1-r)(1-t) + r*t - // eq_round(r, 0) = 1-r, eq_round(r, 1) = r, - // eq_round(r, 2) = 3r-1, eq_round(r, 3) = 5r-2 - // - // S(0) = (1-r)*h(0), S(1) = round_combined_claim - S(0) - // h(1) = S(1)/r, h(3) = 3*h(2) - 3*h(1) + h(0) (degree-2 extrapolation) - // S(2) = (3r-1)*h(2), S(3) = (5r-2)*h(3) - let [raw_h0, raw_h2] = totals; - let total_h0 = &eq_correction * &raw_h0; - let total_h2 = &eq_correction * &raw_h2; - - let one_minus_r = &one - r_round; - let s0 = &one_minus_r * &total_h0; - let s1 = &round_combined_claim - &s0; - - let r_inv = r_round - .inv() - .expect("r_round = 0 is probability 2^{-64} for random challenges"); - let h1 = &s1 * &r_inv; - - let three = FieldElement::::from(3u64); - let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; - - let eq_at_2 = &(&three * r_round) - &one; - let s2 = &eq_at_2 * &total_h2; - - let eq_at_3 = - &(&FieldElement::::from(5u64) * r_round) - &FieldElement::::from(2u64); - let s3 = &eq_at_3 * &h3; - - let poly_evals = vec![s0, s1, s2, s3]; - - let round_poly = RoundPoly::new(poly_evals); - - // Append round polynomial evaluations to transcript - for eval in round_poly.evals() { - transcript.append_field_element(eval); + #[cfg(feature = "parallel")] + let totals: [FieldElement; 2] = if suffix_half >= 256 { + (0..suffix_half) + .into_par_iter() + .fold(zero2, |acc, s| add2(acc, compute_suffix_contribution(s))) + .reduce(zero2, add2) + } else { + (0..suffix_half) + .fold(zero2(), |acc, s| add2(acc, compute_suffix_contribution(s))) + }; + + #[cfg(not(feature = "parallel"))] + let totals: [FieldElement; 2] = (0..suffix_half) + .fold(zero2(), |acc, s| add2(acc, compute_suffix_contribution(s))); + + // Phase 2: Recover S(t) from h(t) and eq_round(r, t). + let [raw_h0, raw_h2] = totals; + let total_h0 = &eq_correction * &raw_h0; + let total_h2 = &eq_correction * &raw_h2; + + let one_minus_r = &one - r_round; + let s0 = &one_minus_r * &total_h0; + let s1 = &round_combined_claim - &s0; + + let r_inv = r_round + .inv() + .expect("r_round = 0 is probability 2^{-64} for random challenges"); + let h1 = &s1 * &r_inv; + + let three = FieldElement::::from(3u64); + let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; + + let eq_at_2 = &(&three * r_round) - &one; + let s2 = &eq_at_2 * &total_h2; + + let eq_at_3 = &(&FieldElement::::from(5u64) * r_round) + - &FieldElement::::from(2u64); + let s3 = &eq_at_3 * &h3; + + let poly_evals = vec![s0, s1, s2, s3]; + let round_poly = RoundPoly::new(poly_evals); + + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + let challenge: FieldElement = transcript.sample_field_element(); + round_combined_claim = round_poly.evaluate(&challenge); + + let eq_update = + &(r_round * &challenge) + &(&one_minus_r * &(&one - &challenge)); + eq_correction = &eq_correction * &eq_update; + + // Fold the four gate tables + let fold_table = |table: &mut Vec>| { + #[cfg(feature = "parallel")] + if half >= 256 { + let folded: Vec> = table + .par_chunks(2) + .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) + .collect(); + *table = folded; + return; + } + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + }; + + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); + + round_polys.push(round_poly); + challenges.push(challenge); } - // Sample challenge for this round - let challenge: FieldElement = transcript.sample_field_element(); + // Transition: absorb the remaining eq_suffix scalar into eq_correction. + // After suffix_len halvings, eq_suffix has been reduced to a single entry. + debug_assert_eq!(eq_suffix.len(), 1); + eq_correction = &eq_correction * &eq_suffix[0]; + + // Phase 2: Prefix rounds (remaining prefix_len rounds) + // Use eq_prefix as the eq_table for standard Dao-Thaler halving. + let mut eq_table = eq_prefix; + + for round_idx in suffix_len..parent_num_vars { + let r_round = ¤t_point[round_idx]; + let half = nl_table.len() / 2; + let one = FieldElement::::one(); + + // Pre-halve eq_table + for j in 0..half { + eq_table[j] = &eq_table[2 * j] + &eq_table[2 * j + 1]; + } + eq_table.truncate(half); + + let compute_pair_sums = |j: usize| -> [FieldElement; 2] { + let eq_rem = &eq_table[j]; + + let nl_l = &nl_table[2 * j]; + let nl_r = &nl_table[2 * j + 1]; + let nr_l = &nr_table[2 * j]; + let nr_r = &nr_table[2 * j + 1]; + let dl_l = &dl_table[2 * j]; + let dl_r = &dl_table[2 * j + 1]; + let dr_l = &dr_table[2 * j]; + let dr_r = &dr_table[2 * j + 1]; + + let gate_0 = + &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let h0 = eq_rem * &gate_0; + + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = + &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + let h2 = eq_rem * &gate_2; + + [h0, h2] + }; + + let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; + let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { + [&a[0] + &b[0], &a[1] + &b[1]] + }; - // Update round_combined_claim for next round - round_combined_claim = round_poly.evaluate(&challenge); - - // Update eq_correction: accumulate eq(r_round, challenge). - // eq(r, c) = r*c + (1-r)*(1-c), the fold factor we skip for eq_table. - let eq_update = &(r_round * &challenge) + &(&one_minus_r * &(&one - &challenge)); - eq_correction = &eq_correction * &eq_update; - - // Fold the four gate tables (eq_table was already halved before inner loop) - // table[j] = table[2j] + challenge * (table[2j+1] - table[2j]) - // - // When half >= 256 we use par_chunks(2) so that all Rayon threads - // participate in folding a single table (instead of just 4-way - // parallelism across the four independent tables). - let fold_table = |table: &mut Vec>| { #[cfg(feature = "parallel")] - if half >= 256 { - let folded: Vec> = table - .par_chunks(2) - .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) - .collect(); - *table = folded; - return; + let totals: [FieldElement; 2] = if half >= 256 { + (0..half) + .into_par_iter() + .fold(zero2, |acc, j| add2(acc, compute_pair_sums(j))) + .reduce(zero2, add2) + } else { + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))) + }; + + #[cfg(not(feature = "parallel"))] + let totals: [FieldElement; 2] = + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))); + + // Phase 2: Recover S(t) from h(t) and eq_round(r, t). + let [raw_h0, raw_h2] = totals; + let total_h0 = &eq_correction * &raw_h0; + let total_h2 = &eq_correction * &raw_h2; + + let one_minus_r = &one - r_round; + let s0 = &one_minus_r * &total_h0; + let s1 = &round_combined_claim - &s0; + + let r_inv = r_round + .inv() + .expect("r_round = 0 is probability 2^{-64} for random challenges"); + let h1 = &s1 * &r_inv; + + let three = FieldElement::::from(3u64); + let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; + + let eq_at_2 = &(&three * r_round) - &one; + let s2 = &eq_at_2 * &total_h2; + + let eq_at_3 = &(&FieldElement::::from(5u64) * r_round) + - &FieldElement::::from(2u64); + let s3 = &eq_at_3 * &h3; + + let poly_evals = vec![s0, s1, s2, s3]; + let round_poly = RoundPoly::new(poly_evals); + + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + let challenge: FieldElement = transcript.sample_field_element(); + round_combined_claim = round_poly.evaluate(&challenge); + + let eq_update = + &(r_round * &challenge) + &(&one_minus_r * &(&one - &challenge)); + eq_correction = &eq_correction * &eq_update; + + let fold_table = |table: &mut Vec>| { + #[cfg(feature = "parallel")] + if half >= 256 { + let folded: Vec> = table + .par_chunks(2) + .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) + .collect(); + *table = folded; + return; + } + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + }; + + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); + + round_polys.push(round_poly); + challenges.push(challenge); + } + } else { + // --- Standard path for small tables (parent_num_vars < SVO_THRESHOLD) --- + let mut eq_table = compute_eq_evals(¤t_point); + + // Verify initial consistency + debug_assert!({ + let mut check_sum = FieldElement::::zero(); + for j in 0..parent_size { + let gate_val = &(&nl_table[j] * &dr_table[j]) + + &(&nr_table[j] * &dl_table[j]) + + &(&lambda * &(&dl_table[j] * &dr_table[j])); + check_sum = &check_sum + &(&eq_table[j] * &gate_val); } - // Sequential fallback (also used when half < 256) + check_sum == combined_claim + }); + + for r_round in current_point.iter().take(parent_num_vars) { + let half = nl_table.len() / 2; + let one = FieldElement::::one(); + + // Pre-halve eq_table for j in 0..half { - let left = &table[2 * j]; - let right = &table[2 * j + 1]; - table[j] = left + &(&challenge * &(right - left)); + eq_table[j] = &eq_table[2 * j] + &eq_table[2 * j + 1]; } - table.truncate(half); - }; + eq_table.truncate(half); + + let compute_pair_sums = |j: usize| -> [FieldElement; 2] { + let eq_rem = &eq_table[j]; + + let nl_l = &nl_table[2 * j]; + let nl_r = &nl_table[2 * j + 1]; + let nr_l = &nr_table[2 * j]; + let nr_r = &nr_table[2 * j + 1]; + let dl_l = &dl_table[2 * j]; + let dl_r = &dl_table[2 * j + 1]; + let dr_l = &dr_table[2 * j]; + let dr_r = &dr_table[2 * j + 1]; + + let gate_0 = + &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let h0 = eq_rem * &gate_0; + + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = + &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + let h2 = eq_rem * &gate_2; + + [h0, h2] + }; - fold_table(&mut nl_table); - fold_table(&mut nr_table); - fold_table(&mut dl_table); - fold_table(&mut dr_table); + let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; + let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { + [&a[0] + &b[0], &a[1] + &b[1]] + }; - round_polys.push(round_poly); - challenges.push(challenge); + #[cfg(feature = "parallel")] + let totals: [FieldElement; 2] = if half >= 256 { + (0..half) + .into_par_iter() + .fold(zero2, |acc, j| add2(acc, compute_pair_sums(j))) + .reduce(zero2, add2) + } else { + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))) + }; + + #[cfg(not(feature = "parallel"))] + let totals: [FieldElement; 2] = + (0..half).fold(zero2(), |acc, j| add2(acc, compute_pair_sums(j))); + + let [raw_h0, raw_h2] = totals; + let total_h0 = &eq_correction * &raw_h0; + let total_h2 = &eq_correction * &raw_h2; + + let one_minus_r = &one - r_round; + let s0 = &one_minus_r * &total_h0; + let s1 = &round_combined_claim - &s0; + + let r_inv = r_round + .inv() + .expect("r_round = 0 is probability 2^{-64} for random challenges"); + let h1 = &s1 * &r_inv; + + let three = FieldElement::::from(3u64); + let h3 = &(&(&three * &total_h2) - &(&three * &h1)) + &total_h0; + + let eq_at_2 = &(&three * r_round) - &one; + let s2 = &eq_at_2 * &total_h2; + + let eq_at_3 = &(&FieldElement::::from(5u64) * r_round) + - &FieldElement::::from(2u64); + let s3 = &eq_at_3 * &h3; + + let poly_evals = vec![s0, s1, s2, s3]; + let round_poly = RoundPoly::new(poly_evals); + + for eval in round_poly.evals() { + transcript.append_field_element(eval); + } + + let challenge: FieldElement = transcript.sample_field_element(); + round_combined_claim = round_poly.evaluate(&challenge); + + let eq_update = + &(r_round * &challenge) + &(&one_minus_r * &(&one - &challenge)); + eq_correction = &eq_correction * &eq_update; + + let fold_table = |table: &mut Vec>| { + #[cfg(feature = "parallel")] + if half >= 256 { + let folded: Vec> = table + .par_chunks(2) + .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) + .collect(); + *table = folded; + return; + } + for j in 0..half { + let left = &table[2 * j]; + let right = &table[2 * j + 1]; + table[j] = left + &(&challenge * &(right - left)); + } + table.truncate(half); + }; + + fold_table(&mut nl_table); + fold_table(&mut nr_table); + fold_table(&mut dl_table); + fold_table(&mut dr_table); + + round_polys.push(round_poly); + challenges.push(challenge); + } } // After all rounds, each table has a single entry: the MLE evaluated @@ -2485,6 +2752,78 @@ mod tests { assert_eq!(verifier_d, expected_d); } + // ==================== SVO (Split-Value Optimization) tests ==================== + + #[test] + fn test_gkr_prove_verify_roundtrip_512_svo() { + // 512 leaves: exercises the SVO path (parent_num_vars = 8 >= SVO_THRESHOLD) + // Tree: 10 layers (sizes 512, 256, 128, 64, 32, 16, 8, 4, 2, 1) + let n = 512; + let nums: Vec = (1..=n).map(|i| FE::from(i as u64)).collect(); + let dens: Vec = (1..=n).map(|i| FE::from(i as u64 + 1000)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + // Prove + let mut prover_transcript = DefaultTranscript::::new(&[0x5F, 0x00]); + let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + + // Verify with a fresh transcript (same seed) + let mut verifier_transcript = DefaultTranscript::::new(&[0x5F, 0x00]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + result.is_ok(), + "GKR verification should succeed for 512 leaves (SVO path): {:?}", + result.err() + ); + let (verifier_point, verifier_n, verifier_d) = result.unwrap(); + + // The verifier's final point must match the prover's + assert_eq!(verifier_point, prover_point); + + // The verifier's leaf claims must match the prover's + assert_eq!(verifier_n, prover_n); + assert_eq!(verifier_d, prover_d); + + // Verify consistency with leaf MLEs + let expected_n = evaluate_mle(&nums, &verifier_point); + let expected_d = evaluate_mle(&dens, &verifier_point); + assert_eq!(verifier_n, expected_n); + assert_eq!(verifier_d, expected_d); + } + + #[test] + fn test_gkr_prove_verify_roundtrip_1024_svo() { + // 1024 leaves: ensures SVO path is exercised at multiple layers + // parent_num_vars = 9 at the bottom layer + let n = 1024; + let nums: Vec = (1..=n).map(|i| FE::from(i as u64 * 3 + 1)).collect(); + let dens: Vec = (1..=n).map(|i| FE::from(i as u64 * 7 + 5)).collect(); + let tree = build_summation_tree(nums.clone(), dens.clone()); + + let mut prover_transcript = DefaultTranscript::::new(&[0xBB]); + let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + + let mut verifier_transcript = DefaultTranscript::::new(&[0xBB]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + result.is_ok(), + "GKR verification should succeed for 1024 leaves (SVO path): {:?}", + result.err() + ); + let (verifier_point, verifier_n, verifier_d) = result.unwrap(); + + assert_eq!(verifier_point, prover_point); + assert_eq!(verifier_n, prover_n); + assert_eq!(verifier_d, prover_d); + + let expected_n = evaluate_mle(&nums, &verifier_point); + let expected_d = evaluate_mle(&dens, &verifier_point); + assert_eq!(verifier_n, expected_n); + assert_eq!(verifier_d, expected_d); + } + // ==================== Batch GKR tests ==================== /// Helper: create a Layer::LogUpGeneric leaf from random-ish fractions. From dd6a2147f28f892fd51c8bdcf016890df8975a10 Mon Sep 17 00:00:00 2001 From: jotabulacios Date: Thu, 9 Apr 2026 20:21:07 -0300 Subject: [PATCH 10/19] Eliminate O(N) combined_claims recompute in batch GKR by evaluating per-instance round polynomials at the challenge point, and fix lint warnings --- crypto/stark/src/fri/fri_functions.rs | 2 - crypto/stark/src/gkr.rs | 208 +++++++++--------- crypto/stark/src/lagrange_kernel.rs | 11 +- crypto/stark/src/lookup.rs | 86 ++++---- crypto/stark/src/proof/stark.rs | 7 +- crypto/stark/src/prover.rs | 56 ++--- crypto/stark/src/sumcheck.rs | 22 +- .../src/tests/bus_tests/soundness_tests.rs | 40 ++-- crypto/stark/src/verifier.rs | 17 +- 9 files changed, 228 insertions(+), 221 deletions(-) diff --git a/crypto/stark/src/fri/fri_functions.rs b/crypto/stark/src/fri/fri_functions.rs index e96cbf5d2..1512e905a 100644 --- a/crypto/stark/src/fri/fri_functions.rs +++ b/crypto/stark/src/fri/fri_functions.rs @@ -17,8 +17,6 @@ pub fn fold_evaluations_in_place, E: IsField>( zeta: &FieldElement, inv_twiddles: &[FieldElement], ) { - let half = evals.len() / 2; - #[cfg(feature = "parallel")] { use rayon::prelude::*; diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 925838343..e3849eb44 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -550,6 +550,7 @@ pub fn gkr_prove( // Phase 1: Suffix rounds (first suffix_len rounds) // Process variables current_point[0..suffix_len]. // eq_suffix is halved each round; eq_prefix stays fixed. + #[allow(clippy::needless_range_loop)] for round_idx in 0..suffix_len { let r_round = ¤t_point[round_idx]; let half = nl_table.len() / 2; @@ -567,42 +568,41 @@ pub fn gkr_prove( // weighted by eq_prefix, then weight by eq_suffix. // // h_raw(t) = Σ_s eq_suffix[s] * Σ_p eq_prefix[p] * gate(t, p*suffix_half + s) - let compute_suffix_contribution = - |suffix_idx: usize| -> [FieldElement; 2] { - let eq_s = &eq_suffix[suffix_idx]; - let mut contrib_h0 = FieldElement::::zero(); - let mut contrib_h2 = FieldElement::::zero(); - - for prefix_idx in 0..prefix_size { - let j = prefix_idx * suffix_half + suffix_idx; - let eq_p = &eq_prefix[prefix_idx]; - - let nl_l = &nl_table[2 * j]; - let nl_r = &nl_table[2 * j + 1]; - let nr_l = &nr_table[2 * j]; - let nr_r = &nr_table[2 * j + 1]; - let dl_l = &dl_table[2 * j]; - let dl_r = &dl_table[2 * j + 1]; - let dr_l = &dr_table[2 * j]; - let dr_r = &dr_table[2 * j + 1]; - - // t=0: gate = nl*dr + dl*(nr + lambda*dr) - let gate_0 = - &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); - contrib_h0 = &contrib_h0 + &(eq_p * &gate_0); - - // t=2: val = 2*right - left - let nl_2 = &(nl_r + nl_r) - nl_l; - let nr_2 = &(nr_r + nr_r) - nr_l; - let dl_2 = &(dl_r + dl_r) - dl_l; - let dr_2 = &(dr_r + dr_r) - dr_l; - let gate_2 = - &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); - contrib_h2 = &contrib_h2 + &(eq_p * &gate_2); - } - - [eq_s * &contrib_h0, eq_s * &contrib_h2] - }; + let compute_suffix_contribution = |suffix_idx: usize| -> [FieldElement; 2] { + let eq_s = &eq_suffix[suffix_idx]; + let mut contrib_h0 = FieldElement::::zero(); + let mut contrib_h2 = FieldElement::::zero(); + + #[allow(clippy::needless_range_loop)] + for prefix_idx in 0..prefix_size { + let j = prefix_idx * suffix_half + suffix_idx; + let eq_p = &eq_prefix[prefix_idx]; + + let nl_l = &nl_table[2 * j]; + let nl_r = &nl_table[2 * j + 1]; + let nr_l = &nr_table[2 * j]; + let nr_r = &nr_table[2 * j + 1]; + let dl_l = &dl_table[2 * j]; + let dl_r = &dl_table[2 * j + 1]; + let dr_l = &dr_table[2 * j]; + let dr_r = &dr_table[2 * j + 1]; + + // t=0: gate = nl*dr + dl*(nr + lambda*dr) + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + contrib_h0 = &contrib_h0 + &(eq_p * &gate_0); + + // t=2: val = 2*right - left + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = + &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + contrib_h2 = &contrib_h2 + &(eq_p * &gate_2); + } + + [eq_s * &contrib_h0, eq_s * &contrib_h2] + }; let zero2 = || [FieldElement::::zero(), FieldElement::::zero()]; let add2 = |a: [FieldElement; 2], b: [FieldElement; 2]| { @@ -699,8 +699,7 @@ pub fn gkr_prove( // Use eq_prefix as the eq_table for standard Dao-Thaler halving. let mut eq_table = eq_prefix; - for round_idx in suffix_len..parent_num_vars { - let r_round = ¤t_point[round_idx]; + for r_round in current_point.iter().take(parent_num_vars).skip(suffix_len) { let half = nl_table.len() / 2; let one = FieldElement::::one(); @@ -722,16 +721,14 @@ pub fn gkr_prove( let dr_l = &dr_table[2 * j]; let dr_r = &dr_table[2 * j + 1]; - let gate_0 = - &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); let h0 = eq_rem * &gate_0; let nl_2 = &(nl_r + nl_r) - nl_l; let nr_2 = &(nr_r + nr_r) - nr_l; let dl_2 = &(dl_r + dl_r) - dl_l; let dr_2 = &(dr_r + dr_r) - dr_l; - let gate_2 = - &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + let gate_2 = &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); let h2 = eq_rem * &gate_2; [h0, h2] @@ -858,16 +855,14 @@ pub fn gkr_prove( let dr_l = &dr_table[2 * j]; let dr_r = &dr_table[2 * j + 1]; - let gate_0 = - &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); let h0 = eq_rem * &gate_0; let nl_2 = &(nl_r + nl_r) - nl_l; let nr_2 = &(nr_r + nr_r) - nr_l; let dl_2 = &(dl_r + dl_r) - dl_l; let dr_2 = &(dr_r + dr_r) - dr_l; - let gate_2 = - &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + let gate_2 = &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); let h2 = eq_rem * &gate_2; [h0, h2] @@ -1473,6 +1468,18 @@ pub fn gkr_prove_batch( sum }; + // Per-instance round polynomial evals, used to update combined_claims + // via O(1) interpolation instead of O(N) table scan. + let mut per_instance_evals: Vec<[FieldElement; 4]> = vec![ + [ + FieldElement::zero(), + FieldElement::zero(), + FieldElement::zero(), + FieldElement::zero() + ]; + per_instance_tables.len() + ]; + for round_idx in 0..max_parent_vars { // For each active instance, compute the round poly contribution let mut batch_s0 = FieldElement::::zero(); @@ -1489,6 +1496,12 @@ pub fn gkr_prove_batch( let half_claim = &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); // S(0) = S(1) = half_claim, S(2) = S(3) = half_claim (constant) + per_instance_evals[idx] = [ + half_claim.clone(), + half_claim.clone(), + half_claim.clone(), + half_claim.clone(), + ]; batch_s0 = &batch_s0 + &(&alpha_pow * &half_claim); batch_s1 = &batch_s1 + &(&alpha_pow * &half_claim); batch_s2 = &batch_s2 + &(&alpha_pow * &half_claim); @@ -1504,6 +1517,12 @@ pub fn gkr_prove_batch( // Already reduced to constant let half_claim = &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + per_instance_evals[idx] = [ + half_claim.clone(), + half_claim.clone(), + half_claim.clone(), + half_claim.clone(), + ]; batch_s0 = &batch_s0 + &(&alpha_pow * &half_claim); batch_s1 = &batch_s1 + &(&alpha_pow * &half_claim); batch_s2 = &batch_s2 + &(&alpha_pow * &half_claim); @@ -1602,6 +1621,7 @@ pub fn gkr_prove_batch( - &FieldElement::::from(2u64); let s3 = &eq_at_3 * &h3; + per_instance_evals[idx] = [s0.clone(), s1.clone(), s2.clone(), s3.clone()]; batch_s0 = &batch_s0 + &(&alpha_pow * &s0); batch_s1 = &batch_s1 + &(&alpha_pow * &s1); batch_s2 = &batch_s2 + &(&alpha_pow * &s2); @@ -1620,27 +1640,24 @@ pub fn gkr_prove_batch( let challenge: FieldElement = transcript.sample_field_element(); _round_combined_claim = round_poly.evaluate(&challenge); - // Update per-instance: fold tables and update eq_correction + // Update per-instance: fold tables, update eq_correction, and + // update combined_claims via O(1) polynomial evaluation (not O(N) table scan). for (idx, tables) in per_instance_tables.iter_mut().enumerate() { + // Update combined_claims from saved per-instance round poly evals. + // S_i(challenge) via degree-3 Lagrange interpolation at {0,1,2,3}. + let [ref si0, ref si1, ref si2, ref si3] = per_instance_evals[idx]; + let instance_poly = + RoundPoly::new(vec![si0.clone(), si1.clone(), si2.clone(), si3.clone()]); + combined_claims[idx] = instance_poly.evaluate(&challenge); + let n_unused = max_parent_vars - tables.parent_num_vars; - if round_idx < n_unused { - // Instance hasn't started yet: its per-instance polynomial is - // constant = claim/2 at all evaluation points. - // p_i(challenge) = claim_i / 2. - combined_claims[idx] = - &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); + if round_idx < n_unused || tables.dl_table.len() / 2 == 0 { continue; } let instance_round = round_idx - n_unused; let half = tables.dl_table.len() / 2; - if half == 0 { - combined_claims[idx] = - &combined_claims[idx] * &FieldElement::::from(2u64).inv().unwrap(); - continue; - } - // Update eq_correction using instance-specific eval point let r_round = &tables.instance_point[instance_round]; let one = FieldElement::::one(); @@ -1650,6 +1667,15 @@ pub fn gkr_prove_batch( // Fold gate tables let fold_table = |table: &mut Vec>| { + #[cfg(feature = "parallel")] + if half >= 256 { + let folded: Vec> = table + .par_chunks(2) + .map(|pair| &pair[0] + &(&challenge * &(&pair[1] - &pair[0]))) + .collect(); + *table = folded; + return; + } for j in 0..half { let left = &table[2 * j]; let right = &table[2 * j + 1]; @@ -1664,39 +1690,6 @@ pub fn gkr_prove_batch( } fold_table(&mut tables.dl_table); fold_table(&mut tables.dr_table); - - // Update per-instance combined claim - // The per-instance round poly evaluated at challenge: - // We need to recompute. Actually, the combined_claim for the batch - // already accounts for all instances. But we need per-instance claims - // for the next round. Let's just evaluate the per-instance sum at the - // challenge point using the folded tables. - // After folding, the tables have `half` elements. The per-instance claim - // for next round = sum over j of eq_table[j] * gate(j) * eq_correction - // But that's expensive. Instead, use the identity: - // new_claim = S_i(challenge) where S_i is this instance's round poly. - // We already computed S_i(0) and S_i(1) = claim_i - S_i(0). - // For non-trivial instances, we can use the round poly evaluation directly. - // Let's recompute per-instance sum from the folded tables. - let mut new_claim = FieldElement::::zero(); - let new_half = tables.dl_table.len(); - if tables.is_singles { - for j in 0..new_half { - let eq_rem = &tables.eq_table[j]; - let gate = &(&tables.dl_table[j] + &tables.dr_table[j]) - + &(&lambda * &(&tables.dl_table[j] * &tables.dr_table[j])); - new_claim = &new_claim + &(eq_rem * &gate); - } - } else { - for j in 0..new_half { - let eq_rem = &tables.eq_table[j]; - let gate = &(&tables.nl_table[j] * &tables.dr_table[j]) - + &(&tables.dl_table[j] - * &(&tables.nr_table[j] + &(&lambda * &tables.dr_table[j]))); - new_claim = &new_claim + &(eq_rem * &gate); - } - } - combined_claims[idx] = &tables.eq_correction * &new_claim; } round_polys.push(round_poly); @@ -2847,9 +2840,8 @@ mod tests { #[test] fn test_batch_gkr_same_size_instances() { // 3 instances, all with 4 leaves (n_vars=2, 3 layers each) - let instances: Vec>> = (0..3) - .map(|_| gen_layers(make_generic_leaf(2))) - .collect(); + let instances: Vec>> = + (0..3).map(|_| gen_layers(make_generic_leaf(2))).collect(); let mut prover_transcript = DefaultTranscript::::new(&[]); let (proof, shared_point, final_claims) = @@ -2895,7 +2887,11 @@ mod tests { let mut verifier_transcript = DefaultTranscript::::new(&[]); let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); - assert!(result.is_ok(), "mixed-size batch verify failed: {:?}", result.err()); + assert!( + result.is_ok(), + "mixed-size batch verify failed: {:?}", + result.err() + ); let (v_point, v_claims) = result.unwrap(); assert_eq!(v_point, shared_point); @@ -2906,10 +2902,10 @@ mod tests { fn test_batch_gkr_mixed_size_with_singles() { // Mix of Generic and Singles leaves with different sizes let instances: Vec>> = vec![ - gen_layers(make_singles_leaf(2)), // 3 layers - gen_layers(make_generic_leaf(3)), // 4 layers - gen_layers(make_singles_leaf(5)), // 6 layers - gen_layers(make_generic_leaf(4)), // 5 layers + gen_layers(make_singles_leaf(2)), // 3 layers + gen_layers(make_generic_leaf(3)), // 4 layers + gen_layers(make_singles_leaf(5)), // 6 layers + gen_layers(make_generic_leaf(4)), // 5 layers ]; let n_layers: Vec = instances.iter().map(|l| l.len() - 1).collect(); @@ -2920,7 +2916,11 @@ mod tests { let mut verifier_transcript = DefaultTranscript::::new(&[]); let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); - assert!(result.is_ok(), "mixed singles/generic batch verify failed: {:?}", result.err()); + assert!( + result.is_ok(), + "mixed singles/generic batch verify failed: {:?}", + result.err() + ); let (v_point, v_claims) = result.unwrap(); assert_eq!(v_point, shared_point); @@ -2946,7 +2946,11 @@ mod tests { let mut verifier_transcript = DefaultTranscript::::new(&[]); let result = gkr_verify_batch(&proof, &n_layers, &mut verifier_transcript); - assert!(result.is_ok(), "many mixed instances batch verify failed: {:?}", result.err()); + assert!( + result.is_ok(), + "many mixed instances batch verify failed: {:?}", + result.err() + ); let (v_point, v_claims) = result.unwrap(); assert_eq!(v_point, shared_point); diff --git a/crypto/stark/src/lagrange_kernel.rs b/crypto/stark/src/lagrange_kernel.rs index 1a5da3816..59d5f9d5e 100644 --- a/crypto/stark/src/lagrange_kernel.rs +++ b/crypto/stark/src/lagrange_kernel.rs @@ -25,6 +25,7 @@ pub fn compute_lagrange_kernel(r: &[FieldElement]) -> Vec(r: &[FieldElement]) -> Vec( - values: &[FieldElement], - r: &[FieldElement], -) -> FieldElement { +pub fn eval_mle(values: &[FieldElement], r: &[FieldElement]) -> FieldElement { let n = r.len(); let big_n = 1usize << n; assert_eq!( @@ -106,10 +104,7 @@ pub fn eval_mle( /// Uses `F * E -> E` multiplication directly (no `to_extension()` conversion). /// /// Panics if `values.len()` is not a power of 2 or if `r.len() != log2(values.len())`. -pub fn eval_mle_base( - values: &[FieldElement], - r: &[FieldElement], -) -> FieldElement +pub fn eval_mle_base(values: &[FieldElement], r: &[FieldElement]) -> FieldElement where F: IsField + IsSubFieldOf, E: IsField, diff --git a/crypto/stark/src/lookup.rs b/crypto/stark/src/lookup.rs index 850d5ce9c..62f5057e4 100644 --- a/crypto/stark/src/lookup.rs +++ b/crypto/stark/src/lookup.rs @@ -127,7 +127,6 @@ pub fn logup_random_point_start(interactions: &[BusInteraction]) -> usize { LOGUP_GAMMA_POWERS_START + extract_column_indices(interactions).len() } - // ============================================================================= // Bus Types // ============================================================================= @@ -865,8 +864,7 @@ impl< // Add bridge running sum constraint for LogUp-GKR tables let mut all_constraints = transition_constraints; if num_interactions > 0 { - let column_indices = - extract_column_indices(&auxiliary_trace_build_data.interactions); + let column_indices = extract_column_indices(&auxiliary_trace_build_data.interactions); let bridge_constraint = LookupBridgeSumConstraint { constraint_idx: all_constraints.len(), column_indices, @@ -1043,8 +1041,7 @@ where let n = rap_challenges.len() - rp_start; let mut l0_expected = FieldElement::::one(); for j in 0..n { - l0_expected *= - FieldElement::::one() - &rap_challenges[rp_start + j]; + l0_expected *= FieldElement::::one() - &rap_challenges[rp_start + j]; } // Aux column 0 is the Lagrange kernel; constrain l[0] = prod(1 - r_j) boundary_constraints.push(BoundaryConstraint::new_aux(0, 0, l0_expected)); @@ -1434,7 +1431,6 @@ where .collect() } - // ============================================================================= // LogUp-GKR Bridge Running Sum // ============================================================================= @@ -1497,7 +1493,7 @@ pub(crate) fn extract_column_indices(interactions: &[BusInteraction]) -> Vec( interaction_a: &BusInteraction, interaction_b: &BusInteraction, @@ -1796,7 +1792,6 @@ where } } - // ============================================================================= // LogUp-GKR Leaf Fraction Computation // ============================================================================= @@ -1882,6 +1877,7 @@ where z - &linear_combination } +#[allow(dead_code)] fn compute_multiplicity_from_step, B: IsField>( step: &TableView, multiplicity: &Multiplicity, @@ -1936,6 +1932,7 @@ fn compute_multiplicity_from_step, B: IsField>( /// Computes the fingerprint for an interaction from a `TableView`. /// /// Returns `z - (bus_id*α^0 + v[0]*α^1 + v[1]*α^2 + ...)` +#[allow(dead_code)] fn compute_fingerprint_from_step, B: IsField>( step: &TableView, interaction: &BusInteraction, @@ -1965,6 +1962,7 @@ fn compute_fingerprint_from_step, B: IsField>( /// Clearing denominators: `c * fp_a * fp_b - sign_a * m_a * fp_b - sign_b * m_b * fp_a = 0` /// /// Degree 3: c (aux) × fp_a (linear in main) × fp_b (linear in main). +#[allow(dead_code)] struct LookupBatchedTermConstraint { interaction_a: BusInteraction, interaction_b: BusInteraction, @@ -1972,6 +1970,7 @@ struct LookupBatchedTermConstraint { constraint_idx: usize, } +#[allow(dead_code)] impl LookupBatchedTermConstraint { pub fn new( interaction_a: BusInteraction, @@ -2090,6 +2089,7 @@ where /// /// For 2 absorbed interactions: /// `(acc_next - acc_curr - Σ terms + L/N) · f₁·f₂ - sign₁·m₁·f₂ - sign₂·m₂·f₁ = 0` (degree 3) +#[allow(dead_code)] struct LookupAccumulatedConstraint { constraint_idx: usize, /// Number of committed term columns (excludes absorbed interactions) @@ -2100,6 +2100,7 @@ struct LookupAccumulatedConstraint { absorbed: Vec, } +#[allow(dead_code)] impl LookupAccumulatedConstraint { pub fn new( constraint_idx: usize, @@ -2327,7 +2328,8 @@ where let mut running_d = FieldElement::::one(); for (k, inter) in interactions.iter().enumerate() { - let fp_k = compute_fingerprint_at_row(inter, main_segment_cols, row, z, &alpha_powers); + let fp_k = + compute_fingerprint_at_row(inter, main_segment_cols, row, z, &alpha_powers); let m_k = compute_multiplicity_at_row(inter, main_segment_cols, row); // F * E -> E multiplication (no to_extension) @@ -2349,7 +2351,7 @@ where // LogUp-GKR Integration // ============================================================================= -use crate::gkr::{build_summation_tree, gen_layers, gkr_prove, GkrProof, Layer}; +use crate::gkr::{GkrProof, Layer, build_summation_tree, gen_layers, gkr_prove}; use crate::lagrange_kernel::{compute_lagrange_kernel, eval_mle_base_with_kernel}; use crypto::fiat_shamir::is_transcript::IsTranscript; @@ -2454,7 +2456,10 @@ pub fn reconstruct_and_verify_gkr_claims( } } Multiplicity::Sum3(a, b, c) => { - if !claim_map.contains_key(a) || !claim_map.contains_key(b) || !claim_map.contains_key(c) { + if !claim_map.contains_key(a) + || !claim_map.contains_key(b) + || !claim_map.contains_key(c) + { return false; } } @@ -2917,6 +2922,7 @@ where } #[cfg(test)] +#[allow(clippy::clone_on_copy, clippy::cloned_ref_to_slice_refs)] mod tests { use super::*; use math::field::goldilocks::GoldilocksField; @@ -2943,11 +2949,8 @@ mod tests { let main_segment_cols = vec![col0.clone()]; // Single sender interaction: bus_id=1, Multiplicity::One, one Direct column - let interaction = BusInteraction::sender( - 1u64, - Multiplicity::One, - Packing::Direct.columns(&[0]), - ); + let interaction = + BusInteraction::sender(1u64, Multiplicity::One, Packing::Direct.columns(&[0])); let interactions = vec![interaction]; // Challenges: z=100, alpha=3 @@ -2955,8 +2958,12 @@ mod tests { let alpha = FE::from(3u64); let challenges = vec![z.clone(), alpha.clone()]; - let (numerators, denominators) = - compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + let (numerators, denominators) = compute_logup_leaf_fractions::( + &interactions, + &main_segment_cols, + trace_len, + &challenges, + ); assert_eq!(numerators.len(), trace_len); assert_eq!(denominators.len(), trace_len); @@ -3009,19 +3016,20 @@ mod tests { let main_segment_cols = vec![col0.clone(), col1.clone()]; // Receiver interaction: bus_id=0, Multiplicity from column 1 - let interaction = BusInteraction::receiver( - 0u64, - Multiplicity::Column(1), - Packing::Direct.columns(&[0]), - ); + let interaction = + BusInteraction::receiver(0u64, Multiplicity::Column(1), Packing::Direct.columns(&[0])); let interactions = vec![interaction]; let z = FE::from(50u64); let alpha = FE::from(7u64); let challenges = vec![z.clone(), alpha.clone()]; - let (numerators, denominators) = - compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + let (numerators, denominators) = compute_logup_leaf_fractions::( + &interactions, + &main_segment_cols, + trace_len, + &challenges, + ); let alpha_powers = compute_alpha_powers(&alpha, 2); @@ -3062,25 +3070,22 @@ mod tests { let main_segment_cols = vec![col0.clone(), col1.clone()]; // Interaction 0: sender, bus_id=0, Multiplicity::One, column 0 - let inter0 = BusInteraction::sender( - 0u64, - Multiplicity::One, - Packing::Direct.columns(&[0]), - ); + let inter0 = BusInteraction::sender(0u64, Multiplicity::One, Packing::Direct.columns(&[0])); // Interaction 1: receiver, bus_id=1, Multiplicity::One, column 1 - let inter1 = BusInteraction::receiver( - 1u64, - Multiplicity::One, - Packing::Direct.columns(&[1]), - ); + let inter1 = + BusInteraction::receiver(1u64, Multiplicity::One, Packing::Direct.columns(&[1])); let interactions = vec![inter0, inter1]; let z = FE::from(200u64); let alpha = FE::from(5u64); let challenges = vec![z.clone(), alpha.clone()]; - let (numerators, denominators) = - compute_logup_leaf_fractions::(&interactions, &main_segment_cols, trace_len, &challenges); + let (numerators, denominators) = compute_logup_leaf_fractions::( + &interactions, + &main_segment_cols, + trace_len, + &challenges, + ); let alpha_powers = compute_alpha_powers(&alpha, 2); @@ -3127,11 +3132,8 @@ mod tests { let col0: Vec = (1..=8).map(|v| FE::from(v as u64)).collect(); let main_segment_cols = vec![col0]; - let interaction = BusInteraction::sender( - 2u64, - Multiplicity::One, - Packing::Direct.columns(&[0]), - ); + let interaction = + BusInteraction::sender(2u64, Multiplicity::One, Packing::Direct.columns(&[0])); let z = FE::from(1000u64); let alpha = FE::from(11u64); diff --git a/crypto/stark/src/proof/stark.rs b/crypto/stark/src/proof/stark.rs index 107f97290..8e440e232 100644 --- a/crypto/stark/src/proof/stark.rs +++ b/crypto/stark/src/proof/stark.rs @@ -5,8 +5,11 @@ use math::field::{ }; use crate::{ - config::Commitment, fri::fri_decommit::FriDecommitment, gkr::{BatchGkrProof, GkrProof}, - lookup::BusPublicInputs, table::Table, + config::Commitment, + fri::fri_decommit::FriDecommitment, + gkr::{BatchGkrProof, GkrProof}, + lookup::BusPublicInputs, + table::Table, }; #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index fc6f34b20..b386b1182 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -37,9 +37,7 @@ use super::constraints::evaluator::ConstraintEvaluator; use super::domain::Domain; use super::fri::fri_decommit::FriDecommitment; use super::grinding; -use super::lookup::{ - BusPublicInputs, LogUpGkrResult, extend_rap_challenges_with_bridge, -}; +use super::lookup::{BusPublicInputs, LogUpGkrResult, extend_rap_challenges_with_bridge}; use super::proof::stark::{DeepPolynomialOpening, LogUpGkrProof, MultiProof, StarkProof}; use super::trace::TraceTable; use super::traits::AIR; @@ -671,7 +669,13 @@ pub trait IsStarkProver< .zip(domains.iter().zip(twiddle_caches.iter())) { let result = Self::reconstruct_round1( - *air, *trace, domain, metadata, &**twiddles, main_pool, aux_pool, + *air, + *trace, + domain, + metadata, + &**twiddles, + main_pool, + aux_pool, ) .expect("reconstruct_round1 failed in debug-checks"); temp_results.push(result); @@ -1696,8 +1700,7 @@ pub trait IsStarkProver< // The instance eval point for this table let n_vars = n_layers_by_instance[i]; - let instance_point = - crate::gkr::instance_eval_point(&shared_random_point, n_vars); + let instance_point = crate::gkr::instance_eval_point(&shared_random_point, n_vars); let air = air_trace_pairs[table_idx].0; let trace = &air_trace_pairs[table_idx].1; @@ -1754,9 +1757,7 @@ pub trait IsStarkProver< #[cfg(feature = "instruments")] let phase_start = Instant::now(); - for ((air, trace, _), gkr_result) in - air_trace_pairs.iter_mut().zip(gkr_results.iter()) - { + for ((air, trace, _), gkr_result) in air_trace_pairs.iter_mut().zip(gkr_results.iter()) { if air.has_trace_interaction() { if let Some(result) = gkr_result { let kernel = &result.lagrange_kernel; @@ -1779,12 +1780,11 @@ pub trait IsStarkProver< // So σ[0] = 0 (start), σ[i+1] = σ[i] + l[i]·batched[i] - Δ. // The circular wrap-around at row N-1 requires σ[0] = σ[N-1] + l[N-1]·batched[N-1] - Δ, // which telescopes to: 0 = Σ l[i]·batched[i] - N·Δ = target - target. - let (bridge_offset, gamma_powers) = - crate::lookup::compute_bridge_params( - &result.column_claims, - &gamma, - trace_len, - ); + let (bridge_offset, gamma_powers) = crate::lookup::compute_bridge_params( + &result.column_claims, + &gamma, + trace_len, + ); let main_cols = trace.columns_main(); // Pre-compute batched values in parallel: batched[i] = Σ_j main_cols[col_j][i] * γ^j @@ -1807,6 +1807,7 @@ pub trait IsStarkProver< // Set σ[0] = 0, then build forward: σ[i+1] = σ[i] + l[i]*batched[i] - Δ trace.set_aux(0, 1, FieldElement::::zero()); let mut sigma = FieldElement::::zero(); + #[allow(clippy::needless_range_loop)] for row in 0..trace_len - 1 { let l_val = trace.get_aux(row, 0); sigma = sigma + l_val * &batched_values[row] - &bridge_offset; @@ -1901,10 +1902,8 @@ pub trait IsStarkProver< let mut metadatas: Vec> = Vec::with_capacity(num_airs); let mut gkr_results_iter = gkr_results.into_iter(); - for ((main_commit, (aux_tree, aux_root)), idx) in main_commits - .into_iter() - .zip(aux_results) - .zip(0..num_airs) + for ((main_commit, (aux_tree, aux_root)), idx) in + main_commits.into_iter().zip(aux_results).zip(0..num_airs) { let gkr_result = gkr_results_iter.next().unwrap(); @@ -2025,14 +2024,17 @@ pub trait IsStarkProver< // Attach LogUp-GKR proof metadata from Phase B'. // In batch mode, gkr_proof is None (the real proof is in MultiProof::batch_gkr_proof). // We still store random_point and column_claims for per-table verification. - proof.logup_gkr_proof = metadata.logup_gkr_result.as_ref().map(|r| LogUpGkrProof { - gkr_proof: r.gkr_proof.clone().unwrap_or_else(|| crate::gkr::GkrProof { - claimed_sum: r.table_contribution.clone(), - layer_proofs: vec![], - }), - random_point: r.random_point.clone(), - column_claims: r.column_claims.clone(), - }); + proof.logup_gkr_proof = + metadata.logup_gkr_result.as_ref().map(|r| LogUpGkrProof { + gkr_proof: r.gkr_proof.clone().unwrap_or_else(|| { + crate::gkr::GkrProof { + claimed_sum: r.table_contribution.clone(), + layer_proofs: vec![], + } + }), + random_point: r.random_point.clone(), + column_claims: r.column_claims.clone(), + }); // Collect per-table sub-op timing via TLS. // Both the store (inside prove_rounds_2_to_4) and this take run on the diff --git a/crypto/stark/src/sumcheck.rs b/crypto/stark/src/sumcheck.rs index 1f36783ed..83eb0a386 100644 --- a/crypto/stark/src/sumcheck.rs +++ b/crypto/stark/src/sumcheck.rs @@ -24,7 +24,10 @@ impl RoundPoly { /// `evals[i]` is the polynomial evaluated at x = i. /// The polynomial has degree `evals.len() - 1`. pub fn new(evals: Vec>) -> Self { - assert!(!evals.is_empty(), "RoundPoly must have at least one evaluation"); + assert!( + !evals.is_empty(), + "RoundPoly must have at least one evaluation" + ); Self { evals } } @@ -91,8 +94,7 @@ impl RoundPoly { // Append point_minus_j values; batch-invert everything in one call to_invert.extend(point_minus_j.iter().cloned()); - FieldElement::inplace_batch_inverse(&mut to_invert) - .expect("All values are nonzero"); + FieldElement::inplace_batch_inverse(&mut to_invert).expect("All values are nonzero"); let (w_inv, pm_inv) = to_invert.split_at(d + 1); @@ -404,7 +406,12 @@ mod tests { fn test_evaluate_at_random_point_cubic() { // p(x) = x^3 + 2x^2 + 3x + 4 // Need 4 evaluation points: 0, 1, 2, 3 - let coeffs = [FE::from(4u64), FE::from(3u64), FE::from(2u64), FE::from(1u64)]; + let coeffs = [ + FE::from(4u64), + FE::from(3u64), + FE::from(2u64), + FE::from(1u64), + ]; let evals: Vec = (0..4) .map(|i| eval_poly_coeffs(&coeffs, &FE::from(i as u64))) .collect(); @@ -493,12 +500,7 @@ mod tests { for x in [0u64, 1, 2, 3, 4, 5, 10, 100, 12345] { let point = FE::from(x); let expected = eval_poly_coeffs(&coeffs, &point); - assert_eq!( - poly.evaluate(&point), - expected, - "Mismatch at x = {}", - x - ); + assert_eq!(poly.evaluate(&point), expected, "Mismatch at x = {}", x); } } diff --git a/crypto/stark/src/tests/bus_tests/soundness_tests.rs b/crypto/stark/src/tests/bus_tests/soundness_tests.rs index 8162909cb..dbd844f7c 100644 --- a/crypto/stark/src/tests/bus_tests/soundness_tests.rs +++ b/crypto/stark/src/tests/bus_tests/soundness_tests.rs @@ -2,6 +2,7 @@ //! //! These tests verify that the verifier correctly rejects proofs that violate //! the bus balance invariant. +#![allow(clippy::clone_on_copy, clippy::type_complexity)] use crypto::fiat_shamir::default_transcript::DefaultTranscript; use math::field::element::FieldElement; @@ -774,14 +775,15 @@ fn test_tampered_gkr_column_claims_rejected() { !gkr_proof.column_claims.is_empty(), "CPU must have column claims" ); - gkr_proof.column_claims[0].1 = - gkr_proof.column_claims[0].1.clone() + FieldElement::one(); + gkr_proof.column_claims[0].1 = gkr_proof.column_claims[0].1.clone() + FieldElement::one(); } else { panic!("CPU table must have a logup_gkr_proof"); } - let air_refs: Vec<&dyn AIR> = - airs.iter().map(|a| a as &dyn AIR).collect(); + let air_refs: Vec<&dyn AIR> = airs + .iter() + .map(|a| a as &dyn AIR) + .collect(); assert!( !Verifier::multi_verify( @@ -814,14 +816,15 @@ fn test_tampered_gkr_claimed_sum_rejected() { "batch proof must have multiple root_claims" ); // Tamper the ADD table's root claim (instance index 1, numerator) - batch_proof.root_claims[1].0 = - batch_proof.root_claims[1].0.clone() + FieldElement::one(); + batch_proof.root_claims[1].0 = batch_proof.root_claims[1].0.clone() + FieldElement::one(); } else { panic!("MultiProof must have a batch_gkr_proof"); } - let air_refs: Vec<&dyn AIR> = - airs.iter().map(|a| a as &dyn AIR).collect(); + let air_refs: Vec<&dyn AIR> = airs + .iter() + .map(|a| a as &dyn AIR) + .collect(); assert!( !Verifier::multi_verify( @@ -847,8 +850,10 @@ fn test_missing_gkr_proof_rejected() { // ADD has bus interactions, so the verifier will reject. multi_proof.proofs[1].logup_gkr_proof = None; - let air_refs: Vec<&dyn AIR> = - airs.iter().map(|a| a as &dyn AIR).collect(); + let air_refs: Vec<&dyn AIR> = airs + .iter() + .map(|a| a as &dyn AIR) + .collect(); assert!( !Verifier::multi_verify( @@ -879,13 +884,16 @@ fn test_tampered_sigma_ood_rejected() { let num_main_cpu = 5usize; let sigma_col_ood_idx = num_main_cpu + 1; // aux column 1 = sigma let cpu_proof = &mut multi_proof.proofs[0]; - let corrupted = *cpu_proof.trace_ood_evaluations.get(0, sigma_col_ood_idx) + FieldElement::one(); + let corrupted = + *cpu_proof.trace_ood_evaluations.get(0, sigma_col_ood_idx) + FieldElement::one(); cpu_proof .trace_ood_evaluations .set(0, sigma_col_ood_idx, corrupted); - let air_refs: Vec<&dyn AIR> = - airs.iter().map(|a| a as &dyn AIR).collect(); + let air_refs: Vec<&dyn AIR> = airs + .iter() + .map(|a| a as &dyn AIR) + .collect(); assert!( !Verifier::multi_verify( @@ -931,8 +939,10 @@ fn test_tampered_lagrange_kernel_random_point_rejected() { panic!("MultiProof must have a batch_gkr_proof"); } - let air_refs: Vec<&dyn AIR> = - airs.iter().map(|a| a as &dyn AIR).collect(); + let air_refs: Vec<&dyn AIR> = airs + .iter() + .map(|a| a as &dyn AIR) + .collect(); assert!( !Verifier::multi_verify( diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index 119934973..9e586092c 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -944,12 +944,8 @@ pub trait IsStarkVerifier< } // Rounds 2-4: verify (bridge is now a transition constraint) - if !Self::verify_rounds_2_to_4( - *air, - proof, - &mut table_transcript, - table_rap_challenges, - ) { + if !Self::verify_rounds_2_to_4(*air, proof, &mut table_transcript, table_rap_challenges) + { #[cfg(not(feature = "test_fiat_shamir"))] error!( "Table {} failed verify_rounds_2_to_4 (num_constraints={}, trace_cols={})", @@ -1179,13 +1175,8 @@ pub trait IsStarkVerifier< #[cfg(feature = "instruments")] let timer1 = Instant::now(); - let challenges = Self::replay_rounds_after_round_1( - air, - proof, - &domain, - transcript, - rap_challenges, - ); + let challenges = + Self::replay_rounds_after_round_1(air, proof, &domain, transcript, rap_challenges); // verify grinding let security_bits = air.context().proof_options.grinding_factor; From a9aa266de98cb011006c2fc3c3cf476e6cbf5376 Mon Sep 17 00:00:00 2001 From: jotabulacios Date: Thu, 9 Apr 2026 20:39:29 -0300 Subject: [PATCH 11/19] save work --- crypto/stark/src/prover.rs | 73 +++++++++++++++++++++++++++----------- 1 file changed, 52 insertions(+), 21 deletions(-) diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index b386b1182..6895e1cb4 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -1661,23 +1661,47 @@ pub trait IsStarkProver< // can replay them. // Step 1: Compute GKR layer trees for each table with bus interactions. - let mut leaf_layers_per_table: Vec>> = Vec::new(); + // Identify which tables participate, then compute layers in parallel. let mut gkr_table_indices: Vec = Vec::new(); - for (idx, (air, trace, _)) in air_trace_pairs.iter().enumerate() { + for (idx, (air, _, _)) in air_trace_pairs.iter().enumerate() { if air.has_trace_interaction() { + gkr_table_indices.push(idx); + } + } + + #[cfg(feature = "parallel")] + let leaf_layers_per_table: Vec>> = gkr_table_indices + .par_iter() + .map(|&idx| { + let (air, trace, _) = &air_trace_pairs[idx]; let interactions = air.bus_interactions(); let main_segment_cols = trace.columns_main(); let trace_len = trace.num_rows(); - let layers = crate::lookup::compute_logup_layers( + crate::lookup::compute_logup_layers( interactions, &main_segment_cols, trace_len, &lookup_challenges, - ); - leaf_layers_per_table.push(layers); - gkr_table_indices.push(idx); - } - } + ) + }) + .collect(); + + #[cfg(not(feature = "parallel"))] + let leaf_layers_per_table: Vec>> = gkr_table_indices + .iter() + .map(|&idx| { + let (air, trace, _) = &air_trace_pairs[idx]; + let interactions = air.bus_interactions(); + let main_segment_cols = trace.columns_main(); + let trace_len = trace.num_rows(); + crate::lookup::compute_logup_layers( + interactions, + &main_segment_cols, + trace_len, + &lookup_challenges, + ) + }) + .collect(); // Step 2: Run batch GKR (single proof for all tables). let batch_gkr_proof = if !leaf_layers_per_table.is_empty() { @@ -1688,17 +1712,13 @@ pub trait IsStarkProver< let (batch_proof, shared_random_point, per_instance_claims) = crate::gkr::gkr_prove_batch(leaf_layers_per_table, transcript); - // Step 3: Distribute batch results to per-table LogUpGkrResult. - let mut gkr_results_vec: Vec>> = - vec![None; num_airs]; - for (i, &table_idx) in gkr_table_indices.iter().enumerate() { + // Step 3: Distribute batch results to per-table LogUpGkrResult (parallel). + // Each table independently computes its Lagrange kernel and column MLE claims. + let finalize_one = |i: usize| -> (usize, LogUpGkrResult) { + let table_idx = gkr_table_indices[i]; let (n_claim, d_claim) = per_instance_claims[i].clone(); - let table_contribution = batch_proof.root_claims[i].clone(); - // Compute table_contribution as n/d (the rational root value). - // The root_claims store (numerator, denominator) of the summation tree root. - let (root_n, root_d) = table_contribution; + let (root_n, root_d) = batch_proof.root_claims[i].clone(); - // The instance eval point for this table let n_vars = n_layers_by_instance[i]; let instance_point = crate::gkr::instance_eval_point(&shared_random_point, n_vars); @@ -1707,10 +1727,6 @@ pub trait IsStarkProver< let interactions = air.bus_interactions(); let main_segment_cols = trace.columns_main(); - // Compute the claimed_sum (table contribution) as a single field element. - // For the bus balance check, we need n/d as an element. - // The summation tree root stores n/d where the claimed_sum = n/d. - // We compute it by inverting d and multiplying by n. let table_contrib = if root_d == FieldElement::one() { root_n } else { @@ -1725,6 +1741,21 @@ pub trait IsStarkProver< d_claim, table_contrib, ); + (table_idx, result) + }; + + #[cfg(feature = "parallel")] + let finalized: Vec<_> = (0..gkr_table_indices.len()) + .into_par_iter() + .map(finalize_one) + .collect(); + + #[cfg(not(feature = "parallel"))] + let finalized: Vec<_> = (0..gkr_table_indices.len()).map(finalize_one).collect(); + + let mut gkr_results_vec: Vec>> = + vec![None; num_airs]; + for (table_idx, result) in finalized { gkr_results_vec[table_idx] = Some(result); } From fb65d0f24b4d53478c2af6e76b853e5a100df2c6 Mon Sep 17 00:00:00 2001 From: jotabulacios Date: Thu, 9 Apr 2026 20:55:57 -0300 Subject: [PATCH 12/19] Port split-value eq optimization (SVO) to batch GKR sumcheck inner loop --- crypto/stark/src/gkr.rs | 174 +++++++++++++++++++++++++++------------- 1 file changed, 119 insertions(+), 55 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index e3849eb44..7c3852736 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -5,6 +5,10 @@ use math::field::{element::FieldElement, traits::IsField}; #[cfg(feature = "parallel")] use rayon::prelude::*; +/// Minimum parent_num_vars for enabling Split-Value Optimization (SVO). +/// Below this threshold, the standard flat eq_table approach is used. +const SVO_THRESHOLD: usize = 8; + // ============================================================================= // Layer enum for gate-specialized GKR // ============================================================================= @@ -521,7 +525,6 @@ pub fn gkr_prove( // // After suffix rounds, eq_suffix is absorbed into eq_correction, and // eq_prefix becomes the eq_table for the remaining prefix rounds. - const SVO_THRESHOLD: usize = 8; let use_svo = parent_num_vars >= SVO_THRESHOLD; if use_svo { @@ -1439,7 +1442,17 @@ pub fn gkr_prove_batch( my_parent_num_vars ); let inst_point = instance_eval_point(¤t_point, my_parent_num_vars); - let eq_table = compute_eq_evals(&inst_point); + let use_svo = my_parent_num_vars >= SVO_THRESHOLD; + let svo_suffix_len = if use_svo { my_parent_num_vars / 2 } else { 0 }; + + let (eq_table, eq_prefix, eq_suffix) = if use_svo { + let suffix = compute_eq_evals(&inst_point[..svo_suffix_len]); + let prefix = + compute_eq_evals(&inst_point[svo_suffix_len..my_parent_num_vars]); + (Vec::new(), prefix, suffix) + } else { + (compute_eq_evals(&inst_point), Vec::new(), Vec::new()) + }; PerInstanceTables { nl_table, @@ -1451,6 +1464,10 @@ pub fn gkr_prove_batch( is_singles, parent_num_vars: my_parent_num_vars, instance_point: inst_point, + use_svo, + svo_suffix_len, + eq_prefix, + eq_suffix, } }) .collect(); @@ -1532,72 +1549,113 @@ pub fn gkr_prove_batch( } // Eq polynomial factoring (same as single-instance prover). - // Use the instance-specific eval point (not the shared current_point) - // so that r_round matches the eq_table built from instance_eval_point. let r_round = tables.instance_point[instance_round].clone(); + let one = FieldElement::::one(); - // Pre-halve eq_table - for j in 0..half { - tables.eq_table[j] = &tables.eq_table[2 * j] + &tables.eq_table[2 * j + 1]; - } - tables.eq_table.truncate(half); + // Helper: compute gate h(0) and h(2) for a generic pair at index j. + let gate_generic = |tables: &PerInstanceTables, + j: usize| + -> [FieldElement; 2] { + let nl_l = &tables.nl_table[2 * j]; + let nl_r = &tables.nl_table[2 * j + 1]; + let nr_l = &tables.nr_table[2 * j]; + let nr_r = &tables.nr_table[2 * j + 1]; + let dl_l = &tables.dl_table[2 * j]; + let dl_r = &tables.dl_table[2 * j + 1]; + let dr_l = &tables.dr_table[2 * j]; + let dr_r = &tables.dr_table[2 * j + 1]; - let one = FieldElement::::one(); + let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); + let nl_2 = &(nl_r + nl_r) - nl_l; + let nr_2 = &(nr_r + nr_r) - nr_l; + let dl_2 = &(dl_r + dl_r) - dl_l; + let dr_2 = &(dr_r + dr_r) - dr_l; + let gate_2 = &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); + [gate_0, gate_2] + }; - // Compute h(0) and h(2) (inner sum without eq_round factor) - let (raw_h0, raw_h2) = if tables.is_singles { - // Singles gate: gate(t) = dl(t) + dr(t) + lambda * dl(t) * dr(t) - // (numerators are all 1, so the fraction sum is 1/dl + 1/dr) - let mut h0 = FieldElement::::zero(); - let mut h2 = FieldElement::::zero(); - for j in 0..half { - let eq_rem = &tables.eq_table[j]; + let gate_singles = + |tables: &PerInstanceTables, j: usize| -> [FieldElement; 2] { let dl_l = &tables.dl_table[2 * j]; let dl_r = &tables.dl_table[2 * j + 1]; let dr_l = &tables.dr_table[2 * j]; let dr_r = &tables.dr_table[2 * j + 1]; - // t=0 let gate_0 = &(dl_l + dr_l) + &(&lambda * &(dl_l * dr_l)); - h0 = &h0 + &(eq_rem * &gate_0); - - // t=2 let dl_2 = &(dl_r + dl_r) - dl_l; let dr_2 = &(dr_r + dr_r) - dr_l; let gate_2 = &(&dl_2 + &dr_2) + &(&lambda * &(&dl_2 * &dr_2)); - h2 = &h2 + &(eq_rem * &gate_2); - } - (h0, h2) - } else { - // Generic gate: nl*dr + dl*(nr + lambda*dr) - let mut h0 = FieldElement::::zero(); - let mut h2 = FieldElement::::zero(); - for j in 0..half { - let eq_rem = &tables.eq_table[j]; - let nl_l = &tables.nl_table[2 * j]; - let nl_r = &tables.nl_table[2 * j + 1]; - let nr_l = &tables.nr_table[2 * j]; - let nr_r = &tables.nr_table[2 * j + 1]; - let dl_l = &tables.dl_table[2 * j]; - let dl_r = &tables.dl_table[2 * j + 1]; - let dr_l = &tables.dr_table[2 * j]; - let dr_r = &tables.dr_table[2 * j + 1]; - - // t=0 - let gate_0 = &(nl_l * dr_l) + &(dl_l * &(nr_l + &(&lambda * dr_l))); - h0 = &h0 + &(eq_rem * &gate_0); - - // t=2 - let nl_2 = &(nl_r + nl_r) - nl_l; - let nr_2 = &(nr_r + nr_r) - nr_l; - let dl_2 = &(dl_r + dl_r) - dl_l; - let dr_2 = &(dr_r + dr_r) - dr_l; - let gate_2 = - &(&nl_2 * &dr_2) + &(&dl_2 * &(&nr_2 + &(&lambda * &dr_2))); - h2 = &h2 + &(eq_rem * &gate_2); - } - (h0, h2) - }; + [gate_0, gate_2] + }; + + // Compute h(0) and h(2) using SVO or standard path. + let (raw_h0, raw_h2) = + if tables.use_svo && instance_round < tables.svo_suffix_len { + // SVO suffix round: nested eq_suffix × (eq_prefix × gate) loop + let suffix_half = tables.eq_suffix.len() / 2; + let prefix_size = tables.eq_prefix.len(); + + // Pre-halve eq_suffix + for j in 0..suffix_half { + tables.eq_suffix[j] = + &tables.eq_suffix[2 * j] + &tables.eq_suffix[2 * j + 1]; + } + tables.eq_suffix.truncate(suffix_half); + + let mut h0 = FieldElement::::zero(); + let mut h2 = FieldElement::::zero(); + for suffix_idx in 0..suffix_half { + let eq_s = &tables.eq_suffix[suffix_idx]; + let mut ch0 = FieldElement::::zero(); + let mut ch2 = FieldElement::::zero(); + #[allow(clippy::needless_range_loop)] + for prefix_idx in 0..prefix_size { + let j = prefix_idx * suffix_half + suffix_idx; + let eq_p = &tables.eq_prefix[prefix_idx]; + let [g0, g2] = if tables.is_singles { + gate_singles(tables, j) + } else { + gate_generic(tables, j) + }; + ch0 = &ch0 + &(eq_p * &g0); + ch2 = &ch2 + &(eq_p * &g2); + } + h0 = &h0 + &(eq_s * &ch0); + h2 = &h2 + &(eq_s * &ch2); + } + (h0, h2) + } else { + // Standard path (non-SVO, or SVO prefix rounds after suffix is exhausted) + if tables.use_svo && instance_round == tables.svo_suffix_len { + // Transition: absorb remaining eq_suffix into eq_correction + // and switch to eq_prefix as eq_table. + debug_assert_eq!(tables.eq_suffix.len(), 1); + tables.eq_correction = &tables.eq_correction * &tables.eq_suffix[0]; + tables.eq_table = std::mem::take(&mut tables.eq_prefix); + tables.use_svo = false; + } + + // Pre-halve eq_table + for j in 0..half { + tables.eq_table[j] = + &tables.eq_table[2 * j] + &tables.eq_table[2 * j + 1]; + } + tables.eq_table.truncate(half); + + let mut h0 = FieldElement::::zero(); + let mut h2 = FieldElement::::zero(); + for j in 0..half { + let eq_rem = &tables.eq_table[j]; + let [g0, g2] = if tables.is_singles { + gate_singles(tables, j) + } else { + gate_generic(tables, j) + }; + h0 = &h0 + &(eq_rem * &g0); + h2 = &h2 + &(eq_rem * &g2); + } + (h0, h2) + }; // Apply eq_correction let total_h0 = &tables.eq_correction * &raw_h0; @@ -1807,6 +1865,12 @@ struct PerInstanceTables { /// The instance-specific evaluation point derived from the shared current_point. /// Used for r_round lookups in the Dao-Thaler eq factoring. instance_point: Vec>, + // SVO (Split-Value Optimization) fields. + // When use_svo is true, eq is split into prefix × suffix for sqrt memory. + use_svo: bool, + svo_suffix_len: usize, + eq_prefix: Vec>, + eq_suffix: Vec>, } // ============================================================================= From 95deda491edad85b8c820e884a603ec24ab36796 Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 10:22:02 -0300 Subject: [PATCH 13/19] Fix three verifier panics in batch GKR --- crypto/stark/src/fri/fri_functions.rs | 1 + 1 file changed, 1 insertion(+) diff --git a/crypto/stark/src/fri/fri_functions.rs b/crypto/stark/src/fri/fri_functions.rs index 1512e905a..b02090533 100644 --- a/crypto/stark/src/fri/fri_functions.rs +++ b/crypto/stark/src/fri/fri_functions.rs @@ -38,6 +38,7 @@ pub fn fold_evaluations_in_place, E: IsField>( #[cfg(not(feature = "parallel"))] { + let half = evals.len() / 2; for j in 0..half { let lo = &evals[2 * j]; let hi = &evals[2 * j + 1]; From 54dfa9751c32aaa8e1d22394a59d513cd6d2350c Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 11:32:55 -0300 Subject: [PATCH 14/19] Prevent verifier DoS via unchecked layer_proofs length and child_claims_by_instance bounds --- crypto/stark/src/gkr.rs | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 7c3852736..779534590 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -1913,6 +1913,16 @@ pub fn gkr_verify_batch( let max_layers = *n_layers_by_instance.iter().max().unwrap(); + if proof.layer_proofs.len() != max_layers { + return Err(GkrError::InvalidTree { + reason: format!( + "expected {} layer proofs but got {}", + max_layers, + proof.layer_proofs.len(), + ), + }); + } + // Track per-instance state let mut n_claims: Vec>> = vec![None; n_instances]; let mut d_claims: Vec>> = vec![None; n_instances]; @@ -1963,6 +1973,17 @@ pub fn gkr_verify_batch( let round_polys = &layer_proof.sumcheck_proof.round_polys; + if layer_proof.child_claims_by_instance.len() < active_instances.len() { + return Err(GkrError::InvalidTree { + reason: format!( + "layer {}: child_claims_by_instance has {} entries but {} active instances", + layer_idx, + layer_proof.child_claims_by_instance.len(), + active_instances.len(), + ), + }); + } + if round_polys.is_empty() { // Trivial layer: no sumcheck needed // Append child claims and sample eta (same as prover) From d04f642aef697109f57d67e86cb0df033fbc05c5 Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 11:33:35 -0300 Subject: [PATCH 15/19] Fix Lagrange kernel soundness, 0-layer bus forgery, and column_claims Fiat-Shamir gap --- crypto/stark/src/lookup.rs | 86 ++++++++++++++++++++++++++---------- crypto/stark/src/prover.rs | 37 +++++++++++++--- crypto/stark/src/verifier.rs | 20 +++++++-- 3 files changed, 110 insertions(+), 33 deletions(-) diff --git a/crypto/stark/src/lookup.rs b/crypto/stark/src/lookup.rs index 62f5057e4..0de9d1de4 100644 --- a/crypto/stark/src/lookup.rs +++ b/crypto/stark/src/lookup.rs @@ -116,15 +116,17 @@ pub const LOGUP_CHALLENGE_GAMMA: usize = 2; pub const LOGUP_BRIDGE_OFFSET_IDX: usize = 3; /// Start index of precomputed gamma powers in the per-table rap_challenges vector. -/// rap_challenges[LOGUP_GAMMA_POWERS_START + j] = γ^j for j = 0, 1, ..., K-1. +/// rap_challenges[LOGUP_GAMMA_POWERS_START + j] = γ^j for j = 0, 1, ..., K. +/// K+1 powers total: γ^0..γ^{K-1} for column inner products, γ^K for the l² self-check. pub const LOGUP_GAMMA_POWERS_START: usize = 4; /// Start index of GKR random_point coordinates in rap_challenges. -/// After gamma_powers[0..K], we append random_point[0..n]. -/// The actual index is LOGUP_GAMMA_POWERS_START + K where K = number of distinct column indices. +/// After gamma_powers[0..K+1] (K+1 powers, one extra for the l² self-check), +/// we append random_point[0..n]. /// Use `logup_random_point_start(interactions)` to compute the concrete index. pub fn logup_random_point_start(interactions: &[BusInteraction]) -> usize { - LOGUP_GAMMA_POWERS_START + extract_column_indices(interactions).len() + // +1 for the extra γ^K power used by the Lagrange kernel l² self-check (BUG-004 fix) + LOGUP_GAMMA_POWERS_START + extract_column_indices(interactions).len() + 1 } // ============================================================================= @@ -1036,7 +1038,8 @@ where // where r_j are the GKR random point coordinates stored in rap_challenges. if self.has_trace_interaction() { let k = extract_column_indices(&self.auxiliary_trace_build_data.interactions).len(); - let rp_start = LOGUP_GAMMA_POWERS_START + k; + // +1: gamma_powers now has K+1 elements (γ^0..γ^K), so random_point starts at 4+K+1 + let rp_start = LOGUP_GAMMA_POWERS_START + k + 1; if rap_challenges.len() > rp_start { let n = rap_challenges.len() - rp_start; let mut l0_expected = FieldElement::::one(); @@ -1640,22 +1643,31 @@ where /// Compute the bridge offset (target/N) and gamma powers from column claims. /// /// Returns (bridge_offset, gamma_powers) where: -/// - bridge_offset = (Σ_j γ^j · c_j) / N -/// - gamma_powers = [γ^0, γ^1, ..., γ^{K-1}] +/// - bridge_offset = (Σ_j γ^j · c_j + γ^K · l_mle_claim) / N +/// - gamma_powers = [γ^0, γ^1, ..., γ^K] (K+1 powers; γ^K is for the l² self-check) +/// +/// The extra γ^K · l_mle_claim term in the target implements the Lagrange kernel +/// self-evaluation check (BUG-004): the bridge constraint forces Σ_i l[i]² = l_mle_claim +/// via Schwartz-Zippel over γ, where l_mle_claim = ∏_k (r_k² + (1-r_k)²) is the +/// expected squared-ℓ₂ norm of the true Lagrange kernel. /// /// Both prover and verifier call this to derive the same values. pub fn compute_bridge_params( column_claims: &[(usize, FieldElement)], gamma: &FieldElement, trace_len: usize, + l_mle_claim: &FieldElement, ) -> (FieldElement, Vec>) { let k = column_claims.len(); - let gamma_powers = compute_alpha_powers(gamma, k); + // K+1 powers: γ^0..γ^{K-1} for column inner products, γ^K for l² self-check + let gamma_powers = compute_alpha_powers(gamma, k + 1); let mut target = FieldElement::::zero(); for ((_, c_j), gp) in column_claims.iter().zip(gamma_powers.iter()) { target += c_j * gp; } + // γ^K · l_mle_claim: enforces Σ_i l[i]² = l_mle_claim via Schwartz-Zippel + target += l_mle_claim * &gamma_powers[k]; let n_inv = FieldElement::::from(trace_len as u64).inv().unwrap(); let bridge_offset = &target * &n_inv; @@ -1669,8 +1681,8 @@ pub fn compute_bridge_params( /// - [0] = z, [1] = α (original) /// - [2] = γ /// - [3] = bridge_offset (target/N) -/// - [4..4+K] = γ^0, γ^1, ..., γ^{K-1} -/// - [4+K..4+K+n] = random_point[0], ..., random_point[n-1] +/// - [4..4+K+1] = γ^0, γ^1, ..., γ^K (K+1 powers; γ^K is for the l² self-check) +/// - [4+K+1..4+K+1+n] = random_point[0], ..., random_point[n-1] pub fn extend_rap_challenges_with_bridge( rap_challenges: &mut Vec>, column_claims: &[(usize, FieldElement)], @@ -1678,14 +1690,25 @@ pub fn extend_rap_challenges_with_bridge( trace_len: usize, random_point: &[FieldElement], ) { - let (bridge_offset, gamma_powers) = compute_bridge_params(column_claims, gamma, trace_len); + // BUG-004: compute l_mle_claim = ∏_k (r_k² + (1-r_k)²). + // When l[i] = eq(bits(i), r), Σ_i l[i]² equals this value. + // Including it in the bridge target (via γ^K) forces the bridge constraint to + // check Σ_i l[i]² = l_mle_claim, binding the committed l column to eq(·, r). + let one = FieldElement::::one(); + let l_mle_claim = random_point.iter().fold(one.clone(), |acc, r_k| { + let one_minus_r = &one - r_k; + acc * (r_k.square() + one_minus_r.square()) + }); + + let (bridge_offset, gamma_powers) = + compute_bridge_params(column_claims, gamma, trace_len, &l_mle_claim); rap_challenges.push(gamma.clone()); // index 2 rap_challenges.push(bridge_offset); // index 3 for gp in gamma_powers { - rap_challenges.push(gp); // indices 4, 5, ... + rap_challenges.push(gp); // indices 4, 5, ..., 4+K } for rp in random_point { - rap_challenges.push(rp.clone()); // indices 4+K, 4+K+1, ... + rap_challenges.push(rp.clone()); // indices 4+K+1, 4+K+2, ... } } @@ -1750,7 +1773,9 @@ where // l (aux column 0) let l_curr = step0.get_aux_evaluation_element(0, 0); - // batched_curr = Σ_j γ^j · col_j_curr using precomputed gamma powers + // batched_curr = Σ_j γ^j · col_j_curr + γ^K · l_curr + // The extra γ^K · l_curr term makes l_curr * batched include γ^K · l_curr², + // enforcing Σ_i l[i]² = l_mle_claim via Schwartz-Zippel over γ (BUG-004 fix). let mut batched = FieldElement::::zero(); for (j, &col_idx) in self.column_indices.iter().enumerate() { let gamma_j = &rap_challenges[LOGUP_GAMMA_POWERS_START + j]; @@ -1758,6 +1783,10 @@ where // F×E→E: base field column × extension field gamma power batched += col_val * gamma_j; } + // γ^K self-check term (l² contribution) + let k = self.column_indices.len(); + let gamma_k = &rap_challenges[LOGUP_GAMMA_POWERS_START + k]; + batched += gamma_k * l_curr; // σ_next - σ_curr - l_curr * batched + bridge_offset transition_evaluations[self.constraint_idx] = @@ -1784,6 +1813,10 @@ where let col_val = step0.get_main_evaluation_element(0, col_idx); batched += col_val * gamma_j; } + // γ^K self-check term (l² contribution, mirrors prover path) + let k = self.column_indices.len(); + let gamma_k = &rap_challenges[LOGUP_GAMMA_POWERS_START + k]; + batched += gamma_k * l_curr; transition_evaluations[self.constraint_idx] = sigma_next - sigma_curr - l_curr * &batched + bridge_offset; @@ -2408,6 +2441,11 @@ pub struct LogUpGkrResult { /// * `interactions` - The AIR's bus interactions /// * `challenges` - LogUp challenges `[z, alpha]` /// +/// # Arguments +/// * `n_layers` - number of GKR summation layers for this instance (0 = single-row table). +/// For 0-layer instances the MLE of any column at the empty point equals its single row value, +/// so the cross-multiplication check is exact and applied unconditionally. +/// /// # Returns /// `true` if verification passes, `false` otherwise. pub fn reconstruct_and_verify_gkr_claims( @@ -2416,6 +2454,7 @@ pub fn reconstruct_and_verify_gkr_claims( column_claims: &[(usize, FieldElement)], interactions: &[BusInteraction], challenges: &[FieldElement], + n_layers: usize, ) -> bool { // Build a map from column index to claimed MLE value let claim_map: std::collections::HashMap> = column_claims @@ -2536,19 +2575,20 @@ pub fn reconstruct_and_verify_gkr_claims( // N(i) = sign * m(i), D(i) = fp(i) // so MLE(N)(r) and MLE(D)(r) can be exactly reconstructed from column MLEs. // - // For multi-interaction tables, the cross-multiplication introduces nonlinear - // terms (products of fingerprints/multiplicities across interactions), so - // MLE(N)(r) != n_recon and MLE(D)(r) != d_recon in general. - // Soundness for these tables is ensured by the bridge running sum constraint. - if interactions.len() == 1 { + // For multi-interaction tables with n_layers > 0, the cross-multiplication + // introduces nonlinear terms (products of fingerprints/multiplicities across + // interactions), so MLE(N)(r) != n_recon and MLE(D)(r) != d_recon in general. + // + // For 0-layer instances (n_layers == 0, single-row tables): the MLE at the empty + // evaluation point equals the single row value exactly — no nonlinearity. The check + // is valid and applied unconditionally regardless of interaction count (BUG-011 fix). + if interactions.len() == 1 || n_layers == 0 { // Direct check: reconstructed values must match GKR output as rational numbers. - // The GKR verifier may return (n_claim, d_claim) in a different representation - // than the prover's raw (numerator, denominator) — e.g. (claimed_sum, 1) instead - // of (root_n, root_d). So we compare as rationals: running_n / running_d == n_claim / d_claim + // Compare as rationals: running_n / running_d == n_claim / d_claim // i.e. running_n * d_claim == n_claim * running_d. &running_n * d_claim == n_claim * &running_d } else { - // Multi-interaction: structural check passed above. + // Multi-interaction with n_layers > 0: structural check passed above. // The bridge constraint (verified during STARK proof) ensures column_claims // are consistent with the committed trace. true diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index 6895e1cb4..cb48cd189 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -1767,10 +1767,18 @@ pub trait IsStarkProver< let (batch_gkr_proof_opt, gkr_results) = batch_gkr_proof; // ===================================================================== - // Round 1, Phase B'': Sample γ for bridge batching (main transcript) + // Round 1, Phase B'': Bind column_claims, then sample γ (BUG-012 fix) // ===================================================================== - // γ is sampled AFTER all GKR messages are bound to the transcript - // but BEFORE forking per table. The verifier replays the same sample. + // column_claims for multi-interaction tables were previously not transcript-bound. + // Append them before sampling γ so that γ is adaptive-prover resistant. + // Verifier must replay in the same order (same gkr_table_indices ordering). + if needs_lookup_challenges { + for result in gkr_results.iter().flatten() { + for (_, claim_val) in &result.column_claims { + transcript.append_field_element(claim_val); + } + } + } let gamma: FieldElement = if needs_lookup_challenges { transcript.sample_field_element() @@ -1807,30 +1815,45 @@ pub trait IsStarkProver< // Column 1: Bridge running sum σ // The constraint checks: σ[i+1] - σ[i] = l[i]·batched[i] - Δ - // where Δ = bridge_offset = target/N. + // where Δ = bridge_offset = target/N and batched[i] = Σ_j γ^j·col_j[i] + γ^K·l[i]. // So σ[0] = 0 (start), σ[i+1] = σ[i] + l[i]·batched[i] - Δ. - // The circular wrap-around at row N-1 requires σ[0] = σ[N-1] + l[N-1]·batched[N-1] - Δ, - // which telescopes to: 0 = Σ l[i]·batched[i] - N·Δ = target - target. + // The circular wrap-around at row N-1 telescopes to: + // 0 = Σ l[i]·batched[i] - N·Δ = target - target. + + // BUG-004 fix: compute l_mle_claim = ∏_k (r_k² + (1-r_k)²) + let one = FieldElement::::one(); + let l_mle_claim = result.random_point.iter().fold(one.clone(), |acc, r_k| { + let one_minus_r = &one - r_k; + acc * (r_k.square() + one_minus_r.square()) + }); + let (bridge_offset, gamma_powers) = crate::lookup::compute_bridge_params( &result.column_claims, &gamma, trace_len, + &l_mle_claim, ); let main_cols = trace.columns_main(); - // Pre-compute batched values in parallel: batched[i] = Σ_j main_cols[col_j][i] * γ^j + // Pre-compute batched values: batched[i] = Σ_j γ^j·col_j[i] + γ^K·l[i] + // The γ^K·l[i] term makes l[i]·batched[i] include γ^K·l[i]², + // enforcing Σ_i l[i]² = l_mle_claim via Schwartz-Zippel (BUG-004 fix). #[cfg(feature = "parallel")] let batched_iter = (0..trace_len).into_par_iter(); #[cfg(not(feature = "parallel"))] let batched_iter = 0..trace_len; let column_claims = &result.column_claims; + let k = column_claims.len(); // γ^K index + let kernel = &result.lagrange_kernel; let batched_values: Vec> = batched_iter .map(|row| { let mut batched = FieldElement::::zero(); for (j, (col_idx, _)) in column_claims.iter().enumerate() { batched += &main_cols[*col_idx][row] * &gamma_powers[j]; } + // γ^K · l[row] (self-check term) + batched += &gamma_powers[k] * &kernel[row]; batched }) .collect(); diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index 9e586092c..e8df6ffed 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -862,13 +862,16 @@ pub trait IsStarkVerifier< return false; } - // Verify that column_claims are consistent with GKR output + // Verify that column_claims are consistent with GKR output. + // Pass n_layers_by_instance[i] so 0-layer multi-interaction tables + // get the direct rational check (BUG-011 fix). if !crate::lookup::reconstruct_and_verify_gkr_claims( n_claim, d_claim, &gkr_proof_data.column_claims, air.bus_interactions(), &lookup_challenges, + n_layers_by_instance[i], ) { #[cfg(not(feature = "test_fiat_shamir"))] error!( @@ -902,9 +905,20 @@ pub trait IsStarkVerifier< } // ===================================================================== - // Phase B'': Sample γ for bridge batching (main transcript) + // Phase B'': Bind column_claims to transcript, then sample γ (BUG-012 fix) // ===================================================================== - // Must match prover: γ sampled AFTER all GKR messages, BEFORE forking. + // column_claims for multi-interaction tables were previously not transcript-bound, + // creating an adaptive-prover gap. Append them before sampling γ so that γ + // depends on the claimed column MLE values. Must match prover ordering exactly. + if needs_lookup_challenges { + for &table_idx in &gkr_table_indices { + if let Some(p) = &multi_proof.proofs[table_idx].logup_gkr_proof { + for (_, claim_val) in &p.column_claims { + transcript.append_field_element(claim_val); + } + } + } + } let gamma: FieldElement = if needs_lookup_challenges { transcript.sample_field_element() From dfd4c53be11d30417cb252aac84511a04ef6757f Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 12:33:21 -0300 Subject: [PATCH 16/19] Return false instead of panicking on zero GKR root denominator --- crypto/stark/src/gkr.rs | 8 ++++++++ crypto/stark/src/verifier.rs | 9 ++++++++- 2 files changed, 16 insertions(+), 1 deletion(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 779534590..19d06d9bc 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -2047,6 +2047,14 @@ pub fn gkr_verify_batch( // instance i's child layer has (n_layers[i] - n_remaining + 1) variables, // so the parent has (n_layers[i] - n_remaining) variables. let parent_num_vars_i = n_layers_by_instance[i] - n_remaining; + if num_rounds < parent_num_vars_i { + return Err(GkrError::InvalidTree { + reason: format!( + "layer {}: num_rounds ({}) < parent_num_vars for instance {} ({})", + layer_idx, num_rounds, i, parent_num_vars_i, + ), + }); + } let sumcheck_n_unused = num_rounds - parent_num_vars_i; // eq evaluation: the prover builds eq from the instance-specific eval point diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index e8df6ffed..9b358fd7e 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -990,7 +990,14 @@ pub trait IsStarkVerifier< let contrib = if *root_d == FieldElement::one() { root_n.clone() } else { - root_n * &root_d.inv().expect("GKR root denominator must be non-zero") + match root_d.inv() { + Ok(inv) => root_n * &inv, + Err(_) => { + #[cfg(not(feature = "test_fiat_shamir"))] + error!("GKR root denominator is zero — invalid proof"); + return false; + } + } }; total = total + &contrib; } From e7761ed94ab0d886ab7f8e7c6cb2bc8e46193d2d Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 12:53:50 -0300 Subject: [PATCH 17/19] Add gate check for trivial layers and fix single-proof API --- crypto/stark/src/gkr.rs | 44 +++++++++++++++++++++++++++++---- crypto/stark/src/proof/stark.rs | 4 +++ crypto/stark/src/prover.rs | 9 +++++-- crypto/stark/src/verifier.rs | 2 +- 4 files changed, 51 insertions(+), 8 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 19d06d9bc..18210064e 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -1085,10 +1085,18 @@ pub fn gkr_verify( if round_polys.is_empty() { // Trivial layer (0 variables in parent): no sumcheck rounds. - // The prover just provides child_claims directly. - // No gate check here --- soundness is enforced by later layers' sumchecks. - // (The verifier's combined_claim may differ from the prover's by a - // scaling factor at the first trivial layer, so we skip the check.) + // Gate check: verify that n_claim * (dl·dr) = nl·dr + nr·dl. + // + // The verifier works in normalized form: n_claim = root_n/root_d, d_claim = 1. + // The prover's gate equation root_n + λ·root_d = nl·dr + nr·dl + λ·dl·dr, + // divided by root_d (= dl·dr), becomes n_claim·(dl·dr) = nl·dr + nr·dl + // (the λ terms cancel). This binds claimed_sum to the actual tree structure. + let [ref nl, ref nr, ref dl, ref dr] = layer_proof.child_claims; + let lhs = &n_claim * &(dl * dr); + let rhs = &(nl * dr) + &(nr * dl); + if lhs != rhs { + return Err(GkrError::GateCheckFailed { layer: layer_idx }); + } } else { // Non-trivial layer: verify sumcheck inline. let num_rounds = round_polys.len(); @@ -1985,7 +1993,33 @@ pub fn gkr_verify_batch( } if round_polys.is_empty() { - // Trivial layer: no sumcheck needed + // Trivial layer: no sumcheck rounds. + // Gate check: verify that the alpha-batched combined claims match the + // alpha-batched gate evaluations (gate_i = nl·dr + nr·dl + λ·dl·dr, scaled + // by 2^n_unused_i to match combined_claims). + { + let mut actual_sum = FieldElement::::zero(); + let mut expected_sum = FieldElement::::zero(); + let mut alpha_pow = FieldElement::::one(); + for (idx, &i) in active_instances.iter().enumerate() { + let [ref nl, ref nr, ref dl, ref dr] = + layer_proof.child_claims_by_instance[idx]; + let gate = &(&(nl * dr) + &(nr * dl)) + &(&lambda * &(dl * dr)); + let n_unused = max_layers - n_layers_by_instance[i]; + let gate_scaled = if n_unused > 0 { + &gate * &FieldElement::::from(1u64 << n_unused) + } else { + gate + }; + actual_sum = &actual_sum + &(&alpha_pow * &combined_claims[idx]); + expected_sum = &expected_sum + &(&alpha_pow * &gate_scaled); + alpha_pow = &alpha_pow * &sumcheck_alpha; + } + if actual_sum != expected_sum { + return Err(GkrError::GateCheckFailed { layer: layer_idx }); + } + } + // Append child claims and sample eta (same as prover) for claims in &layer_proof.child_claims_by_instance { for claim in claims { diff --git a/crypto/stark/src/proof/stark.rs b/crypto/stark/src/proof/stark.rs index 8e440e232..18d50322a 100644 --- a/crypto/stark/src/proof/stark.rs +++ b/crypto/stark/src/proof/stark.rs @@ -85,6 +85,10 @@ pub struct StarkProof, E: IsField, PI> { pub bus_public_inputs: Option>, // LogUp-GKR proof (when using GKR-based LogUp instead of per-row accumulated column) pub logup_gkr_proof: Option>, + // Batch GKR proof shared across all tables in a single-table prove/verify round-trip. + // Populated by `Prover::prove` and consumed by `Verifier::verify`. + #[serde(default)] + pub batch_gkr_proof: Option>, // Public inputs used for boundary constraints pub public_inputs: PI, } diff --git a/crypto/stark/src/prover.rs b/crypto/stark/src/prover.rs index cb48cd189..ab581b275 100644 --- a/crypto/stark/src/prover.rs +++ b/crypto/stark/src/prover.rs @@ -2171,8 +2171,11 @@ pub trait IsStarkProver< PI: Send + Sync + Clone, { let air_trace_pairs = vec![(air, trace, pub_inputs)]; - Self::multi_prove(air_trace_pairs, transcript) - .map(|mut multi_proof| multi_proof.proofs.remove(0)) + Self::multi_prove(air_trace_pairs, transcript).map(|mut multi_proof| { + let mut proof = multi_proof.proofs.remove(0); + proof.batch_gkr_proof = multi_proof.batch_gkr_proof; + proof + }) } // TODO: propagate errors instead of unwrap() in open_deep_composition_poly and FRI operations @@ -2335,6 +2338,8 @@ pub trait IsStarkProver< bus_public_inputs: round_1_result.bus_public_inputs.clone(), // LogUp-GKR proof (not yet used; will replace accumulated column in future) logup_gkr_proof: None, + // Batch GKR proof: populated by Prover::prove after multi_prove completes. + batch_gkr_proof: None, // Public inputs for boundary constraints public_inputs: pub_inputs.clone(), trace_length: domain.interpolation_domain_size, diff --git a/crypto/stark/src/verifier.rs b/crypto/stark/src/verifier.rs index 9b358fd7e..76214cc4c 100644 --- a/crypto/stark/src/verifier.rs +++ b/crypto/stark/src/verifier.rs @@ -1032,7 +1032,7 @@ pub trait IsStarkVerifier< { let multi_proof = MultiProof { proofs: vec![proof.clone()], - batch_gkr_proof: None, + batch_gkr_proof: proof.batch_gkr_proof.clone(), }; Self::multi_verify(&[air], &multi_proof, transcript, &FieldElement::zero()) } From 72ad8cf7e6b65b41625ef51a802e1785b6015dca Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 14:47:23 -0300 Subject: [PATCH 18/19] Delete dead pub, return Result instead of panicking, fix comment and add tests --- crypto/stark/src/gkr.rs | 128 ++++++++++++++---- crypto/stark/src/lookup.rs | 112 +-------------- .../src/tests/bus_tests/completeness_tests.rs | 43 ++++++ 3 files changed, 143 insertions(+), 140 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 18210064e..4450ce173 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -389,12 +389,15 @@ fn evaluate_mle( pub fn gkr_prove( tree: &[SummationLayer], transcript: &mut impl IsTranscript, -) -> ( - GkrProof, - Vec>, - FieldElement, - FieldElement, -) { +) -> Result< + ( + GkrProof, + Vec>, + FieldElement, + FieldElement, + ), + GkrError, +> { assert!(!tree.is_empty(), "tree must have at least one layer"); let num_layers = tree.len(); // layers 0..num_layers-1, root is at num_layers-1 @@ -408,8 +411,10 @@ pub fn gkr_prove( let root_n = &root.numerators[0]; let root_d = &root.denominators[0]; - // Compute claimed_sum = root_n / root_d - let root_d_inv = root_d.inv().expect("root denominator must be nonzero"); + // Compute claimed_sum = root_n / root_d. root_d = 0 means the LogUp challenge z + // collided with a fingerprint denominator (probability ~1/p ≈ 2^{-64}); return + // an error rather than panicking so the caller can retry with a fresh transcript. + let root_d_inv = root_d.inv().map_err(|_| GkrError::ZeroDenominator)?; let claimed_sum = root_n * &root_d_inv; // Append the claimed sum to the transcript @@ -417,7 +422,7 @@ pub fn gkr_prove( // If the tree has only 1 layer (just leaves = root), no reductions needed if num_layers == 1 { - return ( + return Ok(( GkrProof { claimed_sum, layer_proofs: vec![], @@ -425,7 +430,7 @@ pub fn gkr_prove( vec![], root_n.clone(), root_d.clone(), - ); + )); } let mut layer_proofs = Vec::with_capacity(num_layers - 1); @@ -992,7 +997,7 @@ pub fn gkr_prove( } } - ( + Ok(( GkrProof { claimed_sum, layer_proofs, @@ -1000,10 +1005,10 @@ pub fn gkr_prove( current_point, n_claim, d_claim, - ) + )) } -/// Errors that can occur during GKR verification. +/// Errors that can occur during GKR proving or verification. #[derive(Debug, Clone)] pub enum GkrError { /// The summation tree structure is invalid. @@ -1015,6 +1020,9 @@ pub enum GkrError { /// The claimed sum does not match (unused in the verifier itself, /// but available for callers that compare the claimed sum to an external value). ClaimedSumMismatch, + /// The root denominator is zero (LogUp challenge z collided with a fingerprint + /// denominator). Probability ~1/p ≈ 2^{-64}; prover should sample a new transcript. + ZeroDenominator, } impl fmt::Display for GkrError { @@ -1032,6 +1040,12 @@ impl fmt::Display for GkrError { GkrError::ClaimedSumMismatch => { write!(f, "claimed sum mismatch") } + GkrError::ZeroDenominator => { + write!( + f, + "GKR root denominator is zero (LogUp challenge z collided with a fingerprint denominator; probability ~1/p ≈ 2^{{-64}})" + ) + } } } } @@ -1069,7 +1083,7 @@ pub fn gkr_verify( // Step 2: Initialize claims. // The verifier sets n_claim = claimed_sum, d_claim = 1. // This represents the same rational value as root_n/root_d. - // Soundness of the first (trivial) layer is enforced by later sumcheck layers. + // Trivial layers (0 sumcheck rounds) are gate-checked directly in the loop below. let mut n_claim = proof.claimed_sum.clone(); let mut d_claim = FieldElement::::one(); let mut current_point: Vec> = vec![]; @@ -2400,7 +2414,8 @@ mod tests { let tree = build_summation_tree(nums, dens); let mut transcript = DefaultTranscript::::new(&[]); - let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + let (proof, final_point, final_n_claim, final_d_claim) = + gkr_prove(&tree, &mut transcript).unwrap(); // claimed_sum = 68/55 let expected_sum = &FE::from(68u64) * &FE::from(55u64).inv().unwrap(); @@ -2451,7 +2466,8 @@ mod tests { let tree = build_summation_tree(nums.clone(), dens.clone()); let mut transcript = DefaultTranscript::::new(&[]); - let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + let (proof, final_point, final_n_claim, final_d_claim) = + gkr_prove(&tree, &mut transcript).unwrap(); // claimed_sum = root_n / root_d = 1136 / 384 let expected_sum = &FE::from(1136u64) * &FE::from(384u64).inv().unwrap(); @@ -2505,7 +2521,7 @@ mod tests { let expected_sum = root_n * &root_d.inv().unwrap(); let mut transcript = DefaultTranscript::::new(&[]); - let (proof, _, _, _) = gkr_prove(&tree, &mut transcript); + let (proof, _, _, _) = gkr_prove(&tree, &mut transcript).unwrap(); assert_eq!(proof.claimed_sum, expected_sum); } @@ -2520,7 +2536,8 @@ mod tests { let tree = build_summation_tree(nums.clone(), dens.clone()); let mut transcript = DefaultTranscript::::new(&[]); - let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + let (proof, final_point, final_n_claim, final_d_claim) = + gkr_prove(&tree, &mut transcript).unwrap(); // Should have 3 layer proofs assert_eq!(proof.layer_proofs.len(), 3); @@ -2553,7 +2570,8 @@ mod tests { let tree = build_summation_tree(nums.clone(), dens.clone()); let mut transcript = DefaultTranscript::::new(&[0xAB]); - let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + let (proof, final_point, final_n_claim, final_d_claim) = + gkr_prove(&tree, &mut transcript).unwrap(); // Should have 4 layer proofs assert_eq!(proof.layer_proofs.len(), 4); @@ -2586,10 +2604,10 @@ mod tests { let tree = build_summation_tree(nums, dens); let mut t1 = DefaultTranscript::::new(&[0x42]); - let (proof1, point1, n1, d1) = gkr_prove(&tree, &mut t1); + let (proof1, point1, n1, d1) = gkr_prove(&tree, &mut t1).unwrap(); let mut t2 = DefaultTranscript::::new(&[0x42]); - let (proof2, point2, n2, d2) = gkr_prove(&tree, &mut t2); + let (proof2, point2, n2, d2) = gkr_prove(&tree, &mut t2).unwrap(); assert_eq!(proof1.claimed_sum, proof2.claimed_sum); assert_eq!(point1, point2); @@ -2635,7 +2653,7 @@ mod tests { // Re-run the protocol manually to extract the combined_claim at each layer // and verify the sumcheck round poly sum let mut transcript = DefaultTranscript::::new(&[]); - let (proof, _, _, _) = gkr_prove(&tree, &mut transcript); + let (proof, _, _, _) = gkr_prove(&tree, &mut transcript).unwrap(); // Replay transcript to get the same challenges let mut replay = DefaultTranscript::::new(&[]); @@ -2700,7 +2718,8 @@ mod tests { let tree = build_summation_tree(nums, dens); let mut transcript = DefaultTranscript::::new(&[]); - let (proof, final_point, final_n_claim, final_d_claim) = gkr_prove(&tree, &mut transcript); + let (proof, final_point, final_n_claim, final_d_claim) = + gkr_prove(&tree, &mut transcript).unwrap(); assert_eq!(proof.claimed_sum, FE::from(6u64)); // 42/7 = 6 assert_eq!(proof.layer_proofs.len(), 0); @@ -2732,7 +2751,8 @@ mod tests { // Prove let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); - let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + let (proof, prover_point, prover_n, prover_d) = + gkr_prove(&tree, &mut prover_transcript).unwrap(); // Verify with a fresh transcript (same seed) let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); @@ -2772,7 +2792,8 @@ mod tests { // Prove let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); - let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + let (proof, prover_point, prover_n, prover_d) = + gkr_prove(&tree, &mut prover_transcript).unwrap(); // Verify with a fresh transcript (same seed) let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); @@ -2818,7 +2839,7 @@ mod tests { // Prove with correct claimed_sum let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); - let (mut proof, _, _, _) = gkr_prove(&tree, &mut prover_transcript); + let (mut proof, _, _, _) = gkr_prove(&tree, &mut prover_transcript).unwrap(); // Tamper with the claimed_sum proof.claimed_sum = &proof.claimed_sum + &FE::one(); @@ -2836,6 +2857,52 @@ mod tests { ); } + #[test] + fn test_gkr_verify_trivial_layer_gate_check_rejected() { + // A 4-leaf tree produces: + // layer_proofs[0]: trivial (root→layer1, parent_num_vars=0, round_polys=[]) + // layer_proofs[1]: non-trivial (layer1→leaves, parent_num_vars=1) + // + // Tamper with child_claims of the trivial layer_proofs[0] and assert + // that gkr_verify returns GkrError::GateCheckFailed. + let nums = vec![ + FE::from(1u64), + FE::from(3u64), + FE::from(5u64), + FE::from(7u64), + ]; + let dens = vec![ + FE::from(2u64), + FE::from(4u64), + FE::from(6u64), + FE::from(8u64), + ]; + let tree = build_summation_tree(nums, dens); + + let mut prover_transcript = DefaultTranscript::::new(&[0xC0]); + let (mut proof, _, _, _) = gkr_prove(&tree, &mut prover_transcript).unwrap(); + + // layer_proofs[0] is the trivial layer: round_polys must be empty + assert!( + proof.layer_proofs[0].sumcheck_proof.round_polys.is_empty(), + "layer_proofs[0] must be the trivial layer for a 4-leaf tree" + ); + + // Corrupt nl (child_claims[0]) by adding 1 — this breaks the gate equation + // n_claim*(dl*dr) = nl*dr + nr*dl without altering the transcript ordering + proof.layer_proofs[0].child_claims[0] = + proof.layer_proofs[0].child_claims[0].clone() + FE::one(); + + let mut verifier_transcript = DefaultTranscript::::new(&[0xC0]); + let result = gkr_verify(&proof, &mut verifier_transcript); + + assert!( + matches!(result, Err(GkrError::GateCheckFailed { layer: 0 })), + "Tampered trivial-layer child_claims must be rejected with GateCheckFailed {{ layer: 0 }}, got: {:?}", + result + ); + } + #[test] fn test_gkr_prove_verify_roundtrip_16() { // 16 leaves with various fractions @@ -2846,7 +2913,8 @@ mod tests { // Prove let mut prover_transcript = DefaultTranscript::::new(&[0xAA]); - let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + let (proof, prover_point, prover_n, prover_d) = + gkr_prove(&tree, &mut prover_transcript).unwrap(); // Verify with a fresh transcript (same seed) let mut verifier_transcript = DefaultTranscript::::new(&[0xAA]); @@ -2885,7 +2953,8 @@ mod tests { // Prove let mut prover_transcript = DefaultTranscript::::new(&[0x5F, 0x00]); - let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + let (proof, prover_point, prover_n, prover_d) = + gkr_prove(&tree, &mut prover_transcript).unwrap(); // Verify with a fresh transcript (same seed) let mut verifier_transcript = DefaultTranscript::::new(&[0x5F, 0x00]); @@ -2922,7 +2991,8 @@ mod tests { let tree = build_summation_tree(nums.clone(), dens.clone()); let mut prover_transcript = DefaultTranscript::::new(&[0xBB]); - let (proof, prover_point, prover_n, prover_d) = gkr_prove(&tree, &mut prover_transcript); + let (proof, prover_point, prover_n, prover_d) = + gkr_prove(&tree, &mut prover_transcript).unwrap(); let mut verifier_transcript = DefaultTranscript::::new(&[0xBB]); let result = gkr_verify(&proof, &mut verifier_transcript); diff --git a/crypto/stark/src/lookup.rs b/crypto/stark/src/lookup.rs index 0de9d1de4..e45de51ae 100644 --- a/crypto/stark/src/lookup.rs +++ b/crypto/stark/src/lookup.rs @@ -2384,9 +2384,8 @@ where // LogUp-GKR Integration // ============================================================================= -use crate::gkr::{GkrProof, Layer, build_summation_tree, gen_layers, gkr_prove}; +use crate::gkr::{GkrProof, Layer, gen_layers}; use crate::lagrange_kernel::{compute_lagrange_kernel, eval_mle_base_with_kernel}; -use crypto::fiat_shamir::is_transcript::IsTranscript; /// Result of running the LogUp-GKR sub-protocol for a single table. /// @@ -2852,115 +2851,6 @@ where } } -/// Run the LogUp-GKR sub-protocol for a single table's bus interactions. -/// -/// This function: -/// 1. Computes per-row leaf fractions (numerator, denominator) from interactions -/// 2. Builds a binary summation tree over the leaf fractions -/// 3. Runs the GKR protocol to prove the summation tree root -/// 4. Extracts MLE claims for each distinct main trace column at the GKR random point -/// -/// The GKR proof replaces the traditional per-row accumulated column with a -/// logarithmic-depth interactive proof, reducing auxiliary trace columns. -pub fn run_logup_gkr( - interactions: &[BusInteraction], - main_segment_cols: &[Vec>], - trace_len: usize, - challenges: &[FieldElement], - transcript: &mut impl IsTranscript, -) -> LogUpGkrResult -where - F: IsFFTField + IsSubFieldOf + IsPrimeField + Send + Sync, - E: IsField + Send + Sync, -{ - // Step 1: Compute per-row leaf fractions - let (numerators, denominators) = - compute_logup_leaf_fractions(interactions, main_segment_cols, trace_len, challenges); - - // Step 2: Build the summation tree - let tree = build_summation_tree(numerators, denominators); - - // Step 3: Run the GKR protocol - let (gkr_proof, random_point, n_claim, d_claim) = gkr_prove(&tree, transcript); - - let table_contribution = gkr_proof.claimed_sum.clone(); - - // Step 4: Extract column claims — compute MLE at the random point for each - // distinct main trace column index referenced by any interaction. - let mut seen_cols = std::collections::HashSet::new(); - for inter in interactions { - for val in &inter.values { - for col_idx in val.column_indices() { - seen_cols.insert(col_idx); - } - } - // Also collect column indices from multiplicities - match &inter.multiplicity { - Multiplicity::One => {} - Multiplicity::Column(c) => { - seen_cols.insert(*c); - } - Multiplicity::Sum(a, b) => { - seen_cols.insert(*a); - seen_cols.insert(*b); - } - Multiplicity::Negated(c) => { - seen_cols.insert(*c); - } - Multiplicity::Diff(a, b) => { - seen_cols.insert(*a); - seen_cols.insert(*b); - } - Multiplicity::Sum3(a, b, c) => { - seen_cols.insert(*a); - seen_cols.insert(*b); - seen_cols.insert(*c); - } - Multiplicity::Linear(terms) => { - for term in terms { - match term { - LinearTerm::Column { column, .. } => { - seen_cols.insert(*column); - } - LinearTerm::ColumnUnsigned { column, .. } => { - seen_cols.insert(*column); - } - LinearTerm::Constant(_) => {} - } - } - } - } - } - - let mut col_indices: Vec = seen_cols.into_iter().collect(); - col_indices.sort_unstable(); - - // Compute kernel once and reuse for all column claims (and later for aux trace) - let kernel = compute_lagrange_kernel(&random_point); - - #[cfg(feature = "parallel")] - let col_iter = col_indices.into_par_iter(); - #[cfg(not(feature = "parallel"))] - let col_iter = col_indices.into_iter(); - - let column_claims: Vec<(usize, FieldElement)> = col_iter - .map(|col_idx| { - let claim = eval_mle_base_with_kernel(&main_segment_cols[col_idx], &kernel); - (col_idx, claim) - }) - .collect(); - - LogUpGkrResult { - table_contribution, - gkr_proof: Some(gkr_proof), - random_point, - n_claim, - d_claim, - column_claims, - lagrange_kernel: kernel, - } -} - #[cfg(test)] #[allow(clippy::clone_on_copy, clippy::cloned_ref_to_slice_refs)] mod tests { diff --git a/crypto/stark/src/tests/bus_tests/completeness_tests.rs b/crypto/stark/src/tests/bus_tests/completeness_tests.rs index 7ca124fe1..9e842aa76 100644 --- a/crypto/stark/src/tests/bus_tests/completeness_tests.rs +++ b/crypto/stark/src/tests/bus_tests/completeness_tests.rs @@ -531,3 +531,46 @@ fn test_bus_value_features() { &FieldElement::zero(), )); } + +/// Single-table prove+verify via Prover::prove and Verifier::verify (regression for BUG-014). +/// +/// Verifies that batch_gkr_proof is propagated through the single-proof path: +/// Prover::prove must embed it in StarkProof, and Verifier::verify must accept it. +/// +/// The CPU table with all-zero multiplicity columns (add_flag=mul_flag=0) has +/// a GKR claimed_sum of 0 per bus, so the bus balance (sum = 0) holds without +/// a corresponding receiver table. +#[test_log::test] +fn test_single_table_prove_verify_with_gkr() { + let mut cpu_trace = TraceTable::from_columns_main( + vec![ + vec![FE::zero(); 4], // add_flag = 0: no sends to ADD bus + vec![FE::zero(); 4], // mul_flag = 0: no sends to MUL bus + vec![FE::zero(); 4], + vec![FE::zero(); 4], + vec![FE::zero(); 4], + ], + 1, + ); + + let proof_options = ProofOptions::default_test_options(); + let cpu_air = new_cpu_air_with_lookup(&proof_options); + + let proof = Prover::prove( + &cpu_air, + &mut cpu_trace, + &(), + &mut DefaultTranscript::::new(&[]), + ) + .expect("Prover::prove must succeed for a GKR-backed AIR"); + + assert!( + proof.batch_gkr_proof.is_some(), + "batch_gkr_proof must be propagated by Prover::prove (BUG-014 regression)" + ); + + assert!( + Verifier::verify(&proof, &cpu_air, &mut DefaultTranscript::::new(&[]),), + "Verifier::verify must accept a proof produced by Prover::prove for a GKR-backed AIR" + ); +} From 78be304f7f2712686399df39964baf2ee0fac73a Mon Sep 17 00:00:00 2001 From: Nicole Date: Mon, 13 Apr 2026 15:04:57 -0300 Subject: [PATCH 19/19] fix fmt --- crypto/stark/src/gkr.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/crypto/stark/src/gkr.rs b/crypto/stark/src/gkr.rs index 4450ce173..e00d64543 100644 --- a/crypto/stark/src/gkr.rs +++ b/crypto/stark/src/gkr.rs @@ -386,6 +386,7 @@ fn evaluate_mle( /// /// # Panics /// Panics if the tree is empty or has inconsistent layer sizes. +#[allow(clippy::type_complexity)] pub fn gkr_prove( tree: &[SummationLayer], transcript: &mut impl IsTranscript, @@ -2890,8 +2891,7 @@ mod tests { // Corrupt nl (child_claims[0]) by adding 1 — this breaks the gate equation // n_claim*(dl*dr) = nl*dr + nr*dl without altering the transcript ordering - proof.layer_proofs[0].child_claims[0] = - proof.layer_proofs[0].child_claims[0].clone() + FE::one(); + proof.layer_proofs[0].child_claims[0] += FE::one(); let mut verifier_transcript = DefaultTranscript::::new(&[0xC0]); let result = gkr_verify(&proof, &mut verifier_transcript);