Skip to main content

aether_auth/
encrypted_file.rs

1use crate::credential::OAuthCredentialStorage;
2use crate::error::OAuthError;
3use age::scrypt::{Identity, Recipient as ScryptRecipient};
4use age::secrecy::SecretString;
5use age::{Decryptor, Encryptor};
6use async_trait::async_trait;
7use dirs::home_dir;
8use serde::{Deserialize, Serialize};
9use serde_json::Value;
10use std::collections::HashMap;
11use std::env::var;
12use std::io::{Read, Write};
13use std::iter::once;
14use std::path::{Path, PathBuf};
15use std::sync::Mutex;
16
17const DEFAULT_PASSWORD_ENV: &str = "AETHER_CREDENTIALS_PASSWORD";
18
19#[derive(Debug, Clone, Default, Serialize, Deserialize)]
20struct CredentialStore {
21    credentials: HashMap<String, Value>,
22}
23
24pub struct EncryptedFileOAuthCredentialStorage {
25    path: PathBuf,
26    passphrase: String,
27    scrypt_work_factor: Option<u8>,
28    write_guard: Mutex<()>,
29}
30
31impl EncryptedFileOAuthCredentialStorage {
32    pub fn new(path: PathBuf, passphrase: String) -> Self {
33        Self { path, passphrase, scrypt_work_factor: None, write_guard: Mutex::new(()) }
34    }
35
36    pub fn with_scrypt_work_factor(mut self, log_n: u8) -> Self {
37        self.scrypt_work_factor = Some(log_n);
38        self
39    }
40
41    pub fn from_settings(path: Option<PathBuf>, password_env: Option<&str>) -> Result<Self, OAuthError> {
42        let path = path.map_or_else(default_path, Ok)?;
43        let env_var = password_env.unwrap_or(DEFAULT_PASSWORD_ENV);
44        let passphrase = var(env_var).ok().filter(|pass| !pass.is_empty()).ok_or_else(|| {
45            OAuthError::CredentialStore(format!(
46                "Encrypted file credential store requires a passphrase. \
47                     Set the {env_var} environment variable or configure a custom `passwordEnv` in settings."
48            ))
49        })?;
50
51        Ok(Self::new(path, passphrase))
52    }
53
54    fn encrypt(plaintext: &[u8], passphrase: &str, work_factor: Option<u8>) -> Result<Vec<u8>, OAuthError> {
55        let fail = |e| OAuthError::CredentialStore(format!("Encryption failed: {e}"));
56
57        let mut recipient = ScryptRecipient::new(SecretString::from(passphrase));
58        if let Some(log_n) = work_factor {
59            recipient.set_work_factor(log_n);
60        }
61
62        let encryptor = Encryptor::with_recipients(once(&recipient as &dyn age::Recipient))
63            .map_err(|e| OAuthError::CredentialStore(format!("Encryption failed: {e}")))?;
64
65        let mut ciphertext = Vec::new();
66        let mut writer = encryptor.wrap_output(&mut ciphertext).map_err(fail)?;
67        writer.write_all(plaintext).map_err(fail)?;
68        writer.finish().map_err(fail)?;
69
70        Ok(ciphertext)
71    }
72
73    fn decrypt(ciphertext: &[u8], passphrase: &str) -> Result<Vec<u8>, OAuthError> {
74        let decryptor = Decryptor::new(ciphertext)
75            .map_err(|e| OAuthError::CredentialStore(format!("Invalid encrypted file: {e}")))?;
76
77        let mut reader =
78            decryptor.decrypt(once(&Identity::new(SecretString::from(passphrase)) as &dyn age::Identity)).map_err(
79                |e| OAuthError::CredentialStore(format!("Decryption failed — wrong passphrase or corrupted file: {e}")),
80            )?;
81
82        let mut plaintext = Vec::new();
83        reader
84            .read_to_end(&mut plaintext)
85            .map_err(|e| OAuthError::CredentialStore(format!("Decryption failed: {e}")))?;
86
87        Ok(plaintext)
88    }
89
90    fn load_store(&self) -> Result<CredentialStore, OAuthError> {
91        if !self.path.exists() {
92            return Ok(CredentialStore::default());
93        }
94
95        let bytes = std::fs::read(&self.path)?;
96        if bytes.is_empty() {
97            return Ok(CredentialStore::default());
98        }
99
100        let plaintext = Self::decrypt(&bytes, &self.passphrase)?;
101        serde_json::from_slice(&plaintext)
102            .map_err(|e| OAuthError::CredentialStore(format!("Invalid credential data: {e}")))
103    }
104
105    fn update(&self, mutate: impl FnOnce(&mut CredentialStore)) -> Result<(), OAuthError> {
106        let _guard = self
107            .write_guard
108            .lock()
109            .map_err(|_| OAuthError::CredentialStore("Failed to acquire write lock on credential store".to_string()))?;
110
111        let mut store = self.load_store()?;
112        mutate(&mut store);
113
114        let plaintext = serde_json::to_vec(&store)
115            .map_err(|e| OAuthError::CredentialStore(format!("Failed to serialize credentials: {e}")))?;
116
117        let ciphertext = Self::encrypt(&plaintext, &self.passphrase, self.scrypt_work_factor)?;
118        write_atomic(&self.path, &ciphertext)
119    }
120}
121
122fn default_path() -> Result<PathBuf, OAuthError> {
123    home_dir().map(|home| home.join(".aether").join("credentials.enc")).ok_or_else(|| {
124        OAuthError::CredentialStore(
125            "Could not determine the home directory for the encrypted credential file".to_string(),
126        )
127    })
128}
129
130fn write_atomic(path: &Path, data: &[u8]) -> Result<(), OAuthError> {
131    if let Some(parent) = path.parent() {
132        std::fs::create_dir_all(parent)?;
133    }
134
135    let temp_path = path.with_extension("tmp");
136
137    {
138        let mut file = std::fs::File::create(&temp_path)?;
139        file.write_all(data)?;
140        file.sync_all()?;
141    }
142
143    std::fs::rename(&temp_path, path)?;
144
145    Ok(())
146}
147
148#[async_trait]
149impl OAuthCredentialStorage for EncryptedFileOAuthCredentialStorage {
150    async fn load(&self, key: &str) -> Result<Option<serde_json::Value>, OAuthError> {
151        Ok(self.load_store()?.credentials.get(key).cloned())
152    }
153
154    async fn save(&self, key: &str, value: serde_json::Value) -> Result<(), OAuthError> {
155        self.update(|store| {
156            store.credentials.insert(key.to_string(), value);
157        })
158    }
159
160    async fn delete(&self, key: &str) -> Result<(), OAuthError> {
161        self.update(|store| {
162            store.credentials.remove(key);
163        })
164    }
165
166    fn contains(&self, key: &str) -> bool {
167        self.load_store().is_ok_and(|store| store.credentials.contains_key(key))
168    }
169}
170
171#[cfg(test)]
172mod tests {
173    use std::fs::write;
174
175    use super::*;
176    use crate::OAuthCredential;
177
178    /// `N = 2^10`, fast enough that the KDF stops dominating the suite.
179    const TEST_WORK_FACTOR: u8 = 10;
180
181    #[tokio::test]
182    async fn save_then_load_round_trips() {
183        let store = temp_store("correct-passphrase");
184        let cred = test_credential();
185
186        store.save_credential("server-1", cred.clone()).await.unwrap();
187        let loaded = store.load_credential("server-1").await.unwrap().unwrap();
188
189        assert_eq!(loaded.client_id, "client_1");
190        assert_eq!(loaded.access_token, "tok_abc");
191        assert_eq!(loaded.refresh_token.as_deref(), Some("ref_xyz"));
192    }
193
194    #[tokio::test]
195    async fn load_returns_none_for_missing_key() {
196        let store = temp_store("pass");
197        assert!(store.load_credential("nonexistent").await.unwrap().is_none());
198    }
199
200    #[tokio::test]
201    async fn delete_removes_credential() {
202        let store = temp_store("pass");
203        store.save_credential("server-1", test_credential()).await.unwrap();
204        assert!(store.contains("server-1"));
205
206        store.delete("server-1").await.unwrap();
207        assert!(!store.contains("server-1"));
208    }
209
210    #[tokio::test]
211    async fn wrong_passphrase_fails_to_load() {
212        let store = temp_store("correct-pass");
213        store.save_credential("server-1", test_credential()).await.unwrap();
214
215        let wrong_store = EncryptedFileOAuthCredentialStorage::new(store.path.clone(), "wrong-pass".to_string());
216
217        let err = wrong_store.load_credential("server-1").await.unwrap_err();
218        let msg = err.to_string();
219        assert!(msg.contains("Decryption failed"), "Expected decryption error, got: {msg}");
220    }
221
222    #[tokio::test]
223    async fn multiple_credentials_are_isolated() {
224        let store = temp_store("pass");
225
226        let cred_a = OAuthCredential {
227            client_id: "a".to_string(),
228            access_token: "token_a".to_string(),
229            refresh_token: None,
230            expires_at: None,
231        };
232        let cred_b = OAuthCredential {
233            client_id: "b".to_string(),
234            access_token: "token_b".to_string(),
235            refresh_token: None,
236            expires_at: None,
237        };
238
239        store.save_credential("server-a", cred_a).await.unwrap();
240        store.save_credential("server-b", cred_b).await.unwrap();
241
242        let loaded_a = store.load_credential("server-a").await.unwrap().unwrap();
243        let loaded_b = store.load_credential("server-b").await.unwrap().unwrap();
244
245        assert_eq!(loaded_a.access_token, "token_a");
246        assert_eq!(loaded_b.access_token, "token_b");
247    }
248
249    #[tokio::test]
250    async fn save_overwrites_existing_credential() {
251        let store = temp_store("pass");
252        let cred_v1 = OAuthCredential {
253            client_id: "c".to_string(),
254            access_token: "v1".to_string(),
255            refresh_token: None,
256            expires_at: None,
257        };
258        let cred_v2 = OAuthCredential {
259            client_id: "c".to_string(),
260            access_token: "v2".to_string(),
261            refresh_token: None,
262            expires_at: None,
263        };
264
265        store.save_credential("server", cred_v1).await.unwrap();
266        store.save_credential("server", cred_v2).await.unwrap();
267
268        let loaded = store.load_credential("server").await.unwrap().unwrap();
269        assert_eq!(loaded.access_token, "v2");
270    }
271
272    #[tokio::test]
273    async fn loads_existing_encrypted_file_format() {
274        let store = temp_store("pass");
275        let plaintext = serde_json::to_vec(&serde_json::json!({
276            "credentials": {
277                "codex": {
278                    "client_id": "client_1",
279                    "access_token": "tok_abc",
280                    "refresh_token": "ref_xyz",
281                    "expires_at": 9_999_999_999_999_u64
282                }
283            }
284        }))
285        .unwrap();
286        let ciphertext =
287            EncryptedFileOAuthCredentialStorage::encrypt(&plaintext, "pass", Some(TEST_WORK_FACTOR)).unwrap();
288
289        write(&store.path, ciphertext).unwrap();
290        let credential = store.load_credential("codex").await.unwrap().expect("existing credential");
291        assert_eq!(credential.access_token, "tok_abc");
292    }
293
294    #[tokio::test]
295    async fn secrets_round_trip_through_encrypted_file() {
296        let store = temp_store("pass");
297        let secret = serde_json::json!({"client_id": "c", "granted_scopes": ["read"]});
298
299        store.save("mcp:slack", secret.clone()).await.unwrap();
300        store.save_credential("slack", test_credential()).await.unwrap();
301        let reopened = EncryptedFileOAuthCredentialStorage::new(store.path.clone(), "pass".to_string())
302            .with_scrypt_work_factor(TEST_WORK_FACTOR);
303        assert_eq!(reopened.load("mcp:slack").await.unwrap(), Some(secret));
304        assert!(reopened.load_credential("slack").await.unwrap().is_some());
305
306        reopened.delete("mcp:slack").await.unwrap();
307        assert!(store.load("mcp:slack").await.unwrap().is_none());
308    }
309
310    #[test]
311    fn encrypt_decrypt_round_trips() {
312        let plaintext = b"hello, world!";
313        let ciphertext =
314            EncryptedFileOAuthCredentialStorage::encrypt(plaintext, "passphrase", Some(TEST_WORK_FACTOR)).unwrap();
315        let decrypted = EncryptedFileOAuthCredentialStorage::decrypt(&ciphertext, "passphrase").unwrap();
316        assert_eq!(decrypted, plaintext);
317    }
318
319    fn test_credential() -> OAuthCredential {
320        OAuthCredential {
321            client_id: "client_1".to_string(),
322            access_token: "tok_abc".to_string(),
323            refresh_token: Some("ref_xyz".to_string()),
324            expires_at: Some(9_999_999_999_999),
325        }
326    }
327
328    fn temp_store(passphrase: &str) -> EncryptedFileOAuthCredentialStorage {
329        let dir = tempfile::tempdir().unwrap();
330        let path = dir.keep().join("creds.enc");
331        EncryptedFileOAuthCredentialStorage::new(path, passphrase.to_string()).with_scrypt_work_factor(TEST_WORK_FACTOR)
332    }
333}