use std::{
collections::HashMap,
fmt,
sync::{Arc, Mutex},
time::{Duration, SystemTime, UNIX_EPOCH},
};
use base64::{
Engine as _,
engine::general_purpose::{URL_SAFE, URL_SAFE_NO_PAD},
};
use reqwest_middleware::{Middleware, Next};
use serde::Deserialize;
use thiserror::Error;
use url::Url;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Challenge {
pub scheme: String,
pub params: HashMap<String, String>,
}
pub fn parse_challenges(headers: &http::HeaderMap) -> Vec<Challenge> {
headers
.get_all(http::header::WWW_AUTHENTICATE)
.iter()
.filter_map(|value| value.to_str().ok())
.filter_map(|value| http_auth::parse_challenges(value).ok())
.flatten()
.map(|challenge| Challenge {
scheme: challenge.scheme.to_string(),
params: challenge
.params
.iter()
.map(|(key, value)| (key.to_ascii_lowercase(), value.to_unescaped()))
.collect(),
})
.collect()
}
const TOKEN_REFRESH_MARGIN: Duration = Duration::from_secs(60);
#[derive(Clone, Deserialize)]
#[serde(transparent)]
pub struct BearerToken(String);
impl BearerToken {
pub fn new(token: String) -> Self {
Self(token)
}
pub fn secret(&self) -> &str {
&self.0
}
}
impl fmt::Debug for BearerToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("BearerToken").field(&"<redacted>").finish()
}
}
#[derive(Debug, Error)]
#[error("authentication flow failed: {source}")]
pub struct AuthFlowError {
#[source]
source: Box<dyn std::error::Error + Send + Sync + 'static>,
}
impl AuthFlowError {
pub fn new(err: impl Into<Box<dyn std::error::Error + Send + Sync + 'static>>) -> Self {
Self { source: err.into() }
}
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
pub trait AuthFlow: Send + Sync + fmt::Debug {
async fn acquire_token(
&self,
url: &Url,
challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError>;
}
#[derive(Deserialize)]
struct JwtClaims {
exp: Option<u64>,
}
fn jwt_expiration(token: &str) -> Option<SystemTime> {
let mut parts = token.split('.');
let _header = parts.next()?;
let payload = parts.next()?;
let _signature = parts.next()?;
if parts.next().is_some() {
return None;
}
let payload = URL_SAFE_NO_PAD
.decode(payload)
.or_else(|_| URL_SAFE.decode(payload))
.ok()?;
let claims: JwtClaims = serde_json::from_slice(&payload).ok()?;
claims
.exp
.and_then(|exp| UNIX_EPOCH.checked_add(Duration::from_secs(exp)))
}
#[derive(Debug)]
struct CachedToken {
header: reqwest::header::HeaderValue,
expires_at: Option<SystemTime>,
}
impl CachedToken {
fn new(token: &BearerToken) -> Result<Self, reqwest::header::InvalidHeaderValue> {
let mut header =
reqwest::header::HeaderValue::from_str(&format!("Bearer {}", token.secret()))?;
header.set_sensitive(true);
Ok(Self {
header,
expires_at: jwt_expiration(token.secret()),
})
}
fn is_fresh(&self, now: SystemTime) -> bool {
self.expires_at
.is_none_or(|expires_at| now + TOKEN_REFRESH_MARGIN < expires_at)
}
}
#[derive(Debug, Default)]
enum TokenCache {
#[default]
Empty,
Disabled,
Token(CachedToken),
}
#[derive(Debug, Default)]
struct OriginEntry {
state: TokenCache,
gate: Arc<futures::lock::Mutex<()>>,
}
enum CacheLookup {
Empty,
Disabled,
Fresh(reqwest::header::HeaderValue),
}
type OriginKey = (String, String, u16);
fn origin_key(url: &Url) -> Option<OriginKey> {
Some((
url.scheme().to_string(),
url.host_str()?.to_string(),
url.port_or_known_default()?,
))
}
#[derive(Clone, Debug)]
pub struct AuthChallengeMiddleware {
flows: Vec<Arc<dyn AuthFlow>>,
caches: Arc<Mutex<HashMap<OriginKey, OriginEntry>>>,
}
impl Default for AuthChallengeMiddleware {
fn default() -> Self {
Self::new(vec![Arc::new(
crate::trusted_publishing::PrefixAuthAmbientFlow::default(),
)])
}
}
impl AuthChallengeMiddleware {
pub fn new(flows: Vec<Arc<dyn AuthFlow>>) -> Self {
Self {
flows,
caches: Arc::new(Mutex::new(HashMap::new())),
}
}
fn lookup_cache(&self, origin: &OriginKey) -> CacheLookup {
let caches = self
.caches
.lock()
.expect("auth challenge token cache poisoned");
match caches.get(origin).map(|entry| &entry.state) {
Some(TokenCache::Disabled) => CacheLookup::Disabled,
Some(TokenCache::Token(cached)) if cached.is_fresh(SystemTime::now()) => {
CacheLookup::Fresh(cached.header.clone())
}
None | Some(TokenCache::Empty | TokenCache::Token(_)) => CacheLookup::Empty,
}
}
fn store_cache(&self, origin: OriginKey, state: TokenCache) {
self.caches
.lock()
.expect("auth challenge token cache poisoned")
.entry(origin)
.or_default()
.state = state;
}
fn gate(&self, origin: &OriginKey) -> Arc<futures::lock::Mutex<()>> {
self.caches
.lock()
.expect("auth challenge token cache poisoned")
.entry(origin.clone())
.or_default()
.gate
.clone()
}
async fn acquire_serialized(
&self,
origin: OriginKey,
url: &Url,
challenges: &[Challenge],
rejected: Option<&reqwest::header::HeaderValue>,
) -> Option<reqwest::header::HeaderValue> {
let gate = self.gate(&origin);
let _guard = gate.lock().await;
{
let mut caches = self
.caches
.lock()
.expect("auth challenge token cache poisoned");
let entry = caches.entry(origin.clone()).or_default();
match &entry.state {
TokenCache::Disabled => return None,
TokenCache::Token(cached) => {
let ours = rejected.is_some_and(|header| *header == cached.header);
if !ours && cached.is_fresh(SystemTime::now()) {
return Some(cached.header.clone());
}
entry.state = TokenCache::Empty;
}
TokenCache::Empty => {}
}
}
for flow in &self.flows {
match flow.acquire_token(url, challenges).await {
Ok(Some(token)) => match CachedToken::new(&token) {
Ok(cached) => {
let header = cached.header.clone();
self.store_cache(origin, TokenCache::Token(cached));
return Some(header);
}
Err(err) => {
tracing::warn!(
"AuthChallengeMiddleware: {flow:?} returned a token for {url} \
that is not a valid header value ({err}), trying next flow"
);
}
},
Ok(None) => {
tracing::debug!(
"AuthChallengeMiddleware: {flow:?} not applicable for {url}, \
trying next flow"
);
}
Err(err) => {
tracing::warn!(
"AuthChallengeMiddleware: {flow:?} failed to acquire a token \
for {url}: {err}, trying next flow"
);
}
}
}
tracing::debug!(
"AuthChallengeMiddleware: no flow produced a token for {url}, \
disabling its origin"
);
self.store_cache(origin, TokenCache::Disabled);
None
}
}
fn attach_bearer(req: &mut reqwest::Request, header: reqwest::header::HeaderValue) {
req.headers_mut()
.insert(reqwest::header::AUTHORIZATION, header);
}
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
impl Middleware for AuthChallengeMiddleware {
async fn handle(
&self,
mut req: reqwest::Request,
extensions: &mut http::Extensions,
next: Next<'_>,
) -> reqwest_middleware::Result<reqwest::Response> {
let origin = origin_key(req.url());
let Some(origin) = origin else {
return next.run(req, extensions).await;
};
if req.headers().contains_key(reqwest::header::AUTHORIZATION) {
return next.run(req, extensions).await;
}
let cached = self.lookup_cache(&origin);
if matches!(cached, CacheLookup::Disabled) {
return next.run(req, extensions).await;
}
let used_header = if let CacheLookup::Fresh(header) = cached {
attach_bearer(&mut req, header.clone());
Some(header)
} else {
None
};
let url = req.url().clone();
let Some(mut retry_req) = req.try_clone() else {
let response = next.run(req, extensions).await?;
if !parse_challenges(response.headers()).is_empty() {
tracing::warn!(
"AuthChallengeMiddleware: {url} responded with a challenge but the \
request body could not be cloned for replay; returning the \
challenge response unmodified"
);
}
return Ok(response);
};
let response = next.clone().run(req, extensions).await?;
let challenges = parse_challenges(response.headers());
if challenges.is_empty() {
return Ok(response);
}
let Some(header) = self
.acquire_serialized(origin, &url, &challenges, used_header.as_ref())
.await
else {
return Ok(response);
};
attach_bearer(&mut retry_req, header);
next.run(retry_req, extensions).await
}
}
#[cfg(test)]
mod tests {
use std::{
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
time::{Duration, UNIX_EPOCH},
};
use reqwest_middleware::ClientBuilder;
use super::*;
#[derive(Debug)]
struct StaticFlow {
token: Option<&'static str>,
calls: AtomicUsize,
}
impl StaticFlow {
fn new(token: Option<&'static str>) -> Arc<Self> {
Arc::new(Self {
token,
calls: AtomicUsize::new(0),
})
}
}
#[async_trait::async_trait]
impl AuthFlow for StaticFlow {
async fn acquire_token(
&self,
_url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(self.token.map(|t| BearerToken::new(t.to_string())))
}
}
fn protected_router(accept: String, hits: Arc<AtomicUsize>) -> axum::Router {
use axum::{http::StatusCode, response::IntoResponse, routing::get};
axum::Router::new().route(
"/channel/repodata.json",
get(move |headers: axum::http::HeaderMap| {
let hits = hits.clone();
let expected = format!("Bearer {accept}");
async move {
hits.fetch_add(1, Ordering::SeqCst);
match headers.get("authorization").and_then(|v| v.to_str().ok()) {
Some(auth) if auth == expected => (StatusCode::OK, "ok").into_response(),
_ => (
StatusCode::UNAUTHORIZED,
[("www-authenticate", r#"Bearer realm="test""#)],
"unauthorized",
)
.into_response(),
}
}
}),
)
}
async fn spawn_protected_server(accept: &str, hits: Arc<AtomicUsize>) -> Url {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let router = protected_router(accept.to_string(), hits);
tokio::spawn(async move { axum::serve(listener, router).await.unwrap() });
Url::parse(&format!("http://{addr}")).unwrap()
}
async fn spawn_port_token_server(hits: Arc<AtomicUsize>) -> Url {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let router = protected_router(format!("token-{}", addr.port()), hits);
tokio::spawn(async move { axum::serve(listener, router).await.unwrap() });
Url::parse(&format!("http://{addr}")).unwrap()
}
fn client_with(
middleware: AuthChallengeMiddleware,
) -> reqwest_middleware::ClientWithMiddleware {
ClientBuilder::new(reqwest::Client::new())
.with_arc(Arc::new(middleware))
.build()
}
#[tokio::test]
async fn challenge_triggers_mint_and_replay() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let response = client
.get(server_url.join("/channel/repodata.json").unwrap())
.send()
.await
.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
assert_eq!(hits.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn second_request_reuses_cached_token_without_challenge() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 200);
assert_eq!(client.get(url).send().await.unwrap().status(), 200);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
assert_eq!(hits.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn inapplicable_flow_is_negative_cached() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(None);
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 401);
assert_eq!(client.get(url).send().await.unwrap().status(), 401);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
assert_eq!(hits.load(Ordering::SeqCst), 2);
}
fn header_map(values: &[&str]) -> http::HeaderMap {
let mut headers = http::HeaderMap::new();
for v in values {
headers.append(
http::header::WWW_AUTHENTICATE,
http::HeaderValue::from_str(v).unwrap(),
);
}
headers
}
#[test]
fn parses_single_bearer_challenge() {
let challenges = parse_challenges(&header_map(&[r#"Bearer realm="prefix.dev""#]));
assert_eq!(challenges.len(), 1);
assert_eq!(challenges[0].scheme, "Bearer");
assert_eq!(challenges[0].params["realm"], "prefix.dev");
}
#[test]
fn parses_multiple_challenges_in_one_header() {
let challenges = parse_challenges(&header_map(&[
r#"Bearer realm="prefix.dev", error="invalid_token", Basic realm="other""#,
]));
assert_eq!(challenges.len(), 2);
assert_eq!(challenges[0].scheme, "Bearer");
assert_eq!(challenges[0].params["realm"], "prefix.dev");
assert_eq!(challenges[0].params["error"], "invalid_token");
assert_eq!(challenges[1].scheme, "Basic");
assert_eq!(challenges[1].params["realm"], "other");
}
#[test]
fn parses_multiple_headers() {
let challenges =
parse_challenges(&header_map(&[r#"Bearer realm="a""#, r#"Basic realm="b""#]));
assert_eq!(challenges.len(), 2);
assert_eq!(challenges[0].scheme, "Bearer");
assert_eq!(challenges[1].scheme, "Basic");
}
#[test]
fn quoted_commas_do_not_split_challenges() {
let challenges = parse_challenges(&header_map(&[r#"Bearer realm="a,b""#]));
assert_eq!(challenges.len(), 1);
assert_eq!(challenges[0].params["realm"], "a,b");
}
#[test]
fn unquoted_params_and_case_insensitive_keys() {
let challenges = parse_challenges(&header_map(&["Bearer REALM=prefix.dev"]));
assert_eq!(challenges.len(), 1);
assert_eq!(challenges[0].params["realm"], "prefix.dev");
}
#[test]
fn garbage_yields_no_challenges_and_no_panic() {
assert!(parse_challenges(&header_map(&["= = ="])).is_empty());
assert!(parse_challenges(&header_map(&[",,,"])).is_empty());
assert!(parse_challenges(&header_map(&[""])).is_empty());
assert!(parse_challenges(&header_map(&["%%% ###"])).is_empty());
assert!(parse_challenges(&http::HeaderMap::new()).is_empty());
}
#[test]
fn unparsable_header_value_does_not_hide_others() {
let challenges = parse_challenges(&header_map(&["%%% ###", r#"Bearer realm="x""#]));
assert_eq!(challenges.len(), 1);
assert_eq!(challenges[0].scheme, "Bearer");
}
#[test]
fn bearer_token_debug_is_redacted() {
let token = BearerToken::new("supersecret".to_string());
let formatted = format!("{token:?}");
assert!(!formatted.contains("supersecret"));
assert!(formatted.contains("redacted"));
}
fn unsigned_jwt_with_exp(exp: u64) -> String {
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"none","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(format!(r#"{{"exp":{exp}}}"#));
format!("{header}.{payload}.")
}
#[test]
fn jwt_expiration_reads_exp_claim() {
let token = unsigned_jwt_with_exp(1_700_000_000);
assert_eq!(
jwt_expiration(&token),
UNIX_EPOCH.checked_add(Duration::from_secs(1_700_000_000))
);
}
#[test]
fn opaque_token_has_no_expiration() {
assert_eq!(jwt_expiration("not-a-jwt"), None);
}
#[test]
fn cached_jwt_is_stale_inside_refresh_margin() {
let token = BearerToken::new(unsigned_jwt_with_exp(1_700_000_000));
let cached = CachedToken::new(&token).unwrap();
let now = UNIX_EPOCH + Duration::from_secs(1_700_000_000 - 30);
assert!(!cached.is_fresh(now));
let earlier = UNIX_EPOCH + Duration::from_secs(1_700_000_000 - 3600);
assert!(cached.is_fresh(earlier));
}
#[tokio::test]
async fn header_invalid_token_is_not_cached_and_does_not_error() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(Some("bad\ntoken"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 401);
assert_eq!(client.get(url).send().await.unwrap().status(), 401);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
}
#[derive(Debug)]
struct SequenceFlow {
tokens: Mutex<Vec<&'static str>>,
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl AuthFlow for SequenceFlow {
async fn acquire_token(
&self,
_url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
let mut tokens = self.tokens.lock().unwrap();
assert!(
!tokens.is_empty(),
"SequenceFlow exhausted: middleware called acquire_token more times than expected"
);
let token = tokens.remove(0);
Ok(Some(BearerToken::new(token.to_string())))
}
}
#[derive(Debug)]
struct FailingFlow {
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl AuthFlow for FailingFlow {
async fn acquire_token(
&self,
_url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Err(AuthFlowError::new(std::io::Error::other("mint exploded")))
}
}
#[tokio::test]
async fn existing_authorization_header_is_respected() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let response = client
.get(server_url.join("/channel/repodata.json").unwrap())
.header(reqwest::header::AUTHORIZATION, "Bearer user-supplied")
.send()
.await
.unwrap();
assert_eq!(response.status(), 401);
assert_eq!(flow.calls.load(Ordering::SeqCst), 0);
assert_eq!(hits.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn replays_at_most_once() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("never-issued", hits.clone()).await;
let flow = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let response = client
.get(server_url.join("/channel/repodata.json").unwrap())
.send()
.await
.unwrap();
assert_eq!(response.status(), 401);
assert_eq!(hits.load(Ordering::SeqCst), 2);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn stale_cached_token_is_cleared_and_reacquired() {
use axum::{http::StatusCode, response::IntoResponse, routing::get};
let seen: Arc<Mutex<Vec<Option<String>>>> = Arc::new(Mutex::new(Vec::new()));
let seen_in_handler = seen.clone();
let router = axum::Router::new().route(
"/channel/repodata.json",
get(move |headers: axum::http::HeaderMap| {
let seen = seen_in_handler.clone();
async move {
let auth = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.map(str::to_string);
seen.lock().unwrap().push(auth.clone());
if auth.as_deref() == Some("Bearer fresh") {
(StatusCode::OK, "ok").into_response()
} else {
(
StatusCode::UNAUTHORIZED,
[("www-authenticate", r#"Bearer realm="test""#)],
"unauthorized",
)
.into_response()
}
}
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move { axum::serve(listener, router).await.unwrap() });
let server_url = Url::parse(&format!("http://{addr}")).unwrap();
let flow = Arc::new(SequenceFlow {
tokens: Mutex::new(vec!["old", "fresh"]),
calls: AtomicUsize::new(0),
});
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 401);
assert_eq!(client.get(url).send().await.unwrap().status(), 200);
assert_eq!(flow.calls.load(Ordering::SeqCst), 2);
assert_eq!(
*seen.lock().unwrap(),
vec![
None,
Some("Bearer old".to_string()),
Some("Bearer old".to_string()),
Some("Bearer fresh".to_string()),
]
);
}
#[tokio::test]
async fn flow_error_is_swallowed_and_negative_cached() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = Arc::new(FailingFlow {
calls: AtomicUsize::new(0),
});
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 401);
assert_eq!(client.get(url).send().await.unwrap().status(), 401);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
assert_eq!(hits.load(Ordering::SeqCst), 2);
}
#[tokio::test]
#[tracing_test::traced_test]
async fn unclonable_challenged_request_is_returned_unreplayed_with_warning() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let body = reqwest::Body::wrap_stream(futures::stream::once(async {
Ok::<_, std::io::Error>(b"x".to_vec())
}));
let response = client
.get(server_url.join("/channel/repodata.json").unwrap())
.body(body)
.send()
.await
.unwrap();
assert_eq!(response.status(), 401);
assert_eq!(hits.load(Ordering::SeqCst), 1);
assert_eq!(flow.calls.load(Ordering::SeqCst), 0);
assert!(logs_contain("could not be cloned for replay"));
}
#[derive(Debug)]
struct PortTokenFlow {
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl AuthFlow for PortTokenFlow {
async fn acquire_token(
&self,
url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
let port = url.port().expect("test URLs always carry a port");
Ok(Some(BearerToken::new(format!("token-{port}"))))
}
}
#[tokio::test]
async fn tokens_are_cached_per_origin() {
let hits_a = Arc::new(AtomicUsize::new(0));
let hits_b = Arc::new(AtomicUsize::new(0));
let url_a = spawn_port_token_server(hits_a.clone()).await;
let url_b = spawn_port_token_server(hits_b.clone()).await;
let flow = Arc::new(PortTokenFlow {
calls: AtomicUsize::new(0),
});
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let a = url_a.join("/channel/repodata.json").unwrap();
let b = url_b.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(a.clone()).send().await.unwrap().status(), 200);
assert_eq!(client.get(b).send().await.unwrap().status(), 200);
assert_eq!(client.get(a).send().await.unwrap().status(), 200);
assert_eq!(flow.calls.load(Ordering::SeqCst), 2); assert_eq!(hits_a.load(Ordering::SeqCst), 3); assert_eq!(hits_b.load(Ordering::SeqCst), 2); }
#[tokio::test]
async fn flows_are_consulted_in_order_until_one_yields() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let inapplicable = StaticFlow::new(None);
let minting = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![
inapplicable.clone(),
minting.clone(),
]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url.clone()).send().await.unwrap().status(), 200);
assert_eq!(client.get(url).send().await.unwrap().status(), 200);
assert_eq!(inapplicable.calls.load(Ordering::SeqCst), 1);
assert_eq!(minting.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn failing_flow_falls_through_to_next() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let failing = Arc::new(FailingFlow {
calls: AtomicUsize::new(0),
});
let minting = StaticFlow::new(Some("abc123"));
let client = client_with(AuthChallengeMiddleware::new(vec![
failing.clone(),
minting.clone(),
]));
let url = server_url.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(url).send().await.unwrap().status(), 200);
assert_eq!(failing.calls.load(Ordering::SeqCst), 1);
assert_eq!(minting.calls.load(Ordering::SeqCst), 1);
}
#[derive(Debug)]
struct SinglePortFlow {
port: u16,
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl AuthFlow for SinglePortFlow {
async fn acquire_token(
&self,
url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
if url.port() == Some(self.port) {
Ok(Some(BearerToken::new(format!("token-{}", self.port))))
} else {
Ok(None)
}
}
}
#[tokio::test]
async fn negative_cache_is_scoped_per_origin() {
let hits_a = Arc::new(AtomicUsize::new(0));
let hits_b = Arc::new(AtomicUsize::new(0));
let url_a = spawn_port_token_server(hits_a.clone()).await;
let url_b = spawn_port_token_server(hits_b.clone()).await;
let flow = Arc::new(SinglePortFlow {
port: url_b.port().unwrap(),
calls: AtomicUsize::new(0),
});
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let a = url_a.join("/channel/repodata.json").unwrap();
let b = url_b.join("/channel/repodata.json").unwrap();
assert_eq!(client.get(a.clone()).send().await.unwrap().status(), 401);
assert_eq!(client.get(a).send().await.unwrap().status(), 401);
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
assert_eq!(client.get(b).send().await.unwrap().status(), 200);
assert_eq!(flow.calls.load(Ordering::SeqCst), 2);
}
#[derive(Debug)]
struct SlowFlow {
token: &'static str,
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl AuthFlow for SlowFlow {
async fn acquire_token(
&self,
_url: &Url,
_challenges: &[Challenge],
) -> Result<Option<BearerToken>, AuthFlowError> {
self.calls.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(300)).await;
Ok(Some(BearerToken::new(self.token.to_string())))
}
}
#[tokio::test]
async fn concurrent_challenges_mint_once() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let flow = Arc::new(SlowFlow {
token: "abc123",
calls: AtomicUsize::new(0),
});
let client = client_with(AuthChallengeMiddleware::new(vec![flow.clone()]));
let url = server_url.join("/channel/repodata.json").unwrap();
let responses =
futures::future::join_all((0..5).map(|_| client.get(url.clone()).send())).await;
for response in responses {
assert_eq!(response.unwrap().status(), 200);
}
assert_eq!(flow.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn default_middleware_is_inert_for_untrusted_origins() {
let hits = Arc::new(AtomicUsize::new(0));
let server_url = spawn_protected_server("abc123", hits.clone()).await;
let client = client_with(AuthChallengeMiddleware::default());
let response = client
.get(server_url.join("/channel/repodata.json").unwrap())
.send()
.await
.unwrap();
assert_eq!(response.status(), 401);
assert_eq!(hits.load(Ordering::SeqCst), 1);
}
}