diff --git a/lib/foundation/fabro-client/src/auth_store.rs b/lib/foundation/fabro-client/src/auth_store.rs index b6bdf8760..4cde4d48b 100644 --- a/lib/foundation/fabro-client/src/auth_store.rs +++ b/lib/foundation/fabro-client/src/auth_store.rs @@ -18,6 +18,9 @@ use fs2::FileExt; use rand::Rng; use serde::{Deserialize, Serialize}; use thiserror::Error; +#[cfg(unix)] +use tokio::task; +use tokio::task::JoinError; use crate::target::ServerTarget; @@ -106,6 +109,8 @@ pub enum LockError { path: PathBuf, source: std::io::Error, }, + #[error("failed to wait for auth store lock at {path}: {source}")] + Task { path: PathBuf, source: JoinError }, } #[derive(Debug, Clone)] @@ -113,6 +118,11 @@ pub struct AuthStore { path: PathBuf, } +#[cfg(unix)] +pub(crate) struct RefreshLockGuard { + _file: std::fs::File, +} + #[derive(Debug, Default, Serialize, Deserialize)] struct AuthFile { #[serde(default)] @@ -193,6 +203,25 @@ impl AuthStore { }) } + #[cfg(unix)] + pub(crate) async fn acquire_refresh_lock(&self) -> Result { + let store = self.clone(); + let lock_path = self.refresh_lock_path(); + // Another CLI can hold this lock through a network request, so keep + // the blocking wait off the Tokio worker threads. + task::spawn_blocking(move || store.acquire_refresh_lock_blocking()) + .await + .map_err(|source| LockError::Task { + path: lock_path, + source, + })? + } + + #[cfg(not(unix))] + pub(crate) async fn acquire_refresh_lock(&self) -> Result<(), AuthStoreError> { + Err(AuthStoreError::UnsupportedPlatform) + } + fn read_auth_file(&self) -> Result { match fs::read_to_string(&self.path) { Ok(contents) => { @@ -280,6 +309,37 @@ impl AuthStore { self.path.with_extension("lock") } + #[cfg(unix)] + fn refresh_lock_path(&self) -> PathBuf { + self.path.with_extension("refresh.lock") + } + + #[cfg(unix)] + fn acquire_refresh_lock_blocking(&self) -> Result { + self.ensure_parent_dir()?; + let path = self.refresh_lock_path(); + let lock_file = std::fs::OpenOptions::new() + .create(true) + .read(true) + .write(true) + .truncate(false) + .open(&path) + .map_err(|source| LockError::Io { + path: path.clone(), + source, + })?; + match FileExt::try_lock_exclusive(&lock_file) { + Ok(()) => {} + Err(source) if source.kind() == std::io::ErrorKind::WouldBlock => { + lock_file + .lock_exclusive() + .map_err(|source| classify_lock_error(path, source))?; + } + Err(source) => return Err(classify_lock_error(path, source).into()), + } + Ok(RefreshLockGuard { _file: lock_file }) + } + #[cfg(unix)] fn lock_error(&self, source: std::io::Error) -> AuthStoreError { classify_lock_error(self.lock_path(), source).into() diff --git a/lib/foundation/fabro-client/src/client.rs b/lib/foundation/fabro-client/src/client.rs index 531b9f2ba..dc3c86823 100644 --- a/lib/foundation/fabro-client/src/client.rs +++ b/lib/foundation/fabro-client/src/client.rs @@ -449,6 +449,10 @@ impl Client { return Ok(()); } + // Refresh tokens are single-use. Hold the cross-process lock from the + // fresh store read through rotation and persistence. AuthStore methods + // take the shorter auth-file lock inside this guard. + let _refresh_guard = oauth_session.auth_store.acquire_refresh_lock().await?; let Some(entry) = oauth_session.auth_store.get(&oauth_session.target)? else { self.rebuild_with_fallback(oauth_session).await?; return Err(session_expired()); @@ -460,6 +464,10 @@ impl Client { } AuthEntry::OAuth(entry) => entry, }; + if oauth_entry.access_token != failed_access_token { + self.rebuild_client(Some(oauth_entry.access_token)).await?; + return Ok(()); + } if oauth_entry.refresh_token_expires_at <= chrono::Utc::now() { oauth_session.auth_store.remove(&oauth_session.target)?; self.rebuild_with_fallback(oauth_session).await?; @@ -2573,6 +2581,78 @@ mod tests { ]); } + #[cfg(unix)] + #[tokio::test] + async fn concurrent_clients_refresh_a_rotating_token_once() { + let server = MockServer::start_async().await; + let refresh_mock = server + .mock_async(|when, then| { + when.method(POST) + .path("/auth/cli/refresh") + .header("authorization", "Bearer refresh-octocat"); + then.status(200) + .delay(Duration::from_millis(100)) + .header("Content-Type", "application/json") + .json_body(json!({ + "access_token": "access-refreshed", + "access_token_expires_at": (chrono::Utc::now() + + ChronoDuration::minutes(10)) + .to_rfc3339(), + "refresh_token": "refresh-refreshed", + "refresh_token_expires_at": (chrono::Utc::now() + + ChronoDuration::days(30)) + .to_rfc3339(), + "subject": { + "idp_issuer": "https://github.com", + "idp_subject": "12345", + "login": "octocat", + "name": "Name octocat", + "email": "octocat@example.com" + } + })); + }) + .await; + let temp = tempfile::tempdir().unwrap(); + let auth_store = AuthStore::new(temp.path().join("auth.json")); + let target = ServerTarget::http_url(server.base_url()).unwrap(); + let entry = oauth_entry("octocat"); + auth_store + .put(&target, AuthEntry::OAuth(entry.clone())) + .unwrap(); + + let first = Client::builder() + .target(target.clone()) + .credential(Credential::OAuth(entry.clone())) + .oauth_session(OAuthSession::new(target.clone(), auth_store.clone())) + .connect() + .await + .unwrap(); + let second = Client::builder() + .target(target.clone()) + .credential(Credential::OAuth(entry)) + .oauth_session(OAuthSession::new(target, auth_store)) + .connect() + .await + .unwrap(); + + let (first_result, second_result) = tokio::join!( + first.refresh_access_token("access-octocat"), + second.refresh_access_token("access-octocat"), + ); + + first_result.unwrap(); + second_result.unwrap(); + refresh_mock.assert_calls_async(1).await; + assert_eq!( + first.current_state().bearer_token.as_deref(), + Some("access-refreshed") + ); + assert_eq!( + second.current_state().bearer_token.as_deref(), + Some("access-refreshed") + ); + } + #[tokio::test] async fn refresh_access_token_classifies_expired_refresh_tokens() { let server = MockServer::start();