tokio_prompt_orchestrator/
tool_call_parser.rs1use std::collections::HashMap;
8use std::fmt;
9
10#[derive(Debug, Clone)]
16pub struct ToolParameter {
17 pub name: String,
19 pub param_type: String,
21 pub required: bool,
23 pub description: String,
25}
26
27#[derive(Debug, Clone)]
29pub struct ToolSchema {
30 pub name: String,
32 pub description: String,
34 pub parameters: Vec<ToolParameter>,
36}
37
38#[derive(Debug, Clone)]
44pub struct ToolCall {
45 pub tool_name: String,
47 pub arguments: HashMap<String, String>,
49 pub raw_text: String,
51}
52
53#[derive(Debug, Clone, PartialEq)]
59pub enum ParseError {
60 MalformedJson,
62 MissingField(String),
64 UnknownTool(String),
66 InvalidArgType {
68 param: String,
70 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
88pub struct ToolCallParser {
94 pub schemas: HashMap<String, ToolSchema>,
96}
97
98impl ToolCallParser {
99 pub fn new() -> Self {
101 ToolCallParser {
102 schemas: HashMap::new(),
103 }
104 }
105
106 pub fn register_tool(&mut self, schema: ToolSchema) {
108 self.schemas.insert(schema.name.clone(), schema);
109 }
110
111 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 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 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 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 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(¶m.name) {
165 return Err(ParseError::MissingField(param.name.clone()));
166 }
167 }
168 Ok(())
169 }
170
171 pub fn extract_json_objects<'a>(&self, text: &'a str) -> Vec<&'a str> {
176 extract_balanced_braces(text)
177 }
178
179 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
191pub struct ToolCallBuilder {
197 tool_name: String,
198 arguments: HashMap<String, String>,
199}
200
201impl ToolCallBuilder {
202 pub fn new(tool_name: &str) -> Self {
204 ToolCallBuilder {
205 tool_name: tool_name.to_string(),
206 arguments: HashMap::new(),
207 }
208 }
209
210 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 pub fn build(self) -> ToolCall {
218 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
237fn 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
294fn 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
307fn 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 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 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 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 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
350fn 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
386fn 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 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
449fn 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#[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(); 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 let text = r#"{"tool": incomplete"#;
525 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}