use crate::http_signature::{HttpSignatureError, Result};
use base64::{engine::general_purpose::STANDARD, Engine};
use ed25519_dalek::SigningKey;
use pem::{encode, parse, Pem};
use pkcs8::{der::Decode, PrivateKeyInfo};
use rand::rngs::OsRng;
use std::fs;
use std::path::Path;
pub fn load_or_generate_key(path: &Path) -> Result<SigningKey> {
if path.exists() {
let key_str = fs::read_to_string(path)?;
let key_str = key_str.trim();
let key_str = if let Ok(decoded) = STANDARD.decode(key_str) {
String::from_utf8(decoded)?
} else {
key_str.to_string()
};
let pem = parse(&key_str).map_err(|e| HttpSignatureError::Pem(e.to_string()))?;
if pem.tag() != "PRIVATE KEY" {
return Err(HttpSignatureError::Pem("Not a PRIVATE KEY".to_string()));
}
let private_key_info = PrivateKeyInfo::from_der(pem.contents())
.map_err(|e| HttpSignatureError::Pkcs8(e.to_string()))?;
let raw = private_key_info.private_key;
let raw_private_key: [u8; 32] = raw[2..]
.try_into()
.map_err(|_| HttpSignatureError::InvalidPrivateKeyLength)?;
let signing_key = SigningKey::from_bytes(&raw_private_key);
Ok(signing_key)
} else {
let mut csprng = OsRng;
let signing_key = SigningKey::generate(&mut csprng);
let pem = Pem::new(
"PRIVATE KEY",
[
0x30, 0x2E, 0x02, 0x01, 0x00, 0x30, 0x05, 0x06, 0x03, 0x2B, 0x65, 0x70, 0x04, 0x22, 0x04, 0x20, ]
.iter()
.chain(signing_key.to_bytes().iter())
.copied()
.collect::<Vec<u8>>(),
);
let pem_string = encode(&pem);
fs::write(path, pem_string)?;
Ok(signing_key)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::tempdir;
#[test]
fn test_generate_new_key() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("test_key.pem");
assert!(!path.exists());
let key1 = load_or_generate_key(&path).unwrap();
assert!(path.exists());
let key2 = load_or_generate_key(&path).unwrap();
assert_eq!(key1.to_bytes(), key2.to_bytes());
}
#[test]
fn test_load_existing_key() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("test_key.pem");
let original_key = load_or_generate_key(&path).unwrap();
let loaded_key = load_or_generate_key(&path).unwrap();
assert_eq!(original_key.to_bytes(), loaded_key.to_bytes());
}
#[test]
fn test_invalid_pem_format() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("invalid.pem");
let mut file = fs::File::create(&path).unwrap();
file.write_all(b"-----BEGIN INVALID-----\ninvalid content\n-----END INVALID-----")
.unwrap();
let result = load_or_generate_key(&path);
assert!(matches!(result, Err(HttpSignatureError::Pem(_))));
}
#[test]
fn test_wrong_pem_tag() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("wrong_tag.pem");
let mut file = fs::File::create(&path).unwrap();
file.write_all(b"-----BEGIN PUBLIC KEY-----\ninvalid content\n-----END PUBLIC KEY-----")
.unwrap();
let result = load_or_generate_key(&path);
assert!(matches!(result, Err(HttpSignatureError::Pem(_))));
}
#[test]
fn test_base64_encoded_pem() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("test_key.pem");
let original_key = load_or_generate_key(&path).unwrap();
let pem_content = fs::read_to_string(&path).unwrap();
let encoded = STANDARD.encode(pem_content.as_bytes());
fs::write(&path, encoded).unwrap();
let loaded_key = load_or_generate_key(&path).unwrap();
assert_eq!(original_key.to_bytes(), loaded_key.to_bytes());
}
#[test]
fn test_key_persistence() {
let temp_dir = tempdir().unwrap();
let path = temp_dir.path().join("test_key.pem");
let key1 = load_or_generate_key(&path).unwrap();
let content = fs::read_to_string(&path).unwrap();
assert!(content.contains("-----BEGIN PRIVATE KEY-----"));
assert!(content.contains("-----END PRIVATE KEY-----"));
let pem = parse(&content).unwrap();
assert_eq!(pem.tag(), "PRIVATE KEY");
let key2 = load_or_generate_key(&path).unwrap();
assert_eq!(key1.to_bytes(), key2.to_bytes());
}
}