Skip to main content

atman_runtime/
auth_store.rs

1use anyhow::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        Ok(crate::config_hub::ConfigHub::global()?.load_auth()?)
56    }
57
58    pub fn save_to(&self, path: &std::path::Path) -> Result<()> {
59        let hub = crate::config_hub::ConfigHub::from_auth_path(path);
60        hub.update_auth(|store| {
61            *store = self.clone();
62            Ok(())
63        })?;
64        Ok(())
65    }
66
67    pub fn save(&self) -> Result<()> {
68        let dir = config_dir()?;
69        self.save_to(&dir.join(AUTH_FILENAME))
70    }
71
72    pub fn add(&mut self, p: StoredProvider) {
73        self.providers.push(p);
74    }
75
76    pub fn remove(&mut self, id: &str) -> bool {
77        let len_before = self.providers.len();
78        self.providers.retain(|p| p.id != id);
79        self.providers.len() < len_before
80    }
81
82    /// Update the model cache for a provider by ID. Returns false if provider not found.
83    pub fn update_model_cache(&mut self, provider_id: &str, cache: ModelCache) -> bool {
84        if let Some(p) = self.providers.iter_mut().find(|p| p.id == provider_id) {
85            p.model_cache = Some(cache);
86            true
87        } else {
88            false
89        }
90    }
91}
92
93/// Save discovered models as cache for a provider. Reads auth.json, updates, writes back.
94pub fn save_provider_model_cache(
95    provider_id: &str,
96    models: &[crate::provider::DiscoveredModel],
97) -> Result<()> {
98    let cache = ModelCache {
99        fetched_at: chrono::Utc::now().timestamp(),
100        models: models
101            .iter()
102            .map(|model| CachedModel {
103                slug: model.slug.clone(),
104                context_budget: model.context_budget,
105                thinking: model.thinking,
106            })
107            .collect(),
108    };
109    crate::config_hub::ConfigHub::global()?.update_auth_model_cache(provider_id, cache)?;
110    Ok(())
111}
112
113/// Convert cached models to discovered models for registry hydration.
114pub fn cached_to_discovered(cache: &ModelCache) -> Vec<crate::provider::DiscoveredModel> {
115    cache
116        .models
117        .iter()
118        .map(|m| crate::provider::DiscoveredModel {
119            slug: m.slug.clone(),
120            context_budget: m.context_budget,
121            thinking: m.thinking,
122        })
123        .collect()
124}
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129    use tempfile::TempDir;
130
131    #[test]
132    fn load_returns_empty_when_file_missing() {
133        let tmp = TempDir::new().unwrap();
134        let path = tmp.path().join("auth.json");
135        let store: AuthStore = std::fs::read(&path)
136            .ok()
137            .and_then(|b| serde_json::from_slice(&b).ok())
138            .unwrap_or_default();
139        assert!(store.providers.is_empty());
140    }
141
142    #[test]
143    fn save_then_load_round_trips() {
144        let tmp = TempDir::new().unwrap();
145        let path = tmp.path().join("auth.json");
146        let mut store = AuthStore::default();
147        store.add(StoredProvider {
148            id: "test-1".into(),
149            name: "Personal Codex".into(),
150            kind: ProviderKind::Codex,
151            access_token: "tok1".into(),
152            refresh_token: Some("rt1".into()),
153            expires_at: 1761735358,
154            account: Some("x@example.com".into()),
155            enabled: true,
156            model_cache: None,
157        });
158        store.save_to(&path).unwrap();
159
160        let bytes = std::fs::read(&path).unwrap();
161        let loaded: AuthStore = serde_json::from_slice(&bytes).unwrap();
162        assert_eq!(loaded.providers.len(), 1);
163        assert_eq!(loaded.providers[0].name, "Personal Codex");
164    }
165
166    #[test]
167    fn remove_existing_id_returns_true() {
168        let mut store = AuthStore::default();
169        store.add(StoredProvider {
170            id: "keep".into(),
171            name: "A".into(),
172            kind: ProviderKind::Custom,
173            access_token: "t".into(),
174            refresh_token: None,
175            expires_at: 0,
176            account: None,
177            enabled: true,
178            model_cache: None,
179        });
180        store.add(StoredProvider {
181            id: "del".into(),
182            name: "B".into(),
183            kind: ProviderKind::Custom,
184            access_token: "t".into(),
185            refresh_token: None,
186            expires_at: 0,
187            account: None,
188            enabled: true,
189            model_cache: None,
190        });
191        assert!(store.remove("del"));
192        assert_eq!(store.providers.len(), 1);
193        assert_eq!(store.providers[0].id, "keep");
194    }
195
196    #[test]
197    fn remove_missing_id_returns_false() {
198        let mut store = AuthStore::default();
199        assert!(!store.remove("nope"));
200    }
201
202    #[test]
203    fn provider_kind_serde_round_trip() {
204        let kinds = vec![
205            ProviderKind::Codex,
206            ProviderKind::AnthropicOauth,
207            ProviderKind::GitHubCopilot,
208            ProviderKind::Custom,
209        ];
210        for k in kinds {
211            let json = serde_json::to_string(&k).unwrap();
212            let back: ProviderKind = serde_json::from_str(&json).unwrap();
213            assert_eq!(back, k);
214        }
215    }
216
217    #[test]
218    fn auth_store_serde_backward_compat_empty_json() {
219        let store: AuthStore = serde_json::from_str("{}").unwrap();
220        assert!(store.providers.is_empty());
221    }
222
223    #[test]
224    fn auth_store_serde_backward_compat_null_providers() {
225        let json = r#"{"providers": []}"#;
226        let store: AuthStore = serde_json::from_str(json).unwrap();
227        assert!(store.providers.is_empty());
228    }
229}