use std::ops::Deref;
use bls12_381::{batch_inversion::batch_inverse, traits::*, Scalar};
use polynomial::{
domain::Domain,
poly_coeff::{vanishing_poly, PolyCoeff},
CosetFFT,
};
use crate::errors::RSError;
pub(crate) enum ErasurePattern {
BlockSynchronizedErasures(BlockErasureIndices),
#[cfg(test)]
Random { indices: Vec<usize> },
}
type BlockErasureIndex = usize;
#[derive(Debug, Clone, Default)]
pub struct BlockErasureIndices(pub Vec<BlockErasureIndex>);
impl Deref for BlockErasureIndices {
type Target = Vec<BlockErasureIndex>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[derive(Debug)]
pub struct ReedSolomon {
expansion_factor: usize,
poly_len: usize,
evaluation_domain: Domain,
block_size: usize,
num_blocks: usize,
block_size_domain: Domain,
fft_coset_gen: CosetFFT,
}
impl ReedSolomon {
pub fn new(poly_len: usize, expansion_factor: usize, block_size: usize) -> Self {
assert!(
expansion_factor.is_power_of_two()
&& poly_len.is_power_of_two()
&& block_size.is_power_of_two()
);
let evaluation_size = poly_len * expansion_factor;
Self {
expansion_factor,
poly_len,
evaluation_domain: Domain::new(evaluation_size),
block_size,
num_blocks: evaluation_size / block_size,
block_size_domain: Domain::new(block_size),
fft_coset_gen: CosetFFT::new(Scalar::MULTIPLICATIVE_GENERATOR),
}
}
const fn acceptable_num_random_erasures(&self) -> usize {
let total_codeword_len = self.poly_len * self.expansion_factor;
let min_num_evaluations_needed = self.poly_len;
total_codeword_len - min_num_evaluations_needed
}
pub const fn acceptable_num_block_erasures(&self) -> usize {
self.acceptable_num_random_erasures() / self.num_blocks
}
pub const fn codeword_length(&self) -> usize {
self.poly_len * self.expansion_factor
}
pub fn encode(&self, poly_coefficient_form: PolyCoeff) -> Result<Vec<Scalar>, RSError> {
if poly_coefficient_form.len() > self.poly_len {
return Err(RSError::PolynomialHasTooManyCoefficients {
num_coefficients: poly_coefficient_form.len(),
max_num_coefficients: self.poly_len,
});
}
Ok(self.evaluation_domain.fft_scalars(poly_coefficient_form))
}
pub fn recover_polynomial_coefficient(
&self,
codeword_with_erasures: Vec<Scalar>,
erasures: BlockErasureIndices,
) -> Result<PolyCoeff, RSError> {
self.recover_polynomial_coefficient_erasure_pattern(
codeword_with_erasures,
ErasurePattern::BlockSynchronizedErasures(erasures),
)
}
#[cfg(test)]
fn recover_polynomial_coefficient_random_erasure(
&self,
codeword_with_erasures: Vec<Scalar>,
random_erasure: Vec<usize>,
) -> Result<PolyCoeff, RSError> {
self.recover_polynomial_coefficient_erasure_pattern(
codeword_with_erasures,
ErasurePattern::Random {
indices: random_erasure,
},
)
}
fn construct_vanishing_poly_from_block_erasures(
&self,
block_indices: &BlockErasureIndices,
) -> PolyCoeff {
assert!(block_indices.len() != self.block_size, "all of the blocks are missing. This should have been checked by the caller of this method");
let evaluation_domain_size = self.evaluation_domain.roots.len();
let z_x_missing_indices_roots: Vec<_> = block_indices
.iter()
.map(|index| self.block_size_domain.roots[*index])
.collect();
let vanish_poly_first_block = vanishing_poly(&z_x_missing_indices_roots);
let mut z_x = vec![Scalar::ZERO; evaluation_domain_size];
for (i, coeff) in vanish_poly_first_block.0.into_iter().enumerate() {
z_x[i * self.num_blocks] = coeff;
}
z_x.into()
}
fn construct_vanishing_poly_from_erasure_pattern(
&self,
erasures: ErasurePattern,
) -> Result<PolyCoeff, RSError> {
match erasures {
ErasurePattern::BlockSynchronizedErasures(indices) => {
for &block_index in &indices.0 {
if block_index >= self.block_size {
return Err(RSError::InvalidBlockIndex {
block_index,
block_size: self.block_size,
});
}
}
if indices.len() > self.acceptable_num_block_erasures() {
return Err(RSError::TooManyBlockErasures {
num_block_erasures: indices.len(),
max_num_block_erasures_accepted: self.acceptable_num_block_erasures(),
});
}
Ok(self.construct_vanishing_poly_from_block_erasures(&indices))
}
#[cfg(test)]
ErasurePattern::Random { indices } => {
assert!(
indices.len() <= self.acceptable_num_random_erasures(),
"num random erasures = {} but tolerable erasures = {}",
indices.len(),
self.acceptable_num_random_erasures()
);
let roots: Vec<_> = indices
.into_iter()
.map(|index| self.evaluation_domain.roots[index])
.collect();
Ok(vanishing_poly(&roots))
}
}
}
fn recover_polynomial_coefficient_erasure_pattern(
&self,
e_eval: Vec<Scalar>,
erasure: ErasurePattern,
) -> Result<PolyCoeff, RSError> {
let z_x = self.construct_vanishing_poly_from_erasure_pattern(erasure)?;
let z_eval = self.evaluation_domain.fft_scalars(z_x.clone());
let ez_eval: Vec<_> = z_eval.iter().zip(e_eval).map(|(zx, d)| zx * d).collect();
let dz_coeffs = self.evaluation_domain.ifft_scalars(ez_eval);
let dz_coset_eval = self
.evaluation_domain
.coset_fft_scalars(dz_coeffs, &self.fft_coset_gen);
let mut z_inv_coset_eval = self
.evaluation_domain
.coset_fft_scalars(z_x, &self.fft_coset_gen);
batch_inverse(&mut z_inv_coset_eval);
let d_eval: Vec<_> = dz_coset_eval
.iter()
.zip(z_inv_coset_eval)
.map(|(d, zx_inv)| d * zx_inv)
.collect();
let d_coeffs = self
.evaluation_domain
.coset_ifft_scalars(d_eval, &self.fft_coset_gen);
for coefficient in d_coeffs.iter().skip(self.poly_len) {
if *coefficient != Scalar::ZERO {
return Err(RSError::PolynomialHasInvalidLength {
num_coefficients: d_coeffs.len(),
expected_num_coefficients: self.poly_len,
});
}
}
Ok(d_coeffs[..self.poly_len].to_vec().into())
}
}
#[cfg(test)]
mod tests {
use bls12_381::{traits::*, Scalar};
use polynomial::poly_coeff::PolyCoeff;
use crate::{reed_solomon::ErasurePattern, BlockErasureIndices, ReedSolomon};
#[test]
#[should_panic]
fn test_compute_vanishing_panics() {
const POLY_LEN: usize = 16;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 1;
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let block_erasure_indices: Vec<_> = (0..BLOCK_SIZE).collect();
rs.construct_vanishing_poly_from_block_erasures(&BlockErasureIndices(
block_erasure_indices,
));
}
#[test]
fn smoke_test_recovery_no_erasures() {
const POLY_LEN: usize = 16;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 1;
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let poly_coeff = PolyCoeff((0..16).map(|i| -Scalar::from(i)).collect());
let codewords = rs
.encode(poly_coeff.clone())
.expect("polynomial encode failed");
assert_eq!(codewords.len(), 32);
let got_poly_coeff = rs
.recover_polynomial_coefficient(codewords, BlockErasureIndices::default())
.expect("polynomial recovery failed");
assert_eq!(got_poly_coeff.len(), poly_coeff.len());
assert_eq!(got_poly_coeff, poly_coeff);
}
#[test]
fn test_vanishing_poly_erasure_pattern_block_synchronized() {
const POLY_LEN: usize = 512;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 16;
let indices = vec![0, 1, 2, 3];
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let z =
rs.construct_vanishing_poly_from_block_erasures(&BlockErasureIndices(indices.clone()));
assert_eq!(z.len(), POLY_LEN * EXPANSION_FACTOR);
let evals = rs.evaluation_domain.fft_scalars(z);
let blocks: Vec<_> = evals.chunks(BLOCK_SIZE).collect();
assert!(blocks.len() == rs.num_blocks);
for block in &blocks {
for index in 0..BLOCK_SIZE {
if indices.contains(&index) {
assert_eq!(block[index], Scalar::ZERO);
} else {
assert_ne!(block[index], Scalar::ZERO);
}
}
}
}
#[test]
fn test_vanishing_poly_erasure_pattern_equiv_random() {
const POLY_LEN: usize = 64;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 4;
let indices = vec![0, 1];
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let got_z_x =
rs.construct_vanishing_poly_from_block_erasures(&BlockErasureIndices(indices.clone()));
let got_z_x_lagrange_form = rs.evaluation_domain.fft_scalars(got_z_x);
let blocks: Vec<_> = got_z_x_lagrange_form.chunks(BLOCK_SIZE).collect();
let mut all_indices = Vec::new();
for index in indices {
for i in 0..blocks.len() {
all_indices.push(index + i * BLOCK_SIZE);
}
}
let z_x = rs
.construct_vanishing_poly_from_erasure_pattern(ErasurePattern::Random {
indices: all_indices,
})
.expect("failed to create vanishing polynomial");
let expected_z_x_lagrange_form = rs.evaluation_domain.fft_scalars(z_x);
assert_eq!(expected_z_x_lagrange_form, got_z_x_lagrange_form);
}
#[test]
fn smoke_test_recovery_upto_num_acceptable_random_erasures() {
const POLY_LEN: usize = 16;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 1;
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let poly_coeff = PolyCoeff((0..POLY_LEN).map(|i| Scalar::from(i as u64)).collect());
let original_codewords = rs
.encode(poly_coeff.clone())
.expect("polynomial encode failed");
let acceptable_num_erasures: Vec<_> = (0..=rs.acceptable_num_random_erasures()).collect();
for num_erasures in acceptable_num_erasures {
let mut codewords_with_erasures = original_codewords.clone();
let mut missing_indices = Vec::new();
for (index, codewords_with_erasure) in codewords_with_erasures
.iter_mut()
.enumerate()
.take(num_erasures)
{
*codewords_with_erasure = Scalar::ZERO;
missing_indices.push(index);
}
let recovered_poly_coeff = rs
.recover_polynomial_coefficient_random_erasure(
codewords_with_erasures,
missing_indices,
)
.expect("failed to recover polynomial");
assert_eq!(recovered_poly_coeff.len(), poly_coeff.len());
assert_eq!(recovered_poly_coeff, poly_coeff);
}
}
#[test]
fn smoke_test_recovery_upto_num_acceptable_block_erasures() {
const POLY_LEN: usize = 128;
const EXPANSION_FACTOR: usize = 2;
const BLOCK_SIZE: usize = 4;
let rs = ReedSolomon::new(POLY_LEN, EXPANSION_FACTOR, BLOCK_SIZE);
let poly_coeff = PolyCoeff((0..POLY_LEN).map(|i| Scalar::from(i as u64)).collect());
let original_codewords = rs
.encode(poly_coeff.clone())
.expect("polynomial encode failed");
let num_block_erasures: Vec<_> = (0..=BLOCK_SIZE).collect();
for num_block_erasures in num_block_erasures {
let mut blocks: Vec<Vec<Scalar>> = original_codewords
.chunks(BLOCK_SIZE)
.map(<[Scalar]>::to_vec)
.collect();
let mut missing_block_indices = Vec::new();
for index in 0..num_block_erasures {
for block in &mut blocks {
block[index] = Scalar::ZERO;
}
missing_block_indices.push(index);
}
let codeword_with_erasures = blocks.into_iter().flatten().collect();
let maybe_recovered_poly_coeff = rs.recover_polynomial_coefficient(
codeword_with_erasures,
BlockErasureIndices(missing_block_indices),
);
if num_block_erasures <= rs.acceptable_num_block_erasures() {
let recovered_poly_coeff =
maybe_recovered_poly_coeff.expect("polynomial recovery failed");
assert_eq!(recovered_poly_coeff.len(), poly_coeff.len());
assert_eq!(recovered_poly_coeff, poly_coeff);
} else {
assert!(maybe_recovered_poly_coeff.is_err());
}
}
}
}