use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, PoisonError};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use http::header::{AUTHORIZATION, HeaderName};
use http::{HeaderValue, Request, Response, StatusCode};
use tower::{Layer, Service, ServiceExt};
use super::error::TokenError;
use super::token::Token;
use toolkit_http::HttpError;
use toolkit_utils::SecretString;
type GetTokenFn = Arc<dyn Fn() -> Result<SecretString, TokenError> + Send + Sync>;
type InvalidateTokenFn = Arc<dyn Fn() -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
pub const DEFAULT_MIN_INVALIDATION_INTERVAL: Duration = Duration::from_mins(15);
pub type ShouldRefreshFn = Arc<dyn Fn(StatusCode) -> bool + Send + Sync>;
#[derive(Clone)]
pub struct BearerAuthAutoRefreshOpts {
pub min_invalidation_interval: Duration,
pub should_refresh: ShouldRefreshFn,
pub header_name: HeaderName,
}
impl fmt::Debug for BearerAuthAutoRefreshOpts {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BearerAuthAutoRefreshOpts")
.field("min_invalidation_interval", &self.min_invalidation_interval)
.field("header_name", &self.header_name)
.finish_non_exhaustive()
}
}
fn default_should_refresh(status: StatusCode) -> bool {
status == StatusCode::UNAUTHORIZED
}
impl Default for BearerAuthAutoRefreshOpts {
fn default() -> Self {
let should_refresh: ShouldRefreshFn = Arc::new(default_should_refresh);
Self {
min_invalidation_interval: DEFAULT_MIN_INVALIDATION_INTERVAL,
should_refresh,
header_name: AUTHORIZATION,
}
}
}
#[derive(Clone)]
pub struct BearerAuthAutoRefreshLayer {
get_token: GetTokenFn,
invalidate_token: InvalidateTokenFn,
opts: BearerAuthAutoRefreshOpts,
last_invalidation: Arc<Mutex<Option<Instant>>>,
}
impl fmt::Debug for BearerAuthAutoRefreshLayer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BearerAuthAutoRefreshLayer")
.field("opts", &self.opts)
.finish_non_exhaustive()
}
}
impl BearerAuthAutoRefreshLayer {
#[must_use]
pub fn new(token: Token) -> Self {
Self::with_opts(token, BearerAuthAutoRefreshOpts::default())
}
#[must_use]
pub fn with_opts(token: Token, opts: BearerAuthAutoRefreshOpts) -> Self {
let token_for_get = token.clone();
let get_token: GetTokenFn = Arc::new(move || token_for_get.get());
let invalidate_token: InvalidateTokenFn = Arc::new(move || {
let t = token.clone();
Box::pin(async move { t.invalidate().await })
});
Self::from_fns(get_token, invalidate_token, opts)
}
fn from_fns(
get_token: GetTokenFn,
invalidate_token: InvalidateTokenFn,
opts: BearerAuthAutoRefreshOpts,
) -> Self {
Self {
get_token,
invalidate_token,
opts,
last_invalidation: Arc::new(Mutex::new(None)),
}
}
}
impl<S> Layer<S> for BearerAuthAutoRefreshLayer {
type Service = BearerAuthAutoRefreshService<S>;
fn layer(&self, inner: S) -> Self::Service {
BearerAuthAutoRefreshService {
inner,
get_token: Arc::clone(&self.get_token),
invalidate_token: Arc::clone(&self.invalidate_token),
opts: self.opts.clone(),
last_invalidation: Arc::clone(&self.last_invalidation),
}
}
}
#[derive(Clone)]
pub struct BearerAuthAutoRefreshService<S> {
inner: S,
get_token: GetTokenFn,
invalidate_token: InvalidateTokenFn,
opts: BearerAuthAutoRefreshOpts,
last_invalidation: Arc<Mutex<Option<Instant>>>,
}
impl<S: fmt::Debug> fmt::Debug for BearerAuthAutoRefreshService<S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BearerAuthAutoRefreshService")
.field("inner", &self.inner)
.field("opts", &self.opts)
.finish_non_exhaustive()
}
}
fn build_bearer_value(token: &str) -> Result<HeaderValue, http::header::InvalidHeaderValue> {
let raw = zeroize::Zeroizing::new(format!("Bearer {token}"));
let mut value = HeaderValue::from_str(&raw)?;
value.set_sensitive(true);
Ok(value)
}
fn try_acquire_invalidation_slot(
last_invalidation: &Mutex<Option<Instant>>,
min_interval: Duration,
) -> bool {
let mut guard = last_invalidation
.lock()
.unwrap_or_else(PoisonError::into_inner);
let invalidate = match *guard {
Some(last) => last.elapsed() >= min_interval,
None => true,
};
if invalidate {
*guard = Some(Instant::now());
}
invalidate
}
impl<S, B, ResBody> Service<Request<B>> for BearerAuthAutoRefreshService<S>
where
S: Service<Request<B>, Response = Response<ResBody>, Error = HttpError>
+ Clone
+ Send
+ 'static,
S::Future: Send,
B: Clone + Send + 'static,
ResBody: Send + 'static,
{
type Response = Response<ResBody>;
type Error = HttpError;
type Future = Pin<Box<dyn Future<Output = Result<Response<ResBody>, HttpError>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request<B>) -> Self::Future {
if req.headers().contains_key(&self.opts.header_name) {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
return Box::pin(async move { inner.call(req).await });
}
let initial_secret = match (self.get_token)() {
Ok(s) => s,
Err(e) => return Box::pin(async move { Err(HttpError::Transport(Box::new(e))) }),
};
let bearer_value = match build_bearer_value(initial_secret.expose()) {
Ok(v) => v,
Err(e) => return Box::pin(async move { Err(HttpError::InvalidHeaderValue(e)) }),
};
let retry_req = req.clone();
req.headers_mut()
.insert(self.opts.header_name.clone(), bearer_value);
let outer_clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, outer_clone);
let retry_inner = self.inner.clone();
let get_token = Arc::clone(&self.get_token);
let invalidate_token = Arc::clone(&self.invalidate_token);
let opts = self.opts.clone();
let last_invalidation = Arc::clone(&self.last_invalidation);
let initial_secret_for_compare = initial_secret;
Box::pin(async move {
let response = inner.call(req).await?;
if !(opts.should_refresh)(response.status()) {
return Ok(response);
}
let did_invalidate =
try_acquire_invalidation_slot(&last_invalidation, opts.min_invalidation_interval);
if did_invalidate {
tracing::info!(
status = response.status().as_u16(),
header = %opts.header_name,
"OAuth2 auto-refresh: invalidating token after auth-failure response"
);
invalidate_token().await;
}
let new_secret = match get_token() {
Ok(s) => s,
Err(e) => {
tracing::warn!(
error = %e,
status = response.status().as_u16(),
header = %opts.header_name,
"OAuth2 auto-refresh: token unavailable after invalidate; surfacing original response"
);
return Ok(response);
}
};
if new_secret.expose() == initial_secret_for_compare.expose() {
if did_invalidate {
tracing::warn!(
status = response.status().as_u16(),
header = %opts.header_name,
"OAuth2 auto-refresh: token unchanged after invalidate; surfacing original response"
);
}
return Ok(response);
}
let new_value = match build_bearer_value(new_secret.expose()) {
Ok(v) => v,
Err(e) => return Err(HttpError::InvalidHeaderValue(e)),
};
drop(response);
let mut retry_req = retry_req;
retry_req
.headers_mut()
.insert(opts.header_name.clone(), new_value);
retry_inner.oneshot(retry_req).await
})
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
use bytes::Bytes;
use http::{Method, Request, Response, StatusCode};
use http_body_util::Full;
use httpmock::prelude::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use toolkit_utils::SecretString;
use url::Url;
use crate::oauth2::config::OAuthClientConfig;
fn test_config(server: &MockServer) -> OAuthClientConfig {
OAuthClientConfig {
token_endpoint: Some(
Url::parse(&format!("http://localhost:{}/token", server.port())).unwrap(),
),
client_id: "test-client".into(),
client_secret: SecretString::new("test-secret"),
http_config: Some(toolkit_http::HttpClientConfig::for_testing()),
jitter_max: Duration::from_millis(0),
min_refresh_period: Duration::from_millis(100),
..Default::default()
}
}
fn token_json(token: &str, expires_in: u64) -> String {
format!(r#"{{"access_token":"{token}","expires_in":{expires_in},"token_type":"Bearer"}}"#)
}
fn empty_req() -> Request<Full<Bytes>> {
Request::builder()
.method(Method::GET)
.uri("http://example.com/api")
.body(Full::new(Bytes::new()))
.unwrap()
}
type Script = Arc<Mutex<Vec<(Option<String>, StatusCode)>>>;
#[derive(Clone)]
struct ScriptedService {
script: Script,
header_name: HeaderName,
calls: Arc<AtomicUsize>,
}
impl ScriptedService {
fn new(header_name: HeaderName, script: Vec<(Option<&str>, StatusCode)>) -> Self {
let owned: Vec<(Option<String>, StatusCode)> = script
.into_iter()
.map(|(h, s)| (h.map(std::borrow::ToOwned::to_owned), s))
.collect();
Self {
script: Arc::new(Mutex::new(owned)),
header_name,
calls: Arc::new(AtomicUsize::new(0)),
}
}
}
impl Service<Request<Full<Bytes>>> for ScriptedService {
type Response = Response<Full<Bytes>>;
type Error = HttpError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: Request<Full<Bytes>>) -> Self::Future {
self.calls.fetch_add(1, Ordering::SeqCst);
let actual = req
.headers()
.get(&self.header_name)
.map(|v| v.to_str().unwrap().to_owned());
let mut script = self.script.lock().unwrap();
assert!(
!script.is_empty(),
"ScriptedService called more times than scripted"
);
let (expected, status) = script.remove(0);
assert_eq!(actual, expected, "auth header mismatch on call");
Box::pin(async move {
Ok(Response::builder()
.status(status)
.body(Full::new(Bytes::new()))
.unwrap())
})
}
}
#[test]
fn auto_refresh_is_send_sync_clone() {
fn assert_traits<T: Send + Sync + Clone>() {}
assert_traits::<BearerAuthAutoRefreshLayer>();
assert_traits::<BearerAuthAutoRefreshService<ScriptedService>>();
}
#[tokio::test]
async fn happy_path_no_401() {
let server = MockServer::start();
let token_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
let inner =
ScriptedService::new(AUTHORIZATION, vec![(Some("Bearer tok-A"), StatusCode::OK)]);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(calls.load(Ordering::SeqCst), 1, "no retry on 2xx");
assert_eq!(token_mock.calls(), 1, "no extra token fetch on success");
}
#[tokio::test]
async fn refresh_on_401_then_success() {
let server = MockServer::start();
let mut initial_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
assert_eq!(initial_mock.calls(), 1, "initial fetch");
initial_mock.delete();
let refreshed_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-B", 3600));
});
let inner = ScriptedService::new(
AUTHORIZATION,
vec![
(Some("Bearer tok-A"), StatusCode::UNAUTHORIZED),
(Some("Bearer tok-B"), StatusCode::OK),
],
);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(calls.load(Ordering::SeqCst), 2, "exactly one retry");
assert_eq!(refreshed_mock.calls(), 1, "exactly one invalidate fetch");
}
#[tokio::test]
async fn existing_authorization_header_passes_through() {
let server = MockServer::start();
let token_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-unused", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
assert_eq!(token_mock.calls(), 1, "only the initial fetch");
let inner = ScriptedService::new(
AUTHORIZATION,
vec![(Some("Bearer manual"), StatusCode::UNAUTHORIZED)],
);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let mut req = empty_req();
req.headers_mut()
.insert(AUTHORIZATION, HeaderValue::from_static("Bearer manual"));
let resp = Service::call(&mut svc, req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(calls.load(Ordering::SeqCst), 1, "passthrough - no retry");
assert_eq!(token_mock.calls(), 1, "no invalidate when passthrough");
}
#[tokio::test]
async fn unchanged_token_after_invalidate_does_not_loop() {
let server = MockServer::start();
let token_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-same", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
let inner = ScriptedService::new(
AUTHORIZATION,
vec![(Some("Bearer tok-same"), StatusCode::UNAUTHORIZED)],
);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"no retry - original 401 surfaces"
);
assert_eq!(calls.load(Ordering::SeqCst), 1, "inner called exactly once");
assert_eq!(
token_mock.calls(),
2,
"initial fetch + one invalidate fetch"
);
}
#[derive(Clone)]
struct GatedService {
gate: Arc<tokio::sync::Barrier>,
calls: Arc<AtomicUsize>,
}
impl Service<Request<Full<Bytes>>> for GatedService {
type Response = Response<Full<Bytes>>;
type Error = HttpError;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: Request<Full<Bytes>>) -> Self::Future {
self.calls.fetch_add(1, Ordering::SeqCst);
let gate = Arc::clone(&self.gate);
Box::pin(async move {
gate.wait().await;
Ok(Response::builder()
.status(StatusCode::UNAUTHORIZED)
.body(Full::new(Bytes::new()))
.unwrap())
})
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn throttle_blocks_burst_invalidations() {
const BURST: usize = 5;
let server = MockServer::start();
let token_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-burst", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
let inner = GatedService {
gate: Arc::new(tokio::sync::Barrier::new(BURST)),
calls: Arc::new(AtomicUsize::new(0)),
};
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let svc = layer.layer(inner);
let mut handles = Vec::new();
for _ in 0..BURST {
let mut s = svc.clone();
handles.push(tokio::spawn(async move {
Service::call(&mut s, empty_req()).await
}));
}
for h in handles {
let resp = h.await.unwrap().unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
assert_eq!(
calls.load(Ordering::SeqCst),
BURST,
"no retries - token unchanged"
);
assert_eq!(
token_mock.calls(),
2,
"throttle must produce exactly initial + one invalidate fetch"
);
}
#[tokio::test]
async fn custom_predicate_retries_on_403() {
let server = MockServer::start();
let mut initial_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
initial_mock.delete();
let _refreshed_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-B", 3600));
});
let inner = ScriptedService::new(
AUTHORIZATION,
vec![
(Some("Bearer tok-A"), StatusCode::FORBIDDEN),
(Some("Bearer tok-B"), StatusCode::OK),
],
);
let opts = BearerAuthAutoRefreshOpts {
should_refresh: Arc::new(|s| s == StatusCode::FORBIDDEN),
..BearerAuthAutoRefreshOpts::default()
};
let layer = BearerAuthAutoRefreshLayer::with_opts(token, opts);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn custom_predicate_does_not_retry_on_500() {
let server = MockServer::start();
let token_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
let inner = ScriptedService::new(
AUTHORIZATION,
vec![(Some("Bearer tok-A"), StatusCode::INTERNAL_SERVER_ERROR)],
);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(calls.load(Ordering::SeqCst), 1, "no retry on 500");
assert_eq!(token_mock.calls(), 1, "no invalidate on 500");
}
#[tokio::test]
async fn custom_header_name_is_used() {
let server = MockServer::start();
let mut initial_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
initial_mock.delete();
let _refreshed_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-B", 3600));
});
let custom = HeaderName::from_static("x-api-key");
let inner = ScriptedService::new(
custom.clone(),
vec![
(Some("Bearer tok-A"), StatusCode::UNAUTHORIZED),
(Some("Bearer tok-B"), StatusCode::OK),
],
);
let opts = BearerAuthAutoRefreshOpts {
header_name: custom,
..BearerAuthAutoRefreshOpts::default()
};
let layer = BearerAuthAutoRefreshLayer::with_opts(token, opts);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[derive(Clone)]
struct FakeToken {
state: Arc<Mutex<FakeTokenState>>,
}
struct FakeTokenState {
current: Option<String>,
next_swaps: std::collections::VecDeque<Option<String>>,
invalidate_calls: usize,
}
impl FakeToken {
fn new(initial: Option<&str>, swaps: Vec<Option<&str>>) -> Self {
Self {
state: Arc::new(Mutex::new(FakeTokenState {
current: initial.map(str::to_owned),
next_swaps: swaps.into_iter().map(|o| o.map(str::to_owned)).collect(),
invalidate_calls: 0,
})),
}
}
fn invalidate_calls(&self) -> usize {
self.state.lock().unwrap().invalidate_calls
}
fn build_layer(&self) -> BearerAuthAutoRefreshLayer {
self.build_layer_with_opts(BearerAuthAutoRefreshOpts::default())
}
fn build_layer_with_opts(
&self,
opts: BearerAuthAutoRefreshOpts,
) -> BearerAuthAutoRefreshLayer {
let s_get = Arc::clone(&self.state);
let s_inv = Arc::clone(&self.state);
let get_token: GetTokenFn = Arc::new(move || {
let s = s_get.lock().unwrap();
match &s.current {
Some(t) => Ok(SecretString::new(t)),
None => Err(TokenError::Unavailable("fake token unavailable".into())),
}
});
let invalidate_token: InvalidateTokenFn = Arc::new(move || {
let state = Arc::clone(&s_inv);
Box::pin(async move {
let mut s = state.lock().unwrap();
s.invalidate_calls += 1;
if let Some(next) = s.next_swaps.pop_front() {
s.current = next;
}
})
});
BearerAuthAutoRefreshLayer::from_fns(get_token, invalidate_token, opts)
}
}
#[tokio::test]
async fn token_unavailable_on_first_get() {
let token = FakeToken::new(None, vec![]);
let inner = ScriptedService::new(AUTHORIZATION, vec![]);
let calls = Arc::clone(&inner.calls);
let layer = token.build_layer();
let mut svc = layer.layer(inner);
let err = Service::call(&mut svc, empty_req()).await.unwrap_err();
assert!(
matches!(err, HttpError::Transport(_)),
"expected Transport error, got: {err:?}"
);
assert_eq!(calls.load(Ordering::SeqCst), 0, "inner never invoked");
assert_eq!(token.invalidate_calls(), 0, "no invalidate before any send");
}
#[tokio::test]
async fn failed_invalidate_surfaces_original_response() {
let server = MockServer::start();
let mut success_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("tok-A", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
success_mock.delete();
let fail_mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(500)
.header("content-type", "application/json")
.body(r#"{"error":"server_error"}"#);
});
let inner = ScriptedService::new(
AUTHORIZATION,
vec![(Some("Bearer tok-A"), StatusCode::UNAUTHORIZED)],
);
let calls = Arc::clone(&inner.calls);
let layer = BearerAuthAutoRefreshLayer::new(token);
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"original 401 surfaces when invalidate fails"
);
assert_eq!(calls.load(Ordering::SeqCst), 1, "inner called exactly once");
assert!(
fail_mock.calls() >= 1,
"the failed invalidate must have hit the token endpoint"
);
}
#[tokio::test]
async fn token_unavailable_after_invalidate_surfaces_original_response() {
let token = FakeToken::new(Some("tok-A"), vec![None]);
let inner = ScriptedService::new(
AUTHORIZATION,
vec![(Some("Bearer tok-A"), StatusCode::UNAUTHORIZED)],
);
let calls = Arc::clone(&inner.calls);
let layer = token.build_layer();
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"original 401 surfaces when post-invalidate get returns Err"
);
assert_eq!(calls.load(Ordering::SeqCst), 1, "inner called exactly once");
assert_eq!(token.invalidate_calls(), 1, "invalidate fired exactly once");
}
#[tokio::test]
async fn retry_response_also_401_surfaces_without_second_invalidate() {
let token = FakeToken::new(Some("tok-A"), vec![Some("tok-B")]);
let inner = ScriptedService::new(
AUTHORIZATION,
vec![
(Some("Bearer tok-A"), StatusCode::UNAUTHORIZED),
(Some("Bearer tok-B"), StatusCode::UNAUTHORIZED),
],
);
let calls = Arc::clone(&inner.calls);
let layer = token.build_layer();
let mut svc = layer.layer(inner);
let resp = Service::call(&mut svc, empty_req()).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"retry's 401 surfaces verbatim - no second retry"
);
assert_eq!(calls.load(Ordering::SeqCst), 2, "exactly one retry");
assert_eq!(
token.invalidate_calls(),
1,
"no second invalidate after retry-401"
);
}
#[tokio::test]
async fn debug_does_not_reveal_tokens() {
let server = MockServer::start();
let _mock = server.mock(|when, then| {
when.method(POST).path("/token");
then.status(200)
.header("content-type", "application/json")
.body(token_json("super-secret-auto", 3600));
});
let token = Token::new(test_config(&server)).await.unwrap();
let layer = BearerAuthAutoRefreshLayer::new(token);
let dbg = format!("{layer:?}");
assert!(
!dbg.contains("super-secret-auto"),
"Debug must not reveal token: {dbg}"
);
}
#[test]
fn opts_debug_does_not_reveal_predicate() {
let opts = BearerAuthAutoRefreshOpts::default();
let dbg = format!("{opts:?}");
assert!(dbg.contains("min_invalidation_interval"));
assert!(dbg.contains("header_name"));
assert!(!dbg.contains("should_refresh"), "predicate field elided");
}
}