Skip to main content

monoloop_loop/transaction/lifecycle/
task_supervisor.rs

1//! Global task join ownership (v2 §7.3).
2//!
3//! Every runtime task is registered before its start gate is released. Abort
4//! retains the join until the result is observed. Live `JoinHandle` values are
5//! never dropped.
6
7use futures_util::FutureExt;
8use monoloop_contracts::{ExchangeId, ToolExecutionId, TransactionId};
9use std::collections::{HashMap, HashSet};
10use std::future::Future;
11use std::panic::AssertUnwindSafe;
12use tokio::task::{AbortHandle, Id as TokioTaskId, JoinSet};
13
14/// Stable task id within one runtime owner.
15#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
16pub struct TaskId(pub u64);
17
18/// Classification for every supervised spawn.
19#[derive(Clone, Debug, PartialEq, Eq, Hash)]
20pub enum TaskClass {
21    /// Per-transaction coordinator.
22    TransactionCoordinator(TransactionId),
23    /// Event publisher task for one transaction.
24    EventPublisher(TransactionId),
25    /// Connector ownership task.
26    ConnectorOwner(TransactionId, ExchangeId),
27    /// Interpreter pump task.
28    InterpreterOwner(TransactionId, ExchangeId),
29    /// Tool worker.
30    ToolWorker(TransactionId, ToolExecutionId),
31    /// Inner LoopRuntime owner for one transaction (M5).
32    LoopRuntime(TransactionId),
33    /// MCP request task.
34    McpRequest(TransactionId),
35    /// Post-terminal finalizer (Seal + completion publish) for one transaction.
36    Finalizer(TransactionId),
37    /// Runtime-wide service (MCP listener, etc.).
38    RuntimeService,
39}
40
41impl TaskClass {
42    /// Owning transaction when task-scoped.
43    pub fn transaction_id(&self) -> Option<TransactionId> {
44        match self {
45            Self::TransactionCoordinator(t)
46            | Self::EventPublisher(t)
47            | Self::ConnectorOwner(t, _)
48            | Self::InterpreterOwner(t, _)
49            | Self::ToolWorker(t, _)
50            | Self::LoopRuntime(t)
51            | Self::McpRequest(t)
52            | Self::Finalizer(t) => Some(*t),
53            Self::RuntimeService => None,
54        }
55    }
56}
57
58/// Observed task exit kind.
59#[derive(Clone, Copy, Debug, PartialEq, Eq)]
60pub enum TaskExit {
61    /// Future completed normally.
62    Completed,
63    /// Task was aborted / cancelled.
64    Cancelled,
65    /// Join observed a panic.
66    Panicked,
67}
68
69#[derive(Debug)]
70struct TaskMeta {
71    class: TaskClass,
72    abort: AbortHandle,
73    abort_requested: bool,
74    tokio_id: TokioTaskId,
75}
76
77/// Retains every runtime join until observed complete.
78#[derive(Debug)]
79pub struct TaskSupervisor {
80    joins: JoinSet<(TaskId, TaskExit)>,
81    meta: HashMap<TaskId, TaskMeta>,
82    by_transaction: HashMap<TransactionId, HashSet<TaskId>>,
83    /// Maps Tokio task id → supervisor TaskId so cancelled joins never leak meta.
84    by_tokio_id: HashMap<TokioTaskId, TaskId>,
85    next_id: u64,
86}
87
88impl Default for TaskSupervisor {
89    fn default() -> Self {
90        Self::new()
91    }
92}
93
94impl TaskSupervisor {
95    /// Empty supervisor.
96    pub fn new() -> Self {
97        Self {
98            joins: JoinSet::new(),
99            meta: HashMap::new(),
100            by_transaction: HashMap::new(),
101            by_tokio_id: HashMap::new(),
102            next_id: 1,
103        }
104    }
105
106    /// Number of tasks still registered (including abort-requested).
107    pub fn registered_count(&self) -> usize {
108        self.meta.len()
109    }
110
111    /// Whether stopped proof can assert zero owned tasks.
112    pub fn is_empty(&self) -> bool {
113        self.meta.is_empty() && self.joins.is_empty()
114    }
115
116    /// Tasks still associated with a transaction.
117    pub fn tasks_for(&self, tx: &TransactionId) -> Vec<TaskId> {
118        self.by_transaction
119            .get(tx)
120            .map(|s| s.iter().copied().collect())
121            .unwrap_or_default()
122    }
123
124    /// Register then spawn. The future does not run until after registration
125    /// completes (start-gate released at the end of this method).
126    pub fn spawn<F>(&mut self, class: TaskClass, future: F) -> TaskId
127    where
128        F: Future<Output = ()> + Send + 'static,
129    {
130        let id = TaskId(self.next_id);
131        self.next_id = self.next_id.saturating_add(1);
132
133        let (gate_tx, gate_rx) = tokio::sync::oneshot::channel::<()>();
134        let abort = self.joins.spawn(async move {
135            // Fail closed if the gate is dropped before release (abort-before-start).
136            match gate_rx.await {
137                Ok(()) => match AssertUnwindSafe(future).catch_unwind().await {
138                    Ok(()) => (id, TaskExit::Completed),
139                    Err(_) => (id, TaskExit::Panicked),
140                },
141                Err(_) => (id, TaskExit::Cancelled),
142            }
143        });
144        let tokio_id = abort.id();
145
146        if let Some(tx) = class.transaction_id() {
147            self.by_transaction.entry(tx).or_default().insert(id);
148        }
149        self.by_tokio_id.insert(tokio_id, id);
150        self.meta.insert(
151            id,
152            TaskMeta {
153                class,
154                abort,
155                abort_requested: false,
156                tokio_id,
157            },
158        );
159        // Release start gate only after registration is complete.
160        let _ = gate_tx.send(());
161        id
162    }
163
164    /// Request abort; task remains registered until join is observed.
165    pub fn abort(&mut self, id: TaskId) {
166        if let Some(meta) = self.meta.get_mut(&id) {
167            meta.abort_requested = true;
168            meta.abort.abort();
169        }
170    }
171
172    /// Abort all tasks for a transaction.
173    pub fn abort_transaction(&mut self, tx: &TransactionId) {
174        let ids = self.tasks_for(tx);
175        for id in ids {
176            self.abort(id);
177        }
178    }
179
180    /// Abort residual tx work but keep [`TaskClass::Finalizer`] and
181    /// [`TaskClass::EventPublisher`] alive so Seal + completion publication can
182    /// finish (one completion per admission).
183    pub fn abort_transaction_residuals(&mut self, tx: &TransactionId) {
184        let ids = self.tasks_for(tx);
185        for id in ids {
186            let keep = self.meta.get(&id).is_some_and(|m| {
187                matches!(
188                    m.class,
189                    TaskClass::Finalizer(_) | TaskClass::EventPublisher(_)
190                )
191            });
192            if !keep {
193                self.abort(id);
194            }
195        }
196    }
197
198    /// Hard-grace abort: keep only [`TaskClass::Finalizer`] so Seal→completion
199    /// can finish; abort a stuck EventPublisher that already Sealed or timed out.
200    pub fn abort_transaction_except_finalizer(&mut self, tx: &TransactionId) {
201        let ids = self.tasks_for(tx);
202        for id in ids {
203            let keep = self
204                .meta
205                .get(&id)
206                .is_some_and(|m| matches!(m.class, TaskClass::Finalizer(_)));
207            if !keep {
208                self.abort(id);
209            }
210        }
211    }
212
213    /// Abort every registered task (shutdown).
214    pub fn abort_all(&mut self) {
215        let ids: Vec<_> = self.meta.keys().copied().collect();
216        for id in ids {
217            self.abort(id);
218        }
219    }
220
221    /// Abort all tasks and observe joins until the set is empty.
222    ///
223    /// Returns `false` if joins do not complete within the bound — caller MUST
224    /// remain `Quiescing` (never report `Stopped` while owned work remains).
225    /// Live `JoinHandle`s are retained (§7.3 / §21: no drop on deadline).
226    pub async fn abort_and_drain(&mut self) -> bool {
227        self.abort_all();
228        let drain = async { while self.join_next().await.is_some() {} };
229        if tokio::time::timeout(std::time::Duration::from_secs(2), drain)
230            .await
231            .is_err()
232        {
233            return false;
234        }
235        self.meta.clear();
236        self.by_transaction.clear();
237        self.by_tokio_id.clear();
238        true
239    }
240
241    /// Poll for the next finished task and deregister it.
242    pub async fn join_next(&mut self) -> Option<(TaskId, TaskClass, TaskExit)> {
243        let finished = self.joins.join_next_with_id().await?;
244        self.reap_finished(finished)
245    }
246
247    /// Non-blocking reap of already-finished joins.
248    pub fn try_reap_finished(&mut self) -> Vec<(TaskId, TaskClass, TaskExit)> {
249        let mut out = Vec::new();
250        while let Some(finished) = self.joins.try_join_next_with_id() {
251            if let Some(item) = self.reap_finished(finished) {
252                out.push(item);
253            }
254        }
255        out
256    }
257
258    fn reap_finished(
259        &mut self,
260        finished: Result<(TokioTaskId, (TaskId, TaskExit)), tokio::task::JoinError>,
261    ) -> Option<(TaskId, TaskClass, TaskExit)> {
262        match finished {
263            Ok((tokio_id, (id, exit))) => {
264                self.by_tokio_id.remove(&tokio_id);
265                let exit = if self.meta.get(&id).is_some_and(|m| m.abort_requested)
266                    && matches!(exit, TaskExit::Completed)
267                {
268                    TaskExit::Cancelled
269                } else {
270                    exit
271                };
272                let class = self.deregister(id)?;
273                Some((id, class, exit))
274            }
275            Err(err) => {
276                // Abort/panic dropped the future before it returned our tuple.
277                let exit = if err.is_cancelled() {
278                    TaskExit::Cancelled
279                } else {
280                    TaskExit::Panicked
281                };
282                let tokio_id = err.id();
283                let id = self
284                    .by_tokio_id
285                    .remove(&tokio_id)
286                    .or_else(|| self.find_finished_meta())?;
287                let class = self.deregister(id)?;
288                Some((id, class, exit))
289            }
290        }
291    }
292
293    fn find_finished_meta(&self) -> Option<TaskId> {
294        self.meta
295            .iter()
296            .find(|(_, m)| m.abort.is_finished())
297            .map(|(id, _)| *id)
298    }
299
300    fn deregister(&mut self, id: TaskId) -> Option<TaskClass> {
301        let meta = self.meta.remove(&id)?;
302        self.by_tokio_id.remove(&meta.tokio_id);
303        if let Some(tx) = meta.class.transaction_id() {
304            if let Some(set) = self.by_transaction.get_mut(&tx) {
305                set.remove(&id);
306                if set.is_empty() {
307                    self.by_transaction.remove(&tx);
308                }
309            }
310        }
311        Some(meta.class)
312    }
313}