use async_trait::async_trait;
use std::marker::PhantomData;
use std::sync::Arc;
use crate::middleware::AuthContext;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ResourceAction {
Read,
List,
Create,
Update,
Patch,
Delete,
Restore,
HardDelete,
Custom(&'static str),
}
impl ResourceAction {
pub fn name(&self) -> &'static str {
match self {
ResourceAction::Read => "read",
ResourceAction::List => "list",
ResourceAction::Create => "create",
ResourceAction::Update => "update",
ResourceAction::Patch => "patch",
ResourceAction::Delete => "delete",
ResourceAction::Restore => "restore",
ResourceAction::HardDelete => "hard_delete",
ResourceAction::Custom(name) => name,
}
}
}
#[async_trait]
pub trait ResourcePolicy<E: Send + Sync + 'static>: Send + Sync {
fn resource_type() -> &'static str
where
Self: Sized,
{
"resource"
}
fn create_permission() -> &'static str
where
Self: Sized,
{
"create"
}
fn read_permission() -> &'static str
where
Self: Sized,
{
"read"
}
fn list_permission() -> &'static str
where
Self: Sized,
{
"list"
}
fn update_permission() -> &'static str
where
Self: Sized,
{
"update"
}
fn patch_permission() -> &'static str
where
Self: Sized,
{
"patch"
}
fn delete_permission() -> &'static str
where
Self: Sized,
{
"delete"
}
fn restore_permission() -> &'static str
where
Self: Sized,
{
"restore"
}
async fn can(
&self,
action: ResourceAction,
entity: &E,
ctx: &AuthContext,
) -> bool;
fn explicitly_disabled_actions(&self) -> Vec<ResourceAction> {
vec![]
}
}
#[derive(Debug, thiserror::Error)]
#[error("access denied: caller '{caller}' may not perform '{action}' on this resource")]
pub struct AccessDenied {
pub caller: String,
pub action: String,
}
impl AccessDenied {
pub fn new(caller: impl Into<String>, action: &ResourceAction) -> Self {
Self {
caller: caller.into(),
action: action.name().into(),
}
}
}
impl Default for AccessDenied {
fn default() -> Self {
Self {
caller: "anonymous".into(),
action: "unknown".into(),
}
}
}
pub struct PermissionGuard<E> {
policy: Arc<dyn ResourcePolicy<E>>,
}
impl<E: Send + Sync + 'static> PermissionGuard<E> {
pub fn new(policy: Arc<dyn ResourcePolicy<E>>) -> Self {
Self { policy }
}
pub async fn check(
&self,
action: ResourceAction,
entity: &E,
ctx: &AuthContext,
) -> Result<(), AccessDenied> {
if self.policy.explicitly_disabled_actions().contains(&action) {
return Err(AccessDenied::new(&ctx.user_id, &action));
}
if self.policy.can(action, entity, ctx).await {
Ok(())
} else {
Err(AccessDenied::new(&ctx.user_id, &action))
}
}
}
#[async_trait]
pub trait AuthContextProvider: Send + Sync {
async fn current(&self) -> Option<AuthContext>;
}
pub struct ServicePermissionGuard<E, S, P>
where
E: Send + Sync + 'static,
P: ResourcePolicy<E>,
{
service: Arc<S>,
policy: Arc<P>,
_phantom: std::marker::PhantomData<E>,
}
impl<E, S, P> ServicePermissionGuard<E, S, P>
where
E: Send + Sync + 'static,
P: ResourcePolicy<E>,
{
pub fn new(service: Arc<S>, policy: Arc<P>) -> Self {
Self {
service,
policy,
_phantom: std::marker::PhantomData,
}
}
pub fn service(&self) -> &Arc<S> {
&self.service
}
pub fn policy(&self) -> &Arc<P> {
&self.policy
}
pub async fn check(
&self,
action: ResourceAction,
entity: &E,
ctx: &AuthContext,
) -> Result<(), AccessDenied> {
if self.policy.explicitly_disabled_actions().contains(&action) {
return Err(AccessDenied::new(&ctx.user_id, &action));
}
if self.policy.can(action, entity, ctx).await {
Ok(())
} else {
Err(AccessDenied::new(&ctx.user_id, &action))
}
}
}
pub struct PermitAllResourcePolicy<E> {
_phantom: PhantomData<E>,
}
impl<E> PermitAllResourcePolicy<E> {
pub fn new() -> Self {
Self {
_phantom: PhantomData,
}
}
}
impl<E> Default for PermitAllResourcePolicy<E> {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl<E: Send + Sync + 'static> ResourcePolicy<E> for PermitAllResourcePolicy<E> {
async fn can(&self, _action: ResourceAction, _entity: &E, _ctx: &AuthContext) -> bool {
true
}
}
pub struct DenyAllResourcePolicy<E> {
_phantom: PhantomData<E>,
}
impl<E> DenyAllResourcePolicy<E> {
pub fn new() -> Self {
Self {
_phantom: PhantomData,
}
}
}
#[async_trait]
impl<E: Send + Sync + 'static> ResourcePolicy<E> for DenyAllResourcePolicy<E> {
async fn can(&self, _action: ResourceAction, _entity: &E, _ctx: &AuthContext) -> bool {
false
}
}
pub struct RoleRequiredPolicy<E> {
required_roles: Vec<String>,
_phantom: PhantomData<E>,
}
impl<E> RoleRequiredPolicy<E> {
pub fn new(required_roles: Vec<impl Into<String>>) -> Self {
Self {
required_roles: required_roles.into_iter().map(|r| r.into()).collect(),
_phantom: PhantomData,
}
}
}
#[async_trait]
impl<E: Send + Sync + 'static> ResourcePolicy<E> for RoleRequiredPolicy<E> {
async fn can(&self, _action: ResourceAction, _entity: &E, ctx: &AuthContext) -> bool {
self.required_roles
.iter()
.any(|role| ctx.roles.contains(role))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug)]
struct Document {
owner_id: String,
}
struct OwnerPolicy;
#[async_trait]
impl ResourcePolicy<Document> for OwnerPolicy {
async fn can(
&self,
_action: ResourceAction,
entity: &Document,
ctx: &AuthContext,
) -> bool {
ctx.user_id == entity.owner_id
}
}
fn auth_ctx(user_id: &str) -> AuthContext {
AuthContext::new(user_id.to_string())
}
#[tokio::test]
async fn owner_permitted_stranger_denied() {
let guard = PermissionGuard::new(Arc::new(OwnerPolicy));
let doc = Document {
owner_id: "alice".into(),
};
assert!(guard
.check(ResourceAction::Update, &doc, &auth_ctx("alice"))
.await
.is_ok());
assert!(guard
.check(ResourceAction::Update, &doc, &auth_ctx("bob"))
.await
.is_err());
}
#[tokio::test]
async fn permit_all_always_ok() {
let guard: PermissionGuard<Document> =
PermissionGuard::new(Arc::new(PermitAllResourcePolicy::new()));
let doc = Document {
owner_id: "x".into(),
};
assert!(guard
.check(ResourceAction::Delete, &doc, &auth_ctx("anyone"))
.await
.is_ok());
}
}