use std::sync::Arc;
use derivative::Derivative;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use crate::{
air_builders::symbolic::{symbolic_variable::SymbolicVariable, SymbolicConstraintsDag},
prover::stacked_pcs::StackedPcsData,
StarkProtocolConfig, SystemParams,
};
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct TraceWidth {
pub preprocessed: Option<usize>,
pub cached_mains: Vec<usize>,
pub common_main: usize,
}
impl TraceWidth {
pub fn main_widths(&self) -> Vec<usize> {
let mut ret = self.cached_mains.clone();
if self.common_main != 0 {
ret.push(self.common_main);
}
ret
}
pub fn main_width(&self) -> usize {
self.cached_mains.iter().sum::<usize>() + self.common_main
}
pub fn total_width(&self) -> usize {
self.preprocessed.unwrap_or(0) + self.main_width()
}
}
#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, Eq)]
pub struct LinearConstraint {
pub coefficients: Vec<u32>,
pub threshold: u32,
}
impl LinearConstraint {
pub fn is_implied_by(&self, other: &LinearConstraint) -> bool {
self.threshold >= other.threshold
&& self
.coefficients
.iter()
.zip(&other.coefficients)
.all(|(a, b)| a <= b)
}
}
#[derive(Error, Debug)]
pub enum KeygenError {
#[error("AIR {name} has zero main trace width")]
AirWidthZero { name: String },
#[error("AIR {name} must have at least one constraint or interaction")]
AirNoConstraintsOrInteractions { name: String },
#[error("AIR {name} interaction {interaction_index} has a zero-length message")]
InteractionMessageEmpty {
name: String,
interaction_index: usize,
},
#[error("Max constraint degree exceeded for AIR {name}: {degree} > {max_degree}")]
MaxConstraintDegreeExceeded {
name: String,
degree: usize,
max_degree: usize,
},
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[repr(C)]
pub struct StarkVerifyingParams {
pub width: TraceWidth,
pub num_public_values: usize,
pub need_rot: bool,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct VerifierSinglePreprocessedData<Digest> {
pub commit: Digest,
pub hypercube_dim: isize,
pub stacking_width: usize,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[repr(C)]
pub struct StarkVerifyingKey<F, Digest> {
pub preprocessed_data: Option<VerifierSinglePreprocessedData<Digest>>,
pub params: StarkVerifyingParams,
pub symbolic_constraints: SymbolicConstraintsDag<F>,
pub max_constraint_degree: u8,
pub is_required: bool,
pub unused_variables: Vec<SymbolicVariable<F>>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(Clone(bound = ""), Debug(bound = ""))]
#[serde(bound = "")]
pub struct MultiStarkVerifyingKey<SC: StarkProtocolConfig> {
pub inner: MultiStarkVerifyingKey0<SC>,
pub pre_hash: SC::Digest,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(Clone(bound = ""), Debug(bound = ""))]
#[serde(bound = "")]
pub struct MultiStarkVerifyingKey0<SC: StarkProtocolConfig> {
pub params: SystemParams,
pub per_air: Vec<StarkVerifyingKey<SC::F, SC::Digest>>,
pub trace_height_constraints: Vec<LinearConstraint>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(Clone(bound = ""))]
#[serde(bound = "")]
pub struct StarkProvingKey<SC: StarkProtocolConfig> {
pub air_name: String,
pub vk: StarkVerifyingKey<SC::F, SC::Digest>,
pub preprocessed_data: Option<Arc<StackedPcsData<SC::F, SC::Digest>>>,
}
#[derive(Derivative, Serialize, Deserialize)]
#[derivative(Clone(bound = ""))]
#[serde(bound = "")]
pub struct MultiStarkProvingKey<SC: StarkProtocolConfig> {
pub per_air: Vec<StarkProvingKey<SC>>,
pub trace_height_constraints: Vec<LinearConstraint>,
pub max_constraint_degree: usize,
pub params: SystemParams,
pub vk_pre_hash: SC::Digest,
}
impl<Val, Com> StarkVerifyingKey<Val, Com> {
pub fn num_cached_mains(&self) -> usize {
self.params.width.cached_mains.len()
}
pub fn num_parts(&self) -> usize {
1 + self.num_cached_mains() + (self.preprocessed_data.is_some() as usize)
}
pub fn has_interaction(&self) -> bool {
!self.symbolic_constraints.interactions.is_empty()
}
pub fn num_interactions(&self) -> usize {
self.symbolic_constraints.interactions.len()
}
pub fn dag_main_part_index_to_commit_index(&self, index: usize) -> usize {
if index == self.num_cached_mains() {
0
} else {
index + 1 + self.preprocessed_data.is_some() as usize
}
}
}
impl<SC: StarkProtocolConfig> MultiStarkProvingKey<SC> {
pub fn get_vk(&self) -> MultiStarkVerifyingKey<SC> {
MultiStarkVerifyingKey {
inner: self.get_vk0(),
pre_hash: self.vk_pre_hash,
}
}
fn get_vk0(&self) -> MultiStarkVerifyingKey0<SC> {
MultiStarkVerifyingKey0 {
params: self.params.clone(),
per_air: self.per_air.iter().map(|pk| pk.vk.clone()).collect(),
trace_height_constraints: self.trace_height_constraints.clone(),
}
}
}
impl<SC: StarkProtocolConfig> MultiStarkVerifyingKey<SC> {
pub fn max_constraint_degree(&self) -> usize {
self.inner.max_constraint_degree()
}
}
impl<SC: StarkProtocolConfig> MultiStarkVerifyingKey0<SC> {
pub fn max_constraint_degree(&self) -> usize {
self.params.max_constraint_degree
}
}