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 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 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
108pub 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 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#[cfg(feature = "std")]
162pub struct Share {
163 pub index: u8,
165 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#[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#[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
222pub(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 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
272const 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 #[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 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 #[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}