1use std::collections::{HashMap, HashSet, VecDeque};
10
11use async_trait::async_trait;
12
13use crate::plugin::{ContextTransform, Plugin, PluginCapabilities, TransformContext};
14use crate::types::{AgentMessage, AssistantBlock};
15
16#[derive(Debug, Clone, PartialEq)]
18pub struct ToolCallIdNormalization {
19 pub messages: Vec<AgentMessage>,
20 pub renamed_count: usize,
21}
22
23pub fn normalize_tool_call_ids(messages: Vec<AgentMessage>) -> ToolCallIdNormalization {
30 let mut reserved = HashSet::new();
31 for message in &messages {
32 match message {
33 AgentMessage::Assistant { content, .. } => {
34 for block in &content.blocks {
35 if let AssistantBlock::ToolCall(call) = block {
36 reserved.insert(call.id.clone());
37 }
38 }
39 }
40 AgentMessage::ToolResult { tool_call_id, .. } => {
41 reserved.insert(tool_call_id.clone());
42 }
43 _ => {}
44 }
45 }
46
47 let mut used = HashSet::new();
48 let mut next_suffix = 1usize;
49 let mut renamed_count = 0usize;
50 let mut normalized = Vec::with_capacity(messages.len());
51 let mut input = messages.into_iter().peekable();
52
53 while let Some(mut message) = input.next() {
54 let AgentMessage::Assistant { content, .. } = &mut message else {
55 normalized.push(message);
56 continue;
57 };
58
59 let mut result_ids: HashMap<String, VecDeque<String>> = HashMap::new();
60 for block in &mut content.blocks {
61 let AssistantBlock::ToolCall(call) = block else {
62 continue;
63 };
64 let original = call.id.clone();
65 let wire_id = if used.insert(original.clone()) {
66 original.clone()
67 } else {
68 let replacement = loop {
69 let candidate = format!("clark_agent_call_{next_suffix}");
70 next_suffix += 1;
71 if reserved.insert(candidate.clone()) {
72 break candidate;
73 }
74 };
75 used.insert(replacement.clone());
76 call.id = replacement.clone();
77 renamed_count += 1;
78 replacement
79 };
80 result_ids.entry(original).or_default().push_back(wire_id);
81 }
82 normalized.push(message);
83
84 while matches!(input.peek(), Some(AgentMessage::ToolResult { .. })) {
85 let mut result = input.next().expect("peeked tool result");
86 if let AgentMessage::ToolResult { tool_call_id, .. } = &mut result {
87 if let Some(wire_id) = result_ids
88 .get_mut(tool_call_id.as_str())
89 .and_then(VecDeque::pop_front)
90 {
91 *tool_call_id = wire_id;
92 }
93 }
94 normalized.push(result);
95 }
96 }
97
98 ToolCallIdNormalization {
99 messages: normalized,
100 renamed_count,
101 }
102}
103
104#[derive(Debug, Clone, Copy, Default)]
107pub struct UniqueToolCallIds;
108
109impl Plugin for UniqueToolCallIds {
110 fn name(&self) -> &'static str {
111 "unique_tool_call_ids"
112 }
113
114 fn capabilities(&self) -> PluginCapabilities {
115 PluginCapabilities::context_transform()
116 }
117}
118
119#[async_trait]
120impl ContextTransform for UniqueToolCallIds {
121 async fn transform(
122 &self,
123 messages: Vec<AgentMessage>,
124 _cx: &TransformContext<'_>,
125 ) -> Vec<AgentMessage> {
126 let normalized = normalize_tool_call_ids(messages);
127 if normalized.renamed_count > 0 {
128 tracing::warn!(
129 renamed_count = normalized.renamed_count,
130 "normalized duplicate tool-call IDs in provider history"
131 );
132 }
133 normalized.messages
134 }
135}
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140 use crate::tool::ToolCall;
141 use crate::types::{AssistantContent, StopReason, ToolResultContent, UserContent};
142 use serde_json::json;
143
144 fn user(text: &str) -> AgentMessage {
145 AgentMessage::User {
146 content: UserContent::Text(text.to_string()),
147 timestamp: None,
148 }
149 }
150
151 fn assistant(ids: &[&str]) -> AgentMessage {
152 AgentMessage::Assistant {
153 content: AssistantContent {
154 blocks: ids
155 .iter()
156 .map(|id| {
157 AssistantBlock::ToolCall(ToolCall {
158 id: (*id).to_string(),
159 name: "shell".to_string(),
160 arguments: json!({}),
161 })
162 })
163 .collect(),
164 },
165 stop_reason: StopReason::ToolUse,
166 error_message: None,
167 timestamp: None,
168 usage: None,
169 }
170 }
171
172 fn result(id: &str) -> AgentMessage {
173 AgentMessage::ToolResult {
174 tool_call_id: id.to_string(),
175 tool_name: "shell".to_string(),
176 content: ToolResultContent::text("ok"),
177 is_error: false,
178 narration: None,
179 details: None,
180 timestamp: None,
181 }
182 }
183
184 fn call_ids(messages: &[AgentMessage]) -> Vec<&str> {
185 messages
186 .iter()
187 .flat_map(|message| match message {
188 AgentMessage::Assistant { content, .. } => content
189 .blocks
190 .iter()
191 .filter_map(|block| match block {
192 AssistantBlock::ToolCall(call) => Some(call.id.as_str()),
193 _ => None,
194 })
195 .collect(),
196 _ => Vec::new(),
197 })
198 .collect()
199 }
200
201 fn result_ids(messages: &[AgentMessage]) -> Vec<&str> {
202 messages
203 .iter()
204 .filter_map(|message| match message {
205 AgentMessage::ToolResult { tool_call_id, .. } => Some(tool_call_id.as_str()),
206 _ => None,
207 })
208 .collect()
209 }
210
211 #[test]
212 fn reused_shell_id_is_renamed_with_its_result() {
213 let normalized = normalize_tool_call_ids(vec![
214 user("first"),
215 assistant(&["shell:89"]),
216 result("shell:89"),
217 user("second"),
218 assistant(&["shell:89"]),
219 result("shell:89"),
220 ]);
221
222 assert_eq!(normalized.renamed_count, 1);
223 assert_eq!(
224 call_ids(&normalized.messages),
225 vec!["shell:89", "clark_agent_call_1"]
226 );
227 assert_eq!(
228 result_ids(&normalized.messages),
229 vec!["shell:89", "clark_agent_call_1"]
230 );
231 }
232
233 #[test]
234 fn duplicate_ids_in_one_parallel_batch_pair_by_position() {
235 let normalized = normalize_tool_call_ids(vec![
236 user("parallel"),
237 assistant(&["shell:89", "shell:89"]),
238 result("shell:89"),
239 result("shell:89"),
240 ]);
241
242 assert_eq!(
243 call_ids(&normalized.messages),
244 vec!["shell:89", "clark_agent_call_1"]
245 );
246 assert_eq!(
247 result_ids(&normalized.messages),
248 vec!["shell:89", "clark_agent_call_1"]
249 );
250 }
251
252 #[test]
253 fn normalization_is_idempotent_and_avoids_reserved_replacements() {
254 let once = normalize_tool_call_ids(vec![
255 assistant(&["shell:89"]),
256 result("shell:89"),
257 assistant(&["shell:89", "clark_agent_call_1"]),
258 result("shell:89"),
259 result("clark_agent_call_1"),
260 ]);
261 let twice = normalize_tool_call_ids(once.messages.clone());
262
263 assert_eq!(once.renamed_count, 1);
264 assert_eq!(twice.renamed_count, 0);
265 assert_eq!(once.messages, twice.messages);
266 assert_eq!(
267 call_ids(&twice.messages),
268 vec!["shell:89", "clark_agent_call_2", "clark_agent_call_1"]
269 );
270 }
271}