1use std::collections::HashMap;
2use std::sync::Arc;
3use std::time::Instant;
4
5use serde::{Deserialize, Serialize};
6use tokio::sync::broadcast;
7use tokio_util::sync::CancellationToken;
8use uuid::Uuid;
9
10#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq, Hash)]
11#[serde(transparent)]
12pub struct TaskId(pub Uuid);
13
14impl TaskId {
15 pub fn now() -> Self {
16 Self(Uuid::now_v7())
17 }
18}
19
20impl std::fmt::Display for TaskId {
21 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22 self.0.fmt(f)
23 }
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
27#[serde(rename_all = "snake_case")]
28pub enum TaskKind {
29 Bash,
30 Terminal,
31 Flow,
32 Subflow,
33 Agent,
34 Dispatch,
35}
36
37impl TaskKind {
38 pub fn label(self) -> &'static str {
39 match self {
40 TaskKind::Bash => "Bash",
41 TaskKind::Terminal => "Terminal",
42 TaskKind::Flow => "Flow",
43 TaskKind::Subflow => "Subflow",
44 TaskKind::Agent => "Agent",
45 TaskKind::Dispatch => "Dispatch",
46 }
47 }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
51#[serde(rename_all = "snake_case")]
52pub enum TaskStatus {
53 Running,
54 Ok,
55 Err,
56 Killed,
57}
58
59impl TaskStatus {
60 pub fn is_terminal(self) -> bool {
61 matches!(self, TaskStatus::Ok | TaskStatus::Err | TaskStatus::Killed)
62 }
63
64 pub fn is_running(self) -> bool {
65 matches!(self, TaskStatus::Running)
66 }
67}
68
69#[derive(Debug, Clone)]
70pub struct TaskSnapshot {
71 pub id: TaskId,
72 pub kind: TaskKind,
73 pub label: String,
74 pub status: TaskStatus,
75 pub started_at: Instant,
76 pub ended_at: Option<Instant>,
77 pub source_handle: String,
78 pub session_id: String,
79}
80
81impl TaskSnapshot {
82 pub fn elapsed_ms(&self) -> u64 {
83 self.ended_at
84 .unwrap_or_else(Instant::now)
85 .duration_since(self.started_at)
86 .as_millis() as u64
87 }
88
89 pub fn is_running(&self) -> bool {
90 self.status.is_running()
91 }
92}
93
94#[derive(Debug, Clone)]
95pub enum TaskEvent {
96 Registered(TaskSnapshot),
97 StatusChanged {
98 id: TaskId,
99 kind: TaskKind,
100 old: TaskStatus,
101 new: TaskStatus,
102 },
103 Reaped {
104 id: TaskId,
105 },
106}
107
108#[derive(Debug, Clone, Default)]
109pub struct TaskFilter {
110 pub kind: Option<TaskKind>,
111 pub status: Option<TaskStatus>,
112 pub session_id: Option<String>,
113}
114
115impl TaskFilter {
116 pub fn all() -> Self {
117 Self::default()
118 }
119
120 pub fn running() -> Self {
121 Self {
122 status: Some(TaskStatus::Running),
123 ..Default::default()
124 }
125 }
126
127 pub fn matches(&self, snap: &TaskSnapshot) -> bool {
128 if let Some(k) = self.kind
129 && snap.kind != k
130 {
131 return false;
132 }
133 if let Some(s) = self.status
134 && snap.status != s
135 {
136 return false;
137 }
138 if let Some(ref sid) = self.session_id
139 && snap.session_id != *sid
140 {
141 return false;
142 }
143 true
144 }
145}
146
147struct TaskEntry {
148 snapshot: TaskSnapshot,
149 cancel: CancellationToken,
150 kill_hook: Option<std::sync::Arc<dyn Fn() + Send + Sync>>,
151}
152
153#[derive(Clone)]
160pub struct TaskRegistry {
161 inner: Arc<std::sync::Mutex<HashMap<TaskId, TaskEntry>>>,
162 event_tx: broadcast::Sender<TaskEvent>,
163}
164
165impl Default for TaskRegistry {
166 fn default() -> Self {
167 let (event_tx, _) = broadcast::channel(256);
168 Self {
169 inner: Arc::new(std::sync::Mutex::new(HashMap::new())),
170 event_tx,
171 }
172 }
173}
174
175impl TaskRegistry {
176 pub fn new() -> Self {
177 Self::default()
178 }
179
180 pub fn register(
181 &self,
182 kind: TaskKind,
183 label: String,
184 source_handle: String,
185 session_id: String,
186 cancel: CancellationToken,
187 ) -> TaskId {
188 self.register_with_kill_hook(kind, label, source_handle, session_id, cancel, None)
189 }
190
191 pub fn register_with_kill_hook(
192 &self,
193 kind: TaskKind,
194 label: String,
195 source_handle: String,
196 session_id: String,
197 cancel: CancellationToken,
198 kill_hook: Option<std::sync::Arc<dyn Fn() + Send + Sync>>,
199 ) -> TaskId {
200 let id = TaskId::now();
201 let snapshot = TaskSnapshot {
202 id: id.clone(),
203 kind,
204 label,
205 status: TaskStatus::Running,
206 started_at: Instant::now(),
207 ended_at: None,
208 source_handle,
209 session_id,
210 };
211 let entry = TaskEntry {
212 snapshot: snapshot.clone(),
213 cancel,
214 kill_hook,
215 };
216 self.inner.lock().unwrap().insert(id.clone(), entry);
217 let _ = self.event_tx.send(TaskEvent::Registered(snapshot));
218 id
219 }
220
221 pub fn lookup(&self, id: &TaskId) -> Option<TaskSnapshot> {
222 self.inner
223 .lock()
224 .unwrap()
225 .get(id)
226 .map(|e| e.snapshot.clone())
227 }
228
229 pub fn lookup_by_handle(&self, handle: &str) -> Option<TaskSnapshot> {
231 self.inner
232 .lock()
233 .unwrap()
234 .values()
235 .find(|e| e.snapshot.source_handle == handle)
236 .map(|e| e.snapshot.clone())
237 }
238
239 pub fn list(&self, filter: &TaskFilter) -> Vec<TaskSnapshot> {
240 let inner = self.inner.lock().unwrap();
241 let mut out: Vec<TaskSnapshot> = inner
242 .values()
243 .map(|e| e.snapshot.clone())
244 .filter(|s| filter.matches(s))
245 .collect();
246 out.sort_by_key(|s| s.started_at);
247 out
248 }
249
250 pub fn kill(&self, id: &TaskId) -> bool {
251 let inner = self.inner.lock().unwrap();
252 let Some(entry) = inner.get(id) else {
253 return false;
254 };
255 if entry.snapshot.status.is_terminal() {
256 return false;
257 }
258 entry.cancel.cancel();
259 if let Some(hook) = &entry.kill_hook {
260 hook();
261 }
262 true
263 }
264
265 pub fn finish(&self, id: &TaskId, status: TaskStatus) {
268 let mut inner = self.inner.lock().unwrap();
269 let Some(entry) = inner.get_mut(id) else {
270 return;
271 };
272 if entry.snapshot.status.is_terminal() {
273 return;
274 }
275 let old = entry.snapshot.status;
276 entry.snapshot.status = status;
277 entry.snapshot.ended_at = Some(Instant::now());
278 let kind = entry.snapshot.kind;
279 drop(inner);
280 let _ = self.event_tx.send(TaskEvent::StatusChanged {
281 id: id.clone(),
282 kind,
283 old,
284 new: status,
285 });
286 }
287
288 pub fn reap(&self, id: &TaskId) {
289 let mut inner = self.inner.lock().unwrap();
290 let should_remove = inner
291 .get(id)
292 .map(|e| e.snapshot.status.is_terminal())
293 .unwrap_or(false);
294 if should_remove {
295 inner.remove(id);
296 drop(inner);
297 let _ = self.event_tx.send(TaskEvent::Reaped { id: id.clone() });
298 }
299 }
300
301 pub fn subscribe(&self) -> broadcast::Receiver<TaskEvent> {
302 self.event_tx.subscribe()
303 }
304
305 pub fn running_count(&self) -> usize {
306 self.inner
307 .lock()
308 .unwrap()
309 .values()
310 .filter(|e| e.snapshot.status.is_running())
311 .count()
312 }
313}
314
315#[cfg(test)]
316mod tests {
317 use super::*;
318
319 fn cancel() -> CancellationToken {
320 CancellationToken::new()
321 }
322
323 #[test]
324 fn register_and_lookup() {
325 let reg = TaskRegistry::new();
326 let id = reg.register(
327 TaskKind::Bash,
328 "cargo build".into(),
329 "bg_1".into(),
330 "sess".into(),
331 cancel(),
332 );
333 let snap = reg.lookup(&id).expect("found");
334 assert_eq!(snap.kind, TaskKind::Bash);
335 assert_eq!(snap.status, TaskStatus::Running);
336 assert!(snap.ended_at.is_none());
337 }
338
339 #[test]
340 fn lookup_by_handle() {
341 let reg = TaskRegistry::new();
342 let _id = reg.register(
343 TaskKind::Terminal,
344 "vim".into(),
345 "term_1".into(),
346 "sess".into(),
347 cancel(),
348 );
349 let snap = reg.lookup_by_handle("term_1").expect("found");
350 assert_eq!(snap.kind, TaskKind::Terminal);
351 assert!(reg.lookup_by_handle("nope").is_none());
352 }
353
354 #[test]
355 fn list_filters_by_kind_and_status() {
356 let reg = TaskRegistry::new();
357 let b1 = reg.register(
358 TaskKind::Bash,
359 "a".into(),
360 "bg_1".into(),
361 "s".into(),
362 cancel(),
363 );
364 let _t1 = reg.register(
365 TaskKind::Terminal,
366 "vim".into(),
367 "term_1".into(),
368 "s".into(),
369 cancel(),
370 );
371 let _b2 = reg.register(
372 TaskKind::Bash,
373 "ls".into(),
374 "bg_2".into(),
375 "s".into(),
376 cancel(),
377 );
378
379 let bash_only = reg.list(&TaskFilter {
380 kind: Some(TaskKind::Bash),
381 ..Default::default()
382 });
383 assert_eq!(bash_only.len(), 2);
384
385 reg.finish(&b1, TaskStatus::Ok);
386 let running = reg.list(&TaskFilter::running());
387 assert_eq!(running.len(), 2);
388 }
389
390 #[test]
391 fn kill_cancels_token() {
392 let reg = TaskRegistry::new();
393 let tok = cancel();
394 let id = reg.register(
395 TaskKind::Bash,
396 "x".into(),
397 "bg".into(),
398 "s".into(),
399 tok.clone(),
400 );
401 assert!(reg.kill(&id));
402 assert!(tok.is_cancelled());
403 }
404
405 #[test]
406 fn kill_returns_false_for_terminal() {
407 let reg = TaskRegistry::new();
408 let id = reg.register(
409 TaskKind::Bash,
410 "x".into(),
411 "bg".into(),
412 "s".into(),
413 cancel(),
414 );
415 reg.finish(&id, TaskStatus::Ok);
416 assert!(!reg.kill(&id));
417 }
418
419 #[test]
420 fn finish_is_idempotent() {
421 let reg = TaskRegistry::new();
422 let id = reg.register(
423 TaskKind::Bash,
424 "x".into(),
425 "bg".into(),
426 "s".into(),
427 cancel(),
428 );
429 reg.finish(&id, TaskStatus::Ok);
430 reg.finish(&id, TaskStatus::Err);
431 let snap = reg.lookup(&id).unwrap();
432 assert_eq!(snap.status, TaskStatus::Ok);
433 }
434
435 #[test]
436 fn reap_removes_terminal_only() {
437 let reg = TaskRegistry::new();
438 let id = reg.register(
439 TaskKind::Bash,
440 "x".into(),
441 "bg".into(),
442 "s".into(),
443 cancel(),
444 );
445 reg.reap(&id);
446 assert!(reg.lookup(&id).is_some());
447 reg.finish(&id, TaskStatus::Ok);
448 reg.reap(&id);
449 assert!(reg.lookup(&id).is_none());
450 }
451
452 #[test]
453 fn subscribe_receives_registered_event() {
454 let reg = TaskRegistry::new();
455 let mut rx = reg.subscribe();
456 let _id = reg.register(
457 TaskKind::Bash,
458 "x".into(),
459 "bg".into(),
460 "s".into(),
461 cancel(),
462 );
463 let ev = rx.try_recv().expect("got event");
464 match ev {
465 TaskEvent::Registered(s) => assert_eq!(s.kind, TaskKind::Bash),
466 _ => panic!("wrong event"),
467 }
468 }
469
470 #[test]
471 fn subscribe_receives_status_changed() {
472 let reg = TaskRegistry::new();
473 let mut rx = reg.subscribe();
474 let id = reg.register(
475 TaskKind::Bash,
476 "x".into(),
477 "bg".into(),
478 "s".into(),
479 cancel(),
480 );
481 let _ = rx.try_recv();
482 reg.finish(&id, TaskStatus::Ok);
483 let ev = rx.try_recv().expect("got status event");
484 match ev {
485 TaskEvent::StatusChanged { new, .. } => assert_eq!(new, TaskStatus::Ok),
486 _ => panic!("wrong event"),
487 }
488 }
489
490 #[test]
491 fn filter_matches_combines() {
492 let snap = TaskSnapshot {
493 id: TaskId::now(),
494 kind: TaskKind::Terminal,
495 label: "vim".into(),
496 status: TaskStatus::Running,
497 started_at: Instant::now(),
498 ended_at: None,
499 source_handle: "term_1".into(),
500 session_id: "sess_a".into(),
501 };
502 let f = TaskFilter {
503 kind: Some(TaskKind::Terminal),
504 status: Some(TaskStatus::Running),
505 session_id: Some("sess_a".into()),
506 };
507 assert!(f.matches(&snap));
508
509 let f2 = TaskFilter {
510 kind: Some(TaskKind::Bash),
511 ..Default::default()
512 };
513 assert!(!f2.matches(&snap));
514 }
515}