1use std::collections::HashMap;
24use std::sync::{Arc, Mutex, Weak};
25
26use async_trait::async_trait;
27
28use crate::error::Result;
29use crate::session_task::{
30 CreateSessionTask, NewTaskMessage, SessionTask, SessionTaskFilter, SessionTaskRegistry,
31 SessionTaskState, SessionTaskUpdate, TaskMessage, TaskMessageDirection,
32};
33use crate::typed_id::SessionId;
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum TaskTransition {
42 Terminal,
44 AwaitingInput,
46 Message,
48}
49
50impl TaskTransition {
51 pub fn filter_value(&self) -> &'static str {
53 match self {
54 Self::Terminal => "terminal",
55 Self::AwaitingInput => "awaiting_input",
56 Self::Message => "message",
57 }
58 }
59
60 pub fn event_name(&self) -> &'static str {
62 match self {
63 Self::Terminal => "task.terminal",
64 Self::AwaitingInput => "task.awaiting_input",
65 Self::Message => "task.message",
66 }
67 }
68}
69
70#[async_trait]
83pub trait TaskTransitionObserver: Send + Sync + 'static {
84 async fn on_transition(
87 &self,
88 task: &SessionTask,
89 transition: TaskTransition,
90 ) -> anyhow::Result<()>;
91}
92
93type TaskUpdateLocks = HashMap<(SessionId, String), Weak<tokio::sync::Mutex<()>>>;
94
95pub struct ObservingTaskRegistry {
117 inner: Arc<dyn SessionTaskRegistry>,
118 observers: Vec<Arc<dyn TaskTransitionObserver>>,
119 update_locks: Mutex<TaskUpdateLocks>,
120}
121
122impl ObservingTaskRegistry {
123 pub fn new(inner: Arc<dyn SessionTaskRegistry>) -> Self {
124 Self {
125 inner,
126 observers: Vec::new(),
127 update_locks: Mutex::new(HashMap::new()),
128 }
129 }
130
131 pub fn with_observer(mut self, observer: Arc<dyn TaskTransitionObserver>) -> Self {
133 self.observers.push(observer);
134 self
135 }
136
137 pub fn has_observers(&self) -> bool {
139 !self.observers.is_empty()
140 }
141
142 fn task_lock(&self, session_id: SessionId, task_id: &str) -> Arc<tokio::sync::Mutex<()>> {
143 let mut locks = self
144 .update_locks
145 .lock()
146 .expect("task update locks poisoned");
147 locks.retain(|_, lock| lock.strong_count() > 0);
148 let key = (session_id, task_id.to_string());
149 if let Some(lock) = locks.get(&key).and_then(Weak::upgrade) {
150 return lock;
151 }
152 let lock = Arc::new(tokio::sync::Mutex::new(()));
153 locks.insert(key, Arc::downgrade(&lock));
154 lock
155 }
156
157 async fn notify(&self, task: &SessionTask, transition: TaskTransition) {
158 for observer in &self.observers {
159 if let Err(e) = observer.on_transition(task, transition).await {
160 tracing::warn!(
161 task_id = %task.id,
162 session_id = %task.session_id,
163 transition = ?transition,
164 "TaskTransitionObserver failed (best-effort): {e}"
165 );
166 }
167 }
168 }
169}
170
171#[async_trait]
172impl SessionTaskRegistry for ObservingTaskRegistry {
173 async fn create(&self, input: CreateSessionTask) -> Result<SessionTask> {
174 self.inner.create(input).await
175 }
176
177 async fn update(
178 &self,
179 session_id: SessionId,
180 task_id: &str,
181 update: SessionTaskUpdate,
182 ) -> Result<Option<SessionTask>> {
183 let task_lock = self
186 .has_observers()
187 .then(|| self.task_lock(session_id, task_id));
188 let guard = match task_lock.as_ref() {
189 Some(lock) => Some(lock.lock().await),
190 None => None,
191 };
192 let wants_terminal = update.state.is_some_and(|s| s.is_terminal());
195 let wants_awaiting_input =
196 update.input_request.is_some() || update.state == Some(SessionTaskState::AwaitingInput);
197
198 let needs_prior = self.has_observers() && (wants_terminal || wants_awaiting_input);
199 let prior = if needs_prior {
200 self.inner.get(session_id, task_id).await.ok().flatten()
201 } else {
202 None
203 };
204
205 let updated = self.inner.update(session_id, task_id, update).await?;
206 drop(guard);
207
208 if let (Some(task), Some(prior)) = (&updated, &prior) {
209 if wants_terminal && !prior.state.is_terminal() && task.state.is_terminal() {
210 self.notify(task, TaskTransition::Terminal).await;
211 }
212 if wants_awaiting_input
213 && prior.state != SessionTaskState::AwaitingInput
214 && task.state == SessionTaskState::AwaitingInput
215 {
216 self.notify(task, TaskTransition::AwaitingInput).await;
217 }
218 }
219 Ok(updated)
220 }
221
222 async fn get(&self, session_id: SessionId, task_id: &str) -> Result<Option<SessionTask>> {
223 self.inner.get(session_id, task_id).await
224 }
225
226 async fn list(
227 &self,
228 session_id: SessionId,
229 filter: Option<&SessionTaskFilter>,
230 ) -> Result<Vec<SessionTask>> {
231 self.inner.list(session_id, filter).await
232 }
233
234 async fn request_cancel(
235 &self,
236 session_id: SessionId,
237 task_id: &str,
238 ) -> Result<Option<SessionTask>> {
239 self.inner.request_cancel(session_id, task_id).await
240 }
241
242 async fn record_message(
243 &self,
244 session_id: SessionId,
245 task_id: &str,
246 message: NewTaskMessage,
247 ) -> Result<TaskMessage> {
248 let task_lock = self
251 .has_observers()
252 .then(|| self.task_lock(session_id, task_id));
253 let guard = match task_lock.as_ref() {
254 Some(lock) => Some(lock.lock().await),
255 None => None,
256 };
257 let direction = message.direction;
258 let stored = self
259 .inner
260 .record_message(session_id, task_id, message)
261 .await?;
262 let task = if direction == TaskMessageDirection::Outbound && self.has_observers() {
263 self.inner.get(session_id, task_id).await.ok().flatten()
264 } else {
265 None
266 };
267 drop(guard);
268 if let Some(task) = task {
269 self.notify(&task, TaskTransition::Message).await;
270 }
271 Ok(stored)
272 }
273
274 async fn list_messages(
275 &self,
276 session_id: SessionId,
277 task_id: &str,
278 limit: Option<u32>,
279 after_id: Option<&str>,
280 ) -> Result<Vec<TaskMessage>> {
281 self.inner
282 .list_messages(session_id, task_id, limit, after_id)
283 .await
284 }
285}
286
287#[cfg(test)]
288mod tests {
289 use super::*;
290
291 #[test]
292 fn filter_value_and_event_name_are_stable() {
293 assert_eq!(TaskTransition::Terminal.filter_value(), "terminal");
296 assert_eq!(
297 TaskTransition::AwaitingInput.filter_value(),
298 "awaiting_input"
299 );
300 assert_eq!(TaskTransition::Message.filter_value(), "message");
301
302 assert_eq!(TaskTransition::Terminal.event_name(), "task.terminal");
303 assert_eq!(
304 TaskTransition::AwaitingInput.event_name(),
305 "task.awaiting_input"
306 );
307 assert_eq!(TaskTransition::Message.event_name(), "task.message");
308 }
309
310 use crate::session_task::{
313 SessionTaskState, TaskWakePolicy, apply_task_update, new_session_task,
314 };
315 use crate::typed_id::SessionId;
316 use std::collections::HashMap;
317 use std::sync::Mutex;
318
319 #[derive(Default)]
320 struct MemRegistry {
321 tasks: Mutex<HashMap<String, SessionTask>>,
322 yield_after_read: bool,
323 }
324
325 #[async_trait]
326 impl SessionTaskRegistry for MemRegistry {
327 async fn create(&self, input: CreateSessionTask) -> Result<SessionTask> {
328 let task = new_session_task(input, chrono::Utc::now());
329 self.tasks
330 .lock()
331 .unwrap()
332 .insert(task.id.clone(), task.clone());
333 Ok(task)
334 }
335 async fn update(
336 &self,
337 session_id: SessionId,
338 task_id: &str,
339 update: SessionTaskUpdate,
340 ) -> Result<Option<SessionTask>> {
341 let mut tasks = self.tasks.lock().unwrap();
342 let Some(task) = tasks.get_mut(task_id) else {
343 return Ok(None);
344 };
345 if task.session_id != session_id {
346 return Ok(None);
347 }
348 apply_task_update(task, update, chrono::Utc::now());
349 Ok(Some(task.clone()))
350 }
351 async fn get(&self, session_id: SessionId, task_id: &str) -> Result<Option<SessionTask>> {
352 let task = self
353 .tasks
354 .lock()
355 .unwrap()
356 .get(task_id)
357 .filter(|t| t.session_id == session_id)
358 .cloned();
359 if self.yield_after_read {
360 tokio::task::yield_now().await;
361 }
362 Ok(task)
363 }
364 async fn list(
365 &self,
366 _session_id: SessionId,
367 _filter: Option<&SessionTaskFilter>,
368 ) -> Result<Vec<SessionTask>> {
369 Ok(Vec::new())
370 }
371 async fn request_cancel(
372 &self,
373 _session_id: SessionId,
374 _task_id: &str,
375 ) -> Result<Option<SessionTask>> {
376 Ok(None)
377 }
378 async fn record_message(
379 &self,
380 _session_id: SessionId,
381 task_id: &str,
382 message: NewTaskMessage,
383 ) -> Result<TaskMessage> {
384 Ok(TaskMessage {
385 id: "tmsg_x".into(),
386 task_id: task_id.into(),
387 direction: message.direction,
388 content: message.content,
389 in_reply_to: message.in_reply_to,
390 created_at: chrono::Utc::now(),
391 })
392 }
393 async fn list_messages(
394 &self,
395 _session_id: SessionId,
396 _task_id: &str,
397 _limit: Option<u32>,
398 _after_id: Option<&str>,
399 ) -> Result<Vec<TaskMessage>> {
400 Ok(Vec::new())
401 }
402 }
403
404 #[derive(Default)]
405 struct Recorder {
406 seen: Mutex<Vec<TaskTransition>>,
407 snapshots: Mutex<Vec<serde_json::Value>>,
408 }
409
410 #[async_trait]
411 impl TaskTransitionObserver for Recorder {
412 async fn on_transition(
413 &self,
414 task: &SessionTask,
415 transition: TaskTransition,
416 ) -> anyhow::Result<()> {
417 self.snapshots
418 .lock()
419 .unwrap()
420 .push(serde_json::to_value(task).unwrap());
421 self.seen.lock().unwrap().push(transition);
422 Ok(())
423 }
424 }
425
426 async fn seed_running(reg: &MemRegistry, session_id: SessionId) -> String {
427 reg.create(CreateSessionTask {
428 id: None,
429 session_id,
430 kind: "background_tool".into(),
431 display_name: "T".into(),
432 spec: serde_json::Value::Null,
433 state: SessionTaskState::Running,
434 links: Default::default(),
435 wake_policy: TaskWakePolicy::OnActivity,
436 })
437 .await
438 .unwrap()
439 .id
440 }
441
442 #[tokio::test]
443 async fn fires_terminal_once_and_not_on_heartbeat() {
444 let inner = Arc::new(MemRegistry::default());
445 let recorder = Arc::new(Recorder::default());
446 let reg = ObservingTaskRegistry::new(inner.clone()).with_observer(recorder.clone());
447 let session_id = SessionId::new();
448 let task_id = seed_running(&inner, session_id).await;
449
450 reg.update(
452 session_id,
453 &task_id,
454 SessionTaskUpdate {
455 heartbeat_at: Some(chrono::Utc::now()),
456 ..Default::default()
457 },
458 )
459 .await
460 .unwrap();
461 assert!(
462 recorder.seen.lock().unwrap().is_empty(),
463 "heartbeat must not fire a transition"
464 );
465
466 reg.update(
468 session_id,
469 &task_id,
470 SessionTaskUpdate {
471 state: Some(SessionTaskState::Succeeded),
472 ..Default::default()
473 },
474 )
475 .await
476 .unwrap();
477
478 reg.update(
481 session_id,
482 &task_id,
483 SessionTaskUpdate {
484 state: Some(SessionTaskState::Succeeded),
485 ..Default::default()
486 },
487 )
488 .await
489 .unwrap();
490
491 assert_eq!(
492 *recorder.seen.lock().unwrap(),
493 vec![TaskTransition::Terminal],
494 "terminal fires exactly once, never on heartbeat or re-terminal"
495 );
496 }
497
498 #[tokio::test]
499 async fn fires_awaiting_input_only_on_entry() {
500 let inner = Arc::new(MemRegistry::default());
501 let recorder = Arc::new(Recorder::default());
502 let reg = ObservingTaskRegistry::new(inner.clone()).with_observer(recorder.clone());
503 let session_id = SessionId::new();
504 let task_id = seed_running(&inner, session_id).await;
505
506 reg.update(
507 session_id,
508 &task_id,
509 SessionTaskUpdate {
510 input_request: Some(crate::session_task::TaskInputRequest {
511 id: "ir_1".into(),
512 prompt: "approve?".into(),
513 expected: None,
514 }),
515 ..Default::default()
516 },
517 )
518 .await
519 .unwrap();
520
521 assert_eq!(
522 *recorder.seen.lock().unwrap(),
523 vec![TaskTransition::AwaitingInput],
524 "awaiting_input fires once on entry"
525 );
526 reg.update(
527 session_id,
528 &task_id,
529 SessionTaskUpdate {
530 state: Some(SessionTaskState::AwaitingInput),
531 ..Default::default()
532 },
533 )
534 .await
535 .unwrap();
536 assert_eq!(
537 *recorder.seen.lock().unwrap(),
538 vec![TaskTransition::AwaitingInput]
539 );
540 reg.update(
541 session_id,
542 &task_id,
543 SessionTaskUpdate {
544 state: Some(SessionTaskState::Running),
545 ..Default::default()
546 },
547 )
548 .await
549 .unwrap();
550 let updated = reg
551 .update(
552 session_id,
553 &task_id,
554 SessionTaskUpdate {
555 input_request: Some(crate::session_task::TaskInputRequest {
556 id: "ir_2".into(),
557 prompt: "choose again".into(),
558 expected: None,
559 }),
560 ..Default::default()
561 },
562 )
563 .await
564 .unwrap()
565 .unwrap();
566 assert_eq!(
567 *recorder.seen.lock().unwrap(),
568 vec![TaskTransition::AwaitingInput, TaskTransition::AwaitingInput]
569 );
570 assert_eq!(
571 recorder.snapshots.lock().unwrap().last().unwrap(),
572 &serde_json::to_value(updated).unwrap()
573 );
574 }
575 #[tokio::test]
576 async fn competing_terminal_updates_emit_one_transition() {
577 let inner = Arc::new(MemRegistry {
578 yield_after_read: true,
579 ..Default::default()
580 });
581 let recorder = Arc::new(Recorder::default());
582 let reg = ObservingTaskRegistry::new(inner.clone()).with_observer(recorder.clone());
583 let session = SessionId::from_seed(1);
584 let task = seed_running(&inner, session).await;
585 let update = || SessionTaskUpdate {
586 state: Some(SessionTaskState::Succeeded),
587 ..Default::default()
588 };
589 let (first, second) = tokio::join!(
590 reg.update(session, &task, update()),
591 reg.update(session, &task, update())
592 );
593 assert_eq!(first.unwrap().unwrap().state, SessionTaskState::Succeeded);
594 assert_eq!(second.unwrap().unwrap().state, SessionTaskState::Succeeded);
595 assert_eq!(
596 *recorder.seen.lock().unwrap(),
597 vec![TaskTransition::Terminal]
598 );
599 }
600 #[derive(Default)]
601 struct FailingObserver {
602 registry: Mutex<Weak<ObservingTaskRegistry>>,
603 }
604 #[async_trait]
605 impl TaskTransitionObserver for FailingObserver {
606 async fn on_transition(&self, task: &SessionTask, _: TaskTransition) -> anyhow::Result<()> {
607 let registry = self.registry.lock().unwrap().upgrade().unwrap();
608 let lock = registry.task_lock(task.session_id, &task.id);
609 assert!(
610 lock.try_lock().is_ok(),
611 "callbacks must run outside the update lock"
612 );
613 anyhow::bail!("observer unavailable")
614 }
615 }
616
617 #[tokio::test]
618 async fn outbound_messages_preserve_payload_and_survive_observer_failure() {
619 let inner = Arc::new(MemRegistry::default());
620 let recorder = Arc::new(Recorder::default());
621 let failing = Arc::new(FailingObserver::default());
622 let reg = Arc::new(
623 ObservingTaskRegistry::new(inner.clone())
624 .with_observer(failing.clone())
625 .with_observer(recorder.clone()),
626 );
627 *failing.registry.lock().unwrap() = Arc::downgrade(®);
628 let session = SessionId::from_seed(2);
629 let task = seed_running(&inner, session).await;
630 let inbound = reg
631 .record_message(session, &task, NewTaskMessage::inbound_text("answer"))
632 .await
633 .unwrap();
634 assert_eq!(inbound.direction, TaskMessageDirection::Inbound);
635 assert!(recorder.seen.lock().unwrap().is_empty());
636 for text in ["progress α", "finished"] {
637 let mut message = NewTaskMessage::outbound_text(text);
638 message.in_reply_to = Some("request_1".into());
639 let saved = reg.record_message(session, &task, message).await.unwrap();
640 assert_eq!(saved.task_id, task);
641 assert_eq!(saved.direction, TaskMessageDirection::Outbound);
642 assert_eq!(
643 saved.content,
644 vec![crate::session_task::TaskMessagePart::text(text)]
645 );
646 assert_eq!(saved.in_reply_to.as_deref(), Some("request_1"));
647 }
648 assert_eq!(
649 *recorder.seen.lock().unwrap(),
650 vec![TaskTransition::Message, TaskTransition::Message]
651 );
652 let snapshot =
653 serde_json::to_value(inner.get(session, &task).await.unwrap().unwrap()).unwrap();
654 assert_eq!(
655 *recorder.snapshots.lock().unwrap(),
656 vec![snapshot.clone(), snapshot]
657 );
658 let updated = reg
659 .update(
660 session,
661 &task,
662 SessionTaskUpdate {
663 state: Some(SessionTaskState::Failed),
664 summary: Some("failed work".into()),
665 ..Default::default()
666 },
667 )
668 .await
669 .unwrap()
670 .unwrap();
671 assert_eq!(updated.state, SessionTaskState::Failed);
672 assert_eq!(
673 recorder.seen.lock().unwrap().last(),
674 Some(&TaskTransition::Terminal)
675 );
676 assert_eq!(
677 recorder.snapshots.lock().unwrap().last().unwrap(),
678 &serde_json::to_value(updated).unwrap()
679 );
680 }
681
682 #[test]
683 fn task_locks_isolate_keys_and_release_idle_entries() {
684 let reg = ObservingTaskRegistry::new(Arc::new(MemRegistry::default()));
685 let first = reg.task_lock(SessionId::from_seed(1), "task_a");
686 let same = reg.task_lock(SessionId::from_seed(1), "task_a");
687 let other_task = reg.task_lock(SessionId::from_seed(1), "task_b");
688 let other_session = reg.task_lock(SessionId::from_seed(2), "task_a");
689 let guard = first.try_lock().unwrap();
690 assert!(same.try_lock().is_err());
691 assert!(other_task.try_lock().is_ok());
692 assert!(other_session.try_lock().is_ok());
693 drop(guard);
694 drop((first, same, other_task, other_session));
695 let next = reg.task_lock(SessionId::from_seed(3), "task_c");
696 assert_eq!(reg.update_locks.lock().unwrap().len(), 1);
697 assert!(next.try_lock().is_ok());
698 }
699}