use std::pin::Pin;
use std::sync::{Arc, Mutex, Weak};
use async_trait::async_trait;
use google_cloud_auth::credentials::{
Builder as AccessTokenCredentialBuilder, CacheableResource, Credentials, EntityTag,
};
use reqwest::{Request, Response};
use reqwest_middleware::{Middleware, Next, Result as MiddlewareResult};
use tokio::sync::Notify;
use url::Url;
struct CachedResource {
entity_tag: EntityTag,
headers: http::HeaderMap,
}
enum CacheState {
Empty,
Refreshing(Weak<()>),
Ready(CachedResource),
}
struct GCSInner {
credential: Mutex<Option<Credentials>>,
cache: Mutex<CacheState>,
refresh_done: Notify,
}
#[derive(Clone)]
pub struct GCSMiddleware {
inner: Arc<GCSInner>,
}
enum PollResult<'a> {
Wait(Pin<Box<dyn std::future::Future<Output = ()> + Send + 'a>>),
Retry,
Validate {
entity_tag: EntityTag,
headers: http::HeaderMap,
},
StartRefresh(Arc<()>),
}
impl Default for GCSMiddleware {
fn default() -> Self {
Self {
inner: Arc::new(GCSInner {
credential: Mutex::new(None),
cache: Mutex::new(CacheState::Empty),
refresh_done: Notify::new(),
}),
}
}
}
#[cfg(test)]
impl GCSMiddleware {
fn with_credentials(cred: Credentials) -> Self {
Self {
inner: Arc::new(GCSInner {
credential: Mutex::new(Some(cred)),
cache: Mutex::new(CacheState::Empty),
refresh_done: Notify::new(),
}),
}
}
}
#[async_trait]
impl Middleware for GCSMiddleware {
async fn handle(
&self,
mut req: Request,
extensions: &mut http::Extensions,
next: Next<'_>,
) -> MiddlewareResult<Response> {
if req.url().scheme() == "gcs" {
let mut url = req.url().clone();
let bucket_name = url.host_str().ok_or_else(|| {
reqwest_middleware::Error::Middleware(anyhow::anyhow!(
"Host should be present in GCS URL, got: {url}"
))
})?;
let new_url = format!(
"https://storage.googleapis.com/{}{}",
bucket_name,
url.path()
);
url = Url::parse(&new_url).map_err(|e| {
reqwest_middleware::Error::Middleware(anyhow::anyhow!(
"Failed to parse constructed GCS URL '{new_url}': {e}"
))
})?;
*req.url_mut() = url;
req = self.authenticate(req).await?;
}
next.run(req, extensions).await
}
}
impl GCSMiddleware {
async fn authenticate(&self, mut req: Request) -> MiddlewareResult<Request> {
let headers = self.get_or_refresh_token().await?;
req.headers_mut().extend(headers);
Ok(req)
}
async fn get_credential(&self) -> MiddlewareResult<Credentials> {
let mut guard = self.inner.credential.lock().unwrap();
if guard.is_none() {
let scopes = ["https://www.googleapis.com/auth/devstorage.read_only"];
let c = AccessTokenCredentialBuilder::default()
.with_scopes(scopes)
.build()
.map_err(|e| reqwest_middleware::Error::Middleware(anyhow::Error::new(e)))?;
*guard = Some(c);
}
Ok(guard.as_ref().unwrap().clone())
}
fn poll_cache<'a>(&'a self) -> PollResult<'a> {
let mut guard = self.inner.cache.lock().unwrap();
let state = std::mem::replace(&mut *guard, CacheState::Empty);
match state {
CacheState::Refreshing(weak) if weak.upgrade().is_some() => {
*guard = CacheState::Refreshing(weak);
let mut notified = Box::pin(self.inner.refresh_done.notified());
notified.as_mut().enable();
drop(guard);
PollResult::Wait(notified)
}
CacheState::Refreshing(_dead) => PollResult::Retry,
CacheState::Ready(r) => {
let entity_tag = r.entity_tag.clone();
let headers = r.headers.clone();
*guard = CacheState::Ready(r);
PollResult::Validate {
entity_tag,
headers,
}
}
CacheState::Empty => {
let token = Arc::new(());
*guard = CacheState::Refreshing(Arc::downgrade(&token));
PollResult::StartRefresh(token)
}
}
}
async fn get_or_refresh_token(&self) -> MiddlewareResult<http::HeaderMap> {
loop {
match self.poll_cache() {
PollResult::Wait(notified) => {
notified.await;
}
PollResult::Retry => {}
PollResult::Validate {
entity_tag,
headers,
} => {
let cred = self.get_credential().await?;
let mut ext = http::Extensions::new();
ext.insert(entity_tag);
return match cred
.headers(ext)
.await
.map_err(|e| reqwest_middleware::Error::Middleware(anyhow::Error::new(e)))?
{
CacheableResource::NotModified => Ok(headers),
CacheableResource::New { entity_tag, data } => {
*self.inner.cache.lock().unwrap() = CacheState::Ready(CachedResource {
entity_tag,
headers: data.clone(),
});
Ok(data)
}
};
}
PollResult::StartRefresh(token) => {
let mut refresh_guard = RefreshGuard {
inner: Arc::clone(&self.inner),
_token: token,
defused: false,
};
let cred = self.get_credential().await?;
let fetch = cred
.headers(http::Extensions::new())
.await
.map_err(|e| reqwest_middleware::Error::Middleware(anyhow::Error::new(e)));
match fetch {
Ok(CacheableResource::New { entity_tag, data }) => {
let out = data.clone();
*self.inner.cache.lock().unwrap() = CacheState::Ready(CachedResource {
entity_tag,
headers: data,
});
refresh_guard.defused = true;
self.inner.refresh_done.notify_waiters();
return Ok(out);
}
Ok(CacheableResource::NotModified) => unreachable!(
"no entity tag was provided in extensions, \
so NotModified cannot be returned"
),
Err(e) => {
*self.inner.cache.lock().unwrap() = CacheState::Empty;
refresh_guard.defused = true;
self.inner.refresh_done.notify_waiters();
return Err(e);
}
}
}
}
}
}
}
struct RefreshGuard {
inner: Arc<GCSInner>,
_token: Arc<()>,
defused: bool,
}
impl Drop for RefreshGuard {
fn drop(&mut self) {
if !self.defused {
self.inner.refresh_done.notify_waiters();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
use google_cloud_auth::credentials::{CacheableResource, CredentialsProvider, EntityTag};
use google_cloud_auth::errors::CredentialsError;
use reqwest::Client;
use tempfile;
use tokio::sync::Barrier;
use super::*;
type SharedEtag = Arc<std::sync::Mutex<EntityTag>>;
#[derive(Debug)]
struct MockProvider {
refresh_count: Arc<AtomicUsize>,
current_etag: SharedEtag,
barrier: Option<Arc<Barrier>>,
}
impl MockProvider {
fn new() -> (Self, Arc<AtomicUsize>, SharedEtag) {
let count = Arc::new(AtomicUsize::new(0));
let etag: SharedEtag = Arc::new(std::sync::Mutex::new(EntityTag::new()));
let p = Self {
refresh_count: count.clone(),
current_etag: Arc::clone(&etag),
barrier: None,
};
(p, count, etag)
}
fn with_barrier(barrier: Arc<Barrier>) -> (Self, Arc<AtomicUsize>, SharedEtag) {
let (mut p, count, etag) = Self::new();
p.barrier = Some(barrier);
(p, count, etag)
}
}
impl CredentialsProvider for MockProvider {
async fn headers(
&self,
extensions: http::Extensions,
) -> Result<CacheableResource<http::HeaderMap>, CredentialsError> {
let current = self.current_etag.lock().unwrap().clone();
if let Some(caller_tag) = extensions.get::<EntityTag>() {
if *caller_tag == current {
return Ok(CacheableResource::NotModified);
}
}
self.refresh_count.fetch_add(1, Ordering::SeqCst);
if let Some(ref b) = self.barrier {
b.wait().await;
}
let mut map = http::HeaderMap::new();
map.insert(
http::header::AUTHORIZATION,
"Bearer mock-token".parse().unwrap(),
);
Ok(CacheableResource::New {
entity_tag: current,
data: map,
})
}
async fn universe_domain(&self) -> Option<String> {
None
}
}
#[tokio::test]
async fn test_cache_reuses_valid_token() {
let (provider, count, _etag) = MockProvider::new();
let mw = GCSMiddleware::with_credentials(Credentials::from(provider));
let h1 = mw.get_or_refresh_token().await.unwrap();
let h2 = mw.get_or_refresh_token().await.unwrap();
let h3 = mw.get_or_refresh_token().await.unwrap();
assert_eq!(count.load(Ordering::SeqCst), 1);
assert_eq!(h1, h2);
assert_eq!(h2, h3);
}
#[tokio::test]
async fn test_cache_refreshes_on_token_change() {
let (provider, count, etag_handle) = MockProvider::new();
let mw = GCSMiddleware::with_credentials(Credentials::from(provider));
mw.get_or_refresh_token().await.unwrap();
assert_eq!(count.load(Ordering::SeqCst), 1);
*etag_handle.lock().unwrap() = EntityTag::new();
mw.get_or_refresh_token().await.unwrap();
assert_eq!(count.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_singleflight_under_concurrent_load() {
const TASKS: usize = 10;
let barrier = Arc::new(Barrier::new(2));
let (provider, count, _etag) = MockProvider::with_barrier(Arc::clone(&barrier));
let mw = GCSMiddleware::with_credentials(Credentials::from(provider));
let handles: Vec<_> = (0..TASKS)
.map(|_| {
let mw = mw.clone();
tokio::spawn(async move { mw.get_or_refresh_token().await.unwrap() })
})
.collect();
barrier.wait().await;
for handle in handles {
handle.await.unwrap();
}
assert_eq!(
count.load(Ordering::SeqCst),
1,
"expected exactly 1 refresh call, got {}",
count.load(Ordering::SeqCst)
);
}
#[tokio::test]
async fn test_cancellation_during_refresh_does_not_deadlock() {
let barrier = Arc::new(Barrier::new(2));
let (provider, _count, _etag) = MockProvider::with_barrier(Arc::clone(&barrier));
let mw = GCSMiddleware::with_credentials(Credentials::from(provider));
let mw_clone = mw.clone();
let handle = tokio::spawn(async move { mw_clone.get_or_refresh_token().await });
barrier.wait().await;
handle.abort();
let _ = handle.await;
tokio::time::timeout(std::time::Duration::from_secs(5), mw.get_or_refresh_token())
.await
.expect("timed out – RefreshGuard did not unblock callers on cancellation")
.expect("token fetch failed after cancellation recovery");
}
#[tokio::test]
async fn test_gcs_middleware() {
let credentials = match std::env::var("GOOGLE_CLOUD_TEST_KEY_JSON") {
Ok(credentials) if !credentials.is_empty() => credentials,
Err(_) | Ok(_) => {
eprintln!("Skipping test as GOOGLE_CLOUD_TEST_KEY_JSON is not set");
return;
}
};
println!("Running GCS Test");
let key_file = tempfile::NamedTempFile::with_suffix(".json").unwrap();
std::fs::write(&key_file, credentials).unwrap();
let prev_value = std::env::var("GOOGLE_APPLICATION_CREDENTIALS").ok();
std::env::set_var("GOOGLE_APPLICATION_CREDENTIALS", key_file.path());
let client = reqwest_middleware::ClientBuilder::new(Client::new())
.with(GCSMiddleware::default())
.build();
let url = "gcs://test-channel/noarch/repodata.json";
let response = client.get(url).send().await.unwrap();
assert!(response.status().is_success());
let url = "gcs://test-channel-nonexist/noarch/repodata.json";
let response = client.get(url).send().await.unwrap();
assert!(response.status().is_client_error());
if let Some(value) = prev_value {
std::env::set_var("GOOGLE_APPLICATION_CREDENTIALS", value);
} else {
std::env::remove_var("GOOGLE_APPLICATION_CREDENTIALS");
}
}
}