1use std::sync::Arc;
2
3use halo2_base::{
4 gates::circuit::{BaseCircuitParams, CircuitBuilderStage},
5 halo2_proofs::{
6 halo2curves::bn256::G1Affine,
7 plonk::keygen_pk2,
8 poly::{
9 commitment::{CommitmentScheme, Params},
10 kzg::commitment::{KZGCommitmentScheme, ParamsKZG},
11 },
12 },
13};
14use itertools::Itertools;
15use once_cell::sync::Lazy;
16use serde::{Deserialize, Serialize};
17use serde_with::serde_as;
18#[cfg(feature = "evm-prove")]
19use snark_verifier_sdk::snark_verifier::{
20 halo2_base::halo2_proofs::plonk::VerifyingKey, loader::evm::compile_solidity,
21};
22use snark_verifier_sdk::{
23 halo2::aggregation::{AggregationCircuit, AggregationConfigParams, VerifierUniversality},
24 CircuitExt, Snark, SHPLONK,
25};
26
27use crate::{
28 keygen::RawEvmProof,
29 prover::{Halo2Params, Halo2ProvingMetadata, Halo2ProvingPinning},
30};
31
32static SVK: Lazy<G1Affine> = Lazy::new(|| {
35 serde_json::from_str("\"0100000000000000000000000000000000000000000000000000000000000000\"")
36 .unwrap()
37});
38
39static FAKE_KZG_PARAMS: Lazy<Halo2Params> = Lazy::new(|| KZGCommitmentScheme::new_params(1));
42
43pub static KZG_PARAMS_FOR_SVK: Lazy<Halo2Params> = Lazy::new(|| {
44 if std::env::var("RANDOM_SRS").is_ok() {
45 use rand_chacha::{rand_core::SeedableRng, ChaCha20Rng};
47 let mut rng = ChaCha20Rng::seed_from_u64(42);
48 let mut params = ParamsKZG::setup(23, &mut rng);
49 params.downsize(1);
50 params
51 } else {
52 build_kzg_params_for_svk(*SVK)
53 }
54});
55
56fn build_kzg_params_for_svk(g: G1Affine) -> Halo2Params {
57 FAKE_KZG_PARAMS.from_parts(
58 1,
59 vec![g],
60 Some(vec![g]),
61 Default::default(),
62 Default::default(),
63 )
64}
65
66pub trait Halo2ParamsReader {
70 fn read_params(&self, k: usize) -> Arc<Halo2Params>;
71}
72
73#[derive(Debug, Clone, Serialize, Deserialize)]
78pub struct FallbackEvmVerifier {
79 pub sol_code: String,
80 pub artifact: EvmVerifierByteCode,
81}
82
83#[serde_as]
85#[derive(Clone, Debug, Deserialize, Serialize)]
86pub struct EvmVerifierByteCode {
87 pub sol_compiler_version: String,
88 pub sol_compiler_options: String,
89 #[serde_as(as = "serde_with::hex::Hex")]
90 pub bytecode: Vec<u8>,
91}
92
93#[derive(Debug, Clone)]
94pub struct Halo2WrapperProvingKey {
95 pub pinning: Halo2ProvingPinning,
96}
97
98const MIN_ROWS: usize = 20;
99
100fn has_single_advice_column(config_params: &BaseCircuitParams) -> bool {
101 config_params.num_advice_per_phase.as_slice() == [1]
102}
103
104fn assert_single_advice_column(config_params: &BaseCircuitParams) {
105 assert!(
106 has_single_advice_column(config_params),
107 "OpenVM EVM wrapper requires exactly one advice column, got {:?}; increase wrapper_k or leave wrapper_k unset to auto-tune",
108 config_params.num_advice_per_phase
109 );
110}
111
112fn assert_single_instance_column_count(instance_columns: usize) {
113 assert_eq!(
114 instance_columns, 1,
115 "OpenVM EVM wrapper requires exactly one instance column, got {instance_columns}"
116 );
117}
118
119fn assert_single_instance_column_snark(snark: &Snark) {
120 assert_single_instance_column_count(snark.instances.len());
121}
122
123impl Halo2WrapperProvingKey {
124 pub fn keygen_auto_tune(reader: &impl Halo2ParamsReader, dummy_snark: Snark) -> Self {
126 let k = Self::select_k(dummy_snark.clone());
127 tracing::info!("Selected wrapper k: {k}");
128 let params = reader.read_params(k);
129 Self::keygen(¶ms, dummy_snark)
130 }
131
132 pub fn keygen(params: &Halo2Params, dummy_snark: Snark) -> Self {
133 assert_single_instance_column_snark(&dummy_snark);
134
135 let k = params.k();
136 let mut circuit =
137 generate_wrapper_circuit_object(CircuitBuilderStage::Keygen, k as usize, dummy_snark);
138 circuit.calculate_params(Some(MIN_ROWS));
139 let config_params = circuit.builder.config_params.clone();
140 tracing::info!(
141 "Wrapper circuit num advice: {:?}",
142 config_params.num_advice_per_phase
143 );
144 assert_single_advice_column(&config_params);
145 let pk = keygen_pk2(params, &circuit, false).unwrap();
146 let num_pvs = circuit.instances().iter().map(|x| x.len()).collect_vec();
147 Self {
148 pinning: Halo2ProvingPinning {
149 pk,
150 metadata: Halo2ProvingMetadata {
151 config_params,
152 break_points: circuit.break_points(),
153 num_pvs,
154 },
155 },
156 }
157 }
158
159 #[cfg(feature = "evm-verify")]
160 pub fn evm_verify(
162 evm_verifier: &FallbackEvmVerifier,
163 evm_proof: &RawEvmProof,
164 ) -> Result<u64, String> {
165 snark_verifier_sdk::evm::evm_verify(
166 evm_verifier.artifact.bytecode.clone(),
167 vec![evm_proof.instances.clone()],
168 evm_proof.proof.clone(),
169 )
170 }
171
172 #[cfg(feature = "evm-prove")]
173 pub fn generate_fallback_evm_verifier(&self, params: &Halo2Params) -> FallbackEvmVerifier {
175 assert_single_advice_column(&self.pinning.metadata.config_params);
176 assert_single_instance_column_count(self.pinning.metadata.num_pvs.len());
177 assert_eq!(
178 self.pinning.metadata.config_params.k as u32,
179 params.k(),
180 "Provided params don't match circuit config"
181 );
182 gen_evm_verifier(
183 params,
184 self.pinning.pk.get_vk(),
185 self.pinning.metadata.num_pvs.clone(),
186 )
187 }
188
189 #[cfg(feature = "evm-prove")]
190 pub fn prove_for_evm(&self, params: &Halo2Params, snark_to_verify: Snark) -> RawEvmProof {
191 let k = self.pinning.metadata.config_params.k;
192 let prover_circuit = self.generate_circuit_object_for_proving(k, snark_to_verify);
193 let mut pvs = prover_circuit.instances();
194 assert_single_instance_column_count(pvs.len());
195 let proof = snark_verifier_sdk::evm::gen_evm_proof_shplonk(
196 params,
197 &self.pinning.pk,
198 prover_circuit,
199 pvs.clone(),
200 );
201
202 RawEvmProof {
203 instances: pvs.pop().unwrap(),
204 proof,
205 }
206 }
207
208 #[cfg(feature = "evm-prove")]
209 fn generate_circuit_object_for_proving(
210 &self,
211 k: usize,
212 snark_to_verify: Snark,
213 ) -> AggregationCircuit {
214 assert_single_instance_column_snark(&snark_to_verify);
215 assert_single_instance_column_count(self.pinning.metadata.num_pvs.len());
216 assert_eq!(
217 self.pinning.metadata.num_pvs[0],
218 snark_to_verify.instances[0].len() + 12,
219 );
220 assert_single_advice_column(&self.pinning.metadata.config_params);
221 generate_wrapper_circuit_object(CircuitBuilderStage::Prover, k, snark_to_verify)
222 .use_params(
223 self.pinning
224 .metadata
225 .config_params
226 .clone()
227 .try_into()
228 .unwrap(),
229 )
230 .use_break_points(self.pinning.metadata.break_points.clone())
231 }
232
233 pub(crate) fn select_k(dummy_snark: Snark) -> usize {
234 let mut k = 20;
235 let mut first_run = true;
236 loop {
237 let mut circuit = generate_wrapper_circuit_object(
238 CircuitBuilderStage::Keygen,
239 k,
240 dummy_snark.clone(),
241 );
242 circuit.calculate_params(Some(MIN_ROWS));
243 assert_eq!(
244 circuit.builder.config_params.num_advice_per_phase.len(),
245 1,
246 "Snark has multiple phases"
247 );
248 if has_single_advice_column(&circuit.builder.config_params) {
249 circuit.builder.clear();
250 break;
251 }
252 if first_run {
253 k = log2_ceil_usize(
254 circuit.builder.statistics().gate.total_advice_per_phase[0] + MIN_ROWS,
255 );
256 } else {
257 k += 1;
258 }
259 first_run = false;
260 circuit.builder.clear();
262 }
263 k
264 }
265}
266
267fn generate_wrapper_circuit_object(
268 stage: CircuitBuilderStage,
269 k: usize,
270 snark: Snark,
271) -> AggregationCircuit {
272 let config_params = AggregationConfigParams {
273 degree: k as u32,
274 lookup_bits: k - 1,
275 ..Default::default()
276 };
277 let mut circuit = AggregationCircuit::new::<SHPLONK>(
278 stage,
279 config_params,
280 &KZG_PARAMS_FOR_SVK,
281 [snark],
282 VerifierUniversality::None,
283 );
284 circuit.expose_previous_instances(false);
285 circuit
286}
287
288#[cfg(feature = "evm-prove")]
289fn gen_evm_verifier(
290 params: &Halo2Params,
291 vk: &VerifyingKey<G1Affine>,
292 num_instance: Vec<usize>,
293) -> FallbackEvmVerifier {
294 let sol_code = snark_verifier_sdk::evm::gen_evm_verifier_sol_code::<AggregationCircuit, SHPLONK>(
295 params,
296 vk,
297 num_instance,
298 );
299 let byte_code = compile_solidity(&sol_code);
300 FallbackEvmVerifier {
301 sol_code,
302 artifact: EvmVerifierByteCode {
303 sol_compiler_version: "0.8.19".to_string(),
304 sol_compiler_options: "".to_string(),
305 bytecode: byte_code,
306 },
307 }
308}
309
310fn log2_ceil_usize(n: usize) -> usize {
312 assert!(n > 0);
313 (usize::BITS - (n - 1).leading_zeros()) as usize
314}