Skip to main content

claude_codex/providers/cursor/
tool_use_xml.rs

1//! XML tool-use recovery parser.
2//!
3//! Recovers `<tool_use id="..." name="...">{"json"}</tool_use>` from text delta
4//! chunks that arrive incrementally from the upstream. Produces structured
5//! `RecoveredCursorEvent` values for the bridge to emit as Anthropic tool_use
6//! content blocks.
7
8use std::collections::BTreeSet;
9use std::fmt;
10
11/// An event recovered from a text delta chunk.
12#[derive(Debug, Clone, PartialEq)]
13pub enum RecoveredCursorEvent {
14    Text(String),
15    ToolUse(RecoveredCursorToolUse),
16}
17
18/// A recovered `<tool_use>` XML element.
19#[derive(Debug, Clone, PartialEq)]
20pub struct RecoveredCursorToolUse {
21    /// Injected unique identifier (generated by the id factory, not from XML).
22    pub id: String,
23    /// The `id` attribute from the XML element, if present.
24    pub original_id: Option<String>,
25    /// The `name` attribute from the XML element.
26    pub name: String,
27    /// The parsed JSON input object.
28    pub input: serde_json::Map<String, serde_json::Value>,
29}
30
31/// Parser that incrementally recovers `<tool_use>` XML elements from text
32/// delta chunks.
33pub struct CursorToolUseXmlParser {
34    buffer: String,
35    recovered_tool_use: bool,
36    allowed_tool_names: Option<BTreeSet<String>>,
37    id_factory: Box<dyn FnMut() -> String + Send>,
38}
39
40impl fmt::Debug for CursorToolUseXmlParser {
41    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
42        f.debug_struct("CursorToolUseXmlParser")
43            .field("buffer", &self.buffer)
44            .field("recovered_tool_use", &self.recovered_tool_use)
45            .field("allowed_tool_names", &self.allowed_tool_names)
46            .field("id_factory", &"<fn>")
47            .finish()
48    }
49}
50
51impl CursorToolUseXmlParser {
52    /// Create a parser with the default UUID-based id factory.
53    pub fn new(allowed_tool_names: Option<BTreeSet<String>>) -> Self {
54        Self::new_with_id_factory(allowed_tool_names, default_id_factory())
55    }
56
57    /// Create a parser with a custom id factory for deterministic tests.
58    pub fn new_with_id_factory(
59        allowed_tool_names: Option<BTreeSet<String>>,
60        id_factory: impl FnMut() -> String + Send + 'static,
61    ) -> Self {
62        Self {
63            buffer: String::new(),
64            recovered_tool_use: false,
65            allowed_tool_names,
66            id_factory: Box::new(id_factory),
67        }
68    }
69
70    /// Whether any tool_use has been recovered so far.
71    pub fn saw_tool_use(&self) -> bool {
72        self.recovered_tool_use
73    }
74
75    /// Push a new text delta chunk and drain any complete events.
76    pub fn push(&mut self, text: &str) -> Vec<RecoveredCursorEvent> {
77        self.buffer.push_str(text);
78        self.drain(false)
79    }
80
81    /// Flush any remaining buffered text, yielding any leftover events.
82    pub fn flush(&mut self) -> Vec<RecoveredCursorEvent> {
83        self.drain(true)
84    }
85
86    fn drain(&mut self, flush: bool) -> Vec<RecoveredCursorEvent> {
87        let mut events: Vec<RecoveredCursorEvent> = Vec::new();
88
89        loop {
90            if self.buffer.is_empty() {
91                break;
92            }
93
94            // Drop redundant close tags that can appear after recovery.
95            if self.drop_redundant_close_tag() {
96                continue;
97            }
98
99            // Look for the start of a <tool_use tag.
100            let start = self.buffer.find("<tool_use");
101            if start.is_none() {
102                // No tag start found: hold back a prefix of the buffer that
103                // could be a partial <tool_use marker, unless we are flushing.
104                let hold = if flush {
105                    0
106                } else {
107                    tool_use_prefix_suffix_length(&self.buffer)
108                };
109                let end = self.buffer.len().saturating_sub(hold);
110                let text = self.buffer[..end].to_string();
111                if !text.is_empty() {
112                    self.push_text(&mut events, &text);
113                }
114                self.buffer = self.buffer[end..].to_string();
115                break;
116            }
117
118            let start = start.unwrap();
119
120            // Emit any text before the tag.
121            if start > 0 {
122                self.push_text(&mut events, &self.buffer[..start]);
123                self.buffer = self.buffer[start..].to_string();
124                continue;
125            }
126
127            // Look for the close tag.
128            let close_start = self.buffer.find("</tool_use>");
129            if close_start.is_none() {
130                // No close tag yet; if we are flushing, emit the whole buffer
131                // as text. Otherwise hold and wait for more data.
132                if flush {
133                    self.push_text(&mut events, &self.buffer);
134                    self.buffer.clear();
135                }
136                break;
137            }
138
139            let close_start = close_start.unwrap();
140            let close_end = close_start + "</tool_use>".len();
141            let raw = self.buffer[..close_end].to_string();
142
143            // Try to parse a complete <tool_use...>...</tool_use> element.
144            if let Some(parsed) = self.parse_tool_use(&raw) {
145                self.recovered_tool_use = true;
146                events.push(RecoveredCursorEvent::ToolUse(parsed));
147            } else {
148                self.push_text(&mut events, &raw);
149            }
150
151            self.buffer = self.buffer[close_end..].to_string();
152        }
153
154        events
155    }
156
157    /// Try to parse a complete `<tool_use ...>...</tool_use>` element.
158    fn parse_tool_use(&mut self, raw: &str) -> Option<RecoveredCursorToolUse> {
159        let re = regex_lite::Regex::new(r"^<tool_use\b([^>]*)>([\s\S]*?)</tool_use>$").ok()?;
160        let caps = re.captures(raw)?;
161        let attrs_str = caps.get(1).map(|m| m.as_str()).unwrap_or("");
162        let body = caps.get(2).map(|m| m.as_str()).unwrap_or("");
163
164        let attrs = parse_xml_attributes(attrs_str);
165        let name = attrs.get("name")?;
166        let original_id = attrs.get("id").cloned();
167        let original_id = if original_id.as_deref() == Some("") {
168            None
169        } else {
170            original_id
171        };
172
173        // Check allowed tool names.
174        if let Some(ref allowed) = self.allowed_tool_names
175            && !allowed.contains(name)
176        {
177            return None;
178        }
179
180        // Parse the JSON body.
181        let trimmed = body.trim();
182        let input_value: serde_json::Value = if trimmed.is_empty() {
183            serde_json::Value::Object(serde_json::Map::new())
184        } else {
185            match serde_json::from_str(trimmed) {
186                Ok(v) => v,
187                Err(_) => return None,
188            }
189        };
190
191        let input = match input_value {
192            serde_json::Value::Object(map) => map,
193            _ => return None,
194        };
195
196        let id = (self.id_factory)();
197
198        Some(RecoveredCursorToolUse {
199            id,
200            original_id,
201            name: name.clone(),
202            input,
203        })
204    }
205
206    /// Drop a redundant `</tool_use>` close tag that follows an already-
207    /// recovered tool_use (the XML stream may contain duplicate close tags
208    /// after recovery).
209    fn drop_redundant_close_tag(&mut self) -> bool {
210        if !self.recovered_tool_use {
211            return false;
212        }
213        let trimmed_start = self.buffer.len() - self.buffer.trim_start().len();
214        let slice = &self.buffer[trimmed_start..];
215        if let Some(rest) = slice.strip_prefix("</tool_use>") {
216            let prefix = &self.buffer[..trimmed_start];
217            self.buffer = format!("{prefix}{rest}");
218            true
219        } else if let Some(rest) = slice.strip_prefix("</tool_use") {
220            let prefix = &self.buffer[..trimmed_start];
221            self.buffer = format!("{prefix}{rest}");
222            true
223        } else {
224            false
225        }
226    }
227
228    /// Push text into events, filtering empty text and trailing whitespace
229    /// after a recovered tool_use.
230    fn push_text(&self, events: &mut Vec<RecoveredCursorEvent>, text: &str) {
231        if text.is_empty() {
232            return;
233        }
234        if self.recovered_tool_use && text.trim().is_empty() {
235            return;
236        }
237        events.push(RecoveredCursorEvent::Text(text.to_string()));
238    }
239}
240
241fn default_id_factory() -> Box<dyn FnMut() -> String + Send> {
242    Box::new(|| {
243        let id = uuid::Uuid::new_v4().to_string().replace('-', "");
244        format!("call_cursor_{id}")
245    })
246}
247
248// ---------------------------------------------------------------------------
249// XML attribute parsing
250// ---------------------------------------------------------------------------
251
252fn parse_xml_attributes(source: &str) -> std::collections::HashMap<String, String> {
253    let mut attrs = std::collections::HashMap::new();
254    let re = regex_lite::Regex::new(r#"([A-Za-z_][\w:.-]*)\s*=\s*(?:"([^"]*)"|'([^']*)')"#).ok();
255    let re = match re {
256        Some(r) => r,
257        None => return attrs,
258    };
259    for cap in re.captures_iter(source) {
260        let key = cap.get(1).map(|m| m.as_str()).unwrap_or("");
261        let value = cap
262            .get(2)
263            .or_else(|| cap.get(3))
264            .map(|m| m.as_str())
265            .unwrap_or("");
266        if !key.is_empty() {
267            attrs.insert(key.to_string(), decode_xml_attribute(value));
268        }
269    }
270    attrs
271}
272
273fn decode_xml_attribute(value: &str) -> String {
274    value
275        .replace("&quot;", "\"")
276        .replace("&apos;", "'")
277        .replace("&lt;", "<")
278        .replace("&gt;", ">")
279        .replace("&amp;", "&")
280}
281
282/// Compute how many characters at the end of `value` could be a partial
283/// `<tool_use` marker. Returns a length in [0, "<tool_use".len() - 1].
284fn tool_use_prefix_suffix_length(value: &str) -> usize {
285    let marker = "<tool_use";
286    let max = marker.len().saturating_sub(1).min(value.len());
287    for len in (1..=max).rev() {
288        if value.ends_with(&marker[..len]) {
289            return len;
290        }
291    }
292    0
293}
294
295#[cfg(test)]
296mod tests {
297    use super::*;
298
299    fn test_id_factory() -> Box<dyn FnMut() -> String + Send> {
300        let mut counter = 0u64;
301        Box::new(move || {
302            counter += 1;
303            format!("call_cursor_test_{counter}")
304        })
305    }
306
307    #[test]
308    fn recovers_complete_tool_use_xml() {
309        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
310            Some(["Read".to_string()].into_iter().collect()),
311            test_id_factory(),
312        );
313        let input = r#"before <tool_use id="x" name="Read">{"file_path":"a"}</tool_use> after"#;
314        let events = parser.push(input);
315        assert_eq!(events.len(), 3);
316        assert_eq!(events[0], RecoveredCursorEvent::Text("before ".into()));
317        assert!(matches!(&events[1], RecoveredCursorEvent::ToolUse(tool) if tool.name == "Read"));
318        assert_eq!(events[2], RecoveredCursorEvent::Text(" after".into()));
319        assert!(parser.saw_tool_use());
320        assert!(parser.flush().is_empty());
321    }
322
323    #[test]
324    fn recovers_split_tool_use_xml() {
325        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
326            Some(["Read".to_string()].into_iter().collect()),
327            test_id_factory(),
328        );
329        assert_eq!(
330            parser.push("before <tool_"),
331            vec![RecoveredCursorEvent::Text("before ".into())]
332        );
333        let events = parser.push(r#"use id="x" name="Read">{"file_path":"a"}</tool_use> after"#);
334        assert!(matches!(&events[0], RecoveredCursorEvent::ToolUse(tool) if tool.name == "Read"));
335        assert_eq!(events[1], RecoveredCursorEvent::Text(" after".into()));
336        assert!(parser.flush().is_empty());
337    }
338
339    #[test]
340    fn recovers_tool_use_without_original_id() {
341        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
342            Some(["Read".to_string()].into_iter().collect()),
343            test_id_factory(),
344        );
345        let events = parser.push(r#"<tool_use name="Read">{"x":1}</tool_use>"#);
346        assert_eq!(events.len(), 1);
347        if let RecoveredCursorEvent::ToolUse(tool) = &events[0] {
348            assert_eq!(tool.original_id, None);
349            assert_eq!(tool.name, "Read");
350            assert_eq!(tool.input.get("x").and_then(|v| v.as_i64()), Some(1));
351        } else {
352            panic!("expected ToolUse");
353        }
354    }
355
356    #[test]
357    fn recovers_tool_use_with_original_id() {
358        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
359            Some(["Write".to_string()].into_iter().collect()),
360            test_id_factory(),
361        );
362        let events = parser.push(
363            r#"<tool_use id="orig-1" name="Write">{"file_path":"b","content":"hi"}</tool_use>"#,
364        );
365        assert_eq!(events.len(), 1);
366        if let RecoveredCursorEvent::ToolUse(tool) = &events[0] {
367            assert_eq!(tool.original_id.as_deref(), Some("orig-1"));
368            assert_eq!(tool.name, "Write");
369        } else {
370            panic!("expected ToolUse");
371        }
372    }
373
374    #[test]
375    fn filters_disallowed_tool_names() {
376        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
377            Some(["Read".to_string()].into_iter().collect()),
378            test_id_factory(),
379        );
380        // Bash is not in allowed set, so the text is emitted as-is.
381        let events = parser.push(r#"<tool_use id="x" name="Bash">{"command":"pwd"}</tool_use>"#);
382        assert_eq!(events.len(), 1);
383        assert!(matches!(&events[0], RecoveredCursorEvent::Text(t)
384            if t.contains("Bash")));
385        assert!(!parser.saw_tool_use());
386    }
387
388    #[test]
389    fn invalid_json_fallback_to_text() {
390        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
391            Some(["Read".to_string()].into_iter().collect()),
392            test_id_factory(),
393        );
394        let events = parser.push(r#"<tool_use id="x" name="Read">{invalid}</tool_use>"#);
395        assert_eq!(events.len(), 1);
396        assert!(matches!(&events[0], RecoveredCursorEvent::Text(_)));
397        assert!(!parser.saw_tool_use());
398    }
399
400    #[test]
401    fn redundant_close_tag_after_recovery() {
402        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
403            Some(["Read".to_string()].into_iter().collect()),
404            test_id_factory(),
405        );
406        // First recover a tool_use.
407        let events = parser.push(r#"<tool_use name="Read">{}</tool_use>"#);
408        assert_eq!(events.len(), 1);
409        assert!(parser.saw_tool_use());
410
411        // Push a redundant close tag (common in streaming XML).
412        let events = parser.push("</tool_use>");
413        // Should be silently dropped, no text events.
414        assert!(events.is_empty());
415    }
416
417    #[test]
418    fn flush_returns_remaining_text() {
419        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
420            Some(["Read".to_string()].into_iter().collect()),
421            test_id_factory(),
422        );
423        // Push an incomplete tag.
424        let events = parser.push("some text <tool_use name=\"Read\">{\"a\":1}");
425        // No close tag yet, so only text before partial tag should emit.
426        // But the parser holds back the partial prefix, so "some text " emits.
427        assert_eq!(events.len(), 1);
428        assert_eq!(events[0], RecoveredCursorEvent::Text("some text ".into()));
429
430        // Flush should emit the remaining buffer as text since the tag
431        // could not be completed.
432        let events = parser.flush();
433        assert!(!events.is_empty());
434        assert!(matches!(&events[0], RecoveredCursorEvent::Text(t)
435            if t.contains("<tool_use") || t.contains("name")));
436    }
437
438    #[test]
439    fn empty_push_does_nothing() {
440        let mut parser = CursorToolUseXmlParser::new_with_id_factory(None, test_id_factory());
441        let events = parser.push("");
442        assert!(events.is_empty());
443        assert!(parser.flush().is_empty());
444    }
445
446    #[test]
447    fn text_only_chunks() {
448        let mut parser = CursorToolUseXmlParser::new_with_id_factory(None, test_id_factory());
449        assert_eq!(
450            parser.push("hello world"),
451            vec![RecoveredCursorEvent::Text("hello world".into())]
452        );
453        assert!(parser.flush().is_empty());
454    }
455
456    #[test]
457    fn multiple_tool_uses_in_one_push() {
458        let mut parser = CursorToolUseXmlParser::new_with_id_factory(
459            Some(
460                ["Read".to_string(), "Write".to_string()]
461                    .into_iter()
462                    .collect(),
463            ),
464            test_id_factory(),
465        );
466        let events = parser.push(
467            r#"<tool_use name="Read">{"file_path":"a"}</tool_use><tool_use name="Write">{"file_path":"b","content":"c"}</tool_use>"#,
468        );
469        assert_eq!(events.len(), 2);
470        assert!(matches!(&events[0], RecoveredCursorEvent::ToolUse(t) if t.name == "Read"));
471        assert!(matches!(&events[1], RecoveredCursorEvent::ToolUse(t) if t.name == "Write"));
472    }
473
474    #[test]
475    fn parse_xml_attributes_extracts_name_and_id() {
476        let attrs = parse_xml_attributes(r#"id="abc" name="Read""#);
477        assert_eq!(attrs.get("name").map(|s| s.as_str()), Some("Read"));
478        assert_eq!(attrs.get("id").map(|s| s.as_str()), Some("abc"));
479    }
480
481    #[test]
482    fn parse_xml_handles_single_quotes() {
483        let attrs = parse_xml_attributes(r#"name='Read'"#);
484        assert_eq!(attrs.get("name").map(|s| s.as_str()), Some("Read"));
485    }
486
487    #[test]
488    fn parse_xml_decodes_entities() {
489        let attrs = parse_xml_attributes(r#"name="Read &amp; Write""#);
490        assert_eq!(attrs.get("name").map(|s| s.as_str()), Some("Read & Write"));
491    }
492
493    #[test]
494    fn tool_use_prefix_suffix_identifies_partial_tag() {
495        assert_eq!(tool_use_prefix_suffix_length("<tool_"), 6);
496        assert_eq!(tool_use_prefix_suffix_length("<tool_us"), 8);
497        assert_eq!(tool_use_prefix_suffix_length("hello <tool_"), 6);
498        assert_eq!(tool_use_prefix_suffix_length("hello"), 0);
499        assert_eq!(tool_use_prefix_suffix_length("<"), 1); // "<" is a prefix of "<tool_use"
500    }
501
502    #[test]
503    fn saw_tool_use_starts_false() {
504        let parser = CursorToolUseXmlParser::new_with_id_factory(None, test_id_factory());
505        assert!(!parser.saw_tool_use());
506    }
507
508    #[test]
509    fn empty_input_after_flush() {
510        let mut parser = CursorToolUseXmlParser::new_with_id_factory(None, test_id_factory());
511        let events = parser.push("hello");
512        assert_eq!(events.len(), 1);
513        // Flush with no partial tag should be empty.
514        assert!(parser.flush().is_empty());
515        // Another flush should also be empty.
516        assert!(parser.flush().is_empty());
517    }
518}