use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::{Rng, RngCore, distributions::Alphanumeric};
#[derive(Debug, thiserror::Error)]
pub enum KeyDerivationError {
#[error("Invalid input: {0}")]
InvalidInput(String),
#[error("Derivation failed: {0}")]
DerivationFailed(String),
}
pub fn generate_secure_key() -> String {
let mut key_bytes = [0u8; 32];
rand::thread_rng().fill_bytes(&mut key_bytes);
URL_SAFE_NO_PAD.encode(&key_bytes)
}
pub fn generate_secure_key_with_length(bytes: usize) -> String {
let mut key_bytes = vec![0u8; bytes];
rand::thread_rng().fill_bytes(&mut key_bytes);
URL_SAFE_NO_PAD.encode(&key_bytes)
}
pub fn generate_key_id(role: &str) -> String {
let timestamp = chrono::Utc::now().timestamp();
let random: String = rand::thread_rng()
.sample_iter(&Alphanumeric)
.take(8)
.map(char::from)
.collect();
format!("lmcp_{}_{timestamp}_{random}", role.to_lowercase())
}
pub fn derive_key(
input: &str,
salt: &[u8],
iterations: u32,
) -> Result<[u8; 32], KeyDerivationError> {
use pbkdf2::pbkdf2_hmac;
use sha2::Sha256;
if input.is_empty() {
return Err(KeyDerivationError::InvalidInput("Empty input".to_string()));
}
if salt.is_empty() {
return Err(KeyDerivationError::InvalidInput("Empty salt".to_string()));
}
if iterations == 0 {
return Err(KeyDerivationError::InvalidInput(
"Iterations must be > 0".to_string(),
));
}
let mut key = [0u8; 32];
pbkdf2_hmac::<Sha256>(input.as_bytes(), salt, iterations, &mut key);
Ok(key)
}
pub fn generate_master_key() -> Result<[u8; 32], KeyDerivationError> {
generate_master_key_for_application(None)
}
pub fn generate_master_key_for_application(
app_name: Option<&str>,
) -> Result<[u8; 32], KeyDerivationError> {
if let Some(app) = app_name {
let app_specific_var = format!(
"PULSEENGINE_MCP_MASTER_KEY_{}",
app.to_uppercase().replace('-', "_")
);
if let Ok(master_key_b64) = std::env::var(&app_specific_var) {
return decode_master_key(&master_key_b64);
}
}
if let Ok(master_key_b64) = std::env::var("PULSEENGINE_MCP_MASTER_KEY") {
return decode_master_key(&master_key_b64);
}
let mut key = [0u8; 32];
rand::thread_rng().fill_bytes(&mut key);
tracing::warn!(
"Generated new master key. Set PULSEENGINE_MCP_MASTER_KEY={} for persistence",
URL_SAFE_NO_PAD.encode(&key)
);
Ok(key)
}
fn decode_master_key(master_key_b64: &str) -> Result<[u8; 32], KeyDerivationError> {
let key_bytes = URL_SAFE_NO_PAD
.decode(master_key_b64)
.map_err(|e| KeyDerivationError::InvalidInput(format!("Invalid master key: {}", e)))?;
if key_bytes.len() != 32 {
return Err(KeyDerivationError::InvalidInput(format!(
"Master key must be 32 bytes, got {}",
key_bytes.len()
)));
}
let mut key = [0u8; 32];
key.copy_from_slice(&key_bytes);
Ok(key)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_secure_key() {
let key1 = generate_secure_key();
let key2 = generate_secure_key();
assert_ne!(key1, key2);
assert_eq!(key1.len(), 43);
assert!(!key1.contains('+'));
assert!(!key1.contains('/'));
assert!(!key1.contains('='));
}
#[test]
fn test_generate_key_id() {
let id1 = generate_key_id("admin");
let id2 = generate_key_id("admin");
assert_ne!(id1, id2);
assert!(id1.starts_with("lmcp_admin_"));
assert!(id1.matches('_').count() == 3);
}
#[test]
fn test_derive_key() {
let password = "test-password";
let salt = b"test-salt-1234567890";
let key1 = derive_key(password, salt, 1000).unwrap();
let key2 = derive_key(password, salt, 1000).unwrap();
assert_eq!(key1, key2);
let key3 = derive_key(password, b"different-salt", 1000).unwrap();
assert_ne!(key1, key3);
let key4 = derive_key(password, salt, 2000).unwrap();
assert_ne!(key1, key4);
}
#[test]
fn test_derive_key_validation() {
assert!(derive_key("", b"salt", 1000).is_err());
assert!(derive_key("password", b"", 1000).is_err());
assert!(derive_key("password", b"salt", 0).is_err());
}
}