Files
DS4Server/src/dsml.rs
2026-07-29 18:47:14 +02:00

257 lines
8.7 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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: "<DSMLtool_calls>",
tool_end: "</DSMLtool_calls>",
invoke_start: "<DSMLinvoke",
invoke_end: "</DSMLinvoke>",
parameter_start: "<DSMLparameter",
parameter_end: "</DSMLparameter>",
},
Syntax {
tool_start: "<DSMLtool_calls>",
tool_end: "</DSMLtool_calls>",
invoke_start: "<DSMLinvoke",
invoke_end: "</DSMLinvoke>",
parameter_start: "<DSMLparameter",
parameter_end: "</DSMLparameter>",
},
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("&quot;", "\"")
.replace("&gt;", ">")
.replace("&lt;", "<")
.replace("&amp;", "&")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_reference_and_legacy_dsml_syntaxes() {
for (start, end, invoke, invoke_end, parameter, parameter_end) in [
(
"<DSMLtool_calls>",
"</DSMLtool_calls>",
"<DSMLinvoke",
"</DSMLinvoke>",
"<DSMLparameter",
"</DSMLparameter>",
),
(
"<DSMLtool_calls>",
"</DSMLtool_calls>",
"<DSMLinvoke",
"</DSMLinvoke>",
"<DSMLparameter",
"</DSMLparameter>",
),
(
"<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<DSMLtool_calls><DSMLinvoke name=\"read\"><DSMLparameter name=\"path\" string=\"true\">src/main.rs</DSMLparameter></DSMLinvoke>";
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")
);
}
}