claude_codex/providers/codex/auth/
manager.rs1use std::sync::Arc;
2#[cfg(test)]
3use std::sync::Mutex;
4use std::time::{Duration, SystemTime, UNIX_EPOCH};
5use tokio::sync::Mutex as AsyncMutex;
6
7use super::constants::{CLIENT_ID, ISSUER, REFRESH_MARGIN_MS};
8use super::jwt::{TokenResponse, extract_account_id, validate_token_response};
9use super::token_store::{CodexTokenStore, StoredAuth};
10use crate::auth::AuthStorage;
11
12pub struct CodexAuthManager<S: AuthStorage<StoredAuth>> {
13 pub store: CodexTokenStore<S>,
14 #[cfg(test)]
15 test_auth: Arc<Mutex<Option<StoredAuth>>>,
16 refresh_lock: Arc<AsyncMutex<()>>,
17 refresh_client: reqwest::Client,
18 token_endpoint: String,
19}
20
21impl<S: AuthStorage<StoredAuth>> CodexAuthManager<S> {
22 pub fn new(store: CodexTokenStore<S>) -> Self {
23 Self::new_with_token_endpoint(store, format!("{ISSUER}/oauth/token"))
24 }
25
26 fn new_with_token_endpoint(store: CodexTokenStore<S>, token_endpoint: String) -> Self {
27 Self {
28 store,
29 #[cfg(test)]
30 test_auth: Arc::new(Mutex::new(None)),
31 refresh_lock: Arc::new(AsyncMutex::new(())),
32 refresh_client: reqwest::Client::builder()
33 .connect_timeout(Duration::from_secs(15))
34 .timeout(Duration::from_secs(30))
35 .build()
36 .expect("failed to create Codex OAuth refresh client"),
37 token_endpoint,
38 }
39 }
40
41 fn now_ms() -> u64 {
42 SystemTime::now()
43 .duration_since(UNIX_EPOCH)
44 .unwrap_or_default()
45 .as_millis() as u64
46 }
47
48 pub async fn get_auth(&self) -> Result<StoredAuth, anyhow::Error> {
49 let stored = self.load_auth()?.ok_or_else(|| {
50 anyhow::anyhow!(
51 "No Codex credentials. Log in with the Codex CLI (`codex login`) to create ~/.codex/auth.json"
52 )
53 })?;
54
55 if stored.expires > Self::now_ms() + REFRESH_MARGIN_MS {
56 return Ok(stored);
57 }
58
59 self.refresh(false, None).await
60 }
61
62 pub async fn force_refresh(&self, rejected_access: &str) -> Result<StoredAuth, anyhow::Error> {
63 self.refresh(true, Some(rejected_access)).await
64 }
65
66 fn load_auth(&self) -> Result<Option<StoredAuth>, anyhow::Error> {
67 #[cfg(test)]
68 if let Some(auth) = self
69 .test_auth
70 .lock()
71 .map_err(|e| anyhow::anyhow!("{e}"))?
72 .clone()
73 {
74 return Ok(Some(auth));
75 }
76
77 self.store.load_auth()
78 }
79
80 async fn refresh(
81 &self,
82 force: bool,
83 rejected_access: Option<&str>,
84 ) -> Result<StoredAuth, anyhow::Error> {
85 let _refresh_guard = self.refresh_lock.lock().await;
86
87 let current = self
91 .load_auth()?
92 .ok_or_else(|| anyhow::anyhow!("Not authenticated"))?;
93
94 if (!force && current.expires > Self::now_ms() + REFRESH_MARGIN_MS)
95 || rejected_access.is_some_and(|access| current.access != access)
96 {
97 return Ok(current);
98 }
99
100 self.refresh_now(¤t).await
101 }
102
103 async fn refresh_now(&self, current: &StoredAuth) -> Result<StoredAuth, anyhow::Error> {
104 if current.refresh.is_empty() {
105 anyhow::bail!("No refresh token stored; re-authenticate");
106 }
107
108 let form = [
109 ("client_id", CLIENT_ID.to_string()),
110 ("grant_type", "refresh_token".to_string()),
111 ("refresh_token", current.refresh.clone()),
112 ];
113
114 let resp = self
115 .refresh_client
116 .post(&self.token_endpoint)
117 .form(&form)
118 .send()
119 .await
120 .map_err(|e| anyhow::anyhow!("refresh network error: {e}"))?;
121
122 let status = resp.status().as_u16();
123 if status == 401 || status == 403 {
124 if let Some(latest) = self.store.load_auth()?
125 && latest != *current
126 {
127 return Ok(latest);
128 }
129 self.store.clear_auth()?;
130 let err_msg = resp
131 .text()
132 .await
133 .unwrap_or_else(|_| "Token refresh unauthorized".to_string());
134 anyhow::bail!("{err_msg}");
135 }
136
137 if !resp.status().is_success() {
138 anyhow::bail!("Token refresh failed: {status}");
139 }
140
141 let tokens: TokenResponse = resp
142 .json()
143 .await
144 .map_err(|e| anyhow::anyhow!("failed to parse token response: {e}"))?;
145 validate_token_response(&tokens)?;
146 let account_id = extract_account_id(&tokens).or_else(|| current.account_id.clone());
147 let expires = Self::now_ms() + (tokens.expires_in.unwrap_or(3600) * 1000);
148 let next = StoredAuth {
149 access: tokens.access_token,
150 refresh: tokens.refresh_token,
151 expires,
152 account_id,
153 };
154 self.store.save_auth(next.clone())?;
155 Ok(next)
156 }
157
158 pub fn persist_initial_tokens(
159 &self,
160 tokens: &TokenResponse,
161 ) -> Result<StoredAuth, anyhow::Error> {
162 validate_token_response(tokens)?;
163 let account_id = extract_account_id(tokens);
164 let expires = Self::now_ms() + (tokens.expires_in.unwrap_or(3600) * 1000);
165 let auth = StoredAuth {
166 access: tokens.access_token.clone(),
167 refresh: tokens.refresh_token.clone(),
168 expires,
169 account_id,
170 };
171 self.store.save_auth(auth.clone())?;
172 Ok(auth)
173 }
174
175 #[cfg(test)]
176 pub fn set_test_auth(&self, auth: StoredAuth) {
177 if let Ok(mut guard) = self.test_auth.lock() {
178 *guard = Some(auth);
179 }
180 }
181}
182
183#[cfg(test)]
184mod tests {
185 use super::*;
186 use crate::auth::InMemoryAuthStore;
187 use std::io::{Read, Write};
188 use std::net::TcpListener;
189 use std::sync::Arc;
190 use std::sync::atomic::{AtomicUsize, Ordering};
191 use std::thread;
192
193 fn test_store() -> CodexTokenStore<InMemoryAuthStore<StoredAuth>> {
194 CodexTokenStore::new(InMemoryAuthStore::new())
195 }
196
197 #[tokio::test]
198 async fn get_auth_returns_stored() {
199 let store = test_store();
200 let auth = StoredAuth {
201 access: "test_access".into(),
202 refresh: "test_refresh".into(),
203 expires: 9999999999999,
204 account_id: Some("acct_1".into()),
205 };
206 store.save_auth(auth.clone()).unwrap();
207 let manager = CodexAuthManager::new(store);
208 let result = manager.get_auth().await.unwrap();
209 assert_eq!(result.access, "test_access");
210 assert_eq!(result.account_id.as_deref(), Some("acct_1"));
211 }
212
213 #[tokio::test]
214 async fn get_auth_fails_when_no_auth() {
215 let store = test_store();
216 let manager = CodexAuthManager::new(store);
217 assert!(manager.get_auth().await.is_err());
218 assert!(
219 manager
220 .get_auth()
221 .await
222 .unwrap_err()
223 .to_string()
224 .contains("No Codex credentials")
225 );
226 }
227
228 #[tokio::test(flavor = "multi_thread")]
229 async fn concurrent_expired_auth_refreshes_once() {
230 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
231 let addr = listener.local_addr().unwrap();
232 let refreshes = Arc::new(AtomicUsize::new(0));
233 let server_refreshes = refreshes.clone();
234 let server = thread::spawn(move || {
235 let (mut stream, _) = listener.accept().unwrap();
236 let mut request = [0u8; 4096];
237 let read = stream.read(&mut request).unwrap();
238 assert!(read > 0);
239 assert!(String::from_utf8_lossy(&request[..read]).contains("refresh_token=stale"));
240 server_refreshes.fetch_add(1, Ordering::SeqCst);
241
242 let body = br#"{"access_token":"rotated","refresh_token":"rotated-refresh","expires_in":3600}"#;
243 let response = format!(
244 "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n",
245 body.len()
246 );
247 stream.write_all(response.as_bytes()).unwrap();
248 stream.write_all(body).unwrap();
249 });
250
251 let store = test_store();
252 store
253 .save_auth(StoredAuth {
254 access: "expired".into(),
255 refresh: "stale".into(),
256 expires: 0,
257 account_id: Some("acct_1".into()),
258 })
259 .unwrap();
260 let manager = Arc::new(CodexAuthManager::new_with_token_endpoint(
261 store,
262 format!("http://{addr}/oauth/token"),
263 ));
264 let (first, second) = tokio::join!(manager.get_auth(), manager.get_auth());
265 let results = [first.unwrap(), second.unwrap()];
266 server.join().unwrap();
267
268 assert_eq!(refreshes.load(Ordering::SeqCst), 1);
269 assert!(results.iter().all(|auth| auth.access == "rotated"));
270 assert!(results.iter().all(|auth| auth.refresh == "rotated-refresh"));
271 }
272
273 #[tokio::test]
274 async fn stale_401_reuses_already_rotated_auth() {
275 let store = test_store();
276 store
277 .save_auth(StoredAuth {
278 access: "rotated".into(),
279 refresh: "rotated-refresh".into(),
280 expires: u64::MAX,
281 account_id: Some("acct_1".into()),
282 })
283 .unwrap();
284 let manager = CodexAuthManager::new_with_token_endpoint(
285 store,
286 "http://127.0.0.1:1/should-not-be-called".into(),
287 );
288
289 let auth = manager.force_refresh("rejected").await.unwrap();
290 assert_eq!(auth.access, "rotated");
291 assert_eq!(auth.refresh, "rotated-refresh");
292 }
293
294 #[tokio::test]
295 async fn unauthorized_refresh_preserves_changed_refresh_token() {
296 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
297 let addr = listener.local_addr().unwrap();
298 let backing = InMemoryAuthStore::new();
299 let server_backing = backing.clone();
300 let server = thread::spawn(move || {
301 let (mut stream, _) = listener.accept().unwrap();
302 let mut request = [0u8; 4096];
303 assert!(stream.read(&mut request).unwrap() > 0);
304 server_backing
305 .save(StoredAuth {
306 access: "same-access".into(),
307 refresh: "replacement-refresh".into(),
308 expires: u64::MAX,
309 account_id: Some("acct_1".into()),
310 })
311 .unwrap();
312 let body = b"rejected refresh token";
313 let response = format!(
314 "HTTP/1.1 401 Unauthorized\r\ncontent-length: {}\r\nconnection: close\r\n\r\n",
315 body.len()
316 );
317 stream.write_all(response.as_bytes()).unwrap();
318 stream.write_all(body).unwrap();
319 });
320
321 let store = CodexTokenStore::new(backing);
322 store
323 .save_auth(StoredAuth {
324 access: "same-access".into(),
325 refresh: "rejected-refresh".into(),
326 expires: 0,
327 account_id: Some("acct_1".into()),
328 })
329 .unwrap();
330 let manager =
331 CodexAuthManager::new_with_token_endpoint(store, format!("http://{addr}/oauth/token"));
332
333 let auth = manager.get_auth().await.unwrap();
334 server.join().unwrap();
335 assert_eq!(auth.access, "same-access");
336 assert_eq!(auth.refresh, "replacement-refresh");
337 assert_eq!(manager.store.load_auth().unwrap(), Some(auth));
338 }
339
340 #[tokio::test]
341 async fn durable_rotation_and_logout_are_observed_by_shared_manager() {
342 let store = test_store();
343 store
344 .save_auth(StoredAuth {
345 access: "first".into(),
346 refresh: "first-refresh".into(),
347 expires: u64::MAX,
348 account_id: Some("acct_1".into()),
349 })
350 .unwrap();
351 let manager = CodexAuthManager::new(store);
352 assert_eq!(manager.get_auth().await.unwrap().access, "first");
353
354 manager
355 .store
356 .save_auth(StoredAuth {
357 access: "rotated".into(),
358 refresh: "rotated-refresh".into(),
359 expires: u64::MAX,
360 account_id: Some("acct_2".into()),
361 })
362 .unwrap();
363 let rotated = manager.get_auth().await.unwrap();
364 assert_eq!(rotated.access, "rotated");
365 assert_eq!(rotated.account_id.as_deref(), Some("acct_2"));
366
367 manager.store.clear_auth().unwrap();
368 assert!(manager.get_auth().await.is_err());
369 }
370}