Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
chat.cpp967 linesDownload Raw Back to common
1#include "chat.hpp"2#include "chat-template.hpp"3#include "json-schema-to-grammar.h"4#include "log.h"5#include "minja.hpp"6 7std::string common_chat_format_name(common_chat_format format) {8    switch (format) {9        case COMMON_CHAT_FORMAT_CONTENT_ONLY: return "Content-only";10        case COMMON_CHAT_FORMAT_GENERIC: return "Generic";11        case COMMON_CHAT_FORMAT_MISTRAL_NEMO: return "Mistral Nemo";12        case COMMON_CHAT_FORMAT_LLAMA_3_X: return "Llama 3.x";13        case COMMON_CHAT_FORMAT_LLAMA_3_X_WITH_BUILTIN_TOOLS: return "Llama 3.x with builtin tools";14        case COMMON_CHAT_FORMAT_DEEPSEEK_R1: return "DeepSeek R1";15        case COMMON_CHAT_FORMAT_FIREFUNCTION_V2: return "FireFunction v2";16        case COMMON_CHAT_FORMAT_FUNCTIONARY_V3_2: return "Functionary v3.2";17        case COMMON_CHAT_FORMAT_FUNCTIONARY_V3_1_LLAMA_3_1: return "Functionary v3.1 Llama 3.1";18        case COMMON_CHAT_FORMAT_HERMES_2_PRO: return "Hermes 2 Pro";19        case COMMON_CHAT_FORMAT_COMMAND_R7B: return "Command R7B";20        default:21            throw std::runtime_error("Unknown chat format");22    }23}24 25const common_grammar_options grammar_options {26    /* .dotall = */ false,27    /* .compact_spaces = */ false,28    // /* .compact_spaces = */ true,29};30 31static bool parse_json(std::string::const_iterator & it, const std::string::const_iterator & end, json & out) {32    // // https://json.nlohmann.me/features/parsing/sax_interface/33    struct json_error_locator : public nlohmann::json_sax<json> {34        std::size_t position;35        bool found_error;36 37        json_error_locator() : position(0), found_error(false) {}38 39        bool parse_error(std::size_t position, const std::string &, const json::exception &) override {40            this->position = position - 1;41            this->found_error = true;42            return false;43        }44        bool null() override { return true; }45        bool boolean(bool) override { return true; }46        bool number_integer(number_integer_t) override { return true; }47        bool number_unsigned(number_unsigned_t) override { return true; }48        bool number_float(number_float_t, const string_t &) override { return true; }49        bool string(string_t &) override { return true; }50        bool binary(binary_t &) override { return true; }51        bool start_object(std::size_t) override { return true; }52        bool key(string_t &) override { return true; }53        bool end_object() override { return true; }54        bool start_array(std::size_t) override { return true; }55        bool end_array() override { return true; }56    };57    json_error_locator err_loc;58    json::sax_parse(it, end, &err_loc);59 60    std::string::const_iterator temptative_end;61    if (err_loc.found_error) {62        temptative_end = it + err_loc.position;63    } else {64        temptative_end = end;65    }66    std::string json_sub {it, temptative_end};67    try {68        out = json::parse(json_sub);69        it = temptative_end;70        return true;71    } catch (const std::exception &) {72        return false;73    }74}75 76 77/**78 * Takes a prefix regex that must have 1 group to capture the function name, a closing suffix, and expects json parameters in between.79 * Aggregates the prefix, suffix and in-between text into the content.80 */81static common_chat_msg parse_json_tool_calls(82    const std::string& input,83    const std::optional<std::regex> & trigger_opt,84    const std::regex & function_regex,85    const std::regex & close_regex) {86    std::smatch match;87 88    common_chat_msg result;89    result.role = "assistant";90 91 92    auto end = input.end();93    auto it = input.begin();94 95    if (trigger_opt) {96        if (!std::regex_search(it, end, match, *trigger_opt)) {97            result.content = input;98            return result;99        }100        result.content = match.prefix().str();101        it = match.suffix().first;102    }103 104    while (it != end) {105        std::sregex_iterator rend;106        std::sregex_iterator rit(it, end, function_regex);107        if (rit == rend) {108            fprintf(stderr, "No more tool calls found\n");109            result.content += std::string(it, end);110            break;111        }112        auto name = rit->str(1);113        result.content += std::string(it, rit->prefix().second);114        it = rit->suffix().first;115 116        json arguments;117        if (!parse_json(it, end, arguments)) {118            throw std::runtime_error("Failed to parse json tool call arguments");119        }120        if (!std::regex_search(it, end, match, close_regex)) {121            throw std::runtime_error("Malformed input, missing closing pattern");122        }123        it = match.suffix().first;124        result.tool_calls.push_back({name, arguments.is_string() ? arguments.get<std::string>() : arguments.dump(), /* id= */ ""});125    }126    return result;127}128 129static common_chat_msg parse_prefixed_json_tool_call_array(const std::string& input, const std::string & prefix, size_t rstrip_prefix = 0) {130    auto content_end = input.find(prefix);131    size_t tc_start = std::string::npos;132 133    common_chat_msg result;134    result.role = "assistant";135    const auto process_tool_calls = [&](const json & tool_calls) {136        for (const auto & tool_call : tool_calls) {137            const auto & arguments = tool_call["arguments"];138            result.tool_calls.push_back({139                tool_call["name"],140                arguments.is_string() ? arguments.get<std::string>() : arguments.dump(),141                tool_call.contains("id") ? tool_call["id"] : "",142            });143        }144    };145    if (content_end == std::string::npos) {146        result.content = input;147    } else {148        tc_start = content_end + prefix.size() - rstrip_prefix;149        result.content = input.substr(0, content_end);150        auto tool_calls = json::parse(input.substr(tc_start));151        process_tool_calls(tool_calls);152    }153    return result;154}155 156static void foreach_function(const json & tools, const std::function<void(const json &)> & fn) {157    for (const auto & tool : tools) {158        if (!tool.contains("type") || tool["type"] != "function" || !tool.contains("function")) {159            LOG_INF("Skipping tool without function: %s", tool.dump(2).c_str());160            continue;161        }162        fn(tool);163    }164}165 166static std::string apply(167    const common_chat_template & tmpl,168    const nlohmann::ordered_json & messages,169    const nlohmann::ordered_json & tools,170    bool add_generation_prompt,171    const nlohmann::ordered_json & extra_context = nlohmann::ordered_json())172{173    minja::chat_template_inputs tmpl_inputs;174    tmpl_inputs.messages = messages;175    tmpl_inputs.tools = tools;176    tmpl_inputs.add_generation_prompt = add_generation_prompt;177    tmpl_inputs.extra_context = extra_context;178    // TODO: add flag to control date/time, if only for testing purposes.179    // tmpl_inputs.now = std::chrono::system_clock::now();180 181    minja::chat_template_options tmpl_opts;182    tmpl_opts.use_bos_token = false;183    tmpl_opts.use_eos_token = false;184 185    return tmpl.apply(tmpl_inputs, tmpl_opts);186}187 188static common_chat_params common_chat_params_init_generic(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {189    common_chat_params data;190 191    auto tool_call_schemas = json::array();192    foreach_function(inputs.tools, [&](const json & tool) {193        const auto & function = tool["function"];194        auto tool_schema = json {195            {"type", "object"},196            {"properties", {197                {"name", {198                    {"type", "string"},199                    {"const", function["name"]},200                }},201                {"arguments", function["parameters"]},202            }},203            {"required", json::array({"name", "arguments"})},204        };205        if (function.contains("description")) {206            tool_schema["description"] = function["description"];207        }208        if (inputs.parallel_tool_calls) {209            tool_schema["properties"]["id"] = {210                {"type", "string"},211                {"minLength", 4},212            };213            tool_schema["required"].push_back("id");214        }215        tool_call_schemas.emplace_back(tool_schema);216    });217    const auto tool_call =218        inputs.parallel_tool_calls219            ? json {220                {"type", "object"},221                {"properties", {222                    {"tool_calls", {223                        {"type", "array"},224                        {"items", tool_call_schemas.size() == 1 ? tool_call_schemas[0] : json {225                            {"anyOf", tool_call_schemas},226                        }},227                        {"minItems", 1},228                    }},229                }},230                {"required", json::array({"tool_calls"})},231            }232            : json {233                {"type", "object"},234                {"properties", {235                    {"tool_call", tool_call_schemas.size() == 1 ? tool_call_schemas[0] : json {236                        {"anyOf", tool_call_schemas},237                    }},238                }},239                {"required", json::array({"tool_call"})},240            };241    const auto schema =242        inputs.tool_choice != "required"243            ? json {244                {"anyOf", json::array({245                    tool_call,246                    {247                        {"type", "object"},248                        {"properties", {249                            {"response", inputs.json_schema.is_null()250                                ? json {{"type", "string"}}251                                : inputs.json_schema252                            },253                        }},254                        {"required", json::array({"response"})},255                    },256                })}257            }258            : tool_call;259 260    data.grammar_lazy = false;261    data.grammar = build_grammar([&](const common_grammar_builder & builder) {262        builder.add_schema("root", schema);263    }, grammar_options);264 265    auto tweaked_messages = common_chat_template::add_system(266        inputs.messages,267        "Respond in JSON format, either with `tool_call` (a request to call tools) or with `response` reply to the user's request");268 269    data.prompt = apply(tmpl, tweaked_messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);270    data.format = COMMON_CHAT_FORMAT_GENERIC;271    return data;272}273static common_chat_msg common_chat_parse_generic(const std::string & input) {274    json data = json::parse(input);275    common_chat_msg result;276    result.role = "assistant";277    if (data.contains("tool_calls")) {278        for (const auto & tool_call : data["tool_calls"]) {279            result.tool_calls.push_back({280                tool_call["name"],281                tool_call["arguments"].dump(),282                tool_call.contains("id") ? tool_call["id"] : "",283            });284        }285    } else if (data.contains("tool_call")) {286        result.tool_calls.push_back({287            data["tool_call"]["name"],288            data["tool_call"]["arguments"].dump(),289            /* id= */ "",290        });291    } else if (data.contains("response")) {292        const auto & response = data["response"];293        result.content = response.is_string() ? response.get<std::string>() : response.dump(2);294    }295    return result;296}297 298static common_chat_params common_chat_params_init_mistral_nemo(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {299    common_chat_params data;300    data.grammar_lazy = inputs.tool_choice != "required";301    data.grammar = build_grammar([&](const common_grammar_builder & builder) {302        auto schemas = json::array();303        foreach_function(inputs.tools, [&](const json & tool) {304            const auto & function = tool["function"];305            schemas.push_back({306                {"type", "object"},307                {"properties", {308                    // Important note: the model is probably trained to take a JSON stringified arguments value.309                    // It's hard to constrain that for now (while reusing the JSON schema conversion), so we're just expecting a plain object.310                    {"name", {311                        {"type", "string"},312                        {"const", function["name"]},313                    }},314                    {"arguments", function["parameters"]},315                    {"id", {316                        {"type", "string"},317                        // Nemo's template expects a 9-character alphanumeric ID.318                        {"pattern", "^[a-zA-Z0-9]{9}$"},319                    }},320                }},321                {"required", json::array({"name", "arguments", "id"})},322            });323        });324        auto schema = json {325            {"type", "array"},326            {"items", schemas.size() == 1 ? schemas[0] : json {{"anyOf", schemas}}},327            {"minItems", 1},328        };329        if (!inputs.parallel_tool_calls) {330            schema["maxItems"] = 1;331        }332        builder.add_rule("root", "\"[TOOL_CALLS]\" " + builder.add_schema("tool_calls", schema));333    }, grammar_options);334    data.grammar_triggers.push_back({"[TOOL_CALLS]", /* .at_start = */ true});335    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);336    data.format = COMMON_CHAT_FORMAT_MISTRAL_NEMO;337    return data;338}339static common_chat_msg common_chat_parse_mistral_nemo(const std::string & input) {340    return parse_prefixed_json_tool_call_array(input, "[TOOL_CALLS]");341}342 343static common_chat_params common_chat_params_init_command_r7b(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {344    common_chat_params data;345    data.grammar_lazy = inputs.tool_choice != "required";346    data.grammar = build_grammar([&](const common_grammar_builder & builder) {347        auto schemas = json::array();348        foreach_function(inputs.tools, [&](const json & tool) {349            const auto & function = tool["function"];350            schemas.push_back({351                {"type", "object"},352                {"properties", {353                    {"tool_call_id", {354                        {"type", "string"},355                        // Command-R's template expects an integer string.356                        {"pattern", "^[0-9]{1,10}$"},357                    }},358                    {"tool_name", {359                        {"type", "string"},360                        {"const", function["name"]},361                    }},362                    {"parameters", function["parameters"]},363                }},364                {"required", json::array({"tool_call_id", "tool_name", "parameters"})},365            });366        });367        auto schema = json {368            {"type", "array"},369            {"items", schemas.size() == 1 ? schemas[0] : json {{"anyOf", schemas}}},370            {"minItems", 1},371        };372        if (!inputs.parallel_tool_calls) {373            schema["maxItems"] = 1;374        }375        builder.add_rule("root", "\"<|START_ACTION|>\" " + builder.add_schema("tool_calls", schema) + " \"<|END_ACTION|>\"");376    }, grammar_options);377    data.grammar_triggers.push_back({"<|START_ACTION|>", /* .at_start = */ false});378    data.preserved_tokens = {379        "<|START_RESPONSE|>",380        "<|END_RESPONSE|>",381        "<|START_THINKING|>",382        "<|END_THINKING|>",383        "<|END_ACTION|>",384    };385    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);386    data.format = COMMON_CHAT_FORMAT_COMMAND_R7B;387    return data;388}389static common_chat_msg common_chat_parse_command_r7b(const std::string & input) {390    static std::regex response_regex("<\\|START_RESPONSE\\|>([\\s\\S\\n\\r]*?)<\\|END_RESPONSE\\|>");391    static std::regex thought_action_regex("<\\|START_THINKING\\|>([\\s\\S\\n\\r]*?)<\\|END_THINKING\\|><\\|START_ACTION\\|>([\\s\\S\\n\\r]*?)<\\|END_ACTION\\|>");392    std::smatch match;393 394    common_chat_msg result;395    result.role = "assistant";396    if (std::regex_match(input, match, response_regex)) {397        result.content = match[1].str();398    } else if (std::regex_match(input, match, thought_action_regex)) {399        result.tool_plan = match[1].str();400        auto actions_str = match[2].str();401        auto actions = json::parse(actions_str);402        for (const auto & action : actions) {403            result.tool_calls.push_back({404                /* .name = */      action["tool_name"],405                /* .arguments = */ action["parameters"].dump(),406                /* .id = */        action["tool_call_id"],407            });408        }409    } else {410        LOG_ERR("Failed to parse command_r output");411        result.content = input;412    }413    return result;414}415 416static void expect_tool_parameters(const std::string & name, const json & parameters, const std::vector<std::string> & expected_properties) {417    if (!parameters.is_object() || !parameters.contains("type") || parameters["type"] != "object" || !parameters.contains("properties") || !parameters.contains("required")) {418        throw std::runtime_error("Parameters of tool " + name + " must be an object w/ required properties");419    }420    const auto & parameters_properties = parameters.at("properties");421    const auto & parameters_required = parameters.at("required");422    for (const auto & prop : expected_properties) {423        if (!parameters_properties.contains(prop)) {424            throw std::runtime_error("Parameters of tool " + name + " is missing property: " + prop);425        }426        if (std::find(parameters_required.begin(), parameters_required.end(), json(prop)) == parameters_required.end()) {427            throw std::runtime_error("Parameters of tool " + name + " must have property marked as required: " + prop);428        }429    }430    if (parameters_properties.size() != expected_properties.size()) {431        throw std::runtime_error("Parameters of tool " + name + " must only have these properties:" + string_join(expected_properties, ", "));432    }433}434 435static common_chat_params common_chat_params_init_llama_3_1_tool_calls(const common_chat_template & tmpl, const struct common_chat_inputs & inputs, bool allow_python_tag_builtin_tools) {436    auto builtin_tools = json::array();437    common_chat_params data;438    data.grammar_lazy = inputs.tool_choice != "required";439    data.grammar = build_grammar([&](const common_grammar_builder & builder) {440        std::vector<std::string> tool_rules;441 442        auto handle_builtin_tool = [&](const std::string & name, const json & parameters) {443            if (name == "wolfram_alpha") {444                // https://github.com/meta-llama/llama-stack/blob/main/llama_stack/providers/remote/tool_runtime/wolfram_alpha/wolfram_alpha.py445                expect_tool_parameters(name, parameters, {"query"});446            } else if (name == "web_search" || name == "brave_search") {447                // https://github.com/meta-llama/llama-stack/blob/main/llama_stack/providers/remote/tool_runtime/brave_search/brave_search.py448                expect_tool_parameters(name, parameters, {"query"});449            } else if (name == "python" || name == "code_interpreter") {450                // https://github.com/meta-llama/llama-stack/blob/main/llama_stack/providers/inline/tool_runtime/code_interpreter/code_interpreter.py451                expect_tool_parameters(name, parameters, {"code"});452            } else {453                return false;454            }455 456            std::vector<std::string> kvs;457            for (const auto & [key, value] : parameters.at("properties").items()) {458                kvs.push_back("\"" + key + "=\" " + builder.add_schema(name + "-args-" + key, value));459            }460 461            tool_rules.push_back(462                builder.add_rule(463                    name + "-call",464                    "\"<|python_tag|>" + name + ".call(\" " + string_join(kvs, " \", \" ") + " \")\""));465            builtin_tools.push_back(name);466 467            return true;468        };469 470        foreach_function(inputs.tools, [&](const json & tool) {471            const auto & function = tool["function"];472            std::string name = function["name"];473            auto parameters = function["parameters"];474            builder.resolve_refs(parameters);475 476            // https://github.com/meta-llama/llama-stack/tree/main/llama_stack/providers/remote/tool_runtime477            if (allow_python_tag_builtin_tools) {478                handle_builtin_tool(name, parameters);479            }480            tool_rules.push_back(481                builder.add_rule(482                    name + "-call",483                    "\"{\" space "484                    "( \"\\\"type\\\":\" space \"\\\"function\\\",\" space )? "485                    "\"\\\"name\\\": \\\"" + name + "\\\", \\\"parameters\\\": \" " +486                        builder.add_schema(name + "-args", parameters) +487                    " \"}\""));488            data.grammar_triggers.push_back({"{\"name\": \"" + name + "\"", /* .at_start = */ true});489        });490        data.grammar_triggers.push_back({"{\"name\":", /* .at_start = */ true});491        data.grammar_triggers.push_back({"{\n  \"name\":", /* .at_start = */ true});492        data.grammar_triggers.push_back({"{\n    \"name\":", /* .at_start = */ true});493        data.grammar_triggers.push_back({"{\"type\": \"function\"", /* .at_start = */ true});494        data.grammar_triggers.push_back({"{\n  \"type\": \"function\"", /* .at_start = */ true});495        data.grammar_triggers.push_back({"{\n    \"type\": \"function\"", /* .at_start = */ true});496        if (!builtin_tools.empty()) {497            data.grammar_triggers.push_back({"<|python_tag|>", /* .at_start = */ false});498        }499        builder.add_rule("root", string_join(tool_rules, " | "));500    }, grammar_options);501    data.additional_stops.push_back("<|eom_id|>");502    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt, {503        {"tools_in_user_message", false},504        {"builtin_tools", builtin_tools.empty() ? json() : builtin_tools},505    });506    data.format = allow_python_tag_builtin_tools && !builtin_tools.empty()507        ? COMMON_CHAT_FORMAT_LLAMA_3_X_WITH_BUILTIN_TOOLS508        : COMMON_CHAT_FORMAT_LLAMA_3_X;509    return data;510}511static common_chat_msg common_chat_parse_llama_3_1(const std::string & input, bool with_builtin_tools = false) {512    // TODO: tighten & simplify the parser, don't accept leading text context.513    static std::regex function_regex("\\{[\\s\\n\\r]*(?:\"type\"[\\s\\n\\r]*:[\\s\\n\\r]*\"function\"[\\s\\n\\r]*,[\\s\\n\\r]*|[\\s\\n\\r]*)\"name\"[\\s\\n\\r]*:[\\s\\n\\r]*\"([^\"]+)\"[\\s\\n\\r]*,[\\s\\n\\r]*\"parameters\": ");514    static std::regex close_regex("\\}");515    static std::regex builtin_call_regex("<\\|python_tag\\|>([^.(]+)\\.call\\((.*)\\)");516 517    if (with_builtin_tools) {518        std::smatch match;519        if (std::regex_match(input, match, builtin_call_regex)) {520            auto name = match[1].str();521            auto raw_args = match[2].str();522 523            // TODO: if/when builtin tools start accepting more than 1 argument, use parse_json for real parsing.524            auto it_eq = raw_args.find('=');525            auto arg_name = raw_args.substr(0, it_eq);526            auto arg_value_str = raw_args.substr(it_eq + 1);527            auto arg_value = json::parse(arg_value_str);528 529            return {530                /* .role = */ "assistant",531                /* .content = */ match.prefix().str(),532                /* .tool_calls = */ {533                    {534                        /* .name = */ match[1],535                        /* .arguments = */ (json {536                            {arg_name, arg_value},537                        }).dump(),538                        /* .id = */ "",539                    },540                },541            };542        }543    }544    return parse_json_tool_calls(input, std::nullopt, function_regex, close_regex);545}546 547static common_chat_params common_chat_params_init_deepseek_r1(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {548    common_chat_params data;549    data.grammar_lazy = inputs.tool_choice != "required";550    data.grammar = build_grammar([&](const common_grammar_builder & builder) {551        std::vector<std::string> tool_rules;552        foreach_function(inputs.tools, [&](const json & tool) {553            const auto & function = tool["function"];554            std::string name = function["name"];555            auto parameters = function["parameters"];556            auto args_rule = builder.add_schema(name + "-args", parameters);557            tool_rules.push_back(builder.add_rule(name + "-call",558                "\"<|tool▁call▁begin|>function<|tool▁sep|>" + name + "\\n```json\\n\" " + args_rule + " \"```<|tool▁call▁end|>\""));559        });560        data.grammar_triggers.push_back({"<|tool▁calls▁begin|>", /* .at_start = */ false});561        data.preserved_tokens = {562            "<|tool▁sep|>",563            "<|tool▁call▁end|>",564        };565        builder.add_rule("root", "\"<|tool▁calls▁begin|>\" (" + string_join(tool_rules, " | ") + ")" + (inputs.parallel_tool_calls ? "*" : "") + " space");566    }, grammar_options);567    auto prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);568    data.prompt = prompt;569    data.format = COMMON_CHAT_FORMAT_DEEPSEEK_R1;570    return data;571}572static common_chat_msg common_chat_parse_deepseek_r1(const std::string & input) {573    static std::regex trigger_regex("<|tool▁calls▁begin|>");574    static std::regex function_regex("<|tool▁call▁begin|>function<|tool▁sep|>([^\n]+)\n```json\n");575    static std::regex close_regex("```<|tool▁call▁end|>");576    return parse_json_tool_calls(input, trigger_regex, function_regex, close_regex);577}578 579static common_chat_params common_chat_params_init_firefunction_v2(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {580    fprintf(stderr, "%s\n", __func__);581    common_chat_params data;582    data.prompt = apply(tmpl, inputs.messages, /* tools= */ nullptr, inputs.add_generation_prompt, {583        {"datetime", "Jan 29 2025 13:00:00 GMT"},584        {"functions", json(inputs.tools.empty() ? "" : inputs.tools.dump(2))},585    });586    if (!inputs.tools.is_null() && !inputs.tools.empty()) {587        data.grammar_lazy = inputs.tool_choice != "required";588        data.grammar = build_grammar([&](const common_grammar_builder & builder) {589            auto schemas = json::array();590            foreach_function(inputs.tools, [&](const json & tool) {591                const auto & function = tool["function"];592                schemas.push_back({593                    {"type", "object"},594                    {"properties", {595                        {"name", {596                            {"type", "string"},597                            {"const", function["name"]},598                        }},599                        {"arguments", function["parameters"]},600                    }},601                    {"required", json::array({"name", "arguments", "id"})},602                });603            });604            auto schema = json {605                {"type", "array"},606                {"items", schemas.size() == 1 ? schemas[0] : json {{"anyOf", schemas}}},607                {"minItems", 1},608            };609            if (!inputs.parallel_tool_calls) {610                schema["maxItems"] = 1;611            }612            builder.add_rule("root", "\" functools\"? " + builder.add_schema("tool_calls", schema));613        }, grammar_options);614        data.grammar_triggers.push_back({" functools[", /* .at_start = */ false});615        data.format = COMMON_CHAT_FORMAT_FIREFUNCTION_V2;616    } else {617        data.format = COMMON_CHAT_FORMAT_CONTENT_ONLY;618    }619    return data;620}621static common_chat_msg common_chat_parse_firefunction_v2(const std::string & input) {622    return parse_prefixed_json_tool_call_array(input, " functools[", /* rstrip_prefix= */ 1);623}624 625static common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {626    // >>>all\nlet's call functions>>>fn1\n{"arg1": 1...}\n>>>fn2\n{"arg1": 1...}...627    // Using ">>>f1\n", ">>>f2\n"... as trigger words for the grammar628    common_chat_params data;629    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);630    data.format = COMMON_CHAT_FORMAT_FUNCTIONARY_V3_2;631    if (!inputs.tools.is_null() && !inputs.tools.empty()) {632        data.grammar_lazy = inputs.tool_choice != "required";633        data.grammar = build_grammar([&](const common_grammar_builder & builder) {634            std::vector<std::string> first_tool_rules;635            std::vector<std::string> subsequent_tool_rules;636            foreach_function(inputs.tools, [&](const json & tool) {637                const auto & function = tool["function"];638                std::string name = function["name"];639                auto parameters = function["parameters"];640                auto args_rule = builder.add_schema(name + "-args", parameters);641                first_tool_rules.push_back(builder.add_rule(name + "-call", "\"" + name + "\\n\" " + args_rule));642                subsequent_tool_rules.push_back(builder.add_rule(name + "-call2", "\">>>" + name + "\\n\" " + args_rule));643                data.grammar_triggers.push_back({name, /* .at_start = */ true});644                data.grammar_triggers.push_back({">>>" + name, /* .at_start = */ false});645            });646            auto first_rule = first_tool_rules.empty() ? "" : builder.add_rule("first_tool_call", string_join(first_tool_rules, " | ")) + " space";647            if (inputs.parallel_tool_calls) {648                auto subsequent_rule = builder.add_rule("subsequent_tool_call", string_join(subsequent_tool_rules, " | ")) + " space";649                builder.add_rule("root", first_rule + " (" + subsequent_rule + ")*");650            } else {651                builder.add_rule("root", first_rule);652            }653 654        }, grammar_options);655    }656    return data;657}658 659static bool consume(std::string::const_iterator & it, const std::string::const_iterator & end, const std::string & expected) {660    auto expected_it = expected.begin();661    auto tmp_it = it;662    while (tmp_it != end && expected_it != expected.end() && *tmp_it == *expected_it) {663        ++tmp_it;664        ++expected_it;665    }666    if (expected_it == expected.end()) {667        it = tmp_it;668        return true;669    }670    return false;671}672 673static common_chat_msg common_chat_parse_functionary_v3_2(const std::string & input) {674    static std::regex function_regex(R"((?:>>>)?(\w+)\n)");675    static std::regex close_regex(R"($|(?=>>>))");676 677    std::string content;678    auto it = input.begin();679    const auto end = input.end();680 681    if (consume(it, end, "all\n")) {682        std::smatch match;683        if (std::regex_search(it, end, match, function_regex)) {684            auto fun_it = match.prefix().second;685            content = std::string(it, fun_it);686            it = fun_it;687        } else {688            common_chat_msg res;689            res.role = "assistant";690            res.content = std::string(it, end);691            return res;692        }693    }694    // TODO: tighten & simplify.695    try {696        auto res = parse_json_tool_calls(std::string(it, end), std::nullopt, function_regex, close_regex);697        res.content = content + res.content;698        return res;699    } catch (const std::exception & e) {700        LOG_ERR("Failed to parse functionary v3.2 input: %s\n", e.what());701        common_chat_msg res;702        res.role = "assistant";703        res.content = input;704        return res;705    }706}707 708static common_chat_params common_chat_params_init_functionary_v3_1_llama_3_1(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {709    // https://github.com/MeetKai/functionary/blob/main/tests/prompt_test_v3-llama3.1.txt710    common_chat_params data;711    json tools = inputs.tools.is_null() ? inputs.tools : json::array();712    std::string python_code_argument_name;713    auto has_raw_python = false;714 715    data.grammar_lazy = inputs.tool_choice != "required";716    data.grammar = build_grammar([&](const common_grammar_builder & builder) {717        std::vector<std::string> tool_rules;718        foreach_function(inputs.tools, [&](const json & tool) {719            const auto & function = tool["function"];720            const auto & parameters = function["parameters"];721            std::string name = function["name"];722            if (name == "python" || name == "ipython") {723                if (!parameters.contains("type")) {724                    throw std::runtime_error("Missing type in python tool");725                }726                has_raw_python = true;727                auto type = parameters.at("type");728                if (type == "object") {729                    auto properties = parameters.at("properties");730                    for (auto it = properties.begin(); it != properties.end(); ++it) {731                        if (it.value().at("type") == "string") {732                            if (!python_code_argument_name.empty()) {733                                throw std::runtime_error("Multiple string arguments found in python tool");734                            }735                            python_code_argument_name = it.key();736                        }737                    }738                    if (python_code_argument_name.empty()) {739                        throw std::runtime_error("No string argument found in python tool");740                    }741                } else if (type != "string") {742                    throw std::runtime_error("Invalid type in python tool: " + type.dump());743                }744            }745            tool_rules.push_back(builder.add_rule(name + "-call", "\"<function=" + name + ">\" " + builder.add_schema(name + "-args", parameters) + " \"</function>\" space"));746        });747        if (has_raw_python) {748            tool_rules.push_back(builder.add_rule("python-call", "\"<|python_tag|>\" .*"));749            data.grammar_triggers.push_back({"<|python_tag|>", /* .at_start = */ false});750        }751        auto tool_call = builder.add_rule("tool_call", string_join(tool_rules, " | ")) + " space";752        builder.add_rule("root", inputs.parallel_tool_calls ? "(" + tool_call + ")+" : tool_call);753        data.grammar_triggers.push_back({"<function=", /* .at_start = */ false});754    }, grammar_options);755 756    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);757    // TODO: if (has_raw_python)758    data.format = COMMON_CHAT_FORMAT_FUNCTIONARY_V3_1_LLAMA_3_1;759    return data;760}761static common_chat_msg common_chat_parse_functionary_v3_1_llama_3_1(const std::string & input) {762    // This version of Functionary still supports the llama 3.1 tool call format for the python tool.763    static std::regex python_tag_regex(R"(<\|python_tag\|>([\s\S\n]*)$)");764    std::smatch match;765    if (std::regex_search(input, match, python_tag_regex)) {766        auto code = match[1].str();767        return {768            /* .role = */ "assistant",769            /* .content = */ match.prefix().str(),770            /* .tool_calls = */ {771                {772                    /* .name = */ "python",773                    /* .arguments = */ (json {{"code", code}}).dump(),774                    /* .id = */ "",775                },776            }777        };778    }779    static std::regex function_regex(R"(<function=(\w+)>)");780    static std::regex close_regex(R"(</function>)");781    // TODO: tighten & simplify.782    return parse_json_tool_calls(input, std::nullopt, function_regex, close_regex);783}784 785static common_chat_params common_chat_params_init_hermes_2_pro(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {786    common_chat_params data;787    // (content)?(<tool_call>{"name": "foo", "arguments": {"a": 1}}</tool_call>)*788    data.grammar_lazy = inputs.tool_choice != "required";789    data.grammar = build_grammar([&](const common_grammar_builder & builder) {790        std::vector<std::string> tool_rules;791        foreach_function(inputs.tools, [&](const json & tool) {792            const auto & function = tool["function"];793            std::string name = function["name"];794            auto parameters = function["parameters"];795            builder.resolve_refs(parameters);796            tool_rules.push_back(builder.add_schema(name + "-call", {797                {"type", "object"},798                {"properties", json {799                    {"name", json {{"const", name}}},800                    {"arguments", parameters},801                }},802                {"required", json::array({"name", "arguments"})},803            }));804        });805        auto tool_call = "\"<tool_call>\" space " + builder.add_rule("tool_call", string_join(tool_rules, " | ")) + " \"</tool_call>\" space";806        builder.add_rule("root", inputs.parallel_tool_calls ? "(" + tool_call + ")+" : tool_call);807        data.grammar_triggers.push_back({"<tool_call>", /* .at_start = */ false});808        data.preserved_tokens = { "</tool_call>" };809    }, grammar_options);810 811    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);812    data.format = COMMON_CHAT_FORMAT_HERMES_2_PRO;813    return data;814}815static common_chat_msg common_chat_parse_hermes_2_pro(const std::string & input) {816    try {817        std::regex start_pattern(R"([\n\s]*<tool_call>)");818        std::regex middle_pattern(R"([\n\s]*</tool_call>[\n\s]*<tool_call>)");819        std::regex end_pattern(R"([\n\s]*</tool_call>[\n\s]*$)");820 821        auto end = input.end();822        std::sregex_iterator rend;823        std::sregex_iterator rit(input.begin(), end, start_pattern);824        if (rit == rend) {825            return {826                /* .role = */ "assistant",827                /* .content = */ input,828                /* .tool_calls = */ {},829            };830        }831 832        common_chat_msg result;833        result.role = "assistant";834        result.content = rit->prefix();835 836        auto it = rit->suffix().first;837        while (it != end) {838            json call;839            if (!parse_json(it, end, call)) {840                throw std::runtime_error("Failed to parse json tool call");841            }842            const auto & arguments = call["arguments"];843            result.tool_calls.push_back({844                call["name"],845                arguments.dump(),846                // arguments.is_string() ? arguments.get<std::string>() : arguments.dump(),847                /* id= */ "",848            });849            rit = {it, end, middle_pattern};850            if (rit != rend) {851                it = rit->suffix().first;852            } else {853                rit = {it, end, end_pattern};854                if (rit == rend) {855                    throw std::runtime_error("Malformed input, missing </tool_call>");856                }857                break;858            }859        }860        return result;861    } catch (const std::exception & e) {862        return {863            /* .role = */ "assistant",864            /* .content = */ input,865            /* .tool_calls = */ {},866        };867    }868}869 870static common_chat_params common_chat_params_init_without_tools(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {871    common_chat_params data;872    data.prompt = apply(tmpl, inputs.messages, inputs.tools.empty() ? json() : inputs.tools, inputs.add_generation_prompt);873    data.format = COMMON_CHAT_FORMAT_CONTENT_ONLY;874    data.grammar_lazy = false;875    if (!inputs.json_schema.is_null()) {876        if (!inputs.grammar.empty()) {877            throw std::runtime_error("Either \"json_schema\" or \"grammar\" can be specified, but not both");878        }879        data.grammar = json_schema_to_grammar(inputs.json_schema);880    } else {881        data.grammar = inputs.grammar.empty();882    }883    return data;884}885 886common_chat_params common_chat_params_init(const common_chat_template & tmpl, const struct common_chat_inputs & inputs) {887    auto has_tools = !inputs.tools.is_null() && inputs.tool_choice != "none";888    LOG_DBG("[%s] has_tools=%s\n", __func__, has_tools ? "true" : "false");889 890    if (has_tools && !inputs.grammar.empty()) {891        throw std::runtime_error("Cannot specify grammar with tools");892    }893 894    const auto & src = tmpl.source();895    if (src.find(">>>all") != std::string::npos) {896        // Functionary prepends "all\n" to plain content outputs, so we use the parser no matter when897        return common_chat_params_init_functionary_v3_2(tmpl, inputs);898    }899    if (src.find(" functools[") != std::string::npos) {900        // Firefunction v2 requires datetime and functions in the context, even w/o tools.901        return common_chat_params_init_firefunction_v2(tmpl, inputs);902    }903 904    if (!has_tools) {905        return common_chat_params_init_without_tools(tmpl, inputs);906    }907 908    if (src.find("<tool_call>") != std::string::npos) {909        return common_chat_params_init_hermes_2_pro(tmpl, inputs);910    }911    if (src.find("<|start_header_id|>") != std::string::npos912        && src.find("<function=") != std::string::npos) {913        return common_chat_params_init_functionary_v3_1_llama_3_1(tmpl, inputs);914    }915    if (src.find("<|start_header_id|>ipython<|end_header_id|>") != std::string::npos) {916        auto allow_python_tag_builtin_tools = src.find("<|python_tag|>") != std::string::npos;917        return common_chat_params_init_llama_3_1_tool_calls(tmpl, inputs, allow_python_tag_builtin_tools);918    }919    if (src.find("<|tool▁calls▁begin|>") != std::string::npos) {920        return common_chat_params_init_deepseek_r1(tmpl, inputs);921    }922    if (src.find("[TOOL_CALLS]") != std::string::npos) {923        return common_chat_params_init_mistral_nemo(tmpl, inputs);924    }925    if (src.find("<|END_THINKING|><|START_ACTION|>") != std::string::npos) {926        return common_chat_params_init_command_r7b(tmpl, inputs);927    }928    return common_chat_params_init_generic(tmpl, inputs);929}930 931static common_chat_msg common_chat_parse_content_only(const std::string & input) {932    return {933        /* .role = */ "assistant",934        /* .content = */ input,935        /* .tool_calls = */ {},936    };937}938 939common_chat_msg common_chat_parse(const std::string & input, common_chat_format format) {940    switch (format) {941        case COMMON_CHAT_FORMAT_CONTENT_ONLY:942            return common_chat_parse_content_only(input);943        case COMMON_CHAT_FORMAT_GENERIC:944            return common_chat_parse_generic(input);945        case COMMON_CHAT_FORMAT_MISTRAL_NEMO:946            return common_chat_parse_mistral_nemo(input);947        case COMMON_CHAT_FORMAT_LLAMA_3_X:948            return common_chat_parse_llama_3_1(input);949        case COMMON_CHAT_FORMAT_LLAMA_3_X_WITH_BUILTIN_TOOLS:950            return common_chat_parse_llama_3_1(input, /* with_builtin_tools= */ true);951        case COMMON_CHAT_FORMAT_DEEPSEEK_R1:952            return common_chat_parse_deepseek_r1(input);953        case COMMON_CHAT_FORMAT_FUNCTIONARY_V3_2:954            return common_chat_parse_functionary_v3_2(input);955        case COMMON_CHAT_FORMAT_FUNCTIONARY_V3_1_LLAMA_3_1:956            return common_chat_parse_functionary_v3_1_llama_3_1(input);957        case COMMON_CHAT_FORMAT_HERMES_2_PRO:958            return common_chat_parse_hermes_2_pro(input);959        case COMMON_CHAT_FORMAT_FIREFUNCTION_V2:960            return common_chat_parse_firefunction_v2(input);961        case COMMON_CHAT_FORMAT_COMMAND_R7B:962            return common_chat_parse_command_r7b(input);963        default:964            throw std::runtime_error("Unsupported format: " + common_chat_format_name(format));965    }966}967