1use super::{Capability, CapabilityLocalization, CapabilityStatus};
15use crate::session_task::{
16 NewTaskMessage, SessionTask, SessionTaskFilter, SessionTaskRegistry, SessionTaskState,
17 TaskMessage, find_task_executor,
18};
19use crate::tool_types::ToolHints;
20use crate::tools::{Tool, ToolExecutionResult};
21use crate::traits::{ToolContext, ToolContextService};
22use async_trait::async_trait;
23use serde_json::{Value, json};
24use std::sync::Arc;
25use std::time::Duration;
26use tokio::time::{Instant, sleep};
27
28pub const SESSION_TASKS_CAPABILITY_ID: &str = "session_tasks";
29
30const DEFAULT_WAIT_TIMEOUT_SECS: u64 = 300;
31const WAIT_POLL_INTERVAL: Duration = Duration::from_secs(1);
32const WAIT_RECONCILE_EVERY: u64 = 5;
34const GET_TASK_MESSAGE_LIMIT: u32 = 20;
36
37pub struct SessionTasksCapability;
39
40impl Capability for SessionTasksCapability {
41 fn id(&self) -> &str {
42 SESSION_TASKS_CAPABILITY_ID
43 }
44
45 fn name(&self) -> &str {
46 "Session Tasks"
47 }
48
49 fn description(&self) -> &str {
50 "Track, message, cancel, and wait on the session's background tasks (subagents, external agents, background tools)."
51 }
52
53 fn localizations(&self) -> Vec<CapabilityLocalization> {
54 vec![CapabilityLocalization::text(
55 "uk",
56 "Завдання сесії",
57 "Відстежуйте фонові завдання сесії (субагенти, зовнішні агенти, фонові інструменти), надсилайте їм повідомлення, скасовуйте їх та очікуйте на їхнє завершення.",
58 )]
59 }
60
61 fn status(&self) -> CapabilityStatus {
62 CapabilityStatus::Available
63 }
64
65 fn icon(&self) -> Option<&str> {
66 Some("list-checks")
67 }
68
69 fn category(&self) -> Option<&str> {
70 Some("Orchestration")
71 }
72
73 fn features(&self) -> Vec<&'static str> {
74 vec!["session_tasks"]
75 }
76
77 fn system_prompt_addition(&self) -> Option<&str> {
78 Some(SESSION_TASKS_SYSTEM_PROMPT)
79 }
80
81 fn tools(&self) -> Vec<Box<dyn Tool>> {
82 vec![
83 Box::new(ListTasksTool),
84 Box::new(GetTaskTool),
85 Box::new(MessageTaskTool),
86 Box::new(CancelTaskTool),
87 Box::new(WaitTaskTool),
88 ]
89 }
90}
91
92const SESSION_TASKS_SYSTEM_PROMPT: &str = "Every spawned background work item (subagent, external agent, background tool) is a task with a task_id. Use list_tasks/get_task to check status instead of re-spawning. Answer a task in awaiting_input with message_task (set in_reply_to to the pending input request id). Use wait_task only when you have nothing else to do until the task finishes.";
93
94fn require_task_registry(
99 context: &ToolContext,
100) -> Result<&Arc<dyn SessionTaskRegistry>, ToolExecutionResult> {
101 context.session_task_registry.as_ref().ok_or_else(|| {
102 ToolExecutionResult::tool_error(
103 "Session task tools require session_task_registry context (not available in this environment)",
104 )
105 })
106}
107
108use super::util::require_str_trimmed as require_str;
109
110async fn load_task(
111 context: &ToolContext,
112 task_id: &str,
113) -> Result<SessionTask, ToolExecutionResult> {
114 let registry = require_task_registry(context)?;
115 registry
116 .get(context.session_id, task_id)
117 .await
118 .map_err(ToolExecutionResult::internal_error)?
119 .ok_or_else(|| ToolExecutionResult::tool_error(format!("No task found with id: {task_id}")))
120}
121
122fn compact_task_json(task: &SessionTask) -> Value {
124 json!({
125 "id": task.id,
126 "kind": task.kind,
127 "display_name": task.display_name,
128 "state": task.state,
129 "state_detail": task.state_detail,
130 "progress": task.progress,
131 "summary": task.summary,
132 "created_at": task.created_at.to_rfc3339(),
133 "finished_at": task.finished_at.map(|t| t.to_rfc3339()),
134 })
135}
136
137fn message_json(message: &TaskMessage) -> Value {
138 serde_json::to_value(message).unwrap_or_else(|_| json!({}))
139}
140
141fn full_task_json(task: &SessionTask) -> Value {
142 serde_json::to_value(task).unwrap_or_else(|_| json!({}))
143}
144
145pub struct ListTasksTool;
150
151#[async_trait]
152impl Tool for ListTasksTool {
153 fn narrate(
154 &self,
155 tool_call: &crate::tool_types::ToolCall,
156 phase: crate::tool_narration::ToolNarrationPhase,
157 locale: Option<&str>,
158 _ctx: crate::tool_narration::ToolNarrationContext<'_>,
159 ) -> Option<String> {
160 crate::tool_narration::narrate_session_task(
161 self.name(),
162 &tool_call.arguments,
163 phase,
164 locale,
165 )
166 }
167
168 fn name(&self) -> &str {
169 "list_tasks"
170 }
171
172 fn display_name(&self) -> Option<&str> {
173 Some("List Tasks")
174 }
175
176 fn description(&self) -> &str {
177 "List this session's background tasks (subagents, external agents, background tools) with state, progress, and summary."
178 }
179
180 fn parameters_schema(&self) -> Value {
181 json!({
182 "type": "object",
183 "properties": {
184 "state": {
185 "type": "string",
186 "enum": ["queued", "running", "awaiting_input", "succeeded", "failed", "canceled"],
187 "description": "Filter by lifecycle state."
188 },
189 "kind": {
190 "type": "string",
191 "description": "Filter by task kind (e.g. 'subagent', 'external_agent', 'background_tool')."
192 }
193 },
194 "additionalProperties": false
195 })
196 }
197
198 fn hints(&self) -> ToolHints {
199 ToolHints::default()
200 .with_readonly(true)
201 .with_idempotent(true)
202 }
203
204 async fn execute(&self, _arguments: Value) -> ToolExecutionResult {
205 ToolExecutionResult::tool_error("list_tasks requires session context.")
206 }
207
208 async fn execute_with_context(
209 &self,
210 arguments: Value,
211 context: &ToolContext,
212 ) -> ToolExecutionResult {
213 list_tasks_impl(arguments, context)
214 .await
215 .unwrap_or_else(|e| e)
216 }
217
218 fn requires_context(&self) -> bool {
219 true
220 }
221
222 fn required_context_services(&self) -> &'static [ToolContextService] {
223 &[ToolContextService::SessionTaskRegistry]
224 }
225}
226
227async fn list_tasks_impl(
228 arguments: Value,
229 context: &ToolContext,
230) -> Result<ToolExecutionResult, ToolExecutionResult> {
231 let registry = require_task_registry(context)?;
232 let state = match arguments.get("state").and_then(Value::as_str) {
233 Some(raw) => match SessionTaskState::parse(raw) {
234 Some(state) => Some(state),
235 None => {
236 return Ok(ToolExecutionResult::tool_error(format!(
237 "Unknown state filter \"{raw}\". Valid states: queued, running, \
238 awaiting_input, succeeded, failed, canceled."
239 )));
240 }
241 },
242 None => None,
243 };
244 let filter = SessionTaskFilter {
245 kind: arguments
246 .get("kind")
247 .and_then(Value::as_str)
248 .map(str::trim)
249 .filter(|s| !s.is_empty())
250 .map(ToString::to_string),
251 state,
252 };
253 let tasks = registry
254 .list(context.session_id, Some(&filter))
255 .await
256 .map_err(ToolExecutionResult::internal_error)?;
257 let entries = tasks.iter().map(compact_task_json).collect::<Vec<_>>();
258 Ok(ToolExecutionResult::success(json!({
259 "tasks": entries,
260 "count": entries.len(),
261 })))
262}
263
264pub struct GetTaskTool;
269
270#[async_trait]
271impl Tool for GetTaskTool {
272 fn narrate(
273 &self,
274 tool_call: &crate::tool_types::ToolCall,
275 phase: crate::tool_narration::ToolNarrationPhase,
276 locale: Option<&str>,
277 _ctx: crate::tool_narration::ToolNarrationContext<'_>,
278 ) -> Option<String> {
279 crate::tool_narration::narrate_session_task(
280 self.name(),
281 &tool_call.arguments,
282 phase,
283 locale,
284 )
285 }
286
287 fn name(&self) -> &str {
288 "get_task"
289 }
290
291 fn display_name(&self) -> Option<&str> {
292 Some("Get Task")
293 }
294
295 fn description(&self) -> &str {
296 "Get a task's full snapshot (state, progress, input request, result path, error) plus its recent message thread."
297 }
298
299 fn parameters_schema(&self) -> Value {
300 json!({
301 "type": "object",
302 "properties": {
303 "task_id": {
304 "type": "string",
305 "description": "Task ID (task_*)."
306 }
307 },
308 "required": ["task_id"],
309 "additionalProperties": false
310 })
311 }
312
313 fn hints(&self) -> ToolHints {
314 ToolHints::default()
315 .with_readonly(true)
316 .with_idempotent(true)
317 }
318
319 async fn execute(&self, _arguments: Value) -> ToolExecutionResult {
320 ToolExecutionResult::tool_error("get_task requires session context.")
321 }
322
323 async fn execute_with_context(
324 &self,
325 arguments: Value,
326 context: &ToolContext,
327 ) -> ToolExecutionResult {
328 get_task_impl(arguments, context)
329 .await
330 .unwrap_or_else(|e| e)
331 }
332
333 fn requires_context(&self) -> bool {
334 true
335 }
336
337 fn required_context_services(&self) -> &'static [ToolContextService] {
338 &[ToolContextService::SessionTaskRegistry]
339 }
340}
341
342async fn get_task_impl(
343 arguments: Value,
344 context: &ToolContext,
345) -> Result<ToolExecutionResult, ToolExecutionResult> {
346 let task_id = require_str(&arguments, "task_id")?;
347 let task = load_task(context, task_id).await?;
348 let registry = require_task_registry(context)?;
349 let messages = registry
350 .list_messages(
351 context.session_id,
352 task_id,
353 Some(GET_TASK_MESSAGE_LIMIT),
354 None,
355 )
356 .await
357 .unwrap_or_default();
358 Ok(ToolExecutionResult::success(json!({
359 "task": full_task_json(&task),
360 "messages": messages.iter().map(message_json).collect::<Vec<_>>(),
361 })))
362}
363
364pub struct MessageTaskTool;
369
370#[async_trait]
371impl Tool for MessageTaskTool {
372 fn narrate(
373 &self,
374 tool_call: &crate::tool_types::ToolCall,
375 phase: crate::tool_narration::ToolNarrationPhase,
376 locale: Option<&str>,
377 _ctx: crate::tool_narration::ToolNarrationContext<'_>,
378 ) -> Option<String> {
379 crate::tool_narration::narrate_session_task(
380 self.name(),
381 &tool_call.arguments,
382 phase,
383 locale,
384 )
385 }
386
387 fn name(&self) -> &str {
388 "message_task"
389 }
390
391 fn display_name(&self) -> Option<&str> {
392 Some("Message Task")
393 }
394
395 fn description(&self) -> &str {
396 "Send an inbound message to a task. To answer a pending input request, set in_reply_to to the input request id."
397 }
398
399 fn parameters_schema(&self) -> Value {
400 json!({
401 "type": "object",
402 "properties": {
403 "task_id": {
404 "type": "string",
405 "description": "Task ID (task_*)."
406 },
407 "message": {
408 "type": "string",
409 "description": "Message to deliver to the task."
410 },
411 "in_reply_to": {
412 "type": "string",
413 "description": "ID of the pending input request this message answers."
414 }
415 },
416 "required": ["task_id", "message"],
417 "additionalProperties": false
418 })
419 }
420
421 fn hints(&self) -> ToolHints {
422 ToolHints::default().with_long_running(true)
423 }
424
425 async fn execute(&self, _arguments: Value) -> ToolExecutionResult {
426 ToolExecutionResult::tool_error("message_task requires session context.")
427 }
428
429 async fn execute_with_context(
430 &self,
431 arguments: Value,
432 context: &ToolContext,
433 ) -> ToolExecutionResult {
434 message_task_impl(arguments, context)
435 .await
436 .unwrap_or_else(|e| e)
437 }
438
439 fn requires_context(&self) -> bool {
440 true
441 }
442
443 fn required_context_services(&self) -> &'static [ToolContextService] {
444 &[ToolContextService::SessionTaskRegistry]
445 }
446}
447
448async fn message_task_impl(
449 arguments: Value,
450 context: &ToolContext,
451) -> Result<ToolExecutionResult, ToolExecutionResult> {
452 let task_id = require_str(&arguments, "task_id")?.to_string();
453 let message = require_str(&arguments, "message")?.to_string();
454 let in_reply_to = arguments
455 .get("in_reply_to")
456 .and_then(Value::as_str)
457 .map(str::trim)
458 .filter(|s| !s.is_empty())
459 .map(ToString::to_string);
460
461 let task = load_task(context, &task_id).await?;
462 let registry = require_task_registry(context)?;
463
464 let mut new_message = NewTaskMessage::inbound_text(message);
465 new_message.in_reply_to = in_reply_to;
466 let recorded = registry
467 .record_message(context.session_id, &task_id, new_message)
468 .await
469 .map_err(ToolExecutionResult::internal_error)?;
470
471 let delivery = match find_task_executor(&task.kind) {
474 Some(executor) => {
475 let current = registry
478 .get(context.session_id, &task_id)
479 .await
480 .ok()
481 .flatten()
482 .unwrap_or(task);
483 match executor.deliver(¤t, &recorded, context).await {
484 Ok(()) => "delivered".to_string(),
485 Err(e) => format!("failed: {e}"),
486 }
487 }
488 None => format!(
489 "failed: no executor registered for task kind '{}'",
490 task.kind
491 ),
492 };
493
494 Ok(ToolExecutionResult::success(json!({
495 "task_id": task_id,
496 "message_id": recorded.id,
497 "recorded": true,
498 "delivery": delivery,
499 })))
500}
501
502pub struct CancelTaskTool;
507
508#[async_trait]
509impl Tool for CancelTaskTool {
510 fn narrate(
511 &self,
512 tool_call: &crate::tool_types::ToolCall,
513 phase: crate::tool_narration::ToolNarrationPhase,
514 locale: Option<&str>,
515 _ctx: crate::tool_narration::ToolNarrationContext<'_>,
516 ) -> Option<String> {
517 crate::tool_narration::narrate_session_task(
518 self.name(),
519 &tool_call.arguments,
520 phase,
521 locale,
522 )
523 }
524
525 fn name(&self) -> &str {
526 "cancel_task"
527 }
528
529 fn display_name(&self) -> Option<&str> {
530 Some("Cancel Task")
531 }
532
533 fn description(&self) -> &str {
534 "Request cooperative cancellation of a task. The task winds down and may still end succeeded or failed. For a detached `session` task this also cancels the peer session (not just the tracking chip)."
535 }
536
537 fn parameters_schema(&self) -> Value {
538 json!({
539 "type": "object",
540 "properties": {
541 "task_id": {
542 "type": "string",
543 "description": "Task ID (task_*)."
544 }
545 },
546 "required": ["task_id"],
547 "additionalProperties": false
548 })
549 }
550
551 fn hints(&self) -> ToolHints {
552 ToolHints::default().with_idempotent(true)
553 }
554
555 async fn execute(&self, _arguments: Value) -> ToolExecutionResult {
556 ToolExecutionResult::tool_error("cancel_task requires session context.")
557 }
558
559 async fn execute_with_context(
560 &self,
561 arguments: Value,
562 context: &ToolContext,
563 ) -> ToolExecutionResult {
564 cancel_task_impl(arguments, context)
565 .await
566 .unwrap_or_else(|e| e)
567 }
568
569 fn requires_context(&self) -> bool {
570 true
571 }
572
573 fn required_context_services(&self) -> &'static [ToolContextService] {
574 &[ToolContextService::SessionTaskRegistry]
575 }
576}
577
578async fn cancel_task_impl(
579 arguments: Value,
580 context: &ToolContext,
581) -> Result<ToolExecutionResult, ToolExecutionResult> {
582 let task_id = require_str(&arguments, "task_id")?.to_string();
583 let registry = require_task_registry(context)?;
584 let task = match registry.request_cancel(context.session_id, &task_id).await {
585 Ok(Some(task)) => task,
586 Ok(None) => {
587 return Ok(ToolExecutionResult::tool_error(format!(
588 "No task found with id: {task_id}"
589 )));
590 }
591 Err(e) => return Err(ToolExecutionResult::internal_error(e)),
592 };
593
594 let executor_result = if task.state.is_terminal() {
596 "task already terminal".to_string()
597 } else {
598 match find_task_executor(&task.kind) {
599 Some(executor) => match executor.cancel(&task, context).await {
600 Ok(()) => "cancellation requested".to_string(),
601 Err(e) => format!("failed: {e}"),
602 },
603 None => format!(
604 "no executor registered for task kind '{}'; cancel intent recorded",
605 task.kind
606 ),
607 }
608 };
609
610 Ok(ToolExecutionResult::success(json!({
611 "task_id": task_id,
612 "state": task.state,
613 "cancel_requested": true,
614 "executor": executor_result,
615 })))
616}
617
618pub struct WaitTaskTool;
623
624#[async_trait]
625impl Tool for WaitTaskTool {
626 fn narrate(
627 &self,
628 tool_call: &crate::tool_types::ToolCall,
629 phase: crate::tool_narration::ToolNarrationPhase,
630 locale: Option<&str>,
631 _ctx: crate::tool_narration::ToolNarrationContext<'_>,
632 ) -> Option<String> {
633 crate::tool_narration::narrate_session_task(
634 self.name(),
635 &tool_call.arguments,
636 phase,
637 locale,
638 )
639 }
640
641 fn name(&self) -> &str {
642 "wait_task"
643 }
644
645 fn display_name(&self) -> Option<&str> {
646 Some("Wait Task")
647 }
648
649 fn description(&self) -> &str {
650 "Wait until a task reaches a terminal state or asks for input. Returns the latest task snapshot."
651 }
652
653 fn parameters_schema(&self) -> Value {
654 json!({
655 "type": "object",
656 "properties": {
657 "task_id": {
658 "type": "string",
659 "description": "Task ID (task_*)."
660 },
661 "timeout_seconds": {
662 "type": "integer",
663 "minimum": 1,
664 "maximum": 86400,
665 "default": 300,
666 "description": "Maximum seconds to wait before returning the current snapshot."
667 }
668 },
669 "required": ["task_id"],
670 "additionalProperties": false
671 })
672 }
673
674 fn hints(&self) -> ToolHints {
675 ToolHints::default().with_long_running(true)
676 }
677
678 async fn execute(&self, _arguments: Value) -> ToolExecutionResult {
679 ToolExecutionResult::tool_error("wait_task requires session context.")
680 }
681
682 async fn execute_with_context(
683 &self,
684 arguments: Value,
685 context: &ToolContext,
686 ) -> ToolExecutionResult {
687 wait_task_impl(arguments, context)
688 .await
689 .unwrap_or_else(|e| e)
690 }
691
692 fn requires_context(&self) -> bool {
693 true
694 }
695
696 fn required_context_services(&self) -> &'static [ToolContextService] {
697 &[ToolContextService::SessionTaskRegistry]
698 }
699}
700
701async fn wait_task_impl(
702 arguments: Value,
703 context: &ToolContext,
704) -> Result<ToolExecutionResult, ToolExecutionResult> {
705 let task_id = require_str(&arguments, "task_id")?.to_string();
706 let timeout_secs = arguments
707 .get("timeout_seconds")
708 .and_then(Value::as_u64)
709 .unwrap_or(DEFAULT_WAIT_TIMEOUT_SECS);
710 let deadline = Instant::now() + Duration::from_secs(timeout_secs);
711 let mut polls: u64 = 0;
712
713 loop {
714 let task = load_task(context, &task_id).await?;
715 if task.state.is_terminal() || task.state == SessionTaskState::AwaitingInput {
716 return Ok(ToolExecutionResult::success(json!({
717 "task": full_task_json(&task),
718 "timed_out": false,
719 })));
720 }
721 if Instant::now() >= deadline {
722 return Ok(ToolExecutionResult::success(json!({
723 "task": full_task_json(&task),
724 "timed_out": true,
725 "message": format!("Task {task_id} still {} after {timeout_secs}s", task.state),
726 })));
727 }
728 polls += 1;
729 if polls.is_multiple_of(WAIT_RECONCILE_EVERY)
732 && let Some(executor) = find_task_executor(&task.kind)
733 {
734 let _ = executor.reconcile(&task, context).await;
735 }
736 sleep(WAIT_POLL_INTERVAL).await;
737 }
738}
739
740#[cfg(test)]
745pub(crate) mod tests {
746 use super::*;
747 use crate::session_task::{
748 CreateSessionTask, SessionTaskUpdate, TaskError, TaskExecutor, TaskExecutorPlugin,
749 TaskInputRequest, TaskLinks, TaskMessageDirection, TaskMessagePart, TaskWakePolicy,
750 apply_task_update, generate_task_message_id, new_session_task,
751 };
752 use crate::typed_id::SessionId;
753 use chrono::Utc;
754 use std::collections::HashMap;
755 use std::sync::Mutex;
756
757 #[derive(Default)]
760 pub(crate) struct InMemorySessionTaskRegistry {
761 tasks: Mutex<HashMap<String, SessionTask>>,
762 messages: Mutex<HashMap<String, Vec<TaskMessage>>>,
763 }
764
765 #[async_trait]
766 impl SessionTaskRegistry for InMemorySessionTaskRegistry {
767 async fn create(&self, input: CreateSessionTask) -> crate::error::Result<SessionTask> {
768 let mut tasks = self.tasks.lock().unwrap();
769 if let Some(id) = &input.id
770 && let Some(existing) = tasks.get(id)
771 {
772 return Ok(existing.clone());
773 }
774 let task = new_session_task(input, Utc::now());
775 tasks.insert(task.id.clone(), task.clone());
776 Ok(task)
777 }
778
779 async fn update(
780 &self,
781 _session_id: SessionId,
782 task_id: &str,
783 update: SessionTaskUpdate,
784 ) -> crate::error::Result<Option<SessionTask>> {
785 let mut tasks = self.tasks.lock().unwrap();
786 let Some(task) = tasks.get_mut(task_id) else {
787 return Ok(None);
788 };
789 apply_task_update(task, update, Utc::now());
790 Ok(Some(task.clone()))
791 }
792
793 async fn get(
794 &self,
795 _session_id: SessionId,
796 task_id: &str,
797 ) -> crate::error::Result<Option<SessionTask>> {
798 Ok(self.tasks.lock().unwrap().get(task_id).cloned())
799 }
800
801 async fn list(
802 &self,
803 session_id: SessionId,
804 filter: Option<&SessionTaskFilter>,
805 ) -> crate::error::Result<Vec<SessionTask>> {
806 let tasks = self.tasks.lock().unwrap();
807 Ok(tasks
808 .values()
809 .filter(|task| {
810 task.session_id == session_id
811 && filter.is_none_or(|f| {
812 f.kind.as_deref().is_none_or(|kind| task.kind == kind)
813 && f.state.is_none_or(|state| task.state == state)
814 })
815 })
816 .cloned()
817 .collect())
818 }
819
820 async fn request_cancel(
821 &self,
822 _session_id: SessionId,
823 task_id: &str,
824 ) -> crate::error::Result<Option<SessionTask>> {
825 let mut tasks = self.tasks.lock().unwrap();
826 let Some(task) = tasks.get_mut(task_id) else {
827 return Ok(None);
828 };
829 task.cancel_requested_at.get_or_insert_with(Utc::now);
830 task.updated_at = Utc::now();
831 Ok(Some(task.clone()))
832 }
833
834 async fn record_message(
835 &self,
836 session_id: SessionId,
837 task_id: &str,
838 message: NewTaskMessage,
839 ) -> crate::error::Result<TaskMessage> {
840 let stored = {
841 let tasks = self.tasks.lock().unwrap();
842 let Some(task) = tasks.get(task_id) else {
843 return Err(crate::error::AgentLoopError::tool(format!(
844 "no task {task_id}"
845 )));
846 };
847 task.clone()
848 };
849 if let Some(expected) = message.expected_attempt
851 && expected != stored.attempt
852 {
853 return Err(crate::error::AgentLoopError::store(format!(
854 "Stale attempt {expected} for task {task_id} (current attempt {})",
855 stored.attempt
856 )));
857 }
858 let recorded = TaskMessage {
859 id: generate_task_message_id(),
860 task_id: task_id.to_string(),
861 direction: message.direction,
862 content: message.content,
863 in_reply_to: message.in_reply_to,
864 created_at: Utc::now(),
865 };
866 if let Some(in_reply_to) = &recorded.in_reply_to
869 && stored
870 .input_request
871 .as_ref()
872 .is_some_and(|req| &req.id == in_reply_to)
873 {
874 self.update(
875 session_id,
876 task_id,
877 SessionTaskUpdate {
878 state: Some(SessionTaskState::Running),
879 ..Default::default()
880 },
881 )
882 .await?;
883 }
884 self.messages
885 .lock()
886 .unwrap()
887 .entry(task_id.to_string())
888 .or_default()
889 .push(recorded.clone());
890 Ok(recorded)
891 }
892
893 async fn list_messages(
894 &self,
895 _session_id: SessionId,
896 task_id: &str,
897 limit: Option<u32>,
898 after_id: Option<&str>,
899 ) -> crate::error::Result<Vec<TaskMessage>> {
900 let messages = self.messages.lock().unwrap();
901 let all = messages.get(task_id).cloned().unwrap_or_default();
902 let mut iter: Box<dyn Iterator<Item = TaskMessage>> = if let Some(cursor) = after_id {
903 Box::new(all.into_iter().skip_while(move |m| m.id != cursor).skip(1))
904 } else {
905 Box::new(all.into_iter())
906 };
907 let collected: Vec<_> = iter.by_ref().collect();
908 if let Some(limit) = limit {
909 if after_id.is_some() {
910 return Ok(collected.into_iter().take(limit as usize).collect());
911 }
912 let skip = collected.len().saturating_sub(limit as usize);
913 return Ok(collected.into_iter().skip(skip).collect());
914 }
915 Ok(collected)
916 }
917 }
918
919 const TEST_EXECUTOR_KIND: &str = "session_tasks_test";
924 static TEST_DELIVERED: Mutex<Vec<String>> = Mutex::new(Vec::new());
925 static TEST_CANCELED: Mutex<Vec<String>> = Mutex::new(Vec::new());
926
927 fn executor_invocations(log: &Mutex<Vec<String>>, task_id: &str) -> usize {
928 log.lock()
929 .unwrap()
930 .iter()
931 .filter(|id| *id == task_id)
932 .count()
933 }
934
935 struct TestTaskExecutor;
936
937 #[async_trait]
938 impl TaskExecutor for TestTaskExecutor {
939 fn kind(&self) -> &str {
940 TEST_EXECUTOR_KIND
941 }
942
943 async fn deliver(
944 &self,
945 task: &SessionTask,
946 _message: &TaskMessage,
947 _context: &ToolContext,
948 ) -> crate::error::Result<()> {
949 TEST_DELIVERED.lock().unwrap().push(task.id.clone());
950 Ok(())
951 }
952
953 async fn cancel(
954 &self,
955 task: &SessionTask,
956 _context: &ToolContext,
957 ) -> crate::error::Result<()> {
958 TEST_CANCELED.lock().unwrap().push(task.id.clone());
959 Ok(())
960 }
961 }
962
963 inventory::submit! {
964 TaskExecutorPlugin {
965 executor: || Arc::new(TestTaskExecutor),
966 }
967 }
968
969 fn test_context(registry: Arc<InMemorySessionTaskRegistry>) -> ToolContext {
970 ToolContext::new(SessionId::new()).with_session_task_registry(registry)
971 }
972
973 async fn create_task(
974 registry: &InMemorySessionTaskRegistry,
975 context: &ToolContext,
976 kind: &str,
977 state: SessionTaskState,
978 ) -> SessionTask {
979 registry
980 .create(CreateSessionTask {
981 session_id: context.session_id,
982 id: None,
983 kind: kind.to_string(),
984 display_name: "Test Task".to_string(),
985 spec: json!({}),
986 state,
987 links: TaskLinks::default(),
988 wake_policy: TaskWakePolicy::Silent,
989 })
990 .await
991 .unwrap()
992 }
993
994 #[tokio::test]
997 async fn tools_error_without_registry() {
998 let context = ToolContext::new(SessionId::new());
999 let result = ListTasksTool
1000 .execute_with_context(json!({}), &context)
1001 .await;
1002 assert!(matches!(result, ToolExecutionResult::ToolError(_)));
1003 let result = WaitTaskTool
1004 .execute_with_context(json!({"task_id": "task_x"}), &context)
1005 .await;
1006 assert!(matches!(result, ToolExecutionResult::ToolError(_)));
1007 }
1008
1009 #[tokio::test]
1010 async fn list_tasks_returns_compact_entries_with_filters() {
1011 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1012 let context = test_context(registry.clone());
1013 let running = create_task(
1014 ®istry,
1015 &context,
1016 TEST_EXECUTOR_KIND,
1017 SessionTaskState::Running,
1018 )
1019 .await;
1020 create_task(®istry, &context, "other_kind", SessionTaskState::Queued).await;
1021
1022 let result = ListTasksTool
1023 .execute_with_context(json!({}), &context)
1024 .await;
1025 let ToolExecutionResult::Success(value) = result else {
1026 panic!("expected success: {result:?}");
1027 };
1028 assert_eq!(value["count"], 2);
1029
1030 let result = ListTasksTool
1031 .execute_with_context(
1032 json!({"kind": TEST_EXECUTOR_KIND, "state": "running"}),
1033 &context,
1034 )
1035 .await;
1036 let ToolExecutionResult::Success(value) = result else {
1037 panic!("expected success: {result:?}");
1038 };
1039 assert_eq!(value["count"], 1);
1040 let entry = &value["tasks"][0];
1041 assert_eq!(entry["id"], running.id);
1042 assert_eq!(entry["state"], "running");
1043 assert!(entry.get("display_name").is_some());
1044 assert!(entry.get("created_at").is_some());
1045 assert!(entry.get("spec").is_none());
1047 }
1048
1049 #[tokio::test]
1050 async fn get_task_returns_snapshot_and_recent_messages() {
1051 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1052 let context = test_context(registry.clone());
1053 let task = create_task(
1054 ®istry,
1055 &context,
1056 TEST_EXECUTOR_KIND,
1057 SessionTaskState::Running,
1058 )
1059 .await;
1060 for i in 0..25 {
1061 registry
1062 .record_message(
1063 context.session_id,
1064 &task.id,
1065 NewTaskMessage::outbound_text(format!("update {i}")),
1066 )
1067 .await
1068 .unwrap();
1069 }
1070
1071 let result = GetTaskTool
1072 .execute_with_context(json!({"task_id": task.id}), &context)
1073 .await;
1074 let ToolExecutionResult::Success(value) = result else {
1075 panic!("expected success: {result:?}");
1076 };
1077 assert_eq!(value["task"]["id"], task.id);
1078 let messages = value["messages"].as_array().unwrap();
1079 assert_eq!(messages.len(), GET_TASK_MESSAGE_LIMIT as usize);
1080 assert_eq!(messages[0]["content"][0]["text"], "update 5");
1082 assert_eq!(messages[19]["content"][0]["text"], "update 24");
1083 }
1084
1085 #[tokio::test]
1086 async fn get_task_unknown_id_errors() {
1087 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1088 let context = test_context(registry);
1089 let result = GetTaskTool
1090 .execute_with_context(json!({"task_id": "task_missing"}), &context)
1091 .await;
1092 assert!(matches!(result, ToolExecutionResult::ToolError(_)));
1093 }
1094
1095 #[tokio::test]
1096 async fn message_task_records_and_delivers() {
1097 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1098 let context = test_context(registry.clone());
1099 let task = create_task(
1100 ®istry,
1101 &context,
1102 TEST_EXECUTOR_KIND,
1103 SessionTaskState::Running,
1104 )
1105 .await;
1106
1107 let result = MessageTaskTool
1108 .execute_with_context(
1109 json!({"task_id": task.id, "message": "keep going"}),
1110 &context,
1111 )
1112 .await;
1113 let ToolExecutionResult::Success(value) = result else {
1114 panic!("expected success: {result:?}");
1115 };
1116 assert_eq!(value["delivery"], "delivered");
1117 assert_eq!(executor_invocations(&TEST_DELIVERED, &task.id), 1);
1118
1119 let messages = registry
1120 .list_messages(context.session_id, &task.id, None, None)
1121 .await
1122 .unwrap();
1123 assert_eq!(messages.len(), 1);
1124 assert_eq!(messages[0].direction, TaskMessageDirection::Inbound);
1125 assert_eq!(
1126 messages[0].content,
1127 vec![TaskMessagePart::text("keep going")]
1128 );
1129 }
1130
1131 #[tokio::test]
1132 async fn message_task_without_executor_still_records() {
1133 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1134 let context = test_context(registry.clone());
1135 let task = create_task(
1136 ®istry,
1137 &context,
1138 "kind_without_executor",
1139 SessionTaskState::Running,
1140 )
1141 .await;
1142
1143 let result = MessageTaskTool
1144 .execute_with_context(json!({"task_id": task.id, "message": "hello"}), &context)
1145 .await;
1146 let ToolExecutionResult::Success(value) = result else {
1147 panic!("expected success: {result:?}");
1148 };
1149 assert_eq!(value["recorded"], true);
1150 let delivery = value["delivery"].as_str().unwrap();
1151 assert!(delivery.starts_with("failed:"), "delivery: {delivery}");
1152 let messages = registry
1153 .list_messages(context.session_id, &task.id, None, None)
1154 .await
1155 .unwrap();
1156 assert_eq!(messages.len(), 1);
1157 }
1158
1159 #[tokio::test]
1160 async fn message_task_in_reply_to_resumes_awaiting_input() {
1161 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1162 let context = test_context(registry.clone());
1163 let task = create_task(
1164 ®istry,
1165 &context,
1166 TEST_EXECUTOR_KIND,
1167 SessionTaskState::Running,
1168 )
1169 .await;
1170 registry
1171 .update(
1172 context.session_id,
1173 &task.id,
1174 SessionTaskUpdate {
1175 input_request: Some(TaskInputRequest {
1176 id: "req_1".to_string(),
1177 prompt: "Approve?".to_string(),
1178 expected: None,
1179 }),
1180 ..Default::default()
1181 },
1182 )
1183 .await
1184 .unwrap();
1185
1186 let result = MessageTaskTool
1187 .execute_with_context(
1188 json!({"task_id": task.id, "message": "yes", "in_reply_to": "req_1"}),
1189 &context,
1190 )
1191 .await;
1192 assert!(matches!(result, ToolExecutionResult::Success(_)));
1193
1194 let current = registry
1195 .get(context.session_id, &task.id)
1196 .await
1197 .unwrap()
1198 .unwrap();
1199 assert_eq!(current.state, SessionTaskState::Running);
1200 assert!(current.input_request.is_none());
1201 }
1202
1203 #[tokio::test]
1204 async fn cancel_task_records_intent_and_calls_executor() {
1205 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1206 let context = test_context(registry.clone());
1207 let task = create_task(
1208 ®istry,
1209 &context,
1210 TEST_EXECUTOR_KIND,
1211 SessionTaskState::Running,
1212 )
1213 .await;
1214
1215 let result = CancelTaskTool
1216 .execute_with_context(json!({"task_id": task.id}), &context)
1217 .await;
1218 let ToolExecutionResult::Success(value) = result else {
1219 panic!("expected success: {result:?}");
1220 };
1221 assert_eq!(value["cancel_requested"], true);
1222 assert_eq!(value["executor"], "cancellation requested");
1223 assert_eq!(executor_invocations(&TEST_CANCELED, &task.id), 1);
1224
1225 let current = registry
1226 .get(context.session_id, &task.id)
1227 .await
1228 .unwrap()
1229 .unwrap();
1230 assert!(current.cancel_requested_at.is_some());
1231 }
1232
1233 #[tokio::test]
1234 async fn cancel_task_on_terminal_task_skips_executor() {
1235 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1236 let context = test_context(registry.clone());
1237 let task = create_task(
1238 ®istry,
1239 &context,
1240 TEST_EXECUTOR_KIND,
1241 SessionTaskState::Succeeded,
1242 )
1243 .await;
1244
1245 let result = CancelTaskTool
1246 .execute_with_context(json!({"task_id": task.id}), &context)
1247 .await;
1248 let ToolExecutionResult::Success(value) = result else {
1249 panic!("expected success: {result:?}");
1250 };
1251 assert_eq!(value["executor"], "task already terminal");
1252 assert_eq!(executor_invocations(&TEST_CANCELED, &task.id), 0);
1253 }
1254
1255 #[tokio::test]
1256 async fn wait_task_returns_immediately_when_terminal() {
1257 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1258 let context = test_context(registry.clone());
1259 let task = create_task(
1260 ®istry,
1261 &context,
1262 TEST_EXECUTOR_KIND,
1263 SessionTaskState::Running,
1264 )
1265 .await;
1266 registry
1267 .update(
1268 context.session_id,
1269 &task.id,
1270 SessionTaskUpdate {
1271 state: Some(SessionTaskState::Failed),
1272 error: Some(TaskError {
1273 kind: "error".to_string(),
1274 message: "boom".to_string(),
1275 }),
1276 ..Default::default()
1277 },
1278 )
1279 .await
1280 .unwrap();
1281
1282 let result = WaitTaskTool
1283 .execute_with_context(json!({"task_id": task.id}), &context)
1284 .await;
1285 let ToolExecutionResult::Success(value) = result else {
1286 panic!("expected success: {result:?}");
1287 };
1288 assert_eq!(value["timed_out"], false);
1289 assert_eq!(value["task"]["state"], "failed");
1290 assert_eq!(value["task"]["error"]["message"], "boom");
1291 }
1292
1293 #[tokio::test]
1294 async fn wait_task_returns_when_awaiting_input() {
1295 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1296 let context = test_context(registry.clone());
1297 let task = create_task(
1298 ®istry,
1299 &context,
1300 TEST_EXECUTOR_KIND,
1301 SessionTaskState::AwaitingInput,
1302 )
1303 .await;
1304
1305 let result = WaitTaskTool
1306 .execute_with_context(json!({"task_id": task.id}), &context)
1307 .await;
1308 let ToolExecutionResult::Success(value) = result else {
1309 panic!("expected success: {result:?}");
1310 };
1311 assert_eq!(value["timed_out"], false);
1312 assert_eq!(value["task"]["state"], "awaiting_input");
1313 }
1314
1315 #[tokio::test(start_paused = true)]
1316 async fn wait_task_times_out_with_snapshot() {
1317 let registry = Arc::new(InMemorySessionTaskRegistry::default());
1318 let context = test_context(registry.clone());
1319 let task = create_task(
1320 ®istry,
1321 &context,
1322 TEST_EXECUTOR_KIND,
1323 SessionTaskState::Running,
1324 )
1325 .await;
1326
1327 let result = WaitTaskTool
1328 .execute_with_context(json!({"task_id": task.id, "timeout_seconds": 3}), &context)
1329 .await;
1330 let ToolExecutionResult::Success(value) = result else {
1331 panic!("expected success: {result:?}");
1332 };
1333 assert_eq!(value["timed_out"], true);
1334 assert_eq!(value["task"]["state"], "running");
1335 }
1336}