use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::fmt;
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use std::time::Duration;
use futures_util::future::BoxFuture;
use futures_util::FutureExt;
use net_backend_protocol::codes;
use crate::db::DbTx;
use crate::error::AppError;
use crate::http::RequestId;
use crate::state::AppState;
pub trait Event: Send + Sync + 'static {
const NAME: &'static str;
}
#[derive(Debug)]
#[non_exhaustive]
pub enum Decision<E> {
Continue(E),
Reject(AppError),
}
#[derive(Clone, Debug)]
pub struct HookCtx {
state: AppState,
request_id: Option<RequestId>,
}
impl HookCtx {
pub fn new(state: AppState, request_id: Option<RequestId>) -> Self {
Self { state, request_id }
}
pub fn state(&self) -> &AppState {
&self.state
}
pub fn db(&self) -> &crate::db::Db {
self.state.db()
}
pub fn request_id(&self) -> Option<&RequestId> {
self.request_id.as_ref()
}
}
type BeforeFn<E> = Arc<dyn Fn(HookCtx, E) -> BoxFuture<'static, Result<Decision<E>, AppError>> + Send + Sync>;
type AfterFn<E> = Arc<dyn Fn(HookCtx, Arc<E>) -> BoxFuture<'static, Result<(), AppError>> + Send + Sync>;
type InTxFn<E> = Arc<dyn for<'a> Fn(&'a mut DbTx, &'a HookCtx, &'a E) -> BoxFuture<'a, Result<(), AppError>> + Send + Sync>;
type StartFn = Arc<dyn Fn(HookCtx) -> BoxFuture<'static, Result<(), AppError>> + Send + Sync>;
type ShutdownFn = Arc<dyn Fn(HookCtx) -> BoxFuture<'static, ()> + Send + Sync>;
pub struct Hooks {
before: HashMap<TypeId, Box<dyn Any + Send + Sync>>,
in_tx: HashMap<TypeId, Box<dyn Any + Send + Sync>>,
after: HashMap<TypeId, Box<dyn Any + Send + Sync>>,
on_start: Vec<StartFn>,
on_shutdown: Vec<ShutdownFn>,
timeout: Duration,
}
impl Default for Hooks {
fn default() -> Self {
Self {
before: HashMap::new(),
in_tx: HashMap::new(),
after: HashMap::new(),
on_start: Vec::new(),
on_shutdown: Vec::new(),
timeout: Duration::from_secs(2),
}
}
}
impl fmt::Debug for Hooks {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Hooks")
.field("before_events", &self.before.len())
.field("in_tx_events", &self.in_tx.len())
.field("after_events", &self.after.len())
.field("on_start", &self.on_start.len())
.field("on_shutdown", &self.on_shutdown.len())
.field("timeout", &self.timeout)
.finish()
}
}
pub(crate) enum Outcome<T> {
Done(T),
TimedOut,
Panicked,
}
pub(crate) async fn guarded<T>(timeout: Duration, future: BoxFuture<'_, T>) -> Outcome<T> {
match tokio::time::timeout(timeout, AssertUnwindSafe(future).catch_unwind()).await {
Ok(Ok(value)) => Outcome::Done(value),
Ok(Err(_panic)) => Outcome::Panicked,
Err(_) => Outcome::TimedOut,
}
}
impl Hooks {
pub(crate) fn set_timeout(&mut self, timeout: Duration) {
self.timeout = timeout;
}
pub fn before<E, F, Fut>(&mut self, hook: F)
where
E: Event,
F: Fn(HookCtx, E) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Decision<E>, AppError>> + Send + 'static,
{
let hook = Arc::new(hook);
let hook: BeforeFn<E> = Arc::new(move |ctx, event| {
let hook = hook.clone();
Box::pin(async move { hook(ctx, event).await })
});
let list = self.before.entry(TypeId::of::<E>()).or_insert_with(|| Box::new(Vec::<BeforeFn<E>>::new()));
if let Some(list) = list.downcast_mut::<Vec<BeforeFn<E>>>() {
list.push(hook);
}
}
pub fn in_tx<E, F>(&mut self, hook: F)
where
E: Event,
F: for<'a> Fn(&'a mut DbTx, &'a HookCtx, &'a E) -> BoxFuture<'a, Result<(), AppError>> + Send + Sync + 'static,
{
let hook: InTxFn<E> = Arc::new(hook);
let list = self.in_tx.entry(TypeId::of::<E>()).or_insert_with(|| Box::new(Vec::<InTxFn<E>>::new()));
if let Some(list) = list.downcast_mut::<Vec<InTxFn<E>>>() {
list.push(hook);
}
}
pub fn after<E, F, Fut>(&mut self, hook: F)
where
E: Event,
F: Fn(HookCtx, Arc<E>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), AppError>> + Send + 'static,
{
let hook = Arc::new(hook);
let hook: AfterFn<E> = Arc::new(move |ctx, event| {
let hook = hook.clone();
Box::pin(async move { hook(ctx, event).await })
});
let list = self.after.entry(TypeId::of::<E>()).or_insert_with(|| Box::new(Vec::<AfterFn<E>>::new()));
if let Some(list) = list.downcast_mut::<Vec<AfterFn<E>>>() {
list.push(hook);
}
}
pub fn on_start<F, Fut>(&mut self, hook: F)
where
F: Fn(HookCtx) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), AppError>> + Send + 'static,
{
let hook = Arc::new(hook);
self.on_start.push(Arc::new(move |ctx| {
let hook = hook.clone();
Box::pin(async move { hook(ctx).await })
}));
}
pub fn on_shutdown<F, Fut>(&mut self, hook: F)
where
F: Fn(HookCtx) -> Fut + Send + Sync + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
let hook = Arc::new(hook);
self.on_shutdown.push(Arc::new(move |ctx| {
let hook = hook.clone();
Box::pin(async move { hook(ctx).await })
}));
}
pub fn before_count<E: Event>(&self) -> usize {
self.before.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<BeforeFn<E>>>()).map_or(0, Vec::len)
}
pub fn in_tx_count<E: Event>(&self) -> usize {
self.in_tx.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<InTxFn<E>>>()).map_or(0, Vec::len)
}
pub async fn run_in_tx<E: Event>(&self, tx: &mut DbTx, ctx: &HookCtx, event: &E) -> Result<(), AppError> {
let Some(list) = self.in_tx.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<InTxFn<E>>>()) else {
return Ok(());
};
for (index, hook) in list.iter().enumerate() {
let call: BoxFuture<'_, Result<(), AppError>> = Box::pin(async { hook(&mut *tx, ctx, event).await });
match guarded(self.timeout, call).await {
Outcome::Done(Ok(())) => {}
Outcome::Done(Err(error)) => return Err(error),
Outcome::TimedOut => {
tracing::warn!(event = E::NAME, hook = index, "in_tx hook timed out");
return Err(AppError::new(codes::HOOK_TIMEOUT, "a server hook did not answer in time"));
}
Outcome::Panicked => {
tracing::error!(event = E::NAME, hook = index, "in_tx hook panicked");
return Err(AppError::internal_plain());
}
}
}
Ok(())
}
pub fn after_count<E: Event>(&self) -> usize {
self.after.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<AfterFn<E>>>()).map_or(0, Vec::len)
}
pub async fn run_before<E: Event>(&self, ctx: &HookCtx, mut event: E) -> Result<E, AppError> {
let Some(list) = self.before.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<BeforeFn<E>>>()) else {
return Ok(event);
};
for (index, hook) in list.iter().enumerate() {
match guarded(self.timeout, hook(ctx.clone(), event)).await {
Outcome::Done(Ok(Decision::Continue(next))) => event = next,
Outcome::Done(Ok(Decision::Reject(error))) => return Err(error),
Outcome::Done(Err(error)) => return Err(error),
Outcome::TimedOut => {
tracing::warn!(event = E::NAME, hook = index, "before hook timed out");
return Err(AppError::new(codes::HOOK_TIMEOUT, "a server hook did not answer in time"));
}
Outcome::Panicked => {
tracing::error!(event = E::NAME, hook = index, "before hook panicked");
return Err(AppError::internal_plain());
}
}
}
Ok(event)
}
pub async fn run_after<E: Event>(&self, ctx: &HookCtx, event: Arc<E>) {
let Some(list) = self.after.get(&TypeId::of::<E>()).and_then(|l| l.downcast_ref::<Vec<AfterFn<E>>>()) else {
return;
};
for (index, hook) in list.iter().enumerate() {
match guarded(self.timeout, hook(ctx.clone(), event.clone())).await {
Outcome::Done(Ok(())) => {}
Outcome::Done(Err(error)) => tracing::warn!(event = E::NAME, hook = index, %error, "after hook failed"),
Outcome::TimedOut => tracing::warn!(event = E::NAME, hook = index, "after hook timed out"),
Outcome::Panicked => tracing::error!(event = E::NAME, hook = index, "after hook panicked"),
}
}
}
pub(crate) async fn run_start(&self, ctx: &HookCtx) -> Result<(), String> {
for (index, hook) in self.on_start.iter().enumerate() {
match guarded(self.timeout.max(Duration::from_secs(30)), hook(ctx.clone())).await {
Outcome::Done(Ok(())) => {}
Outcome::Done(Err(error)) => return Err(format!("start hook {index} failed: {error}")),
Outcome::TimedOut => return Err(format!("start hook {index} timed out")),
Outcome::Panicked => return Err(format!("start hook {index} panicked")),
}
}
Ok(())
}
pub(crate) async fn run_shutdown(&self, ctx: &HookCtx) {
for (index, hook) in self.on_shutdown.iter().enumerate() {
match guarded(self.timeout.max(Duration::from_secs(10)), hook(ctx.clone())).await {
Outcome::Done(()) => {}
Outcome::TimedOut => tracing::warn!(hook = index, "shutdown hook timed out"),
Outcome::Panicked => tracing::error!(hook = index, "shutdown hook panicked"),
}
}
}
}