Skip to main content

aether_auth/
encrypted_file.rs

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