use crate::common::issuer_utils::IssClaim;
use crate::jwt::log_entry::JwtLogEntry;
use crate::jwt::{IssuerConfig, StatusListCache};
use crate::log::Logger;
use crate::{JwtConfig, LogLevel, LogWriter};
use super::JwtValidator;
use arc_swap::ArcSwap;
use jsonwebtoken::Algorithm;
use std::borrow::Cow;
use std::collections::HashMap;
use std::fmt::Display;
use std::hash::{DefaultHasher, Hash, Hasher};
use std::sync::Arc;
type CachedValidator = Arc<JwtValidator>;
pub(crate) struct JwtValidatorCache {
validators: ArcSwap<HashMap<ValidatorKeyHash, Vec<(OwnedValidatorInfo, CachedValidator)>>>,
}
impl Default for JwtValidatorCache {
fn default() -> Self {
Self {
validators: ArcSwap::from_pointee(HashMap::new()),
}
}
}
impl JwtValidatorCache {
pub(crate) fn init_for_iss(
&self,
iss_config: &IssuerConfig,
jwt_config: &JwtConfig,
status_lists: &StatusListCache,
logger: Option<&Logger>,
) {
let iss = iss_config
.openid_config
.as_ref()
.map_or_else(|| iss_config.policy.iss_claim(), |oidc| oidc.issuer.clone());
for (token_name, tkn_metadata) in &iss_config.policy.token_metadata {
if !tkn_metadata.trusted {
logger.log_any(JwtLogEntry::new(
format!(
"skipping metadata for '{}' from '{}' since `trusted == false`",
token_name, iss_config.issuer_id,
),
Some(LogLevel::INFO),
));
continue;
}
logger.log_any(JwtLogEntry::new(
format!(
"creating validators for token: {token_name} under issuer: {iss} with algorithms: {:?}",
jwt_config.signature_algorithms_supported
),
Some(LogLevel::DEBUG),
));
for algorithm in jwt_config.signature_algorithms_supported.iter().copied() {
let (validator, key) = JwtValidator::new_input_tkn_validator(
Some(&iss),
token_name,
tkn_metadata,
algorithm,
status_lists.clone(),
jwt_config.jwt_sig_validation,
jwt_config.jwt_status_validation,
);
self.insert(&key, validator);
}
for algorithm in jwt_config.signature_algorithms_supported.iter().copied() {
let (validator, key) = JwtValidator::new_multi_issuer_tkn_validator(
Some(&iss),
token_name,
tkn_metadata,
algorithm,
status_lists.clone(),
jwt_config.jwt_sig_validation,
jwt_config.jwt_status_validation,
);
self.insert(&key, validator);
}
}
if jwt_config.jwt_status_validation {
for algorithm in jwt_config.signature_algorithms_supported.iter().copied() {
let status_list_uri = iss_config
.openid_config
.as_ref()
.and_then(|conf| conf.status_list_endpoint.as_ref())
.map(std::string::ToString::to_string);
let (validator, key) = JwtValidator::new_status_list_tkn_validator(
Some(&iss),
status_list_uri,
algorithm,
jwt_config.jwt_sig_validation,
);
self.insert(&key, validator);
}
}
}
fn insert(&self, validator_info: &ValidatorInfo<'_>, validator: JwtValidator) {
let key = validator_info.key_hash();
let owned_info = validator_info.owned();
let cached = Arc::new(validator);
self.validators.rcu(|old| {
let mut new_map = (**old).clone();
new_map
.entry(key.clone())
.or_default()
.push((owned_info.clone(), Arc::clone(&cached)));
new_map
});
}
pub(crate) fn get(&self, validator_info: &ValidatorInfo<'_>) -> Option<Arc<JwtValidator>> {
let guard = self.validators.load();
let validators = guard.get(&validator_info.key_hash())?;
match validators.len() {
0 => None,
1 => Some(validators[0].1.clone()),
_ => validators
.iter()
.find(|(info, _)| info.is_equal_to(validator_info))
.map(|(_, validator)| validator.clone()),
}
}
}
#[derive(Hash, Clone)]
pub(crate) struct ValidatorInfo<'a> {
pub iss: Option<&'a IssClaim>,
pub token_kind: TokenKind<'a>,
pub algorithm: Algorithm,
}
#[derive(Hash, Clone, PartialEq)]
pub(crate) enum TokenKind<'a> {
AuthzRequestInput(&'a str),
StatusList,
AuthorizeMultiIssuer(Cow<'a, str>),
}
impl Display for TokenKind<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TokenKind::StatusList => write!(f, "statuslist+jwt"),
TokenKind::AuthorizeMultiIssuer(tkn_name) => {
write!(f, "{tkn_name}")
},
TokenKind::AuthzRequestInput(tkn_name) => {
write!(f, "{tkn_name}")
},
}
}
}
#[derive(Debug, Clone)]
pub struct OwnedValidatorInfo {
iss: Option<IssClaim>,
token_kind: OwnedTokenKind,
algorithm: Algorithm,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum OwnedTokenKind {
AuthzRequestInput(String),
StatusList,
AuthorizeMultiIssuer(String),
}
impl From<&TokenKind<'_>> for OwnedTokenKind {
fn from(tkn_kind: &TokenKind<'_>) -> Self {
match tkn_kind {
TokenKind::AuthzRequestInput(tkn_name) => {
Self::AuthzRequestInput((*tkn_name).to_string())
},
TokenKind::StatusList => Self::StatusList,
TokenKind::AuthorizeMultiIssuer(tkn_name) => {
Self::AuthorizeMultiIssuer(tkn_name.to_string())
},
}
}
}
impl OwnedTokenKind {
fn is_equal_to(&self, other: &TokenKind<'_>) -> bool {
match (self, other) {
(OwnedTokenKind::StatusList, TokenKind::StatusList) => true,
(
OwnedTokenKind::AuthzRequestInput(tkn_name_string),
TokenKind::AuthzRequestInput(tkn_name_str),
) => tkn_name_string.as_str() == *tkn_name_str,
(
OwnedTokenKind::AuthorizeMultiIssuer(tkn_name_string),
TokenKind::AuthorizeMultiIssuer(tkn_name_str),
) => tkn_name_string.as_str() == tkn_name_str,
_ => false,
}
}
}
#[derive(Clone, Hash, Eq, PartialEq)]
struct ValidatorKeyHash(u64);
impl ValidatorInfo<'_> {
pub(crate) fn owned(&self) -> OwnedValidatorInfo {
OwnedValidatorInfo {
iss: self.iss.map(std::borrow::ToOwned::to_owned),
token_kind: (&self.token_kind).into(),
algorithm: self.algorithm,
}
}
fn key_hash(&self) -> ValidatorKeyHash {
let mut state = DefaultHasher::new();
self.hash(&mut state);
ValidatorKeyHash(state.finish())
}
}
impl OwnedValidatorInfo {
fn is_equal_to(&self, other: &ValidatorInfo<'_>) -> bool {
if self.iss.as_ref() != other.iss {
return false;
}
if !self.token_kind.is_equal_to(&other.token_kind) {
return false;
}
if self.algorithm != other.algorithm {
return false;
}
true
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::common::policy_store::TokenEntityMetadata;
use crate::jwt::validation::JwtValidator;
use std::collections::HashSet;
#[test]
fn test_insert_and_retrieve() {
let store = JwtValidatorCache::default();
let test_iss = IssClaim::new("https://example.com/issuer");
let (validator, info) = JwtValidator::new_input_tkn_validator(
Some(&test_iss),
"access_tkn",
&TokenEntityMetadata {
trusted: true,
entity_type_name: "AccessToken".into(),
token_id: "iss".into(),
required_claims: HashSet::new(),
},
jsonwebtoken::Algorithm::HS256,
StatusListCache::default(),
true,
false,
);
store.insert(&info, validator.clone());
assert!(store.get(&info).is_some());
}
}