Skip to main content

openvm_stark_backend/verifier/
mod.rs

1use core::cmp::Reverse;
2
3use itertools::{izip, Itertools};
4use p3_field::{PrimeCharacteristicRing, TwoAdicField};
5use thiserror::Error;
6
7use crate::{
8    keygen::types::{MultiStarkVerifyingKey, MultiStarkVerifyingKey0},
9    poly_common::Squarable,
10    proof::Proof,
11    verifier::{
12        batch_constraints::{verify_zerocheck_and_logup, BatchConstraintError},
13        proof_shape::{verify_proof_shape, ProofShapeError},
14        stacked_reduction::{verify_stacked_reduction, StackedReductionError},
15        whir::{verify_whir, VerifyWhirError},
16    },
17    FiatShamirTranscript, StarkProtocolConfig,
18};
19
20#[derive(Error, Debug, PartialEq, Eq)]
21pub enum VerifierError<EF: core::fmt::Debug + core::fmt::Display + PartialEq + Eq> {
22    #[error("Protocol and VerifyingKey has mismatch in SystemParams")]
23    SystemParamsMismatch,
24
25    #[error("Trace heights are too large")]
26    TraceHeightsTooLarge,
27
28    #[error("Preprocessed trace height does not match verifier trace data")]
29    PreprocessedTraceHeightMismatch,
30
31    /// A proof without any traces is always considered invalid.
32    #[error("Proof has no traces")]
33    EmptyTraces,
34
35    #[error("Proof shape verification failed: {0}")]
36    ProofShapeError(#[from] ProofShapeError),
37
38    #[error("Batch constraint verification failed: {0}")]
39    BatchConstraintError(#[from] BatchConstraintError<EF>),
40
41    #[error("Stacked reduction verification failed: {0}")]
42    StackedReductionError(#[from] StackedReductionError<EF>),
43
44    #[error("Whir verification failed: {0}")]
45    WhirError(#[from] VerifyWhirError),
46}
47
48pub mod batch_constraints;
49pub mod evaluator;
50pub mod fractional_sumcheck_gkr;
51pub mod proof_shape;
52pub mod stacked_reduction;
53#[cfg(test)]
54mod transcript_extractor;
55pub mod whir;
56
57pub fn verify<SC: StarkProtocolConfig, TS: FiatShamirTranscript<SC>>(
58    config: &SC,
59    mvk: &MultiStarkVerifyingKey<SC>,
60    proof: &Proof<SC>,
61    transcript: &mut TS,
62) -> Result<(), VerifierError<SC::EF>>
63where
64    SC::EF: p3_field::TwoAdicField,
65{
66    if config.params() != &mvk.inner.params {
67        return Err(VerifierError::SystemParamsMismatch);
68    }
69    let &Proof {
70        common_main_commit,
71        trace_vdata,
72        public_values,
73        gkr_proof,
74        batch_constraint_proof,
75        stacking_proof,
76        whir_proof,
77    } = &proof;
78    let &MultiStarkVerifyingKey {
79        inner: mvk,
80        pre_hash: mvk_pre_hash,
81    } = &mvk;
82    let &MultiStarkVerifyingKey0 {
83        params,
84        per_air,
85        trace_height_constraints,
86    } = &mvk;
87    let l_skip = params.l_skip;
88    // Total number of AIRs is in vkey, committed to in vk_pre_hash
89    let num_airs = per_air.len();
90    // Number of present traces. This is implicitly hashed via the
91    // transcript.observe(is_air_present) below.
92    let num_traces = trace_vdata.iter().flatten().collect_vec().len();
93    if num_traces == 0 {
94        return Err(VerifierError::EmptyTraces);
95    }
96    // We verify the proof shape early to return error and prevent later panics
97    let layouts = verify_proof_shape::<SC>(mvk, proof)?;
98
99    let mut trace_id_to_air_id: Vec<usize> = (0..num_airs).collect();
100    trace_id_to_air_id.sort_by_key(|&air_id| {
101        (
102            trace_vdata[air_id].is_none(),
103            trace_vdata[air_id]
104                .as_ref()
105                .map(|vdata| Reverse(vdata.log_height)),
106            air_id,
107        )
108    });
109    trace_id_to_air_id.truncate(num_traces);
110
111    for constraint in trace_height_constraints {
112        let sum = trace_id_to_air_id
113            .iter()
114            .map(|&air_id| {
115                let log_height = trace_vdata[air_id].as_ref().unwrap().log_height;
116                // Proof shape will check n <= n_stack is in bounds
117                (1 << log_height.max(l_skip)) as u64 * constraint.coefficients[air_id] as u64
118            })
119            .sum::<u64>();
120        if sum >= constraint.threshold as u64 {
121            return Err(VerifierError::TraceHeightsTooLarge);
122        }
123    }
124
125    let omega_skip = SC::F::two_adic_generator(l_skip);
126    let omega_skip_pows = omega_skip.powers().take(1 << l_skip).collect_vec();
127
128    // Preamble
129    transcript.observe_commit(*mvk_pre_hash);
130    transcript.observe_commit(proof.common_main_commit);
131
132    for (trace_vdata, avk, pvs) in izip!(&proof.trace_vdata, per_air, &proof.public_values) {
133        let is_air_present = trace_vdata.is_some();
134
135        // Proof shape asserts that vk.is_required => is_air_present (see
136        // ProofShapeVDataError::RequiredAirNoVData)
137        if !avk.is_required {
138            transcript.observe(SC::F::from_bool(is_air_present));
139        }
140        if let Some(trace_vdata) = trace_vdata {
141            if let Some(pdata) = avk.preprocessed_data.as_ref() {
142                if (pdata.hypercube_dim + l_skip as isize) as usize != trace_vdata.log_height {
143                    return Err(VerifierError::PreprocessedTraceHeightMismatch);
144                }
145                transcript.observe_commit(pdata.commit);
146            } else {
147                transcript.observe(SC::F::from_usize(trace_vdata.log_height));
148            }
149            debug_assert_eq!(
150                avk.params.width.cached_mains.len(),
151                trace_vdata.cached_commitments.len()
152            );
153            for commit in &trace_vdata.cached_commitments {
154                transcript.observe_commit(*commit);
155            }
156            debug_assert_eq!(avk.params.num_public_values, pvs.len());
157        }
158        for pv in pvs {
159            transcript.observe(*pv);
160        }
161    }
162
163    // n_per_trace.len() = num_traces
164    let n_per_trace: Vec<isize> = trace_id_to_air_id
165        .iter()
166        .map(|&air_id| trace_vdata[air_id].as_ref().unwrap().log_height as isize - l_skip as isize)
167        .collect();
168    let r = verify_zerocheck_and_logup::<SC, TS>(
169        transcript,
170        mvk,
171        public_values,
172        gkr_proof,
173        batch_constraint_proof,
174        &trace_id_to_air_id,
175        &n_per_trace,
176        &omega_skip_pows,
177    )?;
178
179    let need_rot_per_trace = trace_id_to_air_id
180        .iter()
181        .map(|&air_id| per_air[air_id].params.need_rot)
182        .collect_vec();
183    let mut need_rot_per_commit = vec![need_rot_per_trace];
184    for &air_id in &trace_id_to_air_id {
185        let need_rot = per_air[air_id].params.need_rot;
186        if per_air[air_id].preprocessed_data.is_some() {
187            need_rot_per_commit.push(vec![need_rot]);
188        }
189        let cached_len = trace_vdata[air_id]
190            .as_ref()
191            .unwrap()
192            .cached_commitments
193            .len();
194        for _ in 0..cached_len {
195            need_rot_per_commit.push(vec![need_rot]);
196        }
197    }
198
199    let u_prism = verify_stacked_reduction::<SC, TS>(
200        transcript,
201        stacking_proof,
202        &layouts,
203        &need_rot_per_commit,
204        l_skip,
205        params.n_stack,
206        &proof.batch_constraint_proof.column_openings,
207        &r,
208        &omega_skip_pows,
209    )?;
210
211    let (&u0, u_rest) = u_prism.split_first().unwrap();
212    let u_cube = u0
213        .exp_powers_of_2()
214        .take(l_skip)
215        .chain(u_rest.iter().copied())
216        .collect_vec();
217
218    let mut commits = vec![*common_main_commit];
219    for &air_id in trace_id_to_air_id.iter() {
220        if let Some(preprocessed) = &per_air[air_id].preprocessed_data {
221            commits.push(preprocessed.commit);
222        }
223        commits.extend(&trace_vdata[air_id].as_ref().unwrap().cached_commitments);
224    }
225
226    verify_whir::<SC, TS>(
227        transcript,
228        config,
229        whir_proof,
230        &stacking_proof.stacking_openings,
231        &commits,
232        &u_cube,
233    )?;
234
235    Ok(())
236}