use serde_json::{Map, Value}; const INCOMPLETE_TOOL_CALL: &str = "invalid or incomplete DSML tool call"; #[derive(Clone, Copy)] pub(crate) struct Syntax { pub(crate) tool_start: &'static str, pub(crate) tool_end: &'static str, pub(crate) invoke_start: &'static str, pub(crate) invoke_end: &'static str, pub(crate) parameter_start: &'static str, pub(crate) parameter_end: &'static str, } pub(crate) const SYNTAXES: [Syntax; 3] = [ Syntax { tool_start: "<|DSML|tool_calls>", tool_end: "", invoke_start: "<|DSML|invoke", invoke_end: "", parameter_start: "<|DSML|parameter", parameter_end: "", }, Syntax { tool_start: "", tool_end: "", invoke_start: "", parameter_start: "", }, Syntax { tool_start: "", tool_end: "", invoke_start: "", parameter_start: "", }, ]; pub(crate) fn parse_tool_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String> { let Some((start, syntax)) = SYNTAXES .iter() .filter_map(|syntax| text.find(syntax.tool_start).map(|start| (start, *syntax))) .min_by_key(|(start, _)| *start) else { return Ok((text.to_owned(), Vec::new())); }; let relative_end = text[start..].find(syntax.tool_end); let end = relative_end .map(|relative_end| start + relative_end + syntax.tool_end.len()) .unwrap_or(text.len()); let outer_call_complete = relative_end.is_some(); let content = text[..start].trim_end().to_owned(); let raw = &text[start..end]; let mut cursor = syntax.tool_start.len(); let mut calls = Vec::new(); loop { skip_whitespace(raw, &mut cursor); if raw[cursor..].starts_with(syntax.tool_end) { break; } // DS4 recovers calls whose invokes are complete even if generation // stopped before emitting the outer tool-calls closing tag. if !outer_call_complete && cursor == raw.len() && !calls.is_empty() { break; } let call = (|| { if !raw[cursor..].starts_with(syntax.invoke_start) { return Err(INCOMPLETE_TOOL_CALL.into()); } let tag_end = raw[cursor..] .find('>') .map(|offset| cursor + offset + 1) .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; let name = attribute(&raw[cursor..tag_end], "name") .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; cursor = tag_end; let mut arguments = Map::new(); loop { skip_whitespace(raw, &mut cursor); if raw[cursor..].starts_with(syntax.invoke_end) { cursor += syntax.invoke_end.len(); break; } let (name, value) = parse_parameter(raw, &mut cursor, syntax)? .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; arguments.insert(name, value); } Ok((name, Value::Object(arguments))) })(); match call { Ok(call) => calls.push(call), Err(error) if !calls.is_empty() && error == INCOMPLETE_TOOL_CALL => { break; } Err(error) => return Err(error), } } if calls.is_empty() { return Err(INCOMPLETE_TOOL_CALL.into()); } Ok((content, calls)) } fn parse_parameter( text: &str, cursor: &mut usize, syntax: Syntax, ) -> Result, String> { if !text[*cursor..].starts_with(syntax.parameter_start) { return Ok(None); } let Some(tag_end) = text[*cursor..].find('>').map(|end| *cursor + end + 1) else { return Ok(None); }; let tag = &text[*cursor..tag_end]; let Some(name) = attribute(tag, "name") else { return Ok(None); }; let is_string = attribute(tag, "string"); *cursor = tag_end; let mut nested_start = *cursor; skip_whitespace(text, &mut nested_start); if is_string.is_none() && text[nested_start..].starts_with(syntax.parameter_start) { *cursor = nested_start; let mut nested = Map::new(); loop { skip_whitespace(text, cursor); if !text[*cursor..].starts_with(syntax.parameter_start) { break; } let Some((name, value)) = parse_parameter(text, cursor, syntax)? else { return Ok(None); }; nested.insert(name, value); } skip_whitespace(text, cursor); if !text[*cursor..].starts_with(syntax.parameter_end) { return Ok(None); } *cursor += syntax.parameter_end.len(); return Ok(Some((name, Value::Object(nested)))); } let Some(value_end) = text[*cursor..] .find(syntax.parameter_end) .map(|end| *cursor + end) else { return Ok(None); }; let raw = &text[*cursor..value_end]; *cursor = value_end + syntax.parameter_end.len(); let value = if is_string.as_deref().unwrap_or("true") == "true" { Value::String(unescape(raw)) } else { serde_json::from_str(raw) .map_err(|error| format!("invalid DSML tool arguments: {error}"))? }; Ok(Some((name, value))) } fn attribute(tag: &str, name: &str) -> Option { let start = tag.find(&format!("{name}=\""))? + name.len() + 2; let end = start + tag[start..].find('"')?; Some(unescape(&tag[start..end])) } fn skip_whitespace(text: &str, cursor: &mut usize) { while text[*cursor..] .chars() .next() .is_some_and(char::is_whitespace) { *cursor += text[*cursor..].chars().next().unwrap().len_utf8(); } } fn unescape(text: &str) -> String { text.replace(""", "\"") .replace(">", ">") .replace("<", "<") .replace("&", "&") } #[cfg(test)] mod tests { use super::*; #[test] fn parses_reference_and_legacy_dsml_syntaxes() { for (start, end, invoke, invoke_end, parameter, parameter_end) in [ ( "<|DSML|tool_calls>", "", "<|DSML|invoke", "", "<|DSML|parameter", "", ), ( "", "", "", "", ), ( "", "", "", "", ), ] { let raw = format!( "done{start}{invoke} name=\"read\">{parameter} name=\"path\" string=\"true\">src/main.rs{parameter_end}{invoke_end}{end}" ); let (content, calls) = parse_tool_calls(&raw).unwrap(); assert_eq!(content, "done"); assert_eq!(calls[0].0, "read"); assert_eq!(calls[0].1["path"], "src/main.rs"); } } #[test] fn recovers_complete_invokes_without_the_outer_closing_tag() { let raw = "before<|DSML|tool_calls><|DSML|invoke name=\"read\"><|DSML|parameter name=\"path\" string=\"true\">src/main.rs"; let (content, calls) = parse_tool_calls(raw).unwrap(); assert_eq!(content, "before"); assert_eq!(calls.len(), 1); assert_eq!(calls[0].0, "read"); assert_eq!(calls[0].1["path"], "src/main.rs"); } #[test] fn recovers_complete_invokes_before_an_incomplete_one() { let raw = "truncated"; let (_, calls) = parse_tool_calls(raw).unwrap(); assert_eq!(calls.len(), 1); assert_eq!(calls[0].0, "first"); let invalid = "invalid"; assert!( parse_tool_calls(invalid) .unwrap_err() .starts_with("invalid DSML tool arguments") ); } }