echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
1#include "chat-peg-parser.h"2#include "chat.h"3#include "common.h"4#include "json-schema-to-grammar.h"5#include "peg-parser.h"6#include "testing.h"7#include "peg-parser/simple-tokenize.h"8 9#include <iostream>10#include <numeric>11#include <string>12 13#include "nlohmann/json.hpp"14 15using json = nlohmann::ordered_json;16 17static json create_tools();18static void test_example_native(testing & t);19static void test_example_qwen3_coder(testing & t);20static void test_example_qwen3_non_coder(testing & t);21static void test_command7_parser_compare(testing & t);22static void test_prefix_tool_names(testing & t);23static void test_tagged_peg_parser(testing & t);24 25int main(int argc, char * argv[]) {26 testing t(std::cout);27 if (argc >= 2) {28 t.set_filter(argv[1]);29 }30 31 const char * verbose = getenv("LLAMA_TEST_VERBOSE");32 if (verbose) {33 t.verbose = std::string(verbose) == "1";34 }35 36 t.test("native", test_example_native);37 t.test("qwen3 coder", test_example_qwen3_coder);38 t.test("qwen3 non-coder", test_example_qwen3_non_coder);39 t.test("comparison", test_command7_parser_compare);40 t.test("prefix tool names", test_prefix_tool_names);41 t.test("tagged peg parser", test_tagged_peg_parser);42 43 return t.summary();44}45 46static json create_tools() {47 json tools = json::array();48 49 json tool_weather = {50 { "type", "function" },51 { "function",52 {53 { "name", "get_current_weather" },54 { "description", "Get the current weather in a given location" },55 { "parameters",56 {57 { "type", "object" },58 { "properties",59 { { "location",60 { { "type", "string" }, { "description", "The city and state, e.g. San Francisco, CA" } } },61 { "unit",62 { { "type", "string" },63 { "enum", { "celsius", "fahrenheit" } },64 { "description",65 "The temperature unit to use. Infer this from the users location." } } } } },66 { "required", { "location", "unit" } },67 } },68 } }69 };70 tools.push_back(tool_weather);71 72 json tool_forecast = {73 { "type", "function" },74 { "function",75 {76 { "name", "get_forecast" },77 { "description", "Get the weather forecast for a given location" },78 { "parameters",79 {80 { "type", "object" },81 { "properties",82 { { "location",83 { { "type", "string" }, { "description", "The city and state, e.g. San Francisco, CA" } } },84 { "unit",85 { { "type", "string" },86 { "enum", { "celsius", "fahrenheit" } },87 { "description", "The temperature unit to use. Infer this from the users location." } } },88 { "days",89 { { "type", "integer" },90 { "description", "Number of days to forecast (1-10)" },91 { "minimum", 1 },92 { "maximum", 10 } } } } },93 { "required", { "location", "unit" } },94 } },95 } }96 };97 tools.push_back(tool_forecast);98 99 json tool_search = {100 { "type", "function" },101 { "function",102 { { "name", "search_knowledge_base" },103 { "description", "Search the internal technical documentation knowledge base." },104 { "parameters",105 { { "type", "object" },106 { "properties",107 { { "query", { { "type", "string" }, { "description", "The search query string." } } },108 { "max_results",109 { { "type", "integer" },110 { "description", "The maximum number of results to return." },111 { "default", 5 } } },112 { "category",113 { { "type", "string" },114 { "enum", { "api", "troubleshooting", "billing", "general" } },115 { "description", "Filter search by specific category." } } } } },116 { "required", { "query", "category" } },117 { "additionalProperties", false } } },118 { "strict", true } } }119 };120 tools.push_back(tool_search);121 122 return tools;123}124 125struct tool_argument {126 std::string name;127 std::string type;128 bool is_required;129 json schema;130};131 132struct tool_definition {133 std::string name;134 std::vector<tool_argument> arguments;135 json schema;136};137 138// Test fictitious model output that emits arguments as JSON.139static void test_example_native(testing & t) {140 struct test_case {141 // Parameters142 std::string name;143 json tools;144 common_chat_tool_choice tool_choice;145 common_reasoning_format reasoning_format;146 json json_schema;147 bool parallel_tool_calls;148 std::string generation_prompt;149 std::string input;150 151 // Expect152 std::string expect_reasoning;153 std::string expect_content;154 std::vector<common_chat_tool_call> expect_tool_calls;155 };156 157 auto build_parser = [](const test_case & tc) {158 return build_chat_peg_parser([&](common_chat_peg_builder & p) {159 auto reasoning_in_content = (tc.reasoning_format == COMMON_REASONING_FORMAT_NONE);160 // Always use optional TAG_BASED pattern; generation_prompt is prepended to input161 auto reasoning = p.optional("<think>" + p.reasoning(p.until("</think>")) + "</think>" + p.space());162 163 // tool calling parser164 if (tc.tools.is_array() && !tc.tools.empty()) {165 auto tool_call =166 p.standard_json_tools("<tool_call>[", "]</tool_call>", tc.tools, tc.parallel_tool_calls,167 tc.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED);168 169 return p.sequence({ (reasoning_in_content ? p.eps() : reasoning), p.content(p.until("<tool_call>")),170 p.optional(p.space() + tool_call), p.space(), p.end() });171 }172 173 // response_format parser174 if (tc.json_schema.is_object() && !tc.json_schema.empty()) {175 return p.sequence({ (reasoning_in_content ? p.eps() : reasoning),176 p.content(p.schema(p.json(), "response-output", tc.json_schema)), p.space(),177 p.end() });178 }179 180 // Content-only parser181 return p.sequence({ (reasoning_in_content ? p.eps() : reasoning), p.content(p.rest()), p.end() });182 });183 };184 185 std::vector<test_case> test_cases = std::vector<test_case>{186 {187 /* .name = */ "content with reasoning (no generation_prompt)",188 /* .tools = */ {},189 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,190 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,191 /* .json_schema = */ {},192 /* .parallel_tool_calls = */ false,193 /* .generation_prompt = */ "",194 /* .input = */ ("<think>The user said hello, I must say hello back</think>\nHello"),195 /* .expect_reasoning = */ "The user said hello, I must say hello back",196 /* .expect_content = */ "Hello",197 /* .expect_tool_calls = */ {},198 },199 {200 /* .name = */ "content without reasoning (no generation_prompt)",201 /* .tools = */ {},202 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,203 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,204 /* .json_schema = */ {},205 /* .parallel_tool_calls = */ false,206 /* .generation_prompt = */ "",207 /* .input = */ ("Hello"),208 /* .expect_reasoning = */ "",209 /* .expect_content = */ "Hello",210 /* .expect_tool_calls = */ {},211 },212 {213 /* .name = */ "content with reasoning_format = none (tags appear in content)",214 /* .tools = */ {},215 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,216 /* .reasoning_format = */ COMMON_REASONING_FORMAT_NONE,217 /* .json_schema = */ {},218 /* .parallel_tool_calls = */ false,219 /* .generation_prompt = */ "",220 /* .input = */ ("<think>The user said hello, I must say hello back</think>\nHello"),221 /* .expect_reasoning = */ "",222 /* .expect_content = */ "<think>The user said hello, I must say hello back</think>\nHello",223 /* .expect_tool_calls = */ {},224 },225 {226 /* .name = */ "content with reasoning generation_prompt",227 /* .tools = */ {},228 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,229 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,230 /* .json_schema = */ {},231 /* .parallel_tool_calls = */ false,232 /* .generation_prompt = */ "<think>",233 /* .input = */ ("The user said hello, I must say hello back</think>\nHello"),234 /* .expect_reasoning = */ "The user said hello, I must say hello back",235 /* .expect_content = */ "Hello",236 /* .expect_tool_calls = */ {},237 },238 {239 /* .name = */ "content with reasoning generation_prompt and reasoning_format = none",240 /* .tools = */ {},241 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,242 /* .reasoning_format = */ COMMON_REASONING_FORMAT_NONE,243 /* .json_schema = */ {},244 /* .parallel_tool_calls = */ false,245 /* .generation_prompt = */ "",246 /* .input = */ ("The user said hello, I must say hello back</think>\nHello"),247 /* .expect_reasoning = */ "",248 /* .expect_content = */ "The user said hello, I must say hello back</think>\nHello",249 /* .expect_tool_calls = */ {},250 },251 {252 /* .name = */ "content with closed reasoning generation_prompt (empty reasoning discarded)",253 /* .tools = */ {},254 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,255 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,256 /* .json_schema = */ {},257 /* .parallel_tool_calls = */ false,258 /* .generation_prompt = */ "<think></think>",259 /* .input = */ ("Hello"),260 /* .expect_reasoning = */ "",261 /* .expect_content = */ "Hello",262 /* .expect_tool_calls = */ {},263 },264 {265 /* .name = */ "tools with reasoning generation_prompt",266 /* .tools = */ create_tools(),267 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_AUTO,268 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,269 /* .json_schema = */ {},270 /* .parallel_tool_calls = */ false,271 /* .generation_prompt = */ "<think>",272 /* .input = */273 ("I must get the weather in New York</think>\n"274 "<tool_call>["275 R"({"name": "get_current_weather", "arguments": {"location": "New York City, NY", "unit": "fahrenheit"}})"276 "]</tool_call>"),277 /* .expect_reasoning = */ "I must get the weather in New York",278 /* .expect_content = */ "",279 /* .expect_tool_calls = */280 { {281 /* .name = */ "get_current_weather",282 /* .arguments = */ R"({"location": "New York City, NY", "unit": "fahrenheit"})",283 /* .id = */ "",284 } },285 },286 {287 /* .name = */ "parallel tools with reasoning generation_prompt",288 /* .tools = */ create_tools(),289 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_AUTO,290 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,291 /* .json_schema = */ {},292 /* .parallel_tool_calls = */ true,293 /* .generation_prompt = */ "<think>",294 /* .input = */295 ("I must get the weather in New York and San Francisco and a 3 day forecast of each.</think>\nLet me "296 "search that for you."297 "<tool_call>["298 R"({"name": "get_current_weather", "arguments": {"location": "New York City, NY", "unit": "fahrenheit"}})"299 ", "300 R"({"name": "get_current_weather", "arguments": {"location": "San Francisco, CA", "unit": "fahrenheit"}})"301 ", "302 R"({"name": "get_forecast", "arguments": {"location": "New York City, NY", "unit": "fahrenheit", "days": 3}})"303 ", "304 R"({"name": "get_forecast", "arguments": {"location": "San Francisco, CA", "unit": "fahrenheit", "days": 3}})"305 "]</tool_call>"),306 /* .expect_reasoning = */307 "I must get the weather in New York and San Francisco and a 3 day forecast of each.", /* .expect_content = */ "Let me search that for you.",308 /* .expect_tool_calls = */309 { {310 /* .name = */ "get_current_weather",311 /* .arguments = */ R"({"location": "New York City, NY", "unit": "fahrenheit"})",312 /* .id = */ "",313 },314 {315 /* .name = */ "get_current_weather",316 /* .arguments = */ R"({"location": "San Francisco, CA", "unit": "fahrenheit"})",317 /* .id = */ "",318 },319 {320 /* .name = */ "get_forecast",321 /* .arguments = */ R"({"location": "New York City, NY", "unit": "fahrenheit", "days": 3})",322 /* .id = */ "",323 },324 {325 /* .name = */ "get_forecast",326 /* .arguments = */ R"({"location": "San Francisco, CA", "unit": "fahrenheit", "days": 3})",327 /* .id = */ "",328 } },329 },330 {331 /* .name = */ "response_format with reasoning generation_prompt",332 /* .tools = */ {},333 /* .tool_choice = */ COMMON_CHAT_TOOL_CHOICE_NONE,334 /* .reasoning_format = */ COMMON_REASONING_FORMAT_AUTO,335 /* .json_schema = */336 { { "type", "object" },337 { "properties",338 { { "invoice_number", { { "type", "string" } } },339 { "amount", { { "type", "number" } } },340 { "due_date", { { "type", "string" } } } } },341 { "required", { "invoice_number", "amount", "due_date" } } },342 /* .parallel_tool_calls = */ false,343 /* .generation_prompt = */ "<think>",344 /* .input = */345 ("I must produce the invoice in the requested format</think>\n"346 R"({"invoice_number": "INV-2025-001", "amount": 1250.50, "due_date": "2025-12-31"})"),347 /* .expect_reasoning = */ "I must produce the invoice in the requested format",348 /* .expect_content = */349 R"({"invoice_number": "INV-2025-001", "amount": 1250.50, "due_date": "2025-12-31"})", /* .expect_tool_calls = */ {},350 },351 };352 353 for (const auto & tc : test_cases) {354 t.test(tc.name, [&](testing & t) {355 auto parser = build_parser(tc);356 auto lazy = !tc.tools.empty() && tc.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;357 auto grammar = build_grammar([&](const common_grammar_builder & builder) {358 for (const auto & def : tc.tools) {359 auto function = def.at("function");360 auto parameters = function.at("parameters");361 builder.resolve_refs(parameters);362 };363 parser.build_grammar(builder, lazy);364 });365 366 t.log("Grammar:");367 for (const auto & line : string_split(grammar, "\n")) {368 t.log(line);369 }370 371 std::string effective_input = tc.generation_prompt + tc.input;372 common_peg_parse_context ctx(effective_input);373 auto result = parser.parse(ctx);374 375 t.assert_true("success", result.success());376 377 common_chat_msg msg;378 auto mapper = common_chat_peg_mapper(msg);379 mapper.from_ast(ctx.ast, result);380 381 t.assert_equal("content equal", tc.expect_content, msg.content);382 t.assert_equal("reasoning equal", tc.expect_reasoning, msg.reasoning_content);383 t.assert_equal("number of tool calls", tc.expect_tool_calls.size(), msg.tool_calls.size());384 for (auto i = 0u; i < std::min(tc.expect_tool_calls.size(), msg.tool_calls.size()); i++) {385 t.assert_equal("tool name", tc.expect_tool_calls[i].name, msg.tool_calls[i].name);386 t.assert_equal("tool args", tc.expect_tool_calls[i].arguments, msg.tool_calls[i].arguments);387 }388 });389 }390}391 392static void test_example_qwen3_coder(testing & t) {393 auto tools = create_tools();394 auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {395 auto content = p.rule("content", p.content(p.until("<tool_call>")));396 397 std::vector<common_peg_parser> tool_parsers;398 for (const auto & def : tools) {399 auto function = def.at("function");400 std::string name = function.at("name");401 auto parameters = function.at("parameters");402 auto properties = parameters.at("properties");403 404 std::set<std::string> required_properties;405 if (function.contains("required")) {406 function.at("required").get_to(required_properties);407 }408 409 std::vector<common_peg_parser> arg_parsers;410 for (const auto & [param_name, param_schema] : properties.items()) {411 bool is_required = required_properties.find(param_name) != required_properties.end();412 auto type = param_schema.value("type", "object");413 414 auto arg = p.tool_arg(415 p.sequence({ p.tool_arg_open("<parameter=" + p.tool_arg_name(p.literal(param_name)) + ">"),416 (type == "string" ?417 p.tool_arg_string_value(p.schema(418 p.until_one_of({ "</parameter>\n<parameter=", "</parameter>\n</function>" }),419 "tool-" + name + "-arg-" + param_name + "-schema", param_schema, true)) :420 p.tool_arg_json_value(p.schema(421 p.json(), "tool-" + name + "-arg-" + param_name + "-schema", param_schema))),422 p.tool_arg_close("</parameter>\n" +423 p.peek(p.literal("<parameter=") | p.literal("</function>"))) }));424 425 arg_parsers.push_back(is_required ? p.rule("tool-" + name + "-arg-" + param_name, arg) :426 p.optional(p.rule("tool-" + name + "-arg-" + param_name, arg)));427 }428 429 tool_parsers.push_back(p.rule("tool-" + name, p.tool_open("<function=" + p.tool_name(p.literal(name)) + ">")430 << p.sequence(arg_parsers)431 << p.tool_close(p.literal("</function>"))));432 };433 434 auto tool_call = p.trigger_rule("tool-call", "<tool_call>" << p.choice(tool_parsers) << "</tool_call>");435 436 return content + p.zero_or_more(p.space() + tool_call) + p.end();437 });438 439 auto grammar = build_grammar([&](const common_grammar_builder & builder) {440 for (const auto & def : tools) {441 auto function = def.at("function");442 auto parameters = function.at("parameters");443 builder.resolve_refs(parameters);444 };445 parser.build_grammar(builder);446 });447 448 t.log("Grammar:");449 for (const auto & line : string_split(grammar, "\n")) {450 t.log(line);451 }452 453 t.test("incremental parsing", [&](testing & t) {454 std::string input =455 "Let me search the knowledge base for cat pictures."456 "<tool_call>\n"457 "<function=search_knowledge_base>\n"458 "<parameter=query>cat pictures</parameter>\n"459 "<parameter=category>general</parameter>\n"460 "</function>\n"461 "</tool_call>";462 463 std::vector<std::string> tokens = simple_tokenize(input);464 465 common_chat_msg prev;466 for (auto it = tokens.begin(); it != tokens.end(); it++) {467 std::string in = std::accumulate(tokens.begin(), it + 1, std::string());468 469 common_peg_parse_context ctx(in, (it + 1 < tokens.end()) ? COMMON_PEG_PARSE_FLAG_LENIENT : COMMON_PEG_PARSE_FLAG_NONE);470 471 auto result = parser.parse(ctx);472 if (!t.assert_equal("not fail", false, result.fail())) {473 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));474 }475 476 common_chat_msg msg;477 auto mapper = common_chat_peg_mapper(msg);478 mapper.from_ast(ctx.ast, result);479 480 //t.log("Input: " + input);481 t.log("===========================================");482 t.log("Iteration " + std::to_string(in.size()));483 t.log("Reasoning: " + msg.reasoning_content);484 t.log("Content : " + msg.content);485 for (const auto & tc : msg.tool_calls) {486 t.log("Tool name: " + tc.name);487 t.log("Tool args: " + tc.arguments);488 }489 490 try {491 // This shouldn't emit any runtime errors492 auto diffs = common_chat_msg_diff::compute_diffs(prev, msg);493 } catch (const std::exception & e) {494 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));495 t.assert_true(std::string("failed with ") + e.what(), false);496 }497 498 prev = msg;499 }500 });501}502 503static void test_example_qwen3_non_coder(testing & t) {504 auto tools = create_tools();505 auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {506 // tool calling parser using standard JSON format507 auto tool_call = p.standard_json_tools("<tool_call>", "</tool_call>", tools, true, false);508 509 return p.sequence({ p.content(p.until("<tool_call>")), p.optional(p.space() + tool_call), p.end() });510 });511 512 auto grammar = build_grammar([&](const common_grammar_builder & builder) {513 for (const auto & def : tools) {514 auto function = def.at("function");515 auto parameters = function.at("parameters");516 builder.resolve_refs(parameters);517 };518 parser.build_grammar(builder);519 });520 521 t.log("Grammar:");522 for (const auto & line : string_split(grammar, "\n")) {523 t.log(line);524 }525 526 t.test("tool call parsing", [&](testing & t) {527 std::string input =528 "I need to get the weather.\n"529 "<tool_call>"530 "{\"name\": \"get_current_weather\", \"arguments\": {\"location\": \"New York City, NY\", \"unit\": "531 "\"fahrenheit\"}}"532 "</tool_call>";533 534 common_peg_parse_context ctx(input);535 auto result = parser.parse(ctx);536 537 t.assert_true("success", result.success());538 539 common_chat_msg msg;540 auto mapper = common_chat_peg_mapper(msg);541 mapper.from_ast(ctx.ast, result);542 543 t.assert_equal("content", "I need to get the weather.\n", msg.content);544 t.assert_equal("reasoning", "", msg.reasoning_content);545 t.assert_equal("tool calls count", 1u, msg.tool_calls.size());546 if (!msg.tool_calls.empty()) {547 t.assert_equal("tool name", "get_current_weather", msg.tool_calls[0].name);548 t.assert_equal("tool args", "{\"location\": \"New York City, NY\", \"unit\": \"fahrenheit\"}",549 msg.tool_calls[0].arguments);550 }551 });552 553 t.test("incremental parsing", [&](testing & t) {554 std::string input =555 "I need to get the weather.\n"556 "<tool_call>"557 "{\"name\": \"get_current_weather\", \"arguments\": {\"location\": \"New York City, NY\", \"unit\": "558 "\"fahrenheit\"}}"559 "</tool_call>";560 561 std::vector<std::string> tokens = simple_tokenize(input);562 563 common_chat_msg prev;564 for (auto it = tokens.begin(); it != tokens.end(); it++) {565 std::string in = std::accumulate(tokens.begin(), it + 1, std::string());566 567 common_peg_parse_context ctx(in, (it + 1 < tokens.end()) ? COMMON_PEG_PARSE_FLAG_LENIENT : COMMON_PEG_PARSE_FLAG_NONE);568 569 auto result = parser.parse(ctx);570 if (!t.assert_equal("not fail", false, result.fail())) {571 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));572 }573 574 common_chat_msg msg;575 auto mapper = common_chat_peg_mapper(msg);576 mapper.from_ast(ctx.ast, result);577 578 //t.log("Input: " + input);579 t.log("===========================================");580 t.log("Iteration " + std::to_string(in.size()));581 t.log("Reasoning: " + msg.reasoning_content);582 t.log("Content : " + msg.content);583 for (const auto & tc : msg.tool_calls) {584 t.log("Tool name: " + tc.name);585 t.log("Tool args: " + tc.arguments);586 }587 588 try {589 // This shouldn't emit any runtime errors590 auto diffs = common_chat_msg_diff::compute_diffs(prev, msg);591 } catch (const std::exception & e) {592 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));593 t.assert_true(std::string("failed with ") + e.what(), false);594 }595 596 prev = msg;597 }598 });599}600 601void test_command7_parser_compare(testing & t) {602 auto parser = build_chat_peg_parser([](common_chat_peg_builder & p) {603 auto thinking =604 p.reasoning_block("<|START_THINKING|>" << p.reasoning(p.until("<|END_THINKING|>")) << "<|END_THINKING|>");605 606 auto response = "<|START_RESPONSE|>" << p.content(p.until("<|END_RESPONSE|>")) << "<|END_RESPONSE|>";607 608 auto tool_call_id = p.atomic("\"tool_call_id\"" << (":" << ("\"" + p.tool_id(p.string_content('"')) + "\"")));609 auto tool_call_name =610 p.atomic("\"tool_name\"" << (":" << ("\"" + p.tool_name(p.string_content('"')) + "\"")));611 auto tool_call_args = "\"parameters\"" << (":" << p.tool_args(p.json()));612 613 auto tool_call_fields = p.rule("tool-call-fields", tool_call_id | tool_call_name | tool_call_args);614 auto tool_call =615 p.rule("tool-call", p.tool(p.tool_open(p.literal("{"))616 << tool_call_fields << p.zero_or_more(p.literal(",") << tool_call_fields)617 << p.tool_close(p.literal("}"))));618 619 auto tool_calls = p.rule(620 "tool-calls", "<|START_ACTION|>" << ("[" << tool_call << p.zero_or_more(p.literal(",") << tool_call) << "]")621 << "<|END_ACTION|>");622 623 return p.optional(thinking) << (tool_calls | response) + p.end();624 });625 626 auto test_current = [&](const common_peg_arena & p, const std::string & input, bool is_partial,627 bool print_results) {628 common_peg_parse_context ctx(input, is_partial ? COMMON_PEG_PARSE_FLAG_LENIENT : COMMON_PEG_PARSE_FLAG_NONE);629 auto result = p.parse(ctx);630 631 common_chat_msg msg;632 auto mapper = common_chat_peg_mapper(msg);633 mapper.from_ast(ctx.ast, result);634 635 if (print_results) {636 std::cout << "== Parsed (new) ==\n";637 std::cout << "=== Reasoning ===\n";638 std::cout << msg.reasoning_content << "\n";639 std::cout << "\n\n=== Content ===\n";640 std::cout << msg.content << "\n";641 std::cout << "\n\n=== Tool Calls ===\n";642 for (const auto & tc : msg.tool_calls) {643 std::cout << "id: " << tc.id << "\n";644 std::cout << "name: " << tc.name << "\n";645 std::cout << "args: " << tc.arguments << "\n";646 }647 }648 };649 650 std::string reasoning =651 "To plan an effective trip to Japan that includes both historical sites and modern attractions within a "652 "budget of $4000 for a two-week stay, we need to:\n\n"653 "1. Identify key historical sites and modern attractions in Japan.\n"654 "2. Find affordable accommodation options that provide a balance between comfort and cost.\n"655 "3. Determine the best modes of transportation for getting around Japan.\n"656 "4. Create a day-by-day itinerary that ensures the user gets to see a variety of attractions without "657 "overspending.\n"658 "5. Provide a detailed cost breakdown that includes accommodation, transportation, meals, and entry fees "659 "to attractions.";660 661 std::vector<std::tuple<std::string, std::string, nlohmann::json>> tool_calls = {662 { "call_0", "plan_trip", nlohmann::json::parse(R"({663 "destination": "Japan",664 "duration": 14,665 "budget": 4000,666 "interests": ["historical sites", "modern attractions"],667 "accommodation_preferences": "affordable",668 "transportation_preferences": "efficient",669 "meal_preferences": "local cuisine"670 })") }671 };672 673 std::vector<std::string> tokens;674 675 // Build tokens676 if (!reasoning.empty()) {677 auto tokenized = simple_tokenize(reasoning);678 tokens.emplace_back("<|START_THINKING|>");679 tokens.insert(tokens.end(), tokenized.begin(), tokenized.end());680 tokens.emplace_back("<|END_THINKING|>");681 }682 683 if (!tool_calls.empty()) {684 tokens.emplace_back("<|START_ACTION|>");685 686 auto json = nlohmann::json::array();687 for (const auto & tc : tool_calls) {688 auto tc_json = nlohmann::json::object();689 tc_json["tool_call_id"] = std::get<0>(tc);690 tc_json["tool_name"] = std::get<1>(tc);691 tc_json["parameters"] = std::get<2>(tc);692 json.push_back(tc_json);693 }694 695 auto tokenized = simple_tokenize(json.dump(-1, ' ', true));696 tokens.insert(tokens.end(), tokenized.begin(), tokenized.end());697 698 tokens.emplace_back("<|END_ACTION|>");699 }700 701 std::string input = std::accumulate(tokens.begin(), tokens.end(), std::string());702 703 t.test("current_parse", [&](testing & /* t */) { test_current(parser, input, false, false); });704 t.bench("current_parse_benchmark complete", [&]() { test_current(parser, input, false, false); }, 100);705 t.bench(706 "current_parse_benchmark incremental",707 [&]() {708 std::string in;709 for (auto i = 0u; i < tokens.size(); i++) {710 in += tokens[i];711 test_current(parser, in, i + 1 < tokens.size(), false);712 }713 },714 20);715}716 717// Test that tool names that are proper prefixes of other tool names don't cause718// premature matching during incremental parsing.719// For example, "special_function" should not match when parsing "special_function_with_opt".720static void test_prefix_tool_names(testing & t) {721 // Create tools where one name is a proper prefix of another722 json tools = json::array();723 724 json tool_short = {725 { "type", "function" },726 { "function",727 {728 { "name", "special_function" },729 { "description", "A special function" },730 { "parameters",731 {732 { "type", "object" },733 { "properties",734 {735 { "arg1", { { "type", "integer" } } },736 } },737 { "required", { "arg1" } },738 } },739 } }740 };741 tools.push_back(tool_short);742 743 json tool_long = {744 { "type", "function" },745 { "function",746 {747 { "name", "special_function_with_opt" },748 { "description", "A special function with optional params" },749 { "parameters",750 {751 { "type", "object" },752 { "properties",753 {754 { "arg1", { { "type", "integer" } } },755 { "arg2", { { "type", "integer" } } },756 } },757 { "required", { "arg1" } },758 } },759 } }760 };761 tools.push_back(tool_long);762 763 // Use standard_constructed_tools which had the prefix matching bug764 std::map<std::string, std::string> markers = {765 { "tool_call_start_marker", "<tool_call>" },766 { "tool_call_end_marker", "</tool_call>" },767 { "function_opener", "<function=" },768 { "function_closer", "</function>" },769 { "function_name_suffix", ">" },770 { "parameter_key_prefix", "<param=" },771 { "parameter_key_suffix", ">" },772 { "parameter_closer", "</param>" },773 };774 775 auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {776 auto content = p.rule("content", p.content(p.until("<tool_call>")));777 auto tool_call = p.standard_constructed_tools(markers, tools, false, false);778 return content + p.zero_or_more(p.space() + tool_call) + p.end();779 });780 781 // Test parsing the long tool name - this should NOT trigger the short tool name782 t.test("parse long tool name", [&](testing & t) {783 std::string input =784 "Let me call the function."785 "<tool_call>"786 "<function=special_function_with_opt>"787 "<param=arg1>42</param>"788 "</function>"789 "</tool_call>";790 791 common_peg_parse_context ctx(input);792 auto result = parser.parse(ctx);793 794 t.assert_true("success", result.success());795 796 common_chat_msg msg;797 auto mapper = common_chat_peg_mapper(msg);798 mapper.from_ast(ctx.ast, result);799 800 t.assert_equal("content", "Let me call the function.", msg.content);801 t.assert_equal("tool calls count", 1u, msg.tool_calls.size());802 if (!msg.tool_calls.empty()) {803 t.assert_equal("tool name", "special_function_with_opt", msg.tool_calls[0].name);804 }805 });806 807 // Test incremental parsing - the key test case808 // This ensures that when incrementally parsing "special_function_with_opt",809 // we don't prematurely emit "special_function" as a tool call810 t.test("incremental parse long tool name", [&](testing & t) {811 std::string input =812 "Let me call the function."813 "<tool_call>"814 "<function=special_function_with_opt>"815 "<param=arg1>42</param>"816 "</function>"817 "</tool_call>";818 819 std::vector<std::string> tokens = simple_tokenize(input);820 821 common_chat_msg prev;822 for (auto it = tokens.begin(); it != tokens.end(); it++) {823 std::string in = std::accumulate(tokens.begin(), it + 1, std::string());824 825 common_peg_parse_context ctx(in, (it + 1 < tokens.end()) ? COMMON_PEG_PARSE_FLAG_LENIENT : COMMON_PEG_PARSE_FLAG_NONE);826 auto result = parser.parse(ctx);827 828 if (!t.assert_equal("not fail", false, result.fail())) {829 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));830 return;831 }832 833 common_chat_msg msg;834 auto mapper = common_chat_peg_mapper(msg);835 mapper.from_ast(ctx.ast, result);836 837 // The critical check: during incremental parsing, we should never838 // see "special_function" as the tool name when parsing "special_function_with_opt"839 for (const auto & tc : msg.tool_calls) {840 if (!t.assert_equal("tool name should not be short prefix", false,841 tc.name == "special_function")) {842 t.log("Premature tool name match at input: " + in);843 return;844 }845 }846 847 try {848 auto diffs = common_chat_msg_diff::compute_diffs(prev, msg);849 } catch (const std::exception & e) {850 t.log(in.substr(0, result.end) + "[failed->]" + in.substr(result.end));851 t.assert_true(std::string("diff failed with ") + e.what(), false);852 return;853 }854 855 prev = msg;856 }857 858 // Final check: the complete parse should have the correct tool name859 t.assert_equal("final tool calls count", 1u, prev.tool_calls.size());860 if (!prev.tool_calls.empty()) {861 t.assert_equal("final tool name", "special_function_with_opt", prev.tool_calls[0].name);862 }863 });864 865 // Test parsing the short tool name still works866 t.test("parse short tool name", [&](testing & t) {867 std::string input =868 "Let me call the function."869 "<tool_call>"870 "<function=special_function>"871 "<param=arg1>42</param>"872 "</function>"873 "</tool_call>";874 875 common_peg_parse_context ctx(input);876 auto result = parser.parse(ctx);877 878 t.assert_true("success", result.success());879 880 common_chat_msg msg;881 auto mapper = common_chat_peg_mapper(msg);882 mapper.from_ast(ctx.ast, result);883 884 t.assert_equal("content", "Let me call the function.", msg.content);885 t.assert_equal("tool calls count", 1u, msg.tool_calls.size());886 if (!msg.tool_calls.empty()) {887 t.assert_equal("tool name", "special_function", msg.tool_calls[0].name);888 }889 });890}891 892static void test_tagged_peg_parser(testing & t) {893 t.test("basic tag extraction", [&](testing & t) {894 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {895 return p.tag("greeting", p.until(" ")) + " " + p.tag("name", p.rest()) + p.end();896 });897 898 auto result = parser.parse_and_extract("Hello World");899 t.assert_true("success", result.result.success());900 t.assert_equal("greeting tag", "Hello", result.tags.at("greeting"));901 t.assert_equal("name tag", "World", result.tags.at("name"));902 });903 904 t.test("duplicate tags overwrite", [&](testing & t) {905 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {906 return p.tag("item", p.until(",")) + "," + p.tag("item", p.rest()) + p.end();907 });908 909 auto result = parser.parse_and_extract("first,second");910 t.assert_true("success", result.result.success());911 t.assert_equal("item tag", "second", result.tags.at("item"));912 });913 914 t.test("no tags extracted", [&](testing & t) {915 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {916 return p.rest() + p.end();917 });918 919 auto result = parser.parse_and_extract("Hello");920 t.assert_true("success", result.result.success());921 t.assert_equal("empty tags", 0u, result.tags.size());922 });923 924 t.test("structured extraction", [&](testing & t) {925 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {926 auto header = p.tag("header", p.until("\n"));927 auto body = p.tag("body", p.rest());928 return header + "\n" + body + p.end();929 });930 931 auto result = parser.parse_and_extract("Title\nBody content here");932 t.assert_true("success", result.result.success());933 t.assert_equal("header", "Title", result.tags.at("header"));934 t.assert_equal("body", "Body content here", result.tags.at("body"));935 });936 937 t.test("partial parse", [&](testing & t) {938 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {939 return p.tag("prefix", p.until(":")) + ":" + p.tag("value", p.rest()) + p.end();940 });941 942 auto result = parser.parse_and_extract("key:val", COMMON_PEG_PARSE_FLAG_LENIENT);943 t.assert_true("not fail", !result.result.fail());944 t.assert_equal("prefix tag", "key", result.tags.at("prefix"));945 t.assert_equal("value tag", "val", result.tags.at("value"));946 });947 948 t.test("find in the middle", [&](testing & t) {949 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {950 return p.choice({ p.literal("{"), p.literal(":") }) + p.space() + p.literal("\"") + p.atomic(p.literal("fun_name"));951 });952 953 std::string tpl = "This is a very long jinja template string. We have tools. We will try to call them now: <tool_call>{ \"fun_name\" : { \"arg\" : 1 }</tool_call>";954 auto result = parser.parse_anywhere_and_extract(tpl);955 t.assert_true("success", result.result.success());956 });957 958 t.test("fail find in the middle", [&](testing & t) {959 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {960 return p.choice({ p.literal("{"), p.literal(":") }) + p.space() + p.literal("\"") + p.atomic(p.literal("fun_name"));961 });962 963 std::string tpl = "This is a very long jinja template string. We have tools. We will try to call them now: <tool_call><fun=fun_name><arg name=arg>1</arg></tool_call>";964 auto result = parser.parse_anywhere_and_extract(tpl);965 t.assert_true("failure", result.result.fail());966 });967 968 t.test("find function tag with name", [&](testing &t) {969 std::string haystack = "\n<tool_call>\n<function=foofoo>\n<parameter=first>\nXXXX\n</parameter>\n<parameter=second>\nYYYY\n</parameter>\n</function>\n</tool_call>\n";970 auto parser = build_tagged_peg_parser([](common_peg_parser_builder & p) {971 std::string needle = "foofoo";972 return p.tag("fun_marker", p.choice({973 p.tag("fun_pre", p.literal("<") + p.until_one_of({ ">", needle })) + p.literal(needle) +974 p.tag("fun_post", p.negate(p.space() + p.literal("<")) + p.until(">") + p.literal(">")) + p.space(),975 p.tag("fun_pre", p.literal("[") + p.until_one_of({ "]", needle })) + p.literal(needle) +976 p.tag("fun_post", p.negate(p.space() + p.literal("[") + p.until("]") + p.literal("]")) + p.space()) }));977 });978 auto result = parser.parse_anywhere_and_extract(haystack);979 t.assert_true("success", result.result.success());980 t.assert_equal("fun_pre should be '<function='", "<function=", result.tags["fun_pre"]);981 t.assert_equal("fun_post should be '>'", ">", result.tags["fun_post"]);982 });983}984 