From 068e100ca1373be98fb5e8a5045ba63efc44ac95 Mon Sep 17 00:00:00 2001 From: Nell Date: Sat, 25 Jul 2026 17:19:36 +0200 Subject: [PATCH] init --- src/core/permission_sync.rs | 2 + src/http/mod.rs | 3 + src/http/permissions.rs | 225 ++++++++++++++++++++++++ src/lib.rs | 2 + src/repositories/computed_permission.rs | 87 ++++++--- src/utils/mod.rs | 81 +++++++++ 6 files changed, 371 insertions(+), 29 deletions(-) create mode 100644 src/http/permissions.rs create mode 100644 src/utils/mod.rs diff --git a/src/core/permission_sync.rs b/src/core/permission_sync.rs index 674f5f9..6cc4394 100644 --- a/src/core/permission_sync.rs +++ b/src/core/permission_sync.rs @@ -60,4 +60,6 @@ impl PermissionSyncService { .full_sync_server(server_id) .await; } + + async fn sync_server_permissions_by_user() {} } diff --git a/src/http/mod.rs b/src/http/mod.rs index 5d6b3d1..d8be93d 100644 --- a/src/http/mod.rs +++ b/src/http/mod.rs @@ -7,5 +7,8 @@ pub mod metrics; pub mod middleware; pub mod server; pub mod validation; +pub mod permissions; + +pub use permissions::{RequireServerPermission, RequireChannelPermission}; pub type OxRouter = Router; diff --git a/src/http/permissions.rs b/src/http/permissions.rs new file mode 100644 index 0000000..2701196 --- /dev/null +++ b/src/http/permissions.rs @@ -0,0 +1,225 @@ +// Unused + +use super::context::CurrentUser; +use super::error::HTTPError; +use crate::core::AppState; +use crate::permissions::{ChannelPermission, ServerPermission}; +use axum::extract::FromRequestParts; +use axum::http::request::Parts; +use std::ops::Deref; +use uuid::Uuid; + +/// An Axum extractor that ensures the currently authenticated user has the specified +/// server permission(s) on a target server. +/// +/// The target `server_id` is automatically extracted from path parameters (supporting +/// path parameters named `server_id` or `id`). +/// +/// # Superuser Bypass +/// If the user is a superuser (`is_superuser == true`), the permission check automatically passes. +/// +/// # Usage Example +/// ```rust +/// use axum::extract::State; +/// use uuid::Uuid; +/// use crate::http::permissions::RequireServerPermission; +/// use crate::permissions::ServerPermission; +/// use crate::core::AppState; +/// +/// pub async fn update_server_settings( +/// RequireServerPermission::<{ ServerPermission::MANAGE_SERVER.bits() }>(user): RequireServerPermission<{ ServerPermission::MANAGE_SERVER.bits() }>, +/// State(state): State, +/// Path(server_id): Path, +/// ) -> Result<(), HTTPError> { +/// // User has MANAGE_SERVER or is a superuser +/// Ok(()) +/// } +/// ``` +#[derive(Clone, Debug)] +pub struct RequireServerPermission(pub CurrentUser); + +impl Deref for RequireServerPermission { + type Target = CurrentUser; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl FromRequestParts for RequireServerPermission +where + S: Send + Sync, +{ + type Rejection = HTTPError; + + async fn from_request_parts(parts: &mut Parts, state: &S) -> Result { + // 1. Extract CurrentUser (which validates authentication and returns 401 if missing) + let current_user = CurrentUser::from_request_parts(parts, state).await?; + + // 2. Superuser bypasses all checks + if current_user.is_superuser { + return Ok(RequireServerPermission(current_user)); + } + + // 3. Get AppState from extensions + let app_state = match parts.extensions.get::() { + Some(s) => s.clone(), + None => { + return Err(HTTPError::InternalServerError( + "AppState missing in request extensions".to_string(), + )); + } + }; + + // 4. Extract server_id from path parameters. + let server_id = match extract_path_param_uuid(parts, &["server_id", "id"]) { + Some(id) => id, + None => { + return Err(HTTPError::BadRequest( + "Missing or invalid server_id".to_string(), + )); + } + }; + + // 5. Check user permission via server repository + let permission_result = app_state + .repositories + .server + .get_user_permission(server_id, current_user.id) + .await; + + let permission_bits = match permission_result { + Ok(Some(p)) => p.permissions, + Ok(None) => 0, + Err(e) => return Err(HTTPError::InternalServerError(e.to_string())), + }; + + let required = ServerPermission::from_bits_truncate(PERM); + let granted = ServerPermission::from_bits_truncate(permission_bits as u64); + + if granted.contains(required) { + Ok(RequireServerPermission(current_user)) + } else { + Err(HTTPError::Forbidden) + } + } +} + +/// An Axum extractor that ensures the currently authenticated user has the specified +/// channel permission(s) on a target channel. +/// +/// The target `channel_id` (or `id`) is automatically extracted from path parameters. +/// +/// # Superuser Bypass +/// If the user is a superuser (`is_superuser == true`), the permission check automatically passes. +/// +/// # Usage Example +/// ```rust +/// use axum::extract::State; +/// use uuid::Uuid; +/// use crate::http::permissions::RequireChannelPermission; +/// use crate::permissions::ChannelPermission; +/// use crate::core::AppState; +/// +/// pub async fn read_channel_messages( +/// RequireChannelPermission::<{ ChannelPermission::READ_CHANNEL.bits() }>(user): RequireChannelPermission<{ ChannelPermission::READ_CHANNEL.bits() }>, +/// State(state): State, +/// Path(channel_id): Path, +/// ) -> Result<(), HTTPError> { +/// // User has READ_CHANNEL or is a superuser +/// Ok(()) +/// } +/// ``` +#[derive(Clone, Debug)] +pub struct RequireChannelPermission(pub CurrentUser); + +impl Deref for RequireChannelPermission { + type Target = CurrentUser; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl FromRequestParts for RequireChannelPermission +where + S: Send + Sync, +{ + type Rejection = HTTPError; + + async fn from_request_parts(parts: &mut Parts, state: &S) -> Result { + let current_user = CurrentUser::from_request_parts(parts, state).await?; + + if current_user.is_superuser { + return Ok(RequireChannelPermission(current_user)); + } + + let app_state = match parts.extensions.get::() { + Some(s) => s.clone(), + None => { + return Err(HTTPError::InternalServerError( + "AppState missing in request extensions".to_string(), + )); + } + }; + + let channel_id = match extract_path_param_uuid(parts, &["channel_id", "id"]) { + Some(id) => id, + None => { + return Err(HTTPError::BadRequest( + "Missing or invalid channel_id".to_string(), + )); + } + }; + + let permission_result = app_state + .repositories + .channel + .get_user_permission(channel_id, current_user.id) + .await; + + let permission_bits = match permission_result { + Ok(Some(p)) => p.permissions, + Ok(None) => 0, + Err(e) => return Err(HTTPError::InternalServerError(e.to_string())), + }; + + let required = ChannelPermission::from_bits_truncate(PERM); + let granted = ChannelPermission::from_bits_truncate(permission_bits as u64); + + if granted.contains(required) { + Ok(RequireChannelPermission(current_user)) + } else { + Err(HTTPError::Forbidden) + } + } +} + +/// Helper function to extract a Uuid path parameter matching any of the given key names +/// from Axum request extensions. +fn extract_path_param_uuid(parts: &Parts, keys: &[&str]) -> Option { + if let Some(map) = parts + .extensions + .get::>() + { + for key in keys { + if let Some(val) = map.get(*key) { + if let Ok(uuid) = Uuid::parse_str(val) { + return Some(uuid); + } + } + } + } + + if let Some(params) = parts.extensions.get::>() { + for (k, v) in params { + if keys.contains(&k.as_str()) { + if let Ok(uuid) = Uuid::parse_str(v) { + return Some(uuid); + } + } + } + } + + None +} diff --git a/src/lib.rs b/src/lib.rs index 264b1ca..e0d3913 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -13,3 +13,5 @@ pub mod auth; pub mod metrics; pub mod domain; + +pub mod utils; diff --git a/src/repositories/computed_permission.rs b/src/repositories/computed_permission.rs index 558dc21..dea7d1f 100644 --- a/src/repositories/computed_permission.rs +++ b/src/repositories/computed_permission.rs @@ -5,17 +5,21 @@ use crate::models::{ use crate::permissions::{ChannelPermission, ServerPermission}; use crate::repositories::{AnyResult, RepositoryContext}; -use sea_orm::{ColumnTrait, EntityTrait, QueryFilter, QuerySelect, Set, TransactionTrait}; +use sea_orm::{ + ColumnTrait, EntityTrait, PaginatorTrait, QueryFilter, QuerySelect, Set, TransactionTrait, +}; use std::collections::HashMap; use std::sync::{Arc, OnceLock}; -use tokio::sync::Mutex; use uuid::Uuid; use crate::models::computed_permission::PermissionScopeType; +use crate::utils::ScopedLockManager; -static SYNC_LOCK: OnceLock> = OnceLock::new(); -fn _sync_lock() -> &'static Mutex<()> { - SYNC_LOCK.get_or_init(Mutex::default) +// Instance globale du manager de verrous scopés par Server ID +static PERM_LOCK_MANAGER: OnceLock> = OnceLock::new(); + +fn lock_manager() -> &'static ScopedLockManager { + PERM_LOCK_MANAGER.get_or_init(ScopedLockManager::new) } #[derive(Clone, Debug)] @@ -24,13 +28,29 @@ pub struct ComputedPermissionRepository { } impl ComputedPermissionRepository { + /// Récupère toutes les permissions calculées. pub async fn get_all(&self) -> AnyResult> { Ok(computed_permission::Entity::find() .all(&self.context.db) .await?) } - /// Recalcule le cache de permissions pour tous les utilisateurs du serveur. + /// Vérifie si l'utilisateur possède au moins une entrée de permission sur une ressource. + pub async fn had_perm_on(&self, user_id: Uuid, resource_id: Uuid) -> AnyResult { + Ok(computed_permission::Entity::find() + .filter(computed_permission::Column::UserId.eq(user_id)) + .filter(computed_permission::Column::ResourceId.eq(resource_id)) + .count(&self.context.db) + .await? + > 0) + } + + // ------------------------------------------------------------------------- + // Synchronisations de permissions scopées + // ------------------------------------------------------------------------- + + /// Scope SERVEUR : Recalcule le cache de permissions pour TOUS les utilisateurs du serveur. + /// À n'utiliser que pour les opérations lourdes ou structurelles. pub async fn full_sync_server(&self, server_id: Uuid) -> AnyResult<()> { let user_ids = server_user::Entity::find() .filter(server_user::Column::ServerId.eq(server_id)) @@ -47,20 +67,36 @@ impl ComputedPermissionRepository { Ok(()) } - /// Recalcule le cache de permissions d'un utilisateur sur un serveur. + /// Scope RÔLE : Recalcule le cache uniquement pour les membres d'un rôle spécifique. + pub async fn sync_role_members(&self, role_id: Uuid, server_id: Uuid) -> AnyResult<()> { + let user_ids = role_user::Entity::find() + .filter(role_user::Column::RoleId.eq(role_id)) + .select_only() + .column(role_user::Column::UserId) + .into_tuple::() + .all(&self.context.db) + .await?; + + for user_id in user_ids { + self.full_sync_user(user_id, server_id).await?; + } + + Ok(()) + } + + /// Scope UTILISATEUR : Recalcule le cache de permissions d'un seul utilisateur sur un serveur. /// /// Les permissions effectives sont composées de : - /// /// - permissions serveur accordées aux rôles de l'utilisateur ; /// - permissions serveur accordées directement à l'utilisateur ; /// - permissions de canal accordées aux rôles de l'utilisateur ; /// - permissions directes de l'utilisateur dans les canaux. pub async fn full_sync_user(&self, user_id: Uuid, server_id: Uuid) -> AnyResult<()> { - let _guard = _sync_lock().lock().await; - // --------------------------------------------------------------------- - // Rôles de l'utilisateur - // --------------------------------------------------------------------- + let _server_guard = lock_manager().lock_scope(server_id).await; + // --------------------------------------------------------------------- + // 1. Rôles de l'utilisateur + // --------------------------------------------------------------------- let role_ids = role_user::Entity::find() .filter(role_user::Column::UserId.eq(user_id)) .select_only() @@ -70,9 +106,8 @@ impl ComputedPermissionRepository { .await?; // --------------------------------------------------------------------- - // Permissions serveur des rôles + // 2. Permissions serveur des rôles // --------------------------------------------------------------------- - let mut server_permissions = ServerPermission::empty(); if !role_ids.is_empty() { @@ -89,9 +124,8 @@ impl ComputedPermissionRepository { } // --------------------------------------------------------------------- - // Permissions serveur directes de l'utilisateur + // 3. Permissions serveur directes de l'utilisateur // --------------------------------------------------------------------- - if let Some(permission) = server_user_permission::Entity::find() .filter(server_user_permission::Column::ServerId.eq(server_id)) .filter(server_user_permission::Column::UserId.eq(user_id)) @@ -102,20 +136,18 @@ impl ComputedPermissionRepository { } // --------------------------------------------------------------------- - // Canaux du serveur + // 4. Canaux du serveur // --------------------------------------------------------------------- - let channels = channel::Entity::find() .filter(channel::Column::ServerId.eq(server_id)) .all(&self.context.db) .await?; - let channel_ids: Vec = channels.iter().map(|channel| channel.id).collect(); + let channel_ids: Vec = channels.iter().map(|c| c.id).collect(); // --------------------------------------------------------------------- - // Permissions de rôles pour tous les canaux + // 5. Permissions de rôles pour tous les canaux // --------------------------------------------------------------------- - let role_channel_permissions = if role_ids.is_empty() || channel_ids.is_empty() { Vec::new() } else { @@ -138,9 +170,8 @@ impl ComputedPermissionRepository { } // --------------------------------------------------------------------- - // Permissions directes de l'utilisateur pour tous les canaux + // 6. Permissions directes de l'utilisateur pour tous les canaux // --------------------------------------------------------------------- - let user_channel_permissions = if channel_ids.is_empty() { Vec::new() } else { @@ -161,12 +192,11 @@ impl ComputedPermissionRepository { } // --------------------------------------------------------------------- - // Construction du cache + // 7. Construction des modèles à insérer // --------------------------------------------------------------------- - let mut computed_permissions = Vec::with_capacity(channels.len().saturating_add(1)); - // Permissions au niveau serveur. + // Permissions au niveau serveur computed_permissions.push(computed_permission::ActiveModel { user_id: Set(user_id), server_id: Set(server_id), @@ -176,7 +206,7 @@ impl ComputedPermissionRepository { ..Default::default() }); - // Permissions au niveau canal. + // Permissions au niveau canal for channel in channels { let channel_permissions = permissions_by_channel .remove(&channel.id) @@ -193,9 +223,8 @@ impl ComputedPermissionRepository { } // --------------------------------------------------------------------- - // Remplacement atomique du cache + // 8. Remplacement atomique du cache en BDD // --------------------------------------------------------------------- - self.context .db .transaction::<_, (), anyhow::Error>(|transaction| { diff --git a/src/utils/mod.rs b/src/utils/mod.rs new file mode 100644 index 0000000..7e0db85 --- /dev/null +++ b/src/utils/mod.rs @@ -0,0 +1,81 @@ +use std::collections::HashMap; +use std::hash::Hash; +use std::sync::{Arc, Weak}; +use tokio::sync::{Mutex, OwnedMutexGuard, RwLock, RwLockReadGuard, RwLockWriteGuard}; + +/// Structure gérant les garde-fous pour un verrou scopé. +pub struct ScopedGuard<'a, K: Eq + Hash + Clone> { + _global_read_guard: RwLockReadGuard<'a, ()>, + _local_guard: OwnedMutexGuard<()>, + key: K, + manager: &'a ScopedLockManager, +} + +impl<'a, K: Eq + Hash + Clone> Drop for ScopedGuard<'a, K> { + fn drop(&mut self) { + // Optionnel : Nettoyage des Mutex orphelins dans la HashMap scopée + let mut map = self.manager.scopes.blocking_lock(); + if let Some(weak) = map.get(&self.key) { + if weak.strong_count() == 0 { + map.remove(&self.key); + } + } + } +} + +#[derive(Clone, Default)] +pub struct ScopedLockManager { + global_lock: Arc>, + scopes: Arc>>>>, +} + +impl ScopedLockManager { + pub fn new() -> Self { + Self { + global_lock: Arc::new(RwLock::new(())), + scopes: Arc::new(Mutex::new(HashMap::new())), + } + } + + /// Acquiert le verrou global (Exclusif / Write Lock). + /// Bloque tous les autres verrous globaux et tous les verrous scopés. + pub async fn lock_global(&self) -> RwLockWriteGuard<'_, ()> { + self.global_lock.write().await + } + + /// Acquiert un verrou ciblé sur une clé `key` (ex: server_id ou user_id). + /// Permet des exécutions parallèles sur des `key` différentes, + /// mais garantit la sérialisation pour la même `key`. + pub async fn lock_scope(&self, key: K) -> ScopedGuard<'_, K> { + // 1. Prendre une garde de lecture (Read) sur le verrou global + let global_read = self.global_lock.read().await; + + // 2. Récupérer ou créer le Mutex dédié à la clé `key` + let local_mutex = { + let mut map = self.scopes.lock().await; + if let Some(weak) = map.get(&key) { + if let Some(arc) = weak.upgrade() { + arc + } else { + let arc = Arc::new(Mutex::new(())); + map.insert(key.clone(), Arc::downgrade(&arc)); + arc + } + } else { + let arc = Arc::new(Mutex::new(())); + map.insert(key.clone(), Arc::downgrade(&arc)); + arc + } + }; + + // 3. Verrouiller le Mutex local + let local_guard = local_mutex.lock_owned().await; + + ScopedGuard { + _global_read_guard: global_read, + _local_guard: local_guard, + key, + manager: self, + } + } +}