use atomic_try_update::barrier::ShutdownBarrier;
use atomic_try_update::once::OnceLockFree;
use log::RootLog;
use slog::{debug, error, o, warn, Logger};
use atomic_try_update::stack::Stack;
use std::any::type_name;
use std::borrow::Borrow;
use std::fmt::Display;
use std::sync::atomic::AtomicBool;
use std::sync::Arc;
pub mod log;
#[derive(Debug)]
pub enum FactoryError<E: std::fmt::Debug> {
Cancellation,
CyclicDependency(&'static str),
CouldNotConstruct(&'static str, E),
FieldAlreadySet(&'static str),
DaemonPrepareFailed(&'static str, E),
DaemonRunFailed(&'static str, E),
DaemonStopFailed(&'static str, E),
}
impl<E: std::fmt::Debug> Display for FactoryError<E> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{self:?}")
}
}
impl<E: std::fmt::Debug> std::error::Error for FactoryError<E> {}
pub type FactoryResult<T, E> = std::result::Result<T, FactoryError<E>>;
pub type DaemonResult<T, E> = std::result::Result<T, E>;
pub struct DaemonRef<T: ?Sized> {
inner: std::sync::Arc<T>,
}
impl<T: ?Sized> From<&std::sync::Arc<T>> for DaemonRef<T> {
fn from(val: &std::sync::Arc<T>) -> Self {
DaemonRef { inner: val.clone() }
}
}
impl<T: ?Sized> std::ops::Deref for DaemonRef<T> {
type Target = T;
#[inline]
fn deref(&self) -> &T {
&self.inner
}
}
impl<T: ?Sized> Borrow<T> for DaemonRef<T> {
fn borrow(&self) -> &T {
&self.inner
}
}
#[async_trait::async_trait]
pub trait Daemon<E>: Send + Sync {
fn name(&self) -> &'static str;
async fn prepare(&self) -> DaemonResult<(), E>;
async fn run(&self) -> DaemonResult<(), E>;
async fn stop(&self) -> DaemonResult<(), E>;
}
pub struct DaemonField<T> {
svc: OnceLockFree<Arc<T>>,
}
pub trait DaemonGetter<E>: Sync + Send {
fn get_svc(&self) -> &dyn Daemon<E>;
}
impl<T> Default for DaemonField<T> {
fn default() -> Self {
Self {
svc: Default::default(),
}
}
}
impl<T, E> DaemonGetter<E> for DaemonField<T>
where
T: Daemon<E> + 'static,
{
fn get_svc(&self) -> &dyn Daemon<E> {
self.svc.get_or_prepare_to_set().unwrap().unwrap().as_ref()
}
}
impl<T> Clone for DaemonField<T> {
fn clone(&self) -> Self {
let res = Self::default();
if let Some(val) = self.svc.get_or_seal().unwrap() {
res.svc.set(val.clone()).unwrap();
}
res
}
}
#[allow(clippy::type_complexity)]
pub struct Factory<'a, DaemonBundle: Default, E: Send + Sync> {
log: Logger,
daemon_lifecycles_accumulator: Stack<Box<dyn Send + Sync + Fn(&Self) -> &dyn DaemonGetter<E>>>,
daemon_lifecycles: Vec<Box<dyn Send + Sync + Fn(&Self) -> &dyn DaemonGetter<E>>>,
daemons: DaemonBundle,
stopping: AtomicBool,
}
impl<
'a,
DaemonBundle: Default + Clone + Sync + Send,
E: std::fmt::Debug + Send + Sync + 'static,
> Factory<'a, DaemonBundle, E>
{
pub fn new(root_log: DaemonRef<RootLog>) -> Self {
Self {
log: root_log.get_factory_logger("*"),
daemon_lifecycles_accumulator: Default::default(),
daemon_lifecycles: Default::default(),
daemons: Default::default(),
stopping: false.into(),
}
}
pub fn from_ambient(ambient: &Self, host: String) -> Self {
Self {
log: ambient.log.new(o!("host"=>host)),
daemon_lifecycles: Default::default(),
daemon_lifecycles_accumulator: Default::default(),
daemons: ambient.daemons.clone(),
stopping: false.into(),
}
}
pub fn build<T>(
&self,
field_accessor: impl 'static + Sync + Send + Fn(&DaemonBundle) -> &DaemonField<T>,
constructor: impl Fn(&Self) -> DaemonResult<T, E>,
) -> FactoryResult<DaemonRef<T>, E>
where
T: 'a + Daemon<E>,
DaemonField<T>: DaemonGetter<E>,
{
let svc = field_accessor(&self.daemons)
.svc
.get_or_prepare_to_set()
.map_err(|_| FactoryError::CyclicDependency(std::any::type_name::<T>()))?;
match svc {
Some(ret) => Ok(ret.into()),
None => {
let svc = Arc::new(constructor(self).map_err(|err| {
FactoryError::CouldNotConstruct(std::any::type_name::<T>(), err)
})?);
match field_accessor(&self.daemons).svc.set_prepared(svc) {
Ok(ret) => {
self.daemon_lifecycles_accumulator
.push(Box::new(move |factory| field_accessor(&factory.daemons)));
Ok(ret.into())
}
Err(err) => {
panic!(
"Unexpected failure to set field of type {} error: {err:?}",
type_name::<T>()
);
}
}
}
}
}
pub fn inject<T>(
&self,
field_accessor: impl 'static + Sync + Send + Fn(&DaemonBundle) -> &DaemonField<T>,
val: Arc<T>,
) -> FactoryResult<DaemonRef<T>, E>
where
T: 'a + Daemon<E>,
DaemonField<T>: DaemonGetter<E>,
{
Ok(field_accessor(&self.daemons)
.svc
.set(val)
.map_err(|_| FactoryError::FieldAlreadySet(std::any::type_name::<T>()))?
.into())
}
pub fn finalize_daemons(&mut self) {
_ = self.daemons.clone(); for svc in self.daemon_lifecycles_accumulator.pop_all() {
self.daemon_lifecycles.push(svc);
}
self.daemon_lifecycles.reverse();
}
pub async fn prepare(&self) -> FactoryResult<(), E> {
for svc in &self.daemon_lifecycles {
debug!(self.log, "prepare starts"; "target_daemon"=>svc(self).get_svc().name());
let res = svc(self).get_svc().prepare().await;
if let Err(err) = res {
let err = FactoryError::DaemonPrepareFailed(svc(self).get_svc().name(), err);
return Err(err);
} else {
debug!(self.log, "prepare done"; "target_daemon"=>svc(self).get_svc().name(), "res"=>#?res);
}
}
Ok(())
}
pub fn run<'b>(
&'b self,
scope: &mut async_scoped::Scope<'b, Result<(), FactoryError<E>>, async_scoped::Tokio>,
shutdown: Option<&'b ShutdownBarrier>,
) where
'a: 'b,
{
for svc in &self.daemon_lifecycles {
debug!(self.log, "run"; "target_daemon"=>svc(self).get_svc().name());
scope.spawn(
async move {
match svc(self).get_svc().run().await {
Ok(()) => Ok(()),
Err(e) => {
error!(self.log, "daemon run failed"; "name"=>svc(self).get_svc().name(), "err"=>#?e);
if let Some(shutdown) = shutdown {
warn!(self.log, "daemon run failed; broadcasting shutdown");
_ = shutdown.cancel();
} else {
warn!(self.log, "daemon run failed; not configured for shutdown broadcast");
}
if let Err(e) = self.stop().await {
error!(self.log, "additional error shutting down factory after run error"; "svc"=>svc(self).get_svc().name(), "err"=>#?e);
}
Err(FactoryError::DaemonRunFailed(svc(self).get_svc().name(), e))
}
}
});
}
}
pub async fn stop(&self) -> FactoryResult<(), E> {
if self
.stopping
.swap(true, std::sync::atomic::Ordering::SeqCst)
{
return Ok(());
}
let mut res = Ok(());
for svc_idx in (0..self.daemon_lifecycles.len()).rev() {
let svc_name = self.daemon_lifecycles[svc_idx](self).get_svc().name();
debug!(self.log, "stop"; "target_daemon"=>svc_name);
match self.daemon_lifecycles[svc_idx](self).get_svc().stop().await {
Ok(()) => (),
Err(new_err) => {
let new_err = FactoryError::DaemonStopFailed(svc_name, new_err);
res = match res {
Err(old_err) => Err(old_err), Ok(()) => Err(new_err),
}
}
}
}
res
}
}