use subtle::ConstantTimeEq;
use crate::token::{AuthorizedToken, TokenRejection};
use crate::validator::{CachedAttempt, OAuthValidator};
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum Credential {
StaticToken,
OAuth(AuthorizedToken),
}
pub async fn authenticate<'a>(
candidates: impl IntoIterator<Item = &'a str>,
static_token: Option<&str>,
oauth: Option<&OAuthValidator>,
) -> Result<Credential, TokenRejection> {
let candidates: Vec<&str> = candidates
.into_iter()
.filter(|c| !c.trim().is_empty())
.collect();
if candidates.is_empty() {
return Err(TokenRejection::Missing);
}
let static_token = static_token.filter(|t| !t.is_empty());
if let Some(expected) = static_token {
for candidate in &candidates {
if bool::from(candidate.as_bytes().ct_eq(expected.as_bytes())) {
return Ok(Credential::StaticToken);
}
}
}
let Some(validator) = oauth else {
return Err(TokenRejection::Invalid(
if static_token.is_some() {
"credential does not match the static token"
} else {
"no credential mechanism is configured"
}
.to_string(),
));
};
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)),
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)),
Err(rejection) => *refusal = Some(rejection),
}
}
}
let mut insufficient_scope = false;
let mut first_reason: Option<String> = 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(|| "no credential presented".to_string());
}
}
}
if insufficient_scope {
return Err(TokenRejection::InsufficientScope);
}
Err(TokenRejection::Invalid(first_reason.unwrap_or_else(|| {
"no candidate credential was accepted".to_string()
})))
}
#[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)
);
assert_eq!(
authenticate(["x"], Some(""), None).await,
Err(TokenRejection::Invalid(
"no credential mechanism is configured".into()
))
);
}
#[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();
assert_eq!(
authenticate([token.as_str()], Some(STATIC), None).await,
Err(TokenRejection::Invalid(
"credential does not match the static token".into()
))
);
}
#[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() {
assert_eq!(
authenticate([STATIC], None, None).await,
Err(TokenRejection::Invalid(
"no credential mechanism is configured".into()
))
);
}
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?}"
);
}
}
}