use std::{
any::{Any, TypeId},
collections::HashMap,
future::Future,
panic::AssertUnwindSafe,
pin::Pin,
sync::{
Arc, Mutex,
atomic::{AtomicBool, AtomicU64, Ordering},
},
};
use futures::{FutureExt, future::join_all};
use tokio::sync::watch;
use crate::{
Error, Query, Result, ServiceKey,
service::{Dependency, ServiceEntry, ServiceId, boxed_service},
};
pub(crate) type HandlerFuture = Pin<Box<dyn Future<Output = Result<()>> + Send + 'static>>;
pub(crate) type EventCallback = dyn Fn(Arc<dyn Any + Send + Sync>) -> HandlerFuture + Send + Sync;
pub(crate) type QueryFuture =
Pin<Box<dyn Future<Output = Result<Option<Box<dyn Any + Send + Sync>>>> + Send + 'static>>;
pub(crate) type QueryCallback = dyn Fn(Arc<dyn Any + Send + Sync>) -> QueryFuture + Send + Sync;
async fn catch_handler(future: HandlerFuture) -> Result<()> {
match AssertUnwindSafe(future).catch_unwind().await {
Ok(result) => result,
Err(payload) => Err(Error::panic(payload)),
}
}
pub(crate) struct Listener {
pub id: u64,
pub owner: u64,
pub active: AtomicBool,
pub callback: Arc<EventCallback>,
}
pub(crate) struct QueryListener {
pub id: u64,
pub owner: u64,
pub active: AtomicBool,
pub callback: Arc<QueryCallback>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct DependencySnapshot {
pub revision: u64,
pub generations: Vec<(TypeId, Option<u64>)>,
pub missing: Arc<[&'static str]>,
}
pub(crate) struct Runtime {
pub services: Mutex<HashMap<ServiceId, ServiceEntry>>,
listeners: Mutex<HashMap<TypeId, Vec<Arc<Listener>>>>,
query_listeners: Mutex<HashMap<TypeId, Vec<Arc<QueryListener>>>>,
next_id: AtomicU64,
next_generation: AtomicU64,
service_revision: AtomicU64,
service_changed: watch::Sender<u64>,
}
impl Runtime {
pub fn new() -> Arc<Self> {
let (service_changed, _) = watch::channel(0);
Arc::new(Self {
services: Mutex::new(HashMap::new()),
listeners: Mutex::new(HashMap::new()),
query_listeners: Mutex::new(HashMap::new()),
next_id: AtomicU64::new(1),
next_generation: AtomicU64::new(1),
service_revision: AtomicU64::new(0),
service_changed,
})
}
pub fn next_id(&self) -> u64 {
self.next_id.fetch_add(1, Ordering::Relaxed)
}
pub fn next_service_generation(&self) -> u64 {
self.next_generation.fetch_add(1, Ordering::Relaxed)
}
pub fn subscribe_services(&self) -> watch::Receiver<u64> {
self.service_changed.subscribe()
}
fn notify_service_change(&self) -> u64 {
let revision = self.service_revision.fetch_add(1, Ordering::AcqRel) + 1;
self.service_changed.send_replace(revision);
revision
}
pub fn dependency_snapshot(&self, dependencies: &[Dependency]) -> DependencySnapshot {
let services = self.services.lock().expect("service lock poisoned");
let mut generations = Vec::with_capacity(dependencies.len());
let mut missing = Vec::new();
for dependency in dependencies {
let generation = services
.get(&ServiceId::root(dependency.key))
.filter(|entry| entry.active)
.map(|entry| entry.generation);
if dependency.required && generation.is_none() {
missing.push(dependency.name);
}
generations.push((dependency.key, generation));
}
DependencySnapshot {
revision: self.service_revision.load(Ordering::Acquire),
generations,
missing: missing.into(),
}
}
pub fn remove_service(&self, id: ServiceId, owner: u64, token: u64) -> bool {
let removed_active = {
let mut services = self.services.lock().expect("service lock poisoned");
let matches = services
.get(&id)
.is_some_and(|entry| entry.owner == owner && entry.token == token);
if !matches {
return false;
}
services.remove(&id).is_some_and(|entry| entry.active)
};
if removed_active {
self.notify_service_change();
}
true
}
pub fn replace_service<K: ServiceKey>(
&self,
id: ServiceId,
owner: u64,
token: u64,
value: Arc<K::Value>,
) -> Result<u64> {
let generation = self.next_service_generation();
let active = {
let mut services = self.services.lock().expect("service lock poisoned");
let entry = services
.get_mut(&id)
.ok_or(Error::MissingService { name: K::NAME })?;
if entry.owner != owner || entry.token != token {
return Err(Error::ServiceOwnership { name: K::NAME });
}
entry.value = boxed_service::<K>(value);
entry.generation = generation;
entry.active
};
if active {
self.notify_service_change();
}
Ok(generation)
}
pub fn touch_service<K: ServiceKey>(
&self,
id: ServiceId,
owner: u64,
token: u64,
) -> Result<u64> {
let generation = self.next_service_generation();
let active = {
let mut services = self.services.lock().expect("service lock poisoned");
let entry = services
.get_mut(&id)
.ok_or(Error::MissingService { name: K::NAME })?;
if entry.owner != owner || entry.token != token {
return Err(Error::ServiceOwnership { name: entry.name });
}
entry.generation = generation;
entry.active
};
if active {
self.notify_service_change();
}
Ok(generation)
}
pub fn commit_owner(&self, owner: u64) {
let service_changed = {
let mut services = self.services.lock().expect("service lock poisoned");
let mut changed = false;
for entry in services.values_mut().filter(|entry| entry.owner == owner) {
if !entry.active {
entry.active = true;
changed = true;
}
}
changed
};
for listeners in self
.listeners
.lock()
.expect("listener lock poisoned")
.values()
{
for listener in listeners.iter().filter(|listener| listener.owner == owner) {
listener.active.store(true, Ordering::Release);
}
}
for listeners in self
.query_listeners
.lock()
.expect("query lock poisoned")
.values()
{
for listener in listeners.iter().filter(|listener| listener.owner == owner) {
listener.active.store(true, Ordering::Release);
}
}
if service_changed {
self.notify_service_change();
}
}
pub fn add_listener(&self, event: TypeId, owner: u64, callback: Arc<EventCallback>) -> u64 {
let id = self.next_id();
let listener = Arc::new(Listener {
id,
owner,
active: AtomicBool::new(false),
callback,
});
self.listeners
.lock()
.expect("listener lock poisoned")
.entry(event)
.or_default()
.push(listener);
id
}
pub fn remove_listener(&self, event: TypeId, id: u64) -> bool {
let mut all = self.listeners.lock().expect("listener lock poisoned");
let Some(listeners) = all.get_mut(&event) else {
return false;
};
let mut removed = false;
listeners.retain(|listener| {
if listener.id == id {
listener.active.store(false, Ordering::Release);
removed = true;
false
} else {
true
}
});
if listeners.is_empty() {
all.remove(&event);
}
removed
}
fn event_snapshot(&self, event: TypeId) -> Vec<Arc<Listener>> {
self.listeners
.lock()
.expect("listener lock poisoned")
.get(&event)
.cloned()
.unwrap_or_default()
}
pub async fn emit_serial<E: Send + Sync + 'static>(&self, event: E) -> Result<()> {
let listeners = self.event_snapshot(TypeId::of::<E>());
let event: Arc<dyn Any + Send + Sync> = Arc::new(event);
for listener in listeners {
if listener.active.load(Ordering::Acquire) {
match AssertUnwindSafe((listener.callback)(event.clone()))
.catch_unwind()
.await
{
Ok(result) => result?,
Err(payload) => return Err(Error::panic(payload)),
}
}
}
Ok(())
}
pub async fn emit_parallel<E: Send + Sync + 'static>(&self, event: E) -> Result<()> {
let listeners = self.event_snapshot(TypeId::of::<E>());
let event: Arc<dyn Any + Send + Sync> = Arc::new(event);
let futures = listeners
.into_iter()
.filter(|listener| listener.active.load(Ordering::Acquire))
.map(|listener| catch_handler((listener.callback)(event.clone())));
for result in join_all(futures).await {
result?;
}
Ok(())
}
pub fn add_query_listener(
&self,
query: TypeId,
owner: u64,
callback: Arc<QueryCallback>,
) -> u64 {
let id = self.next_id();
let listener = Arc::new(QueryListener {
id,
owner,
active: AtomicBool::new(false),
callback,
});
self.query_listeners
.lock()
.expect("query lock poisoned")
.entry(query)
.or_default()
.push(listener);
id
}
pub fn remove_query_listener(&self, query: TypeId, id: u64) -> bool {
let mut all = self.query_listeners.lock().expect("query lock poisoned");
let Some(listeners) = all.get_mut(&query) else {
return false;
};
let mut removed = false;
listeners.retain(|listener| {
if listener.id == id {
listener.active.store(false, Ordering::Release);
removed = true;
false
} else {
true
}
});
if listeners.is_empty() {
all.remove(&query);
}
removed
}
pub async fn query<Q: Query>(&self, query: Q) -> Result<Option<Q::Response>> {
let listeners = self
.query_listeners
.lock()
.expect("query lock poisoned")
.get(&TypeId::of::<Q>())
.cloned()
.unwrap_or_default();
let query: Arc<dyn Any + Send + Sync> = Arc::new(query);
for listener in listeners {
if !listener.active.load(Ordering::Acquire) {
continue;
}
let response = match AssertUnwindSafe((listener.callback)(query.clone()))
.catch_unwind()
.await
{
Ok(result) => result?,
Err(payload) => return Err(Error::panic(payload)),
};
if let Some(response) = response {
return response
.downcast::<Q::Response>()
.map(|value| Some(*value))
.map_err(|_| Error::Cleanup("query response type mismatch".into()));
}
}
Ok(None)
}
}