chat : enable tool call in thinking for DS4 (#26269)

This commit is contained in:
Piotr Wilkin (ilintar)
2026-08-01 00:13:07 -05:00
committed by GitHub
parent 876a432116
commit ddd4ec1428
3 changed files with 124 additions and 101 deletions
+58 -52
View File
@@ -1943,18 +1943,6 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
adjusted_messages = deepseek_v4_sort_tool_results(inputs.messages); adjusted_messages = deepseek_v4_sort_tool_results(inputs.messages);
} }
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages);
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages);
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
data.supports_thinking = true;
data.thinking_start_tag = "<think>";
data.thinking_end_tags = {"</think>"};
data.preserved_tokens = {
"DSML",
"<think>",
"</think>",
};
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object(); auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
@@ -1972,6 +1960,18 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
const std::string PARAM_END = "</" + DSML + "parameter>"; const std::string PARAM_END = "</" + DSML + "parameter>";
const std::string GEN_PROMPT = "<Assistant>"; const std::string GEN_PROMPT = "<Assistant>";
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages);
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages);
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
data.supports_thinking = true;
data.thinking_start_tag = THINK_START;
data.thinking_end_tags = {THINK_END, FC_START};
data.preserved_tokens = {
DSML,
THINK_START,
THINK_END,
};
if (inputs.has_continuation()) { if (inputs.has_continuation()) {
const auto & msg = inputs.continue_msg; const auto & msg = inputs.continue_msg;
@@ -1983,35 +1983,16 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
data.prompt += data.generation_prompt; data.prompt += data.generation_prompt;
} }
bool require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
bool has_tool_calls = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) { auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PROMPT); auto generation_prompt = p.literal(GEN_PROMPT);
auto end = p.end(); auto end = p.end();
auto reasoning = p.eps(); // build tool call section first since we might need it in reasoning
if (extract_reasoning && inputs.enable_thinking) {
reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END);
} else if (extract_reasoning) {
// Thinking disabled but reasoning extraction requested: the generation prompt
// contains an empty <think></think> pair (V3.2) or a bare </think> (V4) that
// must still be consumed.
reasoning = is_v4
? p.optional(p.literal(THINK_END))
: p.optional(p.literal(THINK_START) + p.until(THINK_END) + p.literal(THINK_END));
}
if (has_response_format) {
auto response_format = p.rule("response-format",
p.literal("```json") + p.space() +
p.content(p.schema(p.json(), "response-format-schema", inputs.json_schema)) +
p.space() + p.literal("```"));
return generation_prompt + reasoning + response_format + end;
}
if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) {
return generation_prompt + reasoning + p.content(p.rest()) + end;
}
auto tool_choice = p.choice(); auto tool_choice = p.choice();
if (has_tool_calls) {
foreach_function(inputs.tools, [&](const json & tool) { foreach_function(inputs.tools, [&](const json & tool) {
const auto & function = tool.at("function"); const auto & function = tool.at("function");
std::string name = function.at("name"); std::string name = function.at("name");
@@ -2033,14 +2014,11 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
bool is_string = schema_info.resolves_to_string(param_schema); bool is_string = schema_info.resolves_to_string(param_schema);
auto arg = p.tool_arg( auto arg = p.tool_arg(
p.tool_arg_open( p.tool_arg_open(p.literal(PARAM_START + " name=\"") + p.tool_arg_name(p.literal(param_name)) +
p.literal(PARAM_START + " name=\"") +
p.tool_arg_name(p.literal(param_name)) +
p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) + p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
(is_string (is_string ?
? p.tool_arg_string_value(p.until(PARAM_END)) p.tool_arg_string_value(p.until(PARAM_END)) :
: p.tool_arg_json_value(p.schema(p.json(), p.tool_arg_json_value(p.schema(p.json(), "tool-" + name + "-arg-" + param_name + "-schema",
"tool-" + name + "-arg-" + param_name + "-schema",
param_schema, false))) + param_schema, false))) +
p.tool_arg_close(p.literal(PARAM_END))); p.tool_arg_close(p.literal(PARAM_END)));
@@ -2069,16 +2047,13 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
} }
common_peg_parser invoke_body = args_seq; common_peg_parser invoke_body = args_seq;
auto func_parser = p.tool( auto func_parser = p.tool(p.tool_open(p.literal(INVOKE_START + " name=\"") +
p.tool_open(p.literal(INVOKE_START + " name=\"") +
p.tool_name(p.literal(name)) + p.literal("\">\n")) + p.tool_name(p.literal(name)) + p.literal("\">\n")) +
invoke_body + p.space() + invoke_body + p.space() + p.tool_close(p.literal(INVOKE_END)));
p.tool_close(p.literal(INVOKE_END)));
tool_choice |= p.rule("tool-" + name, func_parser); tool_choice |= p.rule("tool-" + name, func_parser);
}); });
}
auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
common_peg_parser tool_calls = p.eps(); common_peg_parser tool_calls = p.eps();
if (inputs.parallel_tool_calls) { if (inputs.parallel_tool_calls) {
@@ -2090,18 +2065,49 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
p.literal(FC_START) + p.space() + tool_choice + p.space() + p.literal(FC_END)); p.literal(FC_START) + p.space() + tool_choice + p.space() + p.literal(FC_END));
} }
auto reasoning = p.eps();
auto reasoning_with_tc = p.eps();
auto obligatory_tool_calls = tool_calls;
bool allow_reasoning_with_tc = false;
if (!require_tools) { if (!require_tools) {
tool_calls = p.optional(tool_calls); tool_calls = p.optional(tool_calls);
} }
auto content_before_tools = p.content(p.until(FC_START)); if (extract_reasoning && inputs.enable_thinking) {
return generation_prompt + reasoning + content_before_tools + tool_calls + end; reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END);
reasoning_with_tc = THINK_START + p.reasoning(p.until_one_of({ FC_START, THINK_END })) + obligatory_tool_calls;
allow_reasoning_with_tc = true;
} else if (extract_reasoning) {
// Thinking disabled but reasoning extraction requested: the generation prompt
// contains an empty <think></think> pair (V3.2) or a bare </think> (V4) that
// must still be consumed.
reasoning = is_v4
? p.optional(p.literal(THINK_END))
: p.optional(p.literal(THINK_START) + p.until(THINK_END) + p.literal(THINK_END));
}
if (has_response_format) {
auto response_format = p.rule("response-format",
p.literal("```json") + p.space() +
p.content(p.schema(p.json(), "response-format-schema", inputs.json_schema)) +
p.space() + p.literal("```"));
return generation_prompt + reasoning + response_format + end;
}
if (!has_tool_calls) {
return generation_prompt + reasoning + p.content(p.rest()) + end;
}
auto content_before_tools = p.negate(p.literal(THINK_START)) + p.content(p.until(FC_START));
return allow_reasoning_with_tc ? generation_prompt + (reasoning_with_tc | (reasoning + content_before_tools + tool_calls)) + end :
generation_prompt + reasoning + content_before_tools + tool_calls + end;
}); });
data.parser = parser.save(); data.parser = parser.save();
if (include_grammar) { if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED)); data.grammar_lazy = has_tools && !require_tools;
data.grammar = build_grammar([&](const common_grammar_builder & builder) { data.grammar = build_grammar([&](const common_grammar_builder & builder) {
foreach_function(inputs.tools, [&](const json & tool) { foreach_function(inputs.tools, [&](const json & tool) {
const auto & function = tool.at("function"); const auto & function = tool.at("function");
+14
View File
@@ -4227,6 +4227,20 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
.expect_reasoning("I'm thinking") .expect_reasoning("I'm thinking")
.expect_content("Hello, world!\nWhat's up?") .expect_content("Hello, world!\nWhat's up?")
.run(); .run();
tst.test(
"Let me check the time\n\n"
"<DSMLtool_calls>\n"
"<DSMLinvoke name=\"get_time\">\n"
"<DSMLparameter name=\"city\" string=\"true\">Tokyo</DSMLparameter>\n"
"</DSMLinvoke>\n"
"</DSMLtool_calls>") // no </think> after the TC close because the grammar will immediately constrain it to end
.enable_thinking(true)
.reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK)
.tools({ get_time_tool })
.expect_reasoning("Let me check the time")
.expect_tool_calls({ { "get_time", R"({"city": "Tokyo"})", {} } })
.run();
} }
// GLM-4.6 tests - format: <tool_call>function_name\n<arg_key>...</arg_key>\n<arg_value>...</arg_value>\n</tool_call> // GLM-4.6 tests - format: <tool_call>function_name\n<arg_key>...</arg_key>\n<arg_value>...</arg_value>\n</tool_call>
+4 -1
View File
@@ -9,6 +9,7 @@
#include "peg-parser.h" #include "peg-parser.h"
#include <fstream> #include <fstream>
#include <iterator>
#include <numeric> #include <numeric>
#include <optional> #include <optional>
#include <sstream> #include <sstream>
@@ -398,7 +399,7 @@ int main(int argc, char ** argv) {
if (std::optional<common_chat_params> spec_tmpl = if (std::optional<common_chat_params> spec_tmpl =
common_chat_try_specialized_template(chat_template, template_source, params)) { common_chat_try_specialized_template(chat_template, template_source, params)) {
LOG_ERR("\n"); LOG_ERR("\n");
LOG_ERR("This template uses a specialized parser, analysis results will not be available."); LOG_ERR("This template uses a specialized parser, analysis results will not be available.\n");
parser_data = *spec_tmpl; parser_data = *spec_tmpl;
} else { } else {
// Render template scenarios if requested // Render template scenarios if requested
@@ -426,7 +427,9 @@ int main(int argc, char ** argv) {
// Generate Parser // Generate Parser
parser_data = autoparser::peg_generator::generate_parser(chat_template, params, analysis); parser_data = autoparser::peg_generator::generate_parser(chat_template, params, analysis);
} }
}
if (!std::empty(parser_data.parser)) {
LOG_ERR("\n=== Generated Parser ===\n"); LOG_ERR("\n=== Generated Parser ===\n");
common_peg_arena arena; common_peg_arena arena;
arena.load(parser_data.parser); arena.load(parser_data.parser);