use argon2::{Algorithm, Argon2, Params, Version};
use zeroize::Zeroize;
use crate::{
errors::KeyDerivationError, memory::SecureKey, profile::SecurityProfile,
report::KeyDerivationReport,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KeyDerivationParams {
pub memory_cost: u32, pub time_cost: u32, pub parallelism: u32, pub key_size: u8, }
impl KeyDerivationParams {
pub fn new(memory_cost: u32, time_cost: u32, parallelism: u32, key_size: u8) -> Self {
Self {
memory_cost,
time_cost,
parallelism,
key_size,
}
}
pub fn standard_defaults() -> Self {
Self {
memory_cost: 46 * 1024, time_cost: 1, parallelism: 1, key_size: 32, }
}
pub fn paranoid_defaults() -> Self {
Self {
memory_cost: 1024 * 1024, time_cost: 10, parallelism: 4, key_size: 32, }
}
pub fn test_defaults() -> Self {
Self {
memory_cost: 1024, time_cost: 1, parallelism: 1, key_size: 32, }
}
pub fn derive_key(
&self,
password: &[u8],
salt: &[u8],
) -> Result<(SecureKey, KeyDerivationReport), KeyDerivationError> {
let start_time = std::time::Instant::now();
let algorithm = Algorithm::Argon2id;
let version = Version::V0x13; let params = Params::new(
self.memory_cost,
self.time_cost,
self.parallelism,
Some(self.key_size as usize),
)
.map_err(|e| {
KeyDerivationError::InvalidParameters(format!("Invalid KDF parameters: {}", e))
})?;
let context = Argon2::new(algorithm, version, params);
let mut buffer = [0u8; 32];
context
.hash_password_into(password, salt, &mut buffer)
.map_err(|e| {
KeyDerivationError::DerivationFailed(format!("Key derivation failed: {}", e))
})?;
let key = SecureKey::new(buffer);
buffer.zeroize();
let duration = start_time.elapsed();
let report = KeyDerivationReport::new(
"Argon2id".to_string(),
format!("{}", version as u8),
self.memory_cost,
self.time_cost,
self.parallelism,
self.key_size,
duration,
);
Ok((key, report))
}
}
impl From<SecurityProfile> for KeyDerivationParams {
fn from(profile: SecurityProfile) -> Self {
match profile {
SecurityProfile::Standard => KeyDerivationParams::standard_defaults(),
SecurityProfile::Paranoid => KeyDerivationParams::paranoid_defaults(),
SecurityProfile::Test => KeyDerivationParams::test_defaults(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_key_derivation_params_defaults() {
let standard = KeyDerivationParams::standard_defaults();
assert_eq!(standard.memory_cost, 46 * 1024);
assert_eq!(standard.time_cost, 1);
assert_eq!(standard.parallelism, 1);
assert_eq!(standard.key_size, 32);
let paranoid = KeyDerivationParams::paranoid_defaults();
assert_eq!(paranoid.memory_cost, 1024 * 1024);
assert_eq!(paranoid.time_cost, 10);
assert_eq!(paranoid.parallelism, 4);
assert_eq!(paranoid.key_size, 32);
let test = KeyDerivationParams::test_defaults();
assert_eq!(test.memory_cost, 1024);
assert_eq!(test.time_cost, 1);
assert_eq!(test.parallelism, 1);
assert_eq!(test.key_size, 32);
}
#[test]
fn test_from_security_profile() {
let standard_params: KeyDerivationParams = SecurityProfile::Standard.into();
assert_eq!(standard_params, KeyDerivationParams::standard_defaults());
let paranoid_params: KeyDerivationParams = SecurityProfile::Paranoid.into();
assert_eq!(paranoid_params, KeyDerivationParams::paranoid_defaults());
let test_params: KeyDerivationParams = SecurityProfile::Test.into();
assert_eq!(test_params, KeyDerivationParams::test_defaults());
}
#[test]
fn test_derive_key_success() {
let params = KeyDerivationParams::test_defaults();
let (key, report) = params
.derive_key(b"test_password", b"test_salt_16_bytes")
.unwrap();
assert_eq!(key.as_bytes().len(), 32);
assert_eq!(report.algorithm, "Argon2id");
assert_eq!(report.algorithm_version, "19");
assert_eq!(report.memory_cost_kib, params.memory_cost);
}
#[test]
fn test_derive_key_deterministic() {
let params = KeyDerivationParams::test_defaults();
let (key1, _) = params
.derive_key(b"test_password", b"test_salt_16_bytes")
.unwrap();
let (key2, _) = params
.derive_key(b"test_password", b"test_salt_16_bytes")
.unwrap();
assert_eq!(key1.as_bytes(), key2.as_bytes());
}
#[test]
fn test_derive_key_different_passwords() {
let params = KeyDerivationParams::test_defaults();
let (key1, _) = params
.derive_key(b"password1", b"test_salt_16_bytes")
.unwrap();
let (key2, _) = params
.derive_key(b"password2", b"test_salt_16_bytes")
.unwrap();
assert_ne!(key1.as_bytes(), key2.as_bytes());
}
#[test]
fn test_derive_key_invalid_parameters() {
let invalid_params = KeyDerivationParams::new(0, 1, 1, 32);
let result = invalid_params.derive_key(b"pw", b"test_salt_16_bytes");
assert!(matches!(
result,
Err(KeyDerivationError::InvalidParameters(_))
));
}
}