use group::Group;
use midnight_circuits::{
instructions::{AssignmentInstructions, BinaryInstructions, PublicInputInstructions},
types::Instantiable,
verifier::{Accumulator, AssignedAccumulator, AssignedVk},
};
use midnight_proofs::{
circuit::{Layouter, Value},
plonk::ConstraintSystem,
poly::EvaluationDomain,
};
use midnight_zk_stdlib::{Relation, ZkStdLib, ZkStdLibArch};
use super::{Ivc, IvcError, C, F, S};
#[derive(Clone, Debug)]
pub struct IvcInstance<T: Ivc> {
pub(crate) vk_repr: F,
pub(crate) state: T::State,
pub(crate) acc: Accumulator<S>,
}
impl<T: Ivc> IvcInstance<T> {
pub fn state(&self) -> &T::State {
&self.state
}
pub fn acc(&self) -> &Accumulator<S> {
&self.acc
}
}
#[derive(Clone, Debug)]
pub struct IvcWitness<T: Ivc> {
pub(crate) prev_state: T::State,
pub(crate) prev_acc: Accumulator<S>,
pub(crate) prev_proof: Vec<u8>,
pub(crate) transition_witness: T::Witness,
}
#[derive(Clone, Debug)]
pub struct IvcCircuit<T: Ivc> {
domain: EvaluationDomain<F>,
cs: ConstraintSystem<F>,
ctx: T::Context,
}
impl<T: Ivc> IvcCircuit<T> {
pub fn new(domain: EvaluationDomain<F>, cs: ConstraintSystem<F>, ctx: T::Context) -> Self {
IvcCircuit { domain, cs, ctx }
}
pub fn ctx(&self) -> &T::Context {
&self.ctx
}
pub fn arch() -> ZkStdLibArch {
let mut arch = T::arch();
arch.bls12_381 = true;
arch.poseidon = true;
arch
}
}
impl<T: Ivc> Relation for IvcCircuit<T> {
type Instance = IvcInstance<T>;
type Witness = IvcWitness<T>;
type Error = IvcError;
fn used_chips(&self) -> ZkStdLibArch {
Self::arch()
}
fn format_instance(instance: &Self::Instance) -> Result<Vec<F>, IvcError> {
Ok([
vec![instance.vk_repr],
T::format_public_input(&instance.state),
AssignedAccumulator::<S>::as_public_input(&instance.acc),
]
.concat())
}
fn circuit(
&self,
std_lib: &ZkStdLib,
layouter: &mut impl Layouter<F>,
instance: Value<Self::Instance>,
witness: Value<Self::Witness>,
) -> Result<(), IvcError> {
let verifier_gadget = std_lib.verifier();
let ivc_gadget = T::new(std_lib.clone(), &self.ctx);
let assigned_self_vk: AssignedVk<S> = verifier_gadget.assign_vk_as_public_input(
layouter,
"self_vk",
&self.domain,
&self.cs,
instance.as_ref().map(|x| x.vk_repr),
)?;
let prev_state_val = witness.as_ref().map(|w| w.prev_state.clone());
let prev_state = ivc_gadget.assign(layouter, prev_state_val)?;
let next_state = ivc_gadget.circuit_transition(
layouter,
&prev_state,
witness.as_ref().map(|w| w.transition_witness.clone()),
)?;
ivc_gadget.constrain_as_public_input(layouter, &next_state)?;
let fixed_base_names = midnight_circuits::verifier::fixed_base_names::<S>(
"self_vk",
self.cs.num_fixed_columns() + self.cs.num_selectors(),
self.cs.permutation().columns.len(),
);
let prev_acc_value = witness.as_ref().map(|w| w.prev_acc.clone());
let prev_acc = verifier_gadget.assign_collapsed_accumulator(
layouter,
&fixed_base_names,
prev_acc_value,
)?;
let prev_proof_pi = [
verifier_gadget.as_public_input(layouter, &assigned_self_vk)?,
ivc_gadget.as_public_input(layouter, &prev_state)?,
verifier_gadget.as_public_input(layouter, &prev_acc)?,
]
.concat();
let id_point = std_lib.bls12_381().assign_fixed(layouter, C::identity())?;
let mut prev_proof_acc = verifier_gadget.prepare(
layouter,
&assigned_self_vk,
&[id_point],
&[&prev_proof_pi],
witness.map(|w| w.prev_proof),
)?;
let is_not_genesis = {
let b = ivc_gadget.circuit_is_genesis(std_lib, layouter, &self.ctx, &prev_state)?;
std_lib.not(layouter, &b)?
};
AssignedAccumulator::scale_by_bit(
layouter,
std_lib.bls12_381().scalar_field_chip(),
&is_not_genesis,
&mut prev_proof_acc,
)?;
let mut next_acc = verifier_gadget.accumulate(layouter, &[prev_proof_acc, prev_acc])?;
next_acc.collapse(
layouter,
std_lib.bls12_381(),
std_lib.bls12_381().scalar_field_chip(),
)?;
verifier_gadget.constrain_as_public_input(layouter, &next_acc)?;
Ok(())
}
fn write_relation<W: std::io::Write>(&self, writer: &mut W) -> std::io::Result<()> {
writer.write_all(&self.domain.k().to_le_bytes())?;
T::write_context(&self.ctx, writer)
}
fn read_relation<R: std::io::Read>(reader: &mut R) -> std::io::Result<Self> {
let mut k_bytes = [0u8; 4];
reader.read_exact(&mut k_bytes)?;
let k = u32::from_le_bytes(k_bytes);
let ctx = T::read_context(reader)?;
let mut cs = ConstraintSystem::default();
ZkStdLib::configure(&mut cs, (Self::arch(), (k - 1) as u8));
let domain = EvaluationDomain::new(cs.degree() as u32, k);
Ok(IvcCircuit { domain, cs, ctx })
}
}