use crate::{
crypto::{CipherSuite, DebugSignature, KeyPair},
MlsError, Result,
};
use saorsa_pqc::api::{MlDsaSecretKey, MlKemSecretKey, SlhDsaSecretKey};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use uuid::Uuid;
#[derive(Clone)]
enum SecretSignatureKey {
MlDsa(MlDsaSecretKey),
SlhDsa(SlhDsaSecretKey),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct MemberId(pub Uuid);
impl MemberId {
pub fn generate() -> Self {
Self(Uuid::new_v4())
}
pub fn from_bytes(bytes: [u8; 16]) -> Self {
Self(Uuid::from_bytes(bytes))
}
pub fn as_bytes(&self) -> &[u8; 16] {
self.0.as_bytes()
}
}
impl fmt::Display for MemberId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
#[derive(Serialize, Deserialize)]
pub struct MemberIdentity {
pub id: MemberId,
pub name: Option<String>,
pub credential: Credential,
pub key_package: KeyPackage,
#[serde(skip)]
signing_key: Option<Arc<SecretSignatureKey>>,
#[serde(skip)]
kem_secret: Option<Arc<MlKemSecretKey>>,
}
impl Clone for MemberIdentity {
fn clone(&self) -> Self {
Self {
id: self.id,
name: self.name.clone(),
credential: self.credential.clone(),
key_package: self.key_package.clone(),
signing_key: self.signing_key.clone(),
kem_secret: self.kem_secret.clone(),
}
}
}
impl PartialEq for MemberIdentity {
fn eq(&self, other: &Self) -> bool {
self.id == other.id
&& self.name == other.name
&& self.credential == other.credential
&& self.key_package == other.key_package
}
}
impl Eq for MemberIdentity {}
impl fmt::Debug for MemberIdentity {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MemberIdentity")
.field("id", &self.id)
.field("name", &self.name)
.field("credential", &self.credential)
.field("key_package", &self.key_package)
.finish_non_exhaustive()
}
}
impl MemberIdentity {
pub fn generate(id: MemberId) -> Result<Self> {
Self::generate_with_suite(id, CipherSuite::default())
}
pub fn generate_with_suite(id: MemberId, suite: CipherSuite) -> Result<Self> {
let keypair = KeyPair::generate(suite);
let signing_key = Arc::new(match &keypair.signature_key {
crate::crypto::SignatureKey::MlDsa { secret, .. } => {
SecretSignatureKey::MlDsa(secret.clone())
}
crate::crypto::SignatureKey::SlhDsa { secret, .. } => {
SecretSignatureKey::SlhDsa(secret.clone())
}
});
let kem_secret = Arc::new(keypair.kem_secret.clone());
let credential = Credential::new_basic(id, None, &keypair, keypair.suite)?;
let key_package = KeyPackage::new(keypair, credential.clone())?;
Ok(Self {
id,
name: None,
credential,
key_package,
signing_key: Some(signing_key),
kem_secret: Some(kem_secret),
})
}
pub fn with_name(name: String) -> Result<Self> {
let id = MemberId::generate();
let mut identity = Self::generate(id)?;
let suite = identity.key_package.cipher_suite;
let keypair = KeyPair::generate(suite);
let signing_key = Arc::new(match &keypair.signature_key {
crate::crypto::SignatureKey::MlDsa { secret, .. } => {
SecretSignatureKey::MlDsa(secret.clone())
}
crate::crypto::SignatureKey::SlhDsa { secret, .. } => {
SecretSignatureKey::SlhDsa(secret.clone())
}
});
let kem_secret = Arc::new(keypair.kem_secret.clone());
identity.name = Some(name.clone());
identity.credential = Credential::new_basic(id, Some(name), &keypair, suite)?;
identity.key_package = KeyPackage::new(keypair, identity.credential.clone())?;
identity.signing_key = Some(signing_key);
identity.kem_secret = Some(kem_secret);
Ok(identity)
}
pub fn cipher_suite(&self) -> CipherSuite {
self.key_package.cipher_suite
}
pub fn signing_key(&self) -> Option<&MlDsaSecretKey> {
self.signing_key.as_deref().map(|key| match key {
SecretSignatureKey::MlDsa(k) => k,
SecretSignatureKey::SlhDsa(_) => panic!("Called signing_key() on SLH-DSA identity"),
})
}
pub fn kem_secret(&self) -> Option<&MlKemSecretKey> {
self.kem_secret.as_deref()
}
pub fn verify_signature(&self, data: &[u8], signature: &crate::crypto::Signature) -> bool {
self.key_package
.verify_signature(data, signature)
.unwrap_or(false)
}
pub fn sign(&self, data: &[u8]) -> Result<crate::crypto::Signature> {
use saorsa_pqc::api::{MlDsa, SlhDsa};
let signing_key = self
.signing_key
.as_ref()
.ok_or_else(|| MlsError::InvalidGroupState("No signing key available".to_string()))?;
match signing_key.as_ref() {
SecretSignatureKey::MlDsa(secret) => {
let ml_dsa = MlDsa::new(self.key_package.cipher_suite.ml_dsa_variant());
let signature = ml_dsa
.sign(secret, data)
.map_err(|e| MlsError::CryptoError(format!("ML-DSA signing failed: {e:?}")))?;
Ok(crate::crypto::Signature::MlDsa(signature))
}
SecretSignatureKey::SlhDsa(secret) => {
let slh_dsa = SlhDsa::new(self.key_package.cipher_suite.slh_dsa_variant());
let signature = slh_dsa
.sign(secret, data)
.map_err(|e| MlsError::CryptoError(format!("SLH-DSA signing failed: {e:?}")))?;
Ok(crate::crypto::Signature::SlhDsa(signature))
}
}
}
pub fn verifying_key_bytes(&self) -> &[u8] {
&self.key_package.verifying_key
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum CredentialType {
Basic = 1,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum Credential {
Basic {
credential_type: CredentialType,
identity: Vec<u8>,
signature: DebugSignature,
},
}
impl Credential {
pub fn new_basic(
member_id: MemberId,
name: Option<String>,
keypair: &KeyPair,
suite: CipherSuite,
) -> Result<Self> {
let mut identity = Vec::new();
identity.extend_from_slice(b"MLS 1.0 Credential");
identity.extend_from_slice(member_id.as_bytes());
if let Some(ref name) = name {
identity.extend_from_slice(name.as_bytes());
}
let suite_bytes =
postcard::to_stdvec(&suite).map_err(|e| MlsError::SerializationError(e.to_string()))?;
identity.extend_from_slice(&suite_bytes);
identity.extend_from_slice(&keypair.verifying_key_bytes());
let signature = keypair.sign(&identity)?;
Ok(Self::Basic {
credential_type: CredentialType::Basic,
identity,
signature: DebugSignature(signature),
})
}
pub fn credential_type(&self) -> CredentialType {
match self {
Self::Basic {
credential_type, ..
} => *credential_type,
}
}
pub fn verify(&self, keypair: &KeyPair) -> bool {
match self {
Self::Basic {
identity,
signature,
..
} => {
let key_bytes = keypair.verifying_key_bytes();
let key_len = key_bytes.len();
if identity.len() < key_len {
return false;
}
let identity_key = &identity[identity.len() - key_len..];
if identity_key != key_bytes.as_slice() {
return false;
}
keypair.verify(identity, &signature.0)
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KeyPackage {
pub version: u16,
pub cipher_suite: CipherSuite,
pub verifying_key: Vec<u8>,
pub agreement_key: Vec<u8>,
pub credential: Credential,
pub extensions: Vec<Extension>,
pub signature: DebugSignature,
}
impl KeyPackage {
pub fn new(keypair: KeyPair, credential: Credential) -> Result<Self> {
if !credential.verify(&keypair) {
return Err(MlsError::InvalidGroupState(
"invalid credential signature".to_string(),
));
}
let mut package = Self {
version: 1,
cipher_suite: keypair.suite,
verifying_key: keypair.verifying_key_bytes(),
agreement_key: keypair.public_key().to_bytes().to_vec(),
credential,
extensions: Vec::new(),
signature: DebugSignature(keypair.sign(&[])?), };
let tbs = package.to_be_signed()?;
package.signature = DebugSignature(keypair.sign(&tbs)?);
Ok(package)
}
pub fn verify_signature(
&self,
data: &[u8],
signature: &crate::crypto::Signature,
) -> Result<bool> {
match signature {
crate::crypto::Signature::MlDsa(sig) => {
use saorsa_pqc::api::{MlDsa, MlDsaPublicKey};
let ml_dsa = MlDsa::new(self.cipher_suite.ml_dsa_variant());
let public_key = MlDsaPublicKey::from_bytes(
self.cipher_suite.ml_dsa_variant(),
&self.verifying_key,
)
.map_err(|e| MlsError::CryptoError(format!("Invalid ML-DSA public key: {e:?}")))?;
ml_dsa.verify(&public_key, data, sig).map_err(|e| {
MlsError::CryptoError(format!("ML-DSA verification failed: {e:?}"))
})
}
crate::crypto::Signature::SlhDsa(sig) => {
use saorsa_pqc::api::{SlhDsa, SlhDsaPublicKey};
let slh_dsa = SlhDsa::new(self.cipher_suite.slh_dsa_variant());
let public_key = SlhDsaPublicKey::from_bytes(
self.cipher_suite.slh_dsa_variant(),
&self.verifying_key,
)
.map_err(|e| MlsError::CryptoError(format!("Invalid SLH-DSA public key: {e:?}")))?;
slh_dsa.verify(&public_key, data, sig).map_err(|e| {
MlsError::CryptoError(format!("SLH-DSA verification failed: {e:?}"))
})
}
}
}
pub fn verify(&self) -> Result<bool> {
let tbs = self.to_be_signed()?;
self.verify_signature(&tbs, &self.signature.0)
}
fn to_be_signed(&self) -> Result<Vec<u8>> {
let mut data = Vec::new();
data.extend_from_slice(&self.version.to_be_bytes());
let suite_bytes = postcard::to_stdvec(&self.cipher_suite)
.map_err(|e| MlsError::SerializationError(e.to_string()))?;
data.extend_from_slice(&suite_bytes);
data.extend_from_slice(&self.verifying_key);
data.extend_from_slice(&self.agreement_key);
let cred_bytes = postcard::to_stdvec(&self.credential)
.map_err(|e| MlsError::SerializationError(e.to_string()))?;
data.extend_from_slice(&cred_bytes);
Ok(data)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum Extension {
ApplicationId(Vec<u8>),
RatchetTree(Vec<u8>),
ExternalPub(Vec<u8>),
ExternalSenders(Vec<u8>),
}
#[derive(Debug, Clone)]
pub struct MemberList {
members: HashMap<MemberId, MemberIdentity>,
}
impl MemberList {
pub fn new() -> Self {
Self {
members: HashMap::new(),
}
}
pub fn add(&mut self, member: MemberIdentity) {
self.members.insert(member.id, member);
}
pub fn remove(&mut self, id: &MemberId) -> Option<MemberIdentity> {
self.members.remove(id)
}
pub fn get(&self, id: &MemberId) -> Option<&MemberIdentity> {
self.members.get(id)
}
pub fn get_mut(&mut self, id: &MemberId) -> Option<&mut MemberIdentity> {
self.members.get_mut(id)
}
pub fn contains(&self, id: &MemberId) -> bool {
self.members.contains_key(id)
}
pub fn len(&self) -> usize {
self.members.len()
}
pub fn is_empty(&self) -> bool {
self.members.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = (&MemberId, &MemberIdentity)> {
self.members.iter()
}
pub fn member_ids(&self) -> Vec<MemberId> {
self.members.keys().copied().collect()
}
}
impl Default for MemberList {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemberState {
pub identity: MemberIdentity,
pub leaf_index: usize,
pub generation: u32,
pub last_update: u64,
}
impl MemberState {
pub fn new(identity: MemberIdentity, leaf_index: usize) -> Self {
Self {
identity,
leaf_index,
generation: 0,
last_update: 0,
}
}
pub fn increment_generation(&mut self) {
self.generation = self.generation.wrapping_add(1);
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LifetimeExtension {
pub not_before: u64,
pub not_after: u64,
}
impl LifetimeExtension {
pub fn new(duration: Duration) -> Self {
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
Self {
not_before: now,
not_after: now + duration.as_secs(),
}
}
pub fn is_valid(&self) -> bool {
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
now >= self.not_before && now <= self.not_after
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct GroupMember {
pub identity: MemberIdentity,
pub index: u32,
pub active: bool,
pub schema_version: u8,
}
impl GroupMember {
pub fn new(identity: MemberIdentity, index: u32) -> Self {
Self {
identity,
index,
active: true,
schema_version: 1,
}
}
pub fn deactivate(&mut self) {
self.active = false;
}
pub fn is_active(&self) -> bool {
self.active
}
pub fn index(&self) -> u32 {
self.index
}
pub fn identity(&self) -> &MemberIdentity {
&self.identity
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemberRegistry {
members: HashMap<u32, GroupMember>,
next_index: u32,
pub schema_version: u8,
}
impl MemberRegistry {
pub fn new() -> Self {
Self {
members: HashMap::new(),
next_index: 0,
schema_version: 1,
}
}
pub fn add_member(&mut self, identity: MemberIdentity) -> Result<u32> {
let index = self.next_index;
let member = GroupMember::new(identity, index);
if self.members.insert(index, member).is_some() {
return Err(MlsError::InvalidGroupState(format!(
"Member index {} already exists",
index
)));
}
self.next_index += 1;
Ok(index)
}
pub fn remove_member(&mut self, index: u32) -> Result<GroupMember> {
self.members.remove(&index).ok_or_else(|| {
let mut uuid_bytes = [0u8; 16];
uuid_bytes[0..4].copy_from_slice(&index.to_be_bytes());
MlsError::MemberNotFound(MemberId::from_bytes(uuid_bytes))
})
}
pub fn get_member(&self, index: u32) -> Option<&GroupMember> {
self.members.get(&index)
}
pub fn get_member_mut(&mut self, index: u32) -> Option<&mut GroupMember> {
self.members.get_mut(&index)
}
pub fn active_members(&self) -> impl Iterator<Item = &GroupMember> {
self.members.values().filter(|m| m.is_active())
}
pub fn total_members(&self) -> usize {
self.members.len()
}
pub fn active_member_count(&self) -> usize {
self.members.values().filter(|m| m.is_active()).count()
}
pub fn is_empty(&self) -> bool {
self.members.is_empty()
}
pub fn member_indices(&self) -> impl Iterator<Item = u32> + '_ {
self.members.keys().copied()
}
pub fn find_member_index(&self, member_id: &MemberId) -> Option<u32> {
self.members
.iter()
.find(|(_, member)| member.identity.id == *member_id)
.map(|(index, _)| *index)
}
}
impl Default for MemberRegistry {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Default)]
pub struct TrustStore {
pub trusted_keys: Vec<Vec<u8>>,
}
impl TrustStore {
pub fn new() -> Self {
Self::default()
}
pub fn add_trusted_key(&mut self, public_key: Vec<u8>) {
self.trusted_keys.push(public_key);
}
pub fn remove_trusted_key(&mut self, public_key: &[u8]) {
self.trusted_keys.retain(|key| key.as_slice() != public_key);
}
pub fn trusted_key_count(&self) -> usize {
self.trusted_keys.len()
}
pub fn is_trusted(&self, public_key: &[u8]) -> bool {
self.trusted_keys
.iter()
.any(|key| key.as_slice() == public_key)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_member_id_generation() {
let id1 = MemberId::generate();
let id2 = MemberId::generate();
assert_ne!(id1, id2);
}
#[test]
fn test_member_identity_creation() {
let identity = MemberIdentity::generate(MemberId::generate()).unwrap();
assert!(identity.name.is_none());
assert_eq!(identity.cipher_suite(), CipherSuite::default());
}
#[test]
fn test_member_identity_with_name() {
let name = "Alice".to_string();
let identity = MemberIdentity::with_name(name.clone()).unwrap();
assert_eq!(identity.name, Some(name));
}
#[test]
fn test_member_list_operations() {
let mut list = MemberList::new();
assert!(list.is_empty());
let id = MemberId::generate();
let member = MemberIdentity::generate(id).unwrap();
list.add(member.clone());
assert_eq!(list.len(), 1);
assert!(list.contains(&id));
assert!(list.get(&id).is_some());
list.remove(&id);
assert!(list.is_empty());
}
#[test]
fn test_credential_verification() {
let keypair = KeyPair::generate(CipherSuite::default());
let credential = Credential::new_basic(
MemberId::generate(),
Some("Test".to_string()),
&keypair,
keypair.suite,
)
.unwrap();
assert!(credential.verify(&keypair));
}
#[test]
fn test_key_package_creation_and_verification() {
let keypair = KeyPair::generate(CipherSuite::default());
let credential =
Credential::new_basic(MemberId::generate(), None, &keypair, keypair.suite).unwrap();
let key_package = KeyPackage::new(keypair, credential).unwrap();
assert!(key_package.verify().unwrap());
}
#[test]
fn test_member_state() {
let identity = MemberIdentity::generate(MemberId::generate()).unwrap();
let mut state = MemberState::new(identity, 0);
assert_eq!(state.generation, 0);
state.increment_generation();
assert_eq!(state.generation, 1);
}
#[test]
fn test_extension_serialization() {
let ext = Extension::ApplicationId(vec![1, 2, 3]);
let serialized = postcard::to_stdvec(&ext).unwrap();
let deserialized: Extension = postcard::from_bytes(&serialized).unwrap();
match deserialized {
Extension::ApplicationId(data) => assert_eq!(data, vec![1, 2, 3]),
_ => panic!("Wrong extension type"),
}
}
#[test]
fn test_lifetime_extension() {
let lifetime = LifetimeExtension::new(Duration::from_secs(3600));
assert!(lifetime.is_valid());
let expired = LifetimeExtension {
not_before: 0,
not_after: 1,
};
assert!(!expired.is_valid());
}
#[test]
fn test_member_identity_update_name() {
let identity1 = MemberIdentity::generate(MemberId::generate()).unwrap();
let identity2 = MemberIdentity::with_name("Bob".to_string()).unwrap();
assert!(identity1.name.is_none());
assert_eq!(identity2.name, Some("Bob".to_string()));
assert_ne!(
identity1.key_package.verifying_key,
identity2.key_package.verifying_key
);
}
#[test]
fn test_member_list_iteration() {
let mut list = MemberList::new();
let id1 = MemberId::generate();
let id2 = MemberId::generate();
list.add(MemberIdentity::generate(id1).unwrap());
list.add(MemberIdentity::generate(id2).unwrap());
let member_ids_list: Vec<MemberId> = list.member_ids();
assert_eq!(member_ids_list.len(), 2);
assert!(member_ids_list.contains(&id1));
assert!(member_ids_list.contains(&id2));
}
}