use crate::hash::Hash;
use crate::merkle::{Proof as LeafProof, Tree};
use crate::utils;
use anyhow::{Result, anyhow};
use starkom_bluesky::Scalar;
use starkom_ff::{Field, PrimeField};
use starkom_poly;
use std::marker::PhantomData;
use std::sync::LazyLock;
type Polynomial = starkom_poly::Polynomial<Scalar>;
static FOLD_DST: LazyLock<Scalar> = LazyLock::new(|| utils::hash_to_scalar(b"starkom/fri/fold"));
trait FoldableTree<H: Hash<Scalar>> {
fn fold(&self) -> Self;
fn fold_all(self, times: usize) -> Vec<Tree<H>>;
}
impl<H: Hash<Scalar>> FoldableTree<H> for Tree<H> {
fn fold(&self) -> Self {
let num_polys = self.num_polys();
let n = self.num_leaves();
assert!(n.is_power_of_two());
let alpha = H::hash_two(*FOLD_DST, self.root_hash(), Scalar::ZERO);
let k = n.trailing_zeros() as usize;
let omega_inv = Scalar::ROOT_OF_UNITY_INV.pow_u64(1u64 << (Scalar::S - k));
let m = n / 2;
let mut omega_inv_i = Scalar::ONE;
let mut leaves = vec![vec![Scalar::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)) * Scalar::TWO_INV;
}
omega_inv_i *= omega_inv;
}
Self::new(leaves)
}
fn fold_all(self, times: usize) -> Vec<Self> {
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<Scalar>,
}
impl Commitment {
pub fn len(&self) -> usize {
self.roots.len()
}
pub fn roots(&self) -> &[Scalar] {
self.roots.as_slice()
}
pub fn root(&self) -> Scalar {
*self.roots.first().unwrap()
}
}
#[derive(Debug, Clone)]
pub struct Query<H: Hash<Scalar>> {
degree_bound: usize,
blowup_log2: usize,
index: usize,
folds: Vec<(LeafProof<H>, LeafProof<H>)>,
_data: PhantomData<H>,
}
impl<H: Hash<Scalar>> Query<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) -> Scalar {
Polynomial::coset_element2(self.index, self.degree_bound << self.blowup_log2)
}
pub fn values(&self) -> (&[Scalar], &[Scalar]) {
(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 k = n.trailing_zeros() as usize;
let folds = self.folds.as_slice();
let num_folds = folds.len();
if num_folds > self.degree_bound.trailing_zeros() as usize + 1 {
return Err(anyhow!("invalid 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 = Scalar::ROOT_OF_UNITY_INV.pow_u64(1u64 << (Scalar::S - k));
for round in 0..num_folds {
let (left, right) = &folds[round];
let root_hash = commitment.roots()[round];
let alpha = H::hash_two(*FOLD_DST, root_hash, Scalar::ZERO);
let neg = right.leaf();
if 1usize << left.len() != n {
return Err(anyhow!(
"invalid left-hand side Merkle proof height (got {}, want {})",
left.len(),
n.trailing_zeros()
));
}
if 1usize << right.len() != n {
return Err(anyhow!(
"invalid right-hand side Merkle proof height (got {}, want {})",
right.len(),
n.trailing_zeros()
));
}
left.check_leaf(pos.as_slice())?;
left.verify(index, root_hash)?;
right.verify((index + n / 2) % n, root_hash)?;
let omega_inv_i = step.pow_small(index);
n /= 2;
index %= n;
for i in 0..pos.len() {
pos[i] =
(pos[i] + neg[i] + alpha * omega_inv_i * (pos[i] - neg[i])) * Scalar::TWO_INV;
}
step = step.square();
}
let (left, right) = folds.last().unwrap();
if !left.is_constant() || !right.is_constant() {
return Err(anyhow!("final folded polynomial is not constant"));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Prover<H: Hash<Scalar>> {
degree_bound: usize,
blowup_log2: usize,
trees: Vec<Tree<H>>,
}
impl<H: Hash<Scalar>> Prover<H> {
pub fn new(polynomials: Vec<Polynomial>, degree_bound: usize, blowup_log2: usize) -> Self {
assert!(degree_bound.is_power_of_two());
assert!(
polynomials
.iter()
.all(|polynomial| degree_bound >= polynomial.degree_bound())
);
let n = degree_bound << blowup_log2;
assert!(n as u64 <= 1u64 << Scalar::S);
let main_tree = Tree::<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) -> Scalar {
self.trees[0].root_hash()
}
pub fn commit(&self) -> Commitment {
Commitment {
roots: self.trees.iter().map(|tree| tree.root_hash()).collect(),
}
}
pub fn query(&self, index: usize) -> Query<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,
_data: Default::default(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hash;
use starkom_bluesky::from_const;
type Poseidon2Hash = hash::Poseidon2Hash<Scalar>;
type Sha2Hash = hash::Sha2Hash<Scalar>;
fn test_prover_impl<H: Hash<Scalar>>(
polynomials: Vec<Polynomial>,
degree_bound: usize,
blowup_log2: usize,
) {
let prover = Prover::<H>::new(polynomials, 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<Polynomial>, degree_bound: usize) {
test_prover_impl::<Sha2Hash>(polynomials.clone(), degree_bound, 1);
test_prover_impl::<Poseidon2Hash>(polynomials.clone(), degree_bound, 1);
test_prover_impl::<Sha2Hash>(polynomials.clone(), degree_bound, 2);
test_prover_impl::<Poseidon2Hash>(polynomials.clone(), degree_bound, 2);
test_prover_impl::<Sha2Hash>(polynomials.clone(), degree_bound, 3);
test_prover_impl::<Poseidon2Hash>(polynomials.clone(), degree_bound, 3);
}
#[test]
fn test_one_constant_polynomial() {
test_prover(vec![Polynomial::with_coefficients(vec![from_const(12)])], 1);
test_prover(vec![Polynomial::with_coefficients(vec![from_const(34)])], 1);
}
#[test]
fn test_two_constant_polynomials() {
test_prover(
vec![
Polynomial::with_coefficients(vec![from_const(12)]),
Polynomial::with_coefficients(vec![from_const(34)]),
],
1,
);
}
#[test]
fn test_three_constant_polynomials() {
test_prover(
vec![
Polynomial::with_coefficients(vec![from_const(34)]),
Polynomial::with_coefficients(vec![from_const(56)]),
Polynomial::with_coefficients(vec![from_const(78)]),
],
1,
);
}
#[test]
fn test_one_polynomial_degree_one() {
test_prover(
vec![Polynomial::with_coefficients(vec![
from_const(12),
from_const(34),
])],
2,
);
test_prover(
vec![Polynomial::with_coefficients(vec![
from_const(56),
from_const(78),
])],
2,
);
}
#[test]
fn test_two_polynomials_degree_one() {
test_prover(
vec![
Polynomial::with_coefficients(vec![from_const(12), from_const(34)]),
Polynomial::with_coefficients(vec![from_const(56), from_const(78)]),
],
2,
);
}
#[test]
fn test_three_polynomials_degree_one() {
test_prover(
vec![
Polynomial::with_coefficients(vec![from_const(34), from_const(56)]),
Polynomial::with_coefficients(vec![from_const(56), from_const(78)]),
Polynomial::with_coefficients(vec![from_const(78), from_const(90)]),
],
2,
);
}
#[test]
fn test_one_polynomial_degree_three() {
test_prover(
vec![Polynomial::with_coefficients(vec![
from_const(12),
from_const(34),
from_const(56),
from_const(78),
])],
4,
);
test_prover(
vec![Polynomial::with_coefficients(vec![
from_const(42),
from_const(43),
from_const(44),
from_const(45),
])],
4,
);
}
#[test]
fn test_two_polynomials_degree_three() {
test_prover(
vec![
Polynomial::with_coefficients(vec![
from_const(12),
from_const(34),
from_const(56),
from_const(78),
]),
Polynomial::with_coefficients(vec![
from_const(42),
from_const(43),
from_const(44),
from_const(45),
]),
],
4,
);
}
#[test]
fn test_three_polynomials_degree_three() {
test_prover(
vec![
Polynomial::with_coefficients(vec![
from_const(42),
from_const(43),
from_const(44),
from_const(45),
]),
Polynomial::with_coefficients(vec![
from_const(12),
from_const(34),
from_const(56),
from_const(78),
]),
Polynomial::with_coefficients(vec![
from_const(34),
from_const(56),
from_const(78),
from_const(90),
]),
],
4,
);
}
}