Skip to main content

claude_codex/providers/codex/auth/
manager.rs

1use 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        // Reload from durable storage after acquiring the single-flight lock.
92        // Another request may have rotated and persisted the token while this
93        // caller was waiting.
94        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(&current).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}