Skip to main content

clark_agent/
history.rs

1//! Provider-facing history invariants.
2//!
3//! Durable transcripts may contain tool-call identifiers chosen by different
4//! provider turns. Some OpenAI-compatible APIs require every identifier in a
5//! request to be globally unique, even when each call/result pair is otherwise
6//! valid. This module repairs that wire-facing invariant without rewriting the
7//! persisted transcript.
8
9use 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/// Result of normalizing one provider-visible history slice.
17#[derive(Debug, Clone, PartialEq)]
18pub struct ToolCallIdNormalization {
19    pub messages: Vec<AgentMessage>,
20    pub renamed_count: usize,
21}
22
23/// Make tool-call IDs unique across a provider-visible history.
24///
25/// The first occurrence keeps its original identifier. Later occurrences are
26/// renamed deterministically, together with positionally corresponding tool
27/// results immediately following that assistant turn. Existing identifiers are
28/// reserved so generated replacements cannot collide with transcript data.
29pub 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/// Stateless context transform that applies [`normalize_tool_call_ids`] before
105/// each provider request.
106#[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}