use alloc::vec;
use alloc::vec::Vec;
use libm::{log2, pow};
use crate::error::ErrorBits;
use crate::ldt::LowDegreeTest;
use crate::proximity::{LDR_M_CAP, alpha_ldr_m, alpha_udr, compute_upper_m, gamma_ldr_m};
use crate::report::{LDT_COMMIT_LABEL, LDT_QUERY_LABEL, SecurityTerm};
use crate::shape::{InstanceShape, StarkAirParams};
#[derive(Copy, Clone, Debug)]
pub struct FriRegime {
pub log_blowup: usize,
pub num_queries: usize,
pub log_final_poly_len: usize,
pub max_log_arity: usize,
pub commit_pow_bits: usize,
pub query_pow_bits: usize,
}
impl FriRegime {
const fn folding_factor(self) -> f64 {
(1usize << self.max_log_arity) as f64
}
}
impl LowDegreeTest for FriRegime {
fn log_blowup(&self) -> usize {
self.log_blowup
}
fn proven_error_udr(&self, air: &StarkAirParams, shape: &InstanceShape) -> ErrorBits {
proven_error_udr(self, air, shape)
}
fn best_ldr(&self, air: &StarkAirParams, shape: &InstanceShape) -> Option<(usize, ErrorBits)> {
best_ldr_m(self, air, shape)
}
fn conjectured_error(&self, shape: &InstanceShape) -> ErrorBits {
conjectured_error(self, shape)
}
fn conjectured_terms(&self, shape: &InstanceShape) -> Vec<SecurityTerm> {
let mut terms = vec![SecurityTerm::new(
LDT_QUERY_LABEL,
conjectured_error(self, shape),
)];
terms.extend(
conjectured_commit_phase_error(self, shape)
.map(|bits| SecurityTerm::new(LDT_COMMIT_LABEL, bits)),
);
terms
}
}
pub fn conjectured_error(regime: &FriRegime, shape: &InstanceShape) -> ErrorBits {
if regime.log_blowup == 0 || shape.modulus_bits == 0 {
return ErrorBits::from_log2(regime.query_pow_bits as f64);
}
let log_blowup_f = regime.log_blowup as f64;
let rho = pow(2.0, -log_blowup_f);
let log2_e_over_rho = core::f64::consts::LOG2_E + log_blowup_f;
let eta = (log2_e_over_rho * rho) / shape.modulus_bits as f64;
let effective = rho + eta;
if effective <= 0.0 || effective >= 1.0 {
return ErrorBits::from_log2(regime.query_pow_bits as f64);
}
let bits_per_query = -log2(effective);
let bits = regime.num_queries as f64 * bits_per_query + regime.query_pow_bits as f64;
ErrorBits::from_log2(bits)
}
pub fn conjectured_commit_phase_error(
regime: &FriRegime,
shape: &InstanceShape,
) -> Option<ErrorBits> {
commit_phase_error_udr(regime, shape)
}
pub fn commit_phase_error_udr(regime: &FriRegime, shape: &InstanceShape) -> Option<ErrorBits> {
let folding_minus_one = regime.folding_factor() - 1.0;
if folding_minus_one <= 0.0 {
return None;
}
let lde_log = shape.log_trace_length + regime.log_blowup;
let num_layers = lde_log.saturating_sub(regime.log_final_poly_len) / regime.max_log_arity;
if num_layers == 0 {
return None;
}
let n = (1u64 << lde_log) as f64;
let bits = shape.modulus_bits as f64 - log2(folding_minus_one * (n + 1.0))
+ regime.commit_pow_bits as f64;
Some(ErrorBits::from_log2(bits.max(0.0)))
}
pub fn commit_phase_error_ldr_m(
regime: &FriRegime,
shape: &InstanceShape,
m: usize,
) -> Option<ErrorBits> {
let rho = pow(2.0, -(regime.log_blowup as f64));
let sqrt_rho = libm::sqrt(rho);
let m_shifted = m as f64 + 0.5;
let pp = gamma_ldr_m(regime.log_blowup, m);
if pp <= 0.0 {
return Some(ErrorBits::from_log2(0.0));
}
let folding_minus_one = regime.folding_factor() - 1.0;
if folding_minus_one <= 0.0 {
return None;
}
let lde_log = shape.log_trace_length + regime.log_blowup;
let n = (1u64 << lde_log) as f64;
let num = (2.0 * pow(m_shifted, 5.0) + 3.0 * m_shifted * pp * rho) * n;
let den = 3.0 * rho * sqrt_rho;
let eps_linear = num / den + m_shifted / sqrt_rho;
let eps_powers = eps_linear * folding_minus_one;
let bits_linear =
shape.modulus_bits as f64 - log2(eps_powers.max(1.0)) + regime.commit_pow_bits as f64;
let bits_n_over_q = shape.modulus_bits as f64
- log2(regime.folding_factor())
- log2(n + 1.0)
- log2(2.0 * m as f64 + 1.0)
+ 0.5 * log2(rho)
+ regime.commit_pow_bits as f64;
Some(ErrorBits::from_log2(
bits_linear.min(bits_n_over_q).max(0.0),
))
}
pub fn query_phase_error(alpha: f64, num_queries: usize, query_pow_bits: usize) -> ErrorBits {
if !alpha.is_finite() || alpha <= 0.0 || alpha >= 1.0 {
return ErrorBits::from_log2(0.0);
}
let bits = query_pow_bits as f64 - log2(pow(alpha, num_queries as f64));
ErrorBits::from_log2(bits)
}
pub fn proven_error_udr(
regime: &FriRegime,
air: &StarkAirParams,
shape: &InstanceShape,
) -> ErrorBits {
if regime.log_blowup == 0 || shape.log_trace_length == 0 || shape.modulus_bits == 0 {
return ErrorBits::from_log2(0.0);
}
let alpha = alpha_udr(shape.log_trace_length, regime.log_blowup, air.max_combo);
let lde = (1u64 << (shape.log_trace_length + regime.log_blowup)) as f64;
let k = (1u64 << shape.log_trace_length) as f64;
if k + air.max_combo as f64 >= alpha * lde {
return ErrorBits::from_log2(0.0);
}
let query = query_phase_error(alpha, regime.num_queries, regime.query_pow_bits);
commit_phase_error_udr(regime, shape).map_or(query, |commit| ErrorBits::min(&[commit, query]))
}
pub fn proven_error_ldr_m(
regime: &FriRegime,
air: &StarkAirParams,
shape: &InstanceShape,
m: usize,
) -> ErrorBits {
if regime.log_blowup == 0 || shape.log_trace_length == 0 || shape.modulus_bits == 0 {
return ErrorBits::from_log2(0.0);
}
let alpha = alpha_ldr_m(regime.log_blowup, m);
if alpha >= 1.0 {
return ErrorBits::from_log2(0.0);
}
let pp = gamma_ldr_m(regime.log_blowup, m);
if pp <= 0.0 {
return ErrorBits::from_log2(0.0);
}
let lde = (1u64 << (shape.log_trace_length + regime.log_blowup)) as f64;
let k = (1u64 << shape.log_trace_length) as f64;
if k + air.max_combo as f64 >= (1.0 - pp) * lde {
return ErrorBits::from_log2(0.0);
}
let query = query_phase_error(alpha, regime.num_queries, regime.query_pow_bits);
commit_phase_error_ldr_m(regime, shape, m)
.map_or(query, |commit| ErrorBits::min(&[commit, query]))
}
pub fn best_ldr_m(
regime: &FriRegime,
air: &StarkAirParams,
shape: &InstanceShape,
) -> Option<(usize, ErrorBits)> {
let trace_length = 1usize << shape.log_trace_length;
let m_max = core::cmp::min(compute_upper_m(trace_length), LDR_M_CAP);
let m_min = 3usize;
if m_max < m_min {
return None;
}
(m_min..=m_max)
.map(|m| (m, proven_error_ldr_m(regime, air, shape, m)))
.max_by(|a, b| {
a.1.bits()
.partial_cmp(&b.1.bits())
.unwrap_or(core::cmp::Ordering::Equal)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::stark::proven_security;
fn benchmark_regime() -> FriRegime {
FriRegime {
log_blowup: 1,
num_queries: 100,
log_final_poly_len: 0,
max_log_arity: 3,
commit_pow_bits: 0,
query_pow_bits: 16,
}
}
fn benchmark_shape() -> InstanceShape {
InstanceShape {
log_trace_length: 20,
modulus_bits: 252,
collision_resistance: 128,
num_batched_functions: 1,
}
}
fn benchmark_air() -> StarkAirParams {
StarkAirParams {
num_constraints: 1,
max_constraint_degree: 2,
max_combo: 2,
}
}
#[test]
fn proven_security_regression_benchmark_high_arity() {
let regime = benchmark_regime();
let air = benchmark_air();
let shape = benchmark_shape();
let udr_ldt = proven_error_udr(®ime, &air, &shape);
let (best_m, ldr_ldt) = best_ldr_m(®ime, &air, &shape).unwrap();
let udr_bits = crate::stark::proven_security_udr(&air, &shape, udr_ldt, &[])
.bits()
.floor() as usize;
let ldr_bits = crate::stark::proven_security_ldr_m(
&air,
&shape,
regime.log_blowup,
best_m,
ldr_ldt,
&[],
)
.bits()
.floor() as usize;
assert_eq!(udr_bits, 57);
assert_eq!(ldr_bits, 65);
let combined = proven_security(
&air,
&shape,
regime.log_blowup,
udr_ldt,
best_m,
ldr_ldt,
&[],
)
.bits()
.floor() as usize;
assert_eq!(combined, 65);
}
#[test]
fn conjectured_commit_phase_binds_below_a_128_bit_target() {
let regime = FriRegime {
log_blowup: 3,
num_queries: 27,
log_final_poly_len: 0,
max_log_arity: 2,
commit_pow_bits: 0,
query_pow_bits: 16,
};
let shape = InstanceShape {
log_trace_length: 23,
modulus_bits: 128,
collision_resistance: 128,
num_batched_functions: 1,
};
let commit = conjectured_commit_phase_error(®ime, &shape).expect("folds occur");
assert!(
(commit.bits() - (128.0 - libm::log2(3.0 * (65_536.0 * 1024.0 + 1.0)))).abs() < 1e-9
);
assert!(commit.bits() < 128.0, "got {}", commit.bits());
assert_eq!(
commit.bits(),
commit_phase_error_udr(®ime, &shape).unwrap().bits()
);
}
#[test]
fn conjectured_commit_phase_credits_only_its_own_grinding() {
let base = benchmark_regime();
let shape = benchmark_shape();
let ground = FriRegime {
commit_pow_bits: 12,
..base
};
let b0 = conjectured_commit_phase_error(&base, &shape).expect("folds occur");
let b12 = conjectured_commit_phase_error(&ground, &shape).expect("folds occur");
assert!((b12.bits() - b0.bits() - 12.0).abs() < 1e-12);
assert_eq!(
conjectured_error(&ground, &shape).bits(),
conjectured_error(&base, &shape).bits()
);
let no_folds = FriRegime {
log_final_poly_len: shape.log_trace_length + base.log_blowup,
..base
};
assert!(conjectured_commit_phase_error(&no_folds, &shape).is_none());
assert_eq!(
LowDegreeTest::conjectured_terms(&no_folds, &shape).len(),
1,
"a fold-free regime reports the query phase only"
);
}
#[test]
fn arity_one_reports_no_commit_round_rather_than_impersonating_arity_two() {
let base = benchmark_regime();
let shape = benchmark_shape();
let air = benchmark_air();
let no_arity = FriRegime {
max_log_arity: 0,
..base
};
assert!(commit_phase_error_udr(&no_arity, &shape).is_none());
assert!(commit_phase_error_ldr_m(&no_arity, &shape, 10).is_none());
assert!(conjectured_commit_phase_error(&no_arity, &shape).is_none());
assert_eq!(
proven_error_udr(&no_arity, &air, &shape).bits(),
query_phase_error(
alpha_udr(shape.log_trace_length, no_arity.log_blowup, air.max_combo),
no_arity.num_queries,
no_arity.query_pow_bits,
)
.bits()
);
}
#[test]
fn conjectured_bounded_by_collision_resistance() {
let regime = FriRegime {
log_blowup: 8,
num_queries: 32,
log_final_poly_len: 0,
max_log_arity: 1,
commit_pow_bits: 0,
query_pow_bits: 0,
};
let shape = InstanceShape {
log_trace_length: 16,
modulus_bits: 128,
collision_resistance: 128,
num_batched_functions: 1,
};
let bits = conjectured_error(®ime, &shape)
.bits()
.min(shape.collision_resistance as f64)
.min(shape.modulus_bits as f64)
.floor() as usize;
assert_eq!(bits, 128);
}
}