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(|provider| !matches!(provider, OAuthProvider::Claude))
78 .filter_map(|provider| {
79 let tokens = credentials::load(&provider)?;
80 Some(SharedLogin {
81 provider: provider_name(provider),
82 expired: credentials::is_expired(&tokens),
83 refreshable: tokens.refresh_token.is_some(),
84 })
85 })
86 .collect()
87}
88
89fn provider_name(provider: OAuthProvider) -> &'static str {
92 match provider {
93 OAuthProvider::ChatGpt => "chatgpt",
94 OAuthProvider::Xai => "grok",
95 OAuthProvider::Claude => "claude",
96 OAuthProvider::Gemini => "gemini",
97 OAuthProvider::Antigravity => "antigravity",
98 OAuthProvider::Copilot => "copilot",
99 OAuthProvider::Kimi => "kimi",
100 }
101}
102
103#[cfg(test)]
104mod tests {
105 use super::*;
106
107 fn now_secs() -> u64 {
108 std::time::SystemTime::now()
109 .duration_since(std::time::UNIX_EPOCH)
110 .unwrap()
111 .as_secs()
112 }
113
114 #[test]
115 fn converts_expiry_seconds_to_milliseconds() {
116 let (access, refresh, expires) = to_apollo(OAuthTokens {
117 access_token: "a".into(),
118 refresh_token: Some("r".into()),
119 expires_at: 1_700_000_000,
120 });
121 assert_eq!(access, "a");
122 assert_eq!(refresh.as_deref(), Some("r"));
123 assert_eq!(expires, 1_700_000_000_000);
124 }
125
126 #[test]
127 fn zero_expiry_stays_zero() {
128 let (_, _, expires) = to_apollo(OAuthTokens {
129 access_token: "a".into(),
130 refresh_token: None,
131 expires_at: 0,
132 });
133 assert_eq!(expires, 0);
134 }
135
136 #[test]
139 fn save_then_load_round_trips_through_a_redirected_store() {
140 let dir = tempfile::tempdir().unwrap();
141 temp_env::with_var(
142 "RS_AI_CREDENTIALS_DIR",
143 Some(dir.path().as_os_str()),
144 || {
145 save(
146 OAuthProvider::ChatGpt,
147 "access",
148 "refresh",
149 (now_secs() as i64 + 3600) * 1000,
150 )
151 .unwrap();
152
153 let (access, refresh, expires) = load(OAuthProvider::ChatGpt).unwrap();
154 assert_eq!(access, "access");
155 assert_eq!(refresh.as_deref(), Some("refresh"));
156 assert!(expires > now_secs() as i64 * 1000);
157
158 let names: Vec<&str> = logins().iter().map(|l| l.provider).collect();
159 assert!(names.contains(&"chatgpt"), "got {names:?}");
160 },
161 );
162 }
163
164 #[test]
165 fn expired_token_without_refresh_reads_as_absent() {
166 let dir = tempfile::tempdir().unwrap();
167 temp_env::with_var(
168 "RS_AI_CREDENTIALS_DIR",
169 Some(dir.path().as_os_str()),
170 || {
171 let tokens = OAuthTokens {
172 access_token: "dead".into(),
173 refresh_token: None,
174 expires_at: 1,
175 };
176 credentials::save(&OAuthProvider::Claude, &tokens).unwrap();
177 assert!(load(OAuthProvider::Claude).is_none());
178 },
179 );
180 }
181
182 #[test]
183 fn expired_token_with_refresh_is_still_returned() {
184 let dir = tempfile::tempdir().unwrap();
185 temp_env::with_var(
186 "RS_AI_CREDENTIALS_DIR",
187 Some(dir.path().as_os_str()),
188 || {
189 let tokens = OAuthTokens {
190 access_token: "stale".into(),
191 refresh_token: Some("r".into()),
192 expires_at: 1,
193 };
194 credentials::save(&OAuthProvider::Claude, &tokens).unwrap();
195 let (access, _, _) = load(OAuthProvider::Claude).unwrap();
196 assert_eq!(access, "stale");
197 },
198 );
199 }
200
201 #[test]
205 fn missing_provider_reads_as_absent() {
206 let dir = tempfile::tempdir().unwrap();
207 temp_env::with_var(
208 "RS_AI_CREDENTIALS_DIR",
209 Some(dir.path().as_os_str()),
210 || {
211 assert!(load(OAuthProvider::Kimi).is_none());
212 assert!(!logins().iter().any(|l| l.provider == "kimi"));
213 },
214 );
215 }
216}