1mod der;
18#[cfg(unix)]
19mod key_file;
20#[cfg(not(unix))]
21mod no_key_file;
22mod seed_form;
23#[cfg(unix)]
24mod stored_identity;
25
26pub use seed_form::KeyFileError;
27#[cfg(unix)]
28pub use stored_identity::{default_identity_dir, identity_path, DEFAULT_IDENTITY_NAME};
29
30use std::fmt;
31
32use aws_lc_rs::rand::SystemRandom;
33use aws_lc_rs::rsa::{KeyPair as RsaKeyPair, KeySize};
34use aws_lc_rs::signature::{
35 KeyPair as _, UnparsedPublicKey, RSA_PSS_2048_8192_SHA384, RSA_PSS_SHA384,
36};
37use macula_mldsa::{PrivateKey, Zeroizing, ML_DSA_87};
38use sha2::{Digest, Sha256, Sha512};
39
40use crate::profile::Profile;
41
42pub const PUZZLE_DIFFICULTY: u32 = 8;
46
47const MLDSA_PUBLIC_KEY_SIZE: usize = 2592;
48const MLDSA_SIGNATURE_SIZE: usize = 4627;
49const RSA_MODULUS_BYTES: usize = 512;
50const COMPOSITE_PREFIX: &[u8] = b"CompositeAlgorithmSignatures2025";
51const COMPOSITE_LABEL: &[u8] = b"COMPSIG-MLDSA87-RSA4096-PSS-SHA512";
52const NODE_ID_LABEL: &[u8] = b"MACULA-NODE-ID-V1";
53const KEY_ID_LABEL: &[u8] = b"MACULA-KEY-ID-V1";
54
55#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
57pub enum Purpose {
58 Identity,
61 Connect,
64}
65
66impl Purpose {
67 pub fn name(self) -> &'static str {
69 match self {
70 Purpose::Identity => "identity",
71 Purpose::Connect => "connect",
72 }
73 }
74}
75
76impl fmt::Display for Purpose {
77 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
78 f.write_str(self.name())
79 }
80}
81
82#[derive(Debug, Clone, PartialEq, Eq)]
84pub enum KeyError {
85 NotAnIdentityKey,
87 DifficultyOutOfRange(u32),
89 RandomnessUnavailable,
91 Generate(&'static str),
93 Sign(&'static str),
95}
96
97impl fmt::Display for KeyError {
98 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
99 match self {
100 KeyError::NotAnIdentityKey => f.write_str("not an identity key"),
101 KeyError::DifficultyOutOfRange(d) => {
102 write!(f, "puzzle difficulty {d} is outside 0 to 256")
103 }
104 KeyError::RandomnessUnavailable => {
105 f.write_str("the operating system gave no randomness")
106 }
107 KeyError::Generate(half) => write!(f, "could not generate the {half} half"),
108 KeyError::Sign(half) => write!(f, "could not sign with the {half} half"),
109 }
110 }
111}
112
113impl std::error::Error for KeyError {}
114
115struct RsaHalf {
117 pair: RsaKeyPair,
118 public_der: Vec<u8>,
120}
121
122pub struct NodeKey {
126 purpose: Purpose,
127 profile: Profile,
128 mldsa_seed: Zeroizing<[u8; 32]>,
129 mldsa_public: Vec<u8>,
130 rsa: Option<RsaHalf>,
131}
132
133impl NodeKey {
134 pub fn generate(purpose: Purpose, profile: Profile) -> Result<NodeKey, KeyError> {
136 let (mldsa_public, mldsa_seed) =
137 macula_mldsa::key_gen_seed(ML_DSA_87).map_err(|_| KeyError::RandomnessUnavailable)?;
138 let rsa = if profile.hybrid() {
139 let pair = RsaKeyPair::generate(KeySize::Rsa4096)
140 .map_err(|_| KeyError::Generate("RSA-4096"))?;
141 let public_der = pair.public_key().as_ref().to_vec();
142 Some(RsaHalf { pair, public_der })
143 } else {
144 None
145 };
146 Ok(NodeKey {
147 purpose,
148 profile,
149 mldsa_seed,
150 mldsa_public,
151 rsa,
152 })
153 }
154
155 pub fn generate_identity(profile: Profile, difficulty: u32) -> Result<NodeKey, KeyError> {
160 if difficulty > 256 {
161 return Err(KeyError::DifficultyOutOfRange(difficulty));
162 }
163 let mut key = NodeKey::generate(Purpose::Identity, profile)?;
164 while !puzzle_solved(&node_id_of(&key.public_key(), profile), difficulty) {
165 let (public, seed) = macula_mldsa::key_gen_seed(ML_DSA_87)
166 .map_err(|_| KeyError::RandomnessUnavailable)?;
167 key.mldsa_public = public;
168 key.mldsa_seed = seed;
169 }
170 Ok(key)
171 }
172
173 pub fn purpose(&self) -> Purpose {
175 self.purpose
176 }
177
178 pub fn profile(&self) -> Profile {
180 self.profile
181 }
182
183 pub fn public_key(&self) -> Vec<u8> {
186 let mut carried = self.mldsa_public.clone();
187 if let Some(rsa) = &self.rsa {
188 carried.extend_from_slice(&rsa.public_der);
189 }
190 carried
191 }
192
193 pub fn node_id(&self) -> Result<[u8; 32], KeyError> {
195 match self.purpose {
196 Purpose::Identity => Ok(node_id_of(&self.public_key(), self.profile)),
197 Purpose::Connect => Err(KeyError::NotAnIdentityKey),
198 }
199 }
200
201 pub fn key_id(&self) -> [u8; 32] {
204 match self.purpose {
205 Purpose::Identity => node_id_of(&self.public_key(), self.profile),
206 Purpose::Connect => key_id_of(&self.public_key(), self.profile),
207 }
208 }
209
210 pub fn sign(&self, message: &[u8]) -> Result<Vec<u8>, KeyError> {
216 let seed = PrivateKey::Seed(&self.mldsa_seed);
217 let Some(rsa) = &self.rsa else {
218 return macula_mldsa::sign(ML_DSA_87, seed, message, &[])
219 .map_err(|_| KeyError::Sign("ML-DSA-87"));
220 };
221 let representative = composite_representative(message);
222 let mut signature = macula_mldsa::sign(ML_DSA_87, seed, &representative, COMPOSITE_LABEL)
223 .map_err(|_| KeyError::Sign("ML-DSA-87"))?;
224 let mut rsa_signature = vec![0u8; rsa.pair.public_modulus_len()];
225 rsa.pair
226 .sign(
227 &RSA_PSS_SHA384,
228 &SystemRandom::new(),
229 &representative,
230 &mut rsa_signature,
231 )
232 .map_err(|_| KeyError::Sign("RSA-PSS"))?;
233 signature.extend_from_slice(&rsa_signature);
234 Ok(signature)
235 }
236}
237
238impl fmt::Display for NodeKey {
239 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
240 write!(
241 f,
242 "{} {} key {}",
243 self.purpose,
244 self.profile,
245 hex_of(&self.key_id())
246 )
247 }
248}
249
250impl fmt::Debug for NodeKey {
251 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
252 fmt::Display::fmt(self, f)
253 }
254}
255
256pub fn verify(message: &[u8], signature: &[u8], carried_key: &[u8], profile: Profile) -> bool {
261 if !profile.hybrid() {
262 return signature.len() == MLDSA_SIGNATURE_SIZE
263 && carried_key.len() == MLDSA_PUBLIC_KEY_SIZE
264 && macula_mldsa::verify(ML_DSA_87, carried_key, message, signature, &[]) == Ok(true);
265 }
266 if signature.len() != signature_size(profile) || !carried_key_well_formed(carried_key, profile)
267 {
268 return false;
269 }
270 let representative = composite_representative(message);
271 let (mldsa_public, rsa_public) = carried_key.split_at(MLDSA_PUBLIC_KEY_SIZE);
272 let (mldsa_signature, rsa_signature) = signature.split_at(MLDSA_SIGNATURE_SIZE);
273 let mldsa_valid = macula_mldsa::verify(
274 ML_DSA_87,
275 mldsa_public,
276 &representative,
277 mldsa_signature,
278 COMPOSITE_LABEL,
279 ) == Ok(true);
280 let rsa_valid = UnparsedPublicKey::new(&RSA_PSS_2048_8192_SHA384, rsa_public)
281 .verify(&representative, rsa_signature)
282 .is_ok();
283 mldsa_valid && rsa_valid
284}
285
286pub fn carried_key_well_formed(key: &[u8], profile: Profile) -> bool {
291 if !profile.hybrid() {
292 return key.len() == MLDSA_PUBLIC_KEY_SIZE;
293 }
294 if key.len() <= MLDSA_PUBLIC_KEY_SIZE {
295 return false;
296 }
297 let der = &key[MLDSA_PUBLIC_KEY_SIZE..];
298 der::rsa_public_key_is_4096_f4(der)
299 && aws_lc_rs::rsa::PublicKey::from_der(der).is_ok_and(|parsed| parsed.as_ref() == der)
300}
301
302pub fn signature_size(profile: Profile) -> usize {
306 if profile.hybrid() {
307 MLDSA_SIGNATURE_SIZE + RSA_MODULUS_BYTES
308 } else {
309 MLDSA_SIGNATURE_SIZE
310 }
311}
312
313pub fn node_id_of(carried_key: &[u8], profile: Profile) -> [u8; 32] {
317 labelled_id(NODE_ID_LABEL, carried_key, profile)
318}
319
320pub fn key_id_of(carried_key: &[u8], profile: Profile) -> [u8; 32] {
323 labelled_id(KEY_ID_LABEL, carried_key, profile)
324}
325
326fn labelled_id(label: &[u8], carried_key: &[u8], profile: Profile) -> [u8; 32] {
329 let name = profile.name();
330 let mut h = Sha256::new();
331 h.update(label);
332 h.update([0, name.len() as u8]);
333 h.update(name.as_bytes());
334 h.update(carried_key);
335 h.finalize().into()
336}
337
338pub fn puzzle_solved(node_id: &[u8; 32], difficulty: u32) -> bool {
341 if difficulty > 256 {
342 return false;
343 }
344 let (whole, rest) = ((difficulty / 8) as usize, difficulty % 8);
345 node_id[..whole].iter().all(|&b| b == 0) && (rest == 0 || node_id[whole] >> (8 - rest) == 0)
346}
347
348fn composite_representative(message: &[u8]) -> Vec<u8> {
351 let mut out = Vec::with_capacity(COMPOSITE_PREFIX.len() + COMPOSITE_LABEL.len() + 1 + 64);
352 out.extend_from_slice(COMPOSITE_PREFIX);
353 out.extend_from_slice(COMPOSITE_LABEL);
354 out.push(0);
355 out.extend_from_slice(&Sha512::digest(message));
356 out
357}
358
359fn hex_of(bytes: &[u8]) -> String {
360 bytes.iter().map(|b| format!("{b:02x}")).collect()
361}