use crate::gf;
use ic_core::traits::{Algorithm, RandomSource, SelfTest};
use ic_core::{ensure, Result, Zeroize};
pub const MAX_SHARES: u8 = 255;
pub struct Shamir;
impl Algorithm for Shamir {
const ID: &'static str = "shamir-gf256";
const NAME: &'static str = "Shamir secret sharing over GF(2^8)";
}
pub fn split<R: RandomSource + ?Sized>(
secret: &[u8],
threshold: u8,
share_count: u8,
rng: &mut R,
out: &mut [u8],
) -> Result<()> {
ensure!(!secret.is_empty(), InvalidLength, "shamir secret is empty");
ensure!(
threshold >= 2,
InvalidParameter,
"shamir threshold must be at least 2"
);
ensure!(
share_count >= threshold,
InvalidParameter,
"shamir share count is below the threshold"
);
let len = secret.len();
let total = len.checked_mul(share_count as usize).ok_or(ic_core::err!(
InvalidLength,
"shamir secret too long for this many shares"
))?;
ensure!(
out.len() == total,
InvalidLength,
"shamir output must be share_count * secret.len() bytes"
);
let mut coefficients = [0u8; MAX_SHARES as usize - 1];
let coefficients = &mut coefficients[..threshold as usize - 1];
for (b, &s) in secret.iter().enumerate() {
if let Err(e) = rng.fill(coefficients) {
coefficients.zeroize();
out.zeroize();
return Err(e);
}
for i in 0..share_count as usize {
let x = (i + 1) as u8;
let mut y = 0u8;
for &c in coefficients.iter().rev() {
y = gf::mul(y, x) ^ c;
}
out[i * len + b] = gf::mul(y, x) ^ s;
}
}
coefficients.zeroize();
Ok(())
}
pub fn combine(shares: &[(u8, &[u8])], out: &mut [u8]) -> Result<()> {
ensure!(
shares.len() >= 2,
InvalidParameter,
"shamir needs at least two shares"
);
for (n, (index, bytes)) in shares.iter().enumerate() {
ensure!(
*index != 0,
InvalidParameter,
"shamir share index 0 is the secret itself"
);
ensure!(
bytes.len() == out.len(),
InvalidLength,
"shamir shares differ in length"
);
ensure!(
shares[..n].iter().all(|(other, _)| other != index),
InvalidParameter,
"shamir share indices repeat"
);
}
out.zeroize();
for (i, (xi, yi)) in shares.iter().enumerate() {
let mut numerator = 1u8;
let mut denominator = 1u8;
for (j, (xj, _)) in shares.iter().enumerate() {
if i != j {
numerator = gf::mul(numerator, *xj);
denominator = gf::mul(denominator, xj ^ xi);
}
}
let basis = gf::mul(numerator, gf::inv(denominator));
for (o, y) in out.iter_mut().zip(yi.iter()) {
*o ^= gf::mul(*y, basis);
}
}
Ok(())
}
#[cfg(feature = "std")]
pub struct Share {
pub index: u8,
pub value: std::vec::Vec<u8>,
}
#[cfg(feature = "std")]
impl Drop for Share {
fn drop(&mut self) {
self.value.zeroize();
}
}
#[cfg(feature = "std")]
impl core::fmt::Debug for Share {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Share")
.field("index", &self.index)
.finish_non_exhaustive()
}
}
#[cfg(feature = "std")]
pub fn split_vec<R: RandomSource + ?Sized>(
secret: &[u8],
threshold: u8,
share_count: u8,
rng: &mut R,
) -> Result<std::vec::Vec<Share>> {
let len = secret.len();
let total = len.checked_mul(share_count as usize).ok_or(ic_core::err!(
InvalidLength,
"shamir secret too long for this many shares"
))?;
let mut all = ic_core::Zeroizing::new(std::vec![0u8; total]);
split(secret, threshold, share_count, rng, all.get_mut())?;
Ok(all
.get()
.chunks(len)
.enumerate()
.map(|(i, bytes)| Share {
index: (i + 1) as u8,
value: bytes.to_vec(),
})
.collect())
}
#[cfg(feature = "std")]
pub fn combine_vec(shares: &[Share]) -> Result<ic_core::Zeroizing<std::vec::Vec<u8>>> {
let len = shares.first().map_or(0, |s| s.value.len());
let parts: std::vec::Vec<(u8, &[u8])> =
shares.iter().map(|s| (s.index, &s.value[..])).collect();
let mut out = ic_core::Zeroizing::new(std::vec![0u8; len]);
combine(&parts, out.get_mut())?;
Ok(out)
}
pub(crate) struct Stream(pub(crate) u8);
impl RandomSource for Stream {
fn fill(&mut self, out: &mut [u8]) -> Result<()> {
for b in out.iter_mut() {
self.0 = self.0.wrapping_mul(29).wrapping_add(7);
*b = self.0;
}
Ok(())
}
}
impl SelfTest for Shamir {
fn self_test() -> Result<()> {
let mut secret = [0u8; 32];
for (i, b) in secret.iter_mut().enumerate() {
*b = i as u8;
}
let mut shares = [0u8; 5 * 32];
split(&secret, 3, 5, &mut Stream(0x11), &mut shares)?;
let mut want = [0u8; 64];
ic_core::codec::hex_decode(KAT_SHARES_1_2, &mut want)?;
ensure!(
ic_core::ct::verify(&shares[..64], &want),
SelfTestFailed,
"shamir: shares differ from the reference"
);
let mut recovered = [0u8; 32];
combine(
&[
(2, &shares[32..64]),
(4, &shares[96..128]),
(5, &shares[128..]),
],
&mut recovered,
)?;
ensure!(
ic_core::ct::verify(&recovered, &secret),
SelfTestFailed,
"shamir: recovery differs"
);
Ok(())
}
}
const KAT_SHARES_1_2: &[u8] = b"5ff2a5a02b96c144f71a6d28237e898c4fa2f5f0fb065154278abd78732e191c69afee93f0ed7a6a5af11d0dc3b352b422ffbec32026b1218aa156eb93552fe4";
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_self_test_passes() {
Shamir::self_test().unwrap();
}
#[test]
fn every_threshold_subset_recovers_the_secret() {
let secret = *b"a 24-byte secret for k=3";
let (k, n) = (3u8, 6u8);
let mut shares = [0u8; 6 * 24];
split(&secret, k, n, &mut Stream(0x5d), &mut shares).unwrap();
let share = |i: usize| ((i + 1) as u8, &shares[i * 24..(i + 1) * 24]);
let mut subsets = 0;
for mask in 1u32..(1 << n) {
let picked: std::vec::Vec<_> = (0..n as usize)
.filter(|i| mask & (1 << i) != 0)
.map(share)
.collect();
if picked.len() < 2 {
continue;
}
let mut out = [0u8; 24];
combine(&picked, &mut out).unwrap();
if picked.len() >= k as usize {
assert_eq!(out, secret, "subset {mask:06b}");
subsets += 1;
} else {
assert_ne!(out, secret, "subset {mask:06b}, below the threshold");
}
}
assert_eq!(subsets, 20 + 15 + 6 + 1);
}
#[test]
fn bad_parameters_and_shares_are_refused() {
let mut out = [0u8; 10];
let mut rng = Stream(1);
assert!(split(b"", 2, 2, &mut rng, &mut []).is_err(), "empty secret");
assert!(
split(b"ab", 1, 5, &mut rng, &mut out).is_err(),
"threshold 1"
);
assert!(
split(b"ab", 3, 2, &mut rng, &mut out[..4]).is_err(),
"count below threshold"
);
assert!(
split(b"ab", 2, 5, &mut rng, &mut out[..9]).is_err(),
"wrong output length"
);
let a = [1u8, 2];
let b = [3u8, 4];
let mut two = [0u8; 2];
assert!(combine(&[(1, &a)], &mut two).is_err(), "one share");
assert!(combine(&[(0, &a), (1, &b)], &mut two).is_err(), "index 0");
assert!(
combine(&[(1, &a), (1, &b)], &mut two).is_err(),
"repeated index"
);
assert!(
combine(&[(1, &a), (2, &b[..1])], &mut two).is_err(),
"length mismatch"
);
combine(&[(1, &a), (2, &b)], &mut two).unwrap();
}
#[test]
fn a_failing_rng_wipes_the_output() {
struct Fails(u8);
impl RandomSource for Fails {
fn fill(&mut self, out: &mut [u8]) -> Result<()> {
if self.0 == 0 {
return Err(ic_core::err!(EntropyFailure, "test"));
}
self.0 -= 1;
out.fill(0x77);
Ok(())
}
}
let mut out = [0u8; 3 * 4];
assert!(split(b"keys", 2, 3, &mut Fails(2), &mut out).is_err());
assert_eq!(out, [0u8; 12]);
}
#[cfg(feature = "std")]
#[test]
fn the_vec_forms_round_trip() {
let shares = split_vec(b"vec secret", 2, 4, &mut Stream(9)).unwrap();
assert_eq!(shares.len(), 4);
assert_eq!(shares[3].index, 4);
let recovered = combine_vec(&shares[2..]).unwrap();
assert_eq!(&recovered.get()[..], b"vec secret");
}
}