use crate::hash::Hasher;
use crate::merkle::{Proof as LeafProof, Tree};
use crate::utils;
use anyhow::{Result, anyhow};
use primitive_types::H256;
use starkom_ff::Field256;
use starkom_poly::Polynomial;
use std::sync::LazyLock;
static FOLD_DST: LazyLock<H256> = LazyLock::new(|| utils::make_dst(b"starkom/fri/fold"));
trait FoldableTree<F: Field256, H: Hasher<F>>: Sized {
fn fold(&self) -> Tree<F, H>;
fn fold_all(self, times: usize) -> Vec<Tree<F, H>>;
}
impl<F: Field256, H: Hasher<F>> FoldableTree<F, H> for Tree<F, H> {
fn fold(&self) -> Tree<F, H> {
let num_polys = self.num_polys();
let n = self.num_leaves();
assert!(n.is_power_of_two());
let alpha = H::challenge(*FOLD_DST, &[self.root_hash()]);
let k = n.trailing_zeros() as usize;
let omega_inv = F::ROOT_OF_UNITY_INV.pow_u64(1u64 << (F::S - k));
let m = n / 2;
let mut omega_inv_i = F::ONE;
let mut leaves = vec![vec![F::ZERO; m]; num_polys];
for i in 0..m {
for j in 0..num_polys {
let pos = self.leaf_value(j, i);
let neg = self.leaf_value(j, i + m);
leaves[j][i] = (pos + neg + alpha * omega_inv_i * (pos - neg)) * F::TWO_INV;
}
omega_inv_i *= omega_inv;
}
Self::new(leaves)
}
fn fold_all(self, times: usize) -> Vec<Tree<F, H>> {
let mut trees = Vec::with_capacity(times + 1);
let mut tree = self;
for _ in 0..times {
let folded = tree.fold();
trees.push(tree);
tree = folded;
}
trees.push(tree);
trees
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Commitment {
roots: Vec<H256>,
}
impl Commitment {
pub fn len(&self) -> usize {
self.roots.len()
}
pub fn roots(&self) -> &[H256] {
self.roots.as_slice()
}
pub fn root(&self) -> H256 {
*self.roots.first().unwrap()
}
}
#[derive(Debug, Clone)]
pub struct Query<F: Field256, H: Hasher<F>> {
degree_bound: usize,
blowup_log2: usize,
index: usize,
folds: Vec<(LeafProof<F, H>, LeafProof<F, H>)>,
}
impl<F: Field256, H: Hasher<F>> Query<F, H> {
pub fn indices(&self) -> (usize, usize) {
let n = self.degree_bound << self.blowup_log2;
(self.index, (self.index + n / 2) % n)
}
pub fn x(&self) -> F {
Polynomial::<F>::coset_element2(self.index, self.degree_bound << self.blowup_log2)
}
pub fn values(&self) -> (&[F], &[F]) {
(self.folds[0].0.leaf(), self.folds[0].1.leaf())
}
pub fn len(&self) -> usize {
self.folds.len()
}
pub fn verify(&self, commitment: &Commitment) -> Result<()> {
let mut n = self.degree_bound << self.blowup_log2;
assert!(n.is_power_of_two());
assert!(self.index < n);
let mut k = n.trailing_zeros() as usize;
let num_folds = self.folds.len();
if num_folds > self.degree_bound.trailing_zeros() as usize + 1 {
return Err(anyhow!("incorrect proof size"));
}
if commitment.len() != num_folds {
return Err(anyhow!("wrong number of folding rounds"));
}
let mut index = self.index;
let mut pos = self.folds[0].0.leaf().to_vec();
let mut step = F::ROOT_OF_UNITY_INV.pow_u64(1u64 << (F::S - k));
for round in 0..num_folds {
let root_hash = commitment.roots()[round];
let (left, right) = &self.folds[round];
if left.len() != k {
return Err(anyhow!(
"invalid left-hand side Merkle proof height (got {}, want {})",
left.len(),
k
));
}
if right.len() != k {
return Err(anyhow!(
"invalid right-hand side Merkle proof height (got {}, want {})",
right.len(),
k
));
}
left.check_leaf(pos.as_slice())?;
left.verify(index, root_hash)?;
right.verify((index + n / 2) % n, root_hash)?;
let omega_inv = step.pow_small(index);
n /= 2;
k -= 1;
index %= n;
let neg = right.leaf();
let alpha = H::challenge(*FOLD_DST, &[root_hash]);
for i in 0..pos.len() {
pos[i] = (pos[i] + neg[i] + alpha * omega_inv * (pos[i] - neg[i])) * F::TWO_INV;
}
step = step.square();
}
let (left, right) = self.folds.last().unwrap();
if !left.is_constant() || !right.is_constant() {
return Err(anyhow!("the final folded polynomial is not constant"));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Prover<F: Field256, H: Hasher<F>> {
degree_bound: usize,
blowup_log2: usize,
trees: Vec<Tree<F, H>>,
}
impl<F: Field256, H: Hasher<F>> Prover<F, H> {
pub fn new(polynomials: Vec<Polynomial<F>>, degree_bound: usize, blowup_log2: usize) -> Self {
assert!(degree_bound.is_power_of_two());
assert!(
polynomials
.iter()
.all(|polynomial| degree_bound >= polynomial.degree_bound())
);
assert!(blowup_log2 > 0);
let n = degree_bound << blowup_log2;
assert!(n as u64 <= 1u64 << F::S);
let main_tree = Tree::<F, H>::new(
polynomials
.into_iter()
.map(|polynomial| polynomial.shift_domain().lde2(n))
.collect(),
);
let trees = main_tree.fold_all(degree_bound.trailing_zeros() as usize);
Self {
degree_bound,
blowup_log2,
trees,
}
}
pub fn degree_bound(&self) -> usize {
self.degree_bound
}
pub fn extended_domain_size(&self) -> usize {
self.degree_bound << self.blowup_log2
}
pub fn size(&self) -> usize {
self.degree_bound << self.blowup_log2
}
pub fn root_hash(&self) -> H256 {
self.trees[0].root_hash()
}
pub fn commit(&self) -> Commitment {
Commitment {
roots: self.trees.iter().map(Tree::root_hash).collect(),
}
}
pub fn query(&self, index: usize) -> Query<F, H> {
let mut n = self.degree_bound << self.blowup_log2;
assert!(index < n);
let mut i = index;
let mut folds = vec![];
for tree in &self.trees {
folds.push((tree.query(i), tree.query((i + n / 2) % n)));
n /= 2;
i %= n;
}
{
let (left, right) = folds.last().unwrap();
assert!(left.is_constant());
assert!(right.is_constant());
}
Query {
degree_bound: self.degree_bound,
blowup_log2: self.blowup_log2,
index,
folds,
}
}
}
#[cfg(all(test, feature = "bluesky", feature = "goldilocks"))]
mod tests {
use super::*;
use crate::hash::{Keccak256Hash, Sha2Hash};
use starkom_bluesky::Scalar as BS;
use starkom_goldilocks::GL4;
#[test]
fn test_fold_dst() {
assert_eq!(
*FOLD_DST,
"0x9ffd3556faeb2cae194ce95adf6b3580f590504daa0dea56966ce4ef233844af"
.parse()
.unwrap()
);
}
fn test_prover_impl<F: Field256, H: Hasher<F>>(
polynomials: &[Polynomial<F>],
degree_bound: usize,
blowup_log2: usize,
) {
let prover = Prover::<F, H>::new(polynomials.to_vec(), degree_bound, blowup_log2);
assert_eq!(prover.degree_bound(), degree_bound);
let n = degree_bound << blowup_log2;
assert_eq!(prover.extended_domain_size(), n);
let commitment = prover.commit();
for i in 0..n {
let query = prover.query(i);
assert_eq!(query.indices(), (i, (i + n / 2) % n));
assert_eq!(query.len(), degree_bound.trailing_zeros() as usize + 1);
assert!(query.verify(&commitment).is_ok());
}
}
fn test_prover(polynomials: Vec<Vec<u64>>, degree_bound: usize) {
let bluesky_polynomials: Vec<Polynomial<BS>> = polynomials
.iter()
.map(|coefficients| {
Polynomial::with_coefficients(
coefficients.iter().copied().map(BS::from_const).collect(),
)
})
.collect();
let goldilocks_polynomials: Vec<Polynomial<GL4>> = polynomials
.into_iter()
.map(|coefficients| {
Polynomial::with_coefficients(
coefficients.into_iter().map(GL4::from_const).collect(),
)
})
.collect();
test_prover_impl::<BS, Sha2Hash<BS>>(&bluesky_polynomials, degree_bound, 1);
test_prover_impl::<GL4, Sha2Hash<GL4>>(&goldilocks_polynomials, degree_bound, 1);
test_prover_impl::<BS, Keccak256Hash<BS>>(&bluesky_polynomials, degree_bound, 1);
test_prover_impl::<GL4, Keccak256Hash<GL4>>(&goldilocks_polynomials, degree_bound, 1);
test_prover_impl::<BS, Sha2Hash<BS>>(&bluesky_polynomials, degree_bound, 2);
test_prover_impl::<GL4, Sha2Hash<GL4>>(&goldilocks_polynomials, degree_bound, 2);
test_prover_impl::<BS, Keccak256Hash<BS>>(&bluesky_polynomials, degree_bound, 2);
test_prover_impl::<GL4, Keccak256Hash<GL4>>(&goldilocks_polynomials, degree_bound, 2);
test_prover_impl::<BS, Sha2Hash<BS>>(&bluesky_polynomials, degree_bound, 3);
test_prover_impl::<GL4, Sha2Hash<GL4>>(&goldilocks_polynomials, degree_bound, 3);
test_prover_impl::<BS, Keccak256Hash<BS>>(&bluesky_polynomials, degree_bound, 3);
test_prover_impl::<GL4, Keccak256Hash<GL4>>(&goldilocks_polynomials, degree_bound, 3);
}
#[test]
fn test_one_constant_polynomial() {
test_prover(vec![vec![12]], 1);
test_prover(vec![vec![34]], 1);
}
#[test]
fn test_two_constant_polynomials() {
test_prover(vec![vec![12], vec![34]], 1);
}
#[test]
fn test_three_constant_polynomials() {
test_prover(vec![vec![34], vec![56], vec![78]], 1);
}
#[test]
fn test_one_polynomial_degree_one() {
test_prover(vec![vec![12, 34]], 2);
test_prover(vec![vec![56, 78]], 2);
}
#[test]
fn test_two_polynomials_degree_one() {
test_prover(vec![vec![12, 34], vec![56, 78]], 2);
}
#[test]
fn test_three_polynomials_degree_one() {
test_prover(vec![vec![34, 56], vec![56, 78], vec![78, 90]], 2);
}
#[test]
fn test_one_polynomial_degree_three() {
test_prover(vec![vec![12, 34, 56, 78]], 4);
test_prover(vec![vec![42, 43, 44, 45]], 4);
}
#[test]
fn test_two_polynomials_degree_three() {
test_prover(vec![vec![12, 34, 56, 78], vec![42, 43, 44, 45]], 4);
}
#[test]
fn test_three_polynomials_degree_three() {
test_prover(
vec![
vec![42, 43, 44, 45],
vec![12, 34, 56, 78],
vec![34, 56, 78, 90],
],
4,
);
}
#[test]
fn test_one_polynomial_degree_seven() {
test_prover(vec![vec![12, 34, 56, 78, 90, 12, 34]], 8);
test_prover(vec![vec![42, 43, 44, 45, 46, 47, 48]], 8);
}
#[test]
fn test_two_polynomials_degree_seven() {
test_prover(
vec![
vec![12, 34, 56, 78, 90, 12, 34],
vec![42, 43, 44, 45, 46, 47, 48],
],
8,
);
}
#[test]
fn test_three_polynomials_degree_seven() {
test_prover(
vec![
vec![42, 43, 44, 45, 46, 47, 48],
vec![12, 34, 56, 78, 90, 12, 34],
vec![34, 56, 78, 90, 78, 56, 34],
],
8,
);
}
const DEGREE_BOUND: usize = 8;
const BLOWUP_LOG2: usize = 2;
const QUERY_INDEX: usize = 3;
fn make_prover(coefficients: [u64; 8]) -> Prover<BS, Sha2Hash<BS>> {
Prover::new(
vec![Polynomial::with_coefficients(
coefficients.into_iter().map(BS::from_const).collect(),
)],
DEGREE_BOUND,
BLOWUP_LOG2,
)
}
fn honest_prover() -> Prover<BS, Sha2Hash<BS>> {
make_prover([12, 34, 56, 78, 90, 12, 34, 56])
}
fn assert_rejected(result: Result<()>, expected: &str) {
let error = result
.expect_err("the tampered proof was accepted")
.to_string();
assert!(error.contains(expected), "unexpected error: {error}");
}
#[test]
fn test_accept_untampered_proof() {
let prover = honest_prover();
assert!(prover.query(QUERY_INDEX).verify(&prover.commit()).is_ok());
}
#[test]
fn test_reject_foreign_commitment() {
let query = honest_prover().query(QUERY_INDEX);
let foreign = make_prover([90, 78, 56, 34, 12, 90, 78, 56]).commit();
assert_rejected(query.verify(&foreign), "root hash mismatch");
}
#[test]
fn test_reject_corrupted_root() {
let prover = honest_prover();
let mut commitment = prover.commit();
commitment.roots[2] = H256::repeat_byte(0xAA);
assert_rejected(
prover.query(QUERY_INDEX).verify(&commitment),
"root hash mismatch",
);
}
#[test]
fn test_reject_truncated_folds() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
query.folds.truncate(3);
assert_rejected(query.verify(&commitment), "wrong number of folding rounds");
}
#[test]
fn test_reject_truncated_commitment() {
let prover = honest_prover();
let mut commitment = prover.commit();
commitment.roots.truncate(3);
assert_rejected(
prover.query(QUERY_INDEX).verify(&commitment),
"wrong number of folding rounds",
);
}
#[test]
fn test_reject_extra_folds() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
let mut donor = prover.query(QUERY_INDEX);
query.folds.push(donor.folds.pop().unwrap());
assert_rejected(query.verify(&commitment), "incorrect proof size");
}
#[test]
fn test_reject_folding_stopped_early() {
let prover = honest_prover();
let mut commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
commitment.roots.truncate(3);
query.folds.truncate(3);
assert_rejected(query.verify(&commitment), "not constant");
}
#[test]
fn test_reject_wrong_left_proof_height() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
let mut donor = prover.query(QUERY_INDEX);
query.folds[0].0 = donor.folds.remove(1).0;
assert_rejected(
query.verify(&commitment),
"invalid left-hand side Merkle proof height",
);
}
#[test]
fn test_reject_wrong_right_proof_height() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
let mut donor = prover.query(QUERY_INDEX);
query.folds[0].1 = donor.folds.remove(1).1;
assert_rejected(
query.verify(&commitment),
"invalid right-hand side Merkle proof height",
);
}
#[test]
fn test_reject_swapped_partners() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
let pair = &mut query.folds[0];
std::mem::swap(&mut pair.0, &mut pair.1);
assert_rejected(query.verify(&commitment), "root hash mismatch");
}
#[test]
fn test_reject_foreign_leaf_proof() {
let prover = honest_prover();
let commitment = prover.commit();
let mut query = prover.query(QUERY_INDEX);
let mut other = prover.query(QUERY_INDEX + 4);
query.folds[1].0 = other.folds.remove(1).0;
assert_rejected(query.verify(&commitment), "leaf value mismatch");
}
}