Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 4 additions & 7 deletions crates/tool_parser/src/parsers/glm4_moe.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
use std::collections::HashMap;

use async_trait::async_trait;
use openai_protocol::common::Tool;
use regex::Regex;
Expand Down Expand Up @@ -172,17 +170,16 @@ impl Glm4MoeParser {
fn parse_arguments(
&self,
args_text: &str,
param_types: &HashMap<String, String>,
param_types: &helpers::ParamTypes<'_>,
) -> serde_json::Map<String, Value> {
let mut arguments = serde_json::Map::new();

for capture in self.arg_extractor.captures_iter(args_text) {
let key = capture.get(1).map_or("", |m| m.as_str()).trim();
let value_str = capture.get(2).map_or("", |m| m.as_str()).trim();

let value =
helpers::coerce_by_schema_type(value_str, param_types.get(key).map(String::as_str))
.unwrap_or_else(|| infer_value(value_str));
let value = helpers::coerce_by_schema_type(value_str, param_types.get(key))
.unwrap_or_else(|| infer_value(value_str));

arguments.insert(key.to_string(), value);
}
Expand All @@ -199,7 +196,7 @@ impl Glm4MoeParser {
// Get arguments text
let args_text = captures.get(2).map_or("", |m| m.as_str());

let param_types = helpers::param_types_for_function(tools, func_name);
let param_types = helpers::ParamTypes::for_function(tools, func_name);
let arguments = self.parse_arguments(args_text, &param_types);

let arguments_str = serde_json::to_string(&arguments)
Expand Down
80 changes: 80 additions & 0 deletions crates/tool_parser/src/parsers/helpers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,43 @@ use crate::{
types::{StreamingParseResult, ToolCallItem},
};

/// Declared types for XML arguments, including unlisted parameter names.
pub(crate) struct ParamTypes<'a> {
schema: Option<&'a Value>,
}

impl<'a> ParamTypes<'a> {
pub(crate) fn for_function(tools: &'a [Tool], func_name: &str) -> Self {
Self {
schema: tools
.iter()
.find(|tool| tool.function.name == func_name)
.map(|tool| &tool.function.parameters),
}
}

/// Read a single declared type without changing union or unknown-type inference.
pub(crate) fn get(&self, name: &str) -> Option<&str> {
Comment thread
ai-jz marked this conversation as resolved.
let root = self.schema?;
let schema = if let Some(property) = root.get("properties").and_then(|p| p.get(name)) {
property
} else {
// additionalProperties excludes keys matched by patternProperties.
// Patterns are not evaluated here; any nonempty patternProperties
// keeps unlisted keys on the existing inference path.
if root
.get("patternProperties")
.and_then(Value::as_object)
.is_some_and(|patterns| !patterns.is_empty())
{
return None;
}
root.get("additionalProperties")?
};
schema.get("type").and_then(Value::as_str)
}
}

/// `param_name -> declared JSON-schema type` for the named function (empty if the
Comment thread
ai-jz marked this conversation as resolved.
/// function or its `properties` are absent). Lets XML-style parsers coerce by the
/// declared type instead of guessing from text (e.g. keep a numeric-looking
Expand Down Expand Up @@ -510,6 +547,49 @@ pub(crate) fn handle_json_tool_streaming(
mod tests {
use super::*;

#[test]
fn test_additional_parameter_types_preserve_explicit_properties() {
let schema = serde_json::json!({
"properties": {"count": {"type": "integer"}, "untyped": {}},
"additionalProperties": {"type": "string"}
});
let types = ParamTypes {
schema: Some(&schema),
};
assert_eq!(types.get("code"), Some("string"));
assert_eq!(types.get("count"), Some("integer"));
assert_eq!(types.get("untyped"), None);
}

#[test]
fn test_additional_parameter_types_keep_unknown_type_fallback() {
for additional in [
Value::Bool(true),
Value::Bool(false),
serde_json::json!({}),
serde_json::json!({"type": ["string", "null"]}),
] {
let schema = serde_json::json!({"additionalProperties": additional});
assert_eq!(
ParamTypes {
schema: Some(&schema)
}
.get("code"),
None
);
}
let schema = serde_json::json!({
"properties": {"count": {"type": "integer"}},
"patternProperties": {"^code": {"type": "number"}},
"additionalProperties": {"type": "string"}
});
let types = ParamTypes {
schema: Some(&schema),
};
assert_eq!(types.get("code"), None);
assert_eq!(types.get("count"), Some("integer"));
}

#[test]
fn test_ends_with_partial_token() {
assert!(ends_with_partial_token("hello <|py", "<|python_tag|>").is_some());
Expand Down
34 changes: 14 additions & 20 deletions crates/tool_parser/src/parsers/minimax_m2.rs
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@ impl MinimaxM2Parser {
fn parse_parameters(
&self,
params_text: &str,
param_types: &HashMap<String, String>,
param_types: &helpers::ParamTypes<'_>,
) -> serde_json::Map<String, Value> {
let mut parameters = serde_json::Map::new();

Expand All @@ -121,11 +121,8 @@ impl MinimaxM2Parser {
let value_str = capture.get(2).map_or("", |m| m.as_str());

let decoded_value = Self::decode_xml_entities(value_str);
let value = helpers::coerce_by_schema_type(
&decoded_value,
param_types.get(key).map(String::as_str),
)
.unwrap_or_else(|| Self::parse_value(&decoded_value));
let value = helpers::coerce_by_schema_type(&decoded_value, param_types.get(key))
.unwrap_or_else(|| Self::parse_value(&decoded_value));

parameters.insert(key.to_string(), value);
}
Expand Down Expand Up @@ -157,7 +154,7 @@ impl MinimaxM2Parser {
let params_text = captures.get(2).map_or("", |m| m.as_str());

// Parse parameters, coerced by this function's declared schema.
let param_types = helpers::param_types_for_function(tools, func_name);
let param_types = helpers::ParamTypes::for_function(tools, func_name);
let parameters = self.parse_parameters(params_text, &param_types);

match serde_json::to_string(&parameters) {
Expand Down Expand Up @@ -212,7 +209,7 @@ impl MinimaxM2Parser {
/// Parse and stream parameters incrementally
fn parse_and_stream_parameters(&mut self, text: &str, tools: &[Tool]) -> Vec<ToolCallItem> {
let mut calls = Vec::new();
let param_types = helpers::param_types_for_function(tools, &self.current_function_name);
let param_types = helpers::ParamTypes::for_function(tools, &self.current_function_name);

// Find all complete parameter patterns in the buffer
let param_matches: Vec<_> = self
Expand All @@ -225,18 +222,15 @@ impl MinimaxM2Parser {

// Coerce by declared type when known; otherwise keep the prior
// JSON-first-then-infer behavior for nested objects/arrays.
let value = helpers::coerce_by_schema_type(
&decoded,
param_types.get(&name).map(String::as_str),
)
.unwrap_or_else(|| {
if decoded.starts_with('{') || decoded.starts_with('[') {
serde_json::from_str::<Value>(&decoded)
.unwrap_or_else(|_| Self::parse_value(&decoded))
} else {
Self::parse_value(&decoded)
}
});
let value = helpers::coerce_by_schema_type(&decoded, param_types.get(&name))
.unwrap_or_else(|| {
if decoded.starts_with('{') || decoded.starts_with('[') {
serde_json::from_str::<Value>(&decoded)
.unwrap_or_else(|_| Self::parse_value(&decoded))
} else {
Self::parse_value(&decoded)
}
});

(name, value)
})
Expand Down
8 changes: 4 additions & 4 deletions crates/tool_parser/src/parsers/qwen_xml.rs
Original file line number Diff line number Diff line change
Expand Up @@ -149,14 +149,14 @@ impl QwenXmlParser {
return Ok(None);
}

let param_types = helpers::param_types_for_function(tools, &function_name);
let param_types = helpers::ParamTypes::for_function(tools, &function_name);
let mut parameters = serde_json::Map::new();

for cap in self.xml_param_pattern.captures_iter(content) {
if let (Some(key_match), Some(value_match)) = (cap.get(1), cap.get(2)) {
let key = key_match.as_str().trim().to_string();
let value = value_match.as_str();
let json_value = coerce_value(value, param_types.get(&key).map(String::as_str));
let json_value = coerce_value(value, param_types.get(&key));
parameters.insert(key, json_value);
}
}
Expand All @@ -180,7 +180,7 @@ impl QwenXmlParser {
parameter_end: usize,
) -> Vec<ToolCallItem> {
let mut calls: Vec<ToolCallItem> = vec![];
let param_types = helpers::param_types_for_function(tools, &self.current_function_name);
let param_types = helpers::ParamTypes::for_function(tools, &self.current_function_name);

// Leave parameters from subsequent coalesced calls for their own iteration.
let mut new_params = serde_json::Map::new();
Expand All @@ -191,7 +191,7 @@ impl QwenXmlParser {
if let (Some(key_match), Some(value_match)) = (cap.get(1), cap.get(2)) {
let key = key_match.as_str().trim().to_string();
let value = value_match.as_str();
let json_value = coerce_value(value, param_types.get(&key).map(String::as_str));
let json_value = coerce_value(value, param_types.get(&key));
new_params.insert(key, json_value);
}
}
Expand Down
66 changes: 66 additions & 0 deletions crates/tool_parser/tests/tool_parser_additional_properties.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
use openai_protocol::common::Tool;
use serde_json::{json, Value};
use tool_parser::{parsers::QwenXmlParser, Glm4MoeParser, MinimaxM2Parser, ToolParser};

#[tokio::test]
async fn test_additional_properties_string_complete_and_streaming() {
let tools: Vec<Tool> = serde_json::from_value(json!([{
"type": "function",
"function": {
"name": "save_config",
"parameters": {"type": "object", "additionalProperties": {"type": "string"}}
}
}]))
.unwrap();
let inputs = [
"<tool_call>save_config<arg_key>code</arg_key><arg_value>42</arg_value><arg_key>empty</arg_key><arg_value>null</arg_value></tool_call>",
"<tool_call><function=save_config><parameter=code>42</parameter><parameter=empty>null</parameter></function></tool_call>",
"<minimax:tool_call><invoke name=\"save_config\"><parameter name=\"code\">42</parameter><parameter name=\"empty\">null</parameter></invoke></minimax:tool_call>",
];
for (dialect, input) in inputs.iter().enumerate() {
for streaming in [false, true] {
let mut parser: Box<dyn ToolParser> = match dialect {
0 => Box::new(Glm4MoeParser::glm47()),
1 => Box::new(QwenXmlParser::new()),
_ => Box::new(MinimaxM2Parser::new()),
};
let arguments = if streaming {
let mut arguments = String::new();
let prefix = if dialect == 2 {
"<minimax:tool_call>"
} else {
"<tool_call>"
};
let chunks = std::iter::once(&input[..prefix.len()]).chain(
input.as_bytes()[prefix.len()..]
.chunks(7)
.map(|chunk| std::str::from_utf8(chunk).unwrap()),
);
for chunk in chunks {
for call in parser.parse_incremental(chunk, &tools).await.unwrap().calls {
assert_eq!(call.tool_index, 0);
arguments.push_str(&call.parameters);
}
}
if let Some(calls) = parser.get_unstreamed_tool_args() {
for call in calls {
arguments.push_str(&call.parameters);
}
}
arguments
} else {
let (_, calls) = parser
.parse_complete_with_tools(input, &tools)
.await
.unwrap();
assert_eq!(calls.len(), 1);
calls[0].function.arguments.clone()
};
assert_eq!(
serde_json::from_str::<Value>(&arguments).unwrap(),
json!({"code": "42", "empty": "null"}),
"dialect={dialect}, streaming={streaming}"
);
}
}
}
Loading