use super::layer::AuthorizationLayer;
use super::pattern::CompiledPattern;
use axum::http::Method;
use stano_security::{Claims, JwtConfig};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
type ValidatorFuture<'a> = Pin<Box<dyn Future<Output = Result<(), String>> + Send + 'a>>;
pub(super) type ClaimsValidator<E> =
Arc<dyn for<'a> Fn(&'a Claims<E>) -> ValidatorFuture<'a> + Send + Sync>;
pub(super) enum Effect<E> {
PermitAll,
Authenticated,
HasRole(Arc<dyn Fn(&E) -> bool + Send + Sync>),
}
impl<E> Clone for Effect<E> {
fn clone(&self) -> Self {
match self {
Effect::PermitAll => Effect::PermitAll,
Effect::Authenticated => Effect::Authenticated,
Effect::HasRole(pred) => Effect::HasRole(Arc::clone(pred)),
}
}
}
pub(super) struct Rule<E> {
pub(super) methods: Option<Vec<Method>>,
pub(super) pattern: CompiledPattern,
pub(super) effect: Effect<E>,
}
#[derive(Debug)]
pub enum AuthorizationBuildError {
MissingAnyRequest,
InvalidPattern(matchit::InsertError),
}
impl std::fmt::Display for AuthorizationBuildError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AuthorizationBuildError::MissingAnyRequest => {
write!(
f,
"authorization chain must end with a terminal any_request() rule"
)
}
AuthorizationBuildError::InvalidPattern(err) => {
write!(f, "invalid request matcher pattern: {err}")
}
}
}
}
impl std::error::Error for AuthorizationBuildError {}
pub struct AuthorizationBuilder<E> {
pub(super) rules: Vec<Rule<E>>,
pub(super) any_request: Option<Effect<E>>,
pub(super) claims_validator: Option<ClaimsValidator<E>>,
pub(super) cookie_name: Option<String>,
poisoned: Option<matchit::InsertError>,
}
impl<E> Default for AuthorizationBuilder<E> {
fn default() -> Self {
Self {
rules: Vec::new(),
any_request: None,
claims_validator: None,
cookie_name: None,
poisoned: None,
}
}
}
impl<E> AuthorizationBuilder<E>
where
E: serde::de::DeserializeOwned + Clone + Send + Sync + 'static,
{
pub fn new() -> Self {
Self::default()
}
pub fn request_matcher(self, pattern: &str) -> RequestMatcherBuilder<E> {
RequestMatcherBuilder {
parent: self,
methods: None,
pattern: PatternSpec::Pattern(pattern.to_string()),
}
}
pub fn any_request(self) -> RequestMatcherBuilder<E> {
RequestMatcherBuilder {
parent: self,
methods: None,
pattern: PatternSpec::AnyRequest,
}
}
pub fn with_claims_validator<F, Fut>(mut self, validator: F) -> Self
where
F: Fn(&Claims<E>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), String>> + Send + 'static,
{
self.claims_validator = Some(Arc::new(move |claims: &Claims<E>| {
Box::pin(validator(claims)) as ValidatorFuture<'_>
}));
self
}
pub fn cookie_name(mut self, name: impl Into<String>) -> Self {
self.cookie_name = Some(name.into());
self
}
pub fn build(
self,
jwt_config: JwtConfig,
) -> Result<AuthorizationLayer, AuthorizationBuildError> {
if let Some(err) = self.poisoned {
return Err(AuthorizationBuildError::InvalidPattern(err));
}
let any_request = self
.any_request
.ok_or(AuthorizationBuildError::MissingAnyRequest)?;
Ok(super::layer::build_layer(
self.rules,
any_request,
jwt_config,
self.claims_validator,
self.cookie_name,
))
}
}
enum PatternSpec {
Pattern(String),
AnyRequest,
}
pub struct RequestMatcherBuilder<E> {
parent: AuthorizationBuilder<E>,
methods: Option<Vec<Method>>,
pattern: PatternSpec,
}
impl<E> RequestMatcherBuilder<E>
where
E: serde::de::DeserializeOwned + Clone + Send + Sync + 'static,
{
pub fn methods(mut self, methods: impl IntoIterator<Item = Method>) -> Self {
self.methods = Some(methods.into_iter().collect());
self
}
pub fn permit_all(self) -> AuthorizationBuilder<E> {
self.finish(Effect::PermitAll)
}
pub fn authenticated(self) -> AuthorizationBuilder<E> {
self.finish(Effect::Authenticated)
}
pub fn has_role(
self,
predicate: impl Fn(&E) -> bool + Send + Sync + 'static,
) -> AuthorizationBuilder<E> {
self.finish(Effect::HasRole(Arc::new(predicate)))
}
pub fn has_any_role<F>(self, predicates: impl IntoIterator<Item = F>) -> AuthorizationBuilder<E>
where
F: Fn(&E) -> bool + Send + Sync + 'static,
{
let predicates: Vec<F> = predicates.into_iter().collect();
self.finish(Effect::HasRole(Arc::new(move |ext: &E| {
predicates.iter().any(|p| p(ext))
})))
}
fn finish(self, effect: Effect<E>) -> AuthorizationBuilder<E> {
let mut parent = self.parent;
match self.pattern {
PatternSpec::AnyRequest => {
parent.any_request = Some(effect);
}
PatternSpec::Pattern(pattern) => {
match CompiledPattern::new(&pattern) {
Ok(compiled) => parent.rules.push(Rule {
methods: self.methods,
pattern: compiled,
effect,
}),
Err(err) => parent.poison(err),
}
}
}
parent
}
}
impl<E> AuthorizationBuilder<E> {
fn poison(&mut self, err: matchit::InsertError) {
self.poisoned.get_or_insert(err);
}
}