1use std::collections::{BTreeMap, VecDeque};
2use std::sync::Arc;
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::time::Duration;
5
6use agentkit_core::{
7 Item, MetadataMap, TaskId, ToolCallId, ToolResultPart, TurnCancellation, TurnId,
8};
9use agentkit_tools_core::{
10 ApprovalRequest, OwnedToolContext, ToolError, ToolExecutionOutcome, ToolExecutor, ToolRequest,
11};
12use async_trait::async_trait;
13use thiserror::Error;
14use tokio::sync::{Mutex, Notify, mpsc, oneshot};
15use tokio::task::JoinHandle;
16
17pub const TOOL_RESULT_FAILURE_KIND_METADATA_KEY: &str = "agentkit.tool.failure_kind";
18pub const TOOL_RESULT_FAILURE_KIND_PERMISSION_DENIED: &str = "permission_denied";
19pub const TOOL_RESULT_NOT_STARTED_METADATA_KEY: &str = "agentkit.tool.not_started";
24
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26pub enum TaskKind {
27 Foreground,
28 Background,
29}
30
31#[derive(Clone, Copy, Debug, PartialEq, Eq)]
32pub enum ContinuePolicy {
33 NotifyOnly,
34 RequestContinue,
35}
36
37#[derive(Clone, Copy, Debug, PartialEq, Eq)]
38pub enum DeliveryMode {
39 ToLoop,
40 Manual,
41}
42
43#[derive(Clone, Debug, PartialEq, Eq)]
44pub struct TaskSnapshot {
45 pub id: TaskId,
46 pub turn_id: TurnId,
47 pub call_id: ToolCallId,
48 pub tool_name: String,
49 pub kind: TaskKind,
50 pub metadata: MetadataMap,
51}
52
53#[derive(Clone, Debug, PartialEq)]
54pub enum TaskEvent {
55 Started(TaskSnapshot),
56 Detached(TaskSnapshot),
57 Completed(TaskSnapshot, ToolResultPart),
58 Cancelled(TaskSnapshot),
59 Failed(TaskSnapshot, ToolError),
60 ContinueRequested,
61}
62
63#[derive(Clone, Debug, PartialEq)]
64pub struct TaskApproval {
65 pub task_id: TaskId,
66 pub tool_request: ToolRequest,
67 pub approval: ApprovalRequest,
68}
69
70#[derive(Clone, Debug, PartialEq)]
71pub enum TaskResolution {
72 Item(Item),
73 Approval(TaskApproval),
74}
75
76#[derive(Clone, Debug, PartialEq)]
77pub enum TaskStartOutcome {
78 Ready(Box<TaskResolution>),
79 Pending { task_id: TaskId, kind: TaskKind },
80}
81
82#[derive(Clone, Debug, PartialEq)]
83pub enum TurnTaskUpdate {
84 Resolution(Box<TaskResolution>),
85 Detached(TaskSnapshot),
86}
87
88#[derive(Clone, Debug, Default, PartialEq)]
89pub struct PendingLoopUpdates {
90 pub resolutions: VecDeque<TaskResolution>,
91}
92
93#[derive(Clone, Debug, Default)]
96pub enum TaskLaunchKind {
97 #[default]
99 Plain,
100 Approved(ApprovalRequest),
102}
103
104#[derive(Clone, Debug)]
105pub struct TaskLaunchRequest {
106 pub task_id: Option<TaskId>,
107 pub request: ToolRequest,
108 pub kind: TaskLaunchKind,
109}
110
111impl TaskLaunchRequest {
112 pub fn plain(task_id: Option<TaskId>, request: ToolRequest) -> Self {
114 Self {
115 task_id,
116 request,
117 kind: TaskLaunchKind::Plain,
118 }
119 }
120
121 pub fn approved(
123 task_id: Option<TaskId>,
124 request: ToolRequest,
125 approval: ApprovalRequest,
126 ) -> Self {
127 Self {
128 task_id,
129 request,
130 kind: TaskLaunchKind::Approved(approval),
131 }
132 }
133}
134
135#[derive(Clone)]
136pub struct TaskStartContext {
137 pub executor: Arc<dyn ToolExecutor>,
138 pub tool_context: OwnedToolContext,
139}
140
141#[derive(Debug, Error, Clone, PartialEq, Eq)]
142pub enum TaskManagerError {
143 #[error("task not found: {0}")]
144 NotFound(TaskId),
145 #[error("task is not running: {0}")]
146 NotRunning(TaskId),
147 #[error("task is already running in the background: {0}")]
148 AlreadyBackground(TaskId),
149 #[error("task manager internal error: {0}")]
150 Internal(String),
151}
152
153pub trait TaskRoutingPolicy: Send + Sync {
154 fn route(&self, request: &ToolRequest) -> RoutingDecision;
155}
156
157impl<F> TaskRoutingPolicy for F
158where
159 F: Fn(&ToolRequest) -> RoutingDecision + Send + Sync,
160{
161 fn route(&self, request: &ToolRequest) -> RoutingDecision {
162 self(request)
163 }
164}
165
166#[derive(Clone, Copy, Debug, PartialEq, Eq)]
167pub enum RoutingDecision {
168 Foreground,
169 Background,
170 ForegroundThenDetachAfter(Duration),
171}
172
173struct DefaultRoutingPolicy;
174
175impl TaskRoutingPolicy for DefaultRoutingPolicy {
176 fn route(&self, _request: &ToolRequest) -> RoutingDecision {
177 RoutingDecision::Foreground
178 }
179}
180
181#[async_trait]
182pub trait TaskManager: Send + Sync {
183 async fn start_task(
184 &self,
185 request: TaskLaunchRequest,
186 ctx: TaskStartContext,
187 ) -> Result<TaskStartOutcome, TaskManagerError>;
188
189 async fn wait_for_turn(
190 &self,
191 turn_id: &TurnId,
192 cancellation: Option<TurnCancellation>,
193 ) -> Result<Option<TurnTaskUpdate>, TaskManagerError>;
194
195 async fn take_pending_loop_updates(&self) -> Result<PendingLoopUpdates, TaskManagerError>;
196
197 async fn wait_for_loop_update(&self) -> Result<(), TaskManagerError> {
202 std::future::pending().await
203 }
204
205 async fn on_turn_interrupted(&self, turn_id: &TurnId) -> Result<(), TaskManagerError>;
206
207 fn handle(&self) -> TaskManagerHandle;
208}
209
210#[async_trait]
211trait TaskManagerControl: Send + Sync {
212 async fn next_event(&self) -> Option<TaskEvent>;
213 async fn cancel(&self, task_id: TaskId) -> Result<(), TaskManagerError>;
214 async fn detach(&self, task_id: TaskId) -> Result<(), TaskManagerError>;
215 async fn list_running(&self) -> Vec<TaskSnapshot>;
216 async fn list_completed(&self) -> Vec<TaskSnapshot>;
217 async fn drain_ready_items(&self) -> Vec<Item>;
218 async fn set_continue_policy(
219 &self,
220 task_id: TaskId,
221 policy: ContinuePolicy,
222 ) -> Result<(), TaskManagerError>;
223 async fn set_delivery_mode(
224 &self,
225 task_id: TaskId,
226 mode: DeliveryMode,
227 ) -> Result<(), TaskManagerError>;
228 async fn wait_for_idle(&self);
229}
230
231#[derive(Clone)]
232pub struct TaskManagerHandle {
233 inner: Arc<dyn TaskManagerControl>,
234}
235
236impl TaskManagerHandle {
237 pub async fn next_event(&self) -> Option<TaskEvent> {
238 self.inner.next_event().await
239 }
240
241 pub async fn cancel(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
242 self.inner.cancel(task_id).await
243 }
244
245 pub async fn detach(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
247 self.inner.detach(task_id).await
248 }
249
250 pub async fn list_running(&self) -> Vec<TaskSnapshot> {
251 self.inner.list_running().await
252 }
253
254 pub async fn list_completed(&self) -> Vec<TaskSnapshot> {
255 self.inner.list_completed().await
256 }
257
258 pub async fn drain_ready_items(&self) -> Vec<Item> {
259 self.inner.drain_ready_items().await
260 }
261
262 pub async fn set_continue_policy(
263 &self,
264 task_id: TaskId,
265 policy: ContinuePolicy,
266 ) -> Result<(), TaskManagerError> {
267 self.inner.set_continue_policy(task_id, policy).await
268 }
269
270 pub async fn set_delivery_mode(
271 &self,
272 task_id: TaskId,
273 mode: DeliveryMode,
274 ) -> Result<(), TaskManagerError> {
275 self.inner.set_delivery_mode(task_id, mode).await
276 }
277
278 pub async fn wait_for_idle(&self) {
280 self.inner.wait_for_idle().await
281 }
282}
283
284pub struct SimpleTaskManager {
285 state: Arc<HandleState>,
286}
287
288impl SimpleTaskManager {
289 pub fn new() -> Self {
290 Self {
291 state: Arc::new(HandleState::default()),
292 }
293 }
294}
295
296impl Default for SimpleTaskManager {
297 fn default() -> Self {
298 Self::new()
299 }
300}
301
302#[async_trait]
303impl TaskManager for SimpleTaskManager {
304 async fn start_task(
305 &self,
306 request: TaskLaunchRequest,
307 ctx: TaskStartContext,
308 ) -> Result<TaskStartOutcome, TaskManagerError> {
309 let task_id = request
310 .task_id
311 .clone()
312 .unwrap_or_else(|| self.state.next_task_id());
313 let outcome = match &request.kind {
314 TaskLaunchKind::Approved(approved) => {
315 ctx.executor
316 .execute_approved_owned(request.request.clone(), approved, ctx.tool_context)
317 .await
318 }
319 TaskLaunchKind::Plain => {
320 ctx.executor
321 .execute_owned(request.request.clone(), ctx.tool_context)
322 .await
323 }
324 };
325 Ok(TaskStartOutcome::Ready(Box::new(
326 map_outcome_to_resolution(Some(task_id), request.request, outcome),
327 )))
328 }
329
330 async fn wait_for_turn(
331 &self,
332 _turn_id: &TurnId,
333 _cancellation: Option<TurnCancellation>,
334 ) -> Result<Option<TurnTaskUpdate>, TaskManagerError> {
335 Ok(None)
336 }
337
338 async fn take_pending_loop_updates(&self) -> Result<PendingLoopUpdates, TaskManagerError> {
339 Ok(PendingLoopUpdates::default())
340 }
341
342 async fn on_turn_interrupted(&self, _turn_id: &TurnId) -> Result<(), TaskManagerError> {
343 Ok(())
344 }
345
346 fn handle(&self) -> TaskManagerHandle {
347 TaskManagerHandle {
348 inner: self.state.clone(),
349 }
350 }
351}
352
353#[derive(Default)]
354struct HandleState {
355 next_task_index: AtomicU64,
356 events_rx: Mutex<Option<mpsc::UnboundedReceiver<TaskEvent>>>,
357}
358
359impl HandleState {
360 fn next_task_id(&self) -> TaskId {
361 let next = self.next_task_index.fetch_add(1, Ordering::SeqCst) + 1;
362 TaskId::new(format!("task-{}", next))
363 }
364}
365
366#[async_trait]
367impl TaskManagerControl for HandleState {
368 async fn next_event(&self) -> Option<TaskEvent> {
369 let mut rx = self.events_rx.lock().await;
370 match rx.as_mut() {
371 Some(inner) => inner.recv().await,
372 None => None,
373 }
374 }
375
376 async fn cancel(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
377 Err(TaskManagerError::NotFound(task_id))
378 }
379
380 async fn detach(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
381 Err(TaskManagerError::NotFound(task_id))
382 }
383
384 async fn list_running(&self) -> Vec<TaskSnapshot> {
385 Vec::new()
386 }
387
388 async fn list_completed(&self) -> Vec<TaskSnapshot> {
389 Vec::new()
390 }
391
392 async fn drain_ready_items(&self) -> Vec<Item> {
393 Vec::new()
394 }
395
396 async fn set_continue_policy(
397 &self,
398 task_id: TaskId,
399 _policy: ContinuePolicy,
400 ) -> Result<(), TaskManagerError> {
401 Err(TaskManagerError::NotFound(task_id))
402 }
403
404 async fn set_delivery_mode(
405 &self,
406 task_id: TaskId,
407 _mode: DeliveryMode,
408 ) -> Result<(), TaskManagerError> {
409 Err(TaskManagerError::NotFound(task_id))
410 }
411
412 async fn wait_for_idle(&self) {}
413}
414
415pub struct AsyncTaskManager {
416 inner: Arc<AsyncInner>,
417 routing: Arc<dyn TaskRoutingPolicy>,
418}
419
420impl AsyncTaskManager {
421 pub fn new() -> Self {
422 let (event_tx, event_rx) = mpsc::unbounded_channel();
423 Self {
424 inner: Arc::new(AsyncInner {
425 state: Mutex::new(AsyncState::default()),
426 host_event_tx: event_tx,
427 host_event_rx: Mutex::new(event_rx),
428 notify: Notify::new(),
429 }),
430 routing: Arc::new(DefaultRoutingPolicy),
431 }
432 }
433
434 pub fn routing(mut self, policy: impl TaskRoutingPolicy + 'static) -> Self {
435 self.routing = Arc::new(policy);
436 self
437 }
438}
439
440impl Default for AsyncTaskManager {
441 fn default() -> Self {
442 Self::new()
443 }
444}
445
446#[derive(Default)]
447struct AsyncState {
448 next_task_index: u64,
449 tasks: BTreeMap<TaskId, TaskRecord>,
450 per_turn_running: BTreeMap<TurnId, usize>,
451 per_turn_updates: BTreeMap<TurnId, VecDeque<TurnTaskUpdate>>,
452 pending_loop_updates: VecDeque<TaskResolution>,
453 manual_ready_items: Vec<Item>,
454}
455
456struct TaskRecord {
457 snapshot: TaskSnapshot,
458 continue_policy: ContinuePolicy,
459 delivery_mode: DeliveryMode,
460 running: bool,
461 completed: bool,
462 join: Option<JoinHandle<()>>,
463}
464
465struct AsyncInner {
466 state: Mutex<AsyncState>,
467 host_event_tx: mpsc::UnboundedSender<TaskEvent>,
468 host_event_rx: Mutex<mpsc::UnboundedReceiver<TaskEvent>>,
469 notify: Notify,
470}
471
472impl AsyncInner {
473 async fn next_task_id(&self) -> TaskId {
474 let mut state = self.state.lock().await;
475 state.next_task_index += 1;
476 TaskId::new(format!("task-{}", state.next_task_index))
477 }
478
479 async fn detach_running_foreground(&self, task_id: &TaskId) -> Result<(), TaskManagerError> {
480 let mut state = self.state.lock().await;
481 let snapshot = {
482 let record = state
483 .tasks
484 .get_mut(task_id)
485 .ok_or_else(|| TaskManagerError::NotFound(task_id.clone()))?;
486 if !record.running {
487 return Err(TaskManagerError::NotRunning(task_id.clone()));
488 }
489 if record.snapshot.kind == TaskKind::Background {
490 return Err(TaskManagerError::AlreadyBackground(task_id.clone()));
491 }
492 record.snapshot.kind = TaskKind::Background;
493 record.snapshot.clone()
494 };
495
496 if let Some(count) = state.per_turn_running.get_mut(&snapshot.turn_id) {
497 *count = count.saturating_sub(1);
498 if *count == 0 {
499 state.per_turn_running.remove(&snapshot.turn_id);
500 }
501 }
502 state
503 .per_turn_updates
504 .entry(snapshot.turn_id.clone())
505 .or_default()
506 .push_back(TurnTaskUpdate::Detached(snapshot.clone()));
507 let _ = self.host_event_tx.send(TaskEvent::Detached(snapshot));
508 self.notify.notify_waiters();
509 Ok(())
510 }
511
512 async fn interrupt_turn(&self, turn_id: &TurnId) {
513 let mut state = self.state.lock().await;
514 let interrupted: Vec<TaskId> = state
515 .tasks
516 .iter()
517 .filter_map(|(id, record)| {
518 (record.snapshot.turn_id == *turn_id
519 && record.snapshot.kind == TaskKind::Foreground
520 && record.running)
521 .then_some(id.clone())
522 })
523 .collect();
524 for task_id in interrupted {
525 if let Some(record) = state.tasks.get_mut(&task_id) {
526 record.running = false;
527 if let Some(join) = record.join.take() {
528 join.abort();
529 }
530 let snapshot = record.snapshot.clone();
531 let _ = self.host_event_tx.send(TaskEvent::Cancelled(snapshot));
532 }
533 }
534 state.per_turn_running.remove(turn_id);
535 self.notify.notify_waiters();
536 }
537}
538
539#[async_trait]
540impl TaskManager for AsyncTaskManager {
541 async fn start_task(
542 &self,
543 request: TaskLaunchRequest,
544 ctx: TaskStartContext,
545 ) -> Result<TaskStartOutcome, TaskManagerError> {
546 let route = self.routing.route(&request.request);
547 let task_id = match request.task_id.clone() {
548 Some(existing) => existing,
549 None => self.inner.next_task_id().await,
550 };
551 let initial_kind = match route {
552 RoutingDecision::Background => TaskKind::Background,
553 _ => TaskKind::Foreground,
554 };
555 let snapshot = TaskSnapshot {
556 id: task_id.clone(),
557 turn_id: request.request.turn_id.clone(),
558 call_id: request.request.call_id.clone(),
559 tool_name: request.request.tool_name.to_string(),
560 kind: initial_kind,
561 metadata: request.request.metadata.clone(),
562 };
563 let mut state = self.inner.state.lock().await;
564 state.tasks.insert(
565 task_id.clone(),
566 TaskRecord {
567 snapshot: snapshot.clone(),
568 continue_policy: ContinuePolicy::NotifyOnly,
569 delivery_mode: DeliveryMode::ToLoop,
570 running: true,
571 completed: false,
572 join: None,
573 },
574 );
575 if initial_kind == TaskKind::Foreground {
576 *state
577 .per_turn_running
578 .entry(snapshot.turn_id.clone())
579 .or_default() += 1;
580 }
581 drop(state);
582 let _ = self
583 .inner
584 .host_event_tx
585 .send(TaskEvent::Started(snapshot.clone()));
586
587 let event_tx = self.inner.host_event_tx.clone();
588 let inner = self.inner.clone();
589 let task_id_for_future = task_id.clone();
590 let turn_id = snapshot.turn_id.clone();
591 let kind = request.kind.clone();
592 let exec_request = request.request.clone();
593 let owned_ctx = ctx.tool_context.clone();
594 let executor = ctx.executor.clone();
595 let route_copy = route;
596 let (start_tx, start_rx) = oneshot::channel();
597 let join = tokio::spawn(async move {
598 if start_rx.await.is_err() {
599 return;
600 }
601 if let RoutingDecision::ForegroundThenDetachAfter(duration) = route_copy {
602 let inner = inner.clone();
603 let task_id = task_id_for_future.clone();
604 tokio::spawn(async move {
605 tokio::time::sleep(duration).await;
606 let _ = inner.detach_running_foreground(&task_id).await;
607 });
608 }
609
610 let outcome = match &kind {
611 TaskLaunchKind::Approved(approval) => {
612 executor
613 .execute_approved_owned(exec_request.clone(), approval, owned_ctx)
614 .await
615 }
616 TaskLaunchKind::Plain => {
617 executor
618 .execute_owned(exec_request.clone(), owned_ctx)
619 .await
620 }
621 };
622
623 let resolution =
624 map_outcome_to_resolution(Some(task_id_for_future.clone()), exec_request, outcome);
625 let completed_result = match &resolution {
626 TaskResolution::Item(item) => item.parts.iter().find_map(|part| match part {
627 agentkit_core::Part::ToolResult(result) => Some(result.clone()),
628 _ => None,
629 }),
630 TaskResolution::Approval(_) => None,
631 };
632
633 let (snapshot, should_request_continue) = {
634 let mut state = inner.state.lock().await;
635 let Some(record) = state.tasks.get_mut(&task_id_for_future) else {
636 return;
637 };
638 record.running = false;
639 record.completed = true;
640 let snapshot = record.snapshot.clone();
641 let continue_policy = record.continue_policy;
642 let delivery_mode = record.delivery_mode;
643 let current_kind = snapshot.kind;
644
645 if current_kind == TaskKind::Foreground {
646 if let Some(count) = state.per_turn_running.get_mut(&turn_id) {
647 *count = count.saturating_sub(1);
648 if *count == 0 {
649 state.per_turn_running.remove(&turn_id);
650 }
651 }
652 state
653 .per_turn_updates
654 .entry(turn_id.clone())
655 .or_default()
656 .push_back(TurnTaskUpdate::Resolution(Box::new(resolution.clone())));
657 } else {
658 match &resolution {
659 TaskResolution::Item(_) if delivery_mode == DeliveryMode::ToLoop => {
660 state.pending_loop_updates.push_back(resolution.clone());
661 }
662 TaskResolution::Approval(_) if delivery_mode == DeliveryMode::ToLoop => {
663 state.pending_loop_updates.push_back(resolution.clone());
664 }
665 TaskResolution::Item(item) => {
666 state.manual_ready_items.push(item.clone());
667 }
668 TaskResolution::Approval(_) => {}
669 }
670 }
671
672 (
673 snapshot,
674 current_kind == TaskKind::Background
675 && delivery_mode == DeliveryMode::ToLoop
676 && continue_policy == ContinuePolicy::RequestContinue,
677 )
678 };
679
680 if let Some(result) = completed_result {
681 let _ = event_tx.send(TaskEvent::Completed(snapshot.clone(), result));
682 }
683 if should_request_continue {
684 let _ = event_tx.send(TaskEvent::ContinueRequested);
685 }
686 inner.notify.notify_waiters();
687 });
688
689 let mut state = self.inner.state.lock().await;
690 let mut join = Some(join);
691 if let Some(record) = state.tasks.get_mut(&task_id)
692 && record.running
693 {
694 record.join = join.take();
695 }
696 drop(state);
697 if let Some(join) = join {
698 join.abort();
699 } else {
700 let _ = start_tx.send(());
701 }
702 Ok(TaskStartOutcome::Pending {
703 task_id,
704 kind: initial_kind,
705 })
706 }
707
708 async fn wait_for_turn(
709 &self,
710 turn_id: &TurnId,
711 cancellation: Option<TurnCancellation>,
712 ) -> Result<Option<TurnTaskUpdate>, TaskManagerError> {
713 loop {
714 let notified = self.inner.notify.notified();
715 tokio::pin!(notified);
716 notified.as_mut().enable();
717 {
718 let mut state = self.inner.state.lock().await;
719 if let Some(queue) = state.per_turn_updates.get_mut(turn_id)
720 && let Some(update) = queue.pop_front()
721 {
722 return Ok(Some(update));
723 }
724 if state
725 .per_turn_running
726 .get(turn_id)
727 .copied()
728 .unwrap_or_default()
729 == 0
730 {
731 return Ok(None);
732 }
733 }
734 if cancellation
735 .as_ref()
736 .is_some_and(TurnCancellation::is_cancelled)
737 {
738 self.inner.interrupt_turn(turn_id).await;
739 continue;
740 }
741 if let Some(cancellation) = cancellation.as_ref() {
742 tokio::select! {
745 biased;
746 _ = &mut notified => {}
747 _ = cancellation.cancelled() => {
748 self.inner.interrupt_turn(turn_id).await;
749 },
750 }
751 } else {
752 notified.await;
753 }
754 }
755 }
756
757 async fn take_pending_loop_updates(&self) -> Result<PendingLoopUpdates, TaskManagerError> {
758 let mut state = self.inner.state.lock().await;
759 Ok(PendingLoopUpdates {
760 resolutions: std::mem::take(&mut state.pending_loop_updates),
761 })
762 }
763
764 async fn wait_for_loop_update(&self) -> Result<(), TaskManagerError> {
765 loop {
766 let notified = self.inner.notify.notified();
769 tokio::pin!(notified);
770 notified.as_mut().enable();
771 {
772 let state = self.inner.state.lock().await;
773 if !state.pending_loop_updates.is_empty() {
774 return Ok(());
775 }
776 }
777 notified.await;
778 }
779 }
780
781 async fn on_turn_interrupted(&self, turn_id: &TurnId) -> Result<(), TaskManagerError> {
782 self.inner.interrupt_turn(turn_id).await;
783 Ok(())
784 }
785
786 fn handle(&self) -> TaskManagerHandle {
787 TaskManagerHandle {
788 inner: self.inner.clone(),
789 }
790 }
791}
792
793#[async_trait]
794impl TaskManagerControl for AsyncInner {
795 async fn next_event(&self) -> Option<TaskEvent> {
796 self.host_event_rx.lock().await.recv().await
797 }
798
799 async fn cancel(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
800 let mut state = self.state.lock().await;
801 let record = state
802 .tasks
803 .get_mut(&task_id)
804 .ok_or_else(|| TaskManagerError::NotFound(task_id.clone()))?;
805 if let Some(join) = record.join.take() {
806 join.abort();
807 }
808 record.running = false;
809 let snapshot = record.snapshot.clone();
810 if record.snapshot.kind == TaskKind::Foreground
811 && let Some(count) = state.per_turn_running.get_mut(&snapshot.turn_id)
812 {
813 *count = count.saturating_sub(1);
814 if *count == 0 {
815 state.per_turn_running.remove(&snapshot.turn_id);
816 }
817 }
818 let _ = self.host_event_tx.send(TaskEvent::Cancelled(snapshot));
819 self.notify.notify_waiters();
820 Ok(())
821 }
822
823 async fn detach(&self, task_id: TaskId) -> Result<(), TaskManagerError> {
824 self.detach_running_foreground(&task_id).await
825 }
826
827 async fn list_running(&self) -> Vec<TaskSnapshot> {
828 let state = self.state.lock().await;
829 state
830 .tasks
831 .values()
832 .filter(|record| record.running)
833 .map(|record| record.snapshot.clone())
834 .collect()
835 }
836
837 async fn list_completed(&self) -> Vec<TaskSnapshot> {
838 let state = self.state.lock().await;
839 state
840 .tasks
841 .values()
842 .filter(|record| record.completed)
843 .map(|record| record.snapshot.clone())
844 .collect()
845 }
846
847 async fn drain_ready_items(&self) -> Vec<Item> {
848 let mut state = self.state.lock().await;
849 std::mem::take(&mut state.manual_ready_items)
850 }
851
852 async fn set_continue_policy(
853 &self,
854 task_id: TaskId,
855 policy: ContinuePolicy,
856 ) -> Result<(), TaskManagerError> {
857 let mut state = self.state.lock().await;
858 let record = state
859 .tasks
860 .get_mut(&task_id)
861 .ok_or_else(|| TaskManagerError::NotFound(task_id.clone()))?;
862 record.continue_policy = policy;
863 Ok(())
864 }
865
866 async fn set_delivery_mode(
867 &self,
868 task_id: TaskId,
869 mode: DeliveryMode,
870 ) -> Result<(), TaskManagerError> {
871 let mut state = self.state.lock().await;
872 let record = state
873 .tasks
874 .get_mut(&task_id)
875 .ok_or_else(|| TaskManagerError::NotFound(task_id.clone()))?;
876 record.delivery_mode = mode;
877 Ok(())
878 }
879
880 async fn wait_for_idle(&self) {
881 loop {
882 {
883 let state = self.state.lock().await;
884 if !state.tasks.values().any(|r| r.running) {
885 return;
886 }
887 }
888 self.notify.notified().await;
889 }
890 }
891}
892
893fn map_outcome_to_resolution(
894 task_id: Option<TaskId>,
895 request: ToolRequest,
896 outcome: ToolExecutionOutcome,
897) -> TaskResolution {
898 match outcome {
899 ToolExecutionOutcome::Completed(result) => TaskResolution::Item(Item {
900 id: None,
901 kind: agentkit_core::ItemKind::Tool,
902 parts: vec![agentkit_core::Part::ToolResult(result.result)],
903 metadata: result.metadata,
904 usage: None,
905 finish_reason: None,
906 created_at: None,
907 }),
908 ToolExecutionOutcome::Interrupted(
909 agentkit_tools_core::ToolInterruption::ApprovalRequired(mut approval),
910 ) => {
911 let task_id = task_id.unwrap_or_default();
912 approval.task_id = Some(task_id.clone());
913 TaskResolution::Approval(TaskApproval {
914 task_id,
915 tool_request: request,
916 approval,
917 })
918 }
919 ToolExecutionOutcome::FailedBeforeInvocation(error) => {
920 let mut metadata = request.metadata;
921 metadata.insert(TOOL_RESULT_NOT_STARTED_METADATA_KEY.into(), true.into());
922 if matches!(error, ToolError::PermissionDenied(_)) {
923 metadata.insert(
924 TOOL_RESULT_FAILURE_KIND_METADATA_KEY.into(),
925 TOOL_RESULT_FAILURE_KIND_PERMISSION_DENIED.into(),
926 );
927 }
928 TaskResolution::Item(Item {
929 id: None,
930 kind: agentkit_core::ItemKind::Tool,
931 parts: vec![agentkit_core::Part::ToolResult(ToolResultPart {
932 call_id: request.call_id,
933 output: agentkit_core::ToolOutput::Text(error.to_string()),
934 is_error: true,
935 metadata,
936 })],
937 metadata: MetadataMap::new(),
938 usage: None,
939 finish_reason: None,
940 created_at: None,
941 })
942 }
943 ToolExecutionOutcome::Failed(error) => {
944 let mut metadata = request.metadata;
945 if matches!(error, ToolError::PermissionDenied(_)) {
946 metadata.insert(
947 TOOL_RESULT_FAILURE_KIND_METADATA_KEY.into(),
948 TOOL_RESULT_FAILURE_KIND_PERMISSION_DENIED.into(),
949 );
950 }
951 TaskResolution::Item(Item {
952 id: None,
953 kind: agentkit_core::ItemKind::Tool,
954 parts: vec![agentkit_core::Part::ToolResult(ToolResultPart {
955 call_id: request.call_id,
956 output: agentkit_core::ToolOutput::Text(error.to_string()),
957 is_error: true,
958 metadata,
959 })],
960 metadata: MetadataMap::new(),
961 usage: None,
962 finish_reason: None,
963 created_at: None,
964 })
965 }
966 }
967}
968
969#[cfg(test)]
970mod tests {
971 use std::collections::BTreeMap;
972 use std::sync::Arc as StdArc;
973 use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
974
975 use agentkit_core::{
976 CancellationController, ItemKind, Part, SessionId, ToolOutput, TurnCancellation,
977 };
978 use agentkit_tools_core::{
979 ApprovalReason, PermissionChecker, PermissionDecision, ToolAnnotations, ToolInterruption,
980 ToolName, ToolResult, ToolSpec,
981 };
982 use serde_json::json;
983 use tokio::sync::Notify;
984 use tokio::time::{Duration, timeout};
985
986 use super::*;
987
988 struct AllowAllPermissions;
989
990 impl PermissionChecker for AllowAllPermissions {
991 fn evaluate(
992 &self,
993 _request: &dyn agentkit_tools_core::PermissionRequest,
994 ) -> PermissionDecision {
995 PermissionDecision::Allow
996 }
997 }
998
999 #[derive(Clone)]
1000 enum TestBehavior {
1001 Block {
1002 entered: StdArc<AtomicBool>,
1003 release: StdArc<Notify>,
1004 output: &'static str,
1005 },
1006 Approval,
1007 }
1008
1009 #[derive(Clone)]
1010 struct TestExecutor {
1011 behaviors: BTreeMap<String, TestBehavior>,
1012 }
1013
1014 impl TestExecutor {
1015 fn new(behaviors: impl IntoIterator<Item = (impl Into<String>, TestBehavior)>) -> Self {
1016 Self {
1017 behaviors: behaviors
1018 .into_iter()
1019 .map(|(name, behavior)| (name.into(), behavior))
1020 .collect(),
1021 }
1022 }
1023 }
1024
1025 #[async_trait]
1026 impl ToolExecutor for TestExecutor {
1027 fn specs(&self) -> Vec<ToolSpec> {
1028 self.behaviors
1029 .keys()
1030 .map(|name| ToolSpec {
1031 name: ToolName::new(name),
1032 description: format!("test tool {name}"),
1033 input_schema: json!({
1034 "type": "object",
1035 "properties": {},
1036 "additionalProperties": false
1037 }),
1038 output_schema: None,
1039 annotations: ToolAnnotations::default(),
1040 metadata: MetadataMap::new(),
1041 })
1042 .collect()
1043 }
1044
1045 async fn execute(
1046 &self,
1047 request: ToolRequest,
1048 _ctx: &mut agentkit_tools_core::ToolContext<'_>,
1049 ) -> ToolExecutionOutcome {
1050 match self.behaviors.get(request.tool_name.0.as_str()) {
1051 Some(TestBehavior::Block {
1052 entered,
1053 release,
1054 output,
1055 }) => {
1056 entered.store(true, AtomicOrdering::SeqCst);
1057 release.notified().await;
1058 ToolExecutionOutcome::Completed(ToolResult {
1059 result: ToolResultPart {
1060 call_id: request.call_id,
1061 output: ToolOutput::Text((*output).into()),
1062 is_error: false,
1063 metadata: request.metadata,
1064 },
1065 duration: None,
1066 metadata: MetadataMap::new(),
1067 })
1068 }
1069 Some(TestBehavior::Approval) => ToolExecutionOutcome::Interrupted(
1070 ToolInterruption::ApprovalRequired(ApprovalRequest {
1071 task_id: None,
1072 call_id: Some(request.call_id.clone()),
1073 id: "approval:test".into(),
1074 request_kind: "tool.test".into(),
1075 reason: ApprovalReason::SensitivePath,
1076 summary: "requires approval".into(),
1077 metadata: MetadataMap::new(),
1078 }),
1079 ),
1080 None => ToolExecutionOutcome::Failed(ToolError::Unavailable(
1081 request.tool_name.0.clone(),
1082 )),
1083 }
1084 }
1085 }
1086
1087 struct NameRoutingPolicy {
1088 routes: BTreeMap<String, RoutingDecision>,
1089 }
1090
1091 impl NameRoutingPolicy {
1092 fn new(routes: impl IntoIterator<Item = (impl Into<String>, RoutingDecision)>) -> Self {
1093 Self {
1094 routes: routes
1095 .into_iter()
1096 .map(|(name, decision)| (name.into(), decision))
1097 .collect(),
1098 }
1099 }
1100 }
1101
1102 impl TaskRoutingPolicy for NameRoutingPolicy {
1103 fn route(&self, request: &ToolRequest) -> RoutingDecision {
1104 self.routes
1105 .get(request.tool_name.0.as_str())
1106 .copied()
1107 .unwrap_or(RoutingDecision::Foreground)
1108 }
1109 }
1110
1111 fn make_request(tool_name: &str, turn_id: &str, call_id: &str) -> ToolRequest {
1112 ToolRequest {
1113 call_id: ToolCallId::new(call_id),
1114 tool_name: ToolName::new(tool_name),
1115 input: json!({}),
1116 session_id: SessionId::new("session-1"),
1117 turn_id: TurnId::new(turn_id),
1118 metadata: MetadataMap::new(),
1119 }
1120 }
1121
1122 fn make_context(
1123 executor: Arc<dyn ToolExecutor>,
1124 turn_id: &TurnId,
1125 cancellation: Option<TurnCancellation>,
1126 ) -> TaskStartContext {
1127 TaskStartContext {
1128 executor,
1129 tool_context: OwnedToolContext {
1130 session_id: SessionId::new("session-1"),
1131 turn_id: turn_id.clone(),
1132 metadata: MetadataMap::new(),
1133 permissions: Arc::new(AllowAllPermissions),
1134 resources: Arc::new(()),
1135 cancellation,
1136 execution_scope: None,
1137 approved_request: None,
1138 },
1139 }
1140 }
1141
1142 async fn next_event(handle: &TaskManagerHandle) -> TaskEvent {
1143 timeout(Duration::from_secs(1), handle.next_event())
1144 .await
1145 .expect("timed out waiting for task event")
1146 .expect("task event stream ended unexpectedly")
1147 }
1148
1149 async fn wait_until_entered(entered: &AtomicBool) {
1150 timeout(Duration::from_secs(1), async {
1151 while !entered.load(AtomicOrdering::SeqCst) {
1152 tokio::task::yield_now().await;
1153 }
1154 })
1155 .await
1156 .expect("task never entered execution");
1157 }
1158
1159 #[tokio::test]
1160 async fn simple_task_manager_executes_inline_and_assigns_task_ids() {
1161 let manager = SimpleTaskManager::new();
1162 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1163 "needs-approval",
1164 TestBehavior::Approval,
1165 )]));
1166 let request = make_request("needs-approval", "turn-1", "call-1");
1167
1168 let outcome = manager
1169 .start_task(
1170 TaskLaunchRequest {
1171 task_id: None,
1172 request: request.clone(),
1173 kind: TaskLaunchKind::Plain,
1174 },
1175 make_context(executor, &request.turn_id, None),
1176 )
1177 .await
1178 .unwrap();
1179
1180 match outcome {
1181 TaskStartOutcome::Ready(resolution) => match *resolution {
1182 TaskResolution::Approval(task) => {
1183 assert!(!task.task_id.0.is_empty());
1184 assert_eq!(task.approval.task_id.as_ref(), Some(&task.task_id));
1185 assert_eq!(task.tool_request.call_id, request.call_id);
1186 }
1187 other => panic!("unexpected task resolution: {other:?}"),
1188 },
1189 other => panic!("unexpected start outcome: {other:?}"),
1190 }
1191
1192 assert!(manager.handle().list_running().await.is_empty());
1193 }
1194
1195 #[tokio::test]
1196 async fn async_manager_interrupt_cancels_foreground_only() {
1197 let fg_release = StdArc::new(Notify::new());
1198 let fg_entered = StdArc::new(AtomicBool::new(false));
1199 let bg_release = StdArc::new(Notify::new());
1200 let bg_entered = StdArc::new(AtomicBool::new(false));
1201 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([
1202 (
1203 "foreground",
1204 TestBehavior::Block {
1205 entered: fg_entered.clone(),
1206 release: fg_release.clone(),
1207 output: "foreground-done",
1208 },
1209 ),
1210 (
1211 "background",
1212 TestBehavior::Block {
1213 entered: bg_entered.clone(),
1214 release: bg_release.clone(),
1215 output: "background-done",
1216 },
1217 ),
1218 ]));
1219 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([
1220 ("foreground", RoutingDecision::Foreground),
1221 ("background", RoutingDecision::Background),
1222 ]));
1223 let handle = manager.handle();
1224 let turn_id = TurnId::new("turn-1");
1225
1226 let foreground = manager
1227 .start_task(
1228 TaskLaunchRequest {
1229 task_id: None,
1230 request: make_request("foreground", "turn-1", "call-fg"),
1231 kind: TaskLaunchKind::Plain,
1232 },
1233 make_context(executor.clone(), &turn_id, None),
1234 )
1235 .await
1236 .unwrap();
1237 let background = manager
1238 .start_task(
1239 TaskLaunchRequest {
1240 task_id: None,
1241 request: make_request("background", "turn-1", "call-bg"),
1242 kind: TaskLaunchKind::Plain,
1243 },
1244 make_context(executor.clone(), &turn_id, None),
1245 )
1246 .await
1247 .unwrap();
1248
1249 assert!(matches!(
1250 foreground,
1251 TaskStartOutcome::Pending {
1252 kind: TaskKind::Foreground,
1253 ..
1254 }
1255 ));
1256 let background_id = match background {
1257 TaskStartOutcome::Pending {
1258 task_id,
1259 kind: TaskKind::Background,
1260 } => task_id,
1261 other => panic!("unexpected background outcome: {other:?}"),
1262 };
1263
1264 let _ = next_event(&handle).await;
1265 let _ = next_event(&handle).await;
1266 wait_until_entered(fg_entered.as_ref()).await;
1267 wait_until_entered(bg_entered.as_ref()).await;
1268
1269 manager.on_turn_interrupted(&turn_id).await.unwrap();
1270
1271 match next_event(&handle).await {
1272 TaskEvent::Cancelled(snapshot) => assert_eq!(snapshot.tool_name, "foreground"),
1273 other => panic!("unexpected event after interrupt: {other:?}"),
1274 }
1275
1276 let running = handle.list_running().await;
1277 assert_eq!(running.len(), 1);
1278 assert_eq!(running[0].id, background_id);
1279 assert_eq!(running[0].tool_name, "background");
1280
1281 bg_release.notify_waiters();
1282 match next_event(&handle).await {
1283 TaskEvent::Completed(snapshot, result) => {
1284 assert_eq!(snapshot.id, background_id);
1285 assert_eq!(result.output, ToolOutput::Text("background-done".into()));
1286 }
1287 other => panic!("unexpected completion event: {other:?}"),
1288 }
1289 }
1290
1291 #[tokio::test]
1292 async fn async_manager_can_manually_detach_a_foreground_task() {
1293 let release = StdArc::new(Notify::new());
1294 let entered = StdArc::new(AtomicBool::new(false));
1295 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1296 "foreground",
1297 TestBehavior::Block {
1298 entered: entered.clone(),
1299 release: release.clone(),
1300 output: "done",
1301 },
1302 )]));
1303 let manager = AsyncTaskManager::new();
1304 let handle = manager.handle();
1305 let request = make_request("foreground", "turn-1", "call-1");
1306
1307 let task_id = match manager
1308 .start_task(
1309 TaskLaunchRequest {
1310 task_id: None,
1311 request: request.clone(),
1312 kind: TaskLaunchKind::Plain,
1313 },
1314 make_context(executor, &request.turn_id, None),
1315 )
1316 .await
1317 .unwrap()
1318 {
1319 TaskStartOutcome::Pending { task_id, .. } => task_id,
1320 other => panic!("unexpected start outcome: {other:?}"),
1321 };
1322
1323 let _ = next_event(&handle).await;
1324 wait_until_entered(entered.as_ref()).await;
1325 handle.detach(task_id.clone()).await.unwrap();
1326
1327 let running = handle.list_running().await;
1328 assert_eq!(running.len(), 1);
1329 assert_eq!(running[0].kind, TaskKind::Background);
1330 match manager.wait_for_turn(&request.turn_id, None).await.unwrap() {
1331 Some(TurnTaskUpdate::Detached(snapshot)) => {
1332 assert_eq!(snapshot.id, task_id);
1333 assert_eq!(snapshot.kind, TaskKind::Background);
1334 }
1335 other => panic!("unexpected turn update: {other:?}"),
1336 }
1337 assert!(
1338 manager
1339 .wait_for_turn(&request.turn_id, None)
1340 .await
1341 .unwrap()
1342 .is_none()
1343 );
1344 match next_event(&handle).await {
1345 TaskEvent::Detached(snapshot) => assert_eq!(snapshot.id, task_id),
1346 other => panic!("unexpected event after detach: {other:?}"),
1347 }
1348
1349 release.notify_waiters();
1350 timeout(Duration::from_secs(1), handle.wait_for_idle())
1351 .await
1352 .expect("wait_for_idle timed out");
1353 }
1354
1355 #[tokio::test]
1356 async fn manual_detach_reports_invalid_task_states() {
1357 let missing = TaskId::new("missing");
1358 let manager = AsyncTaskManager::new();
1359 let handle = manager.handle();
1360 assert_eq!(
1361 handle.detach(missing.clone()).await,
1362 Err(TaskManagerError::NotFound(missing))
1363 );
1364
1365 let approval_executor: Arc<dyn ToolExecutor> =
1366 Arc::new(TestExecutor::new([("approval", TestBehavior::Approval)]));
1367 let approval_request = make_request("approval", "turn-1", "call-approval");
1368 let completed_id = match manager
1369 .start_task(
1370 TaskLaunchRequest {
1371 task_id: None,
1372 request: approval_request.clone(),
1373 kind: TaskLaunchKind::Plain,
1374 },
1375 make_context(approval_executor, &approval_request.turn_id, None),
1376 )
1377 .await
1378 .unwrap()
1379 {
1380 TaskStartOutcome::Pending { task_id, .. } => task_id,
1381 other => panic!("unexpected start outcome: {other:?}"),
1382 };
1383 let _ = next_event(&handle).await;
1384 timeout(Duration::from_secs(1), handle.wait_for_idle())
1385 .await
1386 .expect("wait_for_idle timed out");
1387 assert_eq!(
1388 handle.detach(completed_id.clone()).await,
1389 Err(TaskManagerError::NotRunning(completed_id))
1390 );
1391
1392 let release = StdArc::new(Notify::new());
1393 let entered = StdArc::new(AtomicBool::new(false));
1394 let background_executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1395 "background",
1396 TestBehavior::Block {
1397 entered: entered.clone(),
1398 release,
1399 output: "done",
1400 },
1401 )]));
1402 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1403 "background",
1404 RoutingDecision::Background,
1405 )]));
1406 let handle = manager.handle();
1407 let request = make_request("background", "turn-2", "call-background");
1408 let background_id = match manager
1409 .start_task(
1410 TaskLaunchRequest {
1411 task_id: None,
1412 request: request.clone(),
1413 kind: TaskLaunchKind::Plain,
1414 },
1415 make_context(background_executor, &request.turn_id, None),
1416 )
1417 .await
1418 .unwrap()
1419 {
1420 TaskStartOutcome::Pending { task_id, .. } => task_id,
1421 other => panic!("unexpected start outcome: {other:?}"),
1422 };
1423 let _ = next_event(&handle).await;
1424 wait_until_entered(entered.as_ref()).await;
1425 assert_eq!(
1426 handle.detach(background_id.clone()).await,
1427 Err(TaskManagerError::AlreadyBackground(background_id.clone()))
1428 );
1429 handle.cancel(background_id).await.unwrap();
1430 }
1431
1432 #[tokio::test]
1433 async fn async_manager_can_cancel_background_tasks_by_id() {
1434 let release = StdArc::new(Notify::new());
1435 let entered = StdArc::new(AtomicBool::new(false));
1436 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1437 "background",
1438 TestBehavior::Block {
1439 entered: entered.clone(),
1440 release,
1441 output: "done",
1442 },
1443 )]));
1444 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1445 "background",
1446 RoutingDecision::Background,
1447 )]));
1448 let handle = manager.handle();
1449 let request = make_request("background", "turn-1", "call-1");
1450
1451 let task_id = match manager
1452 .start_task(
1453 TaskLaunchRequest {
1454 task_id: None,
1455 request: request.clone(),
1456 kind: TaskLaunchKind::Plain,
1457 },
1458 make_context(executor, &request.turn_id, None),
1459 )
1460 .await
1461 .unwrap()
1462 {
1463 TaskStartOutcome::Pending { task_id, .. } => task_id,
1464 other => panic!("unexpected start outcome: {other:?}"),
1465 };
1466
1467 let _ = next_event(&handle).await;
1468 wait_until_entered(entered.as_ref()).await;
1469 handle.cancel(task_id.clone()).await.unwrap();
1470
1471 match next_event(&handle).await {
1472 TaskEvent::Cancelled(snapshot) => assert_eq!(snapshot.id, task_id),
1473 other => panic!("unexpected event after cancel: {other:?}"),
1474 }
1475
1476 assert!(handle.list_running().await.is_empty());
1477 }
1478
1479 #[tokio::test]
1480 async fn async_manager_manual_delivery_keeps_results_out_of_loop_updates() {
1481 let release = StdArc::new(Notify::new());
1482 let entered = StdArc::new(AtomicBool::new(false));
1483 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1484 "background",
1485 TestBehavior::Block {
1486 entered: entered.clone(),
1487 release: release.clone(),
1488 output: "manual-done",
1489 },
1490 )]));
1491 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1492 "background",
1493 RoutingDecision::Background,
1494 )]));
1495 let handle = manager.handle();
1496 let request = make_request("background", "turn-1", "call-1");
1497
1498 let task_id = match manager
1499 .start_task(
1500 TaskLaunchRequest {
1501 task_id: None,
1502 request: request.clone(),
1503 kind: TaskLaunchKind::Plain,
1504 },
1505 make_context(executor, &request.turn_id, None),
1506 )
1507 .await
1508 .unwrap()
1509 {
1510 TaskStartOutcome::Pending { task_id, .. } => task_id,
1511 other => panic!("unexpected start outcome: {other:?}"),
1512 };
1513
1514 let _ = next_event(&handle).await;
1515 wait_until_entered(entered.as_ref()).await;
1516 handle
1517 .set_continue_policy(task_id.clone(), ContinuePolicy::RequestContinue)
1518 .await
1519 .unwrap();
1520 handle
1521 .set_delivery_mode(task_id, DeliveryMode::Manual)
1522 .await
1523 .unwrap();
1524
1525 release.notify_waiters();
1526 match next_event(&handle).await {
1527 TaskEvent::Completed(_, result) => {
1528 assert_eq!(result.output, ToolOutput::Text("manual-done".into()))
1529 }
1530 other => panic!("unexpected event: {other:?}"),
1531 }
1532
1533 assert!(
1534 timeout(Duration::from_millis(50), handle.next_event())
1535 .await
1536 .is_err()
1537 );
1538 assert!(
1539 manager
1540 .take_pending_loop_updates()
1541 .await
1542 .unwrap()
1543 .resolutions
1544 .is_empty()
1545 );
1546
1547 let ready_items = handle.drain_ready_items().await;
1548 assert_eq!(ready_items.len(), 1);
1549 assert_eq!(ready_items[0].kind, ItemKind::Tool);
1550 match &ready_items[0].parts[0] {
1551 Part::ToolResult(result) => {
1552 assert_eq!(result.output, ToolOutput::Text("manual-done".into()))
1553 }
1554 other => panic!("unexpected ready item: {other:?}"),
1555 }
1556 }
1557
1558 #[tokio::test]
1559 async fn async_manager_to_loop_delivery_can_request_continue() {
1560 let release = StdArc::new(Notify::new());
1561 let entered = StdArc::new(AtomicBool::new(false));
1562 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1563 "background",
1564 TestBehavior::Block {
1565 entered: entered.clone(),
1566 release: release.clone(),
1567 output: "loop-done",
1568 },
1569 )]));
1570 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1571 "background",
1572 RoutingDecision::Background,
1573 )]));
1574 let handle = manager.handle();
1575 let request = make_request("background", "turn-1", "call-1");
1576
1577 let task_id = match manager
1578 .start_task(
1579 TaskLaunchRequest {
1580 task_id: None,
1581 request: request.clone(),
1582 kind: TaskLaunchKind::Plain,
1583 },
1584 make_context(
1585 executor,
1586 &request.turn_id,
1587 Some(TurnCancellation::new(
1588 CancellationController::new().handle(),
1589 )),
1590 ),
1591 )
1592 .await
1593 .unwrap()
1594 {
1595 TaskStartOutcome::Pending { task_id, .. } => task_id,
1596 other => panic!("unexpected start outcome: {other:?}"),
1597 };
1598
1599 let _ = next_event(&handle).await;
1600 wait_until_entered(entered.as_ref()).await;
1601 handle
1602 .set_continue_policy(task_id, ContinuePolicy::RequestContinue)
1603 .await
1604 .unwrap();
1605
1606 release.notify_waiters();
1607 match next_event(&handle).await {
1608 TaskEvent::Completed(_, result) => {
1609 assert_eq!(result.output, ToolOutput::Text("loop-done".into()))
1610 }
1611 other => panic!("unexpected completion event: {other:?}"),
1612 }
1613 match next_event(&handle).await {
1614 TaskEvent::ContinueRequested => {}
1615 other => panic!("unexpected follow-up event: {other:?}"),
1616 }
1617
1618 let updates = manager.take_pending_loop_updates().await.unwrap();
1619 assert_eq!(updates.resolutions.len(), 1);
1620 assert!(handle.drain_ready_items().await.is_empty());
1621 }
1622
1623 #[tokio::test]
1624 async fn wait_for_loop_update_wakes_without_consuming_host_events() {
1625 let release = StdArc::new(Notify::new());
1626 let entered = StdArc::new(AtomicBool::new(false));
1627 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1628 "background",
1629 TestBehavior::Block {
1630 entered: entered.clone(),
1631 release: release.clone(),
1632 output: "wake-done",
1633 },
1634 )]));
1635 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1636 "background",
1637 RoutingDecision::Background,
1638 )]));
1639 let handle = manager.handle();
1640 let request = make_request("background", "turn-1", "call-1");
1641
1642 manager
1643 .start_task(
1644 TaskLaunchRequest {
1645 task_id: None,
1646 request: request.clone(),
1647 kind: TaskLaunchKind::Plain,
1648 },
1649 make_context(executor, &request.turn_id, None),
1650 )
1651 .await
1652 .unwrap();
1653 wait_until_entered(entered.as_ref()).await;
1654
1655 let waiting = manager.wait_for_loop_update();
1656 tokio::pin!(waiting);
1657 assert!(
1658 timeout(Duration::from_millis(20), &mut waiting)
1659 .await
1660 .is_err()
1661 );
1662 release.notify_waiters();
1663 timeout(Duration::from_secs(1), &mut waiting)
1664 .await
1665 .expect("loop update wake timed out")
1666 .unwrap();
1667
1668 assert!(matches!(next_event(&handle).await, TaskEvent::Started(_)));
1669 assert!(matches!(
1670 next_event(&handle).await,
1671 TaskEvent::Completed(_, _)
1672 ));
1673 assert_eq!(
1674 manager
1675 .take_pending_loop_updates()
1676 .await
1677 .unwrap()
1678 .resolutions
1679 .len(),
1680 1
1681 );
1682 }
1683
1684 #[tokio::test]
1685 async fn wait_for_idle_returns_after_loop_updates_are_queued() {
1686 let release = StdArc::new(Notify::new());
1687 let entered = StdArc::new(AtomicBool::new(false));
1688 let executor: Arc<dyn ToolExecutor> = Arc::new(TestExecutor::new([(
1689 "background",
1690 TestBehavior::Block {
1691 entered: entered.clone(),
1692 release: release.clone(),
1693 output: "idle-done",
1694 },
1695 )]));
1696 let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([(
1697 "background",
1698 RoutingDecision::Background,
1699 )]));
1700 let handle = manager.handle();
1701 let request = make_request("background", "turn-1", "call-1");
1702
1703 let outcome = manager
1704 .start_task(
1705 TaskLaunchRequest {
1706 task_id: None,
1707 request: request.clone(),
1708 kind: TaskLaunchKind::Plain,
1709 },
1710 make_context(executor, &request.turn_id, None),
1711 )
1712 .await
1713 .unwrap();
1714 assert!(matches!(outcome, TaskStartOutcome::Pending { .. }));
1715
1716 let _ = next_event(&handle).await;
1717 wait_until_entered(entered.as_ref()).await;
1718 release.notify_waiters();
1719
1720 timeout(Duration::from_secs(1), handle.wait_for_idle())
1721 .await
1722 .expect("wait_for_idle timed out");
1723
1724 let updates = manager.take_pending_loop_updates().await.unwrap();
1725 assert_eq!(updates.resolutions.len(), 1);
1726 match &updates.resolutions[0] {
1727 TaskResolution::Item(item) => match &item.parts[0] {
1728 Part::ToolResult(result) => {
1729 assert_eq!(result.call_id, request.call_id);
1730 assert_eq!(result.output, ToolOutput::Text("idle-done".into()));
1731 }
1732 other => panic!("unexpected tool item: {other:?}"),
1733 },
1734 other => panic!("unexpected pending update: {other:?}"),
1735 }
1736 }
1737}