use std::collections::{HashMap, VecDeque};
use std::sync::Mutex;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use crate::session_task::{SessionTask, TaskWakePolicy};
use crate::task_observer::{TaskTransition, TaskTransitionObserver};
use crate::typed_id::SessionId;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PendingWake {
pub task_id: String,
pub session_id: SessionId,
pub transition: TaskTransition,
pub text: String,
pub created_at: DateTime<Utc>,
}
pub fn wake_text_for(task: &SessionTask, transition: TaskTransition) -> Option<String> {
match (task.wake_policy, transition) {
(TaskWakePolicy::Silent, _) => None,
(TaskWakePolicy::OnTerminal, TaskTransition::Terminal)
| (TaskWakePolicy::OnActivity, TaskTransition::Terminal) => {
let mut parts = vec![format!(
"Task \"{}\" ({}) finished: {}.",
task.display_name, task.id, task.state
)];
if let Some(summary) = &task.summary {
parts.push(format!("- summary: {summary}"));
}
if let Some(result_path) = &task.result_path {
parts.push(format!("- result_path: {result_path}"));
}
Some(parts.join("\n"))
}
(TaskWakePolicy::OnTerminal, _) => None,
(TaskWakePolicy::OnActivity, TaskTransition::AwaitingInput) => {
let prompt = task
.input_request
.as_ref()
.map(|r| r.prompt.as_str())
.unwrap_or("Task is awaiting input.");
Some(format!(
"Task \"{}\" ({}) is awaiting input: {}",
task.display_name, task.id, prompt
))
}
(TaskWakePolicy::OnActivity, TaskTransition::Message) => {
let detail = task
.state_detail
.as_deref()
.filter(|s| !s.trim().is_empty())
.or_else(|| task.progress.as_ref().and_then(|p| p.label.as_deref()))
.unwrap_or("structured progress update");
Some(format!(
"Task \"{}\" ({}) sent a message: {}",
task.display_name, task.id, detail
))
}
}
}
#[derive(Default)]
pub struct SessionWakeQueue {
queues: Mutex<HashMap<SessionId, VecDeque<PendingWake>>>,
}
impl SessionWakeQueue {
pub fn new() -> Self {
Self::default()
}
pub fn note_transition(&self, task: &SessionTask, transition: TaskTransition) -> bool {
let Some(text) = wake_text_for(task, transition) else {
return false;
};
let wake = PendingWake {
task_id: task.id.clone(),
session_id: task.session_id,
transition,
text,
created_at: Utc::now(),
};
self.queues
.lock()
.expect("wake queue mutex poisoned")
.entry(task.session_id)
.or_default()
.push_back(wake);
true
}
pub fn drain(&self, session_id: SessionId) -> Vec<PendingWake> {
let mut guard = self.queues.lock().expect("wake queue mutex poisoned");
match guard.get_mut(&session_id) {
Some(queue) => queue.drain(..).collect(),
None => Vec::new(),
}
}
pub fn pending_len(&self, session_id: SessionId) -> usize {
self.queues
.lock()
.expect("wake queue mutex poisoned")
.get(&session_id)
.map_or(0, VecDeque::len)
}
pub fn has_pending(&self, session_id: SessionId) -> bool {
self.pending_len(session_id) > 0
}
}
#[async_trait]
impl TaskTransitionObserver for SessionWakeQueue {
async fn on_transition(
&self,
task: &SessionTask,
transition: TaskTransition,
) -> anyhow::Result<()> {
self.note_transition(task, transition);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::session_task::{
CreateSessionTask, SessionTaskState, TaskInputRequest, new_session_task,
};
fn task_with_policy(policy: TaskWakePolicy) -> SessionTask {
let session_id = SessionId::new();
let mut task = new_session_task(
CreateSessionTask {
id: None,
session_id,
kind: "subagent".into(),
display_name: "Test Runner".into(),
spec: serde_json::Value::Null,
state: SessionTaskState::Running,
links: Default::default(),
wake_policy: policy,
},
Utc::now(),
);
task.state = SessionTaskState::Succeeded;
task
}
#[test]
fn silent_never_wakes() {
for transition in [
TaskTransition::Terminal,
TaskTransition::AwaitingInput,
TaskTransition::Message,
] {
assert!(wake_text_for(&task_with_policy(TaskWakePolicy::Silent), transition).is_none());
}
}
#[test]
fn on_terminal_wakes_only_on_terminal() {
let task = task_with_policy(TaskWakePolicy::OnTerminal);
assert!(wake_text_for(&task, TaskTransition::Terminal).is_some());
assert!(wake_text_for(&task, TaskTransition::AwaitingInput).is_none());
assert!(wake_text_for(&task, TaskTransition::Message).is_none());
}
#[test]
fn on_activity_wakes_on_all_three() {
let mut task = task_with_policy(TaskWakePolicy::OnActivity);
task.state = SessionTaskState::AwaitingInput;
task.input_request = Some(TaskInputRequest {
id: "ir_1".into(),
prompt: "pick a branch".into(),
expected: None,
});
let awaiting = wake_text_for(&task, TaskTransition::AwaitingInput).expect("awaiting wakes");
assert!(awaiting.contains("awaiting input"));
assert!(awaiting.contains("pick a branch"));
task.state = SessionTaskState::Running;
task.state_detail = Some("iteration 4/10".into());
let message = wake_text_for(&task, TaskTransition::Message).expect("message wakes");
assert!(message.contains("iteration 4/10"));
task.state = SessionTaskState::Failed;
assert!(wake_text_for(&task, TaskTransition::Terminal).is_some());
}
#[test]
fn terminal_text_includes_summary_and_result_path() {
let mut task = task_with_policy(TaskWakePolicy::OnTerminal);
task.summary = Some("all tests passed".into());
task.result_path = Some("/.tasks/task_x/result.json".into());
let text = wake_text_for(&task, TaskTransition::Terminal).unwrap();
assert!(text.contains("finished: succeeded"));
assert!(text.contains("all tests passed"));
assert!(text.contains("/.tasks/task_x/result.json"));
}
#[test]
fn drain_is_exactly_once() {
let queue = SessionWakeQueue::new();
let task = task_with_policy(TaskWakePolicy::OnTerminal);
let session_id = task.session_id;
assert!(queue.note_transition(&task, TaskTransition::Terminal));
assert_eq!(queue.pending_len(session_id), 1);
let first = queue.drain(session_id);
assert_eq!(first.len(), 1, "first drain returns the wake");
assert_eq!(first[0].task_id, task.id);
let second = queue.drain(session_id);
assert!(
second.is_empty(),
"second drain returns nothing — claimed once"
);
assert_eq!(queue.pending_len(session_id), 0);
}
#[test]
fn silent_transition_enqueues_nothing() {
let queue = SessionWakeQueue::new();
let task = task_with_policy(TaskWakePolicy::Silent);
assert!(!queue.note_transition(&task, TaskTransition::Terminal));
assert_eq!(queue.pending_len(task.session_id), 0);
}
#[test]
fn queues_are_isolated_per_session() {
let queue = SessionWakeQueue::new();
let a = task_with_policy(TaskWakePolicy::OnTerminal);
let b = task_with_policy(TaskWakePolicy::OnTerminal);
queue.note_transition(&a, TaskTransition::Terminal);
queue.note_transition(&b, TaskTransition::Terminal);
assert_eq!(queue.drain(a.session_id).len(), 1);
assert_eq!(
queue.pending_len(b.session_id),
1,
"draining A leaves B intact"
);
}
}