Skip to main content

apollo/
streaming_parser.rs

1//! Streaming tool-call parser — state-machine extraction from partial LLM output.
2//!
3//! Detects `<tool_call>` blocks incrementally as tokens arrive, firing callbacks
4//! as soon as `</tool_call>` is seen. Handles malformed tags and unclosed JSON tolerantly.
5//!
6//! Ported from hermes-rs parser.rs philosophy.
7
8use serde_json::Value;
9use tracing::debug;
10
11// ── Types ─────────────────────────────────────────────────────────────────
12
13/// A detected tool call from streaming output
14#[derive(Debug, Clone)]
15pub struct ToolCall {
16    pub id: String,
17    pub name: String,
18    pub arguments: String,
19}
20
21/// Events emitted by the streaming parser
22#[derive(Debug, Clone)]
23pub enum ParserEvent {
24    /// Text content (between tags)
25    Text(String),
26    /// A complete tool call detected
27    ToolCall { name: String, args: String },
28    /// Parser encountered malformed input
29    Error(String),
30    /// Stream ended
31    End,
32}
33
34/// Internal extracted tool call
35struct ExtractedToolCall {
36    id: String,
37    name: String,
38    args: String,
39}
40
41/// Callback for early tool call detection
42type ToolCallback = Box<dyn Fn(ToolCall) + Send + Sync>;
43
44// ── State machine ─────────────────────────────────────────────────────────
45
46#[derive(Debug, Clone, Copy, PartialEq)]
47enum State {
48    Outside,
49    InsideOpenTag,
50    InsideContent,
51    InsideNestedTag,
52}
53
54/// State-machine streaming parser for XML tool calls
55pub 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    /// Register a callback fired when a complete tool call is parsed
81    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    /// Feed a chunk of streaming output. Returns emitted events.
89    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    /// Flush remaining content. Call when stream ends.
160    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        // ponytail: try full JSON parse. If fails, return None (caller emits Error).
227        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
271/// Convenience: parse a full (non-streaming) response containing tool_call XML.
272pub 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/// Split a completed response into prose and the tool calls it embedded as
281/// `<tool_call>` XML.
282///
283/// Providers without native tool calling emit calls inline in the text. Running
284/// the state machine over the finished body recovers them so callers can treat
285/// them like native tool calls. Returns the text with the tool_call blocks
286/// removed, plus the recovered calls in the order they appeared.
287pub 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// ── Tests ─────────────────────────────────────────────────────────────────
308
309#[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        // Unclosed JSON — parser should emit Error, not panic
391        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        // Either an error or a best-effort parse
395        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}