use crate::common::{
lagrange_interpolate,
share::{shamir::NonRobustShare, ShareError},
SecretSharingScheme, ShamirShare,
};
use ark_ff::{FftField, Zero};
use ark_poly::{
univariate::{DenseOrSparsePolynomial, DensePolynomial},
DenseUVPolynomial, EvaluationDomain, Polynomial,
};
use ark_std::rand::Rng;
use std::collections::HashSet;
use std::marker::PhantomData;
use super::*;
#[derive(Clone, Debug, PartialEq)]
pub struct Robust;
pub type RobustShare<T> = ShamirShare<T, 1, Robust>;
impl<F: FftField> RobustShare<F> {
pub fn new(share: F, id: usize, degree: usize) -> Self {
ShamirShare {
share: [share],
id,
degree,
_sharetype: PhantomData,
}
}
}
impl<F: FftField> From<NonRobustShare<F>> for RobustShare<F> {
fn from(non_robust: NonRobustShare<F>) -> Self {
RobustShare {
share: non_robust.share,
id: non_robust.id,
degree: non_robust.degree,
_sharetype: PhantomData,
}
}
}
impl<F: FftField> SecretSharingScheme<F> for RobustShare<F> {
type SecretType = F;
type Error = InterpolateError;
fn compute_shares(
secret: Self::SecretType,
n: usize,
degree: usize,
_ids: Option<&[usize]>,
rng: &mut impl Rng,
) -> Result<Vec<RobustShare<F>>, InterpolateError> {
if n <= degree {
return Err(InterpolateError::InvalidInput(format!(
"Number of shares ({}) must be greater than threshold ({})",
n, degree
)));
}
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.ok_or_else(|| InterpolateError::NoSuitableDomain(n))?;
let mut poly = DensePolynomial::<F>::rand(degree, rng);
poly.coeffs[0] = secret;
let evals = domain.fft(&poly);
let shares: Vec<RobustShare<F>> = evals
.iter()
.take(n)
.enumerate()
.map(|(i, &eval)| RobustShare::new(eval, i, degree))
.collect();
Ok(shares)
}
fn recover_secret(
shares: &[Self],
n: usize,
t: usize,
) -> Result<(Vec<Self::SecretType>, Self::SecretType), InterpolateError> {
if n < 3 * t + 1 {
return Err(InterpolateError::InvalidInput(format!(
"n ({}) must be >= 3t + 1 ({}) for Byzantine fault tolerance",
n,
3 * t + 1
)));
}
if shares.is_empty() {
return Err(InterpolateError::InvalidInput(
"Share slice is empty".to_string(),
));
}
let degree = shares[0].degree;
if !shares.iter().all(|share| share.degree == degree) {
return Err(InterpolateError::ShareError(ShareError::DegreeMismatch));
};
let mut seen = HashSet::new();
if !shares.iter().all(|s| seen.insert(s.id)) {
return Err(InterpolateError::InvalidInput(
"Duplicate share Ids".to_string(),
));
}
for s in shares {
if s.id >= n {
return Err(InterpolateError::InvalidInput(format!(
"Share id {} is out of range: expected 0 <= id < n (n = {})",
s.id, n
)));
}
}
let share_len = shares.len();
if share_len < degree + t + 1 {
return Err(InterpolateError::InvalidInput(format!(
"Not enough shares provided ({}) to attempt decoding for t={}. At least {} shares are required.",
share_len,
t,
degree + t + 1
)));
}
let mut sorted_shares = shares.to_vec();
sorted_shares.sort_by_key(|i| i.id);
if let Ok(poly) = robust_interpolate_fnt(t, n, &sorted_shares[..degree + t + 1]) {
return Ok((poly.coeffs.clone(), poly.evaluate(&F::zero())));
}
let poly = oec_decode(n, t, sorted_shares.to_vec());
match poly {
Ok(p) => Ok((p.0.coeffs.clone(), p.1)),
Err(e) => Err(e),
}
}
}
fn poly_derivative<F: FftField>(poly: &DensePolynomial<F>) -> DensePolynomial<F> {
if poly.coeffs.len() <= 1 {
return DensePolynomial::from_coefficients_vec(vec![]);
}
let derived_coeffs: Vec<F> = poly
.coeffs
.iter()
.enumerate()
.skip(1)
.map(|(i, coeff)| F::from(i as u64) * coeff)
.collect();
DensePolynomial::from_coefficients_vec(derived_coeffs)
}
fn div_with_remainder<F: FftField>(
numerator: &DensePolynomial<F>,
denominator: &DensePolynomial<F>,
) -> Result<(DensePolynomial<F>, DensePolynomial<F>), InterpolateError> {
let a = DenseOrSparsePolynomial::from(numerator.clone());
let b = DenseOrSparsePolynomial::from(denominator.clone());
let (q, r) = a
.divide_with_q_and_r(&b)
.ok_or_else(|| InterpolateError::PolynomialOperationError("Division failed".to_string()))?;
Ok((DensePolynomial::from(q), DensePolynomial::from(r)))
}
fn robust_interpolate_fnt<F: FftField>(
t: usize,
n: usize,
shares: &[RobustShare<F>],
) -> Result<DensePolynomial<F>, InterpolateError> {
let degree = shares[0].degree;
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.ok_or(InterpolateError::NoSuitableDomain(n))?;
let subset = &shares[..=degree];
let xs: Vec<F> = subset.iter().map(|s| domain.element(s.id)).collect();
let ys: Vec<F> = subset.iter().map(|s| s.share[0]).collect();
let mut a_poly = DensePolynomial::from_coefficients_slice(&[F::one()]);
for &x in &xs {
let xi_poly = DensePolynomial::from_coefficients_slice(&[-x, F::one()]);
a_poly = &a_poly * &xi_poly;
}
let a_derivative = poly_derivative(&a_poly);
let mut interpolated = DensePolynomial::from_coefficients_slice(&[F::zero()]);
for (i, &x_i) in xs.iter().enumerate() {
let denom = a_derivative.evaluate(&x_i);
if denom.is_zero() {
return Err(InterpolateError::PolynomialOperationError(
"Denominator evaluated to zero during interpolation basis calculation".into(),
));
}
let scalar = ys[i] / denom;
let term_divisor = DensePolynomial::from_coefficients_slice(&[-x_i, F::one()]);
let (basis_poly, rem) = div_with_remainder(&a_poly, &term_divisor)?;
if !rem.is_zero() {
return Err(InterpolateError::PolynomialOperationError(
"A(x) not perfectly divisible by (x - x_i)".into(),
));
}
interpolated = &interpolated + &(&basis_poly * scalar);
}
let valid_count = shares
.iter()
.map(|s| (domain.element(s.id), s.share))
.filter(|(x, y)| interpolated.evaluate(x) == y[0])
.count();
if valid_count >= degree + t + 1 {
Ok(interpolated)
} else {
Err(InterpolateError::DecodingError(
"Not enough shares matched the interpolated polynomial".into(),
))
}
}
pub fn batch_recover_secret<F: FftField>(
evals_by_sender: &[(usize, Vec<F>)],
n: usize,
degree: usize,
t: usize,
) -> Result<Vec<Vec<F>>, InterpolateError> {
if n < 3 * t + 1 {
return Err(InterpolateError::InvalidInput(format!(
"n ({}) must be >= 3t + 1 ({}) for Byzantine fault tolerance",
n,
3 * t + 1
)));
}
if evals_by_sender.is_empty() {
return Err(InterpolateError::InvalidInput(
"No evaluations provided".to_string(),
));
}
let batch_len = evals_by_sender[0].1.len();
if batch_len == 0 {
return Err(InterpolateError::InvalidInput("Empty batch".to_string()));
}
if !evals_by_sender.iter().all(|(_, v)| v.len() == batch_len) {
return Err(InterpolateError::InvalidInput(
"Inconsistent batch widths".to_string(),
));
}
let mut sorted: Vec<(usize, &Vec<F>)> =
evals_by_sender.iter().map(|(id, v)| (*id, v)).collect();
sorted.sort_by_key(|(id, _)| *id);
let mut seen = HashSet::new();
for (id, _) in &sorted {
if !seen.insert(*id) {
return Err(InterpolateError::InvalidInput(
"Duplicate sender id".to_string(),
));
}
if *id >= n {
return Err(InterpolateError::InvalidInput(format!(
"Sender id {} out of range (n = {})",
id, n
)));
}
}
let needed = degree + t + 1;
if sorted.len() < needed {
return Err(InterpolateError::InvalidInput(format!(
"Not enough evaluations ({}) for degree {} and t {} (need {})",
sorted.len(),
degree,
t,
needed
)));
}
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.ok_or(InterpolateError::NoSuitableDomain(n))?;
let m = degree + 1;
let subset_xs: Vec<F> = (0..m).map(|i| domain.element(sorted[i].0)).collect();
let mut a_poly = DensePolynomial::from_coefficients_slice(&[F::one()]);
for &x in &subset_xs {
a_poly = &a_poly * &DensePolynomial::from_coefficients_slice(&[-x, F::one()]);
}
let a_derivative = poly_derivative(&a_poly);
let mut basis: Vec<DensePolynomial<F>> = Vec::with_capacity(m);
for &x_i in &subset_xs {
let denom = a_derivative.evaluate(&x_i);
if denom.is_zero() {
return Err(InterpolateError::PolynomialOperationError(
"Denominator evaluated to zero during interpolation basis calculation".into(),
));
}
let inv = F::one() / denom;
let divisor = DensePolynomial::from_coefficients_slice(&[-x_i, F::one()]);
let (basis_poly, rem) = div_with_remainder(&a_poly, &divisor)?;
if !rem.is_zero() {
return Err(InterpolateError::PolynomialOperationError(
"A(x) not perfectly divisible by (x - x_i)".into(),
));
}
basis.push(&basis_poly * inv);
}
let verify_xs: Vec<F> = (0..needed).map(|s| domain.element(sorted[s].0)).collect();
let basis_coeffs: Vec<&[F]> = basis.iter().map(|p| p.coeffs.as_slice()).collect();
let mut verify_matrix = vec![F::zero(); needed * m]; for s in 0..needed {
let xs = verify_xs[s];
let row = &mut verify_matrix[s * m..(s + 1) * m];
for (i, slot) in row.iter_mut().enumerate() {
*slot = basis[i].evaluate(&xs);
}
}
let mut results: Vec<Vec<F>> = Vec::with_capacity(batch_len);
for c in 0..batch_len {
let mut ok = true;
for s in 0..needed {
let row = &verify_matrix[s * m..(s + 1) * m];
let mut acc = F::zero();
for i in 0..m {
acc += row[i] * sorted[i].1[c];
}
if acc != sorted[s].1[c] {
ok = false;
break;
}
}
if ok {
let mut coeffs = vec![F::zero(); degree + 1];
for k in 0..=degree {
let mut acc = F::zero();
for i in 0..m {
let bik = basis_coeffs[i].get(k).copied().unwrap_or(F::zero());
acc += bik * sorted[i].1[c];
}
coeffs[k] = acc;
}
results.push(coeffs);
} else {
let shares: Vec<RobustShare<F>> = sorted
.iter()
.map(|(id, vals)| RobustShare::new(vals[c], *id, degree))
.collect();
let (coeffs, _) = RobustShare::recover_secret(&shares, n, t)?;
results.push(coeffs);
}
}
Ok(results)
}
fn gao_rs_decode<F: FftField>(
received: &[F],
k: usize,
n: usize,
erasure_positions: &[usize],
) -> Result<Vec<F>, InterpolateError> {
if k > n {
return Err(InterpolateError::InvalidInput(format!(
"k ({}) must be less than or equal to n ({})",
k, n
)));
}
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.ok_or(InterpolateError::NoSuitableDomain(n))?;
let s_set: HashSet<usize> = erasure_positions.iter().copied().collect();
let s = s_set.len();
let s_poly = s_set.iter().fold(
DensePolynomial::from_coefficients_slice(&[F::one()]),
|acc, &i| {
let xi = domain.element(i);
&acc * &DensePolynomial::from_coefficients_slice(&[-xi, F::one()])
},
);
let known_points: Vec<_> = (0..n)
.filter(|i| !s_set.contains(i))
.map(|i| (domain.element(i), received[i]))
.collect();
let (x_vals, y_vals): (Vec<F>, Vec<F>) = known_points.iter().cloned().unzip();
let g1 = lagrange_interpolate(&x_vals, &y_vals)?;
let x_a_prod = compute_g0_from_domain(n);
let g0 = &x_a_prod / &s_poly;
let threshold = (n - s + k) / 2;
let (mut r0, mut r1) = (g0.clone(), g1.clone());
let (mut s0, mut s1) = (
DensePolynomial::from_coefficients_slice(&[F::one()]),
DensePolynomial::zero(),
);
let (mut t0, mut t1) = (
DensePolynomial::zero(),
DensePolynomial::from_coefficients_slice(&[F::one()]),
);
while r1.degree() >= threshold {
let q = &r0 / &r1;
let r = &r0 - &q * &r1;
let s = &s0 - &q * &s1;
let t = &t0 - &q * &t1;
r0 = r1;
r1 = r;
s0 = s1;
s1 = s;
t0 = t1;
t1 = t;
}
let g = r1;
let v = t1;
let quotient = &g / &v;
let remainder = &g - "ient * &v;
if remainder.is_zero() && quotient.degree() < k {
Ok(quotient.coeffs.clone())
} else {
Err(InterpolateError::DecodingError(
"Failed to recover message polynomial from g(x)/v(x)".into(),
))
}
}
pub fn compute_g0_from_domain<F: FftField>(n: usize) -> DensePolynomial<F> {
if let Some(cached) = crate::common::get_cached_g0_polynomial::<F>(n) {
return cached;
}
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.expect("Domain of size n must exist over the field");
let evaluation_points: Vec<F> = domain.elements().collect();
let mut g0 = DensePolynomial::from_coefficients_slice(&[F::one()]);
for ai in evaluation_points.iter().take(n) {
let factor = DensePolynomial::from_coefficients_slice(&[-*ai, F::one()]); g0 = &g0 * &factor;
}
crate::common::store_g0_polynomial(n, g0.clone());
g0
}
fn oec_decode<F: FftField>(
n: usize,
t: usize,
shares: Vec<RobustShare<F>>,
) -> Result<(DensePolynomial<F>, F), InterpolateError> {
let domain = crate::common::get_or_create_evaluation_domain::<F>(n)
.ok_or(InterpolateError::NoSuitableDomain(n))?;
let degree = shares[0].degree;
for r in 1..=t {
let required = degree + t + 1 + r;
if shares.len() < required {
break;
}
let subset = &shares[..required];
let mut received = vec![F::zero(); n];
let mut erasures = vec![];
for i in 0..n {
if let Some(val) = subset.iter().find(|s| s.id == i) {
received[i] = val.share[0];
} else {
erasures.push(i);
}
}
if let Ok(coeffs) = gao_rs_decode(&received, degree + 1, n, &erasures) {
let poly = DensePolynomial::from_coefficients_vec(coeffs);
let matched = subset
.iter()
.filter(|s| poly.evaluate(&domain.element(s.id)) == s.share[0])
.count();
if matched >= degree + t + 1 {
return Ok((poly.clone(), poly.evaluate(&F::zero())));
}
}
}
Err(InterpolateError::DecodingError(
"Online Error Correction failed to find a valid polynomial".into(),
))
}
#[cfg(test)]
mod tests {
use super::*;
use ark_bls12_381::Fr;
use ark_poly::GeneralEvaluationDomain;
use ark_std::test_rng;
#[test]
fn test_poly_derivative() {
let coeffs = vec![Fr::from(3), Fr::from(2), Fr::from(1)]; let poly = DensePolynomial::from_coefficients_vec(coeffs);
let deriv = poly_derivative(&poly);
let expected = DensePolynomial::from_coefficients_vec(vec![Fr::from(2), Fr::from(2)]); assert_eq!(deriv, expected);
}
#[test]
fn test_robust_interpolate_fnt_optimistic_case() {
use ark_bls12_381::Fr;
use ark_poly::univariate::DensePolynomial;
let n = 16;
let t = 2;
let domain = GeneralEvaluationDomain::<Fr>::new(n).unwrap();
let coeffs = vec![Fr::from(7u32), Fr::from(3u32), Fr::from(5u32)];
let poly = DensePolynomial::from_coefficients_vec(coeffs);
let shares: Vec<RobustShare<Fr>> = (0..n)
.map(|i| {
let x = domain.element(i);
let y = poly.evaluate(&x);
RobustShare::new(y, i, t)
})
.collect();
let used_shares = shares[..(2 * t + 1)].to_vec();
let result = robust_interpolate_fnt(t, n, &used_shares);
assert!(result.is_ok(), "Optimistic interpolation failed");
let recovered = result.unwrap();
assert_eq!(recovered.coeffs.len(), poly.coeffs.len());
for (a, b) in recovered.coeffs.iter().zip(poly.coeffs.iter()) {
assert_eq!(a, b, "Mismatch in optimistic recovery");
}
}
#[test]
fn test_reed_solomon_erasure() {
let mut rng = test_rng();
let t = 2;
let n = 8;
let secret = Fr::from(42u32);
let ids: Vec<usize> = (0..n).collect();
let shares = RobustShare::compute_shares(secret, n, t, Some(&ids), &mut rng).unwrap();
let mut erased: Vec<Fr> = shares.iter().map(|a| a.share[0]).collect();
let erasures = vec![1, 2];
for &i in &erasures {
erased[i] = Fr::zero();
}
let decoded_erasure = gao_rs_decode(&erased, t + 1, n, &erasures);
let recovered_secret = decoded_erasure.unwrap()[0];
assert_eq!(
recovered_secret, secret,
"Failed to decode with known erasures"
);
}
#[test]
fn test_reed_solomon_error() {
let mut rng = test_rng();
let t = 2;
let n = 10;
let secret = Fr::from(42u32);
let shares = RobustShare::compute_shares(secret, n, t, None, &mut rng).unwrap();
let mut corrupted: Vec<Fr> = shares.iter().map(|a| a.share[0]).collect();
corrupted[2] += Fr::from(5u64);
corrupted[4] += Fr::from(3u64);
let decoded = gao_rs_decode(&corrupted, t + 1, n, &[]);
let recovered_secret = decoded.unwrap()[0];
assert_eq!(
recovered_secret, secret,
"Failed to decode with known erasures"
);
}
#[test]
fn test_reed_solomon_error_all_triples() {
use itertools::Itertools;
let mut rng = test_rng();
let t = 3;
let n = 10;
let secret = Fr::from(42u32);
let shares = RobustShare::compute_shares(secret, n, t, None, &mut rng).unwrap();
for triple in (0..n).combinations(3) {
let mut corrupted: Vec<Fr> = shares.iter().map(|a| a.share[0]).collect();
corrupted[triple[0]] += Fr::from(5u64);
corrupted[triple[1]] += Fr::from(3u64);
corrupted[triple[2]] += Fr::from(3u64);
let decoded = gao_rs_decode(&corrupted, t + 1, n, &[]).unwrap();
let recovered_secret = decoded[0];
assert_eq!(
recovered_secret, secret,
"Failed to decode when corrupting indices {:?}",
triple
);
}
}
#[test]
fn test_oec_protocol() {
use ark_bls12_381::Fr;
use ark_std::test_rng;
let mut rng = test_rng();
let t = 2;
let n = 10;
let secret = Fr::from(42u32);
let ids: Vec<usize> = (0..n).collect();
let mut shares = RobustShare::compute_shares(secret, n, t, Some(&ids), &mut rng).unwrap();
shares[0].share[0] += Fr::from(999u64);
shares[5].share[0] += Fr::from(999u64);
let result = oec_decode(n, t, shares.clone());
assert!(
result.is_ok(),
"Decoding failed despite sufficient honest shares"
);
let (_, recovered_zero) = result.unwrap();
assert_eq!(
recovered_zero, secret,
"Recovered polynomial does not match the original"
);
}
#[test]
fn test_robust_interpolate_full() {
use ark_bls12_381::Fr;
use ark_std::test_rng;
let mut rng = test_rng();
let t = 3;
let n = 10;
let secret = Fr::from(42u32);
let ids: Vec<usize> = (0..n).collect();
let mut shares = RobustShare::compute_shares(secret, n, t, Some(&ids), &mut rng).unwrap();
let corruption_indices = [1, 4];
for &i in &corruption_indices {
shares[i] = (shares[i].clone()
+ RobustShare {
share: [Fr::from(7u64)],
id: i,
degree: t,
_sharetype: PhantomData,
})
.unwrap();
}
let result = RobustShare::recover_secret(&shares, n, t);
assert!(
result.is_ok(),
"robust_interpolate failed despite valid parameters"
);
let (_, val_at_zero) = result.unwrap();
assert_eq!(val_at_zero, secret, "Evaluation at zero incorrect");
}
#[test]
fn test_robust_interpolate_all_corruption_combinations() {
use ark_bls12_381::Fr;
use ark_std::test_rng;
use itertools::Itertools;
let mut rng = test_rng();
let t = 2;
let n = 7;
let secret = Fr::from(42u32);
let base_shares = RobustShare::compute_shares(secret, n, t, None, &mut rng).unwrap();
for k in 1..=t {
for corruption_indices in (0..n).combinations(k) {
let mut shares = base_shares.clone();
for &i in &corruption_indices {
shares[i].share[0] += Fr::from(999u64); }
let result = RobustShare::recover_secret(&shares, n, t);
if result.is_err() {
eprintln!("Failed for corrupted indices: {:?}", corruption_indices);
}
assert!(
result.is_ok(),
"Decoding failed for corrupted indices: {:?}",
corruption_indices
);
let (_, val_at_zero) = result.unwrap();
assert_eq!(
val_at_zero, secret,
"Incorrect recovery at zero for {:?}",
corruption_indices
);
}
}
}
#[test]
fn test_batch_recover_secret_matches_per_chunk() {
use ark_bls12_381::Fr;
use ark_poly::EvaluationDomain;
let mut rng = test_rng();
let n = 10;
let t = 3;
let degree = t;
let batch_len = 16;
let polys: Vec<DensePolynomial<Fr>> = (0..batch_len)
.map(|_| DensePolynomial::<Fr>::rand(degree, &mut rng))
.collect();
let domain = GeneralEvaluationDomain::<Fr>::new(n).unwrap();
let mut evals_by_sender: Vec<(usize, Vec<Fr>)> = (0..n)
.map(|id| {
let x = domain.element(id);
(id, polys.iter().map(|p| p.evaluate(&x)).collect())
})
.collect();
evals_by_sender.reverse();
let batched = batch_recover_secret(&evals_by_sender, n, degree, t).unwrap();
assert_eq!(batched.len(), batch_len);
for c in 0..batch_len {
let shares: Vec<RobustShare<Fr>> = evals_by_sender
.iter()
.map(|(id, vals)| RobustShare::new(vals[c], *id, degree))
.collect();
let (mut per_chunk, _) = RobustShare::recover_secret(&shares, n, t).unwrap();
per_chunk.resize(degree + 1, Fr::zero());
assert_eq!(
batched[c], per_chunk,
"chunk {c} differs from recover_secret"
);
assert_eq!(
batched[c][0], polys[c].coeffs[0],
"chunk {c} secret mismatch"
);
}
}
#[test]
fn test_batch_recover_secret_with_corruption() {
use ark_bls12_381::Fr;
use ark_poly::EvaluationDomain;
let mut rng = test_rng();
let n = 10;
let t = 3;
let degree = t;
let batch_len = 8;
let polys: Vec<DensePolynomial<Fr>> = (0..batch_len)
.map(|_| DensePolynomial::<Fr>::rand(degree, &mut rng))
.collect();
let domain = GeneralEvaluationDomain::<Fr>::new(n).unwrap();
let mut evals_by_sender: Vec<(usize, Vec<Fr>)> = (0..n)
.map(|id| {
let x = domain.element(id);
(id, polys.iter().map(|p| p.evaluate(&x)).collect())
})
.collect();
for bad in 0..t {
for c in 0..batch_len {
evals_by_sender[bad].1[c] += Fr::from((c as u64 + 1) * 7 + bad as u64);
}
}
let batched = batch_recover_secret(&evals_by_sender, n, degree, t).unwrap();
for c in 0..batch_len {
assert_eq!(
batched[c][0], polys[c].coeffs[0],
"corrupted chunk {c} secret mismatch"
);
}
}
}