Skip to main content

ic_cipher/
shamir.rs

1//! Shamir secret sharing over GF(2^8).
2//!
3//! A secret is split into `share_count` shares, any `threshold` of which
4//! recover it, and any fewer of which reveal nothing about it: each byte of
5//! the secret is the constant term of its own random polynomial of degree
6//! `threshold - 1`, and share `i` holds every polynomial's value at `x = i`.
7//! Recovery is Lagrange interpolation at zero.
8//!
9//! The field is the one AES uses, `x^8 + x^4 + x^3 + x + 1`, and the
10//! arithmetic is [`crate::gf`]'s: branch-free and table-free, so neither the
11//! coefficients nor the share bytes reach an address or a branch. Share
12//! indices and counts are public.
13//!
14//! # What it does not do
15//!
16//! Shares carry no integrity. A corrupted share, or fewer than `threshold`
17//! shares, recover a wrong secret rather than an error, because nothing in a
18//! share says what the right answer is. Split a random key rather than data,
19//! and use that key with an AEAD over the data: the AEAD is what detects a bad
20//! recovery. The share's index and the threshold have to travel with it, in
21//! whatever format the caller chooses; nothing here encodes them.
22//!
23//! # Verification
24//!
25//! Against an implementation written from Shamir's construction with its own
26//! field arithmetic (`scripts/gen_shamir_vectors.py`), share for share, and by
27//! recombining every subset of shares this crate's tests enumerate.
28
29use crate::gf;
30use ic_core::traits::{Algorithm, RandomSource, SelfTest};
31use ic_core::{ensure, Result, Zeroize};
32
33/// The most shares a split can make: the nonzero elements of GF(2^8).
34pub const MAX_SHARES: u8 = 255;
35
36/// Shamir secret sharing over GF(2^8), for the ontology and the self-test
37/// table.
38pub struct Shamir;
39
40impl Algorithm for Shamir {
41    const ID: &'static str = "shamir-gf256";
42    const NAME: &'static str = "Shamir secret sharing over GF(2^8)";
43}
44
45/// Split `secret` into `share_count` shares, any `threshold` of which recover
46/// it.
47///
48/// `out` is `share_count * secret.len()` bytes: share `i`, whose index is
49/// `i + 1`, is `out[i * secret.len()..(i + 1) * secret.len()]`. `rng` supplies
50/// `threshold - 1` fresh coefficients for every byte of the secret.
51///
52/// Requires `2 <= threshold <= share_count` (a threshold of one would make
53/// every share the secret) and a non-empty secret. If `rng` fails, `out` is
54/// wiped and the error returned.
55pub fn split<R: RandomSource + ?Sized>(
56    secret: &[u8],
57    threshold: u8,
58    share_count: u8,
59    rng: &mut R,
60    out: &mut [u8],
61) -> Result<()> {
62    ic_core::module::operational()?;
63    ensure!(!secret.is_empty(), InvalidLength, "shamir secret is empty");
64    ensure!(
65        threshold >= 2,
66        InvalidParameter,
67        "shamir threshold must be at least 2"
68    );
69    ensure!(
70        share_count >= threshold,
71        InvalidParameter,
72        "shamir share count is below the threshold"
73    );
74    let len = secret.len();
75    let total = len.checked_mul(share_count as usize).ok_or(ic_core::err!(
76        InvalidLength,
77        "shamir secret too long for this many shares"
78    ))?;
79    ensure!(
80        out.len() == total,
81        InvalidLength,
82        "shamir output must be share_count * secret.len() bytes"
83    );
84
85    let mut coefficients = [0u8; MAX_SHARES as usize - 1];
86    let coefficients = &mut coefficients[..threshold as usize - 1];
87    for (b, &s) in secret.iter().enumerate() {
88        if let Err(e) = rng.fill(coefficients) {
89            coefficients.zeroize();
90            out.zeroize();
91            return Err(e);
92        }
93        for i in 0..share_count as usize {
94            // Horner's rule from the highest coefficient down, ending on the
95            // secret byte as the constant term.
96            let x = (i + 1) as u8;
97            let mut y = 0u8;
98            for &c in coefficients.iter().rev() {
99                y = gf::mul(y, x) ^ c;
100            }
101            out[i * len + b] = gf::mul(y, x) ^ s;
102        }
103    }
104    coefficients.zeroize();
105    Ok(())
106}
107
108/// Recover a secret from shares given as `(index, bytes)`.
109///
110/// Pass at least `threshold` shares; more are fine. With fewer, the result is
111/// a wrong secret rather than an error -- see the module note. Every share
112/// must be `out.len()` bytes, and indices must be nonzero and distinct; at
113/// least two shares are required.
114pub fn combine(shares: &[(u8, &[u8])], out: &mut [u8]) -> Result<()> {
115    ic_core::module::operational()?;
116    ensure!(
117        shares.len() >= 2,
118        InvalidParameter,
119        "shamir needs at least two shares"
120    );
121    for (n, (index, bytes)) in shares.iter().enumerate() {
122        ensure!(
123            *index != 0,
124            InvalidParameter,
125            "shamir share index 0 is the secret itself"
126        );
127        ensure!(
128            bytes.len() == out.len(),
129            InvalidLength,
130            "shamir shares differ in length"
131        );
132        ensure!(
133            shares[..n].iter().all(|(other, _)| other != index),
134            InvalidParameter,
135            "shamir share indices repeat"
136        );
137    }
138
139    out.zeroize();
140    for (i, (xi, yi)) in shares.iter().enumerate() {
141        // The Lagrange basis polynomial for share i, evaluated at zero:
142        // prod over j != i of x_j / (x_j - x_i), subtraction being XOR. It
143        // depends only on the public indices.
144        let mut numerator = 1u8;
145        let mut denominator = 1u8;
146        for (j, (xj, _)) in shares.iter().enumerate() {
147            if i != j {
148                numerator = gf::mul(numerator, *xj);
149                denominator = gf::mul(denominator, xj ^ xi);
150            }
151        }
152        let basis = gf::mul(numerator, gf::inv(denominator));
153        for (o, y) in out.iter_mut().zip(yi.iter()) {
154            *o ^= gf::mul(*y, basis);
155        }
156    }
157    Ok(())
158}
159
160/// A share, with its index, as [`split_vec`] returns them. Wiped on drop.
161#[cfg(feature = "std")]
162pub struct Share {
163    /// The share's index, `1..=255`: its `x`.
164    pub index: u8,
165    /// The share's bytes, as long as the secret.
166    pub value: std::vec::Vec<u8>,
167}
168
169#[cfg(feature = "std")]
170impl Drop for Share {
171    fn drop(&mut self) {
172        self.value.zeroize();
173    }
174}
175
176#[cfg(feature = "std")]
177impl core::fmt::Debug for Share {
178    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
179        f.debug_struct("Share")
180            .field("index", &self.index)
181            .finish_non_exhaustive()
182    }
183}
184
185/// [`split`], returning the shares.
186#[cfg(feature = "std")]
187pub fn split_vec<R: RandomSource + ?Sized>(
188    secret: &[u8],
189    threshold: u8,
190    share_count: u8,
191    rng: &mut R,
192) -> Result<std::vec::Vec<Share>> {
193    let len = secret.len();
194    let total = len.checked_mul(share_count as usize).ok_or(ic_core::err!(
195        InvalidLength,
196        "shamir secret too long for this many shares"
197    ))?;
198    let mut all = ic_core::Zeroizing::new(std::vec![0u8; total]);
199    split(secret, threshold, share_count, rng, all.get_mut())?;
200    Ok(all
201        .get()
202        .chunks(len)
203        .enumerate()
204        .map(|(i, bytes)| Share {
205            index: (i + 1) as u8,
206            value: bytes.to_vec(),
207        })
208        .collect())
209}
210
211/// [`combine`], returning the secret in a buffer wiped on drop.
212#[cfg(feature = "std")]
213pub fn combine_vec(shares: &[Share]) -> Result<ic_core::Zeroizing<std::vec::Vec<u8>>> {
214    let len = shares.first().map_or(0, |s| s.value.len());
215    let parts: std::vec::Vec<(u8, &[u8])> =
216        shares.iter().map(|s| (s.index, &s.value[..])).collect();
217    let mut out = ic_core::Zeroizing::new(std::vec![0u8; len]);
218    combine(&parts, out.get_mut())?;
219    Ok(out)
220}
221
222/// The deterministic coefficient stream `scripts/gen_shamir_vectors.py`
223/// uses: `s = 29 * s + 7 mod 256`. For known-answer tests only.
224pub(crate) struct Stream(pub(crate) u8);
225
226impl RandomSource for Stream {
227    fn fill(&mut self, out: &mut [u8]) -> Result<()> {
228        for b in out.iter_mut() {
229            self.0 = self.0.wrapping_mul(29).wrapping_add(7);
230            *b = self.0;
231        }
232        Ok(())
233    }
234}
235
236impl SelfTest for Shamir {
237    /// A 3-of-5 split of 32 bytes under a fixed coefficient stream, compared
238    /// with the reference implementation's first two shares, then recovered
239    /// from shares 2, 4 and 5.
240    fn self_test() -> Result<()> {
241        let mut secret = [0u8; 32];
242        for (i, b) in secret.iter_mut().enumerate() {
243            *b = i as u8;
244        }
245        let mut shares = [0u8; 5 * 32];
246        split(&secret, 3, 5, &mut Stream(0x11), &mut shares)?;
247        let mut want = [0u8; 64];
248        ic_core::codec::hex_decode(KAT_SHARES_1_2, &mut want)?;
249        ensure!(
250            ic_core::ct::verify(&shares[..64], &want),
251            SelfTestFailed,
252            "shamir: shares differ from the reference"
253        );
254        let mut recovered = [0u8; 32];
255        combine(
256            &[
257                (2, &shares[32..64]),
258                (4, &shares[96..128]),
259                (5, &shares[128..]),
260            ],
261            &mut recovered,
262        )?;
263        ensure!(
264            ic_core::ct::verify(&recovered, &secret),
265            SelfTestFailed,
266            "shamir: recovery differs"
267        );
268        Ok(())
269    }
270}
271
272/// Shares 1 and 2 of the first case in `testvectors/shamir-gf256.json`.
273const KAT_SHARES_1_2: &[u8] = b"5ff2a5a02b96c144f71a6d28237e898c4fa2f5f0fb065154278abd78732e191c69afee93f0ed7a6a5af11d0dc3b352b422ffbec32026b1218aa156eb93552fe4";
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278
279    #[test]
280    fn the_self_test_passes() {
281        Shamir::self_test().unwrap();
282    }
283
284    /// Every subset of `threshold` shares recovers the secret, as does every
285    /// larger subset; a subset one short does not, for this split.
286    #[test]
287    fn every_threshold_subset_recovers_the_secret() {
288        let secret = *b"a 24-byte secret for k=3";
289        let (k, n) = (3u8, 6u8);
290        let mut shares = [0u8; 6 * 24];
291        split(&secret, k, n, &mut Stream(0x5d), &mut shares).unwrap();
292        let share = |i: usize| ((i + 1) as u8, &shares[i * 24..(i + 1) * 24]);
293
294        let mut subsets = 0;
295        for mask in 1u32..(1 << n) {
296            let picked: std::vec::Vec<_> = (0..n as usize)
297                .filter(|i| mask & (1 << i) != 0)
298                .map(share)
299                .collect();
300            if picked.len() < 2 {
301                continue;
302            }
303            let mut out = [0u8; 24];
304            combine(&picked, &mut out).unwrap();
305            if picked.len() >= k as usize {
306                assert_eq!(out, secret, "subset {mask:06b}");
307                subsets += 1;
308            } else {
309                assert_ne!(out, secret, "subset {mask:06b}, below the threshold");
310            }
311        }
312        // C(6,3) + C(6,4) + C(6,5) + C(6,6).
313        assert_eq!(subsets, 20 + 15 + 6 + 1);
314    }
315
316    #[test]
317    fn bad_parameters_and_shares_are_refused() {
318        let mut out = [0u8; 10];
319        let mut rng = Stream(1);
320        assert!(split(b"", 2, 2, &mut rng, &mut []).is_err(), "empty secret");
321        assert!(
322            split(b"ab", 1, 5, &mut rng, &mut out).is_err(),
323            "threshold 1"
324        );
325        assert!(
326            split(b"ab", 3, 2, &mut rng, &mut out[..4]).is_err(),
327            "count below threshold"
328        );
329        assert!(
330            split(b"ab", 2, 5, &mut rng, &mut out[..9]).is_err(),
331            "wrong output length"
332        );
333
334        let a = [1u8, 2];
335        let b = [3u8, 4];
336        let mut two = [0u8; 2];
337        assert!(combine(&[(1, &a)], &mut two).is_err(), "one share");
338        assert!(combine(&[(0, &a), (1, &b)], &mut two).is_err(), "index 0");
339        assert!(
340            combine(&[(1, &a), (1, &b)], &mut two).is_err(),
341            "repeated index"
342        );
343        assert!(
344            combine(&[(1, &a), (2, &b[..1])], &mut two).is_err(),
345            "length mismatch"
346        );
347        combine(&[(1, &a), (2, &b)], &mut two).unwrap();
348    }
349
350    /// A failing random source leaves no partial shares behind.
351    #[test]
352    fn a_failing_rng_wipes_the_output() {
353        struct Fails(u8);
354        impl RandomSource for Fails {
355            fn fill(&mut self, out: &mut [u8]) -> Result<()> {
356                if self.0 == 0 {
357                    return Err(ic_core::err!(EntropyFailure, "test"));
358                }
359                self.0 -= 1;
360                out.fill(0x77);
361                Ok(())
362            }
363        }
364        let mut out = [0u8; 3 * 4];
365        assert!(split(b"keys", 2, 3, &mut Fails(2), &mut out).is_err());
366        assert_eq!(out, [0u8; 12]);
367    }
368
369    #[cfg(feature = "std")]
370    #[test]
371    fn the_vec_forms_round_trip() {
372        let shares = split_vec(b"vec secret", 2, 4, &mut Stream(9)).unwrap();
373        assert_eq!(shares.len(), 4);
374        assert_eq!(shares[3].index, 4);
375        let recovered = combine_vec(&shares[2..]).unwrap();
376        assert_eq!(&recovered.get()[..], b"vec secret");
377    }
378}