257 lines
8.7 KiB
Rust
257 lines
8.7 KiB
Rust
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: "<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_calls>",
|
||
tool_end: "</tool_calls>",
|
||
invoke_start: "<invoke",
|
||
invoke_end: "</invoke>",
|
||
parameter_start: "<parameter",
|
||
parameter_end: "</parameter>",
|
||
},
|
||
];
|
||
|
||
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<Option<(String, Value)>, 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<String> {
|
||
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|tool_calls>",
|
||
"<|DSML|invoke",
|
||
"</|DSML|invoke>",
|
||
"<|DSML|parameter",
|
||
"</|DSML|parameter>",
|
||
),
|
||
(
|
||
"<DSML|tool_calls>",
|
||
"</DSML|tool_calls>",
|
||
"<DSML|invoke",
|
||
"</DSML|invoke>",
|
||
"<DSML|parameter",
|
||
"</DSML|parameter>",
|
||
),
|
||
(
|
||
"<tool_calls>",
|
||
"</tool_calls>",
|
||
"<invoke",
|
||
"</invoke>",
|
||
"<parameter",
|
||
"</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</|DSML|parameter></|DSML|invoke>";
|
||
|
||
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 = "<tool_calls><invoke name=\"first\"></invoke><invoke name=\"second\"><parameter name=\"value\">truncated";
|
||
|
||
let (_, calls) = parse_tool_calls(raw).unwrap();
|
||
|
||
assert_eq!(calls.len(), 1);
|
||
assert_eq!(calls[0].0, "first");
|
||
|
||
let invalid = "<tool_calls><invoke name=\"first\"></invoke><invoke name=\"second\"><parameter name=\"value\" string=\"false\">invalid</parameter></invoke></tool_calls>";
|
||
assert!(
|
||
parse_tool_calls(invalid)
|
||
.unwrap_err()
|
||
.starts_with("invalid DSML tool arguments")
|
||
);
|
||
}
|
||
}
|