monoloop_loop/transaction/lifecycle/
task_supervisor.rs1use 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#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
16pub struct TaskId(pub u64);
17
18#[derive(Clone, Debug, PartialEq, Eq, Hash)]
20pub enum TaskClass {
21 TransactionCoordinator(TransactionId),
23 EventPublisher(TransactionId),
25 ConnectorOwner(TransactionId, ExchangeId),
27 InterpreterOwner(TransactionId, ExchangeId),
29 ToolWorker(TransactionId, ToolExecutionId),
31 LoopRuntime(TransactionId),
33 McpRequest(TransactionId),
35 Finalizer(TransactionId),
37 RuntimeService,
39}
40
41impl TaskClass {
42 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
60pub enum TaskExit {
61 Completed,
63 Cancelled,
65 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#[derive(Debug)]
79pub struct TaskSupervisor {
80 joins: JoinSet<(TaskId, TaskExit)>,
81 meta: HashMap<TaskId, TaskMeta>,
82 by_transaction: HashMap<TransactionId, HashSet<TaskId>>,
83 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 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 pub fn registered_count(&self) -> usize {
108 self.meta.len()
109 }
110
111 pub fn is_empty(&self) -> bool {
113 self.meta.is_empty() && self.joins.is_empty()
114 }
115
116 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 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 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 let _ = gate_tx.send(());
161 id
162 }
163
164 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 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 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 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 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 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 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 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 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}