use std::any::{Any, TypeId};
use std::collections::HashSet;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex, Weak};
use tokio::sync::{mpsc, watch};
use tokio_util::sync::CancellationToken;
use crate::ctx::{Ctx, Shared};
use crate::effect::{Effect, EffectRecord};
use crate::error::{aggregate_arcs, panic_error, CordisError};
use crate::event::{CatchUnwind, Event};
use crate::key::{ScopeId, TypeKey};
use crate::{BoxFuture, Plugin, PluginFactory};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FiberState {
Pending,
Loading,
Active,
Failed,
Disposed,
Unloading,
}
#[derive(Debug, Clone)]
pub struct Snapshot {
pub generation: u64,
pub state: FiberState,
pub error: Option<Arc<CordisError>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct PluginId(pub u64);
#[derive(Debug, Clone)]
pub struct FiberStatusChanged {
pub plugin_id: PluginId,
pub seq: u64,
pub generation: u64,
pub from: FiberState,
pub to: FiberState,
}
impl Event for FiberStatusChanged {
const NAME: &'static str = "rutis::FiberStatusChanged";
type Value = ();
}
pub(crate) enum TaskDone {
Running,
Done(Option<Arc<CordisError>>),
}
pub(crate) struct TransitionTask {
pub done_tx: watch::Sender<TaskDone>,
_keepalive: watch::Receiver<TaskDone>,
}
impl TransitionTask {
pub(crate) fn new() -> Arc<Self> {
let (done_tx, done_rx) = watch::channel(TaskDone::Running);
Arc::new(Self {
done_tx,
_keepalive: done_rx,
})
}
pub(crate) fn complete(&self, err: Option<Arc<CordisError>>) {
let _ = self.done_tx.send(TaskDone::Done(err));
}
}
pub(crate) enum Intent {
RefreshDeps,
RefreshDepsJoin(Arc<TransitionTask>),
Settle(Arc<TransitionTask>),
Restart(Arc<TransitionTask>),
Dispose,
}
impl Intent {
fn task(&self) -> Option<&Arc<TransitionTask>> {
match self {
Intent::RefreshDepsJoin(task) | Intent::Settle(task) | Intent::Restart(task) => {
Some(task)
}
_ => None,
}
}
}
fn complete_intent(intent: &Intent, err: Option<Arc<CordisError>>) {
if let Some(task) = intent.task() {
task.complete(err);
}
}
enum NextState {
Pending,
Disposed,
}
pub(crate) struct Trans {
pub generation: u64,
pub state: FiberState,
pub error: Option<Arc<CordisError>>,
pub seq: u64,
pub status_queue: Vec<FiberStatusChanged>,
pub terminal_task: Option<Arc<TransitionTask>>,
}
pub(crate) struct FiberInner {
pub id: PluginId,
pub name: String,
pub is_root: bool,
pub plugin: Option<Arc<dyn Plugin>>,
pub factory: Option<Arc<dyn ErasedFactory>>,
pub config: Mutex<Option<Arc<dyn Any + Send + Sync>>>,
pub ctx: Ctx,
pub parent_fiber: Option<Weak<FiberInner>>,
pub token: Mutex<CancellationToken>,
pub transition: Mutex<Trans>,
pub snapshot_tx: watch::Sender<Snapshot>,
pub snapshot_rx: watch::Receiver<Snapshot>,
pub effects: Mutex<Vec<Arc<EffectRecord>>>,
pub drained_errors: Mutex<Vec<Arc<CordisError>>>,
pub mount: Mutex<Option<Arc<EffectRecord>>>,
pub declared_injects: Vec<TypeKey>,
pub intents_tx: mpsc::UnboundedSender<Intent>,
pub alive: AtomicBool,
pub provided: Mutex<Vec<(TypeKey, Option<crate::key::ScopeId>)>>,
pub last_deps: Mutex<Option<HashSet<(PluginId, u64, TypeKey, Option<crate::key::ScopeId>)>>>,
}
pub(crate) trait ErasedFactory: Send + Sync + 'static {
fn config_type_id(&self) -> TypeId;
fn name(&self) -> &str;
fn injects(&self) -> &[TypeKey];
fn validate_config_erased(&self, config: &dyn Any) -> Result<(), CordisError>;
fn build_erased(&self, config: &dyn Any) -> Result<Box<dyn Plugin>, CordisError>;
}
struct FactoryAdapter<F, C> {
factory: F,
_marker: std::marker::PhantomData<fn() -> C>,
}
impl<F, C> ErasedFactory for FactoryAdapter<F, C>
where
F: PluginFactory<C>,
C: Send + Sync + 'static,
{
fn config_type_id(&self) -> TypeId {
TypeId::of::<C>()
}
fn name(&self) -> &str {
self.factory.name()
}
fn injects(&self) -> &[TypeKey] {
self.factory.injects()
}
fn validate_config_erased(&self, config: &dyn Any) -> Result<(), CordisError> {
config
.downcast_ref::<C>()
.ok_or_else(|| CordisError::Validation {
issues: vec!["config type mismatch".into()],
})
.and_then(|c| self.factory.validate_config(c))
}
fn build_erased(&self, config: &dyn Any) -> Result<Box<dyn Plugin>, CordisError> {
config
.downcast_ref::<C>()
.ok_or_else(|| CordisError::Validation {
issues: vec!["config type mismatch".into()],
})
.and_then(|c| self.factory.build(c))
}
}
impl FiberInner {
fn current_plugin(&self) -> Result<Arc<dyn Plugin>, CordisError> {
if let Some(plugin) = &self.plugin {
return Ok(plugin.clone());
}
let factory = self
.factory
.as_ref()
.ok_or_else(|| CordisError::PluginFailed("fiber has no plugin or factory".into()))?;
let config = self
.config
.lock()
.unwrap()
.clone()
.ok_or_else(|| CordisError::PluginFailed("factory fiber has no config".into()))?;
let built = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
factory.build_erased(config.as_ref())
}))
.unwrap_or_else(|p| Err(panic_error(p)));
built.map(Arc::from)
}
pub(crate) fn current_token(&self) -> CancellationToken {
self.token.lock().unwrap().clone()
}
pub(crate) fn cancel_current(&self) {
self.token.lock().unwrap().cancel();
}
fn new_generation_token(&self) -> CancellationToken {
let token = CancellationToken::new();
*self.token.lock().unwrap() = token.clone();
token
}
pub(crate) fn post(&self, intent: Intent) -> bool {
if !self.alive.load(Ordering::SeqCst) {
complete_intent(&intent, None);
return false;
}
match self.intents_tx.send(intent) {
Ok(()) => true,
Err(err) => {
complete_intent(&err.0, None);
false
}
}
}
pub(crate) fn post_join(
&self,
task: Arc<TransitionTask>,
make: impl FnOnce(Arc<TransitionTask>) -> Intent,
) -> bool {
if !self.alive.load(Ordering::SeqCst) {
let err = self.transition.lock().unwrap().error.clone();
task.complete(err);
return false;
}
match self.intents_tx.send(make(task.clone())) {
Ok(()) => {
if !self.alive.load(Ordering::SeqCst) {
let err = self.transition.lock().unwrap().error.clone();
task.complete(err);
}
true
}
Err(_) => {
task.complete(None);
false
}
}
}
pub(crate) fn state(&self) -> FiberState {
self.transition.lock().unwrap().state
}
pub(crate) fn state_snapshot(&self) -> Snapshot {
let tr = self.transition.lock().unwrap();
Snapshot {
generation: tr.generation,
state: tr.state,
error: tr.error.clone(),
}
}
fn set_state(&self, tr: &mut Trans, new: FiberState) {
let old = tr.state;
if old == new {
return;
}
tr.state = new;
tr.seq += 1;
tr.status_queue.push(FiberStatusChanged {
plugin_id: self.id,
seq: tr.seq,
generation: tr.generation,
from: old,
to: new,
});
let _ = self.snapshot_tx.send(Snapshot {
generation: tr.generation,
state: new,
error: tr.error.clone(),
});
}
fn flush_status(&self) {
let queue: Vec<FiberStatusChanged> = {
let mut tr = self.transition.lock().unwrap();
std::mem::take(&mut tr.status_queue)
};
for event in queue {
self.ctx.shared().bus.emit(&self.ctx, Arc::new(event));
}
}
fn resolve_deps(
&self,
) -> (
HashSet<(PluginId, u64, TypeKey, Option<ScopeId>)>,
Vec<TypeKey>,
) {
let mut satisfied: HashSet<(PluginId, u64, TypeKey, Option<ScopeId>)> = HashSet::new();
let mut missing: Vec<TypeKey> = Vec::new();
let inject_keys: Vec<TypeKey> = if let Some(plugin) = &self.plugin {
plugin.injects().to_vec()
} else if let Some(factory) = &self.factory {
factory.injects().to_vec()
} else {
return (satisfied, missing);
};
let registry = &self.ctx.shared().registry;
for key in inject_keys {
let scope = self.ctx.scope_for(&key);
match registry.resolve_dep(&key, scope.as_ref()) {
Some((provider_id, provider_gen)) => {
satisfied.insert((provider_id, provider_gen, key, scope));
}
None => missing.push(key),
}
}
(satisfied, missing)
}
async fn refresh_deps(this: &Arc<Self>) {
if this.plugin.is_none() && this.factory.is_none() {
return;
}
if this.state() == FiberState::Disposed {
return;
}
let (satisfied, missing) = this.resolve_deps();
let loaded = !matches!(this.state(), FiberState::Pending | FiberState::Disposed);
if !missing.is_empty() {
if matches!(this.state(), FiberState::Active | FiberState::Loading) {
Self::unload(this, NextState::Pending).await;
}
return;
}
let unchanged = {
let last = this.last_deps.lock().unwrap();
last.as_ref() == Some(&satisfied)
};
if unchanged {
return;
}
if loaded {
Self::unload(this, NextState::Pending).await;
}
Self::load(this, satisfied).await;
}
async fn load(this: &Arc<Self>, deps: HashSet<(PluginId, u64, TypeKey, Option<ScopeId>)>) {
let shared = this.ctx.shared().clone();
this.new_generation_token();
{
let mut tr = this.transition.lock().unwrap();
tr.generation += 1;
tr.error = None;
Self::set_state(this, &mut tr, FiberState::Loading);
}
this.flush_status();
let plugin = match this.current_plugin() {
Ok(p) => p,
Err(e) => {
Self::fail_load(this, e).await;
return;
}
};
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| plugin.validate())) {
Ok(Ok(())) => {}
Ok(Err(e)) => {
Self::fail_load(this, e).await;
return;
}
Err(p) => {
Self::fail_load(this, panic_error(p)).await;
return;
}
}
*this.last_deps.lock().unwrap() = Some(deps.clone());
let ctx = this.ctx.clone();
let outcome = CatchUnwind::new(plugin.apply(&ctx)).await;
let result: Result<Effect, CordisError> = outcome.unwrap_or_else(|p| Err(panic_error(p)));
match result {
Ok(effect) => {
let record = EffectRecord::new(effect, Arc::downgrade(this));
this.effects.lock().unwrap().push(record);
let (fresh, missing) = this.resolve_deps();
if missing.is_empty() {
*this.last_deps.lock().unwrap() = Some(fresh);
}
{
let mut tr = this.transition.lock().unwrap();
tr.error = None;
Self::set_state(this, &mut tr, FiberState::Active);
}
let provided: Vec<TypeKey> = this
.provided
.lock()
.unwrap()
.iter()
.map(|(k, _)| k.clone())
.collect();
for key in provided {
shared.registry.notify_key_changed(&key);
}
}
Err(e) => Self::fail_load(this, e).await,
}
this.flush_status();
}
async fn fail_load(this: &Arc<Self>, error: CordisError) {
{
let mut tr = this.transition.lock().unwrap();
Self::set_state(this, &mut tr, FiberState::Unloading);
}
this.flush_status();
this.cancel_current();
let cleanup_errors = Self::drain_effects(this).await;
let sink = this.ctx.error_sink();
for e in cleanup_errors {
sink(e);
}
let arc = Arc::new(error);
{
let mut tr = this.transition.lock().unwrap();
tr.error = Some(arc);
Self::set_state(this, &mut tr, FiberState::Failed);
}
this.flush_status();
}
async fn drain_effects(this: &Arc<Self>) -> Vec<Arc<CordisError>> {
let handle = this.ctx.handle().clone();
let effects: Vec<Arc<EffectRecord>> = std::mem::take(&mut *this.effects.lock().unwrap());
let mut errors: Vec<Arc<CordisError>> = Vec::new();
for record in effects.into_iter().rev() {
if let Err(e) = record.drain(&handle).await {
errors.push(e);
}
}
errors.extend(std::mem::take(&mut *this.drained_errors.lock().unwrap()));
*this.last_deps.lock().unwrap() = None;
this.provided.lock().unwrap().clear();
errors
}
async fn unload(this: &Arc<Self>, next: NextState) {
{
let mut tr = this.transition.lock().unwrap();
Self::set_state(this, &mut tr, FiberState::Unloading);
}
this.flush_status();
this.cancel_current();
let errors = Self::drain_effects(this).await;
let err = aggregate_arcs(errors);
match next {
NextState::Pending => {
if let Some(e) = err {
(this.ctx.error_sink())(e);
}
let mut tr = this.transition.lock().unwrap();
Self::set_state(this, &mut tr, FiberState::Pending);
}
NextState::Disposed => {
let mut tr = this.transition.lock().unwrap();
tr.error = err;
Self::set_state(this, &mut tr, FiberState::Disposed);
}
}
this.flush_status();
}
}
pub(crate) async fn drive(this: Arc<FiberInner>, mut rx: mpsc::UnboundedReceiver<Intent>) {
while let Some(intent) = rx.recv().await {
match intent {
Intent::RefreshDeps => FiberInner::refresh_deps(&this).await,
Intent::RefreshDepsJoin(task) => {
FiberInner::refresh_deps(&this).await;
complete_task(&this, task);
}
Intent::Settle(task) => {
let err = {
let tr = this.transition.lock().unwrap();
(tr.state == FiberState::Failed)
.then(|| tr.error.clone())
.flatten()
};
task.complete(err);
}
Intent::Restart(task) => {
let state = this.state();
if state == FiberState::Disposed {
if this.is_root {
this.new_generation_token();
{
let mut tr = this.transition.lock().unwrap();
tr.error = None;
tr.terminal_task = None;
FiberInner::set_state(&this, &mut tr, FiberState::Active);
}
this.flush_status();
}
complete_task(&this, task);
continue;
}
if !matches!(state, FiberState::Pending) {
FiberInner::unload(&this, NextState::Pending).await;
}
FiberInner::refresh_deps(&this).await;
complete_task(&this, task);
}
Intent::Dispose => {
if this.state() != FiberState::Disposed {
FiberInner::unload(&this, NextState::Disposed).await;
}
let (task, err) = {
let tr = this.transition.lock().unwrap();
(tr.terminal_task.clone(), tr.error.clone())
};
if let Some(task) = task {
let _ = task.done_tx.send(TaskDone::Done(err));
}
if !this.is_root {
this.alive.store(false, Ordering::SeqCst);
while let Ok(intent) = rx.try_recv() {
complete_intent(&intent, None);
}
release_transient(&this).await;
return; }
}
}
this.flush_status();
}
}
fn complete_task(this: &Arc<FiberInner>, task: Arc<TransitionTask>) {
let err = {
let tr = this.transition.lock().unwrap();
tr.error.clone()
};
task.complete(err);
}
async fn release_transient(this: &Arc<FiberInner>) {
let shared = this.ctx.shared().clone();
shared
.registry
.unregister_injects(this, &this.declared_injects);
let record = this.mount.lock().unwrap().take();
if let Some(record) = record {
let handle = this.ctx.handle().clone();
if let Err(e) = record.drain(&handle).await {
(this.ctx.error_sink())(e);
}
}
}
pub(crate) fn spawn_fiber(
shared: &Arc<Shared>,
parent_ctx: Option<&Ctx>,
plugin: Option<Arc<dyn Plugin>>,
is_root: bool,
) -> Arc<FiberInner> {
spawn_fiber_inner(shared, parent_ctx, plugin, None, None, is_root)
}
pub(crate) fn spawn_factory_fiber<F, C>(
shared: &Arc<Shared>,
parent_ctx: Option<&Ctx>,
factory: F,
config: C,
is_root: bool,
) -> Arc<FiberInner>
where
F: PluginFactory<C>,
C: Send + Sync + 'static,
{
let erased: Arc<dyn ErasedFactory> = Arc::new(FactoryAdapter {
factory,
_marker: std::marker::PhantomData,
});
let boxed: Arc<dyn Any + Send + Sync> = Arc::new(config);
spawn_fiber_inner(shared, parent_ctx, None, Some(erased), Some(boxed), is_root)
}
fn spawn_fiber_inner(
shared: &Arc<Shared>,
parent_ctx: Option<&Ctx>,
plugin: Option<Arc<dyn Plugin>>,
factory: Option<Arc<dyn ErasedFactory>>,
config: Option<Arc<dyn Any + Send + Sync>>,
is_root: bool,
) -> Arc<FiberInner> {
let id = PluginId(
shared
.next_plugin_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst),
);
let initial = if is_root {
FiberState::Active
} else {
FiberState::Pending
};
let declared_injects: Vec<TypeKey> = if let Some(plugin) = &plugin {
plugin.injects().to_vec()
} else if let Some(factory) = &factory {
factory.injects().to_vec()
} else {
Vec::new()
};
let (snapshot_tx, snapshot_rx) = watch::channel(Snapshot {
generation: 0,
state: initial,
error: None,
});
let (tx, rx) = mpsc::unbounded_channel();
let token = CancellationToken::new();
let parent_fiber = parent_ctx.map(|p| p.weak_fiber());
let this = Arc::new_cyclic(|weak: &Weak<FiberInner>| {
let ctx = match parent_ctx {
Some(parent) => Ctx::new_child(shared.clone(), parent, weak.clone(), None),
None => Ctx::new_root(shared.clone(), weak.clone()),
};
FiberInner {
id,
name: plugin
.as_ref()
.map(|p| p.name().to_string())
.or_else(|| factory.as_ref().map(|f| f.name().to_string()))
.unwrap_or_else(|| "root".to_string()),
is_root,
plugin,
factory,
config: Mutex::new(config),
ctx,
parent_fiber,
token: Mutex::new(token),
transition: Mutex::new(Trans {
generation: 0,
state: initial,
error: None,
seq: 0,
status_queue: Vec::new(),
terminal_task: None,
}),
snapshot_tx: snapshot_tx.clone(),
snapshot_rx,
effects: Mutex::new(Vec::new()),
drained_errors: Mutex::new(Vec::new()),
mount: Mutex::new(None),
declared_injects,
intents_tx: tx.clone(),
alive: AtomicBool::new(true),
provided: Mutex::new(Vec::new()),
last_deps: Mutex::new(None),
}
});
for key in &this.declared_injects {
shared.registry.register_inject(key.clone(), &this);
}
shared.handle.spawn(drive(this.clone(), rx));
this
}
pub struct FiberView {
pub id: PluginId,
inner: Arc<FiberInner>,
}
impl Clone for FiberView {
fn clone(&self) -> Self {
Self {
id: self.id,
inner: self.inner.clone(),
}
}
}
impl FiberView {
pub(crate) fn from_inner(inner: Arc<FiberInner>) -> Self {
Self {
id: inner.id,
inner,
}
}
pub fn state(&self) -> Snapshot {
self.inner.state_snapshot()
}
pub fn name(&self) -> &str {
&self.inner.name
}
pub fn watch(&self) -> watch::Receiver<Snapshot> {
self.inner.snapshot_rx.clone()
}
pub fn dispose(&self) -> BoxFuture<'static, Result<(), Arc<CordisError>>> {
let task = {
let mut tr = self.inner.transition.lock().unwrap();
match &tr.terminal_task {
Some(task) => task.clone(),
None => {
self.inner.cancel_current();
let task = TransitionTask::new();
tr.terminal_task = Some(task.clone());
self.inner.post(Intent::Dispose);
task
}
}
};
Box::pin(async move { join_task(&task).await })
}
pub fn restart(&self) -> BoxFuture<'static, Result<(), Arc<CordisError>>> {
let this = self.inner.clone();
Box::pin(async move {
{
let tr = this.transition.lock().unwrap();
if !this.is_root
&& (tr.terminal_task.is_some() || !this.alive.load(Ordering::SeqCst))
{
return Err(Arc::new(CordisError::InactiveEffect));
}
}
this.cancel_current(); let task = TransitionTask::new();
this.post_join(task.clone(), Intent::Restart);
join_task(&task).await
})
}
pub fn update<C: Send + Sync + 'static>(
&self,
new_config: C,
) -> BoxFuture<'static, Result<(), Arc<CordisError>>> {
let this = self.inner.clone();
Box::pin(async move {
{
let tr = this.transition.lock().unwrap();
if !this.is_root
&& (tr.terminal_task.is_some() || !this.alive.load(Ordering::SeqCst))
{
return Err(Arc::new(CordisError::InactiveEffect));
}
}
let factory = this
.factory
.as_ref()
.and_then(|f| (f.config_type_id() == TypeId::of::<C>()).then(|| f.clone()));
let Some(factory) = factory else {
let reason = if this.factory.is_some() {
"config type mismatch"
} else {
"fiber has no factory (static plugin cannot update)"
};
return Err(Arc::new(CordisError::Validation {
issues: vec![reason.into()],
}));
};
let boxed: Arc<dyn Any + Send + Sync> = Arc::new(new_config);
let dry = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
factory.validate_config_erased(boxed.as_ref())?;
let instance = factory.build_erased(boxed.as_ref())?;
instance.validate()?;
Ok::<(), CordisError>(())
}));
match dry {
Ok(Ok(())) => {}
Ok(Err(e)) => return Err(Arc::new(e)),
Err(p) => return Err(Arc::new(panic_error(p))),
}
{
let tr = this.transition.lock().unwrap();
if !this.is_root
&& (tr.terminal_task.is_some() || !this.alive.load(Ordering::SeqCst))
{
return Err(Arc::new(CordisError::InactiveEffect));
}
}
*this.config.lock().unwrap() = Some(boxed);
this.cancel_current(); let task = TransitionTask::new();
this.post_join(task.clone(), Intent::Restart);
join_task(&task).await
})
}
pub fn current_config<C: Send + Sync + 'static>(&self) -> Option<Arc<C>> {
let boxed = self.inner.config.lock().unwrap().clone()?;
boxed.downcast::<C>().ok()
}
}
pub(crate) async fn join_task(task: &Arc<TransitionTask>) -> Result<(), Arc<CordisError>> {
let mut rx = task.done_tx.subscribe();
loop {
match &*rx.borrow() {
TaskDone::Done(err) => return err.clone().map_or(Ok(()), Err),
TaskDone::Running => {}
}
if rx.changed().await.is_err() {
return Err(Arc::new(CordisError::PluginFailed(
"fiber dropped before transition completed".into(),
)));
}
}
}
fn settle(this: &FiberInner) -> BoxFuture<'static, Result<(), Arc<CordisError>>> {
let task = TransitionTask::new();
this.post_join(task.clone(), Intent::Settle);
Box::pin(async move { join_task(&task).await })
}
impl std::future::IntoFuture for FiberView {
type Output = Result<(), Arc<CordisError>>;
type IntoFuture = BoxFuture<'static, Result<(), Arc<CordisError>>>;
fn into_future(self) -> Self::IntoFuture {
settle(&self.inner)
}
}
impl std::future::IntoFuture for &FiberView {
type Output = Result<(), Arc<CordisError>>;
type IntoFuture = BoxFuture<'static, Result<(), Arc<CordisError>>>;
fn into_future(self) -> Self::IntoFuture {
settle(&self.inner)
}
}
#[cfg(test)]
mod transient_tests;