Skip to main content

multi_party_schnorr/keygen/
session.rs

1// Copyright (c) Silence Laboratories Pte. Ltd. All Rights Reserved.
2// This software is licensed under the Silence Laboratories License Agreement.
3
4use 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/// Initialize a server-side keygen session and return the first-round
43/// message to the client.
44///
45/// Builds a `KeygenParty<R1>` and `server_msg1` from the provided
46/// keys/seed. The session state is serialized and passed to
47/// `encrypt_state` as `(ad, payload)` where:
48/// - `ad = hdr || client_msg1 || server_msg1`
49/// - `payload = P1` (the round-1 party state)
50/// - `hdr = [ROUND1, TWO_MSG, offset_le]` with `offset` pointing to `payload`
51///
52/// `encrypt_state` is responsible for encrypting and persisting the
53/// state. It is assumed that caller will use AEAD and store the state
54/// blob in following format: ad | encrypted payload | tag.
55///
56/// Returns `SessionError` on invalid party ID, encoding failure, or
57/// underlying keygen errors.
58#[allow(clippy::too_many_arguments)]
59pub fn server_init<G>(
60    client_msg1: &KeygenMsg1, // c_1 from C
61    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
107/// Given server session state created by `server_init()`, extract
108/// server message to be handled by `server_round1_finish()`.
109///
110/// `encrypted_state` is a encrypted state blob in the following
111/// format: ad | encrypted-payload | tag. This function extects that
112/// prefix of `encrypted_state` is exactly the same as `ad` passed to
113/// `encrypt_state` parameter of `server_init()`.
114///
115/// Returns first server message or ServerError::Decode.
116pub 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
139/// Process third message of first round and call passed callback
140/// with encoded state: ([msg1], P2).
141///
142/// ([msg1, msg2], P1) + msg3 => ([msg1], P2)
143///
144/// `encrypt_state` is passed with 3 parameters:
145/// - `final_session_id`
146/// - `ad` = encoded hdr | msg1
147/// - `payload` = encoded P2
148///
149/// `final_session_id` is additional parameter extracted from second
150/// round message. It could be used as additional input to derive
151/// state id.
152///
153pub fn server_round1_finish<G>(
154    msg3: KeygenMsg1, // a_1 to B or b_1 to A
155    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
200/// Given encypted state created by `server_round1_finish()` or
201/// `server_round2_message()`.
202///
203/// The function is similar to `server_round1_decode_server_message()`
204/// and follows the same convension for `encrypted_state`.
205///
206/// Returns first server message or ServerError::Decode.
207pub 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
233/// Process next message of round2. There are two calls of this
234/// function.  First time when processing KeygenMsg2 from client, and
235/// second time when processing message form other server.
236///
237/// ([msg1], P2) + msg2       => ([msg1, msg2], P2)
238/// or
239/// ([msg1, msg2], P2) + msg3 => Keyshare
240///
241/// `encrypte_state` accepts 3 parameters:
242/// - `final_session_id`
243/// - `ad` = encoded hdr | msg1
244/// - `payload` = encoded P2 or encoded keyshare.
245///
246/// If `ad` is empty than `payload` is encoded keyshare.
247///
248pub 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        // Initializing the keygen for each party.
332        let (party_key_list, party_pubkey_list) = generate_pki(N as usize, &mut rng);
333
334        // Initiate client first round state and create first client's message.
335        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        // Request 1: Client -> Server 1, client receives s1_1
351        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, [c1, s1_msg1])
363                s1_state_1.extend_from_slice(ad);
364                s1_state_1.extend_from_slice(payload);
365
366                // DB[_id] = s1_state_1
367            },
368        )
369        .unwrap();
370
371        // Request 2: Client -> Server 2, , client receives s2_1
372        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                // DB[_id] = s2_state_1
386            },
387        )
388        .unwrap();
389
390        // After communicating with two servers, client could finish
391        // its first round and produce its second message.
392        let (c2, client_msg2) = c1.process(vec![client_msg1, s1_1, s2_1]).unwrap();
393
394        // Servers exchange its first round messages. It could be done
395        // using client or by independent channel.
396
397        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                // DB[_id] = s1_state_2
406            },
407        )
408        .unwrap();
409
410        // We can use `server_round1_decode_server_message()` to
411        // extract `s1_1` from encrypted s1_state_1.
412        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                // DB[_id] = s2_state_2
426            },
427        )
428        .unwrap();
429
430        // Request 3: Client -> Serser 1, client receives message s1_2
431        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                // DB[_id] = s1_state_3
439            },
440        )
441        .unwrap();
442
443        // Request 4: Client -> Serser 2, client receives messages s2_2
444        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                // DB[_id] = s2_state_3
452            },
453        )
454        .unwrap();
455
456        // After execution of requests 3 & 4, client could finish its
457        // second found and calculate its keyshare.
458        let _client_keyshare = c2
459            .process(vec![client_msg2, s1_2.clone(), s2_2.clone()])
460            .unwrap();
461
462        // Now, two servers must exchange its second round messages
463        // and calculate its keyshares.
464        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        // Here we could use `s1_s` or extract the same message from
471        // encrypte state.
472        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}