use crate::middleware::session::{CleaningMemoryStore, session_db::RuniqueSessionStore};
use crate::utils::aliases::{
ADb, ARlockmap, ASecurityCsp, ASecurityHosts, ATera, new, new_registry,
};
use axum::{Router, middleware};
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,
allowed_hosts_middleware, csrf_middleware, dev_no_cache_middleware, error_handler_middleware,
https_redirect_middleware, security_headers_middleware,
};
#[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: Arc<PermissionsPolicy>,
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: new(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(),
}
}
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>()
}
pub fn attach_middlewares(engine: Arc<Self>, router: Router) -> Router {
let mut router = router;
let f = &engine.features;
if engine.config.security.enforce_https {
router = router.layer(middleware::from_fn_with_state(
engine.clone(),
https_redirect_middleware,
));
}
if f.enable_host_validation {
router = router.layer(middleware::from_fn_with_state(
engine.clone(),
allowed_hosts_middleware,
));
}
router = router.layer(middleware::from_fn_with_state(
engine.clone(),
csrf_middleware,
));
if !f.enable_cache {
router = router.layer(middleware::from_fn_with_state(
engine.clone(),
dev_no_cache_middleware,
));
}
if f.enable_csp {
router = router.layer(middleware::from_fn_with_state(
engine.clone(),
security_headers_middleware,
));
}
if f.enable_debug_errors {
router = router.layer(middleware::from_fn(error_handler_middleware));
}
router
}
}