use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::fmt;
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::Arc;
use axum::extract::FromRef;
use net_backend_protocol::UnixMillis;
use crate::config::Config;
use crate::db::Db;
use crate::hooks::Hooks;
use crate::http::client_ip::IpNet;
use crate::shutdown::Shutdown;
use crate::ws::Hub;
pub trait Clock: Send + Sync + 'static {
fn now(&self) -> UnixMillis;
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SystemClock;
impl Clock for SystemClock {
fn now(&self) -> UnixMillis {
UnixMillis::now()
}
}
#[derive(Debug, Default)]
pub struct ManualClock(AtomicI64);
impl ManualClock {
pub fn new(now: UnixMillis) -> Self {
Self(AtomicI64::new(now.get()))
}
pub fn set(&self, now: UnixMillis) {
self.0.store(now.get(), Ordering::SeqCst);
}
pub fn advance(&self, millis: i64) {
let _ = self.0.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |t| Some(t.saturating_add(millis)));
}
}
impl Clock for ManualClock {
fn now(&self) -> UnixMillis {
UnixMillis(self.0.load(Ordering::SeqCst))
}
}
impl<C: Clock> Clock for Arc<C> {
fn now(&self) -> UnixMillis {
(**self).now()
}
}
#[derive(Clone, Default)]
pub(crate) struct Extensions(HashMap<TypeId, Arc<dyn Any + Send + Sync>>);
impl Extensions {
pub(crate) fn insert<T: Send + Sync + 'static>(&mut self, value: T) {
self.0.insert(TypeId::of::<T>(), Arc::new(value));
}
pub(crate) fn get<T: Send + Sync + 'static>(&self) -> Option<Arc<T>> {
self.0.get(&TypeId::of::<T>()).cloned().and_then(|v| v.downcast::<T>().ok())
}
}
struct Inner {
config: Arc<Config>,
db: Db,
clock: Arc<dyn Clock>,
hooks: Arc<Hooks>,
extensions: Extensions,
modules: Vec<&'static str>,
shutdown: Shutdown,
trusted_proxies: Vec<IpNet>,
ws: Hub,
}
#[derive(Clone)]
pub struct AppState(Arc<Inner>);
impl fmt::Debug for AppState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AppState").field("db", &self.0.db).field("modules", &self.0.modules).finish_non_exhaustive()
}
}
impl AppState {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
config: Arc<Config>,
db: Db,
clock: Arc<dyn Clock>,
hooks: Arc<Hooks>,
extensions: Extensions,
modules: Vec<&'static str>,
shutdown: Shutdown,
ws: Hub,
) -> Self {
let trusted_proxies = config.http.trusted_proxies.iter().filter_map(|p| IpNet::parse(p)).collect();
Self(Arc::new(Inner { config, db, clock, hooks, extensions, modules, shutdown, trusted_proxies, ws }))
}
pub fn config(&self) -> &Arc<Config> {
&self.0.config
}
pub fn db(&self) -> &Db {
&self.0.db
}
pub fn clock(&self) -> &Arc<dyn Clock> {
&self.0.clock
}
pub fn now(&self) -> UnixMillis {
self.0.clock.now()
}
pub fn hooks(&self) -> &Arc<Hooks> {
&self.0.hooks
}
pub fn get<T: Send + Sync + 'static>(&self) -> Option<Arc<T>> {
self.0.extensions.get::<T>()
}
pub fn modules(&self) -> &[&'static str] {
&self.0.modules
}
pub(crate) fn trusted_proxies(&self) -> &[IpNet] {
&self.0.trusted_proxies
}
pub fn ws(&self) -> &Hub {
&self.0.ws
}
pub fn shutdown(&self) -> &Shutdown {
&self.0.shutdown
}
}
impl FromRef<AppState> for Db {
fn from_ref(state: &AppState) -> Db {
state.0.db.clone()
}
}
impl FromRef<AppState> for Arc<Config> {
fn from_ref(state: &AppState) -> Arc<Config> {
state.0.config.clone()
}
}
impl FromRef<AppState> for Arc<dyn Clock> {
fn from_ref(state: &AppState) -> Arc<dyn Clock> {
state.0.clock.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn manual_clock() {
let clock = ManualClock::new(UnixMillis(1000));
clock.advance(500);
assert_eq!(clock.now(), UnixMillis(1500));
clock.set(UnixMillis(i64::MAX));
clock.advance(1);
assert_eq!(clock.now(), UnixMillis(i64::MAX));
assert!(SystemClock.now() > UnixMillis(1_700_000_000_000));
}
#[test]
fn extensions_by_type() {
let mut ext = Extensions::default();
ext.insert(5u32);
ext.insert(String::from("x"));
assert_eq!(ext.get::<u32>().as_deref(), Some(&5));
assert_eq!(ext.get::<String>().as_deref().map(String::as_str), Some("x"));
assert!(ext.get::<u64>().is_none());
}
}