fix(auth): serialize CLI token refreshes across processes

This commit is contained in:
Bryan Helmkamp 2026-07-29 10:07:03 -04:00
parent d8434e7672
commit 0e2ee787bc
No known key found for this signature in database
2 changed files with 140 additions and 0 deletions

View file

@ -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<RefreshLockGuard, AuthStoreError> {
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<AuthFile, AuthStoreError> {
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<RefreshLockGuard, AuthStoreError> {
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()

View file

@ -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();