openvm_stark_backend/verifier/
batch_constraints.rs1use 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#[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 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 if !transcript.check_witness(mvk.params.logup.pow_bits, gkr_proof.logup_pow_witness) {
76 return Err(BatchConstraintError::InvalidLogupPowWitness);
77 }
78
79 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 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 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 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 #[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 let batch_s_evals = &sumcheck_round_polys[round];
180 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 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 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 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 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 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 let mut interactions_evals = Vec::new(); let mut constraints_evals = Vec::new(); 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 for (trace_idx, air_openings) in column_openings.iter().enumerate() {
279 let need_rot = need_rot_per_trace[trace_idx];
280 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 for claims in air_openings.iter().skip(1) {
296 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 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}