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: "|DSML|tool_calls>",
invoke_start: "<|DSML|invoke",
invoke_end: "|DSML|invoke>",
parameter_start: "<|DSML|parameter",
parameter_end: "|DSML|parameter>",
},
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