use std::sync::Arc;
use serde::{de::DeserializeOwned, Serialize};
use crate::{
keygen::types::MultiStarkProvingKey,
prover::{
stacked_pcs::StackedPcsData, AirProvingContext, ColMajorMatrix, CommittedTraceData,
CpuColMajorBackend, DeviceMultiStarkProvingKey, ProvingContext,
},
StarkProtocolConfig,
};
pub trait MatrixDimensions {
fn height(&self) -> usize;
fn width(&self) -> usize;
}
pub trait ProverBackend {
const CHALLENGE_EXT_DEGREE: u8;
type Val: Copy + Send + Sync + Serialize + DeserializeOwned;
type Challenge: Copy + Send + Sync + Serialize + DeserializeOwned;
type Commitment: Clone + Send + Sync + Serialize + DeserializeOwned;
type Matrix: MatrixDimensions + Send + Sync;
type OtherAirData: Send + Sync;
type PcsData: Send + Sync;
}
pub trait ProverDevice<PB: ProverBackend, TS>:
TraceCommitter<PB> + MultiRapProver<PB, TS> + OpeningProver<PB, TS>
{
type Error: 'static
+ std::error::Error
+ Send
+ Sync
+ From<<Self as TraceCommitter<PB>>::Error>
+ From<<Self as MultiRapProver<PB, TS>>::Error>
+ From<<Self as OpeningProver<PB, TS>>::Error>;
type DeviceCtx: Clone + Send + Sync;
fn device_ctx(&self) -> &Self::DeviceCtx;
}
pub trait TraceCommitter<PB: ProverBackend> {
type Error: std::fmt::Debug;
fn commit(&self, traces: &[&PB::Matrix]) -> Result<(PB::Commitment, PB::PcsData), Self::Error>;
}
pub trait MultiRapProver<PB: ProverBackend, TS> {
type PartialProof: Clone + Send + Sync + Serialize + DeserializeOwned;
type Artifacts;
type Error: std::fmt::Debug;
fn prove_rap_constraints(
&self,
transcript: &mut TS,
mpk: &DeviceMultiStarkProvingKey<PB>,
ctx: &ProvingContext<PB>,
common_main_pcs_data: &PB::PcsData,
) -> Result<(Self::PartialProof, Self::Artifacts), Self::Error>;
}
pub trait OpeningProver<PB: ProverBackend, TS> {
type OpeningProof: Clone + Send + Sync + Serialize + DeserializeOwned;
type OpeningPoints;
type Error: std::fmt::Debug;
fn prove_openings(
&self,
transcript: &mut TS,
mpk: &DeviceMultiStarkProvingKey<PB>,
ctx: ProvingContext<PB>,
common_main_pcs_data: PB::PcsData,
points: Self::OpeningPoints,
) -> Result<Self::OpeningProof, Self::Error>;
}
pub trait DeviceDataTransporter<SC, PB>
where
SC: StarkProtocolConfig,
PB: ProverBackend<Val = SC::F, Challenge = SC::EF, Commitment = SC::Digest>,
{
fn transport_pk_to_device(
&self,
mpk: &MultiStarkProvingKey<SC>,
) -> DeviceMultiStarkProvingKey<PB>;
fn transport_matrix_to_device(&self, matrix: &ColMajorMatrix<SC::F>) -> PB::Matrix;
fn transport_pcs_data_to_device(
&self,
pcs_data: &StackedPcsData<SC::F, SC::Digest>,
) -> PB::PcsData;
fn transport_committed_trace_data_to_device(
&self,
committed_trace: &CommittedTraceData<CpuColMajorBackend<SC>>,
) -> CommittedTraceData<PB> {
let trace = self.transport_matrix_to_device(&committed_trace.trace);
let data = self.transport_pcs_data_to_device(committed_trace.data.as_ref());
CommittedTraceData {
commitment: committed_trace.commitment,
trace,
data: Arc::new(data),
}
}
fn transport_proving_ctx_to_device(
&self,
ctx: &ProvingContext<CpuColMajorBackend<SC>>,
) -> ProvingContext<PB> {
let per_trace = ctx
.per_trace
.iter()
.map(|(air_idx, trace_ctx)| {
let common_main = self.transport_matrix_to_device(&trace_ctx.common_main);
let cached_mains = trace_ctx
.cached_mains
.iter()
.map(|cd| self.transport_committed_trace_data_to_device(cd))
.collect();
let trace_ctx_gpu = AirProvingContext::new(
cached_mains,
common_main,
trace_ctx.public_values.clone(),
);
(*air_idx, trace_ctx_gpu)
})
.collect();
ProvingContext::new(per_trace)
}
fn transport_matrix_from_device_to_host(&self, matrix: &PB::Matrix) -> ColMajorMatrix<SC::F>;
}