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