mod der;
mod key_file;
pub use key_file::KeyFileError;
use std::fmt;
use aws_lc_rs::rand::SystemRandom;
use aws_lc_rs::rsa::{KeyPair as RsaKeyPair, KeySize};
use aws_lc_rs::signature::{
KeyPair as _, UnparsedPublicKey, RSA_PSS_2048_8192_SHA384, RSA_PSS_SHA384,
};
use macula_mldsa::{PrivateKey, Zeroizing, ML_DSA_87};
use sha2::{Digest, Sha256, Sha512};
use crate::profile::Profile;
pub const PUZZLE_DIFFICULTY: u32 = 8;
const MLDSA_PUBLIC_KEY_SIZE: usize = 2592;
const MLDSA_SIGNATURE_SIZE: usize = 4627;
const RSA_MODULUS_BYTES: usize = 512;
const COMPOSITE_PREFIX: &[u8] = b"CompositeAlgorithmSignatures2025";
const COMPOSITE_LABEL: &[u8] = b"COMPSIG-MLDSA87-RSA4096-PSS-SHA512";
const NODE_ID_LABEL: &[u8] = b"MACULA-NODE-ID-V1";
const KEY_ID_LABEL: &[u8] = b"MACULA-KEY-ID-V1";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Purpose {
Identity,
Connect,
}
impl Purpose {
pub fn name(self) -> &'static str {
match self {
Purpose::Identity => "identity",
Purpose::Connect => "connect",
}
}
}
impl fmt::Display for Purpose {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.name())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum KeyError {
NotAnIdentityKey,
DifficultyOutOfRange(u32),
RandomnessUnavailable,
Generate(&'static str),
Sign(&'static str),
}
impl fmt::Display for KeyError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
KeyError::NotAnIdentityKey => f.write_str("not an identity key"),
KeyError::DifficultyOutOfRange(d) => {
write!(f, "puzzle difficulty {d} is outside 0 to 256")
}
KeyError::RandomnessUnavailable => {
f.write_str("the operating system gave no randomness")
}
KeyError::Generate(half) => write!(f, "could not generate the {half} half"),
KeyError::Sign(half) => write!(f, "could not sign with the {half} half"),
}
}
}
impl std::error::Error for KeyError {}
struct RsaHalf {
pair: RsaKeyPair,
public_der: Vec<u8>,
}
pub struct NodeKey {
purpose: Purpose,
profile: Profile,
mldsa_seed: Zeroizing<[u8; 32]>,
mldsa_public: Vec<u8>,
rsa: Option<RsaHalf>,
}
impl NodeKey {
pub fn generate(purpose: Purpose, profile: Profile) -> Result<NodeKey, KeyError> {
let (mldsa_public, mldsa_seed) =
macula_mldsa::key_gen_seed(ML_DSA_87).map_err(|_| KeyError::RandomnessUnavailable)?;
let rsa = if profile.hybrid() {
let pair = RsaKeyPair::generate(KeySize::Rsa4096)
.map_err(|_| KeyError::Generate("RSA-4096"))?;
let public_der = pair.public_key().as_ref().to_vec();
Some(RsaHalf { pair, public_der })
} else {
None
};
Ok(NodeKey {
purpose,
profile,
mldsa_seed,
mldsa_public,
rsa,
})
}
pub fn generate_identity(profile: Profile, difficulty: u32) -> Result<NodeKey, KeyError> {
if difficulty > 256 {
return Err(KeyError::DifficultyOutOfRange(difficulty));
}
let mut key = NodeKey::generate(Purpose::Identity, profile)?;
while !puzzle_solved(&node_id_of(&key.public_key(), profile), difficulty) {
let (public, seed) = macula_mldsa::key_gen_seed(ML_DSA_87)
.map_err(|_| KeyError::RandomnessUnavailable)?;
key.mldsa_public = public;
key.mldsa_seed = seed;
}
Ok(key)
}
pub fn purpose(&self) -> Purpose {
self.purpose
}
pub fn profile(&self) -> Profile {
self.profile
}
pub fn public_key(&self) -> Vec<u8> {
let mut carried = self.mldsa_public.clone();
if let Some(rsa) = &self.rsa {
carried.extend_from_slice(&rsa.public_der);
}
carried
}
pub fn node_id(&self) -> Result<[u8; 32], KeyError> {
match self.purpose {
Purpose::Identity => Ok(node_id_of(&self.public_key(), self.profile)),
Purpose::Connect => Err(KeyError::NotAnIdentityKey),
}
}
pub fn key_id(&self) -> [u8; 32] {
match self.purpose {
Purpose::Identity => node_id_of(&self.public_key(), self.profile),
Purpose::Connect => key_id_of(&self.public_key(), self.profile),
}
}
pub fn sign(&self, message: &[u8]) -> Result<Vec<u8>, KeyError> {
let seed = PrivateKey::Seed(&self.mldsa_seed);
let Some(rsa) = &self.rsa else {
return macula_mldsa::sign(ML_DSA_87, seed, message, &[])
.map_err(|_| KeyError::Sign("ML-DSA-87"));
};
let representative = composite_representative(message);
let mut signature = macula_mldsa::sign(ML_DSA_87, seed, &representative, COMPOSITE_LABEL)
.map_err(|_| KeyError::Sign("ML-DSA-87"))?;
let mut rsa_signature = vec![0u8; rsa.pair.public_modulus_len()];
rsa.pair
.sign(
&RSA_PSS_SHA384,
&SystemRandom::new(),
&representative,
&mut rsa_signature,
)
.map_err(|_| KeyError::Sign("RSA-PSS"))?;
signature.extend_from_slice(&rsa_signature);
Ok(signature)
}
}
impl fmt::Display for NodeKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"{} {} key {}",
self.purpose,
self.profile,
hex_of(&self.key_id())
)
}
}
impl fmt::Debug for NodeKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(self, f)
}
}
pub fn verify(message: &[u8], signature: &[u8], carried_key: &[u8], profile: Profile) -> bool {
if !profile.hybrid() {
return signature.len() == MLDSA_SIGNATURE_SIZE
&& carried_key.len() == MLDSA_PUBLIC_KEY_SIZE
&& macula_mldsa::verify(ML_DSA_87, carried_key, message, signature, &[]) == Ok(true);
}
if signature.len() != signature_size(profile) || !carried_key_well_formed(carried_key, profile)
{
return false;
}
let representative = composite_representative(message);
let (mldsa_public, rsa_public) = carried_key.split_at(MLDSA_PUBLIC_KEY_SIZE);
let (mldsa_signature, rsa_signature) = signature.split_at(MLDSA_SIGNATURE_SIZE);
let mldsa_valid = macula_mldsa::verify(
ML_DSA_87,
mldsa_public,
&representative,
mldsa_signature,
COMPOSITE_LABEL,
) == Ok(true);
let rsa_valid = UnparsedPublicKey::new(&RSA_PSS_2048_8192_SHA384, rsa_public)
.verify(&representative, rsa_signature)
.is_ok();
mldsa_valid && rsa_valid
}
pub fn carried_key_well_formed(key: &[u8], profile: Profile) -> bool {
if !profile.hybrid() {
return key.len() == MLDSA_PUBLIC_KEY_SIZE;
}
if key.len() <= MLDSA_PUBLIC_KEY_SIZE {
return false;
}
let der = &key[MLDSA_PUBLIC_KEY_SIZE..];
der::rsa_public_key_is_4096_f4(der)
&& aws_lc_rs::rsa::PublicKey::from_der(der).is_ok_and(|parsed| parsed.as_ref() == der)
}
pub fn signature_size(profile: Profile) -> usize {
if profile.hybrid() {
MLDSA_SIGNATURE_SIZE + RSA_MODULUS_BYTES
} else {
MLDSA_SIGNATURE_SIZE
}
}
pub fn node_id_of(carried_key: &[u8], profile: Profile) -> [u8; 32] {
labelled_id(NODE_ID_LABEL, carried_key, profile)
}
pub fn key_id_of(carried_key: &[u8], profile: Profile) -> [u8; 32] {
labelled_id(KEY_ID_LABEL, carried_key, profile)
}
fn labelled_id(label: &[u8], carried_key: &[u8], profile: Profile) -> [u8; 32] {
let name = profile.name();
let mut h = Sha256::new();
h.update(label);
h.update([0, name.len() as u8]);
h.update(name.as_bytes());
h.update(carried_key);
h.finalize().into()
}
pub fn puzzle_solved(node_id: &[u8; 32], difficulty: u32) -> bool {
if difficulty > 256 {
return false;
}
let (whole, rest) = ((difficulty / 8) as usize, difficulty % 8);
node_id[..whole].iter().all(|&b| b == 0) && (rest == 0 || node_id[whole] >> (8 - rest) == 0)
}
fn composite_representative(message: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(COMPOSITE_PREFIX.len() + COMPOSITE_LABEL.len() + 1 + 64);
out.extend_from_slice(COMPOSITE_PREFIX);
out.extend_from_slice(COMPOSITE_LABEL);
out.push(0);
out.extend_from_slice(&Sha512::digest(message));
out
}
fn hex_of(bytes: &[u8]) -> String {
bytes.iter().map(|b| format!("{b:02x}")).collect()
}