Skip to main content

claude_codex/providers/codex/auth/
manager.rs

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