use crate::config::*;
use crate::circuits::circuit_registry::*;
use plonky2::iop::target::Target;
use plonky2::iop::witness::{PartialWitness, WitnessWrite};
use plonky2::plonk::circuit_builder::CircuitBuilder;
use plonky2::plonk::circuit_data::VerifierOnlyCircuitData;
use plonky2::plonk::circuit_data::{CircuitData, VerifierCircuitTarget};
use plonky2::plonk::proof::{ProofWithPublicInputs, ProofWithPublicInputsTarget};
use plonky2::plonk::prover::prove;
use plonky2::util::serialization::gate_serialization::log::Level;
use plonky2::util::timing::TimingTree;
use crate::utils::circuit_helper::*;
#[derive(Debug)]
pub struct RecursiveCircuit {
pub circuit_data: CircuitData<F, C, D>,
inner_circuit_data_verifier: VerifierOnlyCircuitData<C, D>,
inner_circuit_targets: Vec<InnerCircuitTargets>,
}
#[derive(Debug)]
struct InnerCircuitTargets {
proof_target: ProofWithPublicInputsTarget<D>,
verifier_target: VerifierCircuitTarget,
asset_balances: Vec<Target>,
}
impl RecursiveCircuit {
pub fn new(
inner_circuit: &CircuitData<F, C, D>,
asset_count: usize,
) -> RecursiveCircuit {
let config = RECURSIVE_CIRCUIT_CONFIG;
let mut builder = CircuitBuilder::<F, D>::new(config);
let mut inner_targets = Vec::new();
for _ in 0..RECURSIVE_SIZE {
let proof_target = builder.add_virtual_proof_with_pis(&inner_circuit.common);
let verify_target = builder
.add_virtual_verifier_data(inner_circuit.common.config.fri_config.cap_height);
let balances_offset = RecursiveCircuit::get_final_balances_offset(asset_count);
let batch_balance = proof_target.public_inputs[balances_offset].to_vec();
let inner_data = InnerCircuitTargets {
proof_target,
verifier_target: verify_target,
asset_balances: batch_balance,
};
builder.verify_proof::<C>(
&inner_data.proof_target,
&inner_data.verifier_target,
&inner_circuit.common,
);
inner_targets.push(inner_data);
}
let mut final_balances = Vec::new();
for i in 0..asset_count {
final_balances.push(builder.zero());
for inner_data in &inner_targets {
let new_summed_bal = builder.add(inner_data.asset_balances[i], final_balances[i]);
let is_sum1_positive = is_positive(&mut builder, inner_data.asset_balances[i]);
let is_sum2_positive = is_positive(&mut builder, final_balances[i]);
let is_both_positive = builder.and(is_sum1_positive, is_sum2_positive);
let is_result_negative = is_negative(&mut builder, new_summed_bal);
let is_overflow = builder.and(is_both_positive, is_result_negative);
let is_not_overflow = builder.not(is_overflow);
builder.assert_bool(is_not_overflow);
final_balances[i] = new_summed_bal;
}
}
let asset_prices = inner_targets[0].proof_target.public_inputs
[RecursiveCircuit::get_asset_prices_offset(asset_count)]
.to_vec();
for inner_data in inner_targets.iter().take(RECURSIVE_SIZE) {
let inner_asset_prices = inner_data.proof_target.public_inputs
[RecursiveCircuit::get_asset_prices_offset(asset_count)]
.to_vec();
for (j, price) in inner_asset_prices.iter().enumerate().take(asset_count) {
builder.connect(*price, asset_prices[j]);
}
}
let mut concat_hashes = Vec::new();
for inner_data in inner_targets.iter().take(RECURSIVE_SIZE) {
let hash_elements = inner_data.proof_target.public_inputs
[RecursiveCircuit::get_root_hash_offset(asset_count)]
.to_vec();
concat_hashes.extend(hash_elements);
}
let root_hash = builder.hash_n_to_hash_no_pad::<H>(concat_hashes);
builder.register_public_inputs(&final_balances); builder.register_public_inputs(&asset_prices); builder.register_public_inputs(&root_hash.elements);
RecursiveCircuit {
inner_circuit_data_verifier: inner_circuit.verifier_only.clone(),
circuit_data: builder.build::<C>(),
inner_circuit_targets: inner_targets,
}
}
pub fn prove_recursive_circuit(
&self,
subproofs: Vec<ProofWithPublicInputs<F, C, D>>,
) -> ProofWithPublicInputs<F, C, D> {
let mut pw = PartialWitness::new();
for (i, inner_data) in self.inner_circuit_targets.iter().enumerate() {
pw.set_proof_with_pis_target(&inner_data.proof_target, &subproofs[i]).unwrap();
pw.set_verifier_data_target(&inner_data.verifier_target, &self.inner_circuit_data_verifier).unwrap();
}
let mut timing = TimingTree::new("prove recursive", Level::Trace);
let proof = prove::<F, C, D>(
&self.circuit_data.prover_only,
&self.circuit_data.common,
pw,
&mut timing,
)
.unwrap();
timing.print();
proof
}
pub fn prove_empty(
&self,
circuit_registry: &mut CircuitRegistry,
) -> ProofWithPublicInputs<F, C, D> {
let current_digest = self.circuit_data.verifier_only.circuit_digest;
let cached_proof = circuit_registry.get_empty_proof(current_digest);
if let Some(proof) = cached_proof {
return proof.clone();
}
let inner_digest = self.inner_circuit_data_verifier.circuit_digest;
let inner_empty_proof = circuit_registry.get_empty_proof(inner_digest).unwrap();
self.prove_recursive_circuit(vec![inner_empty_proof.clone(); RECURSIVE_SIZE])
}
pub fn get_final_balances_offset(asset_count: usize) -> std::ops::Range<usize> {
let start = 0;
let end = asset_count;
start..end
}
pub fn get_asset_prices_offset(asset_count: usize) -> std::ops::Range<usize> {
let start = asset_count;
let end = asset_count*2;
start..end
}
pub fn get_root_hash_offset(asset_count: usize) -> std::ops::Range<usize> {
let start = asset_count*2;
let end = start + 4;
start..end
}
}