use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use async_trait::async_trait;
use bytes::Bytes;
use hey_sdk::http::{Body, HttpClient, Request, Response, StatusCode};
use hey_sdk::observability::{Hooks, OperationInfo, OperationState};
use hey_sdk::{
AuthStrategy, BearerAuth, Client, Config, Error, ErrorCode, StaticTokenProvider, TokenProvider,
};
use tokio::sync::Notify;
struct Rotating {
token: Mutex<String>,
refreshes: AtomicUsize,
refreshable: bool,
}
impl Rotating {
fn new(refreshable: bool) -> Arc<Rotating> {
Arc::new(Rotating {
token: Mutex::new("token-0".to_string()),
refreshes: AtomicUsize::new(0),
refreshable,
})
}
fn refreshes(&self) -> usize {
self.refreshes.load(Ordering::SeqCst)
}
}
#[async_trait]
impl TokenProvider for Rotating {
async fn access_token(&self) -> Result<String, Error> {
Ok(self.token.lock().unwrap().clone())
}
async fn refresh(&self) -> bool {
let count = self.refreshes.fetch_add(1, Ordering::SeqCst) + 1;
tokio::time::sleep(Duration::from_millis(50)).await;
if self.refreshable {
*self.token.lock().unwrap() = format!("token-{count}");
}
self.refreshable
}
}
struct Stale {
tokens: Mutex<Vec<String>>,
arrived: tokio::sync::Barrier,
}
impl Stale {
fn new(expected: usize) -> Arc<Stale> {
Arc::new(Stale {
tokens: Mutex::new(Vec::new()),
arrived: tokio::sync::Barrier::new(expected),
})
}
fn tokens(&self) -> Vec<String> {
self.tokens.lock().unwrap().clone()
}
}
#[derive(Clone)]
struct Transport(Arc<Stale>);
#[async_trait]
impl HttpClient for Transport {
async fn send(&self, request: Request<Bytes>) -> Result<Response<Body>, Error> {
let token = request.headers()["authorization"]
.to_str()
.unwrap()
.trim_start_matches("Bearer ")
.to_string();
self.0.tokens.lock().unwrap().push(token.clone());
let mut response = Response::new(Body::from("[]"));
if token == "token-0" {
self.0.arrived.wait().await;
*response.status_mut() = StatusCode::UNAUTHORIZED;
}
Ok(response)
}
}
fn client(server: Arc<Stale>, provider: Arc<Rotating>) -> Client {
Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider)
.http_client(Transport(server))
.max_jitter(Duration::ZERO)
.build()
.unwrap()
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_401s_share_one_refresh() {
let server = Stale::new(10);
let provider = Rotating::new(true);
let client = client(server.clone(), provider.clone());
let calls: Vec<_> = (0..10)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.boxes().list().await })
})
.collect();
for call in calls {
call.await.unwrap().unwrap();
}
assert_eq!(provider.refreshes(), 1);
let tokens = server.tokens();
assert_eq!(tokens.len(), 20, "{tokens:?}");
assert_eq!(
tokens.iter().filter(|token| *token == "token-0").count(),
10
);
assert_eq!(
tokens.iter().filter(|token| *token == "token-1").count(),
10
);
}
#[tokio::test]
async fn a_401_that_arrives_after_the_refresh_does_not_refresh_again() {
let server = Stale::new(2);
let provider = Rotating::new(true);
let client = client(server.clone(), provider.clone());
let boxes = client.boxes();
let (first, second) = tokio::join!(boxes.list(), async {
let late = boxes.list();
tokio::time::sleep(Duration::from_millis(200)).await;
late.await
});
first.unwrap();
second.unwrap();
assert_eq!(provider.refreshes(), 1);
assert_eq!(server.tokens().len(), 4);
}
#[tokio::test]
async fn a_caller_cut_off_during_the_refresh_does_not_abandon_it() {
let server = Stale::new(1);
let provider = Rotating::new(true);
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(Transport(server.clone()))
.max_jitter(Duration::ZERO)
.operation_timeout(Duration::from_millis(20))
.build()
.unwrap();
let error = client.boxes().list().await.unwrap_err();
assert!(error.message().contains("timed out"), "{error}");
tokio::time::sleep(Duration::from_millis(100)).await;
assert_eq!(provider.refreshes(), 1);
assert_eq!(
provider.access_token().await.unwrap(),
"token-1",
"the refresh finished without the caller"
);
let unhurried = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(Transport(server.clone()))
.build()
.unwrap();
unhurried.boxes().list().await.unwrap();
assert_eq!(provider.refreshes(), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn a_failed_refresh_leaves_every_caller_with_the_401() {
let server = Stale::new(4);
let provider = Rotating::new(false);
let client = client(server.clone(), provider.clone());
let calls: Vec<_> = (0..4)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.boxes().list().await })
})
.collect();
for call in calls {
let error = call.await.unwrap().unwrap_err();
assert_eq!(error.code(), ErrorCode::Auth);
assert_eq!(error.http_status(), Some(401));
}
assert_eq!(provider.refreshes(), 1);
assert_eq!(server.tokens().len(), 4);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn a_waiter_whose_caller_is_gone_does_not_refresh_for_nobody() {
let server = Stale::new(3);
let provider = Rotating::new(false);
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(Transport(server.clone()))
.max_jitter(Duration::ZERO)
.operation_timeout(Duration::from_millis(20))
.build()
.unwrap();
let calls: Vec<_> = (0..3)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.boxes().list().await })
})
.collect();
for call in calls {
call.await.unwrap().unwrap_err();
}
tokio::time::sleep(Duration::from_millis(300)).await;
assert_eq!(provider.refreshes(), 1);
}
#[derive(Default)]
struct Gate {
open: AtomicBool,
opened: Notify,
}
impl Gate {
fn open(&self) {
self.open.store(true, Ordering::SeqCst);
self.opened.notify_waiters();
}
async fn wait(&self) {
let opened = self.opened.notified();
if !self.open.load(Ordering::SeqCst) {
opened.await;
}
}
}
struct Refusing {
refreshes: AtomicUsize,
started: Arc<Gate>,
release: Arc<Gate>,
}
impl Refusing {
fn new() -> Arc<Refusing> {
Arc::new(Refusing {
refreshes: AtomicUsize::new(0),
started: Arc::default(),
release: Arc::default(),
})
}
fn refreshes(&self) -> usize {
self.refreshes.load(Ordering::SeqCst)
}
}
#[async_trait]
impl TokenProvider for Refusing {
async fn access_token(&self) -> Result<String, Error> {
Ok("token-0".to_string())
}
async fn refresh(&self) -> bool {
self.refreshes.fetch_add(1, Ordering::SeqCst);
self.started.open();
self.release.wait().await;
false
}
}
struct Sequenced {
tokens: Mutex<Vec<String>>,
arrivals: AtomicUsize,
together: usize,
arrived: tokio::sync::Barrier,
held: usize,
gate: Arc<Gate>,
}
impl Sequenced {
fn new(together: usize, held: usize, gate: Arc<Gate>) -> Arc<Sequenced> {
Arc::new(Sequenced {
tokens: Mutex::new(Vec::new()),
arrivals: AtomicUsize::new(0),
together,
arrived: tokio::sync::Barrier::new(together),
held,
gate,
})
}
fn tokens(&self) -> Vec<String> {
self.tokens.lock().unwrap().clone()
}
}
#[derive(Clone)]
struct SequencedTransport(Arc<Sequenced>);
#[async_trait]
impl HttpClient for SequencedTransport {
async fn send(&self, request: Request<Bytes>) -> Result<Response<Body>, Error> {
let token = request.headers()["authorization"]
.to_str()
.unwrap()
.trim_start_matches("Bearer ")
.to_string();
self.0.tokens.lock().unwrap().push(token.clone());
let mut response = Response::new(Body::from("[]"));
if token == "token-0" {
let arrival = self.0.arrivals.fetch_add(1, Ordering::SeqCst) + 1;
if arrival <= self.0.together {
self.0.arrived.wait().await;
}
if arrival == self.0.held {
self.0.gate.wait().await;
}
*response.status_mut() = StatusCode::UNAUTHORIZED;
}
Ok(response)
}
}
struct Ended(Arc<Gate>);
impl Hooks for Ended {
fn on_operation_end(
&self,
_op: &OperationInfo,
_state: OperationState,
_outcome: Result<(), &Error>,
_duration: Duration,
) {
self.0.open();
}
}
#[tokio::test]
async fn a_failed_refresh_is_shared_by_every_request_signed_with_the_credentials_it_was_for() {
let first_failed = Arc::new(Gate::default());
let server = Sequenced::new(2, 2, first_failed.clone());
let provider = Refusing::new();
provider.release.open();
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(SequencedTransport(server.clone()))
.hooks(Ended(first_failed))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let boxes = client.boxes();
let (first, second) = tokio::join!(boxes.list(), boxes.list());
for outcome in [first, second] {
let error = outcome.unwrap_err();
assert_eq!(error.code(), ErrorCode::Auth);
assert_eq!(error.http_status(), Some(401));
}
assert_eq!(
provider.refreshes(),
1,
"one refresh for the one set of credentials"
);
assert_eq!(server.tokens().len(), 2, "and no resend");
let error = boxes.list().await.unwrap_err();
assert_eq!(error.code(), ErrorCode::Auth);
assert_eq!(error.http_status(), Some(401));
assert_eq!(provider.refreshes(), 2);
assert_eq!(server.tokens().len(), 3);
}
#[tokio::test(start_paused = true)]
async fn a_request_whose_401_arrives_during_the_failing_refresh_shares_its_failure() {
let provider = Refusing::new();
let server = Sequenced::new(2, 2, provider.started.clone());
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(SequencedTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let calls: Vec<_> = (0..2)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.boxes().list().await })
})
.collect();
provider.started.wait().await;
tokio::time::sleep(Duration::from_millis(1)).await;
provider.release.open();
for call in calls {
let error = call.await.unwrap().unwrap_err();
assert_eq!(error.code(), ErrorCode::Auth);
assert_eq!(error.http_status(), Some(401));
}
assert_eq!(provider.refreshes(), 1);
assert_eq!(server.tokens().len(), 2);
}
struct Renewing {
signings: AtomicUsize,
refreshes: AtomicUsize,
first_takes: Duration,
}
impl Renewing {
fn new() -> Arc<Renewing> {
Renewing::slow_to_start(Duration::ZERO)
}
fn slow_to_start(first_takes: Duration) -> Arc<Renewing> {
Arc::new(Renewing {
signings: AtomicUsize::new(0),
refreshes: AtomicUsize::new(0),
first_takes,
})
}
fn refreshes(&self) -> usize {
self.refreshes.load(Ordering::SeqCst)
}
}
#[async_trait]
impl TokenProvider for Renewing {
async fn access_token(&self) -> Result<String, Error> {
let signing = self.signings.fetch_add(1, Ordering::SeqCst) + 1;
if signing == 1 {
tokio::time::sleep(self.first_takes).await;
Ok("t0".to_string())
} else if self.refreshes() == 0 {
Ok("t1".to_string())
} else {
Ok("t2".to_string())
}
}
async fn refresh(&self) -> bool {
self.refreshes.fetch_add(1, Ordering::SeqCst);
true
}
}
struct Rejecting {
credentials: Mutex<Vec<String>>,
rejects: Mutex<Vec<String>>,
arrivals: AtomicUsize,
together: usize,
arrived: tokio::sync::Barrier,
}
impl Rejecting {
fn new(together: usize, rejects: &[&str]) -> Arc<Rejecting> {
Arc::new(Rejecting {
credentials: Mutex::new(Vec::new()),
rejects: Mutex::new(rejects.iter().map(ToString::to_string).collect()),
arrivals: AtomicUsize::new(0),
together,
arrived: tokio::sync::Barrier::new(together.max(1)),
})
}
fn credentials(&self) -> Vec<String> {
self.credentials.lock().unwrap().clone()
}
fn reject(&self, tokens: &[&str]) {
*self.rejects.lock().unwrap() = tokens.iter().map(ToString::to_string).collect();
}
}
#[derive(Clone)]
struct RejectingTransport(Arc<Rejecting>);
#[async_trait]
impl HttpClient for RejectingTransport {
async fn send(&self, request: Request<Bytes>) -> Result<Response<Body>, Error> {
let credential = request.headers()["authorization"]
.to_str()
.unwrap()
.to_string();
self.0.credentials.lock().unwrap().push(credential.clone());
let token = credential.trim_start_matches("Bearer ").to_string();
if self.0.arrivals.fetch_add(1, Ordering::SeqCst) < self.0.together {
self.0.arrived.wait().await;
}
let mut response = Response::new(Body::from("[]"));
if self.0.rejects.lock().unwrap().contains(&token) {
*response.status_mut() = StatusCode::UNAUTHORIZED;
}
Ok(response)
}
}
#[tokio::test]
async fn a_token_the_provider_renews_on_its_own_is_a_renewal_and_a_401_on_the_old_one_is_resent() {
let server = Rejecting::new(2, &["t0"]);
let provider = Renewing::new();
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let boxes = client.boxes();
let (first, second) = tokio::join!(boxes.list(), boxes.list());
first.unwrap();
second.unwrap();
assert_eq!(
provider.refreshes(),
0,
"the 401 on t0 was answered by t1, which the provider had already handed over"
);
let credentials = server.credentials();
let mut signed_together = credentials[..2].to_vec();
signed_together.sort();
assert_eq!(
signed_together,
["Bearer t0", "Bearer t1"],
"{credentials:?}"
);
assert_eq!(credentials[2..], ["Bearer t1"], "{credentials:?}");
server.reject(&["t1"]);
boxes.list().await.unwrap();
assert_eq!(
provider.refreshes(),
1,
"a 401 on t1 itself is refreshed, once"
);
assert_eq!(server.credentials()[3..], ["Bearer t1", "Bearer t2"]);
}
#[tokio::test]
async fn a_provider_asked_what_it_would_sign_with_is_not_asked_to_refresh_when_it_has_renewed() {
let server = Rejecting::new(0, &["t0"]);
let provider = Renewing::new();
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
client.boxes().list().await.unwrap();
assert_eq!(
provider.refreshes(),
0,
"the provider had renewed on its own by the time the 401 came back"
);
assert_eq!(server.credentials(), ["Bearer t0", "Bearer t1"]);
server.reject(&["t1"]);
client.boxes().list().await.unwrap();
assert_eq!(
provider.refreshes(),
1,
"a 401 on the token the provider would still sign with is refreshed"
);
assert_eq!(
server.credentials(),
["Bearer t0", "Bearer t1", "Bearer t1", "Bearer t2"]
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn signings_are_recorded_in_the_order_the_provider_issued_in() {
let server = Rejecting::new(0, &["t1"]);
let provider = Renewing::slow_to_start(Duration::from_millis(100));
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let calls: Vec<_> = (0..2)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.boxes().list().await })
})
.collect();
for call in calls {
call.await.unwrap().unwrap();
}
assert_eq!(
provider.refreshes(),
1,
"the 401 on t1, the token the provider issued last, was refreshed"
);
let mut credentials = server.credentials();
credentials.sort();
assert_eq!(
credentials,
["Bearer t0", "Bearer t1", "Bearer t2"],
"each token went out once, and t1 was refreshed to t2 rather than resent"
);
}
#[derive(Default)]
struct Distinct {
signings: AtomicUsize,
refreshes: AtomicUsize,
}
struct DistinctAuth(Arc<Distinct>);
#[async_trait]
impl AuthStrategy for DistinctAuth {
async fn authenticate(&self, request: &mut Request<Bytes>) -> Result<(), Error> {
let signing = self.0.signings.fetch_add(1, Ordering::SeqCst) + 1;
let prefix = if self.0.refreshes.load(Ordering::SeqCst) == 0 {
"signature"
} else {
"renewed"
};
let value = format!("Bearer {prefix}-{signing}").parse().unwrap();
request.headers_mut().insert("authorization", value);
Ok(())
}
async fn refresh(&self) -> bool {
self.0.refreshes.fetch_add(1, Ordering::SeqCst);
true
}
}
#[tokio::test]
async fn a_strategy_of_the_callers_that_signs_every_request_differently_is_not_taken_for_renewing()
{
let server = Rejecting::new(2, &["signature-1"]);
let strategy = Arc::new(Distinct::default());
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.auth_strategy(DistinctAuth(strategy.clone()))
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let boxes = client.boxes();
let (first, second) = tokio::join!(boxes.list(), boxes.list());
first.unwrap();
second.unwrap();
assert_eq!(
strategy.refreshes.load(Ordering::SeqCst),
1,
"the differing signature was not taken for a renewal: the 401 was refreshed"
);
assert_eq!(
server.credentials(),
[
"Bearer signature-1",
"Bearer signature-2",
"Bearer renewed-3"
]
);
}
#[derive(Default)]
struct Failing {
signings: AtomicUsize,
refreshes: AtomicUsize,
}
#[async_trait]
impl TokenProvider for Failing {
async fn access_token(&self) -> Result<String, Error> {
if self.signings.fetch_add(1, Ordering::SeqCst) == 0 {
Ok("t0".to_string())
} else {
Err(Error::auth("the token could not be renewed"))
}
}
async fn refresh(&self) -> bool {
self.refreshes.fetch_add(1, Ordering::SeqCst);
true
}
}
#[tokio::test]
async fn a_provider_that_cannot_hand_over_a_token_fails_the_refresh_without_refreshing() {
let server = Rejecting::new(0, &["t0"]);
let provider = Arc::new(Failing::default());
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.token_provider(provider.clone())
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let error = client.boxes().list().await.unwrap_err();
assert_eq!(error.code(), ErrorCode::Auth, "{error}");
assert_eq!(
provider.refreshes.load(Ordering::SeqCst),
0,
"the failed renewal is the refresh's answer"
);
assert_eq!(server.credentials(), ["Bearer t0"], "and nothing is resent");
}
struct Wrapping {
inner: BearerAuth<StaticTokenProvider>,
signings: AtomicUsize,
refreshes: AtomicUsize,
}
struct WrappingAuth(Arc<Wrapping>);
#[async_trait]
impl AuthStrategy for WrappingAuth {
async fn authenticate(&self, request: &mut Request<Bytes>) -> Result<(), Error> {
self.0.inner.authenticate(request).await?;
let signing = self.0.signings.fetch_add(1, Ordering::SeqCst) + 1;
let signed = format!("Bearer signed-{signing}").parse().unwrap();
request.headers_mut().insert("authorization", signed);
Ok(())
}
async fn refresh(&self) -> bool {
self.0.refreshes.fetch_add(1, Ordering::SeqCst);
true
}
}
#[tokio::test]
async fn a_strategy_of_the_callers_that_signs_through_bearer_auth_is_not_taken_for_renewing() {
let server = Rejecting::new(2, &["signed-1"]);
let strategy = Arc::new(Wrapping {
inner: BearerAuth::new(StaticTokenProvider::new("static")),
signings: AtomicUsize::new(0),
refreshes: AtomicUsize::new(0),
});
let client = Client::builder(Config::default().with_base_url("https://hey.test"))
.auth_strategy(WrappingAuth(strategy.clone()))
.http_client(RejectingTransport(server.clone()))
.max_jitter(Duration::ZERO)
.build()
.unwrap();
let boxes = client.boxes();
let (first, second) = tokio::join!(boxes.list(), boxes.list());
first.unwrap();
second.unwrap();
assert_eq!(
strategy.refreshes.load(Ordering::SeqCst),
1,
"the second request's differing signature was not taken for a renewal"
);
assert_eq!(
server.credentials(),
["Bearer signed-1", "Bearer signed-2", "Bearer signed-3"]
);
}