Skip to main content

wrkflw_secrets/
storage.rs

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/// Encrypted secret storage for sensitive data at rest
11#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct EncryptedSecretStore {
13    /// Encrypted secrets map (base64 encoded, nonce prepended to ciphertext)
14    secrets: HashMap<String, String>,
15    /// Salt for key derivation (base64 encoded)
16    salt: String,
17}
18
19impl EncryptedSecretStore {
20    /// Create a new encrypted secret store with a random key
21    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    /// Create an encrypted secret store from existing data
34    pub fn from_data(secrets: HashMap<String, String>, salt: String) -> Self {
35        Self { secrets, salt }
36    }
37
38    /// Add an encrypted secret
39    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    /// Get and decrypt a secret
46    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    /// Remove a secret
56    pub fn remove_secret(&mut self, name: &str) -> bool {
57        self.secrets.remove(name).is_some()
58    }
59
60    /// List all secret names
61    pub fn list_secrets(&self) -> Vec<String> {
62        self.secrets.keys().cloned().collect()
63    }
64
65    /// Check if a secret exists
66    pub fn has_secret(&self, name: &str) -> bool {
67        self.secrets.contains_key(name)
68    }
69
70    /// Get the number of stored secrets
71    pub fn secret_count(&self) -> usize {
72        self.secrets.len()
73    }
74
75    /// Clear all secrets
76    pub fn clear(&mut self) {
77        self.secrets.clear();
78    }
79
80    /// Encrypt a value with a fresh nonce prepended to the ciphertext
81    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        // Prepend nonce to ciphertext so each secret carries its own nonce
91        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    /// Decrypt a value (nonce is extracted from the ciphertext prefix)
99    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    /// Generate a random salt
123    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    /// Generate a random nonce
130    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    /// Serialize to JSON
137    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    /// Deserialize from JSON.
143    /// Detects the old serialization format (which had a top-level `nonce` field)
144    /// and returns a clear error directing users to re-create their secret store.
145    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        // Detect old format: if the JSON contains a top-level "nonce" field, it's
150        // the pre-per-secret-nonce format that is no longer compatible.
151        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    /// Save to file
165    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    /// Load from file
173    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
190/// Key derivation utilities
191pub struct KeyDerivation;
192
193impl KeyDerivation {
194    /// Derive a key from a password using PBKDF2
195    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    /// Generate a secure random key
207    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        // Add a secret
221        store
222            .add_secret(&key, "test_secret", "secret_value")
223            .unwrap();
224
225        // Retrieve the secret
226        let value = store.get_secret(&key, "test_secret").unwrap();
227        assert_eq!(value, "secret_value");
228
229        // Check metadata
230        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        // Add multiple secrets
243        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        // Retrieve all secrets
248        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        // Add secret with key1
261        store
262            .add_secret(&key1, "test_secret", "secret_value")
263            .unwrap();
264
265        // Try to retrieve with wrong key
266        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        // Serialize to JSON
310        let json = store.to_json().unwrap();
311
312        // Deserialize from JSON
313        let restored_store = EncryptedSecretStore::from_json(&json).unwrap();
314
315        // Verify secrets are still accessible
316        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        // Encrypt the same value twice - should produce different ciphertexts
331        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        // Different nonces means different ciphertexts even for the same plaintext
338        assert_ne!(encrypted_a, encrypted_b);
339
340        // Both should decrypt to the same value
341        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        // Simulate the old serialization format that had a top-level "nonce" field
348        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        // Same password and salt should produce same key
374        assert_eq!(key1, key2);
375
376        // Different salt should produce different key
377        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        // Random keys should be different
388        assert_ne!(key1, key2);
389
390        // Keys should be 32 bytes
391        assert_eq!(key1.len(), 32);
392        assert_eq!(key2.len(), 32);
393    }
394}