1use crate::gf;
30use ic_core::traits::{Algorithm, RandomSource, SelfTest};
31use ic_core::{ensure, Result, Zeroize};
32
33pub const MAX_SHARES: u8 = 255;
35
36pub 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
45pub 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 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
107pub 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 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#[cfg(feature = "std")]
160pub struct Share {
161 pub index: u8,
163 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#[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#[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
220pub(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 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
270const 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 #[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 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 #[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}