use std::any::Any;
use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex, Weak};
use crate::error::CordisError;
use crate::fiber::{FiberInner, FiberState, Intent, PluginId};
use crate::key::{ScopeId, TypeKey};
pub(crate) type CheckFn = Arc<dyn Fn() -> bool + Send + Sync>;
#[derive(Clone)]
pub(crate) struct StoredValue(Arc<Box<dyn Any + Send + Sync>>);
impl StoredValue {
pub(crate) fn new<T: ?Sized + Send + Sync + 'static>(value: Arc<T>) -> Self {
Self(Arc::new(Box::new(value) as Box<dyn Any + Send + Sync>))
}
pub(crate) fn downcast<T: ?Sized + Send + Sync + 'static>(&self) -> Option<Arc<T>> {
(**self.0).downcast_ref::<Arc<T>>().cloned()
}
}
pub(crate) struct Binding {
pub value: StoredValue,
pub provider: Weak<FiberInner>,
pub provider_id: PluginId,
pub provider_gen: u64,
pub check: Option<CheckFn>,
pub removing: std::sync::atomic::AtomicBool,
}
pub(crate) struct Registry {
bindings: Mutex<HashMap<(TypeKey, Option<ScopeId>), Arc<Binding>>>,
inject_index: Mutex<HashMap<TypeKey, Vec<Weak<FiberInner>>>>,
}
impl Registry {
pub(crate) fn new() -> Self {
Self {
bindings: Mutex::new(HashMap::new()),
inject_index: Mutex::new(HashMap::new()),
}
}
pub(crate) fn insert_binding(
&self,
key: TypeKey,
scope: Option<ScopeId>,
binding: Binding,
) -> Result<Arc<Binding>, CordisError> {
let mut bindings = self.bindings.lock().unwrap();
let entry = (key.clone(), scope.clone());
if let Some(existing) = bindings.get(&entry) {
if !existing.removing.load(std::sync::atomic::Ordering::SeqCst) {
let scope_desc = entry.1.as_deref().unwrap_or("<default>");
return Err(CordisError::ServiceExists(format!(
"{} in scope {scope_desc}",
key.describe()
)));
}
}
let stored = Arc::new(binding);
bindings.insert(entry, stored.clone());
Ok(stored)
}
pub(crate) fn lookup(&self, key: &TypeKey, scope: Option<&ScopeId>) -> Option<Arc<Binding>> {
let bindings = self.bindings.lock().unwrap();
bindings.get(&(key.clone(), scope.cloned())).cloned()
}
pub(crate) fn finalize_binding_if(
&self,
key: TypeKey,
scope: Option<ScopeId>,
expected: &Arc<Binding>,
) {
let mut bindings = self.bindings.lock().unwrap();
let still_old = bindings
.get(&(key.clone(), scope.clone()))
.is_some_and(|b| Arc::ptr_eq(b, expected));
if still_old {
bindings.remove(&(key, scope));
}
if bindings.is_empty() {
bindings.shrink_to_fit();
}
}
pub(crate) fn consumers_of(
&self,
key: &TypeKey,
quad: (PluginId, u64, TypeKey, Option<ScopeId>),
) -> Vec<Arc<FiberInner>> {
let index = self.inject_index.lock().unwrap();
let Some(list) = index.get(key) else {
return Vec::new();
};
list.iter()
.filter_map(|weak| {
let fiber = weak.upgrade()?;
let bound = fiber
.last_deps
.lock()
.unwrap()
.as_ref()
.is_some_and(|deps| deps.contains(&quad));
bound.then_some(fiber)
})
.collect()
}
pub(crate) fn register_inject(&self, key: TypeKey, fiber: &Arc<FiberInner>) {
self.inject_index
.lock()
.unwrap()
.entry(key)
.or_default()
.push(Arc::downgrade(fiber));
}
pub(crate) fn unregister_injects(&self, fiber: &Arc<FiberInner>, keys: &[TypeKey]) {
let weak = Arc::downgrade(fiber);
let mut index = self.inject_index.lock().unwrap();
for key in keys {
let mut empty = false;
if let Some(list) = index.get_mut(key) {
list.retain(|w| !Weak::ptr_eq(w, &weak));
empty = list.is_empty();
}
if empty {
index.remove(key);
}
}
if index.is_empty() {
index.shrink_to_fit();
}
}
pub(crate) fn notify_key_changed(&self, key: &TypeKey) {
let fibers: Vec<Arc<FiberInner>> = {
let index = self.inject_index.lock().unwrap();
index
.get(key)
.map(|list| list.iter().filter_map(|w| w.upgrade()).collect())
.unwrap_or_default()
};
for fiber in fibers {
fiber.post(Intent::RefreshDeps);
}
}
pub(crate) fn refresh_all(&self) {
let mut seen: HashSet<PluginId> = HashSet::new();
let index = self.inject_index.lock().unwrap();
for list in index.values() {
for weak in list {
if let Some(fiber) = weak.upgrade() {
if seen.insert(fiber.id) {
fiber.post(Intent::RefreshDeps);
}
}
}
}
}
pub(crate) fn resolve_dep(
&self,
key: &TypeKey,
scope: Option<&ScopeId>,
) -> Option<(PluginId, u64)> {
let binding = self.lookup(key, scope)?;
if binding.removing.load(std::sync::atomic::Ordering::SeqCst) {
return None;
}
let provider = binding.provider.upgrade()?;
if provider.state() != FiberState::Active {
return None;
}
if let Some(check) = &binding.check {
let passed =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| check())).unwrap_or(false);
if !passed {
return None;
}
}
Some((binding.provider_id, binding.provider_gen))
}
}
#[cfg(test)]
mod transient_tests;