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                &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        // 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// ── Tests ─────────────────────────────────────────────────────────────────
281
282#[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        // Unclosed JSON — parser should emit Error, not panic
364        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        // Either an error or a best-effort parse
368        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}