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