Skip to main content

openvm_stark_backend/verifier/
batch_constraints.rs

1use std::{
2    iter::{self, zip},
3    slice,
4};
5
6use itertools::Itertools;
7use p3_field::{batch_multiplicative_inverse, Field, PrimeCharacteristicRing};
8use thiserror::Error;
9use tracing::{debug, instrument};
10
11use crate::{
12    air_builders::symbolic::{symbolic_expression::SymbolicEvaluator, SymbolicConstraints},
13    calculate_n_logup,
14    keygen::types::MultiStarkVerifyingKey0,
15    poly_common::{eval_eq_mle, eval_eq_sharp_uni, eval_eq_uni, UnivariatePoly},
16    proof::{column_openings_by_rot, BatchConstraintProof, GkrProof},
17    verifier::{
18        evaluator::VerifierConstraintEvaluator,
19        fractional_sumcheck_gkr::{verify_gkr, GkrVerificationError},
20    },
21    FiatShamirTranscript, StarkProtocolConfig,
22};
23
24#[derive(Error, Debug, PartialEq, Eq)]
25pub enum BatchConstraintError<EF: core::fmt::Debug + core::fmt::Display + PartialEq + Eq> {
26    #[error("Invalid logup_pow_witness")]
27    InvalidLogupPowWitness,
28
29    #[error("GKR verification failed: {0}")]
30    GkrVerificationFailed(#[from] GkrVerificationError<EF>),
31
32    #[error("GKR numerator evaluation claim {claim} does not match")]
33    GkrNumeratorMismatch { claim: EF },
34
35    #[error("GKR denominator evaluation claim {claim} does not match")]
36    GkrDenominatorMismatch { claim: EF },
37
38    #[error(
39        "`sum_claim` does not equal the sum of `s_0` at all the roots of unity: {sum_claim} != {sum_univ_domain_s_0}"
40    )]
41    SumClaimMismatch {
42        sum_claim: EF,
43        sum_univ_domain_s_0: EF,
44    },
45
46    #[error("Claims are inconsistent")]
47    InconsistentClaims,
48}
49
50/// `public_values` should be in vkey (air_idx) order, including non-present AIRs.
51#[allow(clippy::too_many_arguments)]
52#[instrument(level = "debug", skip_all)]
53pub fn verify_zerocheck_and_logup<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
54    transcript: &mut TS,
55    mvk: &MultiStarkVerifyingKey0<SC>,
56    public_values: &[Vec<SC::F>],
57    gkr_proof: &GkrProof<SC>,
58    batch_proof: &BatchConstraintProof<SC>,
59    trace_id_to_air_id: &[usize],
60    n_per_trace: &[isize],
61    omega_skip_pows: &[SC::F],
62) -> Result<Vec<SC::EF>, BatchConstraintError<SC::EF>> {
63    let l_skip = mvk.params.l_skip;
64    // Proof shape asserts that numerator_term_per_air.len() == denominator_term_per_air.len() ==
65    // num_traces (=: num_airs_present)
66    let BatchConstraintProof {
67        numerator_term_per_air,
68        denominator_term_per_air,
69        univariate_round_coeffs,
70        sumcheck_round_polys,
71        column_openings,
72    } = batch_proof;
73
74    // 1. Check GKR witness
75    if !transcript.check_witness(mvk.params.logup.pow_bits, gkr_proof.logup_pow_witness) {
76        return Err(BatchConstraintError::InvalidLogupPowWitness);
77    }
78
79    // 2. Sample alpha and beta, receive xi, sample lambda
80    let alpha_logup = transcript.sample_ext();
81    let beta_logup = transcript.sample_ext();
82    debug!(%alpha_logup, %beta_logup);
83    let total_interactions = zip(trace_id_to_air_id, n_per_trace)
84        .map(|(&air_idx, &n)| {
85            let n_lift = n.max(0) as usize;
86            let num_interactions = mvk.per_air[air_idx].symbolic_constraints.interactions.len();
87            (num_interactions as u64) << (l_skip + n_lift)
88        })
89        .sum::<u64>();
90    let n_logup: usize = calculate_n_logup(l_skip, total_interactions);
91    debug!(%n_logup);
92
93    let mut xi = Vec::new();
94    let mut p_xi_claim = SC::EF::ZERO;
95    let mut q_xi_claim = alpha_logup;
96    if total_interactions > 0 {
97        (p_xi_claim, q_xi_claim, xi) =
98            verify_gkr::<SC, TS>(gkr_proof, transcript, l_skip + n_logup)?;
99        debug_assert_eq!(xi.len(), l_skip + n_logup);
100    } else if gkr_proof.q0_claim != SC::EF::ONE {
101        return Err(GkrVerificationError::InvalidZeroRoundValue {
102            actual: gkr_proof.q0_claim,
103        }
104        .into());
105    }
106
107    let n_max = n_per_trace.iter().copied().max().unwrap().max(0) as usize;
108    let n_global = n_max.max(n_logup);
109    while xi.len() != l_skip + n_global {
110        xi.push(transcript.sample_ext());
111    }
112    debug!(%n_max);
113    debug!(?xi);
114
115    let lambda = transcript.sample_ext();
116    debug!(%lambda);
117
118    // 3. Observe everything from numerator_per_air and denominator_per_air, compute its sum
119    for (&sum_claim_p, &sum_claim_q) in zip(numerator_term_per_air, denominator_term_per_air) {
120        p_xi_claim -= sum_claim_p;
121        q_xi_claim -= sum_claim_q;
122        transcript.observe_ext(sum_claim_p);
123        transcript.observe_ext(sum_claim_q);
124    }
125    if p_xi_claim != SC::EF::ZERO {
126        return Err(BatchConstraintError::GkrNumeratorMismatch { claim: p_xi_claim });
127    }
128    if q_xi_claim != alpha_logup {
129        return Err(BatchConstraintError::GkrDenominatorMismatch { claim: q_xi_claim });
130    }
131
132    // 4. Sample mu, compute the mu-hash of interleave of numerator_per_air and denominator_per_air
133    let mu = transcript.sample_ext();
134    debug!(%mu);
135
136    let mut sum_claim = SC::EF::ZERO;
137    let mut cur_mu_pow = SC::EF::ONE;
138    for (&sum_claim_p, &sum_claim_q) in zip(numerator_term_per_air, denominator_term_per_air) {
139        sum_claim += sum_claim_p * cur_mu_pow;
140        cur_mu_pow *= mu;
141        sum_claim += sum_claim_q * cur_mu_pow;
142        cur_mu_pow *= mu;
143    }
144
145    // 5. Univariate sumcheck round
146    for &coeff in univariate_round_coeffs {
147        transcript.observe_ext(coeff);
148    }
149
150    let s_deg = mvk.params.max_constraint_degree + 1;
151    let r_0 = transcript.sample_ext();
152    debug!(round = 0, r_round = %r_0);
153    assert_eq!(
154        univariate_round_coeffs.len(),
155        (mvk.max_constraint_degree() + 1) * ((1 << l_skip) - 1) + 1
156    );
157    let s_0 = UnivariatePoly::new(univariate_round_coeffs.clone());
158    let sum_univ_domain_s_0 = s_0
159        .coeffs()
160        .iter()
161        .step_by(1 << l_skip)
162        .copied()
163        .sum::<SC::EF>()
164        * SC::EF::from_usize(1 << l_skip);
165    if sum_claim != sum_univ_domain_s_0 {
166        return Err(BatchConstraintError::SumClaimMismatch {
167            sum_claim,
168            sum_univ_domain_s_0,
169        });
170    }
171    let mut cur_sum = s_0.eval_at_point(r_0);
172    let mut rs = vec![r_0];
173
174    // 6. Multilinear sumcheck rounds
175    #[allow(clippy::needless_range_loop)]
176    for round in 0..n_max {
177        debug!(sumcheck_round = round, sum_claim = %cur_sum, "batch_constraint_sumcheck");
178        // Proof shape asserts that sumcheck_round_polys.len() == n_max
179        let batch_s_evals = &sumcheck_round_polys[round];
180        // Proof shape asserts that batch_s_evals.len() == s_deg
181        for &eval in batch_s_evals.iter() {
182            transcript.observe_ext(eval);
183        }
184        let s_1 = batch_s_evals[0];
185        let s_0 = cur_sum - s_1;
186        let batch_s_evals = iter::once(&s_0).chain(batch_s_evals).collect_vec();
187
188        let mut factorials = vec![SC::F::ONE; s_deg + 1];
189        for i in 1..=s_deg {
190            factorials[i] = factorials[i - 1] * SC::F::from_usize(i);
191        }
192        let invfact = batch_multiplicative_inverse(&factorials);
193
194        let r = transcript.sample_ext();
195        let mut pref_product = vec![SC::EF::ONE; s_deg + 1];
196        let mut suf_product = vec![SC::EF::ONE; s_deg + 1];
197        for i in 0..s_deg {
198            pref_product[i + 1] = pref_product[i] * (r - SC::EF::from_usize(i));
199            suf_product[i + 1] = suf_product[i] * (SC::EF::from_usize(s_deg - i) - r);
200        }
201        cur_sum = (0..=s_deg)
202            .map(|i| {
203                *batch_s_evals[i]
204                    * pref_product[i]
205                    * suf_product[s_deg - i]
206                    * invfact[i]
207                    * invfact[s_deg - i]
208            })
209            .sum::<SC::EF>();
210
211        debug!(round = round + 1, r_round = %r);
212        rs.push(r);
213    }
214
215    // 7. Compute `eq_3b_per_trace`
216    let mut stacked_idx = 0usize;
217    let eq_3b_per_trace = n_per_trace
218        .iter()
219        .enumerate()
220        .map(|(trace_idx, &n)| {
221            let air_idx = trace_id_to_air_id[trace_idx];
222            let interactions = &mvk.per_air[air_idx].symbolic_constraints.interactions;
223            if interactions.is_empty() {
224                return vec![];
225            }
226            // By definition of n_logup, n_lift <= n_logup
227            let n_lift = n.max(0) as usize;
228            let mut b_vec = vec![SC::F::ZERO; n_logup - n_lift];
229            (0..interactions.len())
230                .map(|_| {
231                    debug_assert!(stacked_idx < 1 << (l_skip + n_logup));
232                    debug_assert!(stacked_idx.trailing_zeros() as usize >= l_skip + n_lift);
233                    let mut b_int = stacked_idx >> (l_skip + n_lift);
234                    for b in &mut b_vec {
235                        *b = SC::F::from_bool(b_int & 1 == 1);
236                        b_int >>= 1;
237                    }
238                    stacked_idx += 1 << (l_skip + n_lift);
239                    eval_eq_mle(&xi[l_skip + n_lift..l_skip + n_logup], &b_vec)
240                })
241                .collect_vec()
242        })
243        .collect_vec();
244
245    // 8. Compute `eq_ns` and `eq_sharp_ns`
246    let mut eq_ns = vec![SC::EF::ONE; n_max + 1];
247    let mut eq_sharp_ns = vec![SC::EF::ONE; n_max + 1];
248    eq_ns[0] = eval_eq_uni(l_skip, xi[0], r_0);
249    eq_sharp_ns[0] = eval_eq_sharp_uni(omega_skip_pows, &xi[..l_skip], r_0);
250    debug_assert_eq!(rs.len(), n_max + 1);
251    for (i, r) in rs.iter().enumerate().skip(1) {
252        // xi has length l_skip + n_global >= l_skip + n_max
253        let eq_mle = eval_eq_mle(&[xi[l_skip + i - 1]], slice::from_ref(r));
254        eq_ns[i] = eq_ns[i - 1] * eq_mle;
255        eq_sharp_ns[i] = eq_sharp_ns[i - 1] * eq_mle;
256    }
257    let mut r_rev_prod = rs[n_max];
258    // Product with r_i's to account for \hat{f} vs \tilde{f} for different n's in front-loaded
259    // batch sumcheck.
260    for i in (0..n_max).rev() {
261        eq_ns[i] *= r_rev_prod;
262        eq_sharp_ns[i] *= r_rev_prod;
263        r_rev_prod *= rs[i];
264    }
265
266    // 9. Compute the interaction/constraint evals and their hash
267    let mut interactions_evals = Vec::new(); // len = 2 * num_traces
268    let mut constraints_evals = Vec::new(); // len = num_traces
269    let need_rot_per_trace = trace_id_to_air_id
270        .iter()
271        .map(|&air_idx| mvk.per_air[air_idx].params.need_rot)
272        .collect_vec();
273
274    // Observe common main openings first, and then preprocessed/cached
275    // Proof shape asserts that:
276    // - column_openings.len() == num_traces
277    // - air_openings.len() == vk.num_parts() > 0
278    for (trace_idx, air_openings) in column_openings.iter().enumerate() {
279        let need_rot = need_rot_per_trace[trace_idx];
280        // Proof shape asserts that air_openings[0].len() == width.common_main * (needs_rot ? 2 : 1)
281        for (claim, claim_rot) in column_openings_by_rot(&air_openings[0], need_rot) {
282            transcript.observe_ext(claim);
283            transcript.observe_ext(claim_rot);
284        }
285    }
286
287    for (trace_idx, air_openings) in column_openings.iter().enumerate() {
288        let air_idx = trace_id_to_air_id[trace_idx];
289        let vk = &mvk.per_air[air_idx];
290        let n = n_per_trace[trace_idx];
291        let n_lift = n.max(0) as usize;
292        let need_rot = need_rot_per_trace[trace_idx];
293
294        // claim lengths are checked in proof shape
295        for claims in air_openings.iter().skip(1) {
296            // Proof shape asserts that claims.len() is always multiple of (needs_rot ? 2 : 1)
297            for (claim, claim_rot) in column_openings_by_rot(claims, need_rot) {
298                transcript.observe_ext(claim);
299                transcript.observe_ext(claim_rot);
300            }
301        }
302
303        let has_preprocessed = vk.preprocessed_data.is_some();
304        let common_main = column_openings_by_rot(&air_openings[0], need_rot).collect::<Vec<_>>();
305        let preprocessed = has_preprocessed
306            .then(|| column_openings_by_rot(&air_openings[1], need_rot).collect::<Vec<_>>());
307        let cached_idx = 1 + has_preprocessed as usize;
308        let mut partitioned_main: Vec<_> = air_openings[cached_idx..]
309            .iter()
310            .map(|opening| column_openings_by_rot(opening, need_rot).collect::<Vec<_>>())
311            .collect();
312        partitioned_main.push(common_main);
313        let part_main_slices = partitioned_main
314            .iter()
315            .map(|x| x.as_slice())
316            .collect::<Vec<_>>();
317
318        // We are evaluating the lift, which is the same as evaluating the original with domain
319        // D^{(2^{n})}
320        let (l, rs_n, norm_factor) = if n.is_negative() {
321            (
322                l_skip.wrapping_add_signed(n),
323                &[rs[0].exp_power_of_2(-n as usize)] as &[_],
324                SC::F::from_usize(1 << n.unsigned_abs()).inverse(),
325            )
326        } else {
327            (l_skip, &rs[..=(n as usize)], SC::F::ONE)
328        };
329        let evaluator = VerifierConstraintEvaluator::<SC::F, SC::EF>::new(
330            preprocessed.as_deref(),
331            &part_main_slices,
332            &public_values[air_idx],
333            rs_n,
334            l,
335        );
336
337        let constraints = &vk.symbolic_constraints.constraints;
338        let nodes = evaluator.eval_nodes(&constraints.nodes);
339        let expr = zip(lambda.powers(), &constraints.constraint_idx)
340            .map(|(lambda_pow, idx)| nodes[*idx] * lambda_pow)
341            .sum::<SC::EF>();
342        debug!(%trace_idx, %expr, %air_idx, "constraints_eval");
343        let eq_xi_r = eq_ns[n_lift];
344        debug!(%trace_idx, %eq_xi_r);
345        constraints_evals.push(eq_xi_r * expr);
346
347        let symbolic_constraints = SymbolicConstraints::from(&vk.symbolic_constraints);
348        let interactions = &symbolic_constraints.interactions;
349        let cur_interactions_evals = interactions
350            .iter()
351            .map(|interaction| {
352                let num = evaluator.eval_expr(&interaction.count);
353                let denom = interaction
354                    .message
355                    .iter()
356                    .map(|expr| evaluator.eval_expr(expr))
357                    .chain(std::iter::once(
358                        SC::EF::from_u16(interaction.bus_index) + SC::EF::ONE,
359                    ))
360                    .zip(beta_logup.powers())
361                    .fold(SC::EF::ZERO, |acc, (x, y)| acc + x * y);
362                (num, denom)
363            })
364            .collect_vec();
365        let eq_3bs = &eq_3b_per_trace[trace_idx];
366        let mut num = SC::EF::ZERO;
367        let mut denom = SC::EF::ZERO;
368        for (&eq_3b, (n, d)) in eq_3bs.iter().zip_eq(cur_interactions_evals.iter()) {
369            num += eq_3b * *n;
370            denom += eq_3b * *d;
371        }
372        debug!(%trace_idx, %num, %denom, %air_idx, "interactions_eval");
373        interactions_evals.push(num * norm_factor * eq_sharp_ns[n_lift]);
374        interactions_evals.push(denom * eq_sharp_ns[n_lift]);
375    }
376    let evaluated_claim = interactions_evals
377        .iter()
378        .chain(constraints_evals.iter())
379        .zip(mu.powers())
380        .map(|(x, y)| *x * y)
381        .sum::<SC::EF>();
382    if cur_sum != evaluated_claim {
383        return Err(BatchConstraintError::InconsistentClaims);
384    }
385
386    Ok(rs)
387}