Skip to main content

shadow_crypt_core/v1/
key_ops.rs

1use argon2::{Algorithm, Argon2, Params, Version};
2use zeroize::Zeroize;
3
4use crate::{
5    errors::KeyDerivationError, memory::SecureKey, report::KeyDerivationReport,
6    v1::key::KeyDerivationParams,
7};
8
9pub fn derive_key(
10    password: &[u8],
11    salt: &[u8],
12    kdf_params: &KeyDerivationParams,
13) -> Result<(SecureKey, KeyDerivationReport), KeyDerivationError> {
14    let start_time = std::time::Instant::now();
15    let algorithm = Algorithm::Argon2id;
16    let version = Version::V0x13; // Version 19
17    let params = Params::new(
18        kdf_params.memory_cost,
19        kdf_params.time_cost,
20        kdf_params.parallelism,
21        Some(kdf_params.key_size as usize),
22    )
23    .map_err(|e| KeyDerivationError::InvalidParameters(format!("Invalid KDF parameters: {}", e)))?;
24    let context = Argon2::new(algorithm, version, params);
25
26    let mut buffer = [0u8; 32];
27    context
28        .hash_password_into(password, salt, &mut buffer)
29        .map_err(|e| {
30            KeyDerivationError::DerivationFailed(format!("Key derivation failed: {}", e))
31        })?;
32
33    let key = SecureKey::new(buffer);
34    buffer.zeroize(); // Clear sensitive data from memory
35
36    let duration = start_time.elapsed();
37    let report = KeyDerivationReport::new(
38        "Argon2id".to_string(),
39        format!("{}", version as u8),
40        kdf_params.memory_cost,
41        kdf_params.time_cost,
42        kdf_params.parallelism,
43        kdf_params.key_size,
44        duration,
45    );
46
47    Ok((key, report))
48}
49
50#[cfg(test)]
51mod tests {
52    use super::*;
53    use crate::{profile::SecurityProfile, v1::key::KeyDerivationParams};
54
55    fn get_test_params() -> KeyDerivationParams {
56        let profile = SecurityProfile::Test;
57        KeyDerivationParams::from(profile)
58    }
59
60    #[test]
61    fn test_derive_key_success() {
62        let password = b"test_password";
63        let salt = b"test_salt_16_bytes";
64        let params = get_test_params();
65
66        let result = derive_key(password, salt, &params);
67        assert!(result.is_ok());
68
69        let (key, report) = result.unwrap();
70        assert_eq!(key.as_bytes().len(), 32);
71        assert_eq!(report.algorithm, "Argon2id");
72        assert_eq!(report.algorithm_version, "19");
73        assert_eq!(report.memory_cost_kib, params.memory_cost);
74        assert_eq!(report.time_cost_iterations, params.time_cost);
75        assert_eq!(report.parallelism, params.parallelism);
76        assert_eq!(report.key_size_bytes, params.key_size);
77        assert!(report.duration.as_nanos() > 0);
78    }
79
80    #[test]
81    fn test_derive_key_deterministic() {
82        let password = b"test_password";
83        let salt = b"test_salt_16_bytes";
84        let params = get_test_params();
85
86        let (key1, _) = derive_key(password, salt, &params).unwrap();
87        let (key2, _) = derive_key(password, salt, &params).unwrap();
88
89        assert_eq!(key1.as_bytes(), key2.as_bytes());
90    }
91
92    #[test]
93    fn test_derive_key_different_passwords() {
94        let salt = b"test_salt_16_bytes";
95        let params = get_test_params();
96
97        let (key1, _) = derive_key(b"password1", salt, &params).unwrap();
98        let (key2, _) = derive_key(b"password2", salt, &params).unwrap();
99
100        assert_ne!(key1.as_bytes(), key2.as_bytes());
101    }
102
103    #[test]
104    fn test_derive_key_different_salts() {
105        let password = b"test_password";
106        let params = get_test_params();
107
108        let (key1, _) = derive_key(password, b"salt1_16_bytes_!", &params).unwrap();
109        let (key2, _) = derive_key(password, b"salt2_16_bytes_!", &params).unwrap();
110
111        assert_ne!(key1.as_bytes(), key2.as_bytes());
112    }
113
114    #[test]
115    fn test_derive_key_invalid_parameters() {
116        let password = b"test_password";
117        let salt = b"test_salt_16_bytes";
118
119        // Test with memory_cost = 0 (invalid)
120        let invalid_params = KeyDerivationParams::new(0, 1, 1, 32);
121        let result = derive_key(password, salt, &invalid_params);
122        assert!(result.is_err());
123        if let Err(KeyDerivationError::InvalidParameters(_)) = result {
124            // Expected error
125        } else {
126            panic!("Expected InvalidParameters error");
127        }
128    }
129
130    #[test]
131    fn test_derive_key_invalid_key_size() {
132        let password = b"test_password";
133        let salt = b"test_salt_16_bytes";
134
135        // Test with key_size = 0 (invalid)
136        let invalid_params = KeyDerivationParams::new(1024, 1, 1, 0);
137        let result = derive_key(password, salt, &invalid_params);
138        assert!(result.is_err());
139        if let Err(KeyDerivationError::InvalidParameters(_)) = result {
140            // Expected error
141        } else {
142            panic!("Expected InvalidParameters error");
143        }
144    }
145}