Skip to main content

tokio_prompt_orchestrator/
tool_call_parser.rs

1//! Tool call parsing and schema validation for LLM output.
2//!
3//! Provides [`ToolCallParser`] which scans raw LLM text for JSON tool-call
4//! objects, extracts structured [`ToolCall`] values, and validates them against
5//! registered [`ToolSchema`] definitions.
6
7use std::collections::HashMap;
8use std::fmt;
9
10// ---------------------------------------------------------------------------
11// Parameter / Schema types
12// ---------------------------------------------------------------------------
13
14/// A single parameter in a tool schema.
15#[derive(Debug, Clone)]
16pub struct ToolParameter {
17    /// Parameter name.
18    pub name: String,
19    /// Parameter type (e.g. `"string"`, `"integer"`).
20    pub param_type: String,
21    /// Whether this parameter must be present in every call.
22    pub required: bool,
23    /// Human-readable description for documentation / prompting.
24    pub description: String,
25}
26
27/// Full schema for a single tool.
28#[derive(Debug, Clone)]
29pub struct ToolSchema {
30    /// Canonical tool name matched against `"tool"` in LLM output.
31    pub name: String,
32    /// Short description of what the tool does.
33    pub description: String,
34    /// Declared parameters.
35    pub parameters: Vec<ToolParameter>,
36}
37
38// ---------------------------------------------------------------------------
39// ToolCall
40// ---------------------------------------------------------------------------
41
42/// A parsed (and optionally validated) tool invocation.
43#[derive(Debug, Clone)]
44pub struct ToolCall {
45    /// Name of the tool being invoked.
46    pub tool_name: String,
47    /// Flattened key→value argument map (all values stringified).
48    pub arguments: HashMap<String, String>,
49    /// The raw JSON text that was parsed.
50    pub raw_text: String,
51}
52
53// ---------------------------------------------------------------------------
54// ParseError
55// ---------------------------------------------------------------------------
56
57/// Errors that can occur while parsing or validating a tool call.
58#[derive(Debug, Clone, PartialEq)]
59pub enum ParseError {
60    /// The JSON was structurally invalid.
61    MalformedJson,
62    /// A required field was absent from the JSON object.
63    MissingField(String),
64    /// The tool name is not registered with this parser.
65    UnknownTool(String),
66    /// An argument value has the wrong type.
67    InvalidArgType {
68        /// Parameter name.
69        param: String,
70        /// Expected type string.
71        expected: String,
72    },
73}
74
75impl fmt::Display for ParseError {
76    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77        match self {
78            ParseError::MalformedJson => write!(f, "malformed JSON in tool call"),
79            ParseError::MissingField(field) => write!(f, "missing required field: {}", field),
80            ParseError::UnknownTool(name) => write!(f, "unknown tool: {}", name),
81            ParseError::InvalidArgType { param, expected } => {
82                write!(f, "argument '{}' expected type '{}'", param, expected)
83            }
84        }
85    }
86}
87
88// ---------------------------------------------------------------------------
89// ToolCallParser
90// ---------------------------------------------------------------------------
91
92/// Parses tool calls from LLM text and validates them against registered schemas.
93pub struct ToolCallParser {
94    /// Registered schemas keyed by tool name.
95    pub schemas: HashMap<String, ToolSchema>,
96}
97
98impl ToolCallParser {
99    /// Creates a new parser with no registered schemas.
100    pub fn new() -> Self {
101        ToolCallParser {
102            schemas: HashMap::new(),
103        }
104    }
105
106    /// Registers a tool schema so calls to that tool can be validated.
107    pub fn register_tool(&mut self, schema: ToolSchema) {
108        self.schemas.insert(schema.name.clone(), schema);
109    }
110
111    /// Finds all JSON tool-call objects in `text` and attempts to parse each.
112    ///
113    /// Matches `{"tool": "...", "arguments": {...}}` patterns.  Arguments are
114    /// flattened to `HashMap<String, String>` by converting every value to its
115    /// JSON string representation.
116    pub fn parse_tool_calls(&self, text: &str) -> Vec<Result<ToolCall, ParseError>> {
117        let candidates = self.extract_json_objects(text);
118        candidates
119            .into_iter()
120            .filter_map(|raw| {
121                // Only handle objects that look like tool calls
122                if !raw.contains("\"tool\"") && !raw.contains("'tool'") {
123                    return None;
124                }
125                Some(self.parse_single(raw))
126            })
127            .collect()
128    }
129
130    fn parse_single(&self, raw: &str) -> Result<ToolCall, ParseError> {
131        // Minimal JSON object parser – we only need "tool" and "arguments".
132        let obj = parse_json_object(raw).ok_or(ParseError::MalformedJson)?;
133
134        let tool_name = obj
135            .get("tool")
136            .cloned()
137            .ok_or_else(|| ParseError::MissingField("tool".to_string()))?;
138
139        let args_raw = obj
140            .get("arguments")
141            .cloned()
142            .ok_or_else(|| ParseError::MissingField("arguments".to_string()))?;
143
144        // Parse the arguments sub-object
145        let arguments = parse_flat_object(&args_raw).ok_or(ParseError::MalformedJson)?;
146
147        Ok(ToolCall {
148            tool_name,
149            arguments,
150            raw_text: raw.to_string(),
151        })
152    }
153
154    /// Validates a parsed [`ToolCall`] against registered schemas.
155    ///
156    /// Returns `Ok(())` if the tool is known and all required params are present.
157    pub fn validate_call(&self, call: &ToolCall) -> Result<(), ParseError> {
158        let schema = self
159            .schemas
160            .get(&call.tool_name)
161            .ok_or_else(|| ParseError::UnknownTool(call.tool_name.clone()))?;
162
163        for param in &schema.parameters {
164            if param.required && !call.arguments.contains_key(&param.name) {
165                return Err(ParseError::MissingField(param.name.clone()));
166            }
167        }
168        Ok(())
169    }
170
171    /// Extracts all top-level balanced `{…}` blocks from `text`.
172    ///
173    /// Handles nested braces correctly; does not cross string boundaries in a
174    /// fully general way but is sufficient for typical LLM JSON output.
175    pub fn extract_json_objects<'a>(&self, text: &'a str) -> Vec<&'a str> {
176        extract_balanced_braces(text)
177    }
178
179    /// Formats a tool result in the expected XML-like envelope.
180    pub fn format_tool_result(tool_name: &str, result: &str) -> String {
181        format!("<tool_result name=\"{}\">{}</tool_result>", tool_name, result)
182    }
183}
184
185impl Default for ToolCallParser {
186    fn default() -> Self {
187        Self::new()
188    }
189}
190
191// ---------------------------------------------------------------------------
192// ToolCallBuilder
193// ---------------------------------------------------------------------------
194
195/// Fluent builder for constructing [`ToolCall`] values programmatically.
196pub struct ToolCallBuilder {
197    tool_name: String,
198    arguments: HashMap<String, String>,
199}
200
201impl ToolCallBuilder {
202    /// Creates a builder for the named tool.
203    pub fn new(tool_name: &str) -> Self {
204        ToolCallBuilder {
205            tool_name: tool_name.to_string(),
206            arguments: HashMap::new(),
207        }
208    }
209
210    /// Adds a key/value argument and returns `self` for chaining.
211    pub fn arg(mut self, key: &str, value: &str) -> Self {
212        self.arguments.insert(key.to_string(), value.to_string());
213        self
214    }
215
216    /// Consumes the builder and produces a [`ToolCall`].
217    pub fn build(self) -> ToolCall {
218        // Synthesise a minimal raw JSON representation
219        let args_json: String = self
220            .arguments
221            .iter()
222            .map(|(k, v)| format!("\"{}\":\"{}\"", k, v))
223            .collect::<Vec<_>>()
224            .join(",");
225        let raw_text = format!(
226            "{{\"tool\":\"{}\",\"arguments\":{{{}}}}}",
227            self.tool_name, args_json
228        );
229        ToolCall {
230            tool_name: self.tool_name,
231            arguments: self.arguments,
232            raw_text,
233        }
234    }
235}
236
237// ---------------------------------------------------------------------------
238// Internal helpers
239// ---------------------------------------------------------------------------
240
241/// Returns slices of `text` that are balanced `{…}` blocks.
242fn extract_balanced_braces(text: &str) -> Vec<&str> {
243    let bytes = text.as_bytes();
244    let mut results = Vec::new();
245    let mut depth: usize = 0;
246    let mut start: Option<usize> = None;
247    let mut in_string = false;
248    let mut escape = false;
249    let mut i = 0;
250
251    while i < bytes.len() {
252        let b = bytes[i];
253
254        if escape {
255            escape = false;
256            i += 1;
257            continue;
258        }
259
260        if in_string {
261            match b {
262                b'\\' => escape = true,
263                b'"' => in_string = false,
264                _ => {}
265            }
266            i += 1;
267            continue;
268        }
269
270        match b {
271            b'"' => in_string = true,
272            b'{' => {
273                if depth == 0 {
274                    start = Some(i);
275                }
276                depth += 1;
277            }
278            b'}'
279                if depth > 0 => {
280                    depth -= 1;
281                    if depth == 0 {
282                        if let Some(s) = start.take() {
283                            results.push(&text[s..=i]);
284                        }
285                    }
286                }
287            _ => {}
288        }
289        i += 1;
290    }
291    results
292}
293
294/// Very small JSON object parser: returns a flat `HashMap<String, String>` where
295/// values are their raw JSON representations (string contents are unquoted).
296fn parse_json_object(s: &str) -> Option<HashMap<String, String>> {
297    let s = s.trim();
298    if !s.starts_with('{') || !s.ends_with('}') {
299        return None;
300    }
301    let inner = &s[1..s.len() - 1];
302    let mut map = HashMap::new();
303    parse_kv_pairs(inner, &mut map);
304    Some(map)
305}
306
307/// Parse key-value pairs from a JSON object body (no outer braces).
308fn parse_kv_pairs(s: &str, map: &mut HashMap<String, String>) {
309    let bytes = s.as_bytes();
310    let mut i = 0;
311
312    loop {
313        // Skip whitespace and commas
314        while i < bytes.len() && (bytes[i] == b',' || bytes[i].is_ascii_whitespace()) {
315            i += 1;
316        }
317        if i >= bytes.len() {
318            break;
319        }
320
321        // Expect a quoted key
322        if bytes[i] != b'"' {
323            break;
324        }
325        let (key, next) = match read_string(s, i) {
326            Some(v) => v,
327            None => break,
328        };
329        i = next;
330
331        // Skip whitespace and colon
332        while i < bytes.len() && (bytes[i] == b':' || bytes[i].is_ascii_whitespace()) {
333            i += 1;
334        }
335        if i >= bytes.len() {
336            break;
337        }
338
339        // Read value
340        let (value, next) = match read_value(s, i) {
341            Some(v) => v,
342            None => break,
343        };
344        i = next;
345
346        map.insert(key, value);
347    }
348}
349
350/// Reads a JSON string starting at `pos` (which must point at `"`).
351/// Returns (unescaped content, index after closing quote).
352fn read_string(s: &str, pos: usize) -> Option<(String, usize)> {
353    let bytes = s.as_bytes();
354    debug_assert_eq!(bytes[pos], b'"');
355    let mut i = pos + 1;
356    let mut result = String::new();
357    let mut escape = false;
358
359    while i < bytes.len() {
360        let b = bytes[i];
361        if escape {
362            match b {
363                b'"' => result.push('"'),
364                b'\\' => result.push('\\'),
365                b'n' => result.push('\n'),
366                b'r' => result.push('\r'),
367                b't' => result.push('\t'),
368                _ => {
369                    result.push('\\');
370                    result.push(b as char);
371                }
372            }
373            escape = false;
374        } else if b == b'\\' {
375            escape = true;
376        } else if b == b'"' {
377            return Some((result, i + 1));
378        } else {
379            result.push(b as char);
380        }
381        i += 1;
382    }
383    None
384}
385
386/// Reads a JSON value (string, object, array, or primitive) starting at `pos`.
387/// Returns (raw string representation, index after value).
388fn read_value(s: &str, pos: usize) -> Option<(String, usize)> {
389    let bytes = s.as_bytes();
390    if pos >= bytes.len() {
391        return None;
392    }
393    match bytes[pos] {
394        b'"' => {
395            let (content, next) = read_string(s, pos)?;
396            Some((content, next))
397        }
398        b'{' => {
399            let end = find_matching_close(bytes, pos, b'{', b'}')?;
400            Some((s[pos..=end].to_string(), end + 1))
401        }
402        b'[' => {
403            let end = find_matching_close(bytes, pos, b'[', b']')?;
404            Some((s[pos..=end].to_string(), end + 1))
405        }
406        _ => {
407            // Primitive: read until comma, }, ] or whitespace
408            let start = pos;
409            let mut i = pos;
410            while i < bytes.len() && bytes[i] != b',' && bytes[i] != b'}' && bytes[i] != b']' {
411                i += 1;
412            }
413            Some((s[start..i].trim().to_string(), i))
414        }
415    }
416}
417
418fn find_matching_close(bytes: &[u8], start: usize, open: u8, close: u8) -> Option<usize> {
419    let mut depth = 0usize;
420    let mut in_str = false;
421    let mut escape = false;
422    for (i, &b) in bytes.iter().enumerate().skip(start) {
423        if escape {
424            escape = false;
425            continue;
426        }
427        if in_str {
428            match b {
429                b'\\' => escape = true,
430                b'"' => in_str = false,
431                _ => {}
432            }
433            continue;
434        }
435        if b == b'"' {
436            in_str = true;
437        } else if b == open {
438            depth += 1;
439        } else if b == close {
440            depth -= 1;
441            if depth == 0 {
442                return Some(i);
443            }
444        }
445    }
446    None
447}
448
449/// Parse a JSON object body and return a flat HashMap where nested objects
450/// are kept as their raw JSON string.
451fn parse_flat_object(s: &str) -> Option<HashMap<String, String>> {
452    let s = s.trim();
453    if !s.starts_with('{') || !s.ends_with('}') {
454        return None;
455    }
456    let inner = &s[1..s.len() - 1];
457    let mut map = HashMap::new();
458    parse_kv_pairs(inner, &mut map);
459    Some(map)
460}
461
462// ---------------------------------------------------------------------------
463// Tests
464// ---------------------------------------------------------------------------
465
466#[cfg(test)]
467mod tests {
468    use super::*;
469
470    fn make_parser() -> ToolCallParser {
471        let mut p = ToolCallParser::new();
472        p.register_tool(ToolSchema {
473            name: "search".to_string(),
474            description: "Search the web".to_string(),
475            parameters: vec![
476                ToolParameter {
477                    name: "query".to_string(),
478                    param_type: "string".to_string(),
479                    required: true,
480                    description: "Search query".to_string(),
481                },
482                ToolParameter {
483                    name: "limit".to_string(),
484                    param_type: "integer".to_string(),
485                    required: false,
486                    description: "Max results".to_string(),
487                },
488            ],
489        });
490        p
491    }
492
493    #[test]
494    fn test_parse_valid_call() {
495        let parser = make_parser();
496        let text = r#"Some preamble {"tool":"search","arguments":{"query":"Rust async"}} tail"#;
497        let results = parser.parse_tool_calls(text);
498        assert_eq!(results.len(), 1);
499        let call = results[0].as_ref().unwrap();
500        assert_eq!(call.tool_name, "search");
501        assert_eq!(call.arguments.get("query").unwrap(), "Rust async");
502    }
503
504    #[test]
505    fn test_missing_required_param() {
506        let parser = make_parser();
507        let call = ToolCallBuilder::new("search").build(); // no "query"
508        let err = parser.validate_call(&call).unwrap_err();
509        assert_eq!(err, ParseError::MissingField("query".to_string()));
510    }
511
512    #[test]
513    fn test_unknown_tool() {
514        let parser = make_parser();
515        let call = ToolCallBuilder::new("nonexistent").arg("x", "y").build();
516        let err = parser.validate_call(&call).unwrap_err();
517        assert_eq!(err, ParseError::UnknownTool("nonexistent".to_string()));
518    }
519
520    #[test]
521    fn test_malformed_json_handled() {
522        let parser = make_parser();
523        // Text contains { but it is not a valid tool call JSON
524        let text = r#"{"tool": incomplete"#;
525        // extract_json_objects won't find a balanced block, so no results
526        let results = parser.parse_tool_calls(text);
527        assert!(results.is_empty());
528    }
529
530    #[test]
531    fn test_extract_json_objects_nested_braces() {
532        let parser = ToolCallParser::new();
533        let text = r#"before {"a":{"b":1}} middle {"c":2} after"#;
534        let objs = parser.extract_json_objects(text);
535        assert_eq!(objs.len(), 2);
536        assert!(objs[0].contains("\"a\""));
537        assert!(objs[1].contains("\"c\""));
538    }
539
540    #[test]
541    fn test_builder_and_format() {
542        let call = ToolCallBuilder::new("calculator")
543            .arg("expression", "2+2")
544            .build();
545        assert_eq!(call.tool_name, "calculator");
546        assert_eq!(call.arguments["expression"], "2+2");
547
548        let result = ToolCallParser::format_tool_result("calculator", "4");
549        assert_eq!(result, "<tool_result name=\"calculator\">4</tool_result>");
550    }
551}