use std::panic::AssertUnwindSafe;
use std::sync::{Arc, Mutex, Weak};
use tokio::sync::Notify;
use crate::error::{aggregate_arcs, join_panic_error, panic_error, CordisError};
use crate::fiber::FiberInner;
use crate::BoxFuture;
pub enum Effect {
Done,
Disposer(Box<dyn FnOnce() -> Result<(), CordisError> + Send>),
AsyncDisposer(Box<dyn FnOnce() -> BoxFuture<'static, Result<(), CordisError>> + Send>),
Many(Vec<Effect>),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EffectPhase {
Live,
Draining,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct EffectMeta {
pub label: String,
pub phase: EffectPhase,
pub children: Vec<EffectMeta>,
}
impl EffectMeta {
fn set_phase(&mut self, phase: EffectPhase) {
self.phase = phase;
for child in &mut self.children {
child.set_phase(phase);
}
}
}
impl Effect {
fn kind(&self) -> &'static str {
match self {
Self::Done => "done",
Self::Disposer(_) => "disposer",
Self::AsyncDisposer(_) => "async disposer",
Self::Many(_) => "many",
}
}
fn into_cleanups(self, label: String, out: &mut Vec<Cleanup>) -> EffectMeta {
let mut children = Vec::new();
match self {
Effect::Done => {}
Effect::Disposer(f) => out.push(Cleanup::Sync(f)),
Effect::AsyncDisposer(f) => out.push(Cleanup::Async(f)),
Effect::Many(v) => {
for (index, effect) in v.into_iter().enumerate() {
let child_label = format!("{index}: {}", effect.kind());
children.push(effect.into_cleanups(child_label, out));
}
}
}
EffectMeta {
label,
phase: EffectPhase::Live,
children,
}
}
}
enum Cleanup {
Sync(Box<dyn FnOnce() -> Result<(), CordisError> + Send>),
Async(Box<dyn FnOnce() -> BoxFuture<'static, Result<(), CordisError>> + Send>),
}
pub(crate) struct EffectRecord {
st: Mutex<EffectState>,
metadata: EffectMeta,
notify: Notify,
owner: Option<Weak<FiberInner>>,
}
enum EffectState {
Live(Vec<Cleanup>),
Draining,
Done(Option<Arc<CordisError>>),
}
impl EffectRecord {
pub(crate) fn new(effect: Effect, label: String, owner: Weak<FiberInner>) -> Arc<Self> {
let mut cleanups = Vec::new();
let metadata = effect.into_cleanups(label, &mut cleanups);
Arc::new(Self {
st: Mutex::new(EffectState::Live(cleanups)),
metadata,
notify: Notify::new(),
owner: Some(owner),
})
}
pub(crate) fn snapshot(&self) -> Option<EffectMeta> {
let phase = match &*self.st.lock().unwrap() {
EffectState::Live(_) => EffectPhase::Live,
EffectState::Draining => EffectPhase::Draining,
EffectState::Done(_) => return None,
};
let mut metadata = self.metadata.clone();
metadata.set_phase(phase);
Some(metadata)
}
pub(crate) async fn drain(
self: &Arc<Self>,
handle: &tokio::runtime::Handle,
) -> Result<(), Arc<CordisError>> {
let claimed = {
let mut st = self.st.lock().unwrap();
match &mut *st {
EffectState::Live(_) => {
let live = std::mem::replace(&mut *st, EffectState::Draining);
match live {
EffectState::Live(cleanups) => Some(cleanups),
_ => unreachable!(),
}
}
EffectState::Draining | EffectState::Done(_) => None,
}
};
if let Some(mut cleanups) = claimed {
let this = self.clone();
let runner = handle.clone();
handle.spawn(async move {
let err = run_cleanups(&mut cleanups, &runner).await;
*this.st.lock().unwrap() = EffectState::Done(err.clone());
if let Some(owner) = this.owner.as_ref().and_then(Weak::upgrade) {
let mut effects = owner.effects.lock().unwrap();
let present = effects.iter().any(|r| Arc::ptr_eq(r, &this));
if present {
effects.retain(|r| !Arc::ptr_eq(r, &this));
if effects.capacity() > 64 && effects.len() * 4 < effects.capacity() {
effects.shrink_to_fit();
}
if let Some(e) = &err {
owner.drained_errors.lock().unwrap().push(e.clone());
}
}
drop(effects);
let weak = Arc::downgrade(&this);
let mut index = owner.effect_index.lock().unwrap();
index.retain(|entry| !Weak::ptr_eq(entry, &weak) && entry.strong_count() > 0);
if index.capacity() > 64 && index.len() * 4 < index.capacity() {
index.shrink_to_fit();
}
}
this.notify.notify_waiters();
});
}
self.join().await
}
async fn join(&self) -> Result<(), Arc<CordisError>> {
loop {
let notified = self.notify.notified();
{
let st = self.st.lock().unwrap();
if let EffectState::Done(e) = &*st {
return match e.clone() {
Some(e) => Err(e),
None => Ok(()),
};
}
}
notified.await;
}
}
}
async fn run_cleanups(
cleanups: &mut Vec<Cleanup>,
handle: &tokio::runtime::Handle,
) -> Option<Arc<CordisError>> {
let mut errors: Vec<CordisError> = Vec::new();
while let Some(cleanup) = cleanups.pop() {
let result = match cleanup {
Cleanup::Sync(f) => match std::panic::catch_unwind(AssertUnwindSafe(f)) {
Ok(r) => r,
Err(p) => Err(panic_error(p)),
},
Cleanup::Async(f) => {
let fut = match std::panic::catch_unwind(AssertUnwindSafe(f)) {
Ok(fut) => fut,
Err(p) => {
errors.push(panic_error(p));
continue;
}
};
let joined = handle.spawn(fut);
match joined.await {
Ok(r) => r,
Err(join_err) => Err(join_panic_error(join_err)),
}
}
};
if let Err(e) = result {
errors.push(e);
}
}
aggregate_arcs(errors.into_iter().map(Arc::new).collect())
}
pub struct Disposer {
run: Option<Box<dyn FnOnce() -> BoxFuture<'static, Result<(), Arc<CordisError>>> + Send>>,
}
impl Disposer {
pub(crate) fn new(
run: Box<dyn FnOnce() -> BoxFuture<'static, Result<(), Arc<CordisError>>> + Send>,
) -> Self {
Self { run: Some(run) }
}
pub fn dispose(mut self) -> BoxFuture<'static, Result<(), Arc<CordisError>>> {
match self.run.take() {
Some(run) => run(),
None => Box::pin(async { Ok(()) }),
}
}
}
impl std::fmt::Debug for Disposer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("Disposer")
}
}