KBaba7/llama.cpp
0
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 