1use std::collections::HashMap;
2use std::future::Future;
3use std::path::PathBuf;
4use std::pin::Pin;
5use std::sync::{Arc, LazyLock, Mutex, Weak};
6use std::time::{Instant, SystemTime, UNIX_EPOCH};
7
8use anyhow::Result;
9use base64::Engine;
10use base64::engine::general_purpose::URL_SAFE_NO_PAD;
11use futures::FutureExt;
12use futures::future::{BoxFuture, Shared};
13use rand::RngCore;
14use sha2::{Digest, Sha256};
15
16use crate::auth_store::{AuthCredentialCommit, ProviderKind, StoredProvider};
17use crate::config_hub::{AuthTokenUpdate, ConfigHub};
18use crate::provider::{DiscoveredModel, DiscoveredModelDetails, Provider};
19
20const REFRESH_WINDOW_SECONDS: i64 = 300;
21const REFRESH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(60);
22const REFRESH_LOCK_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(75);
23#[cfg(not(test))]
24const REFRESH_FAILURE_COOLDOWN: std::time::Duration = std::time::Duration::from_secs(5);
25#[cfg(test)]
26const REFRESH_FAILURE_COOLDOWN: std::time::Duration = std::time::Duration::from_millis(50);
27
28type RefreshFuture = Pin<Box<dyn Future<Output = Result<TokenResult>> + Send>>;
29type RefreshFn = dyn Fn(String) -> RefreshFuture + Send + Sync;
30type CredentialResult = std::result::Result<OAuthCredential, OAuthCredentialError>;
31type SharedRefreshFuture = Shared<BoxFuture<'static, SharedRefreshOutcome>>;
32
33#[derive(Debug, Clone, PartialEq, Eq, Hash)]
34struct CredentialGateKey {
35 auth_path: PathBuf,
36 provider_id: String,
37}
38
39static CREDENTIAL_GATES: LazyLock<Mutex<HashMap<CredentialGateKey, CredentialGateEntry>>> =
40 LazyLock::new(|| Mutex::new(HashMap::new()));
41
42struct CredentialGate {
43 key: CredentialGateKey,
44 state: Mutex<CredentialGateState>,
45}
46
47struct CredentialGateEntry {
48 weak: Weak<CredentialGate>,
49 quarantined: Option<Arc<CredentialGate>>,
52}
53
54#[derive(Default)]
55struct CredentialGateState {
56 in_flight: Option<SharedRefreshFuture>,
57 pending_retry_active: bool,
58 failure: Option<CachedRefreshFailure>,
59 pending: Option<PendingCredentialCommit>,
60 quarantine: Option<OAuthCredentialError>,
61}
62
63impl CredentialGate {
64 fn ensure_not_quarantined(&self) -> std::result::Result<(), OAuthCredentialError> {
65 let state = self
66 .state
67 .lock()
68 .unwrap_or_else(std::sync::PoisonError::into_inner);
69 match &state.quarantine {
70 Some(error) => Err(error.clone()),
71 None => Ok(()),
72 }
73 }
74
75 fn quarantine(self: &Arc<Self>) -> OAuthCredentialError {
76 let mut gates = CREDENTIAL_GATES
77 .lock()
78 .unwrap_or_else(std::sync::PoisonError::into_inner);
79 let entry = gates
80 .entry(self.key.clone())
81 .or_insert_with(|| CredentialGateEntry {
82 weak: Arc::downgrade(self),
83 quarantined: None,
84 });
85 entry.weak = Arc::downgrade(self);
86 entry.quarantined = Some(self.clone());
87
88 let mut state = self
89 .state
90 .lock()
91 .unwrap_or_else(std::sync::PoisonError::into_inner);
92 let (error, newly_quarantined) = match &state.quarantine {
93 Some(error) => (error.clone(), false),
94 None => {
95 let error = OAuthCredentialError::Quarantined(self.key.provider_id.clone());
96 state.quarantine = Some(error.clone());
97 (error, true)
98 }
99 };
100 let in_flight = state.in_flight.take();
101 let pending = state.pending.take();
102 state.pending_retry_active = false;
103 state.failure = None;
104 drop(state);
105 drop(gates);
106 drop(in_flight);
107 drop(pending);
108
109 if newly_quarantined {
110 let notification = error.clone();
113 let _ = std::thread::Builder::new()
114 .name("atman-oauth-quarantine-notify".into())
115 .spawn(move || {
116 let _ =
117 crate::panic_capture::blocking(|| crate::notify!(error, "{notification}"));
118 });
119 }
120 error
121 }
122}
123
124struct CachedRefreshFailure {
125 until: Instant,
126 error: OAuthCredentialError,
127}
128
129#[derive(Clone)]
130struct PendingCredentialCommit {
131 snapshot: crate::auth_store::AuthProviderCredentialSnapshot,
132 access_token: String,
133 refresh_token: Option<String>,
134 expires_at: i64,
135 account: Option<String>,
136 _refresh_lock: Arc<RefreshFileLock>,
137}
138
139impl PendingCredentialCommit {
140 fn update(&self) -> AuthTokenUpdate {
141 AuthTokenUpdate {
142 access_token: self.access_token.clone(),
143 refresh_token: self.refresh_token.clone(),
144 expires_at: self.expires_at,
145 account: self.account.clone(),
146 }
147 }
148}
149
150#[derive(Clone)]
151struct RefreshFlightResult {
152 result: CredentialResult,
153 pending: Option<PendingCredentialCommit>,
154}
155
156#[derive(Clone)]
157enum SharedRefreshOutcome {
158 Complete(RefreshFlightResult),
159 Quarantined(OAuthCredentialError),
160}
161
162impl RefreshFlightResult {
163 fn complete(result: CredentialResult) -> Self {
164 Self {
165 result,
166 pending: None,
167 }
168 }
169
170 fn pending(error: OAuthCredentialError, pending: PendingCredentialCommit) -> Self {
171 Self {
172 result: Err(error),
173 pending: Some(pending),
174 }
175 }
176}
177
178pub struct Pkce {
179 pub verifier: String,
180 pub challenge: String,
181}
182
183impl Pkce {
184 pub fn generate() -> Self {
185 let mut bytes = [0u8; 32];
186 rand::thread_rng().fill_bytes(&mut bytes);
187 let verifier = URL_SAFE_NO_PAD.encode(bytes);
188
189 let mut hasher = Sha256::new();
190 hasher.update(verifier.as_bytes());
191 let digest = hasher.finalize();
192 let challenge = URL_SAFE_NO_PAD.encode(digest);
193
194 Pkce {
195 verifier,
196 challenge,
197 }
198 }
199}
200
201pub struct TokenResult {
202 pub access_token: String,
203 pub refresh_token: Option<String>,
204 pub expires_at: i64,
205 pub account: Option<String>,
206}
207
208#[derive(Clone, PartialEq, Eq)]
209pub(crate) struct OAuthCredential {
210 pub access_token: String,
211 pub display_account: Option<String>,
212}
213
214impl std::fmt::Debug for OAuthCredential {
215 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
216 formatter
217 .debug_struct("OAuthCredential")
218 .field("access_token", &"[redacted]")
219 .field(
220 "display_account",
221 &self.display_account.as_ref().map(|_| "[redacted]"),
222 )
223 .finish()
224 }
225}
226
227#[derive(Debug, Clone, thiserror::Error)]
228pub(crate) enum OAuthCredentialError {
229 #[error("OAuth provider `{0}` is not configured")]
230 Missing(String),
231 #[error("OAuth provider `{0}` is disabled")]
232 Disabled(String),
233 #[error("OAuth provider `{0}` changed kind")]
234 KindChanged(String),
235 #[error("OAuth provider `{0}` must be authenticated again")]
236 ReauthenticationRequired(String),
237 #[error("OAuth credential snapshot for `{0}` expired; use a managed provider to refresh it")]
238 SnapshotExpired(String),
239 #[error("load OAuth credentials for `{provider_id}`: {message}")]
240 Load {
241 provider_id: String,
242 message: String,
243 },
244 #[error("lock OAuth credential refresh for `{provider_id}`: {message}")]
245 Lock {
246 provider_id: String,
247 message: String,
248 },
249 #[error("refresh OAuth credentials for `{provider_id}`: {message}")]
250 Refresh {
251 provider_id: String,
252 message: String,
253 },
254 #[error("refresh OAuth credentials for `{0}` timed out")]
255 RefreshTimeout(String),
256 #[error("OAuth provider `{provider_id}` returned invalid credentials: {message}")]
257 InvalidRefresh {
258 provider_id: String,
259 message: &'static str,
260 },
261 #[error("persist OAuth credentials for `{provider_id}`: {message}")]
262 Persist {
263 provider_id: String,
264 message: String,
265 },
266 #[error("OAuth provider `{0}` changed while credentials were refreshing")]
267 Changed(String),
268 #[error("OAuth credential refresh task failed for `{provider_id}`: {message}")]
269 Task {
270 provider_id: String,
271 message: String,
272 },
273 #[error(
274 "OAuth provider `{0}` was quarantined after an internal credential failure; remove it and sign in again or restart atman"
275 )]
276 Quarantined(String),
277 #[error("system clock is before the Unix epoch")]
278 Clock,
279}
280
281#[derive(Clone)]
282pub(crate) struct OAuthCredentialLease {
283 provider_id: String,
284 expected_kind: ProviderKind,
285 hub: ConfigHub,
286 refresher: Arc<RefreshFn>,
287 gate: Arc<CredentialGate>,
288}
289
290impl OAuthCredentialLease {
291 pub(crate) fn new<P: OAuthProvider>(provider_id: impl Into<String>, hub: ConfigHub) -> Self {
292 Self::with_refresher(provider_id, P::KIND.clone(), hub, |refresh_token| {
293 P::refresh_token(&refresh_token)
294 })
295 }
296
297 pub(crate) fn with_refresher(
298 provider_id: impl Into<String>,
299 expected_kind: ProviderKind,
300 hub: ConfigHub,
301 refresher: impl Fn(String) -> RefreshFuture + Send + Sync + 'static,
302 ) -> Self {
303 let provider_id = provider_id.into();
304 let gate = credential_gate(&hub, &provider_id);
305 Self {
306 provider_id,
307 expected_kind,
308 hub,
309 refresher: Arc::new(refresher),
310 gate,
311 }
312 }
313
314 pub(crate) async fn acquire(
315 &self,
316 ) -> std::result::Result<OAuthCredential, OAuthCredentialError> {
317 match crate::panic_capture::future(self.acquire_inner()).await {
318 Ok(result) => result,
319 Err(_) => Err(self.gate.quarantine()),
320 }
321 }
322
323 async fn acquire_inner(&self) -> std::result::Result<OAuthCredential, OAuthCredentialError> {
324 self.gate.ensure_not_quarantined()?;
325 let worker = self.refresh_worker();
326 let (stored, snapshot) = worker.load_state_async().await?;
327 self.discard_stale_pending(&snapshot);
328 self.gate.ensure_not_quarantined()?;
329 if !credentials_need_refresh(&stored)? {
330 self.gate.ensure_not_quarantined()?;
331 return Ok(credentials_from_provider(stored));
332 }
333
334 let result = match self.shared_refresh()?.await {
335 SharedRefreshOutcome::Complete(outcome) => outcome.result,
336 SharedRefreshOutcome::Quarantined(error) => Err(error),
337 };
338 self.gate.ensure_not_quarantined()?;
339 result
340 }
341
342 fn shared_refresh(&self) -> std::result::Result<SharedRefreshFuture, OAuthCredentialError> {
343 let gate = self.gate.clone();
344 let mut state = gate
345 .state
346 .lock()
347 .unwrap_or_else(std::sync::PoisonError::into_inner);
348 if let Some(error) = &state.quarantine {
349 return Err(error.clone());
350 }
351 if state.pending_retry_active {
352 return Err(state
353 .failure
354 .as_ref()
355 .map(|failure| failure.error.clone())
356 .unwrap_or_else(|| OAuthCredentialError::Changed(self.provider_id.clone())));
357 }
358 if let Some(in_flight) = &state.in_flight {
359 return Ok(in_flight.clone());
360 }
361 if let Some(failure) = &state.failure
362 && Instant::now() < failure.until
363 {
364 return Err(failure.error.clone());
365 }
366 state.failure = None;
367
368 let pending = state.pending.clone();
369 let worker = self.refresh_worker();
370 let in_flight = shared_refresh_future(worker.clone(), pending);
371 state.in_flight = Some(in_flight.clone());
372 drop(state);
373
374 let driver = in_flight.clone();
377 std::mem::drop(tokio::spawn(drive_shared_refresh(gate, worker, driver)));
378
379 Ok(in_flight)
380 }
381
382 fn discard_stale_pending(&self, snapshot: &crate::auth_store::AuthProviderCredentialSnapshot) {
383 let mut state = self
384 .gate
385 .state
386 .lock()
387 .unwrap_or_else(std::sync::PoisonError::into_inner);
388 if state
389 .pending
390 .as_ref()
391 .is_some_and(|pending| &pending.snapshot != snapshot)
392 {
393 state.pending = None;
394 state.failure = None;
395 }
396 }
397
398 fn refresh_worker(&self) -> OAuthCredentialRefresh {
399 OAuthCredentialRefresh {
400 provider_id: self.provider_id.clone(),
401 expected_kind: self.expected_kind.clone(),
402 hub: self.hub.clone(),
403 refresher: self.refresher.clone(),
404 gate: self.gate.clone(),
405 }
406 }
407}
408
409fn shared_refresh_future(
410 worker: OAuthCredentialRefresh,
411 pending: Option<PendingCredentialCommit>,
412) -> SharedRefreshFuture {
413 let gate = worker.gate.clone();
414 crate::panic_capture::future(async move { worker.refresh_serialized(pending).await })
415 .map(move |result| match result {
416 Ok(result) => match &result.result {
417 Err(error @ OAuthCredentialError::Quarantined(_)) => {
418 SharedRefreshOutcome::Quarantined(error.clone())
419 }
420 _ => SharedRefreshOutcome::Complete(result),
421 },
422 Err(_) => SharedRefreshOutcome::Quarantined(gate.quarantine()),
423 })
424 .boxed()
425 .shared()
426}
427
428async fn drive_shared_refresh(
429 gate: Arc<CredentialGate>,
430 worker: OAuthCredentialRefresh,
431 in_flight: SharedRefreshFuture,
432) {
433 let mut outcome = match in_flight.await {
434 SharedRefreshOutcome::Complete(outcome) => outcome,
435 SharedRefreshOutcome::Quarantined(_) => return,
436 };
437 loop {
438 let retry_pending = outcome.pending.is_some();
439 let failure = outcome.result.as_ref().err().cloned();
440 {
441 let mut state = gate
442 .state
443 .lock()
444 .unwrap_or_else(std::sync::PoisonError::into_inner);
445 if state.quarantine.is_some() {
446 state.in_flight = None;
447 state.pending_retry_active = false;
448 state.pending = None;
449 state.failure = None;
450 return;
451 }
452 state.in_flight = None;
453 state.pending_retry_active = false;
454 state.pending = outcome.pending;
455 state.failure = failure.map(|error| CachedRefreshFailure {
456 until: Instant::now() + REFRESH_FAILURE_COOLDOWN,
457 error,
458 });
459 }
460 if !retry_pending {
461 return;
462 }
463
464 tokio::time::sleep(REFRESH_FAILURE_COOLDOWN).await;
465 let pending = {
466 let mut state = gate
467 .state
468 .lock()
469 .unwrap_or_else(std::sync::PoisonError::into_inner);
470 if state.in_flight.is_some() {
471 return;
472 }
473 let Some(pending) = state.pending.clone() else {
474 return;
475 };
476 state.pending_retry_active = true;
477 pending
478 };
479 outcome = match shared_refresh_future(worker.clone(), Some(pending)).await {
480 SharedRefreshOutcome::Complete(outcome) => outcome,
481 SharedRefreshOutcome::Quarantined(_) => return,
482 };
483 }
484}
485
486#[derive(Clone)]
487struct OAuthCredentialRefresh {
488 provider_id: String,
489 expected_kind: ProviderKind,
490 hub: ConfigHub,
491 refresher: Arc<RefreshFn>,
492 gate: Arc<CredentialGate>,
493}
494
495enum BlockingOutcome<T> {
496 Complete(T),
497 Quarantined(OAuthCredentialError),
498}
499
500impl OAuthCredentialRefresh {
501 async fn run_blocking<T: Send + 'static>(
502 &self,
503 operation: impl FnOnce() -> T + Send + 'static,
504 ) -> std::result::Result<T, OAuthCredentialError> {
505 self.gate.ensure_not_quarantined()?;
506 let gate = self.gate.clone();
507 let outcome = tokio::task::spawn_blocking(move || {
508 if let Err(error) = gate.ensure_not_quarantined() {
509 return BlockingOutcome::Quarantined(error);
510 }
511 match crate::panic_capture::blocking(operation) {
512 Ok(value) => BlockingOutcome::Complete(value),
513 Err(_) => BlockingOutcome::Quarantined(gate.quarantine()),
514 }
515 })
516 .await;
517 let value = match outcome {
518 Ok(BlockingOutcome::Complete(value)) => value,
519 Ok(BlockingOutcome::Quarantined(error)) => return Err(error),
520 Err(error) if error.is_panic() => {
521 return Err(self.gate.quarantine());
522 }
523 Err(error) => {
524 return Err(OAuthCredentialError::Task {
525 provider_id: self.provider_id.clone(),
526 message: error.to_string(),
527 });
528 }
529 };
530 self.gate.ensure_not_quarantined()?;
531 Ok(value)
532 }
533
534 fn load_state(
535 &self,
536 ) -> std::result::Result<
537 (
538 StoredProvider,
539 crate::auth_store::AuthProviderCredentialSnapshot,
540 ),
541 OAuthCredentialError,
542 > {
543 let Some(state) = self
544 .hub
545 .load_or_create_auth_provider_credential_state(&self.provider_id)
546 .map_err(|error| OAuthCredentialError::Load {
547 provider_id: self.provider_id.clone(),
548 message: error.to_string(),
549 })?
550 else {
551 return Err(OAuthCredentialError::Missing(self.provider_id.clone()));
552 };
553 if !state.0.enabled {
554 return Err(OAuthCredentialError::Disabled(self.provider_id.clone()));
555 }
556 if state.0.kind != self.expected_kind {
557 return Err(OAuthCredentialError::KindChanged(self.provider_id.clone()));
558 }
559 Ok(state)
560 }
561
562 async fn load_state_async(
563 &self,
564 ) -> std::result::Result<
565 (
566 StoredProvider,
567 crate::auth_store::AuthProviderCredentialSnapshot,
568 ),
569 OAuthCredentialError,
570 > {
571 let worker = self.clone();
572 self.run_blocking(move || worker.load_state()).await?
573 }
574
575 async fn persist_pending(
576 &self,
577 pending: &PendingCredentialCommit,
578 ) -> std::result::Result<AuthCredentialCommit, OAuthCredentialError> {
579 let hub = self.hub.clone();
580 let provider_id = self.provider_id.clone();
581 let snapshot = pending.snapshot.clone();
582 let update = pending.update();
583 self.run_blocking(move || {
584 hub.update_auth_tokens_if_current(&provider_id, &snapshot, update)
585 })
586 .await?
587 .map_err(|error| OAuthCredentialError::Persist {
588 provider_id: self.provider_id.clone(),
589 message: error.to_string(),
590 })
591 }
592
593 async fn acquire_refresh_file_lock(
594 &self,
595 ) -> std::result::Result<RefreshFileLock, OAuthCredentialError> {
596 let path = refresh_lock_path(&self.hub, &self.provider_id);
597 let provider_id = self.provider_id.clone();
598 self.run_blocking(move || open_refresh_file_lock(&path))
599 .await?
600 .map(RefreshFileLock)
601 .map_err(|error| OAuthCredentialError::Lock {
602 provider_id,
603 message: error.to_string(),
604 })
605 }
606
607 async fn commit_pending(&self, pending: PendingCredentialCommit) -> RefreshFlightResult {
608 let commit = match self.persist_pending(&pending).await {
609 Ok(commit) => commit,
610 Err(error) => return RefreshFlightResult::pending(error, pending),
611 };
612 match commit {
613 AuthCredentialCommit::Updated { provider, .. } => {
614 if provider.enabled {
615 RefreshFlightResult::complete(Ok(credentials_from_provider(provider)))
616 } else {
617 RefreshFlightResult::complete(Err(OAuthCredentialError::Disabled(
618 self.provider_id.clone(),
619 )))
620 }
621 }
622 AuthCredentialCommit::Missing => RefreshFlightResult::complete(Err(
623 OAuthCredentialError::Missing(self.provider_id.clone()),
624 )),
625 AuthCredentialCommit::Changed => {
626 let (provider, _) = match self.load_state_async().await {
629 Ok(state) => state,
630 Err(error) => return RefreshFlightResult::complete(Err(error)),
631 };
632 match credentials_need_refresh(&provider) {
633 Ok(false) => {
634 RefreshFlightResult::complete(Ok(credentials_from_provider(provider)))
635 }
636 Ok(true) => RefreshFlightResult::complete(Err(OAuthCredentialError::Changed(
637 self.provider_id.clone(),
638 ))),
639 Err(error) => RefreshFlightResult::complete(Err(error)),
640 }
641 }
642 }
643 }
644
645 async fn refresh_serialized(
646 &self,
647 pending: Option<PendingCredentialCommit>,
648 ) -> RefreshFlightResult {
649 if let Some(pending) = pending {
650 return self.commit_pending(pending).await;
653 }
654
655 let refresh_lock = match self.acquire_refresh_file_lock().await {
659 Ok(guard) => Arc::new(guard),
660 Err(error) => {
661 return RefreshFlightResult::complete(Err(error));
662 }
663 };
664
665 let (stored, snapshot) = match self.load_state_async().await {
666 Ok(state) => state,
667 Err(error) => return RefreshFlightResult::complete(Err(error)),
668 };
669 match credentials_need_refresh(&stored) {
670 Ok(false) => {
671 return RefreshFlightResult::complete(Ok(credentials_from_provider(stored)));
672 }
673 Ok(true) => {}
674 Err(error) => return RefreshFlightResult::complete(Err(error)),
675 }
676 let Some(refresh_token) = stored
677 .refresh_token
678 .clone()
679 .filter(|token| !token.trim().is_empty())
680 else {
681 return RefreshFlightResult::complete(Err(
682 OAuthCredentialError::ReauthenticationRequired(self.provider_id.clone()),
683 ));
684 };
685 let tokens =
686 match tokio::time::timeout(REFRESH_TIMEOUT, (self.refresher)(refresh_token)).await {
687 Ok(Ok(tokens)) => tokens,
688 Ok(Err(error)) => {
689 return RefreshFlightResult::complete(Err(OAuthCredentialError::Refresh {
690 provider_id: self.provider_id.clone(),
691 message: error.to_string(),
692 }));
693 }
694 Err(_) => {
695 return RefreshFlightResult::complete(Err(
696 OAuthCredentialError::RefreshTimeout(self.provider_id.clone()),
697 ));
698 }
699 };
700 if let Err(error) = validate_refreshed_tokens(&self.provider_id, &tokens) {
701 return RefreshFlightResult::complete(Err(error));
702 }
703 let TokenResult {
704 access_token,
705 refresh_token,
706 expires_at,
707 account,
708 } = tokens;
709 self.commit_pending(PendingCredentialCommit {
710 snapshot,
711 access_token,
712 refresh_token: refresh_token.filter(|token| !token.trim().is_empty()),
713 expires_at,
714 account: account.filter(|account| !account.trim().is_empty()),
715 _refresh_lock: refresh_lock,
716 })
717 .await
718 }
719}
720
721struct RefreshFileLock(std::fs::File);
722
723impl Drop for RefreshFileLock {
724 fn drop(&mut self) {
725 let _ = fs2::FileExt::unlock(&self.0);
726 }
727}
728
729fn credential_gate(hub: &ConfigHub, provider_id: &str) -> Arc<CredentialGate> {
730 let key = CredentialGateKey {
731 auth_path: normalized_auth_path(hub),
732 provider_id: provider_id.to_string(),
733 };
734 let mut gates = CREDENTIAL_GATES
735 .lock()
736 .unwrap_or_else(std::sync::PoisonError::into_inner);
737 gates.retain(|_, entry| entry.quarantined.is_some() || entry.weak.strong_count() > 0);
738 if let Some(gate) = gates
739 .get(&key)
740 .and_then(|entry| entry.quarantined.clone().or_else(|| entry.weak.upgrade()))
741 {
742 return gate;
743 }
744 let gate = Arc::new(CredentialGate {
745 key: key.clone(),
746 state: Mutex::new(CredentialGateState::default()),
747 });
748 gates.insert(
749 key,
750 CredentialGateEntry {
751 weak: Arc::downgrade(&gate),
752 quarantined: None,
753 },
754 );
755 gate
756}
757
758fn normalized_auth_path(hub: &ConfigHub) -> PathBuf {
759 let auth_path =
760 std::path::absolute(hub.auth_path()).unwrap_or_else(|_| hub.auth_path().to_path_buf());
761 let Some(parent) = auth_path.parent() else {
762 return auth_path;
763 };
764 match std::fs::canonicalize(parent) {
765 Ok(parent) => parent.join(
766 auth_path
767 .file_name()
768 .unwrap_or_else(|| std::ffi::OsStr::new("auth.json")),
769 ),
770 Err(_) => auth_path,
771 }
772}
773
774fn refresh_lock_path(hub: &ConfigHub, provider_id: &str) -> PathBuf {
775 let auth_path = normalized_auth_path(hub);
776 let mut hasher = Sha256::new();
777 hasher.update(auth_path.as_os_str().as_encoded_bytes());
778 hasher.update([0]);
779 hasher.update(provider_id.as_bytes());
780 let digest = hasher.finalize();
781 let suffix = URL_SAFE_NO_PAD.encode(&digest[..16]);
782 hub.auth_path()
783 .parent()
784 .unwrap_or_else(|| std::path::Path::new("."))
785 .join(format!(".oauth-refresh-{suffix}.lock"))
786}
787
788fn open_refresh_file_lock(path: &std::path::Path) -> std::io::Result<std::fs::File> {
789 open_refresh_file_lock_with_timeout(path, REFRESH_LOCK_TIMEOUT)
790}
791
792fn open_refresh_file_lock_with_timeout(
793 path: &std::path::Path,
794 timeout: std::time::Duration,
795) -> std::io::Result<std::fs::File> {
796 if let Some(parent) = path.parent() {
797 std::fs::create_dir_all(parent)?;
798 }
799 let mut options = std::fs::OpenOptions::new();
800 options.read(true).write(true).create(true).truncate(false);
801 #[cfg(unix)]
802 {
803 use std::os::unix::fs::OpenOptionsExt;
804 options.mode(0o600);
805 }
806 let file = options.open(path)?;
807 #[cfg(unix)]
808 {
809 use std::os::unix::fs::PermissionsExt;
810 std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
811 }
812 let deadline = std::time::Instant::now() + timeout;
813 loop {
814 match fs2::FileExt::try_lock_exclusive(&file) {
815 Ok(()) => break,
816 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
817 if std::time::Instant::now() >= deadline {
818 return Err(std::io::Error::new(
819 std::io::ErrorKind::TimedOut,
820 "OAuth refresh lock timed out",
821 ));
822 }
823 std::thread::sleep(std::time::Duration::from_millis(25));
824 }
825 Err(error) => return Err(error),
826 }
827 }
828 Ok(file)
829}
830
831fn credentials_need_refresh(
832 provider: &StoredProvider,
833) -> std::result::Result<bool, OAuthCredentialError> {
834 let now = unix_timestamp()?;
835 if provider.expires_at > now.saturating_add(REFRESH_WINDOW_SECONDS) {
836 return Ok(false);
837 }
838 if provider
839 .refresh_token
840 .as_deref()
841 .is_some_and(|token| !token.trim().is_empty())
842 {
843 return Ok(true);
844 }
845 if provider.expires_at > now {
846 return Ok(false);
847 }
848 Err(OAuthCredentialError::ReauthenticationRequired(
849 provider.id.clone(),
850 ))
851}
852
853fn validate_refreshed_tokens(
854 provider_id: &str,
855 tokens: &TokenResult,
856) -> std::result::Result<(), OAuthCredentialError> {
857 if tokens.access_token.trim().is_empty() {
858 return Err(OAuthCredentialError::InvalidRefresh {
859 provider_id: provider_id.to_string(),
860 message: "access token is empty",
861 });
862 }
863 if tokens.expires_at <= unix_timestamp()? {
864 return Err(OAuthCredentialError::InvalidRefresh {
865 provider_id: provider_id.to_string(),
866 message: "access token is already expired",
867 });
868 }
869 Ok(())
870}
871
872fn unix_timestamp() -> std::result::Result<i64, OAuthCredentialError> {
873 Ok(SystemTime::now()
874 .duration_since(UNIX_EPOCH)
875 .map_err(|_| OAuthCredentialError::Clock)?
876 .as_secs() as i64)
877}
878
879fn credentials_from_provider(provider: StoredProvider) -> OAuthCredential {
880 OAuthCredential {
881 access_token: provider.access_token,
882 display_account: provider.account,
883 }
884}
885
886pub trait OAuthProvider: Provider {
887 const KIND: ProviderKind = ProviderKind::Custom;
888
889 fn authorize_url() -> (String, Pkce, String);
890 fn exchange_code(
891 code: &str,
892 verifier: &str,
893 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<TokenResult>> + Send>>;
894 fn refresh_token(
895 token: &str,
896 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<TokenResult>> + Send>>;
897 fn from_stored(stored: &StoredProvider) -> Self;
898 fn from_managed_stored(_stored: &StoredProvider, _hub: ConfigHub) -> Option<Self>
899 where
900 Self: Sized,
901 {
902 None
903 }
904}
905
906pub fn generate_state() -> String {
907 let mut bytes = [0u8; 16];
908 rand::thread_rng().fill_bytes(&mut bytes);
909 bytes.iter().map(|b| format!("{b:02x}")).collect()
910}
911
912pub fn parse_jwt_exp(token: &str) -> Option<i64> {
913 jwt_payload(token)?.get("exp")?.as_i64()
914}
915
916pub fn extract_account_from_id_token(id_token: &str) -> Option<String> {
917 jwt_payload(id_token)?
918 .get("email")
919 .and_then(|e| e.as_str())
920 .map(|s| s.to_string())
921}
922
923pub fn extract_chatgpt_account_id(access_token: &str) -> Option<String> {
924 let payload = jwt_payload(access_token)?;
925 payload
926 .get("chatgpt_account_id")
927 .and_then(serde_json::Value::as_str)
928 .or_else(|| {
929 payload
930 .get("https://api.openai.com/auth")
931 .and_then(|auth| auth.get("chatgpt_account_id"))
932 .and_then(serde_json::Value::as_str)
933 })
934 .map(str::to_owned)
935}
936
937fn jwt_payload(token: &str) -> Option<serde_json::Value> {
938 let mut parts = token.split('.');
939 let _header = parts.next()?;
940 let payload = parts.next()?;
941 let _signature = parts.next()?;
942 if parts.next().is_some() {
943 return None;
944 }
945 let payload = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
946 serde_json::from_slice(&payload).ok()
947}
948
949pub async fn create_oauth_provider<P: OAuthProvider>(
952 stored: &StoredProvider,
953) -> Result<(Arc<P>, Vec<DiscoveredModel>)> {
954 let provider = create_oauth_provider_from_snapshot_impl::<P>(stored).await?;
955 let models = provider.discover_models().await;
956 Ok((provider, models))
957}
958
959pub async fn create_oauth_provider_with_hub<P: OAuthProvider>(
961 stored: &StoredProvider,
962 hub: ConfigHub,
963) -> Result<(Arc<P>, Vec<DiscoveredModel>)> {
964 let gate = credential_gate(&hub, &stored.id);
965 let provider = create_oauth_provider_impl::<P>(stored, hub).await?;
966 let models = provider.discover_models().await;
967 gate.ensure_not_quarantined()?;
968 Ok((provider, models))
969}
970
971pub async fn create_oauth_provider_with_details<P: OAuthProvider>(
974 stored: &StoredProvider,
975) -> Result<(Arc<P>, Vec<DiscoveredModelDetails>)> {
976 let provider = create_oauth_provider_from_snapshot_impl::<P>(stored).await?;
977 let models = provider
978 .try_discover_models()
979 .await
980 .map_err(|error| anyhow::anyhow!(error))?;
981 Ok((provider, models))
982}
983
984pub async fn create_oauth_provider_with_details_and_hub<P: OAuthProvider>(
986 stored: &StoredProvider,
987 hub: ConfigHub,
988) -> Result<(Arc<P>, Vec<DiscoveredModelDetails>)> {
989 let gate = credential_gate(&hub, &stored.id);
990 let provider = create_oauth_provider_impl::<P>(stored, hub).await?;
991 let models = provider
992 .try_discover_models()
993 .await
994 .map_err(|error| anyhow::anyhow!(error))?;
995 gate.ensure_not_quarantined()?;
996 Ok((provider, models))
997}
998
999pub async fn create_oauth_provider_no_discover<P: OAuthProvider>(
1002 stored: &StoredProvider,
1003) -> Result<Arc<P>> {
1004 create_oauth_provider_from_snapshot_impl::<P>(stored).await
1005}
1006
1007pub async fn create_oauth_provider_no_discover_with_hub<P: OAuthProvider>(
1009 stored: &StoredProvider,
1010 hub: ConfigHub,
1011) -> Result<Arc<P>> {
1012 create_oauth_provider_impl::<P>(stored, hub).await
1013}
1014
1015pub fn create_managed_oauth_provider_from_stored<P: OAuthProvider>(
1021 stored: &StoredProvider,
1022 hub: ConfigHub,
1023) -> Result<Arc<P>> {
1024 ensure_provider_kind::<P>(stored)?;
1025 let gate = credential_gate(&hub, &stored.id);
1026 gate.ensure_not_quarantined()?;
1027 let provider = P::from_managed_stored(stored, hub).ok_or_else(|| {
1028 anyhow::anyhow!(
1029 "provider `{}` does not support managed OAuth credentials",
1030 stored.id
1031 )
1032 })?;
1033 gate.ensure_not_quarantined()?;
1034 Ok(Arc::new(provider))
1035}
1036
1037pub fn create_supported_managed_oauth_provider(
1042 stored: &StoredProvider,
1043 hub: ConfigHub,
1044) -> Result<Arc<dyn Provider>> {
1045 match stored.kind {
1046 ProviderKind::Codex => Ok(create_managed_oauth_provider_from_stored::<
1047 crate::providers::codex::CodexProvider,
1048 >(stored, hub)?),
1049 ref kind => anyhow::bail!("managed OAuth provider kind `{kind:?}` is not supported"),
1050 }
1051}
1052
1053async fn create_oauth_provider_impl<P: OAuthProvider>(
1054 stored: &StoredProvider,
1055 hub: ConfigHub,
1056) -> Result<Arc<P>> {
1057 ensure_provider_kind::<P>(stored)?;
1058 let gate = credential_gate(&hub, &stored.id);
1059 gate.ensure_not_quarantined()?;
1060 let validator = OAuthCredentialRefresh {
1061 provider_id: stored.id.clone(),
1062 expected_kind: P::KIND.clone(),
1063 hub: hub.clone(),
1064 refresher: Arc::new(|refresh_token| P::refresh_token(&refresh_token)),
1065 gate,
1066 };
1067 let (authoritative, _) = validator.load_state_async().await?;
1068 validator.gate.ensure_not_quarantined()?;
1069 let provider = Arc::new(P::from_managed_stored(&authoritative, hub).ok_or_else(|| {
1070 anyhow::anyhow!(
1071 "provider `{}` does not support managed OAuth credentials",
1072 stored.id
1073 )
1074 })?);
1075 validator.gate.ensure_not_quarantined()?;
1076 Ok(provider)
1077}
1078
1079async fn create_oauth_provider_from_snapshot_impl<P: OAuthProvider>(
1080 stored: &StoredProvider,
1081) -> Result<Arc<P>> {
1082 ensure_provider_kind::<P>(stored)?;
1083 if stored.expires_at <= unix_timestamp()? {
1084 return Err(OAuthCredentialError::SnapshotExpired(stored.id.clone()).into());
1085 }
1086 Ok(Arc::new(P::from_stored(stored)))
1087}
1088
1089fn ensure_provider_kind<P: OAuthProvider>(stored: &StoredProvider) -> Result<()> {
1090 if stored.kind != P::KIND {
1091 return Err(OAuthCredentialError::KindChanged(stored.id.clone()).into());
1092 }
1093 Ok(())
1094}
1095
1096pub fn callback_page(ok: bool, title: &str, message: &str) -> String {
1097 let icon = if ok { "✓" } else { "✗" };
1098 let color = if ok { "#0078a0" } else { "#c0392b" };
1099 let acetate = if ok {
1100 "rgba(0,120,160,0.10)"
1101 } else {
1102 "rgba(192,57,43,0.10)"
1103 };
1104 format!(
1105 r#"<!DOCTYPE html>
1106<html lang="zh">
1107<head>
1108<meta charset="utf-8">
1109<meta name="viewport" content="width=device-width,initial-scale=1">
1110<title>{title} — atman</title>
1111<style>
1112 * {{ margin:0; padding:0; box-sizing:border-box; }}
1113 body {{
1114 font-family: "JetBrains Mono","Fira Code",Menlo,Consolas,monospace;
1115 background: linear-gradient(180deg,#f0f0f0 0%,#e8e8ec 100%);
1116 min-height: 100vh; display:flex; align-items:center; justify-content:center;
1117 }}
1118 .card {{
1119 background: #fff; border-radius: 12px; padding: 36px 44px;
1120 box-shadow: 0 2px 8px rgba(0,0,0,.06);
1121 text-align: center; max-width: 520px;
1122 border-top: 3px solid {color};
1123 }}
1124 .logo {{ margin-bottom: 20px; }}
1125 .logo pre {{
1126 font-size: 6.5px; line-height: 1.15; color: #0078a0;
1127 font-family: "JetBrains Mono","Fira Code",Menlo,Consolas,monospace;
1128 }}
1129 .icon {{
1130 font-size: 36px; color: {color}; margin-bottom: 16px;
1131 display: inline-block; width: 56px; height: 56px; line-height: 56px;
1132 border-radius: 50%; background: {acetate};
1133 }}
1134 h1 {{ font-size: 18px; font-weight: 600; color: #1e1e1e; margin-bottom: 8px; }}
1135 p {{ font-size: 13px; color: #606060; line-height: 1.6; }}
1136</style>
1137</head>
1138<body>
1139<div class="card">
1140 <div class="logo"><pre>
1141 ⢀⡤⣾⢿⡿⢿⡿⣷⢤⡀
1142 ⢠⢯⢎⠞⡵⠚⠓⢮⠳⡱⡽⡄
1143 ⡟⡏⡏⣀⣳⣀⣀⣞⣀⡰⢹⢻ ████████╗███╗ ███╗ █████╗ ███╗ ██╗
1144 ⢀⣠⡄⣧⣇⡇⠻⠿⠿⠿⠿⠿⢿⡿⣷⣦⣄⡀ ╚══██╔══╝████╗ ████║██╔══██╗████╗ ██║
1145⢀⡴⡫⡪⠕⠹⡼⡜⡄ ⢠⢢⢮⠍⠺⢗⢝⢦⡀ ██║ ██╔████╔██║███████║██╔██╗ ██║
1146⡞⡞⡞ ⠙⣝⢞⢦⡀⢀⡴⡳⣫⠋ ⢳⢳⢳ ██║ ██║╚██╔╝██║██╔══██║██║╚██╗██║
1147⢧⢧⡣⡀ ⠈⣓⡡⣔⣽⡪⢞⠁ ⢀⢜⡼⡼ ██║ ██║ ██║██║ ██║██║ ╚████║
1148⠈⠓⠿⣾⣿⣿⣿⣿⡿⠿⠛⠙⠾⢷⣿⣿⣿⣿⣷⠿⠚⠁ ╚═╝ ╚═╝ ╚═╝╚═╝ ╚═╝╚═╝ ╚═══╝
1149</pre></div>
1150 <div class="icon">{icon}</div>
1151 <h1>{title}</h1>
1152 <p>{message}</p>
1153</div>
1154</body>
1155</html>"#
1156 )
1157}
1158
1159#[cfg(test)]
1160mod tests {
1161 use super::*;
1162 use base64::Engine;
1163 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
1164 use std::sync::atomic::{AtomicUsize, Ordering};
1165 use tokio::sync::Semaphore;
1166
1167 struct SnapshotOAuthProvider {
1168 access_token: String,
1169 display_account: Option<String>,
1170 }
1171
1172 impl Provider for SnapshotOAuthProvider {
1173 fn name(&self) -> &str {
1174 "snapshot-oauth"
1175 }
1176
1177 fn call<'a>(
1178 &'a self,
1179 _req: crate::provider::LlmRequest,
1180 ) -> crate::tool::BoxFut<
1181 'a,
1182 std::result::Result<crate::provider::AssistantMessage, crate::error::RuntimeError>,
1183 > {
1184 Box::pin(async {
1185 Err(crate::error::RuntimeError::ToolFailed(
1186 "unused test provider".into(),
1187 ))
1188 })
1189 }
1190
1191 fn call_streaming(
1192 &self,
1193 _req: crate::provider::LlmRequest,
1194 ) -> crate::event::Observable<crate::provider::AssistantMessage> {
1195 panic!("unused test provider")
1196 }
1197 }
1198
1199 impl OAuthProvider for SnapshotOAuthProvider {
1200 const KIND: ProviderKind = ProviderKind::Codex;
1201
1202 fn authorize_url() -> (String, Pkce, String) {
1203 panic!("unused test provider")
1204 }
1205
1206 fn exchange_code(_code: &str, _verifier: &str) -> RefreshFuture {
1207 Box::pin(async { panic!("unused test provider") })
1208 }
1209
1210 fn refresh_token(_token: &str) -> RefreshFuture {
1211 Box::pin(async { panic!("snapshot provider must not refresh credentials") })
1212 }
1213
1214 fn from_stored(stored: &StoredProvider) -> Self {
1215 Self {
1216 access_token: stored.access_token.clone(),
1217 display_account: stored.account.clone(),
1218 }
1219 }
1220 }
1221
1222 static MANAGED_REFRESH_CALLS: AtomicUsize = AtomicUsize::new(0);
1223
1224 struct ManagedOAuthProvider {
1225 lease: OAuthCredentialLease,
1226 }
1227
1228 impl ManagedOAuthProvider {
1229 async fn acquire(&self) -> std::result::Result<OAuthCredential, OAuthCredentialError> {
1230 self.lease.acquire().await
1231 }
1232 }
1233
1234 impl Provider for ManagedOAuthProvider {
1235 fn name(&self) -> &str {
1236 "managed-oauth"
1237 }
1238
1239 fn call<'a>(
1240 &'a self,
1241 _req: crate::provider::LlmRequest,
1242 ) -> crate::tool::BoxFut<
1243 'a,
1244 std::result::Result<crate::provider::AssistantMessage, crate::error::RuntimeError>,
1245 > {
1246 Box::pin(async {
1247 Err(crate::error::RuntimeError::ToolFailed(
1248 "unused test provider".into(),
1249 ))
1250 })
1251 }
1252
1253 fn call_streaming(
1254 &self,
1255 _req: crate::provider::LlmRequest,
1256 ) -> crate::event::Observable<crate::provider::AssistantMessage> {
1257 panic!("unused test provider")
1258 }
1259 }
1260
1261 impl OAuthProvider for ManagedOAuthProvider {
1262 const KIND: ProviderKind = ProviderKind::Codex;
1263
1264 fn authorize_url() -> (String, Pkce, String) {
1265 panic!("unused test provider")
1266 }
1267
1268 fn exchange_code(_code: &str, _verifier: &str) -> RefreshFuture {
1269 Box::pin(async { panic!("unused test provider") })
1270 }
1271
1272 fn refresh_token(token: &str) -> RefreshFuture {
1273 assert_eq!(token, "refresh-v1");
1274 MANAGED_REFRESH_CALLS.fetch_add(1, Ordering::SeqCst);
1275 Box::pin(async { Ok(refreshed_tokens()) })
1276 }
1277
1278 fn from_stored(_stored: &StoredProvider) -> Self {
1279 panic!("managed test provider requires a config hub")
1280 }
1281
1282 fn from_managed_stored(stored: &StoredProvider, hub: ConfigHub) -> Option<Self> {
1283 Some(Self {
1284 lease: OAuthCredentialLease::new::<Self>(&stored.id, hub),
1285 })
1286 }
1287 }
1288
1289 struct QuarantiningProvider {
1290 provider_id: String,
1291 hub: ConfigHub,
1292 }
1293
1294 impl Provider for QuarantiningProvider {
1295 fn name(&self) -> &str {
1296 &self.provider_id
1297 }
1298
1299 fn call<'a>(
1300 &'a self,
1301 _req: crate::provider::LlmRequest,
1302 ) -> crate::tool::BoxFut<
1303 'a,
1304 std::result::Result<crate::provider::AssistantMessage, crate::error::RuntimeError>,
1305 > {
1306 Box::pin(async {
1307 Err(crate::error::RuntimeError::ToolFailed(
1308 "unused test provider".into(),
1309 ))
1310 })
1311 }
1312
1313 fn call_streaming(
1314 &self,
1315 _req: crate::provider::LlmRequest,
1316 ) -> crate::event::Observable<crate::provider::AssistantMessage> {
1317 panic!("unused test provider")
1318 }
1319
1320 fn discover_models(&self) -> crate::tool::BoxFut<'static, Vec<DiscoveredModel>> {
1321 let provider_id = self.provider_id.clone();
1322 let hub = self.hub.clone();
1323 Box::pin(async move {
1324 credential_gate(&hub, &provider_id).quarantine();
1325 Vec::new()
1326 })
1327 }
1328
1329 fn try_discover_models(
1330 &self,
1331 ) -> crate::tool::BoxFut<
1332 'static,
1333 std::result::Result<Vec<DiscoveredModelDetails>, crate::provider::ModelDiscoveryError>,
1334 > {
1335 let provider_id = self.provider_id.clone();
1336 let hub = self.hub.clone();
1337 Box::pin(async move {
1338 credential_gate(&hub, &provider_id).quarantine();
1339 Ok(Vec::new())
1340 })
1341 }
1342 }
1343
1344 impl OAuthProvider for QuarantiningProvider {
1345 const KIND: ProviderKind = ProviderKind::Codex;
1346
1347 fn authorize_url() -> (String, Pkce, String) {
1348 panic!("unused test provider")
1349 }
1350
1351 fn exchange_code(_code: &str, _verifier: &str) -> RefreshFuture {
1352 Box::pin(async { panic!("unused test provider") })
1353 }
1354
1355 fn refresh_token(_token: &str) -> RefreshFuture {
1356 Box::pin(async { panic!("unused test provider") })
1357 }
1358
1359 fn from_stored(_stored: &StoredProvider) -> Self {
1360 panic!("unused test provider")
1361 }
1362
1363 fn from_managed_stored(stored: &StoredProvider, hub: ConfigHub) -> Option<Self> {
1364 if stored.id.ends_with("construction") {
1365 credential_gate(&hub, &stored.id).quarantine();
1366 }
1367 Some(Self {
1368 provider_id: stored.id.clone(),
1369 hub,
1370 })
1371 }
1372 }
1373
1374 fn expired_provider(id: &str) -> StoredProvider {
1375 StoredProvider {
1376 id: id.into(),
1377 name: "OAuth account".into(),
1378 kind: ProviderKind::Codex,
1379 access_token: "access-v1".into(),
1380 refresh_token: Some("refresh-v1".into()),
1381 expires_at: chrono::Utc::now().timestamp() - 1,
1382 account: Some("display-v1@example.test".into()),
1383 enabled: true,
1384 model_cache: None,
1385 }
1386 }
1387
1388 fn hub_with_provider(provider: StoredProvider) -> (tempfile::TempDir, ConfigHub) {
1389 let dir = tempfile::tempdir().unwrap();
1390 let hub = ConfigHub::from_config_dir(dir.path());
1391 hub.add_auth_provider(provider).unwrap();
1392 (dir, hub)
1393 }
1394
1395 fn refreshed_tokens() -> TokenResult {
1396 TokenResult {
1397 access_token: "access-v2".into(),
1398 refresh_token: Some("refresh-v2".into()),
1399 expires_at: chrono::Utc::now().timestamp() + 3_600,
1400 account: Some("display-v2@example.test".into()),
1401 }
1402 }
1403
1404 #[test]
1405 fn oauth_credential_debug_redacts_secrets_and_account() {
1406 let credential = OAuthCredential {
1407 access_token: "secret-access-token".into(),
1408 display_account: Some("person@example.test".into()),
1409 };
1410
1411 let debug = format!("{credential:?}");
1412 assert!(debug.contains("[redacted]"));
1413 assert!(!debug.contains("secret-access-token"));
1414 assert!(!debug.contains("person@example.test"));
1415 }
1416
1417 #[test]
1418 fn pkce_generates_43char_verifier_and_valid_challenge() {
1419 let pkce = Pkce::generate();
1420 assert_eq!(pkce.verifier.len(), 43);
1421 assert!(!pkce.challenge.is_empty());
1422 let mut hasher = Sha256::new();
1423 hasher.update(pkce.verifier.as_bytes());
1424 let digest = hasher.finalize();
1425 let expected = URL_SAFE_NO_PAD.encode(digest);
1426 assert_eq!(pkce.challenge, expected);
1427 }
1428
1429 #[test]
1430 fn generate_state_is_32_hex_chars() {
1431 let state = generate_state();
1432 assert_eq!(state.len(), 32);
1433 assert!(state.chars().all(|c| c.is_ascii_hexdigit()));
1434 }
1435
1436 #[test]
1437 fn parse_jwt_exp_works() {
1438 let payload = URL_SAFE_NO_PAD.encode(r#"{"exp":123456789,"other":"data"}"#.as_bytes());
1439 let token = format!("header.{payload}.sig");
1440 assert_eq!(parse_jwt_exp(&token), Some(123456789));
1441 }
1442
1443 #[test]
1444 fn parse_jwt_exp_returns_none_when_missing() {
1445 assert_eq!(parse_jwt_exp("not.a.jwt"), None);
1446 assert_eq!(parse_jwt_exp(""), None);
1447 }
1448
1449 #[test]
1450 fn extract_account_prefers_email() {
1451 let payload = URL_SAFE_NO_PAD.encode(r#"{"email":"a@b.com","sub":"123"}"#.as_bytes());
1452 let token = format!("h.{payload}.sig");
1453 assert_eq!(
1454 extract_account_from_id_token(&token),
1455 Some("a@b.com".to_string())
1456 );
1457 }
1458
1459 #[test]
1460 fn extract_account_returns_none_when_missing() {
1461 let payload = URL_SAFE_NO_PAD.encode(r#"{"name":"John"}"#.as_bytes());
1462 let token = format!("h.{payload}.sig");
1463 assert_eq!(extract_account_from_id_token(&token), None);
1464 }
1465
1466 #[test]
1467 fn extracts_chatgpt_account_id_from_access_token_claims() {
1468 let nested = URL_SAFE_NO_PAD
1469 .encode(r#"{"https://api.openai.com/auth":{"chatgpt_account_id":"account-nested"}}"#);
1470 let top_level = URL_SAFE_NO_PAD.encode(
1471 r#"{"chatgpt_account_id":"account-top","https://api.openai.com/auth":{"chatgpt_account_id":"account-nested"}}"#,
1472 );
1473
1474 assert_eq!(
1475 extract_chatgpt_account_id(&format!("h.{nested}.s")),
1476 Some("account-nested".into())
1477 );
1478 assert_eq!(
1479 extract_chatgpt_account_id(&format!("h.{top_level}.s")),
1480 Some("account-top".into())
1481 );
1482 }
1483
1484 #[test]
1485 fn refresh_file_lock_serializes_independent_openers() {
1486 let dir = tempfile::tempdir().unwrap();
1487 let path = dir.path().join("refresh.lock");
1488 let first = RefreshFileLock(
1489 open_refresh_file_lock_with_timeout(&path, std::time::Duration::from_millis(50))
1490 .unwrap(),
1491 );
1492
1493 let error =
1494 open_refresh_file_lock_with_timeout(&path, std::time::Duration::from_millis(25))
1495 .unwrap_err();
1496 assert_eq!(error.kind(), std::io::ErrorKind::TimedOut);
1497 drop(first);
1498 let _second = RefreshFileLock(
1499 open_refresh_file_lock_with_timeout(&path, std::time::Duration::from_millis(50))
1500 .unwrap(),
1501 );
1502
1503 #[cfg(unix)]
1504 {
1505 use std::os::unix::fs::PermissionsExt;
1506 assert_eq!(
1507 std::fs::metadata(path).unwrap().permissions().mode() & 0o777,
1508 0o600
1509 );
1510 }
1511 }
1512
1513 #[tokio::test]
1514 async fn concurrent_acquire_refreshes_once_and_persists_rotation() {
1515 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1516 let calls = Arc::new(AtomicUsize::new(0));
1517 let started = Arc::new(Semaphore::new(0));
1518 let proceed = Arc::new(Semaphore::new(0));
1519 let lease = OAuthCredentialLease::with_refresher(
1520 "oauth-account",
1521 ProviderKind::Codex,
1522 hub.clone(),
1523 {
1524 let calls = calls.clone();
1525 let started = started.clone();
1526 let proceed = proceed.clone();
1527 move |refresh_token| {
1528 let calls = calls.clone();
1529 let started = started.clone();
1530 let proceed = proceed.clone();
1531 Box::pin(async move {
1532 assert_eq!(refresh_token, "refresh-v1");
1533 calls.fetch_add(1, Ordering::SeqCst);
1534 started.add_permits(1);
1535 proceed.acquire().await.unwrap().forget();
1536 Ok(refreshed_tokens())
1537 })
1538 }
1539 },
1540 );
1541
1542 let first = tokio::spawn({
1543 let lease = lease.clone();
1544 async move { lease.acquire().await }
1545 });
1546 started.acquire().await.unwrap().forget();
1547 let mut rest = Vec::new();
1548 for _ in 0..7 {
1549 let lease = lease.clone();
1550 rest.push(tokio::spawn(async move { lease.acquire().await }));
1551 }
1552 tokio::task::yield_now().await;
1553 assert_eq!(calls.load(Ordering::SeqCst), 1);
1554 proceed.add_permits(8);
1555
1556 let mut credentials = vec![first.await.unwrap().unwrap()];
1557 for task in rest {
1558 credentials.push(task.await.unwrap().unwrap());
1559 }
1560 assert!(credentials.iter().all(|credential| {
1561 credential.access_token == "access-v2"
1562 && credential.display_account.as_deref() == Some("display-v2@example.test")
1563 }));
1564 assert_eq!(calls.load(Ordering::SeqCst), 1);
1565 let stored = hub.load_auth().unwrap().providers.remove(0);
1566 assert_eq!(stored.access_token, "access-v2");
1567 assert_eq!(stored.refresh_token.as_deref(), Some("refresh-v2"));
1568 assert_eq!(stored.account.as_deref(), Some("display-v2@example.test"));
1569 }
1570
1571 #[tokio::test]
1572 async fn concurrent_refresh_failure_is_shared_and_cooled_down() {
1573 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1574 let calls = Arc::new(AtomicUsize::new(0));
1575 let started = Arc::new(Semaphore::new(0));
1576 let proceed = Arc::new(Semaphore::new(0));
1577 let lease =
1578 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, {
1579 let calls = calls.clone();
1580 let started = started.clone();
1581 let proceed = proceed.clone();
1582 move |_| {
1583 let calls = calls.clone();
1584 let started = started.clone();
1585 let proceed = proceed.clone();
1586 Box::pin(async move {
1587 calls.fetch_add(1, Ordering::SeqCst);
1588 started.add_permits(1);
1589 proceed.acquire().await.unwrap().forget();
1590 anyhow::bail!("refresh service unavailable")
1591 })
1592 }
1593 });
1594
1595 let first = tokio::spawn({
1596 let lease = lease.clone();
1597 async move { lease.acquire().await }
1598 });
1599 started.acquire().await.unwrap().forget();
1600 let mut rest = Vec::new();
1601 for _ in 0..7 {
1602 let lease = lease.clone();
1603 rest.push(tokio::spawn(async move { lease.acquire().await }));
1604 }
1605 tokio::task::yield_now().await;
1606 proceed.add_permits(1);
1607
1608 let expected = first.await.unwrap().unwrap_err().to_string();
1609 assert!(expected.contains("refresh service unavailable"));
1610 for task in rest {
1611 assert_eq!(task.await.unwrap().unwrap_err().to_string(), expected);
1612 }
1613 assert_eq!(calls.load(Ordering::SeqCst), 1);
1614
1615 let cooldown_error = lease.acquire().await.unwrap_err().to_string();
1616 assert_eq!(cooldown_error, expected);
1617 assert_eq!(calls.load(Ordering::SeqCst), 1);
1618 }
1619
1620 #[tokio::test]
1621 async fn concurrent_refresh_panic_quarantines_before_shared_completion() {
1622 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1623 let calls = Arc::new(AtomicUsize::new(0));
1624 let started = Arc::new(Semaphore::new(0));
1625 let proceed = Arc::new(Semaphore::new(0));
1626 let lease =
1627 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, {
1628 let calls = calls.clone();
1629 let started = started.clone();
1630 let proceed = proceed.clone();
1631 move |_| {
1632 let calls = calls.clone();
1633 let started = started.clone();
1634 let proceed = proceed.clone();
1635 Box::pin(async move {
1636 calls.fetch_add(1, Ordering::SeqCst);
1637 started.add_permits(1);
1638 proceed.acquire().await.unwrap().forget();
1639 panic!("refresh panic fixture")
1640 })
1641 }
1642 });
1643
1644 let flight = lease.shared_refresh().unwrap();
1645 let joined = Arc::new(Semaphore::new(0));
1646 let mut waiters = Vec::new();
1647 for _ in 0..8 {
1648 let joined = joined.clone();
1649 let mut flight = Box::pin(flight.clone());
1650 waiters.push(tokio::spawn(async move {
1651 let mut announced = false;
1652 std::future::poll_fn(move |context| {
1653 let result = flight.as_mut().poll(context);
1654 if !announced && result.is_pending() {
1655 announced = true;
1656 joined.add_permits(1);
1657 }
1658 result
1659 })
1660 .await
1661 }));
1662 }
1663 started.acquire().await.unwrap().forget();
1664 joined.acquire_many(8).await.unwrap().forget();
1665 proceed.add_permits(1);
1666
1667 let first = waiters.remove(0);
1668 let first_result = tokio::time::timeout(std::time::Duration::from_secs(2), first)
1669 .await
1670 .unwrap()
1671 .unwrap();
1672 assert!(matches!(
1673 first_result,
1674 SharedRefreshOutcome::Quarantined(OAuthCredentialError::Quarantined(provider))
1675 if provider == "oauth-account"
1676 ));
1677 assert!(matches!(
1678 lease.acquire().await,
1679 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1680 ));
1681 for task in waiters {
1682 assert!(matches!(
1683 tokio::time::timeout(std::time::Duration::from_secs(2), task)
1684 .await
1685 .unwrap()
1686 .unwrap(),
1687 SharedRefreshOutcome::Quarantined(OAuthCredentialError::Quarantined(provider))
1688 if provider == "oauth-account"
1689 ));
1690 }
1691 tokio::time::sleep(REFRESH_FAILURE_COOLDOWN * 2).await;
1692 assert!(matches!(
1693 lease.acquire().await,
1694 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1695 ));
1696 assert_eq!(calls.load(Ordering::SeqCst), 1);
1697 }
1698
1699 #[tokio::test]
1700 async fn synchronous_refresher_panic_quarantines_provider() {
1701 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1702 let lease =
1703 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, |_| {
1704 panic!("synchronous refresher panic fixture")
1705 });
1706
1707 assert!(matches!(
1708 lease.acquire().await,
1709 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1710 ));
1711 }
1712
1713 #[tokio::test]
1714 async fn cancelled_waiter_cannot_drop_refresh_panic_quarantine() {
1715 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1716 let calls = Arc::new(AtomicUsize::new(0));
1717 let started = Arc::new(Semaphore::new(0));
1718 let proceed = Arc::new(Semaphore::new(0));
1719 let lease = OAuthCredentialLease::with_refresher(
1720 "oauth-account",
1721 ProviderKind::Codex,
1722 hub.clone(),
1723 {
1724 let calls = calls.clone();
1725 let started = started.clone();
1726 let proceed = proceed.clone();
1727 move |_| {
1728 let calls = calls.clone();
1729 let started = started.clone();
1730 let proceed = proceed.clone();
1731 Box::pin(async move {
1732 calls.fetch_add(1, Ordering::SeqCst);
1733 started.add_permits(1);
1734 proceed.acquire().await.unwrap().forget();
1735 panic!("detached refresh panic fixture")
1736 })
1737 }
1738 },
1739 );
1740
1741 let waiter = tokio::spawn({
1742 let lease = lease.clone();
1743 async move { lease.acquire().await }
1744 });
1745 started.acquire().await.unwrap().forget();
1746 waiter.abort();
1747 assert!(waiter.await.unwrap_err().is_cancelled());
1748 proceed.add_permits(1);
1749
1750 tokio::time::timeout(std::time::Duration::from_secs(2), async {
1751 loop {
1752 if lease.gate.ensure_not_quarantined().is_err() {
1753 break;
1754 }
1755 tokio::task::yield_now().await;
1756 }
1757 })
1758 .await
1759 .unwrap();
1760 drop(lease);
1761 let rebuilt =
1762 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, |_| {
1763 Box::pin(async { Ok(refreshed_tokens()) })
1764 });
1765 assert!(matches!(
1766 rebuilt.acquire().await,
1767 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1768 ));
1769 assert_eq!(calls.load(Ordering::SeqCst), 1);
1770 }
1771
1772 #[tokio::test]
1773 async fn cancelled_blocking_waiter_cannot_drop_quarantine() {
1774 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1775 let lease = OAuthCredentialLease::with_refresher(
1776 "oauth-account",
1777 ProviderKind::Codex,
1778 hub.clone(),
1779 |_| Box::pin(async { Ok(refreshed_tokens()) }),
1780 );
1781 let worker = lease.refresh_worker();
1782 let gate = lease.gate.clone();
1783 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
1784 let (proceed_tx, proceed_rx) = std::sync::mpsc::channel();
1785 let waiter = tokio::spawn(async move {
1786 let result: std::result::Result<(), OAuthCredentialError> = worker
1787 .run_blocking(move || {
1788 let _ = started_tx.send(());
1789 proceed_rx.recv().unwrap();
1790 panic!("blocking credential panic fixture")
1791 })
1792 .await;
1793 result
1794 });
1795 started_rx.await.unwrap();
1796 waiter.abort();
1797 assert!(waiter.await.unwrap_err().is_cancelled());
1798 proceed_tx.send(()).unwrap();
1799
1800 tokio::time::timeout(std::time::Duration::from_secs(2), async {
1801 loop {
1802 if gate.ensure_not_quarantined().is_err() {
1803 break;
1804 }
1805 tokio::task::yield_now().await;
1806 }
1807 })
1808 .await
1809 .unwrap();
1810 let rebuilt =
1811 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, |_| {
1812 Box::pin(async { Ok(refreshed_tokens()) })
1813 });
1814 assert!(matches!(
1815 rebuilt.acquire().await,
1816 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1817 ));
1818 }
1819
1820 #[tokio::test]
1821 async fn quarantine_survives_fresh_credentials_and_is_scoped_by_path_and_id() {
1822 let (dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1823 let lease = OAuthCredentialLease::with_refresher(
1824 "oauth-account",
1825 ProviderKind::Codex,
1826 hub.clone(),
1827 |_| Box::pin(async { Ok(refreshed_tokens()) }),
1828 );
1829 lease.gate.quarantine();
1830 assert!(
1831 hub.update_auth_tokens(
1832 "oauth-account",
1833 AuthTokenUpdate {
1834 access_token: "fresh-access".into(),
1835 refresh_token: Some("fresh-refresh".into()),
1836 expires_at: chrono::Utc::now().timestamp() + 3_600,
1837 account: None,
1838 },
1839 )
1840 .unwrap()
1841 );
1842 drop(lease);
1843
1844 let reloaded_hub = ConfigHub::from_config_dir(dir.path());
1845 let rebuilt = OAuthCredentialLease::with_refresher(
1846 "oauth-account",
1847 ProviderKind::Codex,
1848 reloaded_hub.clone(),
1849 |_| panic!("quarantined refresher must not run"),
1850 );
1851 assert!(matches!(
1852 rebuilt.acquire().await,
1853 Err(OAuthCredentialError::Quarantined(provider)) if provider == "oauth-account"
1854 ));
1855
1856 let mut peer = expired_provider("peer-account");
1857 peer.expires_at = chrono::Utc::now().timestamp() + 3_600;
1858 reloaded_hub.add_auth_provider(peer).unwrap();
1859 let peer = OAuthCredentialLease::with_refresher(
1860 "peer-account",
1861 ProviderKind::Codex,
1862 reloaded_hub,
1863 |_| panic!("fresh peer credentials must not refresh"),
1864 );
1865 assert_eq!(peer.acquire().await.unwrap().access_token, "access-v1");
1866
1867 let mut independent = expired_provider("oauth-account");
1868 independent.expires_at = chrono::Utc::now().timestamp() + 3_600;
1869 let (_other_dir, other_hub) = hub_with_provider(independent);
1870 let independent = OAuthCredentialLease::with_refresher(
1871 "oauth-account",
1872 ProviderKind::Codex,
1873 other_hub,
1874 |_| panic!("fresh independent credentials must not refresh"),
1875 );
1876 assert_eq!(
1877 independent.acquire().await.unwrap().access_token,
1878 "access-v1"
1879 );
1880 }
1881
1882 #[test]
1883 fn healthy_gate_is_not_pinned_by_the_global_registry() {
1884 let dir = tempfile::tempdir().unwrap();
1885 let hub = ConfigHub::from_config_dir(dir.path());
1886 let gate = credential_gate(&hub, "healthy-account");
1887 let weak = Arc::downgrade(&gate);
1888
1889 drop(gate);
1890 assert!(weak.upgrade().is_none());
1891 }
1892
1893 #[tokio::test]
1894 async fn managed_factories_reject_quarantined_provider() {
1895 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1896 credential_gate(&hub, "oauth-account").quarantine();
1897 let stored = hub.load_auth().unwrap().providers.remove(0);
1898
1899 let sync_error =
1900 create_managed_oauth_provider_from_stored::<ManagedOAuthProvider>(&stored, hub.clone())
1901 .err()
1902 .expect("quarantined sync factory succeeded");
1903 assert!(sync_error.to_string().contains("was quarantined"));
1904
1905 let async_error =
1906 create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(&stored, hub)
1907 .await
1908 .err()
1909 .expect("quarantined async factory succeeded");
1910 assert!(async_error.to_string().contains("was quarantined"));
1911 }
1912
1913 #[tokio::test]
1914 async fn managed_factories_recheck_quarantine_after_construction() {
1915 let (_sync_dir, sync_hub) = hub_with_provider(expired_provider("sync-construction"));
1916 let sync_stored = sync_hub.load_auth().unwrap().providers.remove(0);
1917 let sync_error = create_managed_oauth_provider_from_stored::<QuarantiningProvider>(
1918 &sync_stored,
1919 sync_hub,
1920 )
1921 .err()
1922 .expect("sync factory ignored construction-time quarantine");
1923 assert!(sync_error.to_string().contains("was quarantined"));
1924
1925 let (_async_dir, async_hub) = hub_with_provider(expired_provider("async-construction"));
1926 let async_stored = async_hub.load_auth().unwrap().providers.remove(0);
1927 let async_error = create_oauth_provider_no_discover_with_hub::<QuarantiningProvider>(
1928 &async_stored,
1929 async_hub,
1930 )
1931 .await
1932 .err()
1933 .expect("async factory ignored construction-time quarantine");
1934 assert!(async_error.to_string().contains("was quarantined"));
1935 }
1936
1937 #[tokio::test]
1938 async fn managed_discovery_factories_recheck_quarantine_before_returning() {
1939 let (_legacy_dir, legacy_hub) = hub_with_provider(expired_provider("legacy-discovery"));
1940 let legacy_stored = legacy_hub.load_auth().unwrap().providers.remove(0);
1941 let legacy_error =
1942 create_oauth_provider_with_hub::<QuarantiningProvider>(&legacy_stored, legacy_hub)
1943 .await
1944 .err()
1945 .expect("legacy discovery returned a quarantined provider");
1946 assert!(legacy_error.to_string().contains("was quarantined"));
1947
1948 let (_typed_dir, typed_hub) = hub_with_provider(expired_provider("typed-discovery"));
1949 let typed_stored = typed_hub.load_auth().unwrap().providers.remove(0);
1950 let typed_error = create_oauth_provider_with_details_and_hub::<QuarantiningProvider>(
1951 &typed_stored,
1952 typed_hub,
1953 )
1954 .await
1955 .err()
1956 .expect("typed discovery returned a quarantined provider");
1957 assert!(typed_error.to_string().contains("was quarantined"));
1958 }
1959
1960 #[test]
1961 fn quarantine_discards_retry_state_and_releases_refresh_lock() {
1962 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
1963 let gate = credential_gate(&hub, "oauth-account");
1964 let (_, snapshot) = OAuthCredentialRefresh {
1965 provider_id: "oauth-account".into(),
1966 expected_kind: ProviderKind::Codex,
1967 hub: hub.clone(),
1968 refresher: Arc::new(|_| Box::pin(async { Ok(refreshed_tokens()) })),
1969 gate: gate.clone(),
1970 }
1971 .load_state()
1972 .unwrap();
1973 let refresh_lock_path = refresh_lock_path(&hub, "oauth-account");
1974 let pending = PendingCredentialCommit {
1975 snapshot,
1976 access_token: "pending-access".into(),
1977 refresh_token: Some("pending-refresh".into()),
1978 expires_at: chrono::Utc::now().timestamp() + 3_600,
1979 account: None,
1980 _refresh_lock: Arc::new(RefreshFileLock(
1981 open_refresh_file_lock_with_timeout(
1982 &refresh_lock_path,
1983 std::time::Duration::from_millis(50),
1984 )
1985 .unwrap(),
1986 )),
1987 };
1988 let in_flight = futures::future::ready(SharedRefreshOutcome::Complete(
1989 RefreshFlightResult::complete(Err(OAuthCredentialError::Changed(
1990 "oauth-account".into(),
1991 ))),
1992 ))
1993 .boxed()
1994 .shared();
1995 {
1996 let mut state = gate
1997 .state
1998 .lock()
1999 .unwrap_or_else(std::sync::PoisonError::into_inner);
2000 state.in_flight = Some(in_flight);
2001 state.pending_retry_active = true;
2002 state.failure = Some(CachedRefreshFailure {
2003 until: Instant::now() + REFRESH_FAILURE_COOLDOWN,
2004 error: OAuthCredentialError::Changed("oauth-account".into()),
2005 });
2006 state.pending = Some(pending);
2007 }
2008
2009 gate.quarantine();
2010 let state = gate
2011 .state
2012 .lock()
2013 .unwrap_or_else(std::sync::PoisonError::into_inner);
2014 assert!(state.in_flight.is_none());
2015 assert!(!state.pending_retry_active);
2016 assert!(state.failure.is_none());
2017 assert!(state.pending.is_none());
2018 assert!(matches!(
2019 state.quarantine,
2020 Some(OAuthCredentialError::Quarantined(ref provider))
2021 if provider == "oauth-account"
2022 ));
2023 drop(state);
2024 let _reacquired = RefreshFileLock(
2025 open_refresh_file_lock_with_timeout(
2026 &refresh_lock_path,
2027 std::time::Duration::from_millis(50),
2028 )
2029 .unwrap(),
2030 );
2031 }
2032
2033 #[tokio::test]
2034 async fn refresh_persistence_survives_waiter_cancellation() {
2035 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2036 let started = Arc::new(Semaphore::new(0));
2037 let proceed = Arc::new(Semaphore::new(0));
2038 let lease = OAuthCredentialLease::with_refresher(
2039 "oauth-account",
2040 ProviderKind::Codex,
2041 hub.clone(),
2042 {
2043 let started = started.clone();
2044 let proceed = proceed.clone();
2045 move |_| {
2046 let started = started.clone();
2047 let proceed = proceed.clone();
2048 Box::pin(async move {
2049 started.add_permits(1);
2050 proceed.acquire().await.unwrap().forget();
2051 Ok(refreshed_tokens())
2052 })
2053 }
2054 },
2055 );
2056
2057 let waiter = tokio::spawn(async move { lease.acquire().await });
2058 started.acquire().await.unwrap().forget();
2059 waiter.abort();
2060 proceed.add_permits(1);
2061
2062 tokio::time::timeout(std::time::Duration::from_secs(2), async {
2063 loop {
2064 if hub.load_auth().unwrap().providers[0].access_token == "access-v2" {
2065 break;
2066 }
2067 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
2068 }
2069 })
2070 .await
2071 .unwrap();
2072 }
2073
2074 #[tokio::test]
2075 async fn in_flight_refresh_cannot_resurrect_removed_provider() {
2076 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2077 let started = Arc::new(Semaphore::new(0));
2078 let proceed = Arc::new(Semaphore::new(0));
2079 let lease = OAuthCredentialLease::with_refresher(
2080 "oauth-account",
2081 ProviderKind::Codex,
2082 hub.clone(),
2083 {
2084 let started = started.clone();
2085 let proceed = proceed.clone();
2086 move |_| {
2087 let started = started.clone();
2088 let proceed = proceed.clone();
2089 Box::pin(async move {
2090 started.add_permits(1);
2091 proceed.acquire().await.unwrap().forget();
2092 Ok(refreshed_tokens())
2093 })
2094 }
2095 },
2096 );
2097
2098 let acquire = tokio::spawn(async move { lease.acquire().await });
2099 started.acquire().await.unwrap().forget();
2100 assert!(hub.remove_auth_provider("oauth-account").unwrap());
2101 proceed.add_permits(1);
2102
2103 assert!(matches!(
2104 acquire.await.unwrap(),
2105 Err(OAuthCredentialError::Missing(provider)) if provider == "oauth-account"
2106 ));
2107 assert!(hub.load_auth().unwrap().providers.is_empty());
2108 }
2109
2110 #[tokio::test]
2111 async fn in_flight_refresh_rejects_disabled_provider() {
2112 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2113 let started = Arc::new(Semaphore::new(0));
2114 let proceed = Arc::new(Semaphore::new(0));
2115 let lease = OAuthCredentialLease::with_refresher(
2116 "oauth-account",
2117 ProviderKind::Codex,
2118 hub.clone(),
2119 {
2120 let started = started.clone();
2121 let proceed = proceed.clone();
2122 move |_| {
2123 let started = started.clone();
2124 let proceed = proceed.clone();
2125 Box::pin(async move {
2126 started.add_permits(1);
2127 proceed.acquire().await.unwrap().forget();
2128 Ok(refreshed_tokens())
2129 })
2130 }
2131 },
2132 );
2133
2134 let acquire = tokio::spawn(async move { lease.acquire().await });
2135 started.acquire().await.unwrap().forget();
2136 assert!(
2137 hub.set_auth_provider_enabled("oauth-account", false)
2138 .unwrap()
2139 );
2140 proceed.add_permits(1);
2141
2142 assert!(matches!(
2143 acquire.await.unwrap(),
2144 Err(OAuthCredentialError::Disabled(provider)) if provider == "oauth-account"
2145 ));
2146 let stored = hub.load_auth().unwrap().providers.remove(0);
2147 assert!(!stored.enabled);
2148 assert_eq!(stored.access_token, "access-v2");
2149 assert_eq!(stored.refresh_token.as_deref(), Some("refresh-v2"));
2150 }
2151
2152 #[tokio::test]
2153 async fn in_flight_refresh_preserves_concurrent_catalog_update() {
2154 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2155 let started = Arc::new(Semaphore::new(0));
2156 let proceed = Arc::new(Semaphore::new(0));
2157 let lease = OAuthCredentialLease::with_refresher(
2158 "oauth-account",
2159 ProviderKind::Codex,
2160 hub.clone(),
2161 {
2162 let started = started.clone();
2163 let proceed = proceed.clone();
2164 move |_| {
2165 let started = started.clone();
2166 let proceed = proceed.clone();
2167 Box::pin(async move {
2168 started.add_permits(1);
2169 proceed.acquire().await.unwrap().forget();
2170 Ok(refreshed_tokens())
2171 })
2172 }
2173 },
2174 );
2175
2176 let acquire = tokio::spawn(async move { lease.acquire().await });
2177 started.acquire().await.unwrap().forget();
2178 assert!(
2179 hub.update_auth_model_cache_details(
2180 "oauth-account",
2181 "oauth-account@account",
2182 10,
2183 &[crate::provider::DiscoveredModelDetails {
2184 slug: "cached-model".into(),
2185 context_budget: Some(16_384),
2186 capability_knowledge: crate::provider::CapabilityKnowledge::Advertised(
2187 crate::provider::ModelCapabilities::default(),
2188 ),
2189 }],
2190 )
2191 .unwrap()
2192 );
2193 proceed.add_permits(1);
2194
2195 assert_eq!(acquire.await.unwrap().unwrap().access_token, "access-v2");
2196 let stored = hub.load_auth().unwrap().providers.remove(0);
2197 assert_eq!(stored.access_token, "access-v2");
2198 assert_eq!(stored.model_cache.unwrap().models[0].slug, "cached-model");
2199 }
2200
2201 #[tokio::test]
2202 async fn rotated_refresh_token_is_reused_and_preserved_when_omitted() {
2203 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2204 let refresh_inputs = Arc::new(Mutex::new(Vec::new()));
2205 let lease = OAuthCredentialLease::with_refresher(
2206 "oauth-account",
2207 ProviderKind::Codex,
2208 hub.clone(),
2209 {
2210 let refresh_inputs = refresh_inputs.clone();
2211 move |refresh_token| {
2212 let call = {
2213 let mut inputs = refresh_inputs
2214 .lock()
2215 .unwrap_or_else(std::sync::PoisonError::into_inner);
2216 inputs.push(refresh_token);
2217 inputs.len()
2218 };
2219 Box::pin(async move {
2220 Ok(TokenResult {
2221 access_token: format!("access-v{}", call + 1),
2222 refresh_token: (call == 1).then(|| "refresh-v2".into()),
2223 expires_at: chrono::Utc::now().timestamp() + 3_600,
2224 account: None,
2225 })
2226 })
2227 }
2228 },
2229 );
2230
2231 assert_eq!(lease.acquire().await.unwrap().access_token, "access-v2");
2232 assert!(
2233 hub.update_auth_tokens(
2234 "oauth-account",
2235 AuthTokenUpdate {
2236 access_token: "access-v2".into(),
2237 refresh_token: None,
2238 expires_at: chrono::Utc::now().timestamp() - 1,
2239 account: None,
2240 },
2241 )
2242 .unwrap()
2243 );
2244 assert_eq!(lease.acquire().await.unwrap().access_token, "access-v3");
2245
2246 let stored = hub.load_auth().unwrap().providers.remove(0);
2247 assert_eq!(stored.refresh_token.as_deref(), Some("refresh-v2"));
2248 assert_eq!(
2249 *refresh_inputs
2250 .lock()
2251 .unwrap_or_else(std::sync::PoisonError::into_inner),
2252 vec!["refresh-v1", "refresh-v2"]
2253 );
2254 }
2255
2256 #[tokio::test]
2257 async fn in_flight_refresh_adopts_newer_authoritative_credentials() {
2258 let (_dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2259 let started = Arc::new(Semaphore::new(0));
2260 let proceed = Arc::new(Semaphore::new(0));
2261 let lease = OAuthCredentialLease::with_refresher(
2262 "oauth-account",
2263 ProviderKind::Codex,
2264 hub.clone(),
2265 {
2266 let started = started.clone();
2267 let proceed = proceed.clone();
2268 move |_| {
2269 let started = started.clone();
2270 let proceed = proceed.clone();
2271 Box::pin(async move {
2272 started.add_permits(1);
2273 proceed.acquire().await.unwrap().forget();
2274 Ok(refreshed_tokens())
2275 })
2276 }
2277 },
2278 );
2279
2280 let acquire = tokio::spawn(async move { lease.acquire().await });
2281 started.acquire().await.unwrap().forget();
2282 assert!(
2283 hub.update_auth_tokens(
2284 "oauth-account",
2285 AuthTokenUpdate {
2286 access_token: "authoritative-access".into(),
2287 refresh_token: Some("authoritative-refresh".into()),
2288 expires_at: chrono::Utc::now().timestamp() + 3_600,
2289 account: Some("authoritative@example.test".into()),
2290 },
2291 )
2292 .unwrap()
2293 );
2294 proceed.add_permits(1);
2295
2296 let credential = acquire.await.unwrap().unwrap();
2297 assert_eq!(credential.access_token, "authoritative-access");
2298 assert_eq!(
2299 credential.display_account.as_deref(),
2300 Some("authoritative@example.test")
2301 );
2302 assert_eq!(
2303 hub.load_auth().unwrap().providers[0].access_token,
2304 "authoritative-access"
2305 );
2306 }
2307
2308 #[tokio::test]
2309 async fn persistence_failure_holds_refresh_lock_until_detached_retry_commits() {
2310 let (dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2311 let refresh_inputs = Arc::new(Mutex::new(Vec::new()));
2312 let lease = OAuthCredentialLease::with_refresher(
2313 "oauth-account",
2314 ProviderKind::Codex,
2315 hub.clone(),
2316 {
2317 let refresh_inputs = refresh_inputs.clone();
2318 move |refresh_token| {
2319 refresh_inputs
2320 .lock()
2321 .unwrap_or_else(std::sync::PoisonError::into_inner)
2322 .push(refresh_token);
2323 Box::pin(async { Ok(refreshed_tokens()) })
2324 }
2325 },
2326 );
2327 let auth_lock = dir.path().join(".auth.json.lock");
2328 std::fs::remove_file(&auth_lock).unwrap();
2329 std::fs::create_dir(&auth_lock).unwrap();
2330
2331 assert!(matches!(
2332 lease.acquire().await,
2333 Err(OAuthCredentialError::Persist { .. })
2334 ));
2335 let refresh_lock = refresh_lock_path(&hub, "oauth-account");
2336 let lock_error = open_refresh_file_lock_with_timeout(
2337 &refresh_lock,
2338 std::time::Duration::from_millis(25),
2339 )
2340 .unwrap_err();
2341 assert_eq!(lock_error.kind(), std::io::ErrorKind::TimedOut);
2342 assert_eq!(
2343 hub.load_auth().unwrap().providers[0].access_token,
2344 "access-v1"
2345 );
2346 assert!(matches!(
2347 lease.acquire().await,
2348 Err(OAuthCredentialError::Persist { .. })
2349 ));
2350 assert_eq!(
2351 refresh_inputs
2352 .lock()
2353 .unwrap_or_else(std::sync::PoisonError::into_inner)
2354 .len(),
2355 1
2356 );
2357
2358 drop(lease);
2359 std::fs::remove_dir(&auth_lock).unwrap();
2360 let stored = tokio::time::timeout(std::time::Duration::from_secs(2), async {
2361 loop {
2362 let stored = hub.load_auth().unwrap().providers.remove(0);
2363 if stored.access_token == "access-v2" {
2364 break stored;
2365 }
2366 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
2367 }
2368 })
2369 .await
2370 .unwrap();
2371 assert_eq!(stored.access_token, "access-v2");
2372 assert_eq!(stored.refresh_token.as_deref(), Some("refresh-v2"));
2373 let _next_opener = RefreshFileLock(
2374 open_refresh_file_lock_with_timeout(
2375 &refresh_lock,
2376 std::time::Duration::from_millis(50),
2377 )
2378 .unwrap(),
2379 );
2380 assert_eq!(
2381 *refresh_inputs
2382 .lock()
2383 .unwrap_or_else(std::sync::PoisonError::into_inner),
2384 vec!["refresh-v1"]
2385 );
2386 }
2387
2388 #[tokio::test]
2389 async fn active_detached_persistence_retry_returns_cached_error_to_callers() {
2390 let (dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2391 let calls = Arc::new(AtomicUsize::new(0));
2392 let lease = OAuthCredentialLease::with_refresher(
2393 "oauth-account",
2394 ProviderKind::Codex,
2395 hub.clone(),
2396 {
2397 let calls = calls.clone();
2398 move |_| {
2399 calls.fetch_add(1, Ordering::SeqCst);
2400 Box::pin(async { Ok(refreshed_tokens()) })
2401 }
2402 },
2403 );
2404 let auth_lock_path = dir.path().join(".auth.json.lock");
2405 std::fs::remove_file(&auth_lock_path).unwrap();
2406 std::fs::create_dir(&auth_lock_path).unwrap();
2407 assert!(matches!(
2408 lease.acquire().await,
2409 Err(OAuthCredentialError::Persist { .. })
2410 ));
2411
2412 std::fs::remove_dir(&auth_lock_path).unwrap();
2413 let auth_lock = std::fs::OpenOptions::new()
2414 .read(true)
2415 .write(true)
2416 .create(true)
2417 .truncate(false)
2418 .open(&auth_lock_path)
2419 .unwrap();
2420 fs2::FileExt::lock_exclusive(&auth_lock).unwrap();
2421 tokio::time::timeout(std::time::Duration::from_secs(2), async {
2422 loop {
2423 let active = lease
2424 .gate
2425 .state
2426 .lock()
2427 .unwrap_or_else(std::sync::PoisonError::into_inner)
2428 .pending_retry_active;
2429 if active {
2430 break;
2431 }
2432 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
2433 }
2434 })
2435 .await
2436 .unwrap();
2437
2438 assert!(matches!(
2439 tokio::time::timeout(std::time::Duration::from_millis(100), lease.acquire())
2440 .await
2441 .unwrap(),
2442 Err(OAuthCredentialError::Persist { .. })
2443 ));
2444 assert_eq!(calls.load(Ordering::SeqCst), 1);
2445
2446 fs2::FileExt::unlock(&auth_lock).unwrap();
2447 tokio::time::timeout(std::time::Duration::from_secs(2), async {
2448 loop {
2449 if hub.load_auth().unwrap().providers[0].access_token == "access-v2" {
2450 break;
2451 }
2452 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
2453 }
2454 })
2455 .await
2456 .unwrap();
2457 }
2458
2459 #[tokio::test]
2460 async fn persistence_retry_discards_pending_after_credential_change() {
2461 let (dir, hub) = hub_with_provider(expired_provider("oauth-account"));
2462 let calls = Arc::new(AtomicUsize::new(0));
2463 let lease = OAuthCredentialLease::with_refresher(
2464 "oauth-account",
2465 ProviderKind::Codex,
2466 hub.clone(),
2467 {
2468 let calls = calls.clone();
2469 move |_| {
2470 calls.fetch_add(1, Ordering::SeqCst);
2471 Box::pin(async { Ok(refreshed_tokens()) })
2472 }
2473 },
2474 );
2475 let auth_lock = dir.path().join(".auth.json.lock");
2476 std::fs::remove_file(&auth_lock).unwrap();
2477 std::fs::create_dir(&auth_lock).unwrap();
2478
2479 assert!(matches!(
2480 lease.acquire().await,
2481 Err(OAuthCredentialError::Persist { .. })
2482 ));
2483 std::fs::remove_dir(&auth_lock).unwrap();
2484 assert!(
2485 hub.update_auth_tokens(
2486 "oauth-account",
2487 AuthTokenUpdate {
2488 access_token: "authoritative-access".into(),
2489 refresh_token: Some("authoritative-refresh".into()),
2490 expires_at: chrono::Utc::now().timestamp() + 3_600,
2491 account: Some("authoritative@example.test".into()),
2492 },
2493 )
2494 .unwrap()
2495 );
2496
2497 let credential = lease.acquire().await.unwrap();
2498 assert_eq!(credential.access_token, "authoritative-access");
2499 assert_eq!(calls.load(Ordering::SeqCst), 1);
2500 let stored = hub.load_auth().unwrap().providers.remove(0);
2501 assert_eq!(stored.access_token, "authoritative-access");
2502 assert_eq!(
2503 stored.refresh_token.as_deref(),
2504 Some("authoritative-refresh")
2505 );
2506 }
2507
2508 #[tokio::test]
2509 async fn expired_credentials_without_refresh_token_require_authentication() {
2510 let mut provider = expired_provider("oauth-account");
2511 provider.refresh_token = None;
2512 let (_dir, hub) = hub_with_provider(provider);
2513 let lease =
2514 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, |_| {
2515 Box::pin(async { panic!("refresher must not run") })
2516 });
2517
2518 assert!(matches!(
2519 lease.acquire().await,
2520 Err(OAuthCredentialError::ReauthenticationRequired(provider))
2521 if provider == "oauth-account"
2522 ));
2523 }
2524
2525 #[tokio::test]
2526 async fn usable_access_token_without_refresh_token_is_not_rejected_early() {
2527 let mut provider = expired_provider("oauth-account");
2528 provider.refresh_token = None;
2529 provider.expires_at = chrono::Utc::now().timestamp() + 60;
2530 let (_dir, hub) = hub_with_provider(provider);
2531 let lease =
2532 OAuthCredentialLease::with_refresher("oauth-account", ProviderKind::Codex, hub, |_| {
2533 Box::pin(async { panic!("refresher must not run") })
2534 });
2535
2536 assert_eq!(lease.acquire().await.unwrap().access_token, "access-v1");
2537 }
2538
2539 #[tokio::test]
2540 async fn snapshot_factory_uses_valid_credentials_without_refreshing() {
2541 let mut stored = expired_provider("in-memory-oauth");
2542 stored.expires_at = chrono::Utc::now().timestamp() + 60;
2543 let provider = create_oauth_provider_no_discover::<SnapshotOAuthProvider>(&stored)
2544 .await
2545 .unwrap();
2546
2547 assert_eq!(provider.access_token, "access-v1");
2548 assert_eq!(
2549 provider.display_account.as_deref(),
2550 Some("display-v1@example.test")
2551 );
2552 }
2553
2554 #[tokio::test]
2555 async fn snapshot_factory_rejects_expired_credentials_without_refreshing() {
2556 let error = create_oauth_provider_no_discover::<SnapshotOAuthProvider>(&expired_provider(
2557 "in-memory-oauth",
2558 ))
2559 .await
2560 .err()
2561 .unwrap();
2562
2563 assert!(error.to_string().contains("use a managed provider"));
2564 }
2565
2566 #[tokio::test]
2567 async fn inert_managed_factory_does_not_refresh_or_require_an_enabled_provider() {
2568 MANAGED_REFRESH_CALLS.store(0, Ordering::SeqCst);
2569 let mut stored = expired_provider("inert-managed-oauth");
2570 stored.enabled = false;
2571 let (_dir, hub) = hub_with_provider(stored.clone());
2572
2573 let provider =
2574 create_managed_oauth_provider_from_stored::<ManagedOAuthProvider>(&stored, hub)
2575 .unwrap();
2576
2577 assert_eq!(MANAGED_REFRESH_CALLS.load(Ordering::SeqCst), 0);
2578 assert!(matches!(
2579 provider.acquire().await,
2580 Err(OAuthCredentialError::Disabled(provider)) if provider == "inert-managed-oauth"
2581 ));
2582 assert_eq!(MANAGED_REFRESH_CALLS.load(Ordering::SeqCst), 0);
2583 }
2584
2585 #[tokio::test]
2586 async fn factories_reject_wrong_provider_kind_before_construction() {
2587 let mut stored = expired_provider("wrong-kind");
2588 stored.kind = ProviderKind::AnthropicOauth;
2589 stored.expires_at = chrono::Utc::now().timestamp() + 3_600;
2590
2591 let snapshot_error = create_oauth_provider_no_discover::<SnapshotOAuthProvider>(&stored)
2592 .await
2593 .err()
2594 .unwrap();
2595 assert!(matches!(
2596 snapshot_error.downcast_ref::<OAuthCredentialError>(),
2597 Some(OAuthCredentialError::KindChanged(provider)) if provider == "wrong-kind"
2598 ));
2599
2600 let (_dir, hub) = hub_with_provider(stored.clone());
2601 let managed_error =
2602 create_oauth_provider_no_discover_with_hub::<SnapshotOAuthProvider>(&stored, hub)
2603 .await
2604 .err()
2605 .unwrap();
2606 assert!(matches!(
2607 managed_error.downcast_ref::<OAuthCredentialError>(),
2608 Some(OAuthCredentialError::KindChanged(provider)) if provider == "wrong-kind"
2609 ));
2610
2611 let mut requested = expired_provider("authoritative-wrong-kind");
2612 requested.expires_at = chrono::Utc::now().timestamp() + 3_600;
2613 let mut authoritative = requested.clone();
2614 authoritative.kind = ProviderKind::AnthropicOauth;
2615 let (_dir, hub) = hub_with_provider(authoritative);
2616 let authoritative_error =
2617 create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(&requested, hub)
2618 .await
2619 .err()
2620 .unwrap();
2621 assert!(matches!(
2622 authoritative_error.downcast_ref::<OAuthCredentialError>(),
2623 Some(OAuthCredentialError::KindChanged(provider))
2624 if provider == "authoritative-wrong-kind"
2625 ));
2626 }
2627
2628 #[tokio::test]
2629 async fn managed_factory_preserves_pending_rotation_after_persistence_failure() {
2630 MANAGED_REFRESH_CALLS.store(0, Ordering::SeqCst);
2631 let stored = expired_provider("managed-oauth");
2632 let (dir, hub) = hub_with_provider(stored.clone());
2633 let auth_lock = dir.path().join(".auth.json.lock");
2634 std::fs::remove_file(&auth_lock).unwrap();
2635 std::fs::create_dir(&auth_lock).unwrap();
2636
2637 let provider = create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(
2638 &stored,
2639 hub.clone(),
2640 )
2641 .await
2642 .unwrap();
2643 assert_eq!(MANAGED_REFRESH_CALLS.load(Ordering::SeqCst), 0);
2644
2645 assert!(matches!(
2646 provider.acquire().await,
2647 Err(OAuthCredentialError::Persist { .. })
2648 ));
2649 assert_eq!(MANAGED_REFRESH_CALLS.load(Ordering::SeqCst), 1);
2650 assert_eq!(
2651 hub.load_auth().unwrap().providers[0].access_token,
2652 "access-v1"
2653 );
2654
2655 std::fs::remove_dir(&auth_lock).unwrap();
2656 let credential = tokio::time::timeout(std::time::Duration::from_secs(2), async {
2657 loop {
2658 match provider.acquire().await {
2659 Ok(credential) => break credential,
2660 Err(OAuthCredentialError::Persist { .. }) => {
2661 tokio::time::sleep(
2662 REFRESH_FAILURE_COOLDOWN + std::time::Duration::from_millis(25),
2663 )
2664 .await;
2665 }
2666 Err(error) => panic!("unexpected credential error: {error}"),
2667 }
2668 }
2669 })
2670 .await
2671 .unwrap();
2672
2673 assert_eq!(credential.access_token, "access-v2");
2674 assert_eq!(MANAGED_REFRESH_CALLS.load(Ordering::SeqCst), 1);
2675 let persisted = hub.load_auth().unwrap().providers.remove(0);
2676 assert_eq!(persisted.access_token, "access-v2");
2677 assert_eq!(persisted.refresh_token.as_deref(), Some("refresh-v2"));
2678 }
2679
2680 #[tokio::test]
2681 async fn managed_factory_does_not_hide_missing_or_disabled_provider() {
2682 let mut stored = expired_provider("managed-oauth");
2683 stored.expires_at = chrono::Utc::now().timestamp() + 3_600;
2684 let missing_dir = tempfile::tempdir().unwrap();
2685 let missing_hub = ConfigHub::from_config_dir(missing_dir.path());
2686 let missing = create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(
2687 &stored,
2688 missing_hub,
2689 )
2690 .await
2691 .err()
2692 .unwrap();
2693 assert!(matches!(
2694 missing.downcast_ref::<OAuthCredentialError>(),
2695 Some(OAuthCredentialError::Missing(provider)) if provider == "managed-oauth"
2696 ));
2697
2698 stored.enabled = false;
2699 let (_dir, disabled_hub) = hub_with_provider(stored.clone());
2700 let disabled = create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(
2701 &stored,
2702 disabled_hub,
2703 )
2704 .await
2705 .err()
2706 .unwrap();
2707 assert!(matches!(
2708 disabled.downcast_ref::<OAuthCredentialError>(),
2709 Some(OAuthCredentialError::Disabled(provider)) if provider == "managed-oauth"
2710 ));
2711 }
2712
2713 #[tokio::test]
2714 async fn managed_provider_rechecks_authoritative_state_at_request_boundary() {
2715 let mut stored = expired_provider("managed-oauth");
2716 stored.expires_at = chrono::Utc::now().timestamp() + 3_600;
2717
2718 let (_missing_dir, missing_hub) = hub_with_provider(stored.clone());
2719 let missing_provider = create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(
2720 &stored,
2721 missing_hub.clone(),
2722 )
2723 .await
2724 .unwrap();
2725 assert!(missing_hub.remove_auth_provider("managed-oauth").unwrap());
2726 assert!(matches!(
2727 missing_provider.acquire().await,
2728 Err(OAuthCredentialError::Missing(provider)) if provider == "managed-oauth"
2729 ));
2730
2731 let (_disabled_dir, disabled_hub) = hub_with_provider(stored.clone());
2732 let disabled_provider = create_oauth_provider_no_discover_with_hub::<ManagedOAuthProvider>(
2733 &stored,
2734 disabled_hub.clone(),
2735 )
2736 .await
2737 .unwrap();
2738 assert!(
2739 disabled_hub
2740 .set_auth_provider_enabled("managed-oauth", false)
2741 .unwrap()
2742 );
2743 assert!(matches!(
2744 disabled_provider.acquire().await,
2745 Err(OAuthCredentialError::Disabled(provider)) if provider == "managed-oauth"
2746 ));
2747 }
2748
2749 #[tokio::test]
2750 async fn managed_factory_rejects_provider_without_lease_support() {
2751 let mut stored = expired_provider("snapshot-only-oauth");
2752 stored.expires_at = chrono::Utc::now().timestamp() + 3_600;
2753 let (_dir, hub) = hub_with_provider(stored.clone());
2754
2755 let error =
2756 create_oauth_provider_no_discover_with_hub::<SnapshotOAuthProvider>(&stored, hub)
2757 .await
2758 .err()
2759 .unwrap();
2760 assert!(
2761 error
2762 .to_string()
2763 .contains("does not support managed OAuth credentials")
2764 );
2765 }
2766
2767 #[test]
2768 fn callback_page_ok_has_right_icon_and_color() {
2769 let html = callback_page(true, "OK", "done");
2770 assert!(html.contains("✓"));
2771 assert!(html.contains("#0078a0"));
2772 }
2773
2774 #[test]
2775 fn callback_page_err_has_right_icon_and_color() {
2776 let html = callback_page(false, "Err", "fail");
2777 assert!(html.contains("✗"));
2778 assert!(html.contains("#c0392b"));
2779 }
2780}