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
201 changes: 191 additions & 10 deletions crates/tool_parser/src/factory.rs
Original file line number Diff line number Diff line change
Expand Up @@ -239,17 +239,15 @@ impl ParserRegistry {
/// Otherwise → `JsonSchema(schema)` for required/function tool_choice.
/// Returns `Ok(None)` for auto/none tool_choice.
///
/// `reasoning` says the rendered prompt ends inside the model's thinking
/// block (the gateway's `chat_reasoning_starts_in_prefill`). A parser that
/// registered a reasoning prefix then gets its tag wrapped so the forced
/// call follows the reasoning instead of preempting it; every other
/// constraint is unchanged and applies from the first generated token.
/// `parallel_tool_calls` mirrors the OpenAI request field: when `false`,
/// the constraint must allow at most one tool call (`maxItems: 1` on the
/// JSON schema, `stop_after_first` on structural tags).
pub fn generate_tool_constraint(
&self,
configured_parser: Option<&str>,
tools: &[Tool],
tool_choice: &ToolChoice,
reasoning: bool,
parallel_tool_calls: bool,
Comment thread
ighutake-debug marked this conversation as resolved.
) -> Result<Option<ToolConstraint>, String> {
if tools.is_empty() {
return Ok(None);
Expand All @@ -268,8 +266,28 @@ impl ParserRegistry {
if let Some(entry) = entries.get(name) {
if let Some(build_fn) = entry.build_structural_tag.as_ref() {
let mut tag = build_fn(tools, at_least_one);
if let (true, Some(prefix)) = (reasoning, entry.reasoning_prefix) {
tag = wrap_in_reasoning_prefix(tag, prefix())?;
if !parallel_tool_calls {
// triggered_tags dialect: stop after the first complete
// tool call instead of allowing parallel calls. A tag
// without an object-valued `format` cannot carry the
// constraint — fail loudly rather than silently
// delivering an unconstrained tag.
match tag
.get_mut("format")
.and_then(serde_json::Value::as_object_mut)
{
Some(format) => {
format.insert(
"stop_after_first".to_string(),
serde_json::Value::Bool(true),
);
}
None => {
return Err(format!(
"parser '{name}' produced a structural tag without a format object; cannot honor parallel_tool_calls=false"
));
}
}
}
let json_str = serde_json::to_string(&tag)
.map_err(|e| format!("Failed to serialize structural tag: {e}"))?;
Expand All @@ -286,7 +304,7 @@ impl ParserRegistry {
Ok(Some(ToolConstraint::JsonSchema(params_schema)))
}
_ => {
let schema = build_required_array_schema(tools)?;
let schema = build_required_array_schema(tools, parallel_tool_calls)?;
Ok(Some(ToolConstraint::JsonSchema(schema)))
}
}
Expand Down Expand Up @@ -647,7 +665,13 @@ fn wrap_in_reasoning_prefix(
}

/// Build JSON schema for required tool calls (array with minItems: 1).
fn build_required_array_schema(tools: &[Tool]) -> Result<String, String> {
///
/// `parallel_tool_calls=false` adds `maxItems: 1`, so the model can emit at
/// most one tool call (mirrors SGLang's `get_json_schema_constraint`).
fn build_required_array_schema(
tools: &[Tool],
parallel_tool_calls: bool,
) -> Result<String, String> {
let mut any_of_schemas = Vec::with_capacity(tools.len());
for tool in tools {
let tool_schema = json!({
Expand Down Expand Up @@ -692,6 +716,12 @@ fn build_required_array_schema(tools: &[Tool]) -> Result<String, String> {
}
});

if !parallel_tool_calls {
if let serde_json::Value::Object(ref mut obj) = array_schema {
obj.insert("maxItems".to_string(), json!(1));
}
}

if !all_defs.is_empty() {
if let serde_json::Value::Object(ref mut obj) = array_schema {
obj.insert("$defs".to_string(), serde_json::Value::Object(all_defs));
Expand All @@ -701,3 +731,154 @@ fn build_required_array_schema(tools: &[Tool]) -> Result<String, String> {
serde_json::to_string(&array_schema)
.map_err(|e| format!("Failed to serialize tool schema: {e}"))
}

#[cfg(test)]
mod tests {
use openai_protocol::common::{Function, FunctionChoice, Tool, ToolChoice, ToolChoiceValue};
use serde_json::{json, Value};

use super::*;

fn sample_tools() -> Vec<Tool> {
["get_weather", "search"]
.into_iter()
.map(|name| Tool {
tool_type: "function".to_string(),
function: Function {
name: name.to_string(),
description: None,
parameters: json!({"type": "object", "properties": {}}),
strict: None,
},
})
.collect()
}

fn required() -> ToolChoice {
ToolChoice::Value(ToolChoiceValue::Required)
}

fn json_schema_of(constraint: Option<ToolConstraint>) -> Value {
let Some(ToolConstraint::JsonSchema(schema)) = constraint else {
panic!("expected a json_schema constraint");
};
serde_json::from_str(&schema).unwrap()
}

fn structural_tag_of(constraint: Option<ToolConstraint>) -> Value {
let Some(ToolConstraint::StructuralTag(tag)) = constraint else {
panic!("expected a structural_tag constraint");
};
serde_json::from_str(&tag).unwrap()
}

#[test]
fn json_schema_constraint_gets_max_items_when_parallel_disabled() {
let registry = ParserRegistry::new();
let schema = json_schema_of(
registry
.generate_tool_constraint(None, &sample_tools(), &required(), false)
.unwrap(),
);
assert_eq!(schema["maxItems"], json!(1));
assert_eq!(schema["minItems"], json!(1));
}

#[test]
fn json_schema_constraint_omits_max_items_when_parallel_enabled() {
let registry = ParserRegistry::new();
let schema = json_schema_of(
registry
.generate_tool_constraint(None, &sample_tools(), &required(), true)
.unwrap(),
);
assert!(schema.get("maxItems").is_none());
}

#[test]
fn structural_tag_sets_stop_after_first_when_parallel_disabled() {
let factory = ParserFactory::new();
let tag = structural_tag_of(
factory
.registry()
.generate_tool_constraint(Some("mistral"), &sample_tools(), &required(), false)
.unwrap(),
);
assert_eq!(tag["format"]["stop_after_first"], json!(true));
}

#[test]
fn structural_tag_omits_stop_after_first_when_parallel_enabled() {
let factory = ParserFactory::new();
let tag = structural_tag_of(
factory
.registry()
.generate_tool_constraint(Some("mistral"), &sample_tools(), &required(), true)
.unwrap(),
);
assert!(tag["format"].get("stop_after_first").is_none());
}

#[test]
fn function_choice_constraint_is_unchanged_by_parallel_flag() {
// tool_choice = a specific function already means exactly one call;
// the params schema passes through untouched either way.
let registry = ParserRegistry::new();
let choice = ToolChoice::Function {
tool_type: "function".to_string(),
function: FunctionChoice {
name: "get_weather".to_string(),
},
};
for parallel in [true, false] {
let schema = json_schema_of(
registry
.generate_tool_constraint(None, &sample_tools(), &choice, parallel)
.unwrap(),
);
assert_eq!(schema["type"], json!("object"));
assert!(schema.get("maxItems").is_none());
}
}

#[test]
fn auto_choice_still_produces_no_constraint_when_parallel_disabled() {
let registry = ParserRegistry::new();
let constraint = registry
.generate_tool_constraint(
None,
&sample_tools(),
&ToolChoice::Value(ToolChoiceValue::Auto),
false,
)
.unwrap();
assert!(constraint.is_none());
}

#[test]
fn structural_tag_without_format_object_errors_when_parallel_disabled() {
// A builder whose tag has no object-valued "format" cannot carry
// stop_after_first; silently passing it through would break the
// single-call guarantee, so the registry must fail loudly.
let registry = ParserRegistry::new();
registry.register_parser_with_structural_tag(
"bad_tag",
|| Box::new(PassthroughParser::new()),
|_tools, _at_least_one| json!({"unexpected": true}),
);

let err = registry
.generate_tool_constraint(Some("bad_tag"), &sample_tools(), &required(), false)
.expect_err("malformed tag must error when parallel calls are disabled");
assert!(
err.contains("bad_tag"),
"error should name the parser: {err}"
);

// With parallel calls enabled the tag passes through untouched.
let ok = registry
.generate_tool_constraint(Some("bad_tag"), &sample_tools(), &required(), true)
.unwrap();
assert!(ok.is_some());
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -301,7 +301,7 @@ pub(crate) async fn prepare_chat_like(
.as_deref(),
&constraint_tools,
tool_choice,
reasoning,
request.parallel_tool_calls.unwrap_or(true),
)
.map_err(|e| {
error!(function = "ChatPreparationStage::execute", error = %e, "Invalid tool configuration");
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ use async_trait::async_trait;
use axum::response::Response;
use openai_protocol::{
common::{StringOrArray, ToolChoice, ToolChoiceValue},
messages::CreateMessageRequest,
messages::{CreateMessageRequest, ToolChoice as AnthropicToolChoice},
};
use tracing::{debug, error};

Expand Down Expand Up @@ -280,9 +280,24 @@ impl MessagePreparationStage {
}
}

// Step 4: Build tool constraints if tools present. On a thinking
// prompt a parser with a reasoning prefix gets its tag wrapped so a
// forced call follows the reasoning instead of preempting it.
// Step 4: Build tool constraints if tools present
// Anthropic's tool_choice.disable_parallel_tool_use is the Messages
// API equivalent of OpenAI's parallel_tool_calls=false.
let parallel_tool_calls = match request.tool_choice.as_ref() {
Some(
AnthropicToolChoice::Auto {
disable_parallel_tool_use,
}
| AnthropicToolChoice::Any {
disable_parallel_tool_use,
}
| AnthropicToolChoice::Tool {
disable_parallel_tool_use,
..
},
) => !disable_parallel_tool_use.unwrap_or(false),
_ => true,
};
let tool_call_constraint = if let (false, Some(tool_choice)) =
(filtered_tools.is_empty(), chat_tool_choice.as_ref())
{
Expand All @@ -298,7 +313,7 @@ impl MessagePreparationStage {
.as_deref(),
&filtered_tools,
tool_choice,
reasoning,
parallel_tool_calls,
)
.map_err(|e| {
error!(function = "MessagePreparationStage::execute", error = %e, "Invalid tool configuration");
Expand Down
Loading