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