use std::future::Future;
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Weak};
use tokio::runtime::Handle;
use tokio_util::sync::CancellationToken;
use crate::bus::EventBus;
use crate::effect::{Disposer, Effect, EffectRecord};
use crate::error::{default_sink, CordisError, ErrorSink};
use crate::fiber::{
join_task, spawn_fiber, FiberInner, FiberState, FiberView, Intent, TransitionTask,
};
use crate::key::{ScopeId, TypeKey};
use crate::registry::{Binding, CheckFn, Registry, StoredValue};
use crate::Plugin;
use crate::PluginFactory;
pub(crate) struct Shared {
pub handle: Handle,
pub bus: EventBus,
pub registry: Registry,
pub error_sink: ErrorSink,
pub next_plugin_id: AtomicU64,
}
pub(crate) struct CtxInner {
pub(crate) shared: Arc<Shared>,
pub(crate) parent: Option<Ctx>,
pub(crate) fiber: Weak<FiberInner>,
pub(crate) isolate: Option<(TypeKey, ScopeId)>,
}
#[derive(Clone)]
pub struct Ctx(Arc<CtxInner>);
impl Ctx {
pub(crate) fn new_child(
shared: Arc<Shared>,
parent: &Ctx,
fiber: Weak<FiberInner>,
isolate: Option<(TypeKey, ScopeId)>,
) -> Self {
Self(Arc::new(CtxInner {
shared,
parent: Some(parent.clone()),
fiber,
isolate,
}))
}
pub(crate) fn new_root(shared: Arc<Shared>, fiber: Weak<FiberInner>) -> Self {
Self(Arc::new(CtxInner {
shared,
parent: None,
fiber,
isolate: None,
}))
}
pub(crate) fn weak_fiber(&self) -> Weak<FiberInner> {
self.0.fiber.clone()
}
pub(crate) fn shared(&self) -> &Arc<Shared> {
&self.0.shared
}
pub fn handle(&self) -> &Handle {
&self.0.shared.handle
}
pub fn error_sink(&self) -> ErrorSink {
self.0.shared.error_sink.clone()
}
pub fn root() -> Result<Ctx, CordisError> {
let handle = Handle::try_current().map_err(|_| {
CordisError::PluginFailed(
"no tokio runtime in scope; construct inside #[tokio::test]/runtime, or use Ctx::root_with(handle)".into(),
)
})?;
Ok(Self::root_with_sink(handle, default_sink()))
}
pub fn root_with(handle: Handle) -> Ctx {
Self::root_with_sink(handle, default_sink())
}
pub fn root_with_sink(handle: Handle, sink: ErrorSink) -> Ctx {
let shared = Arc::new(Shared {
handle,
bus: EventBus::new(),
registry: Registry::new(),
error_sink: sink,
next_plugin_id: AtomicU64::new(1),
});
let root_fiber = spawn_fiber(&shared, None, None, true);
root_fiber.ctx.clone()
}
pub fn root_view(&self) -> FiberView {
let mut current = self.clone();
while let Some(parent) = current.0.parent.clone() {
current = parent;
}
let fiber = current.0.fiber.upgrade().expect("root fiber alive");
FiberView::from_inner(fiber)
}
pub fn events(&self) -> &EventBus {
&self.0.shared.bus
}
pub(crate) fn scope_for(&self, key: &TypeKey) -> Option<ScopeId> {
let mut current = Some(self.clone());
while let Some(ctx) = current {
if let Some((k, scope)) = &ctx.0.isolate {
if k == key {
return Some(scope.clone());
}
}
current = ctx.0.parent.clone();
}
None
}
pub fn isolate(&self, key: impl Into<TypeKey>, label: &str) -> Ctx {
Ctx::new_child(
self.0.shared.clone(),
self,
self.0.fiber.clone(),
Some((key.into(), Arc::from(label))),
)
}
pub fn get<T: Send + Sync + 'static>(&self) -> Option<Arc<T>> {
self.get_as::<T>(TypeKey::of::<T>())
}
pub fn get_as<T: ?Sized + Send + Sync + 'static>(
&self,
key: impl Into<TypeKey>,
) -> Option<Arc<T>> {
let key = key.into();
let scope = self.scope_for(&key);
let binding = self.0.shared.registry.lookup(&key, scope.as_ref())?;
let provider = binding.provider.upgrade()?;
let self_access = self.in_subtree_of(&provider);
if !self_access {
match self.0.fiber.upgrade() {
None => return None,
Some(accessor)
if matches!(
accessor.state(),
FiberState::Unloading | FiberState::Disposed
) =>
{
return None;
}
_ => {}
}
let visible = provider.state() == FiberState::Active
&& !binding.removing.load(std::sync::atomic::Ordering::SeqCst);
if !visible {
return None;
}
}
binding.value.downcast::<T>()
}
fn in_subtree_of(&self, other: &Arc<FiberInner>) -> bool {
let mut current = self.0.fiber.upgrade();
while let Some(fiber) = current {
if Arc::ptr_eq(&fiber, other) {
return true;
}
current = fiber.parent_fiber.as_ref().and_then(|w| w.upgrade());
}
false
}
pub fn provide<T: Send + Sync + 'static>(&self, value: T) -> Result<Disposer, CordisError> {
self.provide_as::<T>(TypeKey::of::<T>(), Arc::new(value))
}
pub fn provide_as<T: ?Sized + Send + Sync + 'static>(
&self,
key: impl Into<TypeKey>,
value: Arc<T>,
) -> Result<Disposer, CordisError> {
self.provide_inner(key.into(), value, None)
}
pub fn provide_as_with_check<T: ?Sized + Send + Sync + 'static>(
&self,
key: impl Into<TypeKey>,
value: Arc<T>,
check: impl Fn() -> bool + Send + Sync + 'static,
) -> Result<Disposer, CordisError> {
self.provide_inner(key.into(), value, Some(Arc::new(check)))
}
fn provide_inner<T: ?Sized + Send + Sync + 'static>(
&self,
key: TypeKey,
value: Arc<T>,
check: Option<CheckFn>,
) -> Result<Disposer, CordisError> {
let fiber = self.0.fiber.upgrade().ok_or(CordisError::InactiveEffect)?;
{
let tr = fiber.transition.lock().unwrap();
if matches!(tr.state, FiberState::Unloading | FiberState::Disposed) {
return Err(CordisError::InactiveEffect);
}
}
let scope = self.scope_for(&key);
let provider_gen = {
let tr = fiber.transition.lock().unwrap();
tr.generation
};
let stored = self.0.shared.registry.insert_binding(
key.clone(),
scope.clone(),
Binding {
value: StoredValue::new(value),
provider: Arc::downgrade(&fiber),
provider_id: fiber.id,
provider_gen,
check,
removing: std::sync::atomic::AtomicBool::new(false),
},
)?;
fiber
.provided
.lock()
.unwrap()
.push((key.clone(), scope.clone()));
if fiber.state() == FiberState::Active {
self.0.shared.registry.notify_key_changed(&key);
}
let shared = self.0.shared.clone();
let pid = fiber.id;
let provider = Arc::downgrade(&fiber);
let evict_scope = scope.clone();
let evict_key = key.clone();
match self.effect(move || {
Effect::AsyncDisposer(Box::new(move || {
let shared = shared.clone();
let provider = provider.clone();
let scope = evict_scope.clone();
Box::pin(async move {
evict_and_finalize(&shared, provider, pid, provider_gen, evict_key, scope).await
})
}))
}) {
Ok(disposer) => Ok(disposer),
Err(e) => {
self.0
.shared
.registry
.finalize_binding_if(key.clone(), scope, &stored);
Err(e)
}
}
}
pub fn effect(&self, f: impl FnOnce() -> Effect) -> Result<Disposer, CordisError> {
let record = self.register_effect(f)?;
let handle = self.handle().clone();
Ok(Disposer::new(Box::new(move || {
let record = record.clone();
let handle = handle.clone();
Box::pin(async move { record.drain(&handle).await })
})))
}
pub(crate) fn register_effect(
&self,
f: impl FnOnce() -> Effect,
) -> Result<Arc<EffectRecord>, CordisError> {
let fiber = self.0.fiber.upgrade().ok_or(CordisError::InactiveEffect)?;
let record = EffectRecord::new(f(), self.0.fiber.clone());
let handle = self.handle().clone();
{
let tr = fiber.transition.lock().unwrap();
if matches!(tr.state, FiberState::Unloading | FiberState::Disposed) {
drop(tr);
let sink = self.error_sink();
let drain_handle = handle.clone();
handle.spawn(async move {
if let Err(e) = record.drain(&drain_handle).await {
sink(e);
}
});
return Err(CordisError::InactiveEffect);
}
fiber.effects.lock().unwrap().push(record.clone());
}
Ok(record)
}
pub fn plugin(&self, p: impl Plugin) -> FiberView {
let fiber = spawn_fiber(&self.0.shared, Some(self), Some(Arc::new(p)), false);
self.mount_fiber(fiber)
}
pub fn plugin_with<C: Send + Sync + 'static>(
&self,
factory: impl PluginFactory<C>,
config: C,
) -> FiberView {
let fiber =
crate::fiber::spawn_factory_fiber(&self.0.shared, Some(self), factory, config, false);
self.mount_fiber(fiber)
}
pub fn plugin_from<C: Send + Sync + 'static>(
&self,
build: impl Fn(&C) -> Result<Box<dyn Plugin>, CordisError> + Send + Sync + 'static,
config: C,
) -> FiberView {
struct ClosureFactory<F> {
build: F,
}
impl<C, F> PluginFactory<C> for ClosureFactory<F>
where
C: Send + Sync + 'static,
F: Fn(&C) -> Result<Box<dyn Plugin>, CordisError> + Send + Sync + 'static,
{
fn build(&self, config: &C) -> Result<Box<dyn Plugin>, CordisError> {
(self.build)(config)
}
}
self.plugin_with(ClosureFactory { build }, config)
}
fn mount_fiber(&self, fiber: std::sync::Arc<crate::fiber::FiberInner>) -> FiberView {
let view = FiberView::from_inner(fiber.clone());
let child = view.clone();
let sink = self.error_sink();
let registered = self.register_effect(move || {
Effect::AsyncDisposer(Box::new(move || {
let child = child.clone();
let sink = sink.clone();
Box::pin(async move {
let delivered = matches!(child.state().state, FiberState::Disposed);
if let Err(e) = child.dispose().await {
if !delivered {
sink(e);
}
}
Ok(())
})
}))
});
match registered {
Ok(record) => {
*fiber.mount.lock().unwrap() = Some(record);
fiber.post(Intent::RefreshDeps);
}
Err(_) => {
fiber.post(Intent::Dispose);
}
}
view
}
pub fn cancellation_token(&self) -> CancellationToken {
match self.0.fiber.upgrade() {
Some(fiber) => fiber.current_token(),
None => {
let token = CancellationToken::new();
token.cancel();
token
}
}
}
pub fn cancelled(&self) -> impl Future<Output = ()> + Send + 'static {
let token = self.cancellation_token();
async move { token.cancelled().await }
}
pub fn refresh(&self) {
self.0.shared.registry.refresh_all();
}
}
async fn evict_and_finalize(
shared: &Arc<Shared>,
provider: Weak<FiberInner>,
pid: crate::PluginId,
provider_gen: u64,
key: TypeKey,
scope: Option<crate::key::ScopeId>,
) -> Result<(), CordisError> {
let old = shared
.registry
.lookup(&key, scope.as_ref())
.filter(|b| b.provider_id == pid && b.provider_gen == provider_gen);
if let Some(binding) = &old {
binding
.removing
.store(true, std::sync::atomic::Ordering::SeqCst);
}
let consumers: Vec<Arc<FiberInner>> = shared
.registry
.consumers_of(&key, (pid, provider_gen, key.clone(), scope.clone()));
let mut tasks = Vec::new();
for fiber in &consumers {
fiber.cancel_current();
let task = TransitionTask::new();
fiber.post_join(task.clone(), Intent::RefreshDepsJoin);
tasks.push(task);
}
shared.registry.notify_key_changed(&key);
for task in tasks {
let _ = join_task(&task).await;
}
if let Some(binding) = old {
shared
.registry
.finalize_binding_if(key.clone(), scope.clone(), &binding);
}
if let Some(fiber) = provider.upgrade() {
let mut provided = fiber.provided.lock().unwrap();
if let Some(pos) = provided.iter().position(|(k, s)| *k == key && *s == scope) {
provided.swap_remove(pos);
}
}
Ok(())
}