1use std::fmt;
8use std::io::{Read, Write};
9use std::path::Path;
10
11use aws_lc_rs::encoding::AsDer;
12use aws_lc_rs::rsa::KeyPair as RsaKeyPair;
13use aws_lc_rs::signature::KeyPair as _;
14use macula_mldsa::{PrivateKey, Zeroizing, ML_DSA_87};
15
16use super::{der, verify, KeyError, NodeKey, Purpose, RsaHalf};
17use crate::keystore::{KeyStore, KeyStoreError};
18use crate::profile::Profile;
19
20const MAGIC: &[u8] = b"macula-node-key-seed-v1\0";
24
25const MAX_KEY_FILE_BYTES: u64 = 64 * 1024;
27
28const TAG_MLDSA_SEED: u8 = 1;
29const TAG_RSA_PSS: u8 = 2;
30
31#[derive(Debug)]
33pub enum KeyFileError {
34 Io(std::io::Error),
36 NotRegular,
39 Owner,
41 Permissions,
43 TooLarge,
45 BadKeyFile,
47 WrongPurpose(Purpose),
49 WrongProfile(Profile),
51 WrongAlgorithms,
53 WrongKeySize,
55 PrivateKeyInvalid,
57 PublicKeyMismatch,
59 RoundTripFailed,
61 KeyStore(KeyStoreError),
63 Generate(KeyError),
65}
66
67impl fmt::Display for KeyFileError {
68 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
69 match self {
70 KeyFileError::Io(e) => write!(f, "key file: {e}"),
71 KeyFileError::NotRegular => f.write_str("the key file is not a regular file"),
72 KeyFileError::Owner => f.write_str("the key file is owned by another user"),
73 KeyFileError::Permissions => {
74 f.write_str("the key file can be read by its group or others")
75 }
76 KeyFileError::TooLarge => f.write_str("the key file is longer than 64 KiB"),
77 KeyFileError::BadKeyFile => f.write_str("not a key file in the seed form"),
78 KeyFileError::WrongPurpose(p) => write!(f, "the key file holds a key for {p}"),
79 KeyFileError::WrongProfile(p) => write!(f, "the key file holds a key for {p}"),
80 KeyFileError::WrongAlgorithms => f.write_str("the key's halves do not fit its profile"),
81 KeyFileError::WrongKeySize => {
82 f.write_str("the RSA-PSS half is not a 4096-bit key with exponent 65537")
83 }
84 KeyFileError::PrivateKeyInvalid => {
85 f.write_str("the key file's private key is not valid")
86 }
87 KeyFileError::PublicKeyMismatch => {
88 f.write_str("the stored public key is not the one its private key derives")
89 }
90 KeyFileError::RoundTripFailed => f.write_str("the key does not sign and verify"),
91 KeyFileError::KeyStore(e) => write!(f, "key store: {e}"),
92 KeyFileError::Generate(e) => write!(f, "a new key: {e}"),
93 }
94 }
95}
96
97impl std::error::Error for KeyFileError {}
98
99impl From<std::io::Error> for KeyFileError {
100 fn from(e: std::io::Error) -> Self {
101 KeyFileError::Io(e)
102 }
103}
104
105impl NodeKey {
106 pub fn save(&self, path: &Path) -> Result<(), KeyFileError> {
112 let dir = match path.parent() {
113 Some(d) if !d.as_os_str().is_empty() => d,
114 _ => Path::new("."),
115 };
116 create_dir_owner_only(dir, true)?;
117 let base = path
118 .file_name()
119 .ok_or(KeyFileError::NotRegular)?
120 .to_string_lossy();
121 let staging = dir.join(format!(".{base}.saving-{}", random_suffix()?));
122 create_dir_owner_only(&staging, false)?;
123 let result = write_staged(&staging, path, &self.file_bytes()?).and_then(|()| sync_dir(dir));
124 let removed = std::fs::remove_dir_all(&staging);
125 result?;
126 removed.map_err(KeyFileError::from)
127 }
128
129 pub fn load(path: &Path, purpose: Purpose, profile: Profile) -> Result<NodeKey, KeyFileError> {
138 let contents = read_key_file(path)?;
139 let key = parse(&contents, purpose, profile)?;
140 round_trip(&key)?;
141 Ok(key)
142 }
143
144 pub fn load_or_create(path: &Path, profile: Profile) -> Result<NodeKey, KeyFileError> {
149 match std::fs::symlink_metadata(path) {
150 Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
151 let key = NodeKey::generate_identity(profile, super::PUZZLE_DIFFICULTY)
152 .map_err(KeyFileError::Generate)?;
153 key.save(path)?;
154 Ok(key)
155 }
156 _ => NodeKey::load(path, Purpose::Identity, profile),
157 }
158 }
159
160 pub fn save_to_keystore(&self, store: &dyn KeyStore) -> Result<(), KeyFileError> {
163 store
164 .save_key(&self.file_bytes()?)
165 .map_err(KeyFileError::KeyStore)
166 }
167
168 pub fn load_from_keystore(
172 store: &dyn KeyStore,
173 purpose: Purpose,
174 profile: Profile,
175 ) -> Result<NodeKey, KeyFileError> {
176 let contents = store.load_key().map_err(KeyFileError::KeyStore)?;
177 let key = parse(&contents, purpose, profile)?;
178 round_trip(&key)?;
179 Ok(key)
180 }
181
182 fn file_bytes(&self) -> Result<Vec<u8>, KeyFileError> {
184 let mut out = MAGIC.to_vec();
185 out.extend([
186 purpose_tag(self.purpose),
187 profile_tag(self.profile),
188 if self.rsa.is_some() { 2 } else { 1 },
189 ]);
190 append_half(
191 &mut out,
192 TAG_MLDSA_SEED,
193 &self.mldsa_public,
194 &self.mldsa_seed[..],
195 );
196 if let Some(rsa) = &self.rsa {
197 let private = rsa_private_pkcs1(&rsa.pair)?;
198 append_half(&mut out, TAG_RSA_PSS, &rsa.public_der, &private);
199 }
200 Ok(out)
201 }
202}
203
204fn purpose_tag(purpose: Purpose) -> u8 {
205 match purpose {
206 Purpose::Identity => 1,
207 Purpose::Connect => 2,
208 }
209}
210
211fn profile_tag(profile: Profile) -> u8 {
212 match profile {
213 Profile::PqPure => 1,
214 Profile::PqHybrid => 2,
215 }
216}
217
218fn append_half(out: &mut Vec<u8>, tag: u8, public: &[u8], private: &[u8]) {
219 out.push(tag);
220 out.extend((public.len() as u32).to_be_bytes());
221 out.extend_from_slice(public);
222 out.extend((private.len() as u32).to_be_bytes());
223 out.extend_from_slice(private);
224}
225
226fn rsa_private_pkcs1(pair: &RsaKeyPair) -> Result<Zeroizing<Vec<u8>>, KeyFileError> {
229 let pkcs8 = pair.as_der().map_err(|_| KeyFileError::PrivateKeyInvalid)?;
230 der::pkcs1_of_pkcs8(pkcs8.as_ref())
231 .map(Zeroizing::new)
232 .ok_or(KeyFileError::PrivateKeyInvalid)
233}
234
235struct StoredHalf<'a> {
237 tag: u8,
238 public: &'a [u8],
239 private: &'a [u8],
240}
241
242fn parse(bytes: &[u8], purpose: Purpose, profile: Profile) -> Result<NodeKey, KeyFileError> {
244 let rest = bytes.strip_prefix(MAGIC).ok_or(KeyFileError::BadKeyFile)?;
245 let [purpose_byte, profile_byte, count, halves_bytes @ ..] = rest else {
246 return Err(KeyFileError::BadKeyFile);
247 };
248 let stored_purpose = match purpose_byte {
249 1 => Purpose::Identity,
250 2 => Purpose::Connect,
251 _ => return Err(KeyFileError::BadKeyFile),
252 };
253 let stored_profile = match profile_byte {
254 1 => Profile::PqPure,
255 2 => Profile::PqHybrid,
256 _ => return Err(KeyFileError::BadKeyFile),
257 };
258 let halves = parse_halves(halves_bytes)?;
259 if halves.len() != *count as usize {
260 return Err(KeyFileError::BadKeyFile);
261 }
262 if stored_purpose != purpose {
263 return Err(KeyFileError::WrongPurpose(stored_purpose));
264 }
265 if stored_profile != profile {
266 return Err(KeyFileError::WrongProfile(stored_profile));
267 }
268 let fits = match profile {
269 Profile::PqPure => halves.len() == 1 && halves[0].tag == TAG_MLDSA_SEED,
270 Profile::PqHybrid => {
271 halves.len() == 2 && halves[0].tag == TAG_MLDSA_SEED && halves[1].tag == TAG_RSA_PSS
272 }
273 };
274 if !fits {
275 return Err(KeyFileError::WrongAlgorithms);
276 }
277 let (mldsa_seed, mldsa_public) = mldsa_from_half(&halves[0])?;
278 let rsa = if profile.hybrid() {
279 Some(rsa_from_half(&halves[1])?)
280 } else {
281 None
282 };
283 Ok(NodeKey {
284 purpose,
285 profile,
286 mldsa_seed,
287 mldsa_public,
288 rsa,
289 })
290}
291
292fn parse_halves(mut bytes: &[u8]) -> Result<Vec<StoredHalf<'_>>, KeyFileError> {
293 let mut halves = Vec::new();
294 while let Some((&tag, rest)) = bytes.split_first() {
295 if tag != TAG_MLDSA_SEED && tag != TAG_RSA_PSS {
296 return Err(KeyFileError::BadKeyFile);
297 }
298 let (public, rest) = length_prefixed(rest)?;
299 let (private, rest) = length_prefixed(rest)?;
300 halves.push(StoredHalf {
301 tag,
302 public,
303 private,
304 });
305 bytes = rest;
306 }
307 Ok(halves)
308}
309
310fn length_prefixed(bytes: &[u8]) -> Result<(&[u8], &[u8]), KeyFileError> {
311 let (len, rest) = bytes
312 .split_first_chunk::<4>()
313 .ok_or(KeyFileError::BadKeyFile)?;
314 let len = u32::from_be_bytes(*len) as usize;
315 if len > rest.len() {
316 return Err(KeyFileError::BadKeyFile);
317 }
318 Ok(rest.split_at(len))
319}
320
321fn mldsa_from_half(half: &StoredHalf<'_>) -> Result<(Zeroizing<[u8; 32]>, Vec<u8>), KeyFileError> {
322 let seed: [u8; 32] = half
323 .private
324 .try_into()
325 .map_err(|_| KeyFileError::PrivateKeyInvalid)?;
326 let seed = Zeroizing::new(seed);
327 let derived = macula_mldsa::public_key(ML_DSA_87, PrivateKey::Seed(&seed))
328 .map_err(|_| KeyFileError::PrivateKeyInvalid)?;
329 if derived != half.public {
330 return Err(KeyFileError::PublicKeyMismatch);
331 }
332 Ok((seed, derived))
333}
334
335fn rsa_from_half(half: &StoredHalf<'_>) -> Result<RsaHalf, KeyFileError> {
336 let pair = RsaKeyPair::from_der(half.private).map_err(|_| KeyFileError::PrivateKeyInvalid)?;
337 if pair.public_key().as_ref() != half.public {
338 return Err(KeyFileError::PublicKeyMismatch);
339 }
340 if !der::rsa_public_key_is_4096_f4(half.public) {
341 return Err(KeyFileError::WrongKeySize);
342 }
343 Ok(RsaHalf {
344 pair,
345 public_der: half.public.to_vec(),
346 })
347}
348
349fn round_trip(key: &NodeKey) -> Result<(), KeyFileError> {
352 let mut message = [0u8; 32];
353 aws_lc_rs::rand::fill(&mut message).map_err(|_| KeyFileError::RoundTripFailed)?;
354 let signature = key
355 .sign(&message)
356 .map_err(|_| KeyFileError::RoundTripFailed)?;
357 if verify(&message, &signature, &key.public_key(), key.profile) {
358 Ok(())
359 } else {
360 Err(KeyFileError::RoundTripFailed)
361 }
362}
363
364fn random_suffix() -> Result<String, KeyFileError> {
365 let mut bytes = [0u8; 8];
366 aws_lc_rs::rand::fill(&mut bytes)
367 .map_err(|_| KeyFileError::Io(std::io::Error::other("no randomness")))?;
368 Ok(bytes.iter().map(|b| format!("{b:02x}")).collect())
369}
370
371fn write_staged(staging: &Path, path: &Path, contents: &[u8]) -> Result<(), KeyFileError> {
372 let staged = staging.join("key");
373 let mut options = std::fs::OpenOptions::new();
374 options.write(true).create_new(true);
375 #[cfg(unix)]
376 std::os::unix::fs::OpenOptionsExt::mode(&mut options, 0o600);
377 let mut file = options.open(&staged)?;
378 file.write_all(contents)?;
379 file.sync_all()?;
380 drop(file);
381 std::fs::rename(&staged, path)?;
382 Ok(())
383}
384
385fn create_dir_owner_only(dir: &Path, recursive: bool) -> Result<(), KeyFileError> {
386 let mut builder = std::fs::DirBuilder::new();
387 builder.recursive(recursive);
388 #[cfg(unix)]
389 std::os::unix::fs::DirBuilderExt::mode(&mut builder, 0o700);
390 builder.create(dir)?;
391 Ok(())
392}
393
394fn sync_dir(dir: &Path) -> Result<(), KeyFileError> {
395 #[cfg(unix)]
396 std::fs::File::open(dir)?.sync_all()?;
397 #[cfg(not(unix))]
398 let _ = dir;
399 Ok(())
400}
401
402fn read_key_file(path: &Path) -> Result<Vec<u8>, KeyFileError> {
405 if !std::fs::metadata(path)?.is_file() {
406 return Err(KeyFileError::NotRegular);
407 }
408 let mut options = std::fs::OpenOptions::new();
409 options.read(true);
410 #[cfg(unix)]
413 std::os::unix::fs::OpenOptionsExt::custom_flags(
414 &mut options,
415 rustix::fs::OFlags::NONBLOCK.bits() as i32,
416 );
417 let file = options.open(path)?;
418 owner_only(&file.metadata()?)?;
419 let mut contents = Vec::new();
420 file.take(MAX_KEY_FILE_BYTES + 1)
421 .read_to_end(&mut contents)?;
422 if contents.len() as u64 > MAX_KEY_FILE_BYTES {
423 return Err(KeyFileError::TooLarge);
424 }
425 Ok(contents)
426}
427
428fn owner_only(metadata: &std::fs::Metadata) -> Result<(), KeyFileError> {
431 if !metadata.is_file() {
432 return Err(KeyFileError::NotRegular);
433 }
434 #[cfg(unix)]
435 {
436 use std::os::unix::fs::MetadataExt;
437 if metadata.uid() != rustix::process::geteuid().as_raw() {
438 return Err(KeyFileError::Owner);
439 }
440 if metadata.mode() & 0o077 != 0 {
441 return Err(KeyFileError::Permissions);
442 }
443 }
444 Ok(())
445}