use std::collections::HashSet;
use serde_json::{Map, Value};
use crate::config::KeyNamingBuf;
pub(crate) const MAX_TOKEN_BYTES: usize = 16 * 1024;
pub(crate) const MAX_LOGGED_CHARS: usize = 128;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct AuthorizedToken {
pub subject: Option<String>,
pub principal: Option<String>,
pub scopes: Vec<String>,
}
impl AuthorizedToken {
pub fn new(
subject: Option<String>,
principal: Option<String>,
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Self {
let mut deduped: Vec<String> = Vec::new();
for scope in scopes {
let scope = scope.into();
let scope = scope.trim();
if !scope.is_empty() && !deduped.iter().any(|s| s == scope) {
deduped.push(scope.to_string());
}
}
Self {
subject,
principal,
scopes: deduped,
}
}
pub fn has_scope(&self, scope: &str) -> bool {
self.scopes.iter().any(|s| s == scope)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum TokenRejection {
#[error("missing credential")]
Missing,
#[error("invalid token")]
Invalid(String),
#[error("insufficient scope")]
InsufficientScope,
}
pub(crate) fn extract_scopes(claims: &Map<String, Value>, claim_names: &[String]) -> Vec<String> {
let mut seen = HashSet::new();
let mut out = Vec::new();
let mut push = |s: &str| {
let s = s.trim();
if !s.is_empty() && seen.insert(s.to_string()) {
out.push(s.to_string());
}
};
for name in claim_names {
match claims.get(name) {
Some(Value::String(s)) => s.split_whitespace().for_each(&mut push),
Some(Value::Array(items)) => items.iter().filter_map(Value::as_str).for_each(&mut push),
_ => {}
}
}
out
}
pub(crate) fn extract_principal(
claims: &Map<String, Value>,
claim_names: &[String],
) -> Option<String> {
claim_names.iter().find_map(|name| match claims.get(name) {
Some(Value::String(s)) if !s.trim().is_empty() => Some(s.clone()),
_ => None,
})
}
pub(crate) fn check_typ(
typ: Option<&str>,
require_at_jwt: bool,
naming: &KeyNamingBuf,
) -> Result<(), TokenRejection> {
let Some(raw) = typ else {
return if require_at_jwt {
Err(TokenRejection::Invalid(format!(
"token header has no typ and {} is on",
naming.key("require_at_jwt")
)))
} else {
Ok(())
};
};
let lower = raw.trim().to_ascii_lowercase();
let media = lower.strip_prefix("application/").unwrap_or(&lower);
match media {
"at+jwt" => Ok(()),
"jwt" if !require_at_jwt => Ok(()),
_ => Err(TokenRejection::Invalid(format!(
"token typ {:?} is not accepted as an access token{}",
for_log(raw),
if require_at_jwt {
format!(" ({} is on)", naming.key("require_at_jwt"))
} else {
String::new()
}
))),
}
}
pub(crate) fn for_log(s: &str) -> String {
let mut out: String = s.chars().take(MAX_LOGGED_CHARS).collect();
if s.chars().count() > MAX_LOGGED_CHARS {
out.push('…');
}
out
}
pub(crate) fn describe_kid(kid: Option<&str>) -> String {
match kid {
Some(kid) => format!("{:?}", for_log(kid)),
None => "(none)".to_string(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn authorized_token_new_dedupes_scopes_like_a_validation() {
let t = AuthorizedToken::new(Some("sub-1".into()), None, ["b", " a", "b", "", "a", "c"]);
assert_eq!(t.subject.as_deref(), Some("sub-1"));
assert_eq!(t.principal, None);
assert_eq!(t.scopes, ["b", "a", "c"]);
assert!(t.has_scope("a"));
assert!(!t.has_scope(""));
let none = AuthorizedToken::new(None, None, Vec::<String>::new());
assert!(none.scopes.is_empty());
}
#[test]
fn a_rejection_displays_its_category_but_never_the_reason() {
let secret_reason = "token rejected: InvalidAudience";
let invalid = TokenRejection::Invalid(secret_reason.into());
assert_eq!(invalid.to_string(), "invalid token");
assert!(!invalid.to_string().contains("InvalidAudience"));
assert!(format!("{invalid:?}").contains(secret_reason));
assert_eq!(TokenRejection::Missing.to_string(), "missing credential");
assert_eq!(
TokenRejection::InsufficientScope.to_string(),
"insufficient scope"
);
let boxed: Box<dyn std::error::Error + Send + Sync> = Box::new(invalid);
assert_eq!(boxed.to_string(), "invalid token");
}
#[test]
fn logged_values_are_truncated() {
let long = "x".repeat(MAX_LOGGED_CHARS * 3);
assert_eq!(for_log(&long).chars().count(), MAX_LOGGED_CHARS + 1);
assert_eq!(for_log("short"), "short");
}
#[test]
fn typ_rejection_reasons_name_the_setting_per_key_naming() {
let dotted = KeyNamingBuf::Dotted("mcp.oauth".into());
assert_eq!(
check_typ(None, true, &dotted),
Err(TokenRejection::Invalid(
"token header has no typ and mcp.oauth.require_at_jwt is on".into()
))
);
assert_eq!(
check_typ(Some("JWT"), true, &dotted),
Err(TokenRejection::Invalid(
"token typ \"JWT\" is not accepted as an access token \
(mcp.oauth.require_at_jwt is on)"
.into()
))
);
assert_eq!(
check_typ(Some("dpop+jwt"), false, &dotted),
Err(TokenRejection::Invalid(
"token typ \"dpop+jwt\" is not accepted as an access token".into()
))
);
let env = KeyNamingBuf::Env("APP_OAUTH_".into());
assert_eq!(
check_typ(None, true, &env),
Err(TokenRejection::Invalid(
"token header has no typ and APP_OAUTH_REQUIRE_AT_JWT is on".into()
))
);
}
}