claude_codex/providers/cursor/
tool_use_xml.rs1use std::collections::BTreeSet;
9use std::fmt;
10
11#[derive(Debug, Clone, PartialEq)]
13pub enum RecoveredCursorEvent {
14 Text(String),
15 ToolUse(RecoveredCursorToolUse),
16}
17
18#[derive(Debug, Clone, PartialEq)]
20pub struct RecoveredCursorToolUse {
21 pub id: String,
23 pub original_id: Option<String>,
25 pub name: String,
27 pub input: serde_json::Map<String, serde_json::Value>,
29}
30
31pub 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 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 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 pub fn saw_tool_use(&self) -> bool {
72 self.recovered_tool_use
73 }
74
75 pub fn push(&mut self, text: &str) -> Vec<RecoveredCursorEvent> {
77 self.buffer.push_str(text);
78 self.drain(false)
79 }
80
81 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 if self.drop_redundant_close_tag() {
96 continue;
97 }
98
99 let start = self.buffer.find("<tool_use");
101 if start.is_none() {
102 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 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 let close_start = self.buffer.find("</tool_use>");
129 if close_start.is_none() {
130 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 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 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 if let Some(ref allowed) = self.allowed_tool_names
175 && !allowed.contains(name)
176 {
177 return None;
178 }
179
180 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 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 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
248fn 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(""", "\"")
276 .replace("'", "'")
277 .replace("<", "<")
278 .replace(">", ">")
279 .replace("&", "&")
280}
281
282fn 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 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 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 let events = parser.push("</tool_use>");
413 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 let events = parser.push("some text <tool_use name=\"Read\">{\"a\":1}");
425 assert_eq!(events.len(), 1);
428 assert_eq!(events[0], RecoveredCursorEvent::Text("some text ".into()));
429
430 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 & 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); }
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 assert!(parser.flush().is_empty());
515 assert!(parser.flush().is_empty());
517 }
518}