]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
chat : enable tool call in thinking for DS4 (#26269)
authorPiotr Wilkin (ilintar) <redacted>
Sat, 1 Aug 2026 05:13:07 +0000 (07:13 +0200)
committerGitHub <redacted>
Sat, 1 Aug 2026 05:13:07 +0000 (00:13 -0500)
common/chat.cpp
tests/test-chat.cpp
tools/parser/debug-template-parser.cpp

index 7740f35c0edccac34e53f095d183ac7865da6f5d..0e06fa591f9d66f10dcc081817aa7f570920b7c3 100644 (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);
     }
 
-    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_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
     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 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()) {
         const auto & msg = inputs.continue_msg;
 
@@ -1983,102 +1983,77 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
         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 generation_prompt = p.literal(GEN_PROMPT);
-        auto end = p.end();
-
-        auto reasoning = p.eps();
-        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 end               = p.end();
 
+        // build tool call section first since we might need it in reasoning
         auto tool_choice = p.choice();
-        foreach_function(inputs.tools, [&](const json & tool) {
-            const auto & function = tool.at("function");
-            std::string  name     = function.at("name");
-            auto params   = function.contains("parameters") ? function.at("parameters") : json::object();
-            const auto & props    = params.contains("properties") ? params.at("properties") : json::object();
-
-            std::set<std::string> required;
-            if (params.contains("required")) {
-                params.at("required").get_to(required);
-            }
-
-            auto schema_info = common_schema_info();
-            schema_info.resolve_refs(params);
+        if (has_tool_calls) {
+            foreach_function(inputs.tools, [&](const json & tool) {
+                const auto & function = tool.at("function");
+                std::string  name     = function.at("name");
+                auto         params   = function.contains("parameters") ? function.at("parameters") : json::object();
+                const auto & props    = params.contains("properties") ? params.at("properties") : json::object();
 
-            std::vector<common_peg_parser> required_parsers;
-            std::vector<common_peg_parser> optional_parsers;
-            for (const auto & [param_name, param_schema] : props.items()) {
-                bool is_required = required.find(param_name) != required.end();
-                bool is_string   = schema_info.resolves_to_string(param_schema);
-
-                auto arg = p.tool_arg(
-                    p.tool_arg_open(
-                        p.literal(PARAM_START + " name=\"") +
-                        p.tool_arg_name(p.literal(param_name)) +
-                        p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
-                    (is_string
-                         ? p.tool_arg_string_value(p.until(PARAM_END))
-                         : p.tool_arg_json_value(p.schema(p.json(),
-                                                          "tool-" + name + "-arg-" + param_name + "-schema",
-                                                          param_schema, false))) +
-                    p.tool_arg_close(p.literal(PARAM_END)));
-
-                auto named_arg = p.rule("tool-" + name + "-arg-" + param_name, arg);
-                if (is_required) {
-                    required_parsers.push_back(named_arg);
-                } else {
-                    optional_parsers.push_back(named_arg);
+                std::set<std::string> required;
+                if (params.contains("required")) {
+                    params.at("required").get_to(required);
                 }
-            }
 
-            common_peg_parser args_seq = p.eps();
-            for (size_t i = 0; i < required_parsers.size(); i++) {
-                if (i > 0) {
-                    args_seq = args_seq + p.space();
+                auto schema_info = common_schema_info();
+                schema_info.resolve_refs(params);
+
+                std::vector<common_peg_parser> required_parsers;
+                std::vector<common_peg_parser> optional_parsers;
+                for (const auto & [param_name, param_schema] : props.items()) {
+                    bool is_required = required.find(param_name) != required.end();
+                    bool is_string   = schema_info.resolves_to_string(param_schema);
+
+                    auto arg = p.tool_arg(
+                        p.tool_arg_open(p.literal(PARAM_START + " name=\"") + p.tool_arg_name(p.literal(param_name)) +
+                                        p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
+                        (is_string ?
+                             p.tool_arg_string_value(p.until(PARAM_END)) :
+                             p.tool_arg_json_value(p.schema(p.json(), "tool-" + name + "-arg-" + param_name + "-schema",
+                                                            param_schema, false))) +
+                        p.tool_arg_close(p.literal(PARAM_END)));
+
+                    auto named_arg = p.rule("tool-" + name + "-arg-" + param_name, arg);
+                    if (is_required) {
+                        required_parsers.push_back(named_arg);
+                    } else {
+                        optional_parsers.push_back(named_arg);
+                    }
                 }
-                args_seq = args_seq + required_parsers[i];
-            }
 
-            if (!optional_parsers.empty()) {
-                common_peg_parser any_opt = p.choice();
-                for (const auto & opt : optional_parsers) {
-                    any_opt |= opt;
+                common_peg_parser args_seq = p.eps();
+                for (size_t i = 0; i < required_parsers.size(); i++) {
+                    if (i > 0) {
+                        args_seq = args_seq + p.space();
+                    }
+                    args_seq = args_seq + required_parsers[i];
                 }
-                args_seq = args_seq + p.repeat(p.space() + any_opt, 0, -1);
-            }
 
-            common_peg_parser invoke_body = args_seq;
-            auto func_parser = p.tool(
-                p.tool_open(p.literal(INVOKE_START + " name=\"") +
-                            p.tool_name(p.literal(name)) + p.literal("\">\n")) +
-                invoke_body + p.space() +
-                p.tool_close(p.literal(INVOKE_END)));
+                if (!optional_parsers.empty()) {
+                    common_peg_parser any_opt = p.choice();
+                    for (const auto & opt : optional_parsers) {
+                        any_opt |= opt;
+                    }
+                    args_seq = args_seq + p.repeat(p.space() + any_opt, 0, -1);
+                }
 
-            tool_choice |= p.rule("tool-" + name, func_parser);
-        });
+                common_peg_parser invoke_body = args_seq;
+                auto              func_parser = p.tool(p.tool_open(p.literal(INVOKE_START + " name=\"") +
+                                                                   p.tool_name(p.literal(name)) + p.literal("\">\n")) +
+                                                       invoke_body + p.space() + p.tool_close(p.literal(INVOKE_END)));
 
-        auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
+                tool_choice |= p.rule("tool-" + name, func_parser);
+            });
+        }
 
         common_peg_parser tool_calls = p.eps();
         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));
         }
 
+        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) {
             tool_calls = p.optional(tool_calls);
         }
 
-        auto content_before_tools = p.content(p.until(FC_START));
-        return generation_prompt + reasoning + content_before_tools + tool_calls + end;
+        if (extract_reasoning && inputs.enable_thinking) {
+            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();
 
     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) {
             foreach_function(inputs.tools, [&](const json & tool) {
                 const auto & function = tool.at("function");
index 01b07953a6275da7113c919eb389133f3cd25de6..ef02fdde57ef3f7f38320706c5c96ab8d818c952 100644 (file)
@@ -4227,6 +4227,20 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
             .expect_reasoning("I'm thinking")
             .expect_content("Hello, world!\nWhat's up?")
             .run();
+
+        tst.test(
+            "Let me check the time\n\n"
+            "<|DSML|tool_calls>\n"
+            "<|DSML|invoke name=\"get_time\">\n"
+            "<|DSML|parameter name=\"city\" string=\"true\">Tokyo</|DSML|parameter>\n"
+            "</|DSML|invoke>\n"
+            "</|DSML|tool_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>
index 50e8f1efb7d6f8e0cf47c7e97533255615a62f56..8a916f79c78e0dbac35077a4e3af9353cff852dd 100644 (file)
@@ -9,6 +9,7 @@
 #include "peg-parser.h"
 
 #include <fstream>
+#include <iterator>
 #include <numeric>
 #include <optional>
 #include <sstream>
@@ -398,7 +399,7 @@ int main(int argc, char ** argv) {
         if (std::optional<common_chat_params> spec_tmpl =
                 common_chat_try_specialized_template(chat_template, template_source, params)) {
             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;
         } else {
             // Render template scenarios if requested
@@ -426,7 +427,9 @@ int main(int argc, char ** argv) {
                 // Generate Parser
                 parser_data = autoparser::peg_generator::generate_parser(chat_template, params, analysis);
             }
+        }
 
+        if (!std::empty(parser_data.parser)) {
             LOG_ERR("\n=== Generated Parser ===\n");
             common_peg_arena arena;
             arena.load(parser_data.parser);