Skip to main content

telos_agent/tools/executor/
stream.rs

1//! Streaming tool execution path.
2
3use 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    // Done sending, so drop the sender to close the channel when all workers are done.
77    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    // task done, so drop the join set to avoid holding onto any tasks.
108    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}