use std::sync::Arc;
use itertools::Itertools;
use crate::{
air_builders::debug::{debug_constraints_and_interactions, AirProofRawInput},
keygen::{
types::{MultiStarkProvingKey, MultiStarkVerifyingKey},
MultiStarkKeygenBuilder,
},
memory_metering::ProvingMemoryConfig,
proof::*,
prover::{
AirProvingContext, ColMajorMatrix, Coordinator, DeviceDataTransporter,
DeviceMultiStarkProvingKey, MultiRapProver, OpeningProver, Prover, ProverBackend,
ProverDevice, ProvingContext, StridedColMajorMatrixView,
},
verifier::{verify, VerifierError},
AirRef, FiatShamirTranscript, StarkProtocolConfig, SystemParams,
};
#[derive(Debug)]
pub struct VerificationData<SC: StarkProtocolConfig> {
pub vk: MultiStarkVerifyingKey<SC>,
pub proof: Proof<SC>,
}
pub type ProverError<E> =
<<E as StarkEngine>::PD as ProverDevice<<E as StarkEngine>::PB, <E as StarkEngine>::TS>>::Error;
pub type EngineDeviceCtx<E> = <<E as StarkEngine>::PD as ProverDevice<
<E as StarkEngine>::PB,
<E as StarkEngine>::TS,
>>::DeviceCtx;
pub trait StarkEngine
where
<Self::PD as MultiRapProver<Self::PB, Self::TS>>::Artifacts:
Into<<Self::PD as OpeningProver<Self::PB, Self::TS>>::OpeningPoints>,
<Self::PD as MultiRapProver<Self::PB, Self::TS>>::PartialProof:
Into<(GkrProof<Self::SC>, BatchConstraintProof<Self::SC>)>,
<Self::PD as OpeningProver<Self::PB, Self::TS>>::OpeningProof:
Into<(StackingProof<Self::SC>, WhirProof<Self::SC>)>,
{
type SC: StarkProtocolConfig;
type PB: ProverBackend<
Val = <Self::SC as StarkProtocolConfig>::F,
Challenge = <Self::SC as StarkProtocolConfig>::EF,
Commitment = <Self::SC as StarkProtocolConfig>::Digest,
>;
type PD: ProverDevice<Self::PB, Self::TS> + DeviceDataTransporter<Self::SC, Self::PB>;
type TS: FiatShamirTranscript<Self::SC>;
fn new(params: SystemParams) -> Self;
fn config(&self) -> &Self::SC;
fn params(&self) -> &SystemParams {
self.config().params()
}
fn proving_memory_config(&self) -> ProvingMemoryConfig {
ProvingMemoryConfig::from_protocol_config(self.config(), true)
}
fn device(&self) -> &Self::PD;
fn initial_transcript(&self) -> Self::TS;
fn prover_from_transcript(
&self,
transcript: Self::TS,
) -> Coordinator<Self::SC, Self::PB, Self::PD, Self::TS>;
fn prover(&self) -> Coordinator<Self::SC, Self::PB, Self::PD, Self::TS> {
let transcript = self.initial_transcript();
self.prover_from_transcript(transcript)
}
fn keygen(
&self,
airs: &[AirRef<Self::SC>],
) -> (
MultiStarkProvingKey<Self::SC>,
MultiStarkVerifyingKey<Self::SC>,
) {
let mut keygen_builder = MultiStarkKeygenBuilder::new(self.config().clone());
for air in airs {
keygen_builder.add_air(air.clone());
}
let pk = keygen_builder.generate_pk().unwrap();
let vk = pk.get_vk();
(pk, vk)
}
fn prove(
&self,
pk: &DeviceMultiStarkProvingKey<Self::PB>,
ctx: ProvingContext<Self::PB>,
) -> Result<Proof<Self::SC>, ProverError<Self>> {
let mut prover = self.prover();
prover.prove(pk, ctx)
}
fn verify(
&self,
vk: &MultiStarkVerifyingKey<Self::SC>,
proof: &Proof<Self::SC>,
) -> Result<(), VerifierError<<Self::SC as StarkProtocolConfig>::EF>> {
let mut transcript = self.initial_transcript();
verify(self.config(), vk, proof, &mut transcript)
}
fn debug(&self, airs: &[AirRef<Self::SC>], ctx: &ProvingContext<Self::PB>) {
let mut keygen_builder = MultiStarkKeygenBuilder::new(self.config().clone());
for air in airs {
keygen_builder.add_air(air.clone());
}
let pk = keygen_builder.generate_pk().unwrap();
let transpose = |mat: ColMajorMatrix<<Self::SC as StarkProtocolConfig>::F>| {
let row_major = StridedColMajorMatrixView::from(mat.as_view()).to_row_major_matrix();
Arc::new(row_major)
};
let (inputs, used_airs, used_pks): (Vec<_>, Vec<_>, Vec<_>) = ctx
.per_trace
.iter()
.map(|(air_id, trace_ctx)| {
let common_main = self
.device()
.transport_matrix_from_device_to_host(&trace_ctx.common_main);
let cached_mains = trace_ctx
.cached_mains
.iter()
.map(|cd| {
transpose(
self.device()
.transport_matrix_from_device_to_host(&cd.trace),
)
})
.collect_vec();
let common_main = Some(transpose(common_main));
let public_values = trace_ctx.public_values.clone();
(
AirProofRawInput {
cached_mains,
common_main,
public_values,
},
airs[*air_id].clone(),
&pk.per_air[*air_id],
)
})
.multiunzip();
debug_constraints_and_interactions(&used_airs, &used_pks, &inputs);
}
fn run_test(
&self,
airs: Vec<AirRef<Self::SC>>,
ctxs: Vec<AirProvingContext<Self::PB>>,
) -> RunTestResult<Self::SC, Self::PB, Self::PD, Self::TS> {
let (pk, vk) = self.keygen(&airs);
let device = self.prover().device;
let d_pk = device.transport_pk_to_device(&pk);
let ctx = ProvingContext::new(ctxs.into_iter().enumerate().collect());
let proof = self.prove(&d_pk, ctx).map_err(StarkTestError::Prover)?;
self.verify(&vk, &proof).map_err(StarkTestError::Verifier)?;
Ok(VerificationData { vk, proof })
}
}
type RunTestResult<SC, PB, PD, TS> = Result<
VerificationData<SC>,
StarkTestError<<PD as ProverDevice<PB, TS>>::Error, <SC as StarkProtocolConfig>::EF>,
>;
#[derive(Debug)]
pub enum StarkTestError<
PE: std::error::Error,
EF: std::fmt::Debug + std::fmt::Display + PartialEq + Eq,
> {
Prover(PE),
Verifier(VerifierError<EF>),
}
impl<PE: std::error::Error, EF: core::fmt::Debug + core::fmt::Display + PartialEq + Eq>
core::fmt::Display for StarkTestError<PE, EF>
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Prover(e) => write!(f, "Prover error: {e:?}"),
Self::Verifier(e) => write!(f, "Verifier error: {e}"),
}
}
}
impl<PE: std::error::Error, EF: core::fmt::Debug + core::fmt::Display + PartialEq + Eq>
std::error::Error for StarkTestError<PE, EF>
{
}