use std::sync::{Arc, OnceLock};
use zisk_common::{HashMode, ProgramVK, ZiskPaths};
use zisk_prover_backend::{CircomCircuit, GuestProgram};
use crate::{Result, SdkError};
#[derive(Clone)]
pub struct Recurser {
pub(crate) recurser_id: String,
pub(crate) templates: zisk_recurser::CircomTemplates,
pub(crate) setup_dir: String,
pub(crate) output_dir: String,
pub(crate) vk_cache: Arc<OnceLock<ProgramVK>>,
}
impl Recurser {
pub fn recurser_id(&self) -> &str {
&self.recurser_id
}
pub fn n_free(&self) -> usize {
self.templates.n_free()
}
pub fn vk(&self) -> Result<ProgramVK> {
if let Some(vk) = self.vk_cache.get() {
return Ok(vk.clone());
}
let artifacts = zisk_recurser::RecurserArtifacts::new(&self.output_dir, &self.recurser_id);
let limbs = artifacts.read_verkey().map_err(|e| {
SdkError::Recurser(format!(
"failed to read recurser verkey ({e}). \
Did `client.setup(&agg).run()` complete?"
))
})?;
let hash_mode = read_setup_hash_mode(&self.setup_dir)?;
let vk = ProgramVK { vk: limbs.to_vec(), hash_mode };
let _ = self.vk_cache.set(vk.clone());
Ok(vk)
}
}
fn read_setup_hash_mode(setup_dir: &str) -> Result<HashMode> {
zisk_recurser::setup::read_proving_key_hash(setup_dir)
.map_err(SdkError::backend)?
.parse::<HashMode>()
.map_err(SdkError::backend)
}
fn expect_template_decl(circuit: &CircomCircuit, template: &str) -> Result<()> {
let needle = format!("template {template}(");
match circuit.source().matches(&needle).count() {
1 => Ok(()),
n => Err(SdkError::Recurser(format!(
"circuit '{}' must define `template {template}(...)` exactly once, found {n}",
circuit.name(),
))),
}
}
fn derive_program_vks(programs: &[&GuestProgram]) -> Result<Vec<[String; 4]>> {
let mut vks: Vec<[String; 4]> = Vec::with_capacity(programs.len());
for prog in programs {
let pvk = prog.vk().map_err(|e| {
SdkError::Recurser(format!("failed to derive VK for program '{}': {e}", prog.name()))
})?;
let limbs: [u64; 4] = <[u64; 4]>::try_from(pvk.vk.as_slice()).map_err(|_| {
SdkError::Recurser(format!(
"program VK for '{}' did not decode into 4 u64 limbs",
prog.name()
))
})?;
let limbs_str: [String; 4] = limbs.map(|w| w.to_string());
if let Some(prior) = vks.iter().position(|existing| existing == &limbs_str) {
return Err(SdkError::Recurser(format!(
"duplicate program VK at index {} ('{}'); already registered at index {}",
vks.len(),
prog.name(),
prior,
)));
}
vks.push(limbs_str);
}
Ok(vks)
}
pub struct AggregationProgramBuilder<'a> {
aggregate: CircomCircuit,
n_free: usize,
n_publics_agg: usize,
normalize: Option<CircomCircuit>,
programs: Vec<&'a GuestProgram>,
}
impl<'a> AggregationProgramBuilder<'a> {
pub fn new(aggregate: impl Into<CircomCircuit>, n_publics_agg: usize) -> Self {
Self {
aggregate: aggregate.into(),
n_free: 0,
n_publics_agg,
normalize: None,
programs: Vec::new(),
}
}
#[must_use]
pub fn programs(mut self, programs: &[&'a GuestProgram]) -> Self {
self.programs = programs.to_vec();
self
}
#[must_use]
pub fn free_inputs(mut self, n_free: usize) -> Self {
self.n_free = n_free;
self
}
#[must_use]
pub fn normalize(mut self, circuit: impl Into<CircomCircuit>) -> Self {
self.normalize = Some(circuit.into());
self
}
pub fn build(self) -> Result<Recurser> {
expect_template_decl(&self.aggregate, "AggregatePublics")?;
if let Some(circuit) = &self.normalize {
expect_template_decl(circuit, "NormalizePublics")?;
}
let n_publics_agg = self.n_publics_agg;
let max_publics = zisk_recurser::templates::ZISK_PUBLICS;
if n_publics_agg == 0 || n_publics_agg > max_publics {
return Err(SdkError::Recurser(format!(
"n_publics_agg must be in 1..={max_publics}, got {n_publics_agg}"
)));
}
let normalize = self
.normalize
.as_ref()
.map(|c| zisk_recurser::NormalizeCircuit { body: c.source().to_string() });
let program_vks = derive_program_vks(&self.programs)?;
let templates = zisk_recurser::CircomTemplates {
normalize: normalize.clone(),
aggregate_publics: self.aggregate.source().to_string(),
n_free: self.n_free,
n_publics_agg,
program_vks: program_vks.clone(),
};
let setup_dir = ZiskPaths::global()
.home
.to_str()
.ok_or_else(|| SdkError::Recurser("default ~/.zisk path is not valid UTF-8".into()))?
.to_string();
let output_dir = ZiskPaths::global()
.home
.join("recurser")
.to_str()
.ok_or_else(|| SdkError::Recurser("~/.zisk/recurser path is not valid UTF-8".into()))?
.to_string();
let zisk_vk = zisk_recurser::setup::read_vadcop_final_verkey(&setup_dir).map_err(|e| {
SdkError::Recurser(format!(
"failed to locate local vadcop_final verkey ({e}). \
Run `cargo-zisk setup --recursive` on this machine \
(required even when using a remote coordinator)."
))
})?;
let inputs = zisk_recurser::RecurserManifestInputs::new(
zisk_vk,
program_vks,
normalize.as_ref(),
&templates.aggregate_publics,
templates.n_free,
templates.n_publics_agg,
);
let recurser_id = inputs.compute_id();
Ok(Recurser {
recurser_id,
templates,
setup_dir,
output_dir,
vk_cache: Arc::new(OnceLock::new()),
})
}
}
pub struct AggregationProgram(std::sync::LazyLock<Recurser>);
impl AggregationProgram {
pub const fn new(init: fn() -> Recurser) -> Self {
Self(std::sync::LazyLock::new(init))
}
}
impl std::ops::Deref for AggregationProgram {
type Target = Recurser;
fn deref(&self) -> &Recurser {
&self.0
}
}
#[macro_export]
macro_rules! load_aggregation_program {
($name:literal) => {{
#[cfg(zisk_skip_guest_build)]
{
$crate::AggregationProgram::new(|| {
panic!(concat!(
"aggregation program `",
$name,
"` is unavailable: the guest build was skipped"
))
})
}
#[cfg(not(zisk_skip_guest_build))]
{
$crate::AggregationProgram::new(|| {
include!(env!(
concat!("ZISK_AGG_", $name),
concat!(
"no aggregation program named `",
$name,
"` was processed by `build_program` — expected \
`<programs>/aggregations/",
$name,
".toml` (after creating the aggregations dir, trigger \
one rebuild, e.g. touch build.rs)"
)
))
.build()
.expect(concat!(
"failed to build aggregation program `",
$name,
"`"
))
})
}
}};
}
#[cfg(test)]
mod tests {
use super::*;
fn dummy_agg() -> Recurser {
Recurser {
recurser_id: "rid".into(),
templates: zisk_recurser::CircomTemplates {
normalize: None,
aggregate_publics: "// body".into(),
n_free: 0,
n_publics_agg: 6,
program_vks: vec![],
},
setup_dir: "/tmp/zisk-test-setup".into(),
output_dir: "/tmp/zisk-test-output".into(),
vk_cache: Arc::new(OnceLock::new()),
}
}
#[test]
fn vk_cache_is_shared_across_clones() {
let agg = dummy_agg();
let agg_clone = agg.clone();
let _ = agg_clone.vk_cache.set(ProgramVK { vk: vec![1, 2, 3, 4], ..Default::default() });
assert_eq!(agg.vk_cache.get().map(|v| v.vk.clone()), Some(vec![1, 2, 3, 4]));
assert_eq!(agg_clone.vk_cache.get().map(|v| v.vk.clone()), Some(vec![1, 2, 3, 4]));
assert!(agg
.vk_cache
.set(ProgramVK { vk: vec![9, 9, 9, 9], ..Default::default() })
.is_err());
assert_eq!(agg.vk_cache.get().map(|v| v.vk.clone()), Some(vec![1, 2, 3, 4]));
}
}