use core::fmt::Debug;
use getset::Getters;
use p3_field::{BasedVectorSpace, ExtensionField, PrimeField64, TwoAdicField};
use serde::{Deserialize, Serialize};
use crate::{hasher::MerkleHasher, interaction::LogUpSecurityParameters};
pub trait StarkProtocolConfig: 'static + Clone + Send + Sync {
type F: TwoAdicField + PrimeField64;
type EF: TwoAdicField + ExtensionField<Self::F>;
type Digest: Copy
+ Send
+ Sync
+ Debug
+ Default
+ PartialEq
+ Eq
+ Serialize
+ for<'de> Deserialize<'de>;
type Hasher: MerkleHasher<F = Self::F, Digest = Self::Digest>;
const D_EF: usize = <Self::EF as BasedVectorSpace<Self::F>>::DIMENSION;
fn params(&self) -> &SystemParams;
fn hasher(&self) -> &Self::Hasher;
}
pub type Val<SC> = <SC as StarkProtocolConfig>::F;
pub type Com<SC> = <SC as StarkProtocolConfig>::Digest;
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq, Getters)]
pub struct SystemParams {
pub l_skip: usize,
pub n_stack: usize,
pub w_stack: usize,
pub log_blowup: usize,
#[getset(get = "pub")]
pub whir: WhirConfig,
pub logup: LogUpSecurityParameters,
pub max_constraint_degree: usize,
}
impl SystemParams {
pub fn logup_pow_bits(&self) -> usize {
self.logup.pow_bits
}
pub fn k_whir(&self) -> usize {
self.whir.k
}
#[inline]
pub fn log_stacked_height(&self) -> usize {
self.l_skip + self.n_stack
}
#[inline]
pub fn log_final_poly_len(&self) -> usize {
self.whir.log_final_poly_len(self.log_stacked_height())
}
#[inline]
pub fn num_whir_rounds(&self) -> usize {
self.whir.num_whir_rounds()
}
#[inline]
pub fn num_whir_sumcheck_rounds(&self) -> usize {
self.whir.num_sumcheck_rounds()
}
#[allow(clippy::too_many_arguments)]
pub fn new(
log_blowup: usize,
l_skip: usize,
n_stack: usize,
w_stack: usize,
log_final_poly_len: usize,
folding_pow_bits: usize,
mu_pow_bits: usize,
proximity: WhirProximityStrategy,
security_bits: usize,
logup: LogUpSecurityParameters,
max_constraint_degree: usize,
whir_query_phase_pow_bits: usize,
k_whir: usize,
) -> SystemParams {
let log_stacked_height = l_skip + n_stack;
SystemParams {
l_skip,
n_stack,
w_stack,
log_blowup,
whir: WhirConfig::new(
log_blowup,
log_stacked_height,
WhirParams {
k: k_whir,
log_final_poly_len,
query_phase_pow_bits: whir_query_phase_pow_bits,
proximity,
folding_pow_bits,
mu_pow_bits,
},
security_bits,
),
logup,
max_constraint_degree,
}
}
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct WhirParams {
pub k: usize,
pub log_final_poly_len: usize,
pub query_phase_pow_bits: usize,
pub proximity: WhirProximityStrategy,
pub folding_pow_bits: usize,
pub mu_pow_bits: usize,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct WhirConfig {
pub k: usize,
pub rounds: Vec<WhirRoundConfig>,
pub mu_pow_bits: usize,
pub query_phase_pow_bits: usize,
pub folding_pow_bits: usize,
pub proximity: WhirProximityStrategy,
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct WhirRoundConfig {
pub num_queries: usize,
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub enum WhirProximityStrategy {
UniqueDecoding,
SplitUniqueList { m: usize, list_start_round: usize },
ListDecoding { m: usize },
}
impl WhirProximityStrategy {
pub fn initial_round(&self) -> ProximityRegime {
self.in_round(0)
}
pub fn in_round(&self, whir_round: usize) -> ProximityRegime {
match *self {
Self::UniqueDecoding => ProximityRegime::UniqueDecoding,
Self::SplitUniqueList {
m,
list_start_round,
} => {
if whir_round < list_start_round {
ProximityRegime::UniqueDecoding
} else {
ProximityRegime::ListDecoding { m }
}
}
Self::ListDecoding { m } => ProximityRegime::ListDecoding { m },
}
}
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub enum ProximityRegime {
UniqueDecoding,
ListDecoding { m: usize },
}
impl ProximityRegime {
pub fn max_agreement(&self, log_inv_rate: usize) -> f64 {
let rho = 2.0_f64.powf(-(log_inv_rate as f64));
let max_agreement = match *self {
ProximityRegime::UniqueDecoding => (1.0 + rho) / 2.0,
ProximityRegime::ListDecoding { m } => {
let m = m.max(1) as f64;
rho.sqrt() * (1.0 + 1.0 / (2.0 * m))
}
};
max_agreement.clamp(f64::MIN_POSITIVE, 1.0)
}
pub fn whir_query_security_bits(&self, num_queries: usize, log_inv_rate: usize) -> f64 {
-(num_queries as f64) * self.max_agreement(log_inv_rate).log2()
}
pub fn whir_per_query_security_bits(&self, log_inv_rate: usize) -> f64 {
self.whir_query_security_bits(1, log_inv_rate)
}
}
impl WhirConfig {
pub fn new(
log_blowup: usize,
log_stacked_height: usize,
whir_params: WhirParams,
security_bits: usize,
) -> Self {
let query_phase_pow_bits = whir_params.query_phase_pow_bits;
let protocol_security_level = security_bits.saturating_sub(query_phase_pow_bits);
let k_whir = whir_params.k;
let num_rounds = log_stacked_height
.saturating_sub(whir_params.log_final_poly_len)
.div_ceil(k_whir);
let mut log_inv_rate = log_blowup;
let proximity = whir_params.proximity;
let mut round_parameters = Vec::with_capacity(num_rounds);
for round in 0..num_rounds {
let next_rate = log_inv_rate + (k_whir - 1);
let num_queries = Self::queries(
proximity.in_round(round),
protocol_security_level,
log_inv_rate,
);
round_parameters.push(WhirRoundConfig { num_queries });
log_inv_rate = next_rate;
}
Self {
k: k_whir,
rounds: round_parameters,
mu_pow_bits: whir_params.mu_pow_bits,
query_phase_pow_bits,
folding_pow_bits: whir_params.folding_pow_bits,
proximity,
}
}
#[inline]
pub fn log_final_poly_len(&self, log_stacked_height: usize) -> usize {
log_stacked_height - self.num_whir_rounds() * self.k
}
pub fn num_whir_rounds(&self) -> usize {
self.rounds.len()
}
#[inline]
pub fn num_sumcheck_rounds(&self) -> usize {
self.num_whir_rounds() * self.k
}
pub fn queries(
proximity_regime: ProximityRegime,
protocol_security_level: usize,
log_inv_rate: usize,
) -> usize {
let per_query_bits = proximity_regime.whir_per_query_security_bits(log_inv_rate);
let num_queries_f = (protocol_security_level as f64) / per_query_bits;
num_queries_f.ceil() as usize
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_whir_list_decoding_query_bits_monotone_in_num_queries() {
let regime = ProximityRegime::ListDecoding { m: 2 };
let sec_10 = regime.whir_query_security_bits(10, 1);
let sec_20 = regime.whir_query_security_bits(20, 1);
assert!(sec_20 > sec_10);
assert!(sec_10 > 0.0);
}
}