Skip to main content

wisp/components/
tool_call_statuses.rs

1use acp_utils::notifications::SubAgentProgressParams;
2use agent_client_protocol::schema as acp;
3use std::collections::HashMap;
4use std::time::Instant;
5
6use crate::components::sub_agent_tracker::SubAgentTracker;
7use crate::components::tool_call_status_view::{ToolCallStatus, diff_preview_from_acp, render_tool_tree};
8use crate::components::tracked_tool_call::{TrackedToolCall, raw_input_fragment, upsert_tracked_tool_call};
9use tui::{Frame, ViewContext};
10
11#[derive(Debug, Clone, PartialEq, Eq)]
12pub enum PromptTermination {
13    EndTurn,
14    Cancelled,
15    Failed(String),
16}
17
18impl PromptTermination {
19    pub(crate) fn terminal_status(&self) -> ToolCallStatus {
20        match self {
21            PromptTermination::EndTurn => ToolCallStatus::Success,
22            PromptTermination::Cancelled => ToolCallStatus::Error("cancelled".to_string()),
23            PromptTermination::Failed(msg) => ToolCallStatus::Error(format!("failed: {msg}")),
24        }
25    }
26}
27
28/// Tracks active tool calls and produces status lines for the frame.
29#[derive(Clone)]
30pub struct ToolCallStatuses {
31    /// Ordered list of tool call IDs (insertion order)
32    tool_order: Vec<String>,
33    /// Tool call info by ID
34    tool_calls: HashMap<String, TrackedToolCall>,
35    /// Sub-agent states keyed by parent tool call ID
36    sub_agents: SubAgentTracker,
37    /// Animation tick for the spinner on running tool calls
38    tick: u16,
39}
40
41impl ToolCallStatuses {
42    pub fn new() -> Self {
43        Self { tool_order: Vec::new(), tool_calls: HashMap::new(), sub_agents: SubAgentTracker::default(), tick: 0 }
44    }
45
46    pub fn running_any(&self) -> bool {
47        self.tool_calls.values().any(|tc| matches!(tc.status, ToolCallStatus::Running)) || self.sub_agents.any_running()
48    }
49
50    /// Advance the animation state. Call this on tick events.
51    pub fn on_tick(&mut self, _now: Instant) {
52        if self.running_any() {
53            self.tick = self.tick.wrapping_add(1);
54        }
55    }
56
57    /// Handle a new tool call from ACP `SessionUpdate::ToolCall`.
58    pub fn on_tool_call(&mut self, tool_call: &acp::ToolCall) {
59        let id = tool_call.tool_call_id.0.to_string();
60        let arguments = tool_call.raw_input.as_ref().map(raw_input_fragment).unwrap_or_default();
61
62        let tracked = upsert_tracked_tool_call(
63            &mut self.tool_order,
64            &mut self.tool_calls,
65            &id,
66            &tool_call.title,
67            arguments.clone(),
68        );
69        tracked.update_name(&tool_call.title);
70        tracked.arguments = arguments;
71        tracked.status = ToolCallStatus::Running;
72    }
73
74    /// Handle a tool call update from ACP `SessionUpdate::ToolCallUpdate`.
75    pub fn on_tool_call_update(&mut self, update: &acp::ToolCallUpdate) {
76        let id = update.tool_call_id.0.to_string();
77
78        if let Some(tc) = self.tool_calls.get_mut(&id) {
79            if let Some(title) = &update.fields.title {
80                tc.update_name(title);
81            }
82            if let Some(raw_input) = &update.fields.raw_input {
83                tc.append_arguments(&raw_input_fragment(raw_input));
84            }
85            if let Some(meta) = &update.meta
86                && let Some(dv) = meta.get("display_value").and_then(|v| v.as_str())
87            {
88                tc.display_value = Some(dv.to_string());
89            }
90            if let Some(content) = &update.fields.content {
91                for item in content {
92                    if let acp::ToolCallContent::Diff(diff) = item {
93                        tc.diff_preview = Some(diff_preview_from_acp(diff));
94                    }
95                }
96            }
97            if let Some(status) = update.fields.status {
98                tc.apply_status(status);
99            }
100        }
101    }
102
103    pub fn finalize_running(&mut self, termination: &PromptTermination) {
104        let terminal_status = termination.terminal_status();
105
106        for tool_call in self.tool_calls.values_mut() {
107            if matches!(tool_call.status, ToolCallStatus::Running) {
108                tool_call.status = terminal_status.clone();
109            }
110        }
111
112        self.sub_agents.finalize_running(termination);
113    }
114
115    pub fn has_tool(&self, id: &str) -> bool {
116        self.tool_calls.contains_key(id)
117    }
118
119    #[cfg(test)]
120    pub fn is_tool_running(&self, id: &str) -> bool {
121        self.tool_calls.get(id).is_some_and(|tc| matches!(tc.status, ToolCallStatus::Running))
122    }
123
124    /// Handle a sub-agent progress notification.
125    pub fn on_sub_agent_progress(&mut self, notification: &SubAgentProgressParams) {
126        self.sub_agents.on_progress(notification);
127    }
128
129    #[cfg(test)]
130    pub fn remove_tool(&mut self, id: &str) {
131        self.tool_calls.remove(id);
132        self.tool_order.retain(|tool_id| tool_id != id);
133        self.sub_agents.remove(id);
134    }
135
136    pub fn render_tool(&self, id: &str, context: &ViewContext) -> Frame {
137        render_tool_tree(id, &self.tool_calls, &self.sub_agents, self.tick, context)
138    }
139
140    /// Clear all tracked tool calls (e.g., after pushing to scrollback).
141    pub fn clear(&mut self) {
142        self.tool_order.clear();
143        self.tool_calls.clear();
144        self.sub_agents.clear();
145    }
146}
147
148impl Default for ToolCallStatuses {
149    fn default() -> Self {
150        Self::new()
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use super::*;
157    use acp_utils::notifications::{SubAgentEvent, SubAgentProgressParams};
158    use tui::{DiffLine, DiffPreview, DiffTag, SplitDiffEntry, SplitDiffRow};
159
160    fn ctx() -> ViewContext {
161        ViewContext::new((80, 24))
162    }
163
164    fn make_tool_call(id: &str, title: &str, raw_input: Option<&str>) -> acp::ToolCall {
165        let mut tc = acp::ToolCall::new(id.to_string(), title);
166        if let Some(input) = raw_input {
167            tc = tc.raw_input(serde_json::from_str::<serde_json::Value>(input).unwrap());
168        }
169        tc
170    }
171
172    fn make_tool_call_update(id: &str, status: acp::ToolCallStatus) -> acp::ToolCallUpdate {
173        acp::ToolCallUpdate::new(id.to_string(), acp::ToolCallUpdateFields::new().status(status))
174    }
175
176    fn make_sub_agent_notification(parent_tool_id: &str, agent_name: &str, event_json: &str) -> SubAgentProgressParams {
177        make_sub_agent_notification_with_task_id(parent_tool_id, agent_name, agent_name, event_json)
178    }
179
180    fn make_sub_agent_notification_with_task_id(
181        parent_tool_id: &str,
182        task_id: &str,
183        agent_name: &str,
184        event_json: &str,
185    ) -> SubAgentProgressParams {
186        let json = format!(
187            r#"{{"parent_tool_id":"{parent_tool_id}","task_id":"{task_id}","agent_name":"{agent_name}","event":{event_json}}}"#,
188        );
189        serde_json::from_str(&json).unwrap()
190    }
191
192    #[test]
193    fn progress_reports_sub_agent_running_tools() {
194        let mut statuses = ToolCallStatuses::new();
195        statuses.on_tool_call(&make_tool_call("parent-1", "spawn_subagent", None));
196        statuses.on_tool_call_update(&make_tool_call_update("parent-1", acp::ToolCallStatus::Completed));
197        statuses.on_sub_agent_progress(&make_sub_agent_notification(
198            "parent-1",
199            "explorer",
200            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{}"},"model_name":"m"}}"#,
201        ));
202
203        assert!(statuses.running_any());
204    }
205
206    #[test]
207    fn remove_tool_cleans_up_sub_agent_state() {
208        let mut statuses = ToolCallStatuses::new();
209        statuses.on_tool_call(&make_tool_call("parent-1", "spawn_subagent", None));
210        statuses.on_sub_agent_progress(&make_sub_agent_notification(
211            "parent-1",
212            "explorer",
213            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{}"},"model_name":"m"}}"#,
214        ));
215
216        statuses.remove_tool("parent-1");
217        assert!(!statuses.running_any());
218        assert!(statuses.render_tool("parent-1", &ctx()).lines().is_empty());
219    }
220
221    #[test]
222    fn clear_removes_sub_agent_state() {
223        let mut statuses = ToolCallStatuses::new();
224        statuses.on_tool_call(&make_tool_call("parent-1", "spawn_subagent", None));
225        statuses.on_sub_agent_progress(&make_sub_agent_notification(
226            "parent-1",
227            "explorer",
228            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{}"},"model_name":"m"}}"#,
229        ));
230
231        statuses.clear();
232        assert!(!statuses.running_any());
233    }
234
235    #[test]
236    fn deserialize_tool_call_event() {
237        let n = make_sub_agent_notification(
238            "p1",
239            "explorer",
240            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{\"pattern\":\"test\"}"},"model_name":"m"}}"#,
241        );
242        assert!(matches!(n.event, SubAgentEvent::ToolCall { .. }));
243    }
244
245    #[test]
246    fn deserialize_tool_call_update_event() {
247        let n = make_sub_agent_notification(
248            "p1",
249            "explorer",
250            r#"{"ToolCallUpdate":{"update":{"id":"c1","chunk":"{\"pattern\":\"updated\"}"},"model_name":"m"}}"#,
251        );
252        assert!(matches!(n.event, SubAgentEvent::ToolCallUpdate { .. }));
253    }
254
255    #[test]
256    fn deserialize_tool_result_event() {
257        let n = make_sub_agent_notification(
258            "p1",
259            "explorer",
260            r#"{"ToolResult":{"result":{"id":"c1","name":"grep","arguments":"{}","result":"ok"},"model_name":"m"}}"#,
261        );
262        assert!(matches!(n.event, SubAgentEvent::ToolResult { .. }));
263    }
264
265    #[test]
266    fn deserialize_done_event() {
267        let n = make_sub_agent_notification("p1", "explorer", r#""Done""#);
268        assert!(matches!(n.event, SubAgentEvent::Done));
269    }
270
271    #[test]
272    fn deserialize_other_variant() {
273        let n = make_sub_agent_notification("p1", "explorer", r#""Other""#);
274        assert!(matches!(n.event, SubAgentEvent::Other));
275    }
276
277    #[test]
278    fn test_diff_preview_rendered_on_success() {
279        let mut statuses = ToolCallStatuses::new();
280        statuses.on_tool_call(&make_tool_call("tool-1", "Edit", None));
281
282        let tc = statuses.tool_calls.get_mut("tool-1").unwrap();
283        tc.status = ToolCallStatus::Success;
284        tc.diff_preview = Some(DiffPreview {
285            lines: vec![
286                DiffLine { tag: DiffTag::Removed, content: "old line".to_string() },
287                DiffLine { tag: DiffTag::Added, content: "new line".to_string() },
288            ],
289            rows: vec![SplitDiffRow {
290                left: Some(SplitDiffEntry::new(DiffTag::Removed, "old line", Some(1))),
291                right: Some(SplitDiffEntry::new(DiffTag::Added, "new line", Some(1))),
292            }],
293            lang_hint: "rs".to_string(),
294            start_line: Some(1),
295        });
296
297        let frame = statuses.render_tool("tool-1", &ctx());
298        let lines = frame.lines();
299        assert!(lines.len() > 1);
300        let all_text: String = lines.iter().map(tui::Line::plain_text).collect();
301        assert!(all_text.contains("old line"), "Expected removed line: {all_text}");
302        assert!(all_text.contains("new line"), "Expected added line: {all_text}");
303    }
304
305    #[test]
306    fn test_diff_preview_not_rendered_while_running() {
307        let mut statuses = ToolCallStatuses::new();
308        statuses.on_tool_call(&make_tool_call("tool-1", "Edit", None));
309
310        let tc = statuses.tool_calls.get_mut("tool-1").unwrap();
311        tc.diff_preview = Some(DiffPreview {
312            lines: vec![DiffLine { tag: DiffTag::Added, content: "new line".to_string() }],
313            rows: vec![SplitDiffRow {
314                left: None,
315                right: Some(SplitDiffEntry::new(DiffTag::Added, "new line", Some(1))),
316            }],
317            lang_hint: "rs".to_string(),
318            start_line: Some(1),
319        });
320
321        let frame = statuses.render_tool("tool-1", &ctx());
322        assert_eq!(frame.lines().len(), 1, "Should only have status line while running");
323    }
324
325    #[test]
326    fn finalize_running_marks_top_level_tools_terminal() {
327        let mut statuses = ToolCallStatuses::new();
328        statuses.on_tool_call(&make_tool_call("tool-1", "Read", None));
329
330        statuses.finalize_running(&PromptTermination::EndTurn);
331
332        assert!(!statuses.is_tool_running("tool-1"));
333        assert!(!statuses.running_any());
334        let frame = statuses.render_tool("tool-1", &ctx());
335        assert!(frame.lines()[0].plain_text().contains('✓'));
336    }
337
338    #[test]
339    fn finalize_running_marks_sub_agent_tools_terminal() {
340        let mut statuses = ToolCallStatuses::new();
341        statuses.on_sub_agent_progress(&make_sub_agent_notification(
342            "parent-1",
343            "explorer",
344            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{}"},"model_name":"m"}}"#,
345        ));
346
347        assert!(statuses.running_any());
348
349        statuses.finalize_running(&PromptTermination::Cancelled);
350
351        assert!(!statuses.running_any());
352    }
353
354    #[test]
355    fn finalize_running_failed_variant_preserves_cause() {
356        let mut statuses = ToolCallStatuses::new();
357        statuses.on_tool_call(&make_tool_call("tool-1", "Read", None));
358        assert!(statuses.is_tool_running("tool-1"));
359
360        statuses.finalize_running(&PromptTermination::Failed("error".to_string()));
361
362        assert!(!statuses.is_tool_running("tool-1"));
363        assert!(!statuses.running_any());
364    }
365
366    #[test]
367    fn finalize_running_cancelled_vs_failed_both_error() {
368        let mut cancelled = ToolCallStatuses::new();
369        cancelled.on_tool_call(&make_tool_call("t1", "test", None));
370        cancelled.finalize_running(&PromptTermination::Cancelled);
371
372        let mut failed = ToolCallStatuses::new();
373        failed.on_tool_call(&make_tool_call("t1", "test", None));
374        failed.finalize_running(&PromptTermination::Failed("error".to_string()));
375
376        let cancelled_frame = cancelled.render_tool("t1", &ctx());
377        let failed_frame = failed.render_tool("t1", &ctx());
378        assert!(!cancelled_frame.lines()[0].plain_text().contains('\u{2713}'));
379        assert!(!failed_frame.lines()[0].plain_text().contains('\u{2713}'));
380    }
381
382    #[test]
383    fn running_any_false_when_sub_agent_done() {
384        let mut statuses = ToolCallStatuses::new();
385        statuses.on_sub_agent_progress(&make_sub_agent_notification(
386            "parent-1",
387            "explorer",
388            r#"{"ToolCall":{"request":{"id":"c1","name":"grep","arguments":"{}"},"model_name":"m"}}"#,
389        ));
390        assert!(statuses.running_any());
391
392        statuses.on_sub_agent_progress(&make_sub_agent_notification("parent-1", "explorer", r#""Done""#));
393        assert!(statuses.running_any());
394
395        statuses.finalize_running(&PromptTermination::EndTurn);
396        assert!(!statuses.running_any());
397    }
398}