use std::fmt;
use aws_lc_rs::encoding::AsDer;
use aws_lc_rs::rsa::KeyPair as RsaKeyPair;
use aws_lc_rs::signature::KeyPair as _;
use macula_mldsa::{PrivateKey, Zeroizing, ML_DSA_87};
use super::{der, verify, KeyError, NodeKey, Purpose, RsaHalf};
use crate::keystore::{KeyStore, KeyStoreError};
use crate::profile::Profile;
const MAGIC: &[u8] = b"macula-node-key-seed-v1\0";
const STORE_MAGIC: &[u8] = b"macula-node-key-private-v1\0";
const LAYOUT_CAPACITY: usize = 8 * 1024;
#[derive(Clone, Copy, PartialEq, Eq)]
enum Form {
File,
Store,
}
const TAG_MLDSA_SEED: u8 = 1;
const TAG_RSA_PSS: u8 = 2;
#[derive(Debug)]
pub enum KeyFileError {
Io(std::io::Error),
NotRegular,
Owner,
Permissions,
TooLarge,
BadKeyFile,
WrongPurpose(Purpose),
WrongProfile(Profile),
WrongAlgorithms,
WrongKeySize,
PrivateKeyInvalid,
PublicKeyMismatch,
RoundTripFailed,
KeyStore(KeyStoreError),
Generate(KeyError),
NoKeyFile,
KeptInTheKeyFileForm,
}
impl fmt::Display for KeyFileError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
KeyFileError::Io(e) => write!(f, "key file: {e}"),
KeyFileError::NotRegular => f.write_str("the key file is not a regular file"),
KeyFileError::Owner => f.write_str("the key file is owned by another user"),
KeyFileError::Permissions => {
f.write_str("the key file can be read by its group or others")
}
KeyFileError::TooLarge => f.write_str("the key file is longer than 64 KiB"),
KeyFileError::BadKeyFile => f.write_str("not a key file in the seed form"),
KeyFileError::WrongPurpose(p) => write!(f, "the key file holds a key for {p}"),
KeyFileError::WrongProfile(p) => write!(f, "the key file holds a key for {p}"),
KeyFileError::WrongAlgorithms => f.write_str("the key's halves do not fit its profile"),
KeyFileError::WrongKeySize => {
f.write_str("the RSA-PSS half is not a 4096-bit key with exponent 65537")
}
KeyFileError::PrivateKeyInvalid => {
f.write_str("the key file's private key is not valid")
}
KeyFileError::PublicKeyMismatch => {
f.write_str("the stored public key is not the one its private key derives")
}
KeyFileError::RoundTripFailed => f.write_str("the key does not sign and verify"),
KeyFileError::KeyStore(e) => write!(f, "key store: {e}"),
KeyFileError::Generate(e) => write!(f, "a new key: {e}"),
KeyFileError::KeptInTheKeyFileForm => f.write_str(
"the key store holds a key in the key-file form macula-rust 0.7.0 kept; \
it no longer loads: create the identity again \
(NodeKey::generate_identity, then save_to_keystore)",
),
KeyFileError::NoKeyFile => f.write_str(
"no key file on this platform: keep the key in Credential Manager \
through keystore::KeyringStore (NodeKey::save_to_keystore)",
),
}
}
}
impl std::error::Error for KeyFileError {}
impl From<std::io::Error> for KeyFileError {
fn from(e: std::io::Error) -> Self {
KeyFileError::Io(e)
}
}
impl NodeKey {
pub fn save_to_keystore(&self, store: &dyn KeyStore) -> Result<(), KeyFileError> {
store
.save_key(&self.laid_out(Form::Store)?)
.map_err(KeyFileError::KeyStore)
}
pub fn load_from_keystore(
store: &dyn KeyStore,
purpose: Purpose,
profile: Profile,
) -> Result<NodeKey, KeyFileError> {
let contents = store.load_key().map_err(KeyFileError::KeyStore)?;
if contents.starts_with(MAGIC) {
return Err(KeyFileError::KeptInTheKeyFileForm);
}
let key = parse_form(&contents, Form::Store, purpose, profile)?;
round_trip(&key)?;
Ok(key)
}
pub(super) fn file_bytes(&self) -> Result<Zeroizing<Vec<u8>>, KeyFileError> {
self.laid_out(Form::File)
}
fn laid_out(&self, form: Form) -> Result<Zeroizing<Vec<u8>>, KeyFileError> {
let (magic, with_public) = match form {
Form::File => (MAGIC, true),
Form::Store => (STORE_MAGIC, false),
};
let public = |key: &[u8]| -> Vec<u8> {
match with_public {
true => key.to_vec(),
false => Vec::new(),
}
};
let mut out = Zeroizing::new(Vec::with_capacity(LAYOUT_CAPACITY));
out.extend_from_slice(magic);
out.extend([
purpose_tag(self.purpose),
profile_tag(self.profile),
if self.rsa.is_some() { 2 } else { 1 },
]);
append_half(
&mut out,
TAG_MLDSA_SEED,
&public(&self.mldsa_public),
&self.mldsa_seed[..],
);
if let Some(rsa) = &self.rsa {
let private = rsa_private_pkcs1(&rsa.pair)?;
append_half(&mut out, TAG_RSA_PSS, &public(&rsa.public_der), &private);
}
Ok(out)
}
}
fn purpose_tag(purpose: Purpose) -> u8 {
match purpose {
Purpose::Identity => 1,
Purpose::Connect => 2,
}
}
fn profile_tag(profile: Profile) -> u8 {
match profile {
Profile::PqPure => 1,
Profile::PqHybrid => 2,
}
}
fn append_half(out: &mut Vec<u8>, tag: u8, public: &[u8], private: &[u8]) {
out.push(tag);
out.extend((public.len() as u32).to_be_bytes());
out.extend_from_slice(public);
out.extend((private.len() as u32).to_be_bytes());
out.extend_from_slice(private);
}
fn rsa_private_pkcs1(pair: &RsaKeyPair) -> Result<Zeroizing<Vec<u8>>, KeyFileError> {
let pkcs8 = pair.as_der().map_err(|_| KeyFileError::PrivateKeyInvalid)?;
der::pkcs1_of_pkcs8(pkcs8.as_ref())
.map(Zeroizing::new)
.ok_or(KeyFileError::PrivateKeyInvalid)
}
struct StoredHalf<'a> {
tag: u8,
public: &'a [u8],
private: &'a [u8],
}
pub(super) fn parse(
bytes: &[u8],
purpose: Purpose,
profile: Profile,
) -> Result<NodeKey, KeyFileError> {
parse_form(bytes, Form::File, purpose, profile)
}
fn parse_form(
bytes: &[u8],
form: Form,
purpose: Purpose,
profile: Profile,
) -> Result<NodeKey, KeyFileError> {
let magic = match form {
Form::File => MAGIC,
Form::Store => STORE_MAGIC,
};
let rest = bytes.strip_prefix(magic).ok_or(KeyFileError::BadKeyFile)?;
let [purpose_byte, profile_byte, count, halves_bytes @ ..] = rest else {
return Err(KeyFileError::BadKeyFile);
};
let stored_purpose = match purpose_byte {
1 => Purpose::Identity,
2 => Purpose::Connect,
_ => return Err(KeyFileError::BadKeyFile),
};
let stored_profile = match profile_byte {
1 => Profile::PqPure,
2 => Profile::PqHybrid,
_ => return Err(KeyFileError::BadKeyFile),
};
let halves = parse_halves(halves_bytes)?;
if halves.len() != *count as usize {
return Err(KeyFileError::BadKeyFile);
}
if stored_purpose != purpose {
return Err(KeyFileError::WrongPurpose(stored_purpose));
}
if stored_profile != profile {
return Err(KeyFileError::WrongProfile(stored_profile));
}
let fits = match profile {
Profile::PqPure => halves.len() == 1 && halves[0].tag == TAG_MLDSA_SEED,
Profile::PqHybrid => {
halves.len() == 2 && halves[0].tag == TAG_MLDSA_SEED && halves[1].tag == TAG_RSA_PSS
}
};
if !fits {
return Err(KeyFileError::WrongAlgorithms);
}
let (mldsa_seed, mldsa_public) = mldsa_from_half(&halves[0], form)?;
let rsa = if profile.hybrid() {
Some(rsa_from_half(&halves[1], form)?)
} else {
None
};
Ok(NodeKey {
purpose,
profile,
mldsa_seed,
mldsa_public,
rsa,
})
}
fn parse_halves(mut bytes: &[u8]) -> Result<Vec<StoredHalf<'_>>, KeyFileError> {
let mut halves = Vec::new();
while let Some((&tag, rest)) = bytes.split_first() {
if tag != TAG_MLDSA_SEED && tag != TAG_RSA_PSS {
return Err(KeyFileError::BadKeyFile);
}
let (public, rest) = length_prefixed(rest)?;
let (private, rest) = length_prefixed(rest)?;
halves.push(StoredHalf {
tag,
public,
private,
});
bytes = rest;
}
Ok(halves)
}
fn length_prefixed(bytes: &[u8]) -> Result<(&[u8], &[u8]), KeyFileError> {
let (len, rest) = bytes
.split_first_chunk::<4>()
.ok_or(KeyFileError::BadKeyFile)?;
let len = u32::from_be_bytes(*len) as usize;
if len > rest.len() {
return Err(KeyFileError::BadKeyFile);
}
Ok(rest.split_at(len))
}
fn stored_public_fits(stored: &[u8], derived: &[u8], form: Form) -> Result<(), KeyFileError> {
let fits = match form {
Form::File => stored == derived,
Form::Store => stored.is_empty(),
};
match fits {
true => Ok(()),
false => Err(KeyFileError::PublicKeyMismatch),
}
}
fn mldsa_from_half(
half: &StoredHalf<'_>,
form: Form,
) -> Result<(Zeroizing<[u8; 32]>, Vec<u8>), KeyFileError> {
let seed: [u8; 32] = half
.private
.try_into()
.map_err(|_| KeyFileError::PrivateKeyInvalid)?;
let seed = Zeroizing::new(seed);
let derived = macula_mldsa::public_key(ML_DSA_87, PrivateKey::Seed(&seed))
.map_err(|_| KeyFileError::PrivateKeyInvalid)?;
stored_public_fits(half.public, &derived, form)?;
Ok((seed, derived))
}
fn rsa_from_half(half: &StoredHalf<'_>, form: Form) -> Result<RsaHalf, KeyFileError> {
let pair = RsaKeyPair::from_der(half.private).map_err(|_| KeyFileError::PrivateKeyInvalid)?;
let derived = pair.public_key().as_ref().to_vec();
stored_public_fits(half.public, &derived, form)?;
if !der::rsa_public_key_is_4096_f4(&derived) {
return Err(KeyFileError::WrongKeySize);
}
Ok(RsaHalf {
pair,
public_der: derived,
})
}
pub(super) fn round_trip(key: &NodeKey) -> Result<(), KeyFileError> {
let mut message = [0u8; 32];
aws_lc_rs::rand::fill(&mut message).map_err(|_| KeyFileError::RoundTripFailed)?;
let signature = key
.sign(&message)
.map_err(|_| KeyFileError::RoundTripFailed)?;
if verify(&message, &signature, &key.public_key(), key.profile) {
Ok(())
} else {
Err(KeyFileError::RoundTripFailed)
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
const CREDENTIAL_MANAGER_MAX: usize = 2560;
const MAX_RSA4096_PKCS1: usize = 4 + 3 + (4 + 513) + (2 + 3) + (4 + 513) + 5 * (4 + 257);
#[derive(Default)]
struct Memory(Mutex<Option<Vec<u8>>>);
impl KeyStore for Memory {
fn save_key(&self, key: &[u8]) -> Result<(), KeyStoreError> {
*self.0.lock().unwrap() = Some(key.to_vec());
Ok(())
}
fn load_key(&self) -> Result<Zeroizing<Vec<u8>>, KeyStoreError> {
let held = self.0.lock().unwrap().clone();
held.map(Zeroizing::new).ok_or(KeyStoreError::NotFound)
}
fn delete_key(&self) -> Result<(), KeyStoreError> {
*self.0.lock().unwrap() = None;
Ok(())
}
}
fn kept(store: &Memory) -> usize {
store.0.lock().unwrap().as_ref().map_or(0, Vec::len)
}
fn kept_at_worst(profile: Profile) -> usize {
let key = NodeKey::generate_identity(profile, 0).unwrap();
let store = Memory::default();
key.save_to_keystore(&store).unwrap();
let loaded = NodeKey::load_from_keystore(&store, Purpose::Identity, profile).unwrap();
assert_eq!(loaded.public_key(), key.public_key(), "{profile:?}");
let Some(rsa) = &key.rsa else {
return kept(&store);
};
let pkcs1 = rsa_private_pkcs1(&rsa.pair).unwrap().len();
assert!(pkcs1 <= MAX_RSA4096_PKCS1, "{pkcs1}");
kept(&store) - pkcs1 + MAX_RSA4096_PKCS1
}
#[test]
fn a_kept_key_fits_credential_manager_at_the_worst_rsa_size() {
for profile in [Profile::PqPure, Profile::PqHybrid] {
let worst = kept_at_worst(profile);
assert!(
worst <= CREDENTIAL_MANAGER_MAX,
"{profile:?}: {worst} bytes"
);
}
}
#[test]
fn a_key_kept_in_the_0_7_0_form_is_refused_naming_the_fix() {
let key = NodeKey::generate_identity(Profile::PqPure, 0).unwrap();
let store = Memory::default();
store.save_key(&key.file_bytes().unwrap()).unwrap();
let Err(e) = NodeKey::load_from_keystore(&store, Purpose::Identity, Profile::PqPure) else {
panic!("the 0.7.0 form is refused");
};
assert!(matches!(e, KeyFileError::KeptInTheKeyFileForm), "{e:?}");
assert!(e.to_string().contains("create the identity again"), "{e}");
}
}