use std::sync::Arc;
use futures_util::future::BoxFuture;
use utoipa_axum::router::OpenApiRouter;
use crate::auth::Authenticator;
use crate::command::AppCommand;
use crate::config::Config;
use crate::db::{Db, Dialect};
use crate::error::Error;
use crate::hooks::Hooks;
use crate::migrate::Migration;
use crate::rate_limit::RateLimiter;
use crate::state::{AppState, Extensions};
use crate::ws::WsHandlers;
pub const RESERVED_MODULE_NAMES: &[&str] = &["app", "core", "nbs"];
pub const MAX_MODULE_NAME_BYTES: usize = 32;
pub trait Module: Send + Sync + 'static {
fn name(&self) -> &'static str;
fn depends_on(&self) -> &'static [&'static str] {
&[]
}
fn migrations(&self, dialect: Dialect) -> Vec<Migration> {
let _ = dialect;
Vec::new()
}
fn routes(&self) -> OpenApiRouter<AppState> {
OpenApiRouter::new()
}
fn openapi(&self) -> Option<utoipa::openapi::OpenApi> {
None
}
fn register_hooks(&self, hooks: &mut Hooks) {
let _ = hooks;
}
fn setup(&self, setup: &mut Setup<'_>) -> Result<(), Error> {
let _ = setup;
Ok(())
}
fn ws_handlers(&self, handlers: &mut WsHandlers) {
let _ = handlers;
}
fn commands(&self) -> Vec<Arc<dyn AppCommand>> {
Vec::new()
}
fn start<'a>(&'a self, state: &'a AppState) -> BoxFuture<'a, Result<(), Error>> {
let _ = state;
Box::pin(async { Ok(()) })
}
fn shutdown<'a>(&'a self, state: &'a AppState) -> BoxFuture<'a, ()> {
let _ = state;
Box::pin(async {})
}
}
pub struct Setup<'a> {
pub(crate) config: &'a Config,
pub(crate) db: &'a Db,
pub(crate) extensions: &'a mut Extensions,
pub(crate) authenticators: &'a mut Vec<Arc<dyn Authenticator>>,
pub(crate) rate_limiters: &'a mut Vec<Arc<dyn RateLimiter>>,
}
impl std::fmt::Debug for Setup<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Setup").field("authenticators", &self.authenticators.len()).field("rate_limiters", &self.rate_limiters.len()).finish_non_exhaustive()
}
}
impl Setup<'_> {
pub fn config(&self) -> &Config {
self.config
}
pub fn db(&self) -> &Db {
self.db
}
pub fn insert_state<T: Send + Sync + 'static>(&mut self, value: T) {
self.extensions.insert(value);
}
pub fn has_state<T: Send + Sync + 'static>(&self) -> bool {
self.extensions.get::<T>().is_some()
}
pub fn add_authenticator(&mut self, authenticator: Arc<dyn Authenticator>) {
self.authenticators.push(authenticator);
}
pub fn add_rate_limiter(&mut self, limiter: Arc<dyn RateLimiter>) {
self.rate_limiters.push(limiter);
}
}
pub fn validate_module_name(name: &str) -> Result<(), String> {
let mut bytes = name.bytes();
let first_ok = bytes.next().is_some_and(|b| b.is_ascii_lowercase());
if !first_ok || !name.bytes().all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_') {
return Err(format!("module name `{name}` must match [a-z][a-z0-9_]*"));
}
if name.len() > MAX_MODULE_NAME_BYTES {
return Err(format!("module name `{name}` is longer than {MAX_MODULE_NAME_BYTES} bytes"));
}
if RESERVED_MODULE_NAMES.contains(&name) {
return Err(format!("module name `{name}` is reserved"));
}
Ok(())
}
#[derive(Default)]
pub(crate) struct ModuleSet(Vec<Box<dyn Module>>);
impl ModuleSet {
pub(crate) fn push(&mut self, module: Box<dyn Module>) {
self.0.push(module);
}
pub(crate) fn validate(&self) -> Result<(), Error> {
let mut problems = Vec::new();
let mut seen = std::collections::HashSet::new();
for module in &self.0 {
let name = module.name();
if let Err(problem) = validate_module_name(name) {
problems.push(problem);
}
for dependency in module.depends_on() {
if !seen.contains(dependency) {
let later = self.0.iter().any(|m| m.name() == *dependency);
problems.push(if later {
format!("module `{name}` needs `{dependency}` registered before it (register `{dependency}` first)")
} else {
format!("module `{name}` needs the module `{dependency}` (register it first)")
});
}
}
if !seen.insert(name) {
problems.push(format!("module `{name}` is registered twice"));
}
}
if problems.is_empty() {
Ok(())
} else {
Err(Error::Module(problems.join("; ")))
}
}
pub(crate) fn iter(&self) -> impl DoubleEndedIterator<Item = &dyn Module> {
self.0.iter().map(|m| m.as_ref())
}
pub(crate) fn names(&self) -> Vec<&'static str> {
self.0.iter().map(|m| m.name()).collect()
}
pub(crate) fn get(&self, name: &str) -> Option<&dyn Module> {
self.iter().find(|m| m.name() == name)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Named(&'static str);
impl Module for Named {
fn name(&self) -> &'static str {
self.0
}
}
#[test]
fn names() {
for ok in ["chat", "storage", "a", "game_2", "abcdefghijabcdefghijabcdefghij12"] {
assert!(validate_module_name(ok).is_ok(), "{ok}");
}
for bad in ["", "Chat", "2chat", "_x", "chat-x", "chat.x", "app", "core", "nbs", "abcdefghijabcdefghijabcdefghij123", "ch at"] {
assert!(validate_module_name(bad).is_err(), "{bad}");
}
}
#[test]
fn order_and_duplicates() {
let mut set = ModuleSet::default();
for name in ["zeta", "alpha", "mid"] {
set.push(Box::new(Named(name)));
}
assert!(set.validate().is_ok());
assert_eq!(set.names(), ["zeta", "alpha", "mid"], "registration order, not sorted");
assert_eq!(set.iter().rev().map(|m| m.name()).collect::<Vec<_>>(), ["mid", "alpha", "zeta"]);
set.push(Box::new(Named("alpha")));
set.push(Box::new(Named("App")));
let error = set.validate().err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("`alpha` is registered twice") && error.contains("`App`"), "{error}");
}
struct Needs(&'static str, &'static [&'static str]);
impl Module for Needs {
fn name(&self) -> &'static str {
self.0
}
fn depends_on(&self) -> &'static [&'static str] {
self.1
}
}
#[test]
fn dependencies_come_first() {
let mut ok = ModuleSet::default();
ok.push(Box::new(Needs("auth", &[])));
ok.push(Box::new(Needs("chat", &["auth"])));
assert!(ok.validate().is_ok());
let mut late = ModuleSet::default();
late.push(Box::new(Needs("chat", &["auth"])));
late.push(Box::new(Needs("auth", &[])));
let error = late.validate().err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("needs `auth` registered before it"), "{error}");
let mut missing = ModuleSet::default();
missing.push(Box::new(Needs("chat", &["auth"])));
let error = missing.validate().err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("needs the module `auth`"), "{error}");
}
}