use jsonwebtoken::Algorithm;
use serde::{Deserialize, Serialize};
use std::{collections::HashSet, num::NonZeroUsize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct JwtConfig {
pub jwks: Option<String>,
pub jwt_sig_validation: bool,
pub jwt_status_validation: bool,
pub signature_algorithms_supported: HashSet<Algorithm>,
pub token_cache_max_ttl_secs: usize,
pub token_cache_capacity: usize,
pub token_cache_earliest_expiration_eviction: bool,
pub trusted_issuer_loader: TrustedIssuerLoaderConfig,
pub jwks_refresh_interval: Option<u64>,
pub jwks_refresh_min_interval: u64,
pub status_list_refresh_interval_max: u64,
}
pub(crate) const DEFAULT_JWKS_REFRESH_INTERVAL_SECS: u64 = 3600;
pub(crate) const MIN_JWKS_REFRESH_SECS: u64 = 5;
pub(crate) const MIN_STATUS_LIST_REFRESH_SECS: u64 = 5;
#[must_use]
pub(crate) fn normalize_status_list_refresh_interval_max(value: u64) -> u64 {
match value {
0 => JwtConfig::DEFAULT_STATUS_LIST_REFRESH_INTERVAL_MAX_SECS,
v => v.max(MIN_STATUS_LIST_REFRESH_SECS),
}
}
impl Default for JwtConfig {
fn default() -> Self {
let config = Self {
jwks: None,
jwt_sig_validation: true,
jwt_status_validation: true,
signature_algorithms_supported: HashSet::new(),
token_cache_capacity: Self::DEFAULT_TOKEN_CACHE_CAPACITY,
token_cache_earliest_expiration_eviction: true,
token_cache_max_ttl_secs: Self::DEFAULT_TOKEN_CACHE_MAX_TTL_SECS,
trusted_issuer_loader: TrustedIssuerLoaderConfig::default(),
jwks_refresh_interval: None,
jwks_refresh_min_interval: Self::DEFAULT_JWKS_REFRESH_MIN_INTERVAL,
status_list_refresh_interval_max:
JwtConfig::DEFAULT_STATUS_LIST_REFRESH_INTERVAL_MAX_SECS,
};
config.allow_all_algorithms()
}
}
impl JwtConfig {
pub const DEFAULT_TOKEN_CACHE_CAPACITY: usize = 100;
pub const DEFAULT_TOKEN_CACHE_MAX_TTL_SECS: usize = 5;
pub const DEFAULT_JWKS_REFRESH_MIN_INTERVAL: u64 = 30;
pub const DEFAULT_STATUS_LIST_REFRESH_INTERVAL_MAX_SECS: u64 = 300;
pub(crate) fn normalize(&mut self) {
self.status_list_refresh_interval_max =
normalize_status_list_refresh_interval_max(self.status_list_refresh_interval_max);
}
#[must_use]
pub fn new_without_validation() -> Self {
Self {
jwks: None,
jwt_sig_validation: false,
jwt_status_validation: false,
signature_algorithms_supported: HashSet::new(),
..Default::default()
}
.allow_all_algorithms()
}
pub(crate) fn supported_algorithms() -> HashSet<Algorithm> {
HashSet::from_iter([
Algorithm::HS256,
Algorithm::HS384,
Algorithm::HS512,
Algorithm::ES256,
Algorithm::ES384,
Algorithm::RS256,
Algorithm::RS384,
Algorithm::RS512,
Algorithm::PS256,
Algorithm::PS384,
Algorithm::PS512,
Algorithm::EdDSA,
])
}
#[must_use]
pub fn allow_all_algorithms(mut self) -> Self {
self.signature_algorithms_supported = Self::supported_algorithms();
self
}
}
#[derive(Debug, Default, Clone, PartialEq, Serialize, Deserialize)]
pub enum TrustedIssuerLoaderTypeRaw {
#[default]
#[serde(rename = "SYNC")]
Sync,
#[serde(rename = "ASYNC")]
Async,
}
impl TrustedIssuerLoaderTypeRaw {
pub(crate) fn to_config(&self, workers: WorkersCount) -> TrustedIssuerLoaderConfig {
match self {
TrustedIssuerLoaderTypeRaw::Sync => TrustedIssuerLoaderConfig::Sync { workers },
TrustedIssuerLoaderTypeRaw::Async => TrustedIssuerLoaderConfig::Async { workers },
}
}
}
#[derive(Debug, Copy, Clone, PartialEq, Serialize, Deserialize)]
pub enum TrustedIssuerLoaderConfig {
Sync {
workers: WorkersCount,
},
Async {
workers: WorkersCount,
},
}
impl Default for TrustedIssuerLoaderConfig {
fn default() -> Self {
Self::Sync {
workers: WorkersCount::MIN,
}
}
}
#[derive(
Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, derive_more::Deref, serde::Serialize,
)]
pub struct WorkersCount(NonZeroUsize);
#[cfg(not(target_arch = "wasm32"))]
impl WorkersCount {
pub const MAX: WorkersCount = WorkersCount(NonZeroUsize::new(1000).unwrap());
pub const DEFAULT: WorkersCount = WorkersCount(NonZeroUsize::new(10).unwrap());
}
#[cfg(target_arch = "wasm32")]
impl WorkersCount {
pub const MAX: WorkersCount = WorkersCount(NonZeroUsize::new(6).unwrap());
pub const DEFAULT: WorkersCount = WorkersCount(NonZeroUsize::new(2).unwrap());
}
impl WorkersCount {
pub const MIN: WorkersCount = WorkersCount(NonZeroUsize::MIN);
}
impl WorkersCount {
#[must_use]
pub fn new(value: usize) -> Self {
let value = NonZeroUsize::new(value)
.unwrap_or(NonZeroUsize::MIN)
.min(Self::MAX.0);
Self(value)
}
}
impl Default for WorkersCount {
fn default() -> Self {
Self::DEFAULT
}
}
impl<'de> serde::Deserialize<'de> for WorkersCount {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = usize::deserialize(deserializer)?;
Ok(Self::new(value))
}
}
impl PartialEq<usize> for WorkersCount {
fn eq(&self, other: &usize) -> bool {
self.0.get() == *other
}
}