Skip to main content

acp_utils/conversation/
tool_calls.rs

1use crate::notifications::{SubAgentEvent, SubAgentProgressParams};
2use agent_client_protocol::schema::{MaybeUndefined, v2 as acp};
3use schemars::JsonSchema;
4use serde::Serialize;
5
6/// A tracked tool call within a sub-agent.
7#[derive(Debug, Clone, PartialEq, Eq, Serialize, JsonSchema)]
8#[serde(rename_all = "camelCase")]
9pub struct SubAgentToolCall {
10    pub id: String,
11    pub name: String,
12    pub raw_input: String,
13    pub display_value: Option<String>,
14    pub status: ToolStatus,
15    #[serde(skip)]
16    kind: ToolKind,
17}
18
19impl SubAgentToolCall {
20    pub fn bash_command(&self) -> Option<String> {
21        bash_command(self.kind, &self.raw_input)
22    }
23}
24
25/// Per-sub-agent state: tracks its tool calls in arrival order.
26#[derive(Debug, Clone, PartialEq, Eq, Serialize, JsonSchema)]
27#[serde(rename_all = "camelCase")]
28pub struct SubAgentState {
29    pub task_id: String,
30    pub agent_name: String,
31    pub done: bool,
32    pub tool_calls: Vec<SubAgentToolCall>,
33}
34
35impl SubAgentState {
36    fn tool_call_mut(&mut self, id: &str) -> Option<&mut SubAgentToolCall> {
37        self.tool_calls.iter_mut().find(|call| call.id == id)
38    }
39
40    fn upsert(&mut self, id: &str, name: &str, arguments: String) -> &mut SubAgentToolCall {
41        let index = self.tool_calls.iter().position(|call| call.id == id).unwrap_or_else(|| {
42            self.tool_calls.push(SubAgentToolCall {
43                id: id.to_string(),
44                name: name.to_string(),
45                raw_input: arguments,
46                display_value: None,
47                status: ToolStatus::Running,
48                kind: tool_kind(name),
49            });
50            self.tool_calls.len() - 1
51        });
52        &mut self.tool_calls[index]
53    }
54}
55
56/// A tool call as the merge of every update the agent sent for it.
57#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
58#[serde(rename_all = "camelCase")]
59pub struct ToolCall {
60    pub status: ToolStatus,
61    #[serde(skip_serializing_if = "Option::is_none")]
62    pub error: Option<String>,
63    pub sub_agents: Vec<SubAgentState>,
64    #[serde(rename = "toolCall")]
65    protocol: Box<acp::ToolCallUpdate>,
66}
67
68impl ToolCall {
69    pub(super) fn from_update(update: &acp::ToolCallUpdate) -> Self {
70        let mut tool = Self {
71            status: ToolStatus::Running,
72            error: None,
73            sub_agents: Vec::new(),
74            protocol: Box::new(update.clone()),
75        };
76        tool.refresh_status();
77        tool
78    }
79
80    pub fn title(&self) -> &str {
81        self.protocol.title.value().map_or("", String::as_str)
82    }
83
84    pub fn raw_input(&self) -> String {
85        self.protocol.raw_input.value().map_or_else(String::new, raw_input_fragment)
86    }
87
88    pub fn display_value(&self) -> Option<&str> {
89        self.meta_str("display_value")
90    }
91
92    pub fn content(&self) -> &[acp::ToolCallContent] {
93        self.protocol.content.value().map_or(&[], Vec::as_slice)
94    }
95
96    pub fn diffs(&self) -> impl Iterator<Item = &acp::Diff> {
97        self.content().iter().filter_map(|content| match content {
98            acp::ToolCallContent::Diff(diff) => Some(diff),
99            _ => None,
100        })
101    }
102
103    pub(super) fn apply_update(&mut self, update: &acp::ToolCallUpdate) {
104        self.protocol.apply_update(update.clone());
105        self.refresh_status();
106    }
107
108    pub(super) fn append_content(&mut self, content: acp::ToolCallContent) {
109        match &mut self.protocol.content {
110            MaybeUndefined::Value(items) => items.push(content),
111            value => *value = MaybeUndefined::Value(vec![content]),
112        }
113    }
114
115    pub(super) fn apply_sub_agent_progress(&mut self, notification: &SubAgentProgressParams) {
116        apply_sub_agent_progress(&mut self.sub_agents, notification);
117    }
118
119    pub(super) fn finalize(&mut self, status: ToolStatus, error: Option<&str>) {
120        if self.status == ToolStatus::Running {
121            self.status = status;
122            self.error = error.map(str::to_owned);
123        }
124        for agent in &mut self.sub_agents {
125            agent.done = true;
126            for call in &mut agent.tool_calls {
127                if call.status == ToolStatus::Running {
128                    call.status = status;
129                }
130            }
131        }
132    }
133
134    pub fn bash_command(&self) -> Option<String> {
135        bash_command(self.kind(), &self.raw_input())
136    }
137
138    pub(super) fn is_running(&self) -> bool {
139        self.status == ToolStatus::Running
140            || self
141                .sub_agents
142                .iter()
143                .any(|agent| !agent.done || agent.tool_calls.iter().any(|call| call.status == ToolStatus::Running))
144    }
145
146    pub(super) fn rendering_final(&self) -> bool {
147        !self.is_running() && (self.kind() != ToolKind::SpawnSubagent || !self.sub_agents.is_empty())
148    }
149
150    fn kind(&self) -> ToolKind {
151        tool_kind(self.protocol.name.value().map_or_else(|| self.title(), String::as_str))
152    }
153
154    fn refresh_status(&mut self) {
155        self.status = match self.protocol.status.value() {
156            Some(acp::ToolCallStatus::Completed) => ToolStatus::Success,
157            Some(acp::ToolCallStatus::Failed) => ToolStatus::Failed,
158            Some(acp::ToolCallStatus::Cancelled) => ToolStatus::Cancelled,
159            _ => ToolStatus::Running,
160        };
161        self.error = None;
162    }
163
164    fn meta_str(&self, key: &str) -> Option<&str> {
165        self.protocol.meta.value().and_then(|meta| meta.get(key)).and_then(serde_json::Value::as_str)
166    }
167}
168
169#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, JsonSchema)]
170#[serde(rename_all = "snake_case")]
171pub enum ToolStatus {
172    Running,
173    Success,
174    Cancelled,
175    Failed,
176}
177
178fn apply_sub_agent_progress(states: &mut Vec<SubAgentState>, notification: &SubAgentProgressParams) {
179    let index = states.iter().position(|agent| agent.task_id == notification.task_id).unwrap_or_else(|| {
180        states.push(SubAgentState {
181            task_id: notification.task_id.clone(),
182            agent_name: notification.agent_name.clone(),
183            done: false,
184            tool_calls: Vec::new(),
185        });
186        states.len() - 1
187    });
188    let agent = &mut states[index];
189
190    match &notification.event {
191        SubAgentEvent::ToolCall { request } => {
192            let call = agent.upsert(&request.id, &request.name, request.arguments.clone());
193            update_title(&mut call.name, &request.name);
194            call.kind = tool_kind(&request.name);
195            call.raw_input.clone_from(&request.arguments);
196            call.status = ToolStatus::Running;
197        }
198        SubAgentEvent::ToolCallUpdate { update } => {
199            let call = agent.upsert(&update.id, "tool", String::new());
200            call.raw_input.push_str(&update.chunk);
201            call.status = ToolStatus::Running;
202        }
203        SubAgentEvent::ToolResult { result } => {
204            if let Some(call) = agent.tool_call_mut(&result.id) {
205                call.status = ToolStatus::Success;
206                if let Some(result_meta) = &result.result_meta {
207                    call.name.clone_from(&result_meta.display.title);
208                    call.display_value = Some(result_meta.display.value.clone());
209                }
210            }
211        }
212        SubAgentEvent::ToolError { error } => {
213            if let Some(call) = agent.tool_call_mut(&error.id) {
214                call.status = ToolStatus::Failed;
215            }
216        }
217        SubAgentEvent::Done => agent.done = true,
218        SubAgentEvent::Other => {}
219    }
220}
221
222fn update_title(current: &mut String, new_title: &str) {
223    if !new_title.is_empty() {
224        current.clear();
225        current.push_str(new_title);
226    }
227}
228
229fn raw_input_fragment(raw_input: &serde_json::Value) -> String {
230    raw_input.as_str().map_or_else(|| raw_input.to_string(), str::to_string)
231}
232
233#[derive(Debug, Clone, Copy, PartialEq, Eq)]
234enum ToolKind {
235    Bash,
236    SpawnSubagent,
237    Other,
238}
239
240fn tool_kind(tool_name: &str) -> ToolKind {
241    let name = tool_name.rsplit("__").next().unwrap_or(tool_name);
242    if name.eq_ignore_ascii_case("bash") {
243        ToolKind::Bash
244    } else if name.eq_ignore_ascii_case("spawn_subagent") {
245        ToolKind::SpawnSubagent
246    } else {
247        ToolKind::Other
248    }
249}
250
251fn bash_command(kind: ToolKind, raw_input: &str) -> Option<String> {
252    if kind != ToolKind::Bash {
253        return None;
254    }
255    serde_json::from_str::<serde_json::Value>(raw_input).ok()?.get("command")?.as_str().map(str::to_string)
256}