use crate::middleware::session::{CleaningMemoryStore, session_db::RuniqueSessionStore};
use crate::utils::aliases::{
ADb, APermissionsPolicy, ARlockmap, ASecurityCsp, ASecurityHosts, ATera, new, new_registry,
};
use std::any::TypeId;
use std::collections::HashMap;
use std::sync::{Arc, LazyLock, RwLock};
use tera::Tera;
use crate::config::RuniqueConfig;
use crate::middleware::{
HostPolicy, MiddlewareConfig, PermissionsPolicy, SecurityPolicy, TrustedProxies,
};
#[cfg(feature = "orm")]
use sea_orm::DatabaseConnection;
#[derive(Debug)]
pub struct RuniqueEngine {
pub config: RuniqueConfig,
pub tera: ATera,
#[cfg(feature = "orm")]
pub db: ADb,
pub url_registry: ARlockmap,
pub features: MiddlewareConfig,
pub security_csp: ASecurityCsp,
pub security_hosts: ASecurityHosts,
pub csrf_exempt_paths: Arc<Vec<String>>,
pub permissions_policy: APermissionsPolicy,
pub trusted_proxies: Arc<TrustedProxies>,
pub session_store: LazyLock<RwLock<Option<Arc<CleaningMemoryStore>>>>,
pub session_db_store: LazyLock<RwLock<Option<Arc<RuniqueSessionStore>>>>,
pub extensions: HashMap<TypeId, Arc<dyn std::any::Any + Send + Sync>>,
}
impl RuniqueEngine {
#[cfg(feature = "orm")]
pub fn new(config: RuniqueConfig, tera: Tera, db: DatabaseConnection) -> Self {
let features = MiddlewareConfig::from_env();
let security_csp = SecurityPolicy::default();
let security_hosts = HostPolicy::default();
Self {
config,
tera: new(tera),
db: ADb::from_connection(db),
url_registry: new_registry(),
features,
security_csp: new(security_csp),
security_hosts: new(security_hosts),
csrf_exempt_paths: Arc::new(vec![]),
permissions_policy: Arc::new(PermissionsPolicy::default()),
trusted_proxies: Arc::new(TrustedProxies::default()),
session_store: LazyLock::new(|| RwLock::new(None)),
session_db_store: LazyLock::new(|| RwLock::new(None)),
extensions: HashMap::new(),
}
}
#[must_use]
pub fn session_store_saturated(&self) -> bool {
self.session_store
.read()
.ok()
.and_then(|g| g.as_ref().map(|s| s.is_saturated()))
.unwrap_or(false)
}
pub async fn close_user_sessions(&self, user_id: crate::utils::pk::Pk) {
let db_store = self.session_db_store.read().ok().and_then(|g| g.clone());
if let Some(store) = db_store
&& let Err(e) = store.invalidate_all(user_id).await
{
tracing::error!(
user_id = %user_id,
error = %e,
"closing the user's sessions in the database failed"
);
}
let memory_store = self.session_store.read().ok().and_then(|g| g.clone());
if let Some(store) = memory_store {
store.invalidate_user_sessions(user_id).await;
}
}
pub fn extension<T: std::any::Any + Send + Sync + 'static>(&self) -> Option<Arc<T>> {
self.extensions
.get(&TypeId::of::<T>())
.and_then(|arc| arc.clone().downcast::<T>().ok())
}
pub fn custom_db<T: std::any::Any + Send + Sync + 'static>(&self) -> Option<Arc<T>> {
self.extension::<T>()
}
}