telos_agent/tools/executor/
stream.rs1use crate::config::AgentConfig;
4use crate::diagnostics::ToolFailureKind;
5use crate::model::message::{ToolCall, ToolResult};
6use crate::tools::api::{ToolProgress, ToolRegistry};
7use async_stream::stream;
8use futures_core::stream::Stream;
9use futures_util::FutureExt;
10use std::panic::AssertUnwindSafe;
11use std::sync::Arc;
12use tracing::warn;
13
14use super::batch::build_batches;
15use super::invoke::{invoke_tool, json_error_payload, record_tool_failure, tool_detail};
16use super::types::{PreparedCall, ToolExecutionEvent, ToolExecutionStreamItem};
17
18enum WorkerMessage {
19 Event(ToolExecutionEvent),
20 Done { index: usize, result: ToolResult, feedback: Vec<String> },
21}
22
23pub fn execute_tool_calls_stream<'a>(
24 calls: Vec<ToolCall>,
25 tools: &'a ToolRegistry,
26 config: &'a AgentConfig,
27 session_id: &'a str,
28 turn_id: u64,
29 messages: Arc<Vec<crate::model::message::Message>>,
30 read_file_state: crate::tools::api::FileReadState,
31) -> impl Stream<Item = ToolExecutionStreamItem> + 'a {
32 let batches =
33 build_batches(calls, tools, config, session_id, turn_id, messages, read_file_state);
34
35 stream! {
36 for batch in batches {
37 let items = if batch.concurrency_safe && config.tool_concurrency_limit > 1 {
38 execute_concurrent_batch(batch.calls, tools.clone(), config.clone()).await
39 } else {
40 execute_sequential_batch(batch.calls, tools.clone(), config.clone()).await
41 };
42 for item in items {
43 yield item;
44 }
45 }
46 }
47}
48
49async fn execute_concurrent_batch(
50 calls: Vec<PreparedCall>,
51 tools: ToolRegistry,
52 config: AgentConfig,
53) -> Vec<ToolExecutionStreamItem> {
54 let limit = config.tool_concurrency_limit;
55 let mut queued = calls.into_iter().peekable();
56 let mut active = 0usize;
57 let mut completed = Vec::new();
58 let mut items = Vec::new();
59 let (send, mut recv) = tokio::sync::mpsc::unbounded_channel::<WorkerMessage>();
60 let worker_tx = send.clone();
61 let mut join_set = tokio::task::JoinSet::new();
62
63 while active < limit {
64 let Some(prepared) = queued.next() else {
65 break;
66 };
67 active += 1;
68 spawn_tool_event_worker(
69 &mut join_set,
70 prepared,
71 tools.clone(),
72 config.clone(),
73 worker_tx.clone(),
74 );
75 }
76 drop(send);
78
79 loop {
80 match recv.recv().await {
81 Some(WorkerMessage::Event(event)) => {
82 items.push(ToolExecutionStreamItem::Event(event));
83 }
84 Some(WorkerMessage::Done { index, result, feedback }) => {
85 active -= 1;
86 completed.push((index, result, feedback));
87
88 if let Some(prepared) = queued.next() {
89 active += 1;
90 spawn_tool_event_worker(
91 &mut join_set,
92 prepared,
93 tools.clone(),
94 config.clone(),
95 worker_tx.clone(),
96 );
97 }
98
99 if active == 0 {
100 break;
101 }
102 }
103 None => break,
104 }
105 }
106
107 drop(join_set);
109
110 completed.sort_by_key(|(index, _, _)| *index);
111 for (_, result, feedback) in completed {
112 items.push(ToolExecutionStreamItem::Result { result, feedback });
113 }
114
115 items
116}
117
118async fn execute_sequential_batch(
119 calls: Vec<PreparedCall>,
120 tools: ToolRegistry,
121 config: AgentConfig,
122) -> Vec<ToolExecutionStreamItem> {
123 let mut items = Vec::new();
124 for prepared in calls {
125 let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<WorkerMessage>();
126 let mut join_set = tokio::task::JoinSet::new();
127 spawn_tool_event_worker(&mut join_set, prepared, tools.clone(), config.clone(), tx);
128 while let Some(msg) = rx.recv().await {
129 match msg {
130 WorkerMessage::Event(event) => items.push(ToolExecutionStreamItem::Event(event)),
131 WorkerMessage::Done { index: _, result, feedback } => {
132 items.push(ToolExecutionStreamItem::Result { result, feedback })
133 }
134 }
135 }
136 }
137 items
138}
139
140fn spawn_tool_event_worker(
141 join_set: &mut tokio::task::JoinSet<()>,
142 prepared: PreparedCall,
143 tools: ToolRegistry,
144 config: AgentConfig,
145 tx: tokio::sync::mpsc::UnboundedSender<WorkerMessage>,
146) {
147 let index = prepared.index;
148 let tool_call_id = prepared.call.id.clone();
149 let name = prepared.call.name.clone();
150 let panic_call = prepared.call.clone();
151 let panic_context = prepared.context.clone();
152 let panic_config = config.clone();
153
154 join_set.spawn(async move {
155 let result = AssertUnwindSafe(run_tool_with_event_forwarding(prepared, tools, config, tx.clone()))
156 .catch_unwind()
157 .await;
158
159 if let Err(ref err) = result {
160 let message = if let Some(s) = err.downcast_ref::<String>() {
161 s.clone()
162 } else if let Some(s) = err.downcast_ref::<&str>() {
163 s.to_string()
164 } else {
165 "tool invocation panicked".to_string()
166 };
167 warn!(tool = %name, tool_call_id = %tool_call_id, "tool invocation panicked: {message}");
168 record_tool_failure(
169 &panic_config,
170 &panic_context,
171 &panic_call,
172 ToolFailureKind::ExecutionPanic,
173 &message,
174 )
175 .await;
176 let result = ToolResult {
177 tool_call_id, name, content: json_error_payload("execution_panic", message),
178 is_error: true,
179 };
180 let _ = tx.send(WorkerMessage::Done { index, result, feedback: Vec::new() });
181 }
182 });
183}
184
185async fn run_tool_with_event_forwarding(
186 prepared: PreparedCall,
187 tools: ToolRegistry,
188 config: AgentConfig,
189 tx: tokio::sync::mpsc::UnboundedSender<WorkerMessage>,
190) {
191 let index = prepared.index;
192 let detail = tool_detail(&tools, &prepared.call.name, &prepared.call.arguments);
193 let _ = tx.send(WorkerMessage::Event(ToolExecutionEvent::ToolStarted {
194 tool_call_id: prepared.call.id.clone(),
195 name: prepared.call.name.clone(),
196 detail,
197 }));
198
199 let (progress_tx, mut progress_rx) = tokio::sync::mpsc::unbounded_channel::<ToolProgress>();
200 let mut context = prepared.context;
201 context.progress = Some(progress_tx);
202
203 let call = prepared.call.clone();
204 let tool = tools.get(&prepared.call.name);
205 let context_for_not_found = context.clone();
206 let result_task = async move {
207 match tool {
208 Ok(tool) => invoke_tool(call, tool, context, &config, &tools).await,
209 Err(err) => {
210 record_tool_failure(
211 &config,
212 &context_for_not_found,
213 &call,
214 ToolFailureKind::ToolNotFound,
215 &err.to_string(),
216 )
217 .await;
218 (
219 Vec::new(),
220 ToolResult {
221 tool_call_id: call.id.clone(),
222 name: call.name.clone(),
223 content: json_error_payload("tool_not_found", err.to_string()),
224 is_error: true,
225 },
226 Vec::new(),
227 )
228 }
229 }
230 };
231 tokio::pin!(result_task);
232
233 let (approval_events, result, feedback) = loop {
234 tokio::select! {
235 maybe_progress = progress_rx.recv() => {
236 if let Some(progress) = maybe_progress {
237 let _ = tx.send(WorkerMessage::Event(ToolExecutionEvent::ToolProgress {
238 tool_call_id: progress.tool_call_id,
239 name: prepared.call.name.clone(),
240 message: progress.message,
241 data: progress.data,
242 }));
243 }
244 }
245 result = &mut result_task => {
246 while let Ok(progress) = progress_rx.try_recv() {
247 let _ = tx.send(WorkerMessage::Event(ToolExecutionEvent::ToolProgress {
248 tool_call_id: progress.tool_call_id,
249 name: prepared.call.name.clone(),
250 message: progress.message,
251 data: progress.data,
252 }));
253 }
254 break result;
255 }
256 }
257 };
258
259 for event in approval_events {
260 let _ = tx.send(WorkerMessage::Event(event));
261 }
262
263 let _ = tx.send(WorkerMessage::Done { index, result, feedback });
264}