use std::sync::Arc;
use std::time::Duration;
use zisk_common::Proof;
use crate::job_handle::{subscriber_list_from, JobHandle, Subscriber};
use crate::prove::{JobEvent, ProveResult};
use crate::recurser::Recurser;
use crate::{Client, Result, SdkError};
pub struct AggregationInput<'a> {
pub(crate) proof: &'a Proof,
pub(crate) free_inputs: Vec<u64>,
}
impl<'a> From<&'a Proof> for AggregationInput<'a> {
fn from(proof: &'a Proof) -> Self {
Self { proof, free_inputs: Vec::new() }
}
}
pub trait ProofExt {
fn with_free_inputs(&self, inputs: impl Into<Vec<u64>>) -> AggregationInput<'_>;
}
impl ProofExt for Proof {
fn with_free_inputs(&self, inputs: impl Into<Vec<u64>>) -> AggregationInput<'_> {
AggregationInput { proof: self, free_inputs: inputs.into() }
}
}
pub struct AggregateProofsRequest<'a, C> {
client: &'a C,
agg: &'a Recurser,
input_a: AggregationInput<'a>,
input_b: AggregationInput<'a>,
root_c_recurser_agg: Option<[u64; 4]>,
timeout: Option<Duration>,
subscribers: Vec<Subscriber>,
}
#[allow(private_bounds)]
impl<'a, C: Client> AggregateProofsRequest<'a, C> {
pub(crate) fn new(
client: &'a C,
agg: &'a Recurser,
input_a: AggregationInput<'a>,
input_b: AggregationInput<'a>,
) -> Self {
Self {
client,
agg,
input_a,
input_b,
root_c_recurser_agg: None,
timeout: None,
subscribers: Vec::new(),
}
}
#[must_use]
pub fn root_c_recurser_agg(mut self, limbs: [u64; 4]) -> Self {
self.root_c_recurser_agg = Some(limbs);
self
}
#[must_use]
pub fn timeout(mut self, duration: Duration) -> Self {
self.timeout = Some(duration);
self
}
#[must_use]
pub fn on(mut self, event: JobEvent, cb: impl Fn(JobEvent) + Send + Sync + 'static) -> Self {
self.subscribers.push((event, Arc::new(cb)));
self
}
pub fn run(self) -> Result<JobHandle<ProveResult>> {
let n_free = self.agg.n_free();
let validate = |side: char, input: &AggregationInput<'_>| -> Result<()> {
let got = input.free_inputs.len();
if got != n_free {
return Err(SdkError::Recurser(format!(
"proof_{side} supplies {got} free inputs but the recurser \
consumes exactly {n_free} (a plain &Proof supplies 0; use \
`.with_free_inputs(..)` to attach exactly {n_free})",
)));
}
Ok(())
};
validate('a', &self.input_a)?;
validate('b', &self.input_b)?;
let free_a = self.input_a.free_inputs;
let free_b = self.input_b.free_inputs;
let subs = subscriber_list_from(self.subscribers);
self.client.run_aggregate_proofs(
self.agg,
self.input_a.proof,
self.input_b.proof,
&free_a,
&free_b,
self.root_c_recurser_agg,
self.timeout,
subs,
)
}
}