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 ensure!(
75 out.len() == len * share_count as usize,
76 InvalidLength,
77 "shamir output must be share_count * secret.len() bytes"
78 );
79
80 let mut coefficients = [0u8; MAX_SHARES as usize - 1];
81 let coefficients = &mut coefficients[..threshold as usize - 1];
82 for (b, &s) in secret.iter().enumerate() {
83 if let Err(e) = rng.fill(coefficients) {
84 coefficients.zeroize();
85 out.zeroize();
86 return Err(e);
87 }
88 for i in 0..share_count as usize {
89 let x = (i + 1) as u8;
92 let mut y = 0u8;
93 for &c in coefficients.iter().rev() {
94 y = gf::mul(y, x) ^ c;
95 }
96 out[i * len + b] = gf::mul(y, x) ^ s;
97 }
98 }
99 coefficients.zeroize();
100 Ok(())
101}
102
103pub fn combine(shares: &[(u8, &[u8])], out: &mut [u8]) -> Result<()> {
110 ensure!(
111 shares.len() >= 2,
112 InvalidParameter,
113 "shamir needs at least two shares"
114 );
115 for (n, (index, bytes)) in shares.iter().enumerate() {
116 ensure!(
117 *index != 0,
118 InvalidParameter,
119 "shamir share index 0 is the secret itself"
120 );
121 ensure!(
122 bytes.len() == out.len(),
123 InvalidLength,
124 "shamir shares differ in length"
125 );
126 ensure!(
127 shares[..n].iter().all(|(other, _)| other != index),
128 InvalidParameter,
129 "shamir share indices repeat"
130 );
131 }
132
133 out.zeroize();
134 for (i, (xi, yi)) in shares.iter().enumerate() {
135 let mut numerator = 1u8;
139 let mut denominator = 1u8;
140 for (j, (xj, _)) in shares.iter().enumerate() {
141 if i != j {
142 numerator = gf::mul(numerator, *xj);
143 denominator = gf::mul(denominator, xj ^ xi);
144 }
145 }
146 let basis = gf::mul(numerator, gf::inv(denominator));
147 for (o, y) in out.iter_mut().zip(yi.iter()) {
148 *o ^= gf::mul(*y, basis);
149 }
150 }
151 Ok(())
152}
153
154#[cfg(feature = "std")]
156pub struct Share {
157 pub index: u8,
159 pub value: std::vec::Vec<u8>,
161}
162
163#[cfg(feature = "std")]
164impl Drop for Share {
165 fn drop(&mut self) {
166 self.value.zeroize();
167 }
168}
169
170#[cfg(feature = "std")]
171impl core::fmt::Debug for Share {
172 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
173 f.debug_struct("Share")
174 .field("index", &self.index)
175 .finish_non_exhaustive()
176 }
177}
178
179#[cfg(feature = "std")]
181pub fn split_vec<R: RandomSource + ?Sized>(
182 secret: &[u8],
183 threshold: u8,
184 share_count: u8,
185 rng: &mut R,
186) -> Result<std::vec::Vec<Share>> {
187 let len = secret.len();
188 let mut all = ic_core::Zeroizing::new(std::vec![0u8; len * share_count as usize]);
189 split(secret, threshold, share_count, rng, all.get_mut())?;
190 Ok(all
191 .get()
192 .chunks(len)
193 .enumerate()
194 .map(|(i, bytes)| Share {
195 index: (i + 1) as u8,
196 value: bytes.to_vec(),
197 })
198 .collect())
199}
200
201#[cfg(feature = "std")]
203pub fn combine_vec(shares: &[Share]) -> Result<ic_core::Zeroizing<std::vec::Vec<u8>>> {
204 let len = shares.first().map_or(0, |s| s.value.len());
205 let parts: std::vec::Vec<(u8, &[u8])> =
206 shares.iter().map(|s| (s.index, &s.value[..])).collect();
207 let mut out = ic_core::Zeroizing::new(std::vec![0u8; len]);
208 combine(&parts, out.get_mut())?;
209 Ok(out)
210}
211
212pub(crate) struct Stream(pub(crate) u8);
215
216impl RandomSource for Stream {
217 fn fill(&mut self, out: &mut [u8]) -> Result<()> {
218 for b in out.iter_mut() {
219 self.0 = self.0.wrapping_mul(29).wrapping_add(7);
220 *b = self.0;
221 }
222 Ok(())
223 }
224}
225
226impl SelfTest for Shamir {
227 fn self_test() -> Result<()> {
231 let mut secret = [0u8; 32];
232 for (i, b) in secret.iter_mut().enumerate() {
233 *b = i as u8;
234 }
235 let mut shares = [0u8; 5 * 32];
236 split(&secret, 3, 5, &mut Stream(0x11), &mut shares)?;
237 let mut want = [0u8; 64];
238 ic_core::codec::hex_decode(KAT_SHARES_1_2, &mut want)?;
239 ensure!(
240 ic_core::ct::verify(&shares[..64], &want),
241 SelfTestFailed,
242 "shamir: shares differ from the reference"
243 );
244 let mut recovered = [0u8; 32];
245 combine(
246 &[
247 (2, &shares[32..64]),
248 (4, &shares[96..128]),
249 (5, &shares[128..]),
250 ],
251 &mut recovered,
252 )?;
253 ensure!(
254 ic_core::ct::verify(&recovered, &secret),
255 SelfTestFailed,
256 "shamir: recovery differs"
257 );
258 Ok(())
259 }
260}
261
262const KAT_SHARES_1_2: &[u8] = b"5ff2a5a02b96c144f71a6d28237e898c4fa2f5f0fb065154278abd78732e191c69afee93f0ed7a6a5af11d0dc3b352b422ffbec32026b1218aa156eb93552fe4";
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 #[test]
270 fn the_self_test_passes() {
271 Shamir::self_test().unwrap();
272 }
273
274 #[test]
277 fn every_threshold_subset_recovers_the_secret() {
278 let secret = *b"a 24-byte secret for k=3";
279 let (k, n) = (3u8, 6u8);
280 let mut shares = [0u8; 6 * 24];
281 split(&secret, k, n, &mut Stream(0x5d), &mut shares).unwrap();
282 let share = |i: usize| ((i + 1) as u8, &shares[i * 24..(i + 1) * 24]);
283
284 let mut subsets = 0;
285 for mask in 1u32..(1 << n) {
286 let picked: std::vec::Vec<_> = (0..n as usize)
287 .filter(|i| mask & (1 << i) != 0)
288 .map(share)
289 .collect();
290 if picked.len() < 2 {
291 continue;
292 }
293 let mut out = [0u8; 24];
294 combine(&picked, &mut out).unwrap();
295 if picked.len() >= k as usize {
296 assert_eq!(out, secret, "subset {mask:06b}");
297 subsets += 1;
298 } else {
299 assert_ne!(out, secret, "subset {mask:06b}, below the threshold");
300 }
301 }
302 assert_eq!(subsets, 20 + 15 + 6 + 1);
304 }
305
306 #[test]
307 fn bad_parameters_and_shares_are_refused() {
308 let mut out = [0u8; 10];
309 let mut rng = Stream(1);
310 assert!(split(b"", 2, 2, &mut rng, &mut []).is_err(), "empty secret");
311 assert!(
312 split(b"ab", 1, 5, &mut rng, &mut out).is_err(),
313 "threshold 1"
314 );
315 assert!(
316 split(b"ab", 3, 2, &mut rng, &mut out[..4]).is_err(),
317 "count below threshold"
318 );
319 assert!(
320 split(b"ab", 2, 5, &mut rng, &mut out[..9]).is_err(),
321 "wrong output length"
322 );
323
324 let a = [1u8, 2];
325 let b = [3u8, 4];
326 let mut two = [0u8; 2];
327 assert!(combine(&[(1, &a)], &mut two).is_err(), "one share");
328 assert!(combine(&[(0, &a), (1, &b)], &mut two).is_err(), "index 0");
329 assert!(
330 combine(&[(1, &a), (1, &b)], &mut two).is_err(),
331 "repeated index"
332 );
333 assert!(
334 combine(&[(1, &a), (2, &b[..1])], &mut two).is_err(),
335 "length mismatch"
336 );
337 combine(&[(1, &a), (2, &b)], &mut two).unwrap();
338 }
339
340 #[test]
342 fn a_failing_rng_wipes_the_output() {
343 struct Fails(u8);
344 impl RandomSource for Fails {
345 fn fill(&mut self, out: &mut [u8]) -> Result<()> {
346 if self.0 == 0 {
347 return Err(ic_core::err!(EntropyFailure, "test"));
348 }
349 self.0 -= 1;
350 out.fill(0x77);
351 Ok(())
352 }
353 }
354 let mut out = [0u8; 3 * 4];
355 assert!(split(b"keys", 2, 3, &mut Fails(2), &mut out).is_err());
356 assert_eq!(out, [0u8; 12]);
357 }
358
359 #[cfg(feature = "std")]
360 #[test]
361 fn the_vec_forms_round_trip() {
362 let shares = split_vec(b"vec secret", 2, 4, &mut Stream(9)).unwrap();
363 assert_eq!(shares.len(), 4);
364 assert_eq!(shares[3].index, 4);
365 let recovered = combine_vec(&shares[2..]).unwrap();
366 assert_eq!(&recovered.get()[..], b"vec secret");
367 }
368}