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#[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 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 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}