Skip to main content

everruns_core/
task_observer.rs

1// Embeddable in-process task-transition observer (EVE-729).
2//
3// A `TaskTransition` names a lifecycle transition a `SessionTaskRegistry` fires
4// on: reaching a terminal state, entering `awaiting_input`, or emitting an
5// outbound message. A `TaskTransitionObserver` receives those transitions in
6// process — the same seam the server's webhook dispatcher uses, minus HTTP.
7//
8// Design Decision: the enum + trait live in `everruns-core` (not the server) so
9// `everruns-host` embedders can observe task transitions without depending on
10// the control-plane server or making HTTP calls. The server webhook dispatcher
11// (`DirectTaskWebhookNotifier`) is one implementation of this trait; in-process
12// embedders provide their own. A `SessionTaskRegistry` fires each real
13// transition once to every registered observer, so an in-process observer sees
14// exactly the same transitions the webhook path fires (see the parity test in
15// `crates/server/src/storage/session_task_store.rs`).
16//
17// Filter semantics: `Terminal` is the regression-safe default (org webhooks only
18// ever fire on it); `AwaitingInput` and `Message` are the non-terminal
19// transitions that are opt-in per delivery target via `event_filter` (EVE-682).
20// The `filter_value` / `event_name` strings are shared with webhook payloads so
21// the two paths stay byte-for-byte aligned.
22
23use 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/// A task lifecycle transition an observer can be notified of.
36///
37/// `Terminal` is the only transition org webhooks ever fire on (regression-safe).
38/// `AwaitingInput` and `Message` are opt-in per delivery target via
39/// `event_filter` (EVE-682).
40#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum TaskTransition {
42    /// Task reached a terminal state (succeeded / failed / canceled).
43    Terminal,
44    /// Task transitioned into `awaiting_input`.
45    AwaitingInput,
46    /// Task emitted an outbound message.
47    Message,
48}
49
50impl TaskTransition {
51    /// The `event_filter` member string that enables this transition.
52    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    /// The `event` field value in a delivered webhook payload.
61    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/// Receive task-transition notifications in process.
71///
72/// A `SessionTaskRegistry` invokes `on_transition` once per real transition for
73/// every registered observer. Implementations must treat delivery as
74/// best-effort: the registry logs errors and never fails the underlying task
75/// operation because an observer returned `Err`. Observers must not block for
76/// long — the registry dispatches them off the task-update path, but a slow
77/// observer still delays its own delivery.
78///
79/// The server webhook dispatcher (`DirectTaskWebhookNotifier`) is one
80/// implementation. Embedders of `everruns-host` implement this trait to get
81/// in-process callbacks with the same transition semantics, without HTTP.
82#[async_trait]
83pub trait TaskTransitionObserver: Send + Sync + 'static {
84    /// Handle one task transition. Best-effort: returning `Err` is logged and
85    /// never fails the task operation that produced the transition.
86    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
95/// A [`SessionTaskRegistry`] decorator that fans real task transitions out to
96/// registered [`TaskTransitionObserver`]s.
97///
98/// This is the reusable, storage-agnostic form of the fan-out the server's
99/// `DbSessionTaskRegistry` performs inline (EVE-729): it wraps *any* inner
100/// registry (in-memory, SQLite, gRPC) so an embedder — e.g. `everruns-host`
101/// with a [`crate::wake_queue::SessionWakeQueue`] — gets the same transition
102/// notifications without depending on the control-plane server.
103///
104/// Transition detection mirrors `DbSessionTaskRegistry` exactly so mid-turn and
105/// between-turn delivery agree on *when* a wake fires:
106///
107///   * `Terminal` — fired once when an update moves a non-terminal task into a
108///     terminal state, gated on this update's own intent (an update that does
109///     not set a terminal state never fires it, so a racing heartbeat cannot
110///     wake on another writer's transition).
111///   * `AwaitingInput` — fired once on the transition into `awaiting_input`.
112///   * `Message` — fired for each outbound message.
113///
114/// Observers are awaited in registration order (fast, in-process consumers); a
115/// failing observer is logged and never fails the underlying task op.
116pub 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    /// Register an observer to receive every real transition.
132    pub fn with_observer(mut self, observer: Arc<dyn TaskTransitionObserver>) -> Self {
133        self.observers.push(observer);
134        self
135    }
136
137    /// Whether any observer is registered (fan-out is otherwise a no-op).
138    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        // Keep the prior snapshot and mutation together, so competing writers
184        // cannot both claim the same transition. Callbacks run after unlocking.
185        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        // Only this update's own intent can trigger a wake, so a racing
193        // heartbeat/progress update never fires on another writer's transition.
194        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        // Inbound replies can leave awaiting_input, so message writes share
249        // the task update lock even when they do not notify observers.
250        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        // These strings are a wire contract shared with webhook payloads and the
294        // per-task `event_filter`; changing them silently breaks delivery.
295        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    // ---- ObservingTaskRegistry fan-out gating -----------------------------
311
312    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        // A heartbeat-only update (no state change) must not fire any transition.
451        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        // Transition to terminal fires exactly one Terminal.
467        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        // A redundant terminal-state update on an already-terminal task must not
479        // re-fire (prior is already terminal).
480        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(&reg);
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}