Skip to main content

atman_runtime/
auth_store.rs

1use anyhow::{Context, Result};
2use serde::{Deserialize, Serialize};
3
4use crate::storage::config_dir;
5
6const AUTH_FILENAME: &str = "auth.json";
7
8#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
9#[serde(rename_all = "kebab-case")]
10pub enum ProviderKind {
11    Codex,
12    AnthropicOauth,
13    GitHubCopilot,
14    Custom,
15}
16
17#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct ModelCache {
19    pub fetched_at: i64,
20    pub models: Vec<CachedModel>,
21}
22
23#[derive(Debug, Clone, Serialize, Deserialize)]
24pub struct CachedModel {
25    pub slug: String,
26    #[serde(skip_serializing_if = "Option::is_none")]
27    pub context_budget: Option<u64>,
28    pub thinking: bool,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
32pub struct StoredProvider {
33    pub id: String,
34    pub name: String,
35    pub kind: ProviderKind,
36    pub access_token: String,
37    #[serde(skip_serializing_if = "Option::is_none")]
38    pub refresh_token: Option<String>,
39    pub expires_at: i64,
40    #[serde(skip_serializing_if = "Option::is_none")]
41    pub account: Option<String>,
42    pub enabled: bool,
43    #[serde(skip_serializing_if = "Option::is_none")]
44    pub model_cache: Option<ModelCache>,
45}
46
47#[derive(Debug, Clone, Default, Serialize, Deserialize)]
48pub struct AuthStore {
49    #[serde(default, skip_serializing_if = "Vec::is_empty")]
50    pub providers: Vec<StoredProvider>,
51}
52
53impl AuthStore {
54    pub fn load() -> Result<Self> {
55        let dir = config_dir().context("resolve config dir for auth.json")?;
56        let path = dir.join(AUTH_FILENAME);
57        if !path.exists() {
58            return Ok(Self::default());
59        }
60        let bytes = std::fs::read(&path).with_context(|| format!("read {}", path.display()))?;
61        let store: Self =
62            serde_json::from_slice(&bytes).with_context(|| format!("parse {}", path.display()))?;
63        Ok(store)
64    }
65
66    pub fn save_to(&self, path: &std::path::Path) -> Result<()> {
67        if let Some(parent) = path.parent() {
68            std::fs::create_dir_all(parent)
69                .with_context(|| format!("mkdir {}", parent.display()))?;
70        }
71        let tmp = path.with_file_name(format!(".{}.tmp", AUTH_FILENAME));
72        let json = serde_json::to_vec_pretty(self).context("serialize auth store")?;
73        std::fs::write(&tmp, &json).with_context(|| format!("write {}", tmp.display()))?;
74        std::fs::rename(&tmp, path)
75            .with_context(|| format!("rename {} -> {}", tmp.display(), path.display()))?;
76        #[cfg(unix)]
77        {
78            use std::os::unix::fs::PermissionsExt;
79            let _ = std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600));
80        }
81        Ok(())
82    }
83
84    pub fn save(&self) -> Result<()> {
85        let dir = config_dir().context("resolve config dir for auth.json")?;
86        std::fs::create_dir_all(&dir).with_context(|| format!("mkdir {}", dir.display()))?;
87        let path = dir.join(AUTH_FILENAME);
88        self.save_to(&path)
89    }
90
91    pub fn add(&mut self, p: StoredProvider) {
92        self.providers.push(p);
93    }
94
95    pub fn remove(&mut self, id: &str) -> bool {
96        let len_before = self.providers.len();
97        self.providers.retain(|p| p.id != id);
98        self.providers.len() < len_before
99    }
100
101    /// Update the model cache for a provider by ID. Returns false if provider not found.
102    pub fn update_model_cache(&mut self, provider_id: &str, cache: ModelCache) -> bool {
103        if let Some(p) = self.providers.iter_mut().find(|p| p.id == provider_id) {
104            p.model_cache = Some(cache);
105            true
106        } else {
107            false
108        }
109    }
110}
111
112/// Save discovered models as cache for a provider. Reads auth.json, updates, writes back.
113pub fn save_provider_model_cache(
114    provider_id: &str,
115    models: &[crate::provider::DiscoveredModel],
116) -> Result<()> {
117    let mut store = AuthStore::load().unwrap_or_default();
118    let cache = ModelCache {
119        fetched_at: chrono::Utc::now().timestamp(),
120        models: models
121            .iter()
122            .map(|m| CachedModel {
123                slug: m.slug.clone(),
124                context_budget: m.context_budget,
125                thinking: m.thinking,
126            })
127            .collect(),
128    };
129    store.update_model_cache(provider_id, cache);
130    store.save()
131}
132
133/// Convert cached models to discovered models for registry hydration.
134pub fn cached_to_discovered(cache: &ModelCache) -> Vec<crate::provider::DiscoveredModel> {
135    cache
136        .models
137        .iter()
138        .map(|m| crate::provider::DiscoveredModel {
139            slug: m.slug.clone(),
140            context_budget: m.context_budget,
141            thinking: m.thinking,
142        })
143        .collect()
144}
145
146#[cfg(test)]
147mod tests {
148    use super::*;
149    use tempfile::TempDir;
150
151    #[test]
152    fn load_returns_empty_when_file_missing() {
153        let tmp = TempDir::new().unwrap();
154        let path = tmp.path().join("auth.json");
155        let store: AuthStore = std::fs::read(&path)
156            .ok()
157            .and_then(|b| serde_json::from_slice(&b).ok())
158            .unwrap_or_default();
159        assert!(store.providers.is_empty());
160    }
161
162    #[test]
163    fn save_then_load_round_trips() {
164        let tmp = TempDir::new().unwrap();
165        let path = tmp.path().join("auth.json");
166        let mut store = AuthStore::default();
167        store.add(StoredProvider {
168            id: "test-1".into(),
169            name: "Personal Codex".into(),
170            kind: ProviderKind::Codex,
171            access_token: "tok1".into(),
172            refresh_token: Some("rt1".into()),
173            expires_at: 1761735358,
174            account: Some("x@example.com".into()),
175            enabled: true,
176            model_cache: None,
177        });
178        store.save_to(&path).unwrap();
179
180        let bytes = std::fs::read(&path).unwrap();
181        let loaded: AuthStore = serde_json::from_slice(&bytes).unwrap();
182        assert_eq!(loaded.providers.len(), 1);
183        assert_eq!(loaded.providers[0].name, "Personal Codex");
184    }
185
186    #[test]
187    fn remove_existing_id_returns_true() {
188        let mut store = AuthStore::default();
189        store.add(StoredProvider {
190            id: "keep".into(),
191            name: "A".into(),
192            kind: ProviderKind::Custom,
193            access_token: "t".into(),
194            refresh_token: None,
195            expires_at: 0,
196            account: None,
197            enabled: true,
198            model_cache: None,
199        });
200        store.add(StoredProvider {
201            id: "del".into(),
202            name: "B".into(),
203            kind: ProviderKind::Custom,
204            access_token: "t".into(),
205            refresh_token: None,
206            expires_at: 0,
207            account: None,
208            enabled: true,
209            model_cache: None,
210        });
211        assert!(store.remove("del"));
212        assert_eq!(store.providers.len(), 1);
213        assert_eq!(store.providers[0].id, "keep");
214    }
215
216    #[test]
217    fn remove_missing_id_returns_false() {
218        let mut store = AuthStore::default();
219        assert!(!store.remove("nope"));
220    }
221
222    #[test]
223    fn provider_kind_serde_round_trip() {
224        let kinds = vec![
225            ProviderKind::Codex,
226            ProviderKind::AnthropicOauth,
227            ProviderKind::GitHubCopilot,
228            ProviderKind::Custom,
229        ];
230        for k in kinds {
231            let json = serde_json::to_string(&k).unwrap();
232            let back: ProviderKind = serde_json::from_str(&json).unwrap();
233            assert_eq!(back, k);
234        }
235    }
236
237    #[test]
238    fn auth_store_serde_backward_compat_empty_json() {
239        let store: AuthStore = serde_json::from_str("{}").unwrap();
240        assert!(store.providers.is_empty());
241    }
242
243    #[test]
244    fn auth_store_serde_backward_compat_null_providers() {
245        let json = r#"{"providers": []}"#;
246        let store: AuthStore = serde_json::from_str(json).unwrap();
247        assert!(store.providers.is_empty());
248    }
249}