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