1use serde_json::Value;
9use tracing::debug;
10
11#[derive(Debug, Clone)]
15pub struct ToolCall {
16 pub id: String,
17 pub name: String,
18 pub arguments: String,
19}
20
21#[derive(Debug, Clone)]
23pub enum ParserEvent {
24 Text(String),
26 ToolCall { name: String, args: String },
28 Error(String),
30 End,
32}
33
34struct ExtractedToolCall {
36 id: String,
37 name: String,
38 args: String,
39}
40
41type ToolCallback = Box<dyn Fn(ToolCall) + Send + Sync>;
43
44#[derive(Debug, Clone, Copy, PartialEq)]
47enum State {
48 Outside,
49 InsideOpenTag,
50 InsideContent,
51 InsideNestedTag,
52}
53
54pub struct StreamingToolCallParser {
56 state: State,
57 buffer: String,
58 tag_buffer: String,
59 nested_depth: usize,
60 in_tool_call: bool,
61 position: usize,
62 on_tool_call: Option<ToolCallback>,
63 call_counter: usize,
64}
65
66impl StreamingToolCallParser {
67 pub fn new() -> Self {
68 Self {
69 state: State::Outside,
70 buffer: String::new(),
71 tag_buffer: String::new(),
72 nested_depth: 0,
73 in_tool_call: false,
74 position: 0,
75 on_tool_call: None,
76 call_counter: 0,
77 }
78 }
79
80 pub fn on_tool_call<F>(&mut self, callback: F)
82 where
83 F: Fn(ToolCall) + Send + Sync + 'static,
84 {
85 self.on_tool_call = Some(Box::new(callback));
86 }
87
88 pub fn feed(&mut self, chunk: &str) -> Vec<ParserEvent> {
90 let mut events = Vec::new();
91
92 for ch in chunk.chars() {
93 self.position += 1;
94 match self.state {
95 State::Outside => {
96 if ch == '<' {
97 if !self.buffer.is_empty() {
98 let text = std::mem::take(&mut self.buffer);
99 events.push(ParserEvent::Text(text));
100 }
101 self.state = State::InsideOpenTag;
102 self.tag_buffer.clear();
103 } else {
104 self.buffer.push(ch);
105 }
106 }
107
108 State::InsideOpenTag => {
109 if ch == '>' {
110 let tag = self.tag_buffer.trim().to_string();
111 if tag.starts_with("tool_call") {
112 self.in_tool_call = true;
113 self.state = State::InsideContent;
114 if tag.ends_with('/') || tag.starts_with("tool_call/") {
115 self.finish_tool_call(&mut events);
116 }
117 } else if tag.starts_with('/') && tag[1..].trim() == "tool_call" {
118 self.finish_tool_call(&mut events);
119 } else if self.in_tool_call {
120 self.nested_depth += 1;
121 self.buffer.push('<');
122 self.buffer.push_str(&tag);
123 self.buffer.push('>');
124 self.state = State::InsideNestedTag;
125 } else {
126 self.buffer.push('<');
127 self.buffer.push_str(&tag);
128 self.buffer.push('>');
129 self.state = State::Outside;
130 }
131 } else {
132 self.tag_buffer.push(ch);
133 }
134 }
135
136 State::InsideContent => {
137 if ch == '<' {
138 self.state = State::InsideOpenTag;
139 self.tag_buffer.clear();
140 } else {
141 self.buffer.push(ch);
142 }
143 }
144
145 State::InsideNestedTag => {
146 if ch == '>' {
147 self.buffer.push('>');
148 self.state = State::InsideContent;
149 } else {
150 self.buffer.push(ch);
151 }
152 }
153 }
154 }
155
156 events
157 }
158
159 pub fn end(&mut self) -> Vec<ParserEvent> {
161 let mut events = Vec::new();
162
163 if self.in_tool_call && !self.buffer.is_empty() {
164 let content = std::mem::take(&mut self.buffer);
165 if let Some(tc) = self.parse_json_content(&content) {
166 if let Some(ref cb) = self.on_tool_call {
167 cb(ToolCall {
168 id: tc.id.clone(),
169 name: tc.name.clone(),
170 arguments: tc.args.clone(),
171 });
172 }
173 events.push(ParserEvent::ToolCall {
174 name: tc.name,
175 args: tc.args,
176 });
177 } else {
178 events.push(ParserEvent::Text(format!(
179 "<tool_call>{}</tool_call>",
180 content
181 )));
182 }
183 } else if !self.buffer.is_empty() {
184 events.push(ParserEvent::Text(std::mem::take(&mut self.buffer)));
185 }
186
187 self.in_tool_call = false;
188 self.state = State::Outside;
189 self.nested_depth = 0;
190 events.push(ParserEvent::End);
191 events
192 }
193
194 fn finish_tool_call(&mut self, events: &mut Vec<ParserEvent>) {
195 self.in_tool_call = false;
196 let content = std::mem::take(&mut self.buffer).trim().to_string();
197
198 if let Some(tc) = self.parse_json_content(&content) {
199 if let Some(ref cb) = self.on_tool_call {
200 cb(ToolCall {
201 id: tc.id.clone(),
202 name: tc.name.clone(),
203 arguments: tc.args.clone(),
204 });
205 }
206 events.push(ParserEvent::ToolCall {
207 name: tc.name,
208 args: tc.args,
209 });
210 } else {
211 debug!(
212 "unparseable tool_call: {:?}",
213 crate::text::truncate_chars(&content, 100)
214 );
215 events.push(ParserEvent::Error(format!(
216 "Malformed tool_call content: {}",
217 crate::text::truncate_chars(&content, 100)
218 )));
219 }
220
221 self.state = State::Outside;
222 self.nested_depth = 0;
223 }
224
225 fn parse_json_content(&mut self, content: &str) -> Option<ExtractedToolCall> {
226 if let Ok(val) = serde_json::from_str::<Value>(content) {
228 let name = val
229 .get("name")
230 .or_else(|| val.get("function"))
231 .and_then(|v| v.as_str())
232 .unwrap_or("unknown")
233 .to_string();
234 let args = val
235 .get("arguments")
236 .or_else(|| val.get("input"))
237 .and_then(|v| {
238 if v.is_string() {
239 v.as_str().map(|s| s.to_string())
240 } else {
241 Some(v.to_string())
242 }
243 })
244 .unwrap_or_else(|| "{}".to_string());
245 self.call_counter += 1;
246 return Some(ExtractedToolCall {
247 id: format!("tool_{}", self.call_counter),
248 name,
249 args,
250 });
251 }
252 None
253 }
254
255 pub fn reset(&mut self) {
256 self.state = State::Outside;
257 self.buffer.clear();
258 self.tag_buffer.clear();
259 self.nested_depth = 0;
260 self.in_tool_call = false;
261 self.position = 0;
262 }
263}
264
265impl Default for StreamingToolCallParser {
266 fn default() -> Self {
267 Self::new()
268 }
269}
270
271pub fn parse_tool_calls(content: &str) -> Vec<ParserEvent> {
273 let mut parser = StreamingToolCallParser::new();
274 let mut events = parser.feed(content);
275 events.extend(parser.end());
276 events.retain(|e| !matches!(e, ParserEvent::End));
277 events
278}
279
280pub fn recover_tool_calls(content: &str) -> (String, Vec<ToolCall>) {
288 let mut text = String::new();
289 let mut calls = Vec::new();
290
291 for event in parse_tool_calls(content) {
292 match event {
293 ParserEvent::Text(t) => text.push_str(&t),
294 ParserEvent::ToolCall { name, args } => calls.push(ToolCall {
295 id: format!("xml_{}", calls.len()),
296 name,
297 arguments: args,
298 }),
299 ParserEvent::Error(e) => debug!("streaming parser: {}", e),
300 ParserEvent::End => {}
301 }
302 }
303
304 (text, calls)
305}
306
307#[cfg(test)]
310mod tests {
311 use super::*;
312
313 #[test]
314 fn test_simple_tool_call() {
315 let input = r#"<tool_call>{"name":"read","arguments":{"path":"/tmp/x"}}</tool_call>"#;
316 let events = parse_tool_calls(input);
317 let names: Vec<&str> = events
318 .iter()
319 .filter_map(|e| {
320 if let ParserEvent::ToolCall { name, .. } = e {
321 Some(name.as_str())
322 } else {
323 None
324 }
325 })
326 .collect();
327 assert_eq!(names, vec!["read"]);
328 }
329
330 #[test]
331 fn test_mixed_text_and_tool_calls() {
332 let input = concat!(
333 "Let me check. ",
334 r#"<tool_call>{"name":"read","arguments":{"path":"x"}}</tool_call>"#,
335 " Found it."
336 );
337 let events = parse_tool_calls(input);
338 assert_eq!(events.len(), 3);
339 assert!(matches!(&events[0], ParserEvent::Text(t) if t == "Let me check. "));
340 assert!(matches!(&events[1], ParserEvent::ToolCall { name, .. } if name == "read"));
341 assert!(matches!(&events[2], ParserEvent::Text(t) if t == " Found it."));
342 }
343
344 #[test]
345 fn test_streaming_chunks() {
346 let mut parser = StreamingToolCallParser::new();
347 let chunks = vec![
348 "Hello. ",
349 "<tool_call",
350 ">",
351 r#"{"name":"search","arguments":{"q":"hello"}}"#,
352 "</tool_call>",
353 " Done.",
354 ];
355 let mut all_events = Vec::new();
356 for chunk in chunks {
357 all_events.extend(parser.feed(chunk));
358 }
359 all_events.extend(parser.end());
360 let texts: Vec<&str> = all_events
361 .iter()
362 .filter_map(|e| {
363 if let ParserEvent::Text(t) = e {
364 Some(t.as_str())
365 } else {
366 None
367 }
368 })
369 .collect();
370 assert!(texts.contains(&"Hello. "));
371 assert!(texts.contains(&" Done."));
372 }
373
374 #[test]
375 fn test_multiple_tool_calls() {
376 let input = concat!(
377 r#"<tool_call>{"name":"read","arguments":{"path":"a"}}</tool_call>"#,
378 r#"<tool_call>{"name":"read","arguments":{"path":"b"}}</tool_call>"#
379 );
380 let events = parse_tool_calls(input);
381 let tc_count = events
382 .iter()
383 .filter(|e| matches!(e, ParserEvent::ToolCall { .. }))
384 .count();
385 assert_eq!(tc_count, 2);
386 }
387
388 #[test]
389 fn test_malformed_json_fallback() {
390 let input = r#"<tool_call>{"name": "shell", "arguments": {"cmd": "ls"}</tool_call>"#;
392 let events = parse_tool_calls(input);
393 let has_error = events.iter().any(|e| matches!(e, ParserEvent::Error(_)));
394 assert!(
396 has_error
397 || events
398 .iter()
399 .any(|e| matches!(e, ParserEvent::ToolCall { .. }))
400 );
401 }
402
403 #[test]
404 fn test_callback_on_complete() {
405 let mut parser = StreamingToolCallParser::new();
406 let called = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
407 let c = called.clone();
408 parser.on_tool_call(move |_tc| {
409 c.store(true, std::sync::atomic::Ordering::SeqCst);
410 });
411 parser.feed(r#"<tool_call>{"name":"test","arguments":{}}</tool_call>"#);
412 assert!(called.load(std::sync::atomic::Ordering::SeqCst));
413 }
414
415 #[test]
416 fn test_plain_text() {
417 let mut parser = StreamingToolCallParser::new();
418 parser.feed("Just text");
419 let events = parser.end();
420 assert!(events
421 .iter()
422 .any(|e| matches!(e, ParserEvent::Text(t) if t == "Just text")));
423 }
424
425 #[test]
426 fn test_reset() {
427 let mut parser = StreamingToolCallParser::new();
428 parser.feed("<tool_call>{\"name\":\"x\"");
429 parser.reset();
430 let events = parser.feed(r#"<tool_call>{"name":"y","arguments":{}}</tool_call>"#);
431 assert!(!events.is_empty());
432 let names: Vec<&str> = events
433 .iter()
434 .filter_map(|e| {
435 if let ParserEvent::ToolCall { name, .. } = e {
436 Some(name.as_str())
437 } else {
438 None
439 }
440 })
441 .collect();
442 assert_eq!(names, vec!["y"]);
443 }
444
445 #[test]
446 fn recover_splits_prose_from_tool_calls() {
447 let input = concat!(
448 "Let me check that.",
449 r#"<tool_call>{"name":"read","arguments":{"path":"/tmp/x"}}</tool_call>"#,
450 "Done."
451 );
452 let (text, calls) = recover_tool_calls(input);
453
454 assert_eq!(calls.len(), 1);
455 assert_eq!(calls[0].name, "read");
456 assert_eq!(calls[0].id, "xml_0");
457 assert!(text.contains("Let me check that."));
458 assert!(text.contains("Done."));
459 assert!(!text.contains("tool_call"));
460 }
461
462 #[test]
463 fn recover_ids_are_unique_per_call() {
464 let input = concat!(
465 r#"<tool_call>{"name":"a","arguments":{}}</tool_call>"#,
466 r#"<tool_call>{"name":"b","arguments":{}}</tool_call>"#
467 );
468 let (_, calls) = recover_tool_calls(input);
469
470 assert_eq!(calls.len(), 2);
471 assert_eq!(calls[0].id, "xml_0");
472 assert_eq!(calls[1].id, "xml_1");
473 }
474
475 #[test]
476 fn recover_returns_no_calls_for_plain_text() {
477 let (text, calls) = recover_tool_calls("just a normal answer");
478 assert!(calls.is_empty());
479 assert_eq!(text, "just a normal answer");
480 }
481}