use std::sync::Arc;
use crypto_bigint::subtle::ConstantTimeEq;
use crypto_box::{PublicKey, SecretKey};
use zeroize::Zeroizing;
use crate::{
common::{
ser::Serializable,
traits::{GroupElem, Round, ScalarReduce},
},
keygen::{KeyRefreshData, KeygenError, KeygenMsg1, KeygenMsg2, KeygenParty, R0, R1, R2},
};
#[derive(Debug)]
pub enum SessionError {
InvalidPartyId,
InvalidSessionId,
StateDecode,
Dkg(KeygenError),
Encode,
Decode,
}
const N: u8 = 3;
const T: u8 = 2;
const ROUND1: u8 = 1;
const ROUND2: u8 = 2;
const ONE_MSG: u8 = 1;
const TWO_MSG: u8 = 2;
impl From<KeygenError> for SessionError {
fn from(value: KeygenError) -> Self {
SessionError::Dkg(value)
}
}
#[allow(clippy::too_many_arguments)]
pub fn server_init<G>(
client_msg1: &KeygenMsg1, party_id: u8,
decryption_key: Arc<SecretKey>,
encyption_keys: Vec<(u8, PublicKey)>,
refresh_data: Option<KeyRefreshData<G>>,
key_id: Option<[u8; 32]>,
seed: [u8; 32],
extra_data: Option<Vec<u8>>,
encrypt_state: impl FnOnce(&[u8], &[u8]),
) -> Result<KeygenMsg1, SessionError>
where
G: GroupElem,
G::Scalar: ScalarReduce<[u8; 32]>,
G::Scalar: Serializable,
{
let p0 = KeygenParty::<R0, G>::new(
T,
N,
party_id,
decryption_key,
encyption_keys,
refresh_data,
key_id,
seed,
extra_data,
)?;
let (p1, server_msg1) = p0.process(())?;
let mut buffer = Zeroizing::new(vec![ROUND1, TWO_MSG, 0, 0]);
ciborium::into_writer(&client_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
ciborium::into_writer(&server_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
let offset = buffer.len();
buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
ciborium::into_writer(&p1, &mut *buffer).map_err(|_| SessionError::Encode)?;
let (ad, plantext) = buffer.split_at(offset);
encrypt_state(ad, plantext);
Ok(server_msg1)
}
pub fn server_round1_decode_server_message(
encrypted_state: &[u8],
) -> Result<KeygenMsg1, SessionError> {
let (hdr, payload) = encrypted_state
.split_first_chunk::<4>()
.ok_or(SessionError::Decode)?;
if hdr[0] != ROUND1 || hdr[1] != TWO_MSG {
return Err(SessionError::Decode);
}
let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
let mut ad = payload
.get(..offset.wrapping_sub(4))
.ok_or(SessionError::Decode)?;
let _msg1: KeygenMsg1 = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
let _msg2: KeygenMsg1 = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
Ok(_msg2)
}
pub fn server_round1_finish<G>(
msg3: KeygenMsg1, session_id: &[u8],
decrypted_state: &[u8],
encrypt_state: impl FnOnce(&[u8], &[u8], &[u8]),
) -> Result<KeygenMsg2<G>, SessionError>
where
G: GroupElem,
G::Scalar: ScalarReduce<[u8; 32]>,
G::Scalar: Serializable,
{
let (hdr, mut payload) = decrypted_state
.split_first_chunk::<4>()
.ok_or(SessionError::Decode)?;
if hdr[0] != ROUND1 || hdr[1] != TWO_MSG {
return Err(SessionError::Decode);
}
let msg1: KeygenMsg1 = ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
let msg2: KeygenMsg1 = ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
let p1: KeygenParty<R1<G>, G> =
ciborium::from_reader(&mut payload).map_err(|_| SessionError::Decode)?;
if session_id.ct_ne(msg1.session_id.as_slice()).into() {
return Err(SessionError::InvalidSessionId);
}
let (p2, server_msg1) = p1.process(vec![msg1, msg2, msg3])?;
let mut buffer = Zeroizing::new(vec![ROUND2, ONE_MSG, 0, 0]);
ciborium::into_writer(&server_msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
let offset = buffer.len();
buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
ciborium::into_writer(&p2, &mut *buffer).map_err(|_| SessionError::Encode)?;
let (ad, plantext) = buffer.split_at(offset);
encrypt_state(server_msg1.session_id.as_slice(), ad, plantext);
Ok(server_msg1)
}
pub fn server_round2_decode_server_message<G>(
encrypted_state: &[u8],
) -> Result<KeygenMsg2<G>, SessionError>
where
G: GroupElem,
G::Scalar: ScalarReduce<[u8; 32]> + Serializable,
{
let (hdr, payload) = encrypted_state
.split_first_chunk::<4>()
.ok_or(SessionError::Decode)?;
if hdr[0] != ROUND2 || !(hdr[1] == ONE_MSG || hdr[1] == TWO_MSG) {
return Err(SessionError::Decode);
}
let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
let mut ad = payload
.get(..offset.wrapping_sub(4))
.ok_or(SessionError::Decode)?;
let _msg1: KeygenMsg2<G> = ciborium::from_reader(&mut ad).map_err(|_| SessionError::Decode)?;
Ok(_msg1)
}
pub fn server_round2_message<G>(
msg3: KeygenMsg2<G>,
decrypted_state: &[u8],
encrypt_state: impl FnOnce(&[u8], &[u8], &[u8]),
) -> Result<(), SessionError>
where
G: GroupElem,
G::Scalar: ScalarReduce<[u8; 32]> + Serializable,
{
let (hdr, payload) = decrypted_state
.split_first_chunk::<4>()
.ok_or(SessionError::Decode)?;
if hdr[0] != ROUND2 {
return Err(SessionError::Decode);
}
let offset = u16::from_le_bytes([hdr[2], hdr[3]]) as usize;
let (mut msgs, payload) = payload
.split_at_checked(offset.wrapping_sub(4))
.ok_or(SessionError::Decode)?;
let p2: KeygenParty<R2, G> =
ciborium::from_reader(payload).map_err(|_| SessionError::Decode)?;
match hdr[1] {
ONE_MSG => {
let msg1: KeygenMsg2<G> =
ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
let mut buffer = Zeroizing::new(vec![ROUND2, TWO_MSG, 0, 0]);
ciborium::into_writer(&msg1, &mut *buffer).map_err(|_| SessionError::Encode)?;
ciborium::into_writer(&msg3, &mut *buffer).map_err(|_| SessionError::Encode)?;
let offset = buffer.len();
buffer[2..4].copy_from_slice(&(offset as u16).to_le_bytes());
ciborium::into_writer(&p2, &mut *buffer).map_err(|_| SessionError::Encode)?;
encrypt_state(&msg3.session_id, &buffer[..offset], &buffer[offset..]);
}
TWO_MSG => {
let msg1: KeygenMsg2<G> =
ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
let msg2: KeygenMsg2<G> =
ciborium::from_reader(&mut msgs).map_err(|_| SessionError::Decode)?;
let final_session_id = msg3.session_id;
let share = p2
.process(vec![msg1, msg2, msg3])
.map_err(|_| SessionError::Decode)?;
let mut buffer = Zeroizing::new(vec![]);
ciborium::into_writer(&share, &mut *buffer).map_err(|_| SessionError::Encode)?;
encrypt_state(&final_session_id, &[], &buffer);
}
_ => return Err(SessionError::Decode),
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::keygen::utils::generate_pki;
fn server_session<G>()
where
G: GroupElem,
G::Scalar: ScalarReduce<[u8; 32]>,
G::Scalar: Serializable,
{
let mut rng = rand::thread_rng();
let (party_key_list, party_pubkey_list) = generate_pki(N as usize, &mut rng);
let (c1, client_msg1) = KeygenParty::<R0, G>::new(
T,
N,
0,
party_key_list[0].clone(),
party_pubkey_list.clone(),
None,
None,
[1; 32],
None,
)
.unwrap()
.process(())
.unwrap();
let mut s1_state_1 = vec![];
let s1_1 = server_init::<G>(
&client_msg1,
1,
party_key_list[1].clone(),
party_pubkey_list.clone(),
None,
None,
[2; 32],
None,
|ad, payload| {
s1_state_1.extend_from_slice(ad);
s1_state_1.extend_from_slice(payload);
},
)
.unwrap();
let mut s2_state_1 = vec![];
let s2_1 = server_init::<G>(
&client_msg1,
2,
party_key_list[2].clone(),
party_pubkey_list.clone(),
None,
None,
[3; 32],
None,
|ad, payload| {
s2_state_1.extend_from_slice(ad);
s2_state_1.extend_from_slice(payload);
},
)
.unwrap();
let (c2, client_msg2) = c1.process(vec![client_msg1, s1_1, s2_1]).unwrap();
let mut s1_state_2 = vec![];
let s1_2 = server_round1_finish::<G>(
s2_1,
&client_msg1.session_id,
&s1_state_1,
|_final_session_id, ad, payload| {
s1_state_2.extend_from_slice(ad);
s1_state_2.extend_from_slice(payload);
},
)
.unwrap();
let s1_1_decoded = server_round1_decode_server_message(&s1_state_1).unwrap();
assert_eq!(s1_1.session_id, s1_1_decoded.session_id);
assert_eq!(s1_1.commitment, s1_1_decoded.commitment);
let mut s2_state_2 = vec![];
let s2_2 = server_round1_finish::<G>(
s1_1_decoded,
&client_msg1.session_id,
&s2_state_1,
|_final_session_id, ad, payload| {
s2_state_2.extend_from_slice(ad);
s2_state_2.extend_from_slice(payload);
},
)
.unwrap();
let mut s1_state_3 = vec![];
server_round2_message::<G>(
client_msg2.clone(),
&s1_state_2,
|_final_session_id, ad, payload| {
s1_state_3.extend_from_slice(ad);
s1_state_3.extend_from_slice(payload);
},
)
.unwrap();
let mut s2_state_3 = vec![];
server_round2_message::<G>(
client_msg2.clone(),
&s2_state_2,
|_final_session_id, ad, payload| {
s2_state_3.extend_from_slice(ad);
s2_state_3.extend_from_slice(payload);
},
)
.unwrap();
let _client_keyshare = c2
.process(vec![client_msg2, s1_2.clone(), s2_2.clone()])
.unwrap();
let mut s1_share = vec![];
server_round2_message::<G>(s2_2, &s1_state_3, |_final_session_id, _ad, share| {
s1_share.extend_from_slice(share);
})
.unwrap();
let s1_2_decoded = server_round2_decode_server_message::<G>(&s1_state_3).unwrap();
let mut s2_share = vec![];
server_round2_message::<G>(
s1_2_decoded,
&s2_state_3,
|_final_session_id, _ad, share| {
s2_share.extend_from_slice(share);
},
)
.unwrap();
}
#[cfg(feature = "eddsa")]
#[test]
fn session_curve25519() {
use curve25519_dalek::EdwardsPoint;
server_session::<EdwardsPoint>();
}
#[cfg(feature = "taproot")]
#[test]
fn session_taproot() {
use k256::ProjectivePoint;
server_session::<ProjectivePoint>();
}
}