use std::{collections::HashMap, future::Future, sync::Arc, time::Duration};
use tokio::sync::{Notify, RwLock, Semaphore, mpsc, oneshot};
use tokio::task::{JoinError, JoinHandle};
use tokio_util::sync::CancellationToken;
use crate::core::actor::{ActorExitReason, TaskActor, TaskActorParams};
use crate::core::outcome::TaskOutcome;
use crate::events::{Bus, Event, EventKind};
use crate::identity::TaskId;
use crate::reasons;
use crate::tasks::TaskSpec;
pub(crate) type OutcomeTx = oneshot::Sender<TaskOutcome>;
pub(crate) enum RegistryCommand {
Add(TaskId, TaskSpec, Option<OutcomeTx>),
Remove(TaskId),
}
struct Handle {
join: JoinHandle<ActorExitReason>,
cancel: CancellationToken,
label: Arc<str>,
done: Option<OutcomeTx>,
}
#[derive(Default)]
struct Inner {
tasks: HashMap<TaskId, Handle>,
by_label: HashMap<Arc<str>, TaskId>,
}
#[derive(Default)]
struct PendingInner {
counts: HashMap<TaskId, usize>,
labels: HashMap<TaskId, Arc<str>>,
}
#[derive(Default)]
struct PendingJoins {
inner: std::sync::Mutex<PendingInner>,
drained: Notify,
}
impl PendingJoins {
fn inc(&self, id: TaskId) {
let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
*g.counts.entry(id).or_insert(0) += 1;
}
fn label(&self, id: TaskId, label: Arc<str>) {
let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
if g.counts.contains_key(&id) {
g.labels.insert(id, label);
}
}
fn dec(&self, id: TaskId) {
let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
if let Some(n) = g.counts.get_mut(&id) {
*n -= 1;
if *n == 0 {
g.counts.remove(&id);
g.labels.remove(&id);
}
}
if g.counts.is_empty() {
self.drained.notify_waiters();
}
}
fn contains(&self, id: TaskId) -> bool {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.counts
.contains_key(&id)
}
fn is_empty(&self) -> bool {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.counts
.is_empty()
}
fn pending_labels(&self) -> Vec<Arc<str>> {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.labels
.values()
.cloned()
.collect()
}
async fn wait_drained(&self) {
loop {
let notified = self.drained.notified();
if self.is_empty() {
return;
}
notified.await;
}
}
}
pub(crate) struct Registry {
state: RwLock<Inner>,
bus: Bus,
runtime_token: CancellationToken,
semaphore: Option<Arc<Semaphore>>,
grace: Duration,
empty_notify: Notify,
cmd_rx: std::sync::Mutex<Option<mpsc::UnboundedReceiver<RegistryCommand>>>,
pending_joins: Arc<PendingJoins>,
listener_handle: std::sync::Mutex<Option<JoinHandle<()>>>,
}
impl Registry {
pub fn new(
bus: Bus,
runtime_token: CancellationToken,
semaphore: Option<Arc<Semaphore>>,
grace: Duration,
cmd_rx: mpsc::UnboundedReceiver<RegistryCommand>,
) -> Arc<Self> {
Arc::new(Self {
state: RwLock::new(Inner::default()),
bus,
runtime_token,
semaphore,
grace,
empty_notify: Notify::new(),
cmd_rx: std::sync::Mutex::new(Some(cmd_rx)),
pending_joins: Arc::new(PendingJoins::default()),
listener_handle: std::sync::Mutex::new(None),
})
}
pub async fn is_terminated(&self, id: TaskId) -> bool {
if self.state.read().await.tasks.contains_key(&id) {
return false;
}
!self.pending_joins.contains(id)
}
pub async fn wait_joins_within(&self, grace: Duration) -> Vec<Arc<str>> {
let _ = tokio::time::timeout(grace, self.pending_joins.wait_drained()).await;
self.pending_joins.pending_labels()
}
#[inline]
fn notify_after_remove(&self, len_after: usize) {
if len_after == 0 {
self.empty_notify.notify_one();
}
}
pub async fn wait_until_empty(&self) {
loop {
let notified = self.empty_notify.notified();
if self.is_empty().await {
return;
}
notified.await;
}
}
pub fn spawn_listener(self: Arc<Self>) {
let mut cmd_rx = self
.cmd_rx
.lock()
.unwrap_or_else(|e| e.into_inner())
.take()
.expect("spawn_listener called exactly once");
let mut bus_rx = self.bus.subscribe();
let rt = self.runtime_token.clone();
let me = self.clone();
let handle = tokio::spawn(async move {
loop {
tokio::select! {
biased;
_ = rt.cancelled() => break,
cmd = cmd_rx.recv() => match cmd {
Some(RegistryCommand::Add(id, spec, done)) => {
me.guarded("registry", me.spawn_and_register(id, spec, done))
.await;
}
Some(RegistryCommand::Remove(id)) => {
me.guarded("registry", me.remove_task(id)).await;
}
None => break,
},
msg = bus_rx.recv() => match msg {
Ok(ev) => me.guarded("registry", me.handle_bus_event(&ev)).await,
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
me.guarded("registry", me.reap_finished()).await;
continue;
}
}
}
}
cmd_rx.close();
while let Some(cmd) = cmd_rx.recv().await {
match cmd {
RegistryCommand::Add(id, spec, done) => {
me.guarded("registry", me.spawn_and_register(id, spec, done))
.await;
}
RegistryCommand::Remove(id) => {
me.guarded("registry", me.remove_task(id)).await;
}
}
}
me.cancel_all_within(Duration::ZERO).await;
me.pending_joins.wait_drained().await;
});
*self
.listener_handle
.lock()
.unwrap_or_else(|e| e.into_inner()) = Some(handle);
}
pub async fn join_listener(&self) {
let handle = self
.listener_handle
.lock()
.unwrap_or_else(|e| e.into_inner())
.take();
if let Some(handle) = handle {
let _ = handle.await;
}
}
async fn guarded(&self, who: &'static str, fut: impl Future<Output = ()>) {
if let Err(msg) = crate::core::panic_guard::guarded(fut).await {
self.bus.publish(Event::subscriber_panicked(
who,
format!("listener panic: {msg}"),
));
}
}
async fn handle_bus_event(&self, event: &Event) {
match event.kind {
EventKind::ActorExhausted | EventKind::ActorDead => {
if let Some(id) = event.id {
self.cleanup_task(id).await;
}
}
_ => {}
}
}
pub async fn list(&self) -> Vec<(TaskId, Arc<str>)> {
let st = self.state.read().await;
let mut v: Vec<(TaskId, Arc<str>)> = st
.tasks
.iter()
.map(|(id, h)| (*id, h.label.clone()))
.collect();
v.sort_by_key(|(id, _)| *id);
v
}
pub async fn contains(&self, id: TaskId) -> bool {
self.state.read().await.tasks.contains_key(&id)
}
pub async fn id_for_label(&self, name: &str) -> Option<TaskId> {
self.state.read().await.by_label.get(name).copied()
}
pub async fn is_empty(&self) -> bool {
self.state.read().await.tasks.is_empty()
}
pub async fn cancel_all_within(&self, grace: Duration) -> Vec<Arc<str>> {
let grace = grace.min(Duration::from_secs(60 * 60 * 24 * 365 * 30));
let handles: Vec<(TaskId, Handle)> = {
let mut st = self.state.write().await;
st.by_label.clear();
let drained = st.tasks.drain().collect::<Vec<_>>();
self.empty_notify.notify_waiters();
drained
};
for (id, h) in &handles {
self.pending_joins.inc(*id);
h.cancel.cancel();
}
let deadline = tokio::time::Instant::now() + grace;
let mut stuck = Vec::new();
for (id, h) in handles {
let label = h.label.clone();
let mut join = h.join;
match tokio::time::timeout_at(deadline, &mut join).await {
Ok(res) => {
self.pending_joins.dec(id);
Self::report_join(&self.bus, id, &label, res, h.done);
}
Err(_elapsed) => {
join.abort();
let _ = join.await;
self.pending_joins.dec(id);
if let Some(done) = h.done {
let _ = done.send(TaskOutcome::ForceAborted);
}
self.bus.publish(
Event::new(EventKind::TaskRemoved)
.with_task(Arc::clone(&label))
.with_id(id)
.with_reason("force_terminated_after_grace"),
);
stuck.push(label);
}
}
}
let _ = tokio::time::timeout_at(deadline, self.pending_joins.wait_drained()).await;
stuck
}
async fn spawn_and_register(&self, id: TaskId, spec: TaskSpec, done: Option<OutcomeTx>) {
let label: Arc<str> = Arc::from(spec.task().name());
let mut st = self.state.write().await;
if st.by_label.contains_key(&label) {
drop(st);
if let Some(done) = done {
let _ = done.send(TaskOutcome::Rejected {
reason: Arc::from(reasons::ALREADY_EXISTS),
});
}
self.bus.publish(
Event::new(EventKind::TaskAddFailed)
.with_task(label)
.with_id(id)
.with_reason(reasons::ALREADY_EXISTS),
);
return;
}
let task_token = self.runtime_token.child_token();
let actor = TaskActor::new(
self.bus.clone(),
label.clone(),
spec.task().clone(),
TaskActorParams {
restart: spec.restart(),
backoff: spec.backoff(),
timeout: spec.timeout(),
max_retries: spec.max_retries(),
},
self.semaphore.clone(),
id,
);
let task_token_clone = task_token.clone();
let join_handle = tokio::spawn(async move { actor.run(task_token_clone).await });
st.tasks.insert(
id,
Handle {
join: join_handle,
cancel: task_token,
label: label.clone(),
done,
},
);
st.by_label.insert(label.clone(), id);
drop(st);
self.bus.publish(
Event::new(EventKind::TaskAdded)
.with_task(label)
.with_id(id),
);
}
async fn remove_task(&self, id: TaskId) {
self.pending_joins.inc(id);
if let Some((handle, len_after)) = self.take_handle(id).await {
self.notify_after_remove(len_after);
handle.cancel.cancel();
self.spawn_join_report(id, handle.label, handle.join, Some(self.grace), handle.done);
} else {
self.pending_joins.dec(id);
self.bus.publish(
Event::new(EventKind::TaskRemoved)
.with_id(id)
.with_reason("task_not_found"),
);
}
}
async fn cleanup_task(&self, id: TaskId) {
self.pending_joins.inc(id);
if let Some((handle, len_after)) = self.take_handle(id).await {
self.notify_after_remove(len_after);
self.spawn_join_report(id, handle.label, handle.join, Some(self.grace), handle.done);
} else {
self.pending_joins.dec(id);
}
}
async fn take_handle(&self, id: TaskId) -> Option<(Handle, usize)> {
let mut st = self.state.write().await;
let h = st.tasks.remove(&id)?;
st.by_label.remove(&h.label);
let len_after = st.tasks.len();
Some((h, len_after))
}
fn spawn_join_report(
&self,
id: TaskId,
name: Arc<str>,
join: JoinHandle<ActorExitReason>,
force_after: Option<Duration>,
done: Option<OutcomeTx>,
) {
let bus = self.bus.clone();
let pending = Arc::clone(&self.pending_joins);
pending.label(id, Arc::clone(&name));
tokio::spawn(async move {
let mut join = join;
match force_after {
Some(grace) => match tokio::time::timeout(grace, &mut join).await {
Ok(res) => {
Self::report_join(&bus, id, &name, res, done);
pending.dec(id);
}
Err(_) => {
join.abort();
let _ = join.await;
if let Some(done) = done {
let _ = done.send(TaskOutcome::ForceAborted);
}
bus.publish(
Event::new(EventKind::TaskRemoved)
.with_task(name)
.with_id(id)
.with_reason("force_terminated_after_grace"),
);
pending.dec(id);
}
},
None => {
let res = join.await;
Self::report_join(&bus, id, &name, res, done);
pending.dec(id);
}
}
});
}
async fn reap_finished(&self) {
let finished: Vec<TaskId> = {
let st = self.state.read().await;
st.tasks
.iter()
.filter(|(_, h)| h.join.is_finished())
.map(|(id, _)| *id)
.collect()
};
for id in finished {
self.cleanup_task(id).await;
}
}
fn report_join(
bus: &Bus,
id: TaskId,
name: &str,
res: Result<ActorExitReason, JoinError>,
done: Option<OutcomeTx>,
) {
if let Err(e) = &res
&& e.is_panic()
{
bus.publish(
Event::new(EventKind::ActorDead)
.with_task(name)
.with_id(id)
.with_reason("actor_panic"),
);
}
if let Some(done) = done {
let _ = done.send(Self::outcome_of(res));
}
bus.publish(
Event::new(EventKind::TaskRemoved)
.with_task(name)
.with_id(id),
);
}
fn outcome_of(res: Result<ActorExitReason, JoinError>) -> TaskOutcome {
match res {
Ok(ActorExitReason::Completed) => TaskOutcome::Completed,
Ok(ActorExitReason::Canceled) => TaskOutcome::Canceled,
Ok(ActorExitReason::Exhausted {
reason,
exit_code,
source,
}) => TaskOutcome::Failed {
reason,
exit_code,
source,
},
Ok(ActorExitReason::Fatal {
reason,
exit_code,
source,
}) => TaskOutcome::Fatal {
reason,
exit_code,
source,
},
Err(e) if e.is_panic() => TaskOutcome::Panicked,
Err(_aborted) => TaskOutcome::ForceAborted,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn pending_wait_drained_resolves_after_last_dec() {
let p = Arc::new(PendingJoins::default());
let a = TaskId::next();
let b = TaskId::next();
p.inc(a);
p.inc(b);
assert!(!p.is_empty());
let p2 = Arc::clone(&p);
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
p2.dec(a);
p2.dec(b);
});
tokio::time::timeout(Duration::from_secs(1), p.wait_drained())
.await
.expect("wait_drained must resolve once every join is decremented");
assert!(p.is_empty(), "no joins should remain after draining");
}
#[tokio::test]
async fn pending_wait_drained_returns_immediately_when_empty() {
let p = PendingJoins::default();
tokio::time::timeout(Duration::from_millis(100), p.wait_drained())
.await
.expect("an empty PendingJoins must resolve immediately");
}
fn registry() -> Arc<Registry> {
let bus = Bus::new(64);
let token = CancellationToken::new();
let (_tx, rx) = mpsc::unbounded_channel();
Registry::new(bus, token, None, Duration::from_secs(5), rx)
}
#[tokio::test]
async fn wait_joins_within_reports_stuck_labels_then_drains() {
let reg = registry();
assert!(
reg.wait_joins_within(Duration::from_millis(50))
.await
.is_empty(),
"an empty join set must drain immediately"
);
let id = TaskId::next();
reg.pending_joins.inc(id);
reg.pending_joins.label(id, Arc::from("stuck-task"));
let stuck = reg.wait_joins_within(Duration::from_millis(30)).await;
assert_eq!(
stuck,
vec![Arc::<str>::from("stuck-task")],
"an in-flight join must be reported with its label on timeout"
);
let p = Arc::clone(®.pending_joins);
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
p.dec(id);
});
assert!(
reg.wait_joins_within(Duration::from_secs(1))
.await
.is_empty(),
"must drain once the in-flight join is decremented"
);
}
#[tokio::test]
async fn reap_finished_removes_completed_handles() {
let reg = registry();
let join = tokio::spawn(async { ActorExitReason::Completed });
while !join.is_finished() {
tokio::task::yield_now().await;
}
reg.state.write().await.tasks.insert(
TaskId::next(),
Handle {
join,
cancel: CancellationToken::new(),
label: Arc::from("done"),
done: None,
},
);
assert!(!reg.is_empty().await);
reg.reap_finished().await;
assert!(
reg.is_empty().await,
"reap_finished must drop the completed handle"
);
}
#[tokio::test]
async fn reap_finished_keeps_running_handles() {
let reg = registry();
let cancel = CancellationToken::new();
let child = cancel.clone();
let join = tokio::spawn(async move {
child.cancelled().await;
ActorExitReason::Canceled
});
reg.state.write().await.tasks.insert(
TaskId::next(),
Handle {
join,
cancel,
label: Arc::from("running"),
done: None,
},
);
reg.reap_finished().await;
assert!(!reg.is_empty().await, "a running actor must not be reaped");
}
#[tokio::test(flavor = "current_thread")]
async fn lag_recovery_keeps_retained_terminal_event() {
let bus = Bus::new(1);
let runtime_token = CancellationToken::new();
let (_cmd_tx, cmd_rx) = mpsc::unbounded_channel();
let reg = Registry::new(
bus.clone(),
runtime_token.clone(),
None,
Duration::from_millis(50),
cmd_rx,
);
let id = TaskId::next();
let label: Arc<str> = Arc::from("running-during-lag");
let actor_token = CancellationToken::new();
let actor_wait = actor_token.clone();
let join = tokio::spawn(async move {
actor_wait.cancelled().await;
ActorExitReason::Canceled
});
{
let mut state = reg.state.write().await;
state.by_label.insert(Arc::clone(&label), id);
state.tasks.insert(
id,
Handle {
join,
cancel: actor_token.clone(),
label: Arc::clone(&label),
done: None,
},
);
}
reg.clone().spawn_listener();
bus.publish(Event::new(EventKind::TaskStarting).with_task("lag-seed"));
bus.publish(
Event::new(EventKind::ActorExhausted)
.with_task(label)
.with_id(id),
);
let recovered =
tokio::time::timeout(Duration::from_millis(250), reg.wait_until_empty()).await;
actor_token.cancel();
runtime_token.cancel();
tokio::time::timeout(Duration::from_secs(1), reg.join_listener())
.await
.expect("registry listener must stop after cancellation");
assert!(
recovered.is_ok(),
"lag recovery must preserve and process the retained terminal event"
);
}
#[tokio::test]
async fn shutdown_drains_buffered_command_and_never_silently_drops() {
use crate::{TaskContext, TaskError, TaskFn, TaskRef};
let bus = Bus::new(64);
let token = CancellationToken::new();
let (tx, rx) = mpsc::unbounded_channel();
let reg = Registry::new(bus, token.clone(), None, Duration::from_millis(50), rx);
reg.clone().spawn_listener();
let task: TaskRef = TaskFn::arc("buffered", |ctx: TaskContext| async move {
ctx.cancelled().await;
Err(TaskError::Canceled)
});
let (done_tx, done_rx) = oneshot::channel();
let id = TaskId::next();
tx.send(RegistryCommand::Add(
id,
TaskSpec::restartable(task),
Some(done_tx),
))
.expect("channel is open before shutdown");
token.cancel();
tokio::time::timeout(Duration::from_secs(2), reg.join_listener())
.await
.expect("join_listener must not hang");
let outcome = tokio::time::timeout(Duration::from_secs(1), done_rx)
.await
.expect("watcher must resolve")
.expect("watcher sender must not be dropped — the buffered Add must be acted on");
assert!(
matches!(outcome, TaskOutcome::Canceled | TaskOutcome::ForceAborted),
"a buffered task drained at shutdown must terminate, got {outcome:?}"
);
assert!(
reg.pending_joins.is_empty(),
"wait_drained must leave no in-flight joins after shutdown"
);
assert!(
tx.send(RegistryCommand::Remove(TaskId::next())).is_err(),
"after shutdown the command channel is closed; sends must return Err"
);
}
}