1use crate::{SecretError, SecretResult};
2use aes_gcm::{
3 aead::{Aead, KeyInit, OsRng},
4 Aes256Gcm, Key, Nonce,
5};
6use base64::{engine::general_purpose, Engine as _};
7use serde::{Deserialize, Serialize};
8use std::collections::HashMap;
9
10#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct EncryptedSecretStore {
13 secrets: HashMap<String, String>,
15 salt: String,
17}
18
19impl EncryptedSecretStore {
20 pub fn new() -> SecretResult<(Self, [u8; 32])> {
22 let key = Aes256Gcm::generate_key(&mut OsRng);
23 let salt = Self::generate_salt();
24
25 let store = Self {
26 secrets: HashMap::new(),
27 salt: general_purpose::STANDARD.encode(salt),
28 };
29
30 Ok((store, key.into()))
31 }
32
33 pub fn from_data(secrets: HashMap<String, String>, salt: String) -> Self {
35 Self { secrets, salt }
36 }
37
38 pub fn add_secret(&mut self, key: &[u8; 32], name: &str, value: &str) -> SecretResult<()> {
40 let encrypted = Self::encrypt_value(key, value)?;
41 self.secrets.insert(name.to_string(), encrypted);
42 Ok(())
43 }
44
45 pub fn get_secret(&self, key: &[u8; 32], name: &str) -> SecretResult<String> {
47 let encrypted = self
48 .secrets
49 .get(name)
50 .ok_or_else(|| SecretError::not_found(name))?;
51
52 Self::decrypt_value(key, encrypted)
53 }
54
55 pub fn remove_secret(&mut self, name: &str) -> bool {
57 self.secrets.remove(name).is_some()
58 }
59
60 pub fn list_secrets(&self) -> Vec<String> {
62 self.secrets.keys().cloned().collect()
63 }
64
65 pub fn has_secret(&self, name: &str) -> bool {
67 self.secrets.contains_key(name)
68 }
69
70 pub fn secret_count(&self) -> usize {
72 self.secrets.len()
73 }
74
75 pub fn clear(&mut self) {
77 self.secrets.clear();
78 }
79
80 fn encrypt_value(key: &[u8; 32], value: &str) -> SecretResult<String> {
82 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(key));
83 let nonce_bytes = Self::generate_nonce();
84 let nonce = Nonce::from_slice(&nonce_bytes);
85
86 let ciphertext = cipher
87 .encrypt(nonce, value.as_bytes())
88 .map_err(|e| SecretError::EncryptionError(format!("Encryption failed: {}", e)))?;
89
90 let mut combined = Vec::with_capacity(nonce_bytes.len() + ciphertext.len());
92 combined.extend_from_slice(&nonce_bytes);
93 combined.extend_from_slice(&ciphertext);
94
95 Ok(general_purpose::STANDARD.encode(&combined))
96 }
97
98 fn decrypt_value(key: &[u8; 32], encrypted: &str) -> SecretResult<String> {
100 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(key));
101 let combined = general_purpose::STANDARD
102 .decode(encrypted)
103 .map_err(|e| SecretError::EncryptionError(format!("Invalid ciphertext: {}", e)))?;
104
105 if combined.len() < 12 {
106 return Err(SecretError::EncryptionError(
107 "Ciphertext too short to contain nonce".to_string(),
108 ));
109 }
110
111 let (nonce_bytes, ciphertext) = combined.split_at(12);
112 let nonce = Nonce::from_slice(nonce_bytes);
113
114 let plaintext = cipher
115 .decrypt(nonce, ciphertext)
116 .map_err(|e| SecretError::EncryptionError(format!("Decryption failed: {}", e)))?;
117
118 String::from_utf8(plaintext)
119 .map_err(|e| SecretError::EncryptionError(format!("Invalid UTF-8: {}", e)))
120 }
121
122 fn generate_salt() -> [u8; 32] {
124 let mut salt = [0u8; 32];
125 rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut salt);
126 salt
127 }
128
129 fn generate_nonce() -> [u8; 12] {
131 let mut nonce = [0u8; 12];
132 rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut nonce);
133 nonce
134 }
135
136 pub fn to_json(&self) -> SecretResult<String> {
138 serde_json::to_string_pretty(self)
139 .map_err(|e| SecretError::internal(format!("Serialization failed: {}", e)))
140 }
141
142 pub fn from_json(json: &str) -> SecretResult<Self> {
146 let raw: serde_json::Value = serde_json::from_str(json)
147 .map_err(|e| SecretError::internal(format!("Deserialization failed: {}", e)))?;
148
149 if raw.get("nonce").is_some() {
152 return Err(SecretError::InvalidFormat(
153 "This secret store uses the old serialization format (shared nonce). \
154 It is incompatible with the current version which uses per-secret nonces. \
155 Please re-create your secret store. See BREAKING_CHANGES.md for details."
156 .to_string(),
157 ));
158 }
159
160 serde_json::from_value(raw)
161 .map_err(|e| SecretError::internal(format!("Deserialization failed: {}", e)))
162 }
163
164 pub async fn save_to_file(&self, path: &str) -> SecretResult<()> {
166 let json = self.to_json()?;
167 tokio::fs::write(path, json)
168 .await
169 .map_err(SecretError::IoError)
170 }
171
172 pub async fn load_from_file(path: &str) -> SecretResult<Self> {
174 let json = tokio::fs::read_to_string(path)
175 .await
176 .map_err(SecretError::IoError)?;
177 Self::from_json(&json)
178 }
179}
180
181impl Default for EncryptedSecretStore {
182 fn default() -> Self {
183 Self {
184 secrets: HashMap::new(),
185 salt: general_purpose::STANDARD.encode(Self::generate_salt()),
186 }
187 }
188}
189
190pub struct KeyDerivation;
192
193impl KeyDerivation {
194 pub fn derive_key_from_password(password: &str, salt: &[u8], iterations: u32) -> [u8; 32] {
196 let mut key = [0u8; 32];
197 let _ = pbkdf2::pbkdf2::<hmac::Hmac<sha2::Sha256>>(
198 password.as_bytes(),
199 salt,
200 iterations,
201 &mut key,
202 );
203 key
204 }
205
206 pub fn generate_random_key() -> [u8; 32] {
208 Aes256Gcm::generate_key(&mut OsRng).into()
209 }
210}
211
212#[cfg(test)]
213mod tests {
214 use super::*;
215
216 #[tokio::test]
217 async fn test_encrypted_secret_store_basic() {
218 let (mut store, key) = EncryptedSecretStore::new().unwrap();
219
220 store
222 .add_secret(&key, "test_secret", "secret_value")
223 .unwrap();
224
225 let value = store.get_secret(&key, "test_secret").unwrap();
227 assert_eq!(value, "secret_value");
228
229 assert!(store.has_secret("test_secret"));
231 assert_eq!(store.secret_count(), 1);
232
233 let secrets = store.list_secrets();
234 assert_eq!(secrets.len(), 1);
235 assert!(secrets.contains(&"test_secret".to_string()));
236 }
237
238 #[tokio::test]
239 async fn test_encrypted_secret_store_multiple_secrets() {
240 let (mut store, key) = EncryptedSecretStore::new().unwrap();
241
242 store.add_secret(&key, "secret1", "value1").unwrap();
244 store.add_secret(&key, "secret2", "value2").unwrap();
245 store.add_secret(&key, "secret3", "value3").unwrap();
246
247 assert_eq!(store.get_secret(&key, "secret1").unwrap(), "value1");
249 assert_eq!(store.get_secret(&key, "secret2").unwrap(), "value2");
250 assert_eq!(store.get_secret(&key, "secret3").unwrap(), "value3");
251
252 assert_eq!(store.secret_count(), 3);
253 }
254
255 #[tokio::test]
256 async fn test_encrypted_secret_store_wrong_key() {
257 let (mut store, key1) = EncryptedSecretStore::new().unwrap();
258 let (_, key2) = EncryptedSecretStore::new().unwrap();
259
260 store
262 .add_secret(&key1, "test_secret", "secret_value")
263 .unwrap();
264
265 let result = store.get_secret(&key2, "test_secret");
267 assert!(result.is_err());
268 }
269
270 #[tokio::test]
271 async fn test_encrypted_secret_store_not_found() {
272 let (store, key) = EncryptedSecretStore::new().unwrap();
273
274 let result = store.get_secret(&key, "nonexistent");
275 assert!(result.is_err());
276
277 match result.unwrap_err() {
278 SecretError::NotFound { name } => {
279 assert_eq!(name, "nonexistent");
280 }
281 _ => panic!("Expected NotFound error"),
282 }
283 }
284
285 #[tokio::test]
286 async fn test_encrypted_secret_store_remove() {
287 let (mut store, key) = EncryptedSecretStore::new().unwrap();
288
289 store
290 .add_secret(&key, "test_secret", "secret_value")
291 .unwrap();
292 assert!(store.has_secret("test_secret"));
293
294 let removed = store.remove_secret("test_secret");
295 assert!(removed);
296 assert!(!store.has_secret("test_secret"));
297
298 let removed_again = store.remove_secret("test_secret");
299 assert!(!removed_again);
300 }
301
302 #[tokio::test]
303 async fn test_encrypted_secret_store_serialization() {
304 let (mut store, key) = EncryptedSecretStore::new().unwrap();
305
306 store.add_secret(&key, "secret1", "value1").unwrap();
307 store.add_secret(&key, "secret2", "value2").unwrap();
308
309 let json = store.to_json().unwrap();
311
312 let restored_store = EncryptedSecretStore::from_json(&json).unwrap();
314
315 assert_eq!(
317 restored_store.get_secret(&key, "secret1").unwrap(),
318 "value1"
319 );
320 assert_eq!(
321 restored_store.get_secret(&key, "secret2").unwrap(),
322 "value2"
323 );
324 }
325
326 #[tokio::test]
327 async fn test_each_secret_gets_unique_nonce() {
328 let (mut store, key) = EncryptedSecretStore::new().unwrap();
329
330 store.add_secret(&key, "secret_a", "same_value").unwrap();
332 store.add_secret(&key, "secret_b", "same_value").unwrap();
333
334 let encrypted_a = store.secrets.get("secret_a").unwrap();
335 let encrypted_b = store.secrets.get("secret_b").unwrap();
336
337 assert_ne!(encrypted_a, encrypted_b);
339
340 assert_eq!(store.get_secret(&key, "secret_a").unwrap(), "same_value");
342 assert_eq!(store.get_secret(&key, "secret_b").unwrap(), "same_value");
343 }
344
345 #[test]
346 fn test_old_format_with_nonce_field_rejected() {
347 let old_format_json = r#"{
349 "secrets": {},
350 "salt": "dGVzdHNhbHQ=",
351 "nonce": "dGVzdG5vbmNl"
352 }"#;
353
354 let result = EncryptedSecretStore::from_json(old_format_json);
355 assert!(result.is_err());
356 let err_msg = result.unwrap_err().to_string();
357 assert!(
358 err_msg.contains("old serialization format"),
359 "Expected old-format error, got: {}",
360 err_msg
361 );
362 }
363
364 #[test]
365 fn test_key_derivation() {
366 let password = "test_password";
367 let salt = b"test_salt_bytes_32_chars_long!!";
368 let iterations = 10000;
369
370 let key1 = KeyDerivation::derive_key_from_password(password, salt, iterations);
371 let key2 = KeyDerivation::derive_key_from_password(password, salt, iterations);
372
373 assert_eq!(key1, key2);
375
376 let different_salt = b"different_salt_bytes_32_chars!";
378 let key3 = KeyDerivation::derive_key_from_password(password, different_salt, iterations);
379 assert_ne!(key1, key3);
380 }
381
382 #[test]
383 fn test_random_key_generation() {
384 let key1 = KeyDerivation::generate_random_key();
385 let key2 = KeyDerivation::generate_random_key();
386
387 assert_ne!(key1, key2);
389
390 assert_eq!(key1.len(), 32);
392 assert_eq!(key2.len(), 32);
393 }
394}