apollo/providers/
shared_credentials.rs1use rs_ai_oauth::credentials;
18use rs_ai_oauth::{OAuthProvider, OAuthTokens};
19
20pub type ApolloTokens = (String, Option<String>, i64);
27
28fn to_apollo(tokens: OAuthTokens) -> ApolloTokens {
29 let expires_at_ms = (tokens.expires_at as i64).saturating_mul(1000);
30 (tokens.access_token, tokens.refresh_token, expires_at_ms)
31}
32
33pub fn load(provider: OAuthProvider) -> Option<ApolloTokens> {
40 let tokens = credentials::load(&provider)?;
41 if credentials::is_expired(&tokens) && tokens.refresh_token.is_none() {
42 return None;
43 }
44 Some(to_apollo(tokens))
45}
46
47pub fn save(
49 provider: OAuthProvider,
50 access: &str,
51 refresh: &str,
52 expires_at_ms: i64,
53) -> anyhow::Result<()> {
54 let tokens = OAuthTokens {
55 access_token: access.to_string(),
56 refresh_token: Some(refresh.to_string()),
57 expires_at: (expires_at_ms / 1000).max(0) as u64,
58 };
59 credentials::save(&provider, &tokens)?;
60 Ok(())
61}
62
63pub struct SharedLogin {
65 pub provider: &'static str,
66 pub expired: bool,
67 pub refreshable: bool,
68}
69
70pub fn logins() -> Vec<SharedLogin> {
75 credentials::logged_in_providers()
76 .into_iter()
77 .filter_map(|provider| {
78 let tokens = credentials::load(&provider)?;
79 Some(SharedLogin {
80 provider: provider_name(provider),
81 expired: credentials::is_expired(&tokens),
82 refreshable: tokens.refresh_token.is_some(),
83 })
84 })
85 .collect()
86}
87
88fn provider_name(provider: OAuthProvider) -> &'static str {
91 match provider {
92 OAuthProvider::ChatGpt => "chatgpt",
93 OAuthProvider::Xai => "grok",
94 OAuthProvider::Claude => "claude",
95 OAuthProvider::Gemini => "gemini",
96 OAuthProvider::Antigravity => "antigravity",
97 OAuthProvider::Copilot => "copilot",
98 OAuthProvider::Kimi => "kimi",
99 }
100}
101
102#[cfg(test)]
103mod tests {
104 use super::*;
105
106 fn now_secs() -> u64 {
107 std::time::SystemTime::now()
108 .duration_since(std::time::UNIX_EPOCH)
109 .unwrap()
110 .as_secs()
111 }
112
113 #[test]
114 fn converts_expiry_seconds_to_milliseconds() {
115 let (access, refresh, expires) = to_apollo(OAuthTokens {
116 access_token: "a".into(),
117 refresh_token: Some("r".into()),
118 expires_at: 1_700_000_000,
119 });
120 assert_eq!(access, "a");
121 assert_eq!(refresh.as_deref(), Some("r"));
122 assert_eq!(expires, 1_700_000_000_000);
123 }
124
125 #[test]
126 fn zero_expiry_stays_zero() {
127 let (_, _, expires) = to_apollo(OAuthTokens {
128 access_token: "a".into(),
129 refresh_token: None,
130 expires_at: 0,
131 });
132 assert_eq!(expires, 0);
133 }
134
135 #[test]
138 fn save_then_load_round_trips_through_a_redirected_store() {
139 let dir = tempfile::tempdir().unwrap();
140 temp_env::with_var(
141 "RS_AI_CREDENTIALS_DIR",
142 Some(dir.path().as_os_str()),
143 || {
144 save(
145 OAuthProvider::Claude,
146 "access",
147 "refresh",
148 (now_secs() as i64 + 3600) * 1000,
149 )
150 .unwrap();
151
152 let (access, refresh, expires) = load(OAuthProvider::Claude).unwrap();
153 assert_eq!(access, "access");
154 assert_eq!(refresh.as_deref(), Some("refresh"));
155 assert!(expires > now_secs() as i64 * 1000);
156
157 let names: Vec<&str> = logins().iter().map(|l| l.provider).collect();
158 assert!(names.contains(&"claude"), "got {names:?}");
159 },
160 );
161 }
162
163 #[test]
164 fn expired_token_without_refresh_reads_as_absent() {
165 let dir = tempfile::tempdir().unwrap();
166 temp_env::with_var(
167 "RS_AI_CREDENTIALS_DIR",
168 Some(dir.path().as_os_str()),
169 || {
170 let tokens = OAuthTokens {
171 access_token: "dead".into(),
172 refresh_token: None,
173 expires_at: 1,
174 };
175 credentials::save(&OAuthProvider::Claude, &tokens).unwrap();
176 assert!(load(OAuthProvider::Claude).is_none());
177 },
178 );
179 }
180
181 #[test]
182 fn expired_token_with_refresh_is_still_returned() {
183 let dir = tempfile::tempdir().unwrap();
184 temp_env::with_var(
185 "RS_AI_CREDENTIALS_DIR",
186 Some(dir.path().as_os_str()),
187 || {
188 let tokens = OAuthTokens {
189 access_token: "stale".into(),
190 refresh_token: Some("r".into()),
191 expires_at: 1,
192 };
193 credentials::save(&OAuthProvider::Claude, &tokens).unwrap();
194 let (access, _, _) = load(OAuthProvider::Claude).unwrap();
195 assert_eq!(access, "stale");
196 },
197 );
198 }
199
200 #[test]
204 fn missing_provider_reads_as_absent() {
205 let dir = tempfile::tempdir().unwrap();
206 temp_env::with_var(
207 "RS_AI_CREDENTIALS_DIR",
208 Some(dir.path().as_os_str()),
209 || {
210 assert!(load(OAuthProvider::Kimi).is_none());
211 assert!(!logins().iter().any(|l| l.provider == "kimi"));
212 },
213 );
214 }
215}