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 &content[..content.len().min(100)]
214 );
215 events.push(ParserEvent::Error(format!(
216 "Malformed tool_call content: {}",
217 &content[..content.len().min(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
280#[cfg(test)]
283mod tests {
284 use super::*;
285
286 #[test]
287 fn test_simple_tool_call() {
288 let input = r#"<tool_call>{"name":"read","arguments":{"path":"/tmp/x"}}</tool_call>"#;
289 let events = parse_tool_calls(input);
290 let names: Vec<&str> = events
291 .iter()
292 .filter_map(|e| {
293 if let ParserEvent::ToolCall { name, .. } = e {
294 Some(name.as_str())
295 } else {
296 None
297 }
298 })
299 .collect();
300 assert_eq!(names, vec!["read"]);
301 }
302
303 #[test]
304 fn test_mixed_text_and_tool_calls() {
305 let input = concat!(
306 "Let me check. ",
307 r#"<tool_call>{"name":"read","arguments":{"path":"x"}}</tool_call>"#,
308 " Found it."
309 );
310 let events = parse_tool_calls(input);
311 assert_eq!(events.len(), 3);
312 assert!(matches!(&events[0], ParserEvent::Text(t) if t == "Let me check. "));
313 assert!(matches!(&events[1], ParserEvent::ToolCall { name, .. } if name == "read"));
314 assert!(matches!(&events[2], ParserEvent::Text(t) if t == " Found it."));
315 }
316
317 #[test]
318 fn test_streaming_chunks() {
319 let mut parser = StreamingToolCallParser::new();
320 let chunks = vec![
321 "Hello. ",
322 "<tool_call",
323 ">",
324 r#"{"name":"search","arguments":{"q":"hello"}}"#,
325 "</tool_call>",
326 " Done.",
327 ];
328 let mut all_events = Vec::new();
329 for chunk in chunks {
330 all_events.extend(parser.feed(chunk));
331 }
332 all_events.extend(parser.end());
333 let texts: Vec<&str> = all_events
334 .iter()
335 .filter_map(|e| {
336 if let ParserEvent::Text(t) = e {
337 Some(t.as_str())
338 } else {
339 None
340 }
341 })
342 .collect();
343 assert!(texts.contains(&"Hello. "));
344 assert!(texts.contains(&" Done."));
345 }
346
347 #[test]
348 fn test_multiple_tool_calls() {
349 let input = concat!(
350 r#"<tool_call>{"name":"read","arguments":{"path":"a"}}</tool_call>"#,
351 r#"<tool_call>{"name":"read","arguments":{"path":"b"}}</tool_call>"#
352 );
353 let events = parse_tool_calls(input);
354 let tc_count = events
355 .iter()
356 .filter(|e| matches!(e, ParserEvent::ToolCall { .. }))
357 .count();
358 assert_eq!(tc_count, 2);
359 }
360
361 #[test]
362 fn test_malformed_json_fallback() {
363 let input = r#"<tool_call>{"name": "shell", "arguments": {"cmd": "ls"}</tool_call>"#;
365 let events = parse_tool_calls(input);
366 let has_error = events.iter().any(|e| matches!(e, ParserEvent::Error(_)));
367 assert!(
369 has_error
370 || events
371 .iter()
372 .any(|e| matches!(e, ParserEvent::ToolCall { .. }))
373 );
374 }
375
376 #[test]
377 fn test_callback_on_complete() {
378 let mut parser = StreamingToolCallParser::new();
379 let called = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
380 let c = called.clone();
381 parser.on_tool_call(move |_tc| {
382 c.store(true, std::sync::atomic::Ordering::SeqCst);
383 });
384 parser.feed(r#"<tool_call>{"name":"test","arguments":{}}</tool_call>"#);
385 assert!(called.load(std::sync::atomic::Ordering::SeqCst));
386 }
387
388 #[test]
389 fn test_plain_text() {
390 let mut parser = StreamingToolCallParser::new();
391 parser.feed("Just text");
392 let events = parser.end();
393 assert!(events
394 .iter()
395 .any(|e| matches!(e, ParserEvent::Text(t) if t == "Just text")));
396 }
397
398 #[test]
399 fn test_reset() {
400 let mut parser = StreamingToolCallParser::new();
401 parser.feed("<tool_call>{\"name\":\"x\"");
402 parser.reset();
403 let events = parser.feed(r#"<tool_call>{"name":"y","arguments":{}}</tool_call>"#);
404 assert!(!events.is_empty());
405 let names: Vec<&str> = events
406 .iter()
407 .filter_map(|e| {
408 if let ParserEvent::ToolCall { name, .. } = e {
409 Some(name.as_str())
410 } else {
411 None
412 }
413 })
414 .collect();
415 assert_eq!(names, vec!["y"]);
416 }
417}