use rayon::prelude::*;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Instant;
use crate::types::*;
use crate::utils::logger::*;
use crate::{
circuits::batch_circuit::BatchCircuit,
circuits::circuit_registry::CircuitRegistry,
circuits::recursive_circuit::RecursiveCircuit,
merkle_tree::{MerkleTree, Node},
utils::util::*,
config::{BATCH_SIZE, RECURSIVE_SIZE, F, C, D},
*,
};
use anyhow::Result;
use plonky2::util::serialization::DefaultGateSerializer;
use plonky2::hash::hash_types::HashOut;
use plonky2::plonk::proof::ProofWithPublicInputs;
use plonky2::plonk::circuit_data::VerifierCircuitData;
use plonky2::plonk::config::GenericHashOut;
use zstd;
fn prove_recursively(
inner_circuit_digest: Option<HashOut<F>>,
asset_count: usize,
mut inner_proofs: Vec<ProofWithPublicInputs<F, C, D>>,
mut merkle_tree: MerkleTree,
mut merkle_depth: Option<usize>,
circuit_registry: &mut CircuitRegistry,
progress: &mut ProveProgress,
) -> (ProofWithPublicInputs<F, C, D>, MerkleTree) {
progress.print_progress_bar();
let inner_circuit;
if let Some(inner_circuit_digest) = inner_circuit_digest {
inner_circuit = &circuit_registry
.get_recursive_circuit(inner_circuit_digest)
.unwrap()
.circuit
.circuit_data;
} else {
inner_circuit = &circuit_registry.get_batch_circuit().circuit_data;
}
if merkle_depth.is_none() {
merkle_depth = Some(merkle_tree.depth - 2); }
let build_circuit_time = Instant::now();
let recursive_circuit = RecursiveCircuit::new(inner_circuit, asset_count);
progress.update_recursive_circuit_progress();
if cfg!(debug_assertions) {
let elapsed = build_circuit_time.elapsed();
progress.clear_bar();
log_warning!(
"Recursive circuit at depth {} build time: {:?}",
merkle_depth.unwrap(),
elapsed
);
progress.print_progress_bar();
}
let empty_proof = circuit_registry
.get_empty_proof(inner_circuit.verifier_only.circuit_digest)
.unwrap();
pad_recursive_proofs(&mut inner_proofs, empty_proof);
let mut count = 0;
for node in merkle_tree.get_nodes_from_depth(merkle_depth.unwrap() + 1) {
if node.hash().is_some() {
count += 1;
continue; }
let hash_offset = RecursiveCircuit::get_root_hash_offset(asset_count);
let hash_elements = inner_proofs[count].public_inputs[hash_offset].to_vec();
let hash_bytes = pis_to_hash_bytes::<F, D>(&hash_elements);
node.set_hash(hash_bytes.clone());
count += 1;
}
let subproofs = inner_proofs.chunks(RECURSIVE_SIZE);
let mut recursive_proofs = Vec::new();
for chunk in subproofs {
let timer = Instant::now();
let proof = recursive_circuit.prove_recursive_circuit(chunk.to_vec());
recursive_proofs.push(proof);
if cfg!(debug_assertions) {
let elapsed = timer.elapsed();
progress.clear_bar();
log_warning!("Recursive proof time: {:?}", elapsed);
progress.print_progress_bar();
}
progress.update_recursive_progress();
}
let inner_circuit_digest = recursive_circuit.circuit_data.verifier_only.circuit_digest;
circuit_registry.add_recursive_circuit(recursive_circuit, merkle_depth.unwrap());
let nodes = &mut merkle_tree.get_nodes_from_depth(merkle_depth.unwrap());
let mut count = 0;
for node in nodes {
if count >= recursive_proofs.len() {
break; }
let hash_offset = RecursiveCircuit::get_root_hash_offset(asset_count);
let hash_elements = recursive_proofs[count].public_inputs[hash_offset].to_vec();
let hash_bytes = pis_to_hash_bytes::<F, D>(&hash_elements);
node.set_hash(hash_bytes.clone());
count += 1;
}
if recursive_proofs.len() > 1 {
prove_recursively(
Some(inner_circuit_digest),
asset_count,
recursive_proofs,
merkle_tree,
Some(merkle_depth.unwrap() - 1),
circuit_registry,
progress,
)
} else {
(recursive_proofs[0].clone(), merkle_tree)
}
}
pub fn prove_global(mut ledger: Ledger) -> Result<(FinalProof, MerkleTree, Vec<u64>)> {
let asset_count = ledger.asset_names.len();
pad_accounts(
&mut ledger.account_balances,
&mut ledger.hashes,
asset_count,
BATCH_SIZE,
)?;
let mut progress = ProveProgress::new(ledger.account_balances.len() / BATCH_SIZE);
log_info!("Creating batch circuit and proving all accounts...");
progress.print_progress_bar();
let batch_circuit = BatchCircuit::new(asset_count);
let mut batch_proofs = Vec::new();
let mut merkle_leafs = Vec::new();
let mut account_nonces = Vec::new();
let mut count = 0;
for chunk in ledger.account_balances.chunks(BATCH_SIZE) {
let circuit_ref = &batch_circuit;
let batch_time = Instant::now();
let mut leaf_hashes = Vec::new();
for i in 0..chunk.len() {
let userhash = ledger.hashes[count * BATCH_SIZE + i].clone();
let balances = chunk[i].clone();
let nonce = rand::random::<u64>();
account_nonces.push(nonce);
let hash = hash_account(&balances, userhash, nonce);
leaf_hashes.push(hash);
}
let proof = circuit_ref
.prove_batch_circuit(&ledger.asset_prices, chunk, &leaf_hashes)
.unwrap();
merkle_leafs.push(leaf_hashes);
progress.update_batch_progress();
if cfg!(debug_assertions) {
let elapsed = batch_time.elapsed();
progress.clear_bar();
log_warning!("Batch {} took {:?}", count, elapsed);
progress.print_progress_bar();
}
batch_proofs.push(proof);
count += 1;
}
progress.clear_bar(); log_success!("Proved all batch circuits successfully!");
progress.print_progress_bar();
let mut leaf_nodes = Vec::new();
for leaf_hashes in merkle_leafs {
for hash in leaf_hashes {
let node = Node::new(Some(hash.to_bytes()));
leaf_nodes.push(node);
}
}
let mut merkle_tree = MerkleTree::new_from_leafs(leaf_nodes, 1, true);
let batch_circuit_digest = batch_circuit.circuit_data.verifier_only.circuit_digest;
let mut circuit_registry = CircuitRegistry::new(batch_circuit, &ledger.asset_prices);
let batch_nodes = merkle_tree.get_nodes_from_depth(merkle_tree.depth - 1);
let mut count = 0;
let batch_proofs_length = batch_proofs.len();
for node in batch_nodes {
let proof = {
if count >= batch_proofs_length {
circuit_registry
.get_empty_proof(batch_circuit_digest)
.unwrap()
} else {
&batch_proofs[count]
}
};
let hash_offset = BatchCircuit::get_root_hash_offset(asset_count);
let hash_elements = proof.public_inputs[hash_offset.clone()].to_vec();
let hash_bytes = pis_to_hash_bytes::<F, D>(&hash_elements);
node.set_hash(hash_bytes.clone());
count += 1;
}
progress.clear_bar();
log_success!(
"Created merkle tree structure with {} levels (1 accounts, 1 batch, {} recursive)",
merkle_tree.depth,
merkle_tree.depth - 2
);
progress.print_progress_bar();
progress.clear_bar();
log_info!("Starting the recursive proving...");
progress.print_progress_bar();
let (root_proof, merkle_tree) = prove_recursively(
None,
asset_count,
batch_proofs,
merkle_tree,
None,
&mut circuit_registry,
&mut progress,
);
progress.clear_bar();
log_success!("Proved all recursive circuits successfully!");
log_info!("Creating final proof...");
let asset_prices = ledger.asset_prices;
let root_circuit_verifier_data: VerifierCircuitData<F, C, D> = circuit_registry
.get_recursive_circuit_by_depth(1)
.unwrap()
.circuit
.circuit_data
.verifier_data()
.clone();
let final_proof = FinalProof {
proof: root_proof,
batch_size: BATCH_SIZE,
recursive_size: RECURSIVE_SIZE,
asset_prices: asset_prices.clone(),
asset_names: ledger.asset_names.clone(),
asset_decimals: ledger.asset_decimals.clone(),
tree_depth: merkle_tree.depth,
root_circuit_verifier_data: root_circuit_verifier_data
.to_bytes(&DefaultGateSerializer)
.unwrap(),
timestamp: ledger.timestamp,
prover_version: format!("v{}", env!("CARGO_PKG_VERSION")),
};
log_success!("Created final proof successfully!");
Ok((final_proof, merkle_tree, account_nonces))
}
pub fn prove_user_inclusion(
user_index: usize,
user_hash: String,
nonce: u64,
merkle_tree: &MerkleTree,
ledger: &Ledger,
) -> Result<InclusionProof> {
let user_balances = ledger.account_balances[user_index].clone();
let user_node_path = merkle_tree.get_nth_leaf_path(user_index).unwrap();
let merkle_proof = merkle_tree.prove_inclusion(user_node_path);
let inclusion_proof = InclusionProof {
user_hash,
user_balances: user_balances.clone(),
merkle_proof,
root_hash: merkle_tree.root.hash().clone().unwrap(),
nonce,
};
Ok(inclusion_proof)
}
pub fn prove_user_inclusion_by_hash(
user_hash: String,
merkle_tree: &MerkleTree,
nonces: &[u64],
ledger: &Ledger,
) -> Result<InclusionProof> {
let user_index = ledger.hashes.iter().position(|x| *x == user_hash);
if user_index.is_none() {
return Err(anyhow::anyhow!("User hash not found in ledger"));
}
let user_index = user_index.unwrap();
let user_nonce = nonces[user_index];
prove_user_inclusion(user_index, user_hash, user_nonce, merkle_tree, ledger)
}
pub fn prove_inclusion_all_batched(
ledger: &Ledger,
merkle_tree: &MerkleTree,
nonces: Vec<u64>,
) -> Result<()> {
let total_hashes = ledger.hashes.len();
let num_cpus = rayon::current_num_threads();
log_info!(
"Processing {} hashes in batches grouped by first 3 characters using {} threads...",
total_hashes,
num_cpus
);
let mut groups: HashMap<String, Vec<(usize, &String)>> = HashMap::new();
for (index, userhash) in ledger.hashes.iter().enumerate() {
let prefix = userhash.chars().take(3).collect::<String>();
groups
.entry(prefix)
.or_insert_with(Vec::new)
.push((index, userhash));
}
let total_groups = groups.len();
let processed_groups = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let processed_hashes = Arc::new(std::sync::atomic::AtomicUsize::new(0));
log_info!(
"Created {} groups based on first 3 characters",
total_groups
);
std::fs::create_dir_all("inclusion_proofs")?;
let processing_result: Result<()> = groups.par_iter().try_for_each(
|(prefix, group)| -> Result<()> {
let group_result: Result<HashMap<String, InclusionProof>> = group
.par_iter()
.map(|(index, userhash)| -> Result<(String, InclusionProof)> {
let inclusion_proof = prove_user_inclusion(
*index,
(*userhash).clone(),
nonces[*index],
merkle_tree,
ledger,
)?;
Ok(((*userhash).clone(), inclusion_proof))
})
.collect();
match group_result {
Ok(inclusion_proofs_map) => {
let bundle_filename =
format!("inclusion_proofs/inclusion_proofs_{prefix}.json.zst");
let bundle_json = serde_json::to_string(&inclusion_proofs_map)?;
let compressed_data = zstd::encode_all(bundle_json.as_bytes(), 3)?; std::fs::write(&bundle_filename, compressed_data)?;
let completed_groups =
processed_groups.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
let completed_hashes = processed_hashes
.fetch_add(group.len(), std::sync::atomic::Ordering::Relaxed)
+ group.len();
if completed_groups % 10 == 0 || completed_groups == total_groups {
log_success!(
"Completed group '{}' ({}/{} groups, {}/{} hashes) - compressed to {}",
prefix,
completed_groups,
total_groups,
completed_hashes,
total_hashes,
bundle_filename
);
}
Ok(())
}
Err(e) => Err(e),
}
},
);
processing_result?;
log_success!(
"Successfully processed all {} groups with {} total inclusion proofs!",
total_groups,
total_hashes
);
Ok(())
}
pub fn prove_inclusion_all(
ledger: &Ledger,
merkle_tree: &MerkleTree,
nonces: Vec<u64>,
) -> Result<()> {
let total_hashes = ledger.hashes.len();
let progress = Arc::new(Mutex::new(ProveInclusionProgress::new(total_hashes)));
{
let prog = progress.lock().unwrap(); prog.print_progress_bar();
}
let processing_result: Result<()> = ledger
.hashes
.par_iter() .enumerate()
.try_for_each(|(index, userhash)| {
let inclusion_proof =
prove_user_inclusion(index, userhash.clone(), nonces[index], merkle_tree, ledger)?;
let inclusion_filename = format!("inclusion_proofs/inclusion_proof_{userhash}.json");
let inclusion_proof_json = serde_json::to_string(&inclusion_proof)?; std::fs::write(inclusion_filename, inclusion_proof_json)?;
{
let mut prog = progress.lock().unwrap(); prog.update_progress(1);
}
Ok(())
});
{
let prog = progress.lock().unwrap(); prog.clear_bar();
}
processing_result
}