1use std::sync::Arc;
5
6use crypto_bigint::subtle::ConstantTimeEq;
7use crypto_box::{PublicKey, SecretKey};
8use zeroize::Zeroizing;
9
10use crate::{
11 common::{
12 ser::Serializable,
13 traits::{GroupElem, Round, ScalarReduce},
14 },
15 keygen::{KeyRefreshData, KeygenError, KeygenMsg1, KeygenMsg2, KeygenParty, R0, R1, R2},
16};
17
18#[derive(Debug)]
19pub enum SessionError {
20 InvalidPartyId,
21 InvalidSessionId,
22 StateDecode,
23 Dkg(KeygenError),
24 Encode,
25 Decode,
26}
27
28const N: u8 = 3;
29const T: u8 = 2;
30
31const ROUND1: u8 = 1;
32const ROUND2: u8 = 2;
33const ONE_MSG: u8 = 1;
34const TWO_MSG: u8 = 2;
35
36impl From<KeygenError> for SessionError {
37 fn from(value: KeygenError) -> Self {
38 SessionError::Dkg(value)
39 }
40}
41
42#[allow(clippy::too_many_arguments)]
59pub fn server_init<G>(
60 client_msg1: &KeygenMsg1, party_id: u8,
62 decryption_key: Arc<SecretKey>,
63 encyption_keys: Vec<(u8, PublicKey)>,
64 refresh_data: Option<KeyRefreshData<G>>,
65 key_id: Option<[u8; 32]>,
66 seed: [u8; 32],
67 extra_data: Option<Vec<u8>>,
68 encrypt_state: impl FnOnce(&[u8], &[u8]),
69) -> Result<KeygenMsg1, SessionError>
70where
71 G: GroupElem,
72 G::Scalar: ScalarReduce<[u8; 32]>,
73 G::Scalar: Serializable,
74{
75 let p0 = KeygenParty::<R0, G>::new(
76 T,
77 N,
78 party_id,
79 decryption_key,
80 encyption_keys,
81 refresh_data,
82 key_id,
83 seed,
84 extra_data,
85 )?;
86
87 let (p1, server_msg1) = p0.process(())?;
88
89 let mut buffer = Zeroizing::new(vec![ROUND1, TWO_MSG, 0, 0]);
90
91 ciborium::into_writer(&client_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
92 ciborium::into_writer(&server_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
93
94 let offset = buffer.len();
95
96 buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
97
98 ciborium::into_writer(&p1, &mut *buffer).map_err(|_| SessionError::Encode)?;
99
100 let (ad, plantext) = buffer.split_at(offset);
101
102 encrypt_state(ad, plantext);
103
104 Ok(server_msg1)
105}
106
107pub fn server_round1_decode_server_message(
117 encrypted_state: &[u8],
118) -> Result<KeygenMsg1, SessionError> {
119 let (hdr, payload) = encrypted_state
120 .split_first_chunk::<4>()
121 .ok_or(SessionError::Decode)?;
122
123 if hdr[0] != ROUND1 || hdr[1] != TWO_MSG {
124 return Err(SessionError::Decode);
125 }
126
127 let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
128
129 let mut ad = payload
130 .get(..offset.wrapping_sub(4))
131 .ok_or(SessionError::Decode)?;
132
133 let _msg1: KeygenMsg1 = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
134 let _msg2: KeygenMsg1 = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
135
136 Ok(_msg2)
137}
138
139pub fn server_round1_finish<G>(
154 msg3: KeygenMsg1, session_id: &[u8],
156 decrypted_state: &[u8],
157 encrypt_state: impl FnOnce(&[u8], &[u8], &[u8]),
158) -> Result<KeygenMsg2<G>, SessionError>
159where
160 G: GroupElem,
161 G::Scalar: ScalarReduce<[u8; 32]>,
162 G::Scalar: Serializable,
163{
164 let (hdr, mut payload) = decrypted_state
165 .split_first_chunk::<4>()
166 .ok_or(SessionError::Decode)?;
167
168 if hdr[0] != ROUND1 || hdr[1] != TWO_MSG {
169 return Err(SessionError::Decode);
170 }
171
172 let msg1: KeygenMsg1 = ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
173 let msg2: KeygenMsg1 = ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
174 let p1: KeygenParty<R1<G>, G> =
175 ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
176
177 if session_id.ct_ne(msg1.session_id.as_slice()).into() {
178 return Err(SessionError::InvalidSessionId);
179 }
180
181 let (p2, server_msg1) = p1.process(vec![msg1, msg2, msg3])?;
182
183 let mut buffer = Zeroizing::new(vec![ROUND2, ONE_MSG, 0, 0]);
184
185 ciborium::into_writer(&server_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
186
187 let offset = buffer.len();
188
189 buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
190
191 ciborium::into_writer(&p2, &mut *buffer).map_err(|_| SessionError::Encode)?;
192
193 let (ad, plantext) = buffer.split_at(offset);
194
195 encrypt_state(server_msg1.session_id.as_slice(), ad, plantext);
196
197 Ok(server_msg1)
198}
199
200pub fn server_round2_decode_server_message<G>(
208 encrypted_state: &[u8],
209) -> Result<KeygenMsg2<G>, SessionError>
210where
211 G: GroupElem,
212 G::Scalar: ScalarReduce<[u8; 32]> + Serializable,
213{
214 let (hdr, payload) = encrypted_state
215 .split_first_chunk::<4>()
216 .ok_or(SessionError::Decode)?;
217
218 if hdr[0] != ROUND2 || !(hdr[1] == ONE_MSG || hdr[1] == TWO_MSG) {
219 return Err(SessionError::Decode);
220 }
221
222 let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
223
224 let mut ad = payload
225 .get(..offset.wrapping_sub(4))
226 .ok_or(SessionError::Decode)?;
227
228 let _msg1: KeygenMsg2<G> = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
229
230 Ok(_msg1)
231}
232
233pub fn server_round2_message<G>(
249 msg3: KeygenMsg2<G>,
250 decrypted_state: &[u8],
251 encrypt_state: impl FnOnce(&[u8], &[u8], &[u8]),
252) -> Result<(), SessionError>
253where
254 G: GroupElem,
255 G::Scalar: ScalarReduce<[u8; 32]> + Serializable,
256{
257 let (hdr, payload) = decrypted_state
258 .split_first_chunk::<4>()
259 .ok_or(SessionError::Decode)?;
260
261 if hdr[0] != ROUND2 {
262 return Err(SessionError::Decode);
263 }
264
265 let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
266
267 let (mut msgs, payload) = payload
268 .split_at_checked(offset.wrapping_sub(4))
269 .ok_or(SessionError::Decode)?;
270
271 let p2: KeygenParty<R2, G> =
272 ciborium::from_reader(payload).map_err(|_| SessionError::Decode)?;
273
274 match hdr[1] {
275 ONE_MSG => {
276 let msg1: KeygenMsg2<G> =
277 ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
278
279 let mut buffer = Zeroizing::new(vec![ROUND2, TWO_MSG, 0, 0]);
280
281 ciborium::into_writer(&msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
282 ciborium::into_writer(&msg3, &mut *buffer).map_err(|_| SessionError::Encode)?;
283
284 let offset = buffer.len();
285
286 buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
287
288 ciborium::into_writer(&p2, &mut *buffer).map_err(|_| SessionError::Encode)?;
289
290 encrypt_state(&msg3.session_id, &buffer[..offset], &buffer[offset..]);
291 }
292
293 TWO_MSG => {
294 let msg1: KeygenMsg2<G> =
295 ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
296
297 let msg2: KeygenMsg2<G> =
298 ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
299
300 let final_session_id = msg3.session_id;
301
302 let share = p2
303 .process(vec![msg1, msg2, msg3])
304 .map_err(|_| SessionError::Decode)?;
305
306 let mut buffer = Zeroizing::new(vec![]);
307
308 ciborium::into_writer(&share, &mut *buffer).map_err(|_| SessionError::Encode)?;
309
310 encrypt_state(&final_session_id, &[], &buffer);
311 }
312
313 _ => return Err(SessionError::Decode),
314 }
315
316 Ok(())
317}
318
319#[cfg(test)]
320mod tests {
321 use super::*;
322 use crate::keygen::utils::generate_pki;
323
324 fn server_session<G>()
325 where
326 G: GroupElem,
327 G::Scalar: ScalarReduce<[u8; 32]>,
328 G::Scalar: Serializable,
329 {
330 let mut rng = rand::thread_rng();
331 let (party_key_list, party_pubkey_list) = generate_pki(N as usize, &mut rng);
333
334 let (c1, client_msg1) = KeygenParty::<R0, G>::new(
336 T,
337 N,
338 0,
339 party_key_list[0].clone(),
340 party_pubkey_list.clone(),
341 None,
342 None,
343 [1; 32],
344 None,
345 )
346 .unwrap()
347 .process(())
348 .unwrap();
349
350 let mut s1_state_1 = vec![];
352 let s1_1 = server_init::<G>(
353 &client_msg1,
354 1,
355 party_key_list[1].clone(),
356 party_pubkey_list.clone(),
357 None,
358 None,
359 [2; 32],
360 None,
361 |ad, payload| {
362 s1_state_1.extend_from_slice(ad);
364 s1_state_1.extend_from_slice(payload);
365
366 },
368 )
369 .unwrap();
370
371 let mut s2_state_1 = vec![];
373 let s2_1 = server_init::<G>(
374 &client_msg1,
375 2,
376 party_key_list[2].clone(),
377 party_pubkey_list.clone(),
378 None,
379 None,
380 [3; 32],
381 None,
382 |ad, payload| {
383 s2_state_1.extend_from_slice(ad);
384 s2_state_1.extend_from_slice(payload);
385 },
387 )
388 .unwrap();
389
390 let (c2, client_msg2) = c1.process(vec![client_msg1, s1_1, s2_1]).unwrap();
393
394 let mut s1_state_2 = vec![];
398 let s1_2 = server_round1_finish::<G>(
399 s2_1,
400 &client_msg1.session_id,
401 &s1_state_1,
402 |_final_session_id, ad, payload| {
403 s1_state_2.extend_from_slice(ad);
404 s1_state_2.extend_from_slice(payload);
405 },
407 )
408 .unwrap();
409
410 let s1_1_decoded = server_round1_decode_server_message(&s1_state_1).unwrap();
413
414 assert_eq!(s1_1.session_id, s1_1_decoded.session_id);
415 assert_eq!(s1_1.commitment, s1_1_decoded.commitment);
416
417 let mut s2_state_2 = vec![];
418 let s2_2 = server_round1_finish::<G>(
419 s1_1_decoded,
420 &client_msg1.session_id,
421 &s2_state_1,
422 |_final_session_id, ad, payload| {
423 s2_state_2.extend_from_slice(ad);
424 s2_state_2.extend_from_slice(payload);
425 },
427 )
428 .unwrap();
429
430 let mut s1_state_3 = vec![];
432 server_round2_message::<G>(
433 client_msg2.clone(),
434 &s1_state_2,
435 |_final_session_id, ad, payload| {
436 s1_state_3.extend_from_slice(ad);
437 s1_state_3.extend_from_slice(payload);
438 },
440 )
441 .unwrap();
442
443 let mut s2_state_3 = vec![];
445 server_round2_message::<G>(
446 client_msg2.clone(),
447 &s2_state_2,
448 |_final_session_id, ad, payload| {
449 s2_state_3.extend_from_slice(ad);
450 s2_state_3.extend_from_slice(payload);
451 },
453 )
454 .unwrap();
455
456 let _client_keyshare = c2
459 .process(vec![client_msg2, s1_2.clone(), s2_2.clone()])
460 .unwrap();
461
462 let mut s1_share = vec![];
465 server_round2_message::<G>(s2_2, &s1_state_3, |_final_session_id, _ad, share| {
466 s1_share.extend_from_slice(share);
467 })
468 .unwrap();
469
470 let s1_2_decoded = server_round2_decode_server_message::<G>(&s1_state_3).unwrap();
473
474 let mut s2_share = vec![];
475 server_round2_message::<G>(
476 s1_2_decoded,
477 &s2_state_3,
478 |_final_session_id, _ad, share| {
479 s2_share.extend_from_slice(share);
480 },
481 )
482 .unwrap();
483 }
484
485 #[cfg(feature = "eddsa")]
486 #[test]
487 fn session_curve25519() {
488 use curve25519_dalek::EdwardsPoint;
489
490 server_session::<EdwardsPoint>();
491 }
492
493 #[cfg(feature = "taproot")]
494 #[test]
495 fn session_taproot() {
496 use k256::ProjectivePoint;
497
498 server_session::<ProjectivePoint>();
499 }
500}