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    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/// Central management layer for all observable tasks.
154///
155/// Sub-registries (BgRegistry, TermRegistry, Executor, agent_ctrl) register
156/// tasks here on spawn and call `finish` when the task ends. Typed operations
157/// (bash.output, term.input, term.capture) stay on the sub-registries and are
158/// looked up by `source_handle`.
159#[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    /// Find a task by its source handle (e.g. "bg_1", "term_2", run_id).
230    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    /// Transition a task to a terminal status. Called by the owning
266    /// sub-registry when the task finishes.
267    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}