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