use std::fmt;
use subtle::{Choice, ConditionallySelectable, ConstantTimeEq};
use zeroize::Zeroizing;
use crate::token::{AuthorizedToken, InvalidToken, InvalidTokenKind, TokenRejection};
use crate::validator::{CachedAttempt, OAuthValidator};
#[cfg(test)]
thread_local! {
pub(crate) static STATIC_COMPARISONS: std::cell::Cell<usize> =
const { std::cell::Cell::new(0) };
}
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum Credential {
StaticToken,
OAuth(AuthorizedToken),
}
#[derive(Clone, Default)]
pub struct StaticTokens {
entries: Vec<StaticEntry>,
}
#[derive(Clone)]
struct StaticEntry {
label: Option<String>,
secret: Zeroizing<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum StaticTokensError {
#[error("the static token at position {index} is empty or blank")]
#[non_exhaustive]
BlankSecret {
index: usize,
},
#[error("the static token at position {index} repeats the one at position {first}")]
#[non_exhaustive]
DuplicateSecret {
index: usize,
first: usize,
},
#[error(
"the label of the static token at position {index} must be 1 to 64 visible ASCII \
characters (no spaces)"
)]
#[non_exhaustive]
InvalidLabel {
index: usize,
},
#[error(
"the label {label:?} of the static token at position {index} is already used at \
position {first}"
)]
#[non_exhaustive]
DuplicateLabel {
index: usize,
first: usize,
label: String,
},
}
impl StaticTokens {
pub const MAX_LABEL_LEN: usize = 64;
pub fn new() -> Self {
Self::default()
}
pub fn single(secret: impl Into<String>) -> Result<Self, StaticTokensError> {
Self::new().with(None, secret)
}
pub fn with(
mut self,
label: Option<&str>,
secret: impl Into<String>,
) -> Result<Self, StaticTokensError> {
let secret = Zeroizing::new(secret.into());
let index = self.entries.len();
if secret.trim().is_empty() {
return Err(StaticTokensError::BlankSecret { index });
}
if let Some(label) = label {
if !is_log_safe_label(label) {
return Err(StaticTokensError::InvalidLabel { index });
}
if let Some(first) = self
.entries
.iter()
.position(|e| e.label.as_deref() == Some(label))
{
return Err(StaticTokensError::DuplicateLabel {
index,
first,
label: label.to_string(),
});
}
}
if let Some(first) = self.position_of(&secret) {
return Err(StaticTokensError::DuplicateSecret { index, first });
}
self.entries.push(StaticEntry {
label: label.map(str::to_string),
secret,
});
Ok(self)
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn labels(&self) -> impl Iterator<Item = Option<&str>> {
self.entries.iter().map(|e| e.label.as_deref())
}
#[cfg(feature = "tower")]
pub(crate) fn merged(set: Option<Self>, token: Option<Zeroizing<String>>) -> Option<Self> {
let mut merged = set.unwrap_or_default();
if let Some(token) = token.filter(|t| !t.trim().is_empty())
&& merged.position_of(&token).is_none()
{
merged.entries.push(StaticEntry {
label: None,
secret: token,
});
}
(!merged.is_empty()).then_some(merged)
}
#[cfg(feature = "env")]
pub(crate) fn push_checked(&mut self, label: &str, secret: Zeroizing<String>) {
debug_assert!(is_log_safe_label(label) && !secret.trim().is_empty());
self.entries.push(StaticEntry {
label: Some(label.to_string()),
secret,
});
}
#[cfg(feature = "env")]
pub(crate) fn contains(&self, secret: &str) -> bool {
self.position_of(secret).is_some()
}
fn position_of(&self, secret: &str) -> Option<usize> {
let secrets: Vec<&str> = self.secrets();
find_static(&[secret], &secrets)
}
fn secrets(&self) -> Vec<&str> {
self.entries.iter().map(|e| e.secret.as_str()).collect()
}
fn match_at(&self, index: usize) -> StaticTokenMatch {
StaticTokenMatch {
label: self.entries.get(index).and_then(|e| e.label.clone()),
}
}
}
impl fmt::Debug for StaticTokens {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StaticTokens")
.field("len", &self.entries.len())
.field("labels", &self.labels().collect::<Vec<_>>())
.field("secrets", &"<redacted>")
.finish()
}
}
fn is_log_safe_label(label: &str) -> bool {
!label.is_empty()
&& label.len() <= StaticTokens::MAX_LABEL_LEN
&& label.bytes().all(|b| b.is_ascii_graphic())
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct StaticTokenMatch {
label: Option<String>,
}
impl StaticTokenMatch {
pub fn label(&self) -> Option<&str> {
self.label.as_deref()
}
#[cfg(feature = "tower")]
pub(crate) fn unlabeled() -> Self {
Self { label: None }
}
}
fn counted_ct_eq(candidate: &str, secret: &str) -> Choice {
#[cfg(test)]
STATIC_COMPARISONS.with(|n| n.set(n.get() + 1));
candidate.as_bytes().ct_eq(secret.as_bytes())
}
fn find_static(candidates: &[&str], secrets: &[&str]) -> Option<usize> {
let mut found = Choice::from(0);
let mut index: u64 = 0;
for candidate in candidates {
for (i, secret) in (0u64..).zip(secrets) {
let equal = counted_ct_eq(candidate, secret);
index.conditional_assign(&i, equal & !found);
found |= equal;
}
}
if bool::from(found) {
usize::try_from(index).ok()
} else {
None
}
}
pub async fn authenticate<'a>(
candidates: impl IntoIterator<Item = &'a str>,
static_token: Option<&str>,
oauth: Option<&OAuthValidator>,
) -> Result<Credential, TokenRejection> {
let one: [&str; 1];
let secrets: &[&str] = match static_token.filter(|t| !t.is_empty()) {
Some(token) => {
one = [token];
&one
}
None => &[],
};
check_candidates(candidates, secrets, oauth)
.await
.map(|(credential, _)| credential)
}
pub async fn authenticate_with_static_tokens<'a>(
candidates: impl IntoIterator<Item = &'a str>,
static_tokens: Option<&StaticTokens>,
oauth: Option<&OAuthValidator>,
) -> Result<(Credential, Option<StaticTokenMatch>), TokenRejection> {
let secrets = static_tokens.map(StaticTokens::secrets).unwrap_or_default();
let (credential, index) = check_candidates(candidates, &secrets, oauth).await?;
let matched = index.zip(static_tokens).map(|(i, set)| set.match_at(i));
Ok((credential, matched))
}
async fn check_candidates<'a>(
candidates: impl IntoIterator<Item = &'a str>,
secrets: &[&str],
oauth: Option<&OAuthValidator>,
) -> Result<(Credential, Option<usize>), TokenRejection> {
let candidates: Vec<&str> = candidates
.into_iter()
.filter(|c| !c.trim().is_empty())
.collect();
if candidates.is_empty() {
return Err(TokenRejection::Missing);
}
if let Some(index) = find_static(&candidates, secrets) {
return Ok((Credential::StaticToken, Some(index)));
}
let Some(validator) = oauth else {
return Err(match secrets.len() {
0 => TokenRejection::invalid(
InvalidTokenKind::NoMechanism,
"no credential mechanism is configured",
),
1 => TokenRejection::invalid(
InvalidTokenKind::StaticTokenMismatch,
"credential does not match the static token",
),
_ => TokenRejection::invalid(
InvalidTokenKind::StaticTokenMismatch,
"credential does not match any static token",
),
});
};
let mut refusals: Vec<Option<TokenRejection>> = Vec::with_capacity(candidates.len());
for candidate in &candidates {
match validator.validate_cached(candidate).await {
CachedAttempt::Decided(Ok(token)) => return Ok((Credential::OAuth(token), None)),
CachedAttempt::Decided(Err(rejection)) => refusals.push(Some(rejection)),
CachedAttempt::NeedsKeyFetch => refusals.push(None),
}
}
for (candidate, refusal) in candidates.iter().zip(refusals.iter_mut()) {
if refusal.is_none() {
match validator.validate(candidate).await {
Ok(token) => return Ok((Credential::OAuth(token), None)),
Err(rejection) => *refusal = Some(rejection),
}
}
}
let mut insufficient_scope = false;
let mut first_reason: Option<InvalidToken> = None;
for refusal in refusals.into_iter().flatten() {
match refusal {
TokenRejection::InsufficientScope => insufficient_scope = true,
TokenRejection::Invalid(reason) => {
first_reason.get_or_insert(reason);
}
TokenRejection::Missing => {
first_reason.get_or_insert_with(|| {
InvalidToken::new(InvalidTokenKind::Other, "no credential presented")
});
}
}
}
if insufficient_scope {
return Err(TokenRejection::InsufficientScope);
}
Err(TokenRejection::Invalid(first_reason.unwrap_or_else(|| {
InvalidToken::new(
InvalidTokenKind::Other,
"no candidate credential was accepted",
)
})))
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::testing;
const STATIC: &str = "static-secret";
fn unreachable_validator() -> Arc<OAuthValidator> {
Arc::new(OAuthValidator::new(&testing::resolved_config("http://127.0.0.1:1/jwks")).unwrap())
}
async fn live_validator() -> (testing::FakeJwksServer, Arc<OAuthValidator>) {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = Arc::new(OAuthValidator::new(&testing::resolved_config(&jwks.url)).unwrap());
(jwks, v)
}
fn unscoped_token() -> String {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE, "sub": "user-2",
"exp": testing::now() + 3600, "scope": "openid profile",
}),
)
}
fn expired_token() -> String {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE,
"exp": testing::now() - 3600, "scope": "mcp:read",
}),
)
}
#[tokio::test]
async fn a_static_match_wins_without_touching_oauth() {
let v = unreachable_validator();
assert_eq!(
authenticate([STATIC], Some(STATIC), Some(&v)).await,
Ok(Credential::StaticToken)
);
assert_eq!(
authenticate([STATIC], Some(STATIC), None).await,
Ok(Credential::StaticToken)
);
}
#[tokio::test]
async fn an_oauth_match_returns_the_authorized_token() {
let (jwks, v) = live_validator().await;
let token = testing::valid_token();
for static_token in [None, Some(STATIC)] {
match authenticate([token.as_str()], static_token, Some(&v)).await {
Ok(Credential::OAuth(t)) => {
assert_eq!(t.subject.as_deref(), Some("user-1"));
assert!(t.has_scope("mcp:read"));
}
other => panic!("expected an OAuth credential, got {other:?}"),
}
}
assert!(jwks.hits.load(std::sync::atomic::Ordering::SeqCst) >= 1);
}
#[tokio::test]
async fn a_valid_token_lacking_scope_is_insufficient_scope() {
let (_jwks, v) = live_validator().await;
let token = unscoped_token();
assert_eq!(
authenticate([token.as_str()], Some(STATIC), Some(&v)).await,
Err(TokenRejection::InsufficientScope)
);
}
#[tokio::test]
async fn no_candidate_is_missing_whatever_is_configured() {
let v = unreachable_validator();
for (static_token, oauth) in [
(Some(STATIC), None),
(None, Some(&*v)),
(Some(STATIC), Some(&*v)),
(None, None),
] {
assert_eq!(
authenticate(std::iter::empty(), static_token, oauth).await,
Err(TokenRejection::Missing),
"static={static_token:?} oauth={}",
oauth.is_some()
);
}
}
#[tokio::test]
async fn blank_candidates_count_as_absent() {
let v = unreachable_validator();
assert_eq!(
authenticate(["", " ", "\t"], Some(STATIC), Some(&v)).await,
Err(TokenRejection::Missing)
);
assert_eq!(
authenticate(["", " ", STATIC], Some(STATIC), Some(&v)).await,
Ok(Credential::StaticToken)
);
}
#[tokio::test]
async fn a_blank_static_token_never_matches() {
assert_eq!(
authenticate([""], Some(""), None).await,
Err(TokenRejection::Missing)
);
crate::token::assert_invalid(
authenticate(["x"], Some(""), None).await,
InvalidTokenKind::NoMechanism,
"no credential mechanism is configured",
"",
);
}
#[tokio::test]
async fn an_invalid_credential_carries_the_first_reason() {
let (_jwks, v) = live_validator().await;
let expired = expired_token();
match authenticate([expired.as_str(), "not-a-jwt"], Some(STATIC), Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.starts_with("token rejected:"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
match authenticate(["not-a-jwt", expired.as_str()], Some(STATIC), Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.starts_with("credential is not a JWT"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
}
#[tokio::test]
async fn static_only_refuses_a_jwt_without_validating_it() {
let token = testing::valid_token();
crate::token::assert_invalid(
authenticate([token.as_str()], Some(STATIC), None).await,
InvalidTokenKind::StaticTokenMismatch,
"credential does not match the static token",
"",
);
}
#[tokio::test]
async fn oauth_only_refuses_the_static_value() {
let v = unreachable_validator();
match authenticate([STATIC], None, Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.starts_with("credential is not a JWT"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
}
#[tokio::test]
async fn neither_mechanism_configured_accepts_nothing() {
crate::token::assert_invalid(
authenticate([STATIC], None, None).await,
InvalidTokenKind::NoMechanism,
"no credential mechanism is configured",
"",
);
}
async fn eager_refetch_validator() -> (testing::FakeJwksServer, OAuthValidator) {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = OAuthValidator::build(
&testing::resolved_config(&jwks.url),
std::time::Duration::ZERO,
)
.unwrap();
(jwks, v)
}
fn foreign_token() -> String {
testing::mint(
testing::KEY_B_PEM,
"proxy-key",
&serde_json::json!({
"iss": "https://proxy.example.test/", "aud": "proxy",
"exp": testing::now() + 3600,
}),
)
}
fn hits(jwks: &testing::FakeJwksServer) -> usize {
jwks.hits.load(std::sync::atomic::Ordering::SeqCst)
}
#[tokio::test]
async fn a_foreign_kid_does_not_trigger_a_refetch_when_another_candidate_is_cached() {
let (jwks, v) = eager_refetch_validator().await;
let valid = testing::valid_token();
let foreign = foreign_token();
assert!(v.validate(&valid).await.is_ok());
assert_eq!(hits(&jwks), 1);
for candidates in [
[foreign.as_str(), valid.as_str()],
[valid.as_str(), foreign.as_str()],
] {
assert!(matches!(
authenticate(candidates, None, Some(&v)).await,
Ok(Credential::OAuth(_))
));
}
assert_eq!(hits(&jwks), 1, "no refetch for the foreign kid");
match authenticate([foreign.as_str(), "garbage"], None, Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.contains("proxy-key"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
assert_eq!(hits(&jwks), 2);
}
#[tokio::test]
async fn a_cold_cache_still_fetches_for_the_only_candidate() {
let (jwks, v) = eager_refetch_validator().await;
let valid = testing::valid_token();
assert_eq!(hits(&jwks), 0);
assert!(matches!(
authenticate([valid.as_str()], None, Some(&v)).await,
Ok(Credential::OAuth(_))
));
assert_eq!(hits(&jwks), 1);
let (_jwks, v) = eager_refetch_validator().await;
match authenticate([foreign_token().as_str(), "not-a-jwt"], None, Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.contains("proxy-key"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
}
#[tokio::test]
async fn mixed_candidates_any_success_wins_in_either_order() {
let (_jwks, v) = live_validator().await;
let valid = testing::valid_token();
let unscoped = unscoped_token();
let expired = expired_token();
for candidates in [
vec![expired.as_str(), STATIC],
vec![STATIC, expired.as_str()],
vec!["garbage", STATIC],
vec![unscoped.as_str(), STATIC],
] {
assert_eq!(
authenticate(candidates.iter().copied(), Some(STATIC), Some(&v)).await,
Ok(Credential::StaticToken),
"{candidates:.30?}"
);
}
for candidates in [
vec!["wrong-static", valid.as_str()],
vec![valid.as_str(), "wrong-static"],
vec![unscoped.as_str(), valid.as_str()],
vec![expired.as_str(), valid.as_str()],
] {
assert!(
matches!(
authenticate(candidates.iter().copied(), Some(STATIC), Some(&v)).await,
Ok(Credential::OAuth(_))
),
"{candidates:.30?}"
);
}
for candidates in [
vec![unscoped.as_str(), "garbage"],
vec!["garbage", unscoped.as_str()],
vec![expired.as_str(), unscoped.as_str()],
] {
assert_eq!(
authenticate(candidates.iter().copied(), Some(STATIC), Some(&v)).await,
Err(TokenRejection::InsufficientScope),
"{candidates:.30?}"
);
}
}
fn rotation() -> StaticTokens {
StaticTokens::new()
.with(Some("current"), "key-current")
.and_then(|t| t.with(Some("next"), "key-next"))
.and_then(|t| t.with(None, "key-unlabeled"))
.unwrap()
}
fn label_of(
result: Result<(Credential, Option<StaticTokenMatch>), TokenRejection>,
) -> Option<String> {
match result {
Ok((Credential::StaticToken, Some(m))) => m.label().map(str::to_string),
other => panic!("expected a static match, got {other:?}"),
}
}
#[tokio::test]
async fn a_one_entry_set_is_identical_to_the_single_token() {
let (_jwks, live) = live_validator().await;
let valid = testing::valid_token();
let unscoped = unscoped_token();
let expired = expired_token();
let single = StaticTokens::single(STATIC).unwrap();
let candidate_sets: Vec<Vec<&str>> = vec![
vec![],
vec!["", " "],
vec![STATIC],
vec!["wrong"],
vec!["wrong", STATIC],
vec![valid.as_str()],
vec![unscoped.as_str()],
vec![expired.as_str(), "garbage"],
vec!["static-secret "],
];
for oauth in [None, Some(&*live)] {
for candidates in &candidate_sets {
let old = authenticate(candidates.iter().copied(), Some(STATIC), oauth).await;
let new = authenticate_with_static_tokens(
candidates.iter().copied(),
Some(&single),
oauth,
)
.await;
match (&old, &new) {
(Ok(Credential::StaticToken), Ok((Credential::StaticToken, Some(m)))) => {
assert_eq!(m.label(), None);
}
(Ok(Credential::OAuth(a)), Ok((Credential::OAuth(b), None))) => {
assert_eq!(a.subject, b.subject);
}
(Err(a), Err(b)) => assert_eq!(a, b, "{candidates:.30?}"),
_ => panic!("{candidates:.30?}: {old:?} vs {new:?}"),
}
}
}
for set in [None, Some(&StaticTokens::new())] {
crate::token::assert_invalid(
authenticate_with_static_tokens(["x"], set, None).await,
InvalidTokenKind::NoMechanism,
"no credential mechanism is configured",
"",
);
}
}
#[tokio::test]
async fn every_entry_is_accepted_and_named() {
let set = rotation();
let v = unreachable_validator();
for (secret, label) in [
("key-current", Some("current")),
("key-next", Some("next")),
("key-unlabeled", None),
] {
for oauth in [None, Some(&*v)] {
let result = authenticate_with_static_tokens([secret], Some(&set), oauth).await;
assert_eq!(label_of(result).as_deref(), label, "{secret}");
}
let result = authenticate_with_static_tokens(["junk", secret], Some(&set), None).await;
assert_eq!(label_of(result).as_deref(), label);
}
let result =
authenticate_with_static_tokens(["key-next", "key-current"], Some(&set), None).await;
assert_eq!(label_of(result).as_deref(), Some("next"));
}
#[tokio::test]
async fn a_wrong_token_is_refused() {
let set = rotation();
for candidate in ["key-", "key-current ", "KEY-CURRENT", "key-nextx", "other"] {
crate::token::assert_invalid(
authenticate_with_static_tokens([candidate], Some(&set), None).await,
InvalidTokenKind::StaticTokenMismatch,
"credential does not match any static token",
candidate,
);
}
assert_eq!(
authenticate_with_static_tokens([" "], Some(&set), None).await,
Err(TokenRejection::Missing)
);
let v = unreachable_validator();
match authenticate_with_static_tokens(["other"], Some(&set), Some(&v)).await {
Err(TokenRejection::Invalid(reason)) => {
assert!(reason.starts_with("credential is not a JWT"), "{reason}");
}
other => panic!("expected Invalid, got {other:?}"),
}
}
#[test]
fn a_blank_secret_is_refused_at_construction() {
for blank in ["", " ", "\t\n", " "] {
assert_eq!(
StaticTokens::single(blank).unwrap_err(),
StaticTokensError::BlankSecret { index: 0 }
);
}
assert_eq!(
StaticTokens::single("a")
.unwrap()
.with(Some("x"), " ")
.unwrap_err(),
StaticTokensError::BlankSecret { index: 1 }
);
}
#[test]
fn duplicates_and_bad_labels_are_refused() {
let base = || StaticTokens::single("one").unwrap();
assert_eq!(
base().with(Some("again"), "one").unwrap_err(),
StaticTokensError::DuplicateSecret { index: 1, first: 0 }
);
let labeled = StaticTokens::new().with(Some("a"), "one").unwrap();
assert_eq!(
labeled.clone().with(Some("a"), "two").unwrap_err(),
StaticTokensError::DuplicateLabel {
index: 1,
first: 0,
label: "a".into()
}
);
let ok = labeled
.with(None, "two")
.and_then(|t| t.with(None, "three"))
.and_then(|t| t.with(Some("b"), "four"))
.unwrap();
assert_eq!(ok.len(), 4);
assert!(!ok.is_empty() && StaticTokens::new().is_empty());
assert_eq!(
ok.labels().collect::<Vec<_>>(),
[Some("a"), None, None, Some("b")]
);
let longest = "l".repeat(StaticTokens::MAX_LABEL_LEN);
assert!(base().with(Some(&longest), "two").is_ok());
assert!(base().with(Some("!~"), "two").is_ok());
let too_long = "l".repeat(StaticTokens::MAX_LABEL_LEN + 1);
for bad in [
"",
"has space",
"tab\t",
"new\nline",
"caf\u{e9}",
"\u{1b}[31m",
&too_long,
] {
assert_eq!(
base().with(Some(bad), "two").unwrap_err(),
StaticTokensError::InvalidLabel { index: 1 },
"{bad:?}"
);
}
let err = StaticTokens::single("hunter2")
.and_then(|t| t.with(Some("again"), "hunter2"))
.unwrap_err();
assert!(!err.to_string().contains("hunter2") && !format!("{err:?}").contains("hunter2"));
}
#[test]
fn debug_never_prints_a_secret() {
let set = StaticTokens::new()
.with(Some("current"), "hunter2-current")
.and_then(|t| t.with(None, "hunter2-other"))
.unwrap();
let rendered = format!("{set:?} {set:#?}");
assert!(!rendered.contains("hunter2"), "{rendered}");
assert!(rendered.contains("current") && rendered.contains("<redacted>"));
assert!(rendered.contains("len: 2"), "{rendered}");
let matched = StaticTokenMatch {
label: Some("current".into()),
};
assert!(!format!("{matched:?}").contains("hunter2"));
}
fn comparisons() -> usize {
STATIC_COMPARISONS.with(std::cell::Cell::get)
}
fn reset_comparisons() {
STATIC_COMPARISONS.with(|n| n.set(0));
}
#[tokio::test(flavor = "current_thread")]
async fn every_entry_is_compared_even_when_the_first_matches() {
let set = rotation();
for (candidates, matched) in [
(vec!["key-current"], Some(Some("current"))),
(vec!["key-unlabeled"], Some(None)),
(vec!["nothing"], None),
(vec!["key-current", "junk"], Some(Some("current"))),
(vec!["junk", "key-next"], Some(Some("next"))),
] {
reset_comparisons();
let result =
authenticate_with_static_tokens(candidates.iter().copied(), Some(&set), None).await;
assert_eq!(
comparisons(),
candidates.len() * set.len(),
"{candidates:?}"
);
match matched {
Some(label) => assert_eq!(label_of(result).as_deref(), label),
None => assert!(result.is_err()),
}
}
reset_comparisons();
assert!(
authenticate([STATIC, "junk"], Some(STATIC), None)
.await
.is_ok()
);
assert_eq!(comparisons(), 2);
}
#[test]
fn find_static_picks_the_first_matching_candidate() {
let secrets = ["a1", "b22", "c333"];
assert_eq!(find_static(&["c333"], &secrets), Some(2));
assert_eq!(find_static(&["a1"], &secrets), Some(0));
assert_eq!(find_static(&["b22", "a1"], &secrets), Some(1));
assert_eq!(find_static(&["zz"], &secrets), None);
assert_eq!(find_static(&["a1"], &[]), None);
}
}