Skip to main content

atman_runtime/
task_registry.rs

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