use std::sync::Arc;
use async_trait::async_trait;
use crate::error::Result;
use crate::session_task::{
CreateSessionTask, NewTaskMessage, SessionTask, SessionTaskFilter, SessionTaskRegistry,
SessionTaskState, SessionTaskUpdate, TaskMessage, TaskMessageDirection,
};
use crate::typed_id::SessionId;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TaskTransition {
Terminal,
AwaitingInput,
Message,
}
impl TaskTransition {
pub fn filter_value(&self) -> &'static str {
match self {
Self::Terminal => "terminal",
Self::AwaitingInput => "awaiting_input",
Self::Message => "message",
}
}
pub fn event_name(&self) -> &'static str {
match self {
Self::Terminal => "task.terminal",
Self::AwaitingInput => "task.awaiting_input",
Self::Message => "task.message",
}
}
}
#[async_trait]
pub trait TaskTransitionObserver: Send + Sync + 'static {
async fn on_transition(
&self,
task: &SessionTask,
transition: TaskTransition,
) -> anyhow::Result<()>;
}
pub struct ObservingTaskRegistry {
inner: Arc<dyn SessionTaskRegistry>,
observers: Vec<Arc<dyn TaskTransitionObserver>>,
}
impl ObservingTaskRegistry {
pub fn new(inner: Arc<dyn SessionTaskRegistry>) -> Self {
Self {
inner,
observers: Vec::new(),
}
}
pub fn with_observer(mut self, observer: Arc<dyn TaskTransitionObserver>) -> Self {
self.observers.push(observer);
self
}
pub fn has_observers(&self) -> bool {
!self.observers.is_empty()
}
async fn notify(&self, task: &SessionTask, transition: TaskTransition) {
for observer in &self.observers {
if let Err(e) = observer.on_transition(task, transition).await {
tracing::warn!(
task_id = %task.id,
session_id = %task.session_id,
transition = ?transition,
"TaskTransitionObserver failed (best-effort): {e}"
);
}
}
}
}
#[async_trait]
impl SessionTaskRegistry for ObservingTaskRegistry {
async fn create(&self, input: CreateSessionTask) -> Result<SessionTask> {
self.inner.create(input).await
}
async fn update(
&self,
session_id: SessionId,
task_id: &str,
update: SessionTaskUpdate,
) -> Result<Option<SessionTask>> {
let wants_terminal = update.state.is_some_and(|s| s.is_terminal());
let wants_awaiting_input =
update.input_request.is_some() || update.state == Some(SessionTaskState::AwaitingInput);
let needs_prior = self.has_observers() && (wants_terminal || wants_awaiting_input);
let prior = if needs_prior {
self.inner.get(session_id, task_id).await.ok().flatten()
} else {
None
};
let updated = self.inner.update(session_id, task_id, update).await?;
if let (Some(task), Some(prior)) = (&updated, &prior) {
if wants_terminal && !prior.state.is_terminal() && task.state.is_terminal() {
self.notify(task, TaskTransition::Terminal).await;
}
if wants_awaiting_input
&& prior.state != SessionTaskState::AwaitingInput
&& task.state == SessionTaskState::AwaitingInput
{
self.notify(task, TaskTransition::AwaitingInput).await;
}
}
Ok(updated)
}
async fn get(&self, session_id: SessionId, task_id: &str) -> Result<Option<SessionTask>> {
self.inner.get(session_id, task_id).await
}
async fn list(
&self,
session_id: SessionId,
filter: Option<&SessionTaskFilter>,
) -> Result<Vec<SessionTask>> {
self.inner.list(session_id, filter).await
}
async fn request_cancel(
&self,
session_id: SessionId,
task_id: &str,
) -> Result<Option<SessionTask>> {
self.inner.request_cancel(session_id, task_id).await
}
async fn record_message(
&self,
session_id: SessionId,
task_id: &str,
message: NewTaskMessage,
) -> Result<TaskMessage> {
let direction = message.direction;
let stored = self
.inner
.record_message(session_id, task_id, message)
.await?;
if direction == TaskMessageDirection::Outbound
&& self.has_observers()
&& let Ok(Some(task)) = self.inner.get(session_id, task_id).await
{
self.notify(&task, TaskTransition::Message).await;
}
Ok(stored)
}
async fn list_messages(
&self,
session_id: SessionId,
task_id: &str,
limit: Option<u32>,
after_id: Option<&str>,
) -> Result<Vec<TaskMessage>> {
self.inner
.list_messages(session_id, task_id, limit, after_id)
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn filter_value_and_event_name_are_stable() {
assert_eq!(TaskTransition::Terminal.filter_value(), "terminal");
assert_eq!(
TaskTransition::AwaitingInput.filter_value(),
"awaiting_input"
);
assert_eq!(TaskTransition::Message.filter_value(), "message");
assert_eq!(TaskTransition::Terminal.event_name(), "task.terminal");
assert_eq!(
TaskTransition::AwaitingInput.event_name(),
"task.awaiting_input"
);
assert_eq!(TaskTransition::Message.event_name(), "task.message");
}
use crate::session_task::{
SessionTaskState, TaskWakePolicy, apply_task_update, new_session_task,
};
use crate::typed_id::SessionId;
use std::collections::HashMap;
use std::sync::Mutex;
#[derive(Default)]
struct MemRegistry {
tasks: Mutex<HashMap<String, SessionTask>>,
}
#[async_trait]
impl SessionTaskRegistry for MemRegistry {
async fn create(&self, input: CreateSessionTask) -> Result<SessionTask> {
let task = new_session_task(input, chrono::Utc::now());
self.tasks
.lock()
.unwrap()
.insert(task.id.clone(), task.clone());
Ok(task)
}
async fn update(
&self,
session_id: SessionId,
task_id: &str,
update: SessionTaskUpdate,
) -> Result<Option<SessionTask>> {
let mut tasks = self.tasks.lock().unwrap();
let Some(task) = tasks.get_mut(task_id) else {
return Ok(None);
};
if task.session_id != session_id {
return Ok(None);
}
apply_task_update(task, update, chrono::Utc::now());
Ok(Some(task.clone()))
}
async fn get(&self, session_id: SessionId, task_id: &str) -> Result<Option<SessionTask>> {
Ok(self
.tasks
.lock()
.unwrap()
.get(task_id)
.filter(|t| t.session_id == session_id)
.cloned())
}
async fn list(
&self,
_session_id: SessionId,
_filter: Option<&SessionTaskFilter>,
) -> Result<Vec<SessionTask>> {
Ok(Vec::new())
}
async fn request_cancel(
&self,
_session_id: SessionId,
_task_id: &str,
) -> Result<Option<SessionTask>> {
Ok(None)
}
async fn record_message(
&self,
_session_id: SessionId,
task_id: &str,
message: NewTaskMessage,
) -> Result<TaskMessage> {
Ok(TaskMessage {
id: "tmsg_x".into(),
task_id: task_id.into(),
direction: message.direction,
content: message.content,
in_reply_to: message.in_reply_to,
created_at: chrono::Utc::now(),
})
}
async fn list_messages(
&self,
_session_id: SessionId,
_task_id: &str,
_limit: Option<u32>,
_after_id: Option<&str>,
) -> Result<Vec<TaskMessage>> {
Ok(Vec::new())
}
}
#[derive(Default)]
struct Recorder {
seen: Mutex<Vec<TaskTransition>>,
}
#[async_trait]
impl TaskTransitionObserver for Recorder {
async fn on_transition(
&self,
_task: &SessionTask,
transition: TaskTransition,
) -> anyhow::Result<()> {
self.seen.lock().unwrap().push(transition);
Ok(())
}
}
async fn seed_running(reg: &MemRegistry, session_id: SessionId) -> String {
reg.create(CreateSessionTask {
id: None,
session_id,
kind: "background_tool".into(),
display_name: "T".into(),
spec: serde_json::Value::Null,
state: SessionTaskState::Running,
links: Default::default(),
wake_policy: TaskWakePolicy::OnActivity,
})
.await
.unwrap()
.id
}
#[tokio::test]
async fn fires_terminal_once_and_not_on_heartbeat() {
let inner = Arc::new(MemRegistry::default());
let recorder = Arc::new(Recorder::default());
let reg = ObservingTaskRegistry::new(inner.clone()).with_observer(recorder.clone());
let session_id = SessionId::new();
let task_id = seed_running(&inner, session_id).await;
reg.update(
session_id,
&task_id,
SessionTaskUpdate {
heartbeat_at: Some(chrono::Utc::now()),
..Default::default()
},
)
.await
.unwrap();
assert!(
recorder.seen.lock().unwrap().is_empty(),
"heartbeat must not fire a transition"
);
reg.update(
session_id,
&task_id,
SessionTaskUpdate {
state: Some(SessionTaskState::Succeeded),
..Default::default()
},
)
.await
.unwrap();
reg.update(
session_id,
&task_id,
SessionTaskUpdate {
state: Some(SessionTaskState::Succeeded),
..Default::default()
},
)
.await
.unwrap();
assert_eq!(
*recorder.seen.lock().unwrap(),
vec![TaskTransition::Terminal],
"terminal fires exactly once, never on heartbeat or re-terminal"
);
}
#[tokio::test]
async fn fires_awaiting_input_only_on_entry() {
let inner = Arc::new(MemRegistry::default());
let recorder = Arc::new(Recorder::default());
let reg = ObservingTaskRegistry::new(inner.clone()).with_observer(recorder.clone());
let session_id = SessionId::new();
let task_id = seed_running(&inner, session_id).await;
reg.update(
session_id,
&task_id,
SessionTaskUpdate {
input_request: Some(crate::session_task::TaskInputRequest {
id: "ir_1".into(),
prompt: "approve?".into(),
expected: None,
}),
..Default::default()
},
)
.await
.unwrap();
assert_eq!(
*recorder.seen.lock().unwrap(),
vec![TaskTransition::AwaitingInput],
"awaiting_input fires once on entry"
);
}
}