mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-24 13:37:01 +02:00
common/peg : handle invalid utf-8 sequences in the AST (#29161)
* common/peg : handle invalid utf-8 sequences in the AST * cont : return maximal subpart per Unicode recommendations * cont : remove strict argument
This commit is contained in:
@@ -318,13 +318,13 @@ void common_chat_peg_mapper::map(const common_peg_ast_node & node) {
|
||||
bool is_content = node.tag == common_chat_peg_builder::CONTENT;
|
||||
|
||||
if (is_reasoning) { // GPT OSS can have more than 1 reasoning block, so concatenate here
|
||||
result.reasoning_content += std::string(node.text);
|
||||
result.reasoning_content += node.sanitized_text();
|
||||
}
|
||||
|
||||
if (is_content) {
|
||||
// Concatenate content from multiple content nodes (e.g., when reasoning markers
|
||||
// are preserved before content markers in reasoning_format=NONE mode)
|
||||
result.content += std::string(node.text);
|
||||
result.content += node.sanitized_text();
|
||||
}
|
||||
|
||||
// Handle tool-related tags (supporting both JSON and tagged formats)
|
||||
@@ -1058,12 +1058,12 @@ void common_chat_peg_gemma4_mapper::visit(const common_peg_ast_arena & arena, co
|
||||
const auto & node = arena.get(id);
|
||||
|
||||
if (node.tag == "reasoning") {
|
||||
result.reasoning_content += std::string(node.text);
|
||||
result.reasoning_content += node.sanitized_text();
|
||||
return;
|
||||
}
|
||||
|
||||
if (node.tag == "content") {
|
||||
result.content += std::string(node.text);
|
||||
result.content += node.sanitized_text();
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -1206,12 +1206,12 @@ void common_chat_peg_minimax_m3_mapper::visit(const common_peg_ast_arena & arena
|
||||
const auto & node = arena.get(id);
|
||||
|
||||
if (node.tag == common_chat_peg_builder::REASONING) {
|
||||
result.reasoning_content += std::string(node.text);
|
||||
result.reasoning_content += node.sanitized_text();
|
||||
return;
|
||||
}
|
||||
|
||||
if (node.tag == common_chat_peg_builder::CONTENT) {
|
||||
result.content += std::string(node.text);
|
||||
result.content += node.sanitized_text();
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
+50
-24
@@ -166,6 +166,25 @@ common_peg_ast_id common_peg_ast_arena::find_by_rule(const common_peg_ast_node &
|
||||
return COMMON_PEG_INVALID_AST_ID;
|
||||
}
|
||||
|
||||
std::string common_peg_ast_node::sanitized_text() const {
|
||||
if (invalid_utf8.empty()) {
|
||||
return std::string(text);
|
||||
}
|
||||
|
||||
std::string out;
|
||||
out.reserve(text.size() + 2 * invalid_utf8.size());
|
||||
|
||||
size_t seg_start = start;
|
||||
for (const auto & invalid : invalid_utf8) {
|
||||
out.append(text.data() + (seg_start - start), invalid.pos - seg_start);
|
||||
out.append("\xEF\xBF\xBD");
|
||||
seg_start = invalid.pos + invalid.len;
|
||||
}
|
||||
out.append(text.data() + (seg_start - start), end - seg_start);
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
void common_peg_ast_arena::visit(common_peg_ast_id id, const common_peg_ast_visitor & visitor) const {
|
||||
if (id == COMMON_PEG_INVALID_AST_ID) {
|
||||
return;
|
||||
@@ -282,6 +301,7 @@ struct parser_executor {
|
||||
|
||||
auto pos = start_pos;
|
||||
std::vector<common_peg_ast_id> nodes;
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8;
|
||||
|
||||
for (size_t i = 0; i < p.children.size(); i++) {
|
||||
const auto & child_id = p.children[i];
|
||||
@@ -306,13 +326,14 @@ struct parser_executor {
|
||||
if (!result.nodes.empty()) {
|
||||
nodes.insert(nodes.end(), result.nodes.begin(), result.nodes.end());
|
||||
}
|
||||
invalid_utf8.insert(invalid_utf8.end(), result.invalid_utf8.begin(), result.invalid_utf8.end());
|
||||
|
||||
if (result.need_more_input()) {
|
||||
ctx.parse_depth--;
|
||||
if (ctx.is_debug()) {
|
||||
fprintf(stderr, "%sSEQ -> NEED_MORE\n", debug_indent().c_str());
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, result.end, std::move(nodes));
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, result.end, std::move(nodes), std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
pos = result.end;
|
||||
@@ -322,7 +343,7 @@ struct parser_executor {
|
||||
if (ctx.is_debug()) {
|
||||
fprintf(stderr, "%sSEQ -> SUCCESS at %zu->%zu\n", debug_indent().c_str(), start_pos, pos);
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, std::move(nodes));
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, std::move(nodes), std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
common_peg_parse_result operator()(const common_peg_choice_parser & p) {
|
||||
@@ -370,6 +391,7 @@ struct parser_executor {
|
||||
auto pos = start_pos;
|
||||
int match_count = 0;
|
||||
std::vector<common_peg_ast_id> nodes;
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8;
|
||||
|
||||
// Try to match up to max_count times (or unlimited if max_count is -1)
|
||||
while (p.max_count == -1 || match_count < p.max_count) {
|
||||
@@ -400,6 +422,7 @@ struct parser_executor {
|
||||
if (!result.nodes.empty()) {
|
||||
nodes.insert(nodes.end(), result.nodes.begin(), result.nodes.end());
|
||||
}
|
||||
invalid_utf8.insert(invalid_utf8.end(), result.invalid_utf8.begin(), result.invalid_utf8.end());
|
||||
|
||||
pos = result.end;
|
||||
match_count++;
|
||||
@@ -410,13 +433,14 @@ struct parser_executor {
|
||||
if (!result.nodes.empty()) {
|
||||
nodes.insert(nodes.end(), result.nodes.begin(), result.nodes.end());
|
||||
}
|
||||
invalid_utf8.insert(invalid_utf8.end(), result.invalid_utf8.begin(), result.invalid_utf8.end());
|
||||
|
||||
ctx.parse_depth--;
|
||||
if (ctx.is_debug()) {
|
||||
fprintf(stderr, "%sREPEAT -> NEED_MORE (count=%d, nodes=%zu)\n", debug_indent().c_str(),
|
||||
match_count, nodes.size());
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, result.end, std::move(nodes));
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, result.end, std::move(nodes), std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
// Child failed - stop trying
|
||||
@@ -434,7 +458,7 @@ struct parser_executor {
|
||||
fprintf(stderr, "%sREPEAT -> NEED_MORE (not enough matches: %d < %d)\n", debug_indent().c_str(),
|
||||
match_count, p.min_count);
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, pos, std::move(nodes));
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, pos, std::move(nodes), std::move(invalid_utf8));
|
||||
}
|
||||
if (ctx.is_debug()) {
|
||||
fprintf(stderr, "%sREPEAT -> FAIL (not enough matches: %d < %d)\n", debug_indent().c_str(), match_count,
|
||||
@@ -448,7 +472,7 @@ struct parser_executor {
|
||||
fprintf(stderr, "%sREPEAT -> SUCCESS (count=%d, nodes=%zu)\n", debug_indent().c_str(), match_count,
|
||||
nodes.size());
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, std::move(nodes));
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, std::move(nodes), std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
common_peg_parse_result operator()(const common_peg_and_parser & p) {
|
||||
@@ -664,23 +688,23 @@ struct parser_executor {
|
||||
// Scan input and check for delimiters
|
||||
size_t pos = start_pos;
|
||||
size_t last_valid_pos = start_pos;
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8;
|
||||
|
||||
while (pos < ctx.input.size()) {
|
||||
auto utf8_result = common_parse_utf8_codepoint(ctx.input, pos);
|
||||
|
||||
if (utf8_result.status == utf8_parse_result::INCOMPLETE) {
|
||||
// Incomplete UTF-8 sequence
|
||||
if (!ctx.is_lenient()) {
|
||||
// Input is complete but UTF-8 is incomplete = malformed
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_FAIL, start_pos);
|
||||
}
|
||||
// Return what we have so far (before incomplete sequence)
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, last_valid_pos);
|
||||
if (utf8_result.status == utf8_parse_result::INCOMPLETE && ctx.is_lenient()) {
|
||||
// The rest of the sequence may still arrive, return what we have so far
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, last_valid_pos, {}, std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
if (utf8_result.status == utf8_parse_result::INVALID) {
|
||||
// Malformed UTF-8
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_FAIL, start_pos);
|
||||
if (utf8_result.status != utf8_parse_result::SUCCESS) {
|
||||
// Malformed UTF-8, or a sequence truncated by the end of a complete input.
|
||||
// A delimiter cannot start inside bytes that fail to decode, so consume them and move on
|
||||
invalid_utf8.push_back({pos, utf8_result.bytes_consumed});
|
||||
pos += utf8_result.bytes_consumed;
|
||||
last_valid_pos = pos;
|
||||
continue;
|
||||
}
|
||||
|
||||
// Check if a delimiter starts at this position
|
||||
@@ -688,12 +712,12 @@ struct parser_executor {
|
||||
|
||||
if (match == common_trie::COMPLETE_MATCH) {
|
||||
// Found a complete delimiter, return everything before it
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos);
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, {}, std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
if (match == common_trie::PARTIAL_MATCH) {
|
||||
// Found a partial match extending to end of input, return everything before it
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos);
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, pos, {}, std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
pos += utf8_result.bytes_consumed;
|
||||
@@ -702,9 +726,9 @@ struct parser_executor {
|
||||
|
||||
if (last_valid_pos == ctx.input.size() && ctx.is_lenient()) {
|
||||
// Reached the end of a partial stream, there might still be more input that we need to consume.
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, last_valid_pos);
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT, start_pos, last_valid_pos, {}, std::move(invalid_utf8));
|
||||
}
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, last_valid_pos);
|
||||
return common_peg_parse_result(COMMON_PEG_PARSE_RESULT_SUCCESS, start_pos, last_valid_pos, {}, std::move(invalid_utf8));
|
||||
}
|
||||
|
||||
common_peg_parse_result operator()(const common_peg_schema_parser & p) {
|
||||
@@ -728,10 +752,11 @@ struct parser_executor {
|
||||
result.end,
|
||||
text,
|
||||
std::move(result.nodes),
|
||||
result.need_more_input()
|
||||
result.need_more_input(),
|
||||
result.invalid_utf8
|
||||
);
|
||||
|
||||
return common_peg_parse_result(result.type, result.start, result.end, { node_id });
|
||||
return common_peg_parse_result(result.type, result.start, result.end, { node_id }, std::move(result.invalid_utf8));
|
||||
}
|
||||
|
||||
return result;
|
||||
@@ -757,10 +782,11 @@ struct parser_executor {
|
||||
result.end,
|
||||
text,
|
||||
std::move(result.nodes),
|
||||
result.need_more_input()
|
||||
result.need_more_input(),
|
||||
result.invalid_utf8
|
||||
);
|
||||
|
||||
return common_peg_parse_result(result.type, result.start, result.end, { node_id });
|
||||
return common_peg_parse_result(result.type, result.start, result.end, { node_id }, std::move(result.invalid_utf8));
|
||||
}
|
||||
|
||||
return result;
|
||||
|
||||
+21
-4
@@ -72,6 +72,12 @@ enum common_peg_parse_result_type {
|
||||
|
||||
const char * common_peg_parse_result_type_name(common_peg_parse_result_type type);
|
||||
|
||||
// A run of input bytes that does not decode as UTF-8
|
||||
struct common_peg_invalid_utf8 {
|
||||
size_t pos;
|
||||
size_t len;
|
||||
};
|
||||
|
||||
struct common_peg_ast_node {
|
||||
common_peg_ast_id id;
|
||||
std::string rule;
|
||||
@@ -82,6 +88,12 @@ struct common_peg_ast_node {
|
||||
std::vector<common_peg_ast_id> children;
|
||||
|
||||
bool is_partial = false;
|
||||
|
||||
// Invalid UTF-8 inside the node, in ascending order
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8;
|
||||
|
||||
// Returns the text with every invalid run replaced by U+FFFD
|
||||
std::string sanitized_text() const;
|
||||
};
|
||||
|
||||
struct common_peg_parse_result;
|
||||
@@ -98,10 +110,11 @@ class common_peg_ast_arena {
|
||||
size_t end,
|
||||
std::string_view text,
|
||||
std::vector<common_peg_ast_id> children,
|
||||
bool is_partial = false
|
||||
bool is_partial = false,
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8 = {}
|
||||
) {
|
||||
common_peg_ast_id id = nodes_.size();
|
||||
nodes_.push_back({id, rule, tag, start, end, text, std::move(children), is_partial});
|
||||
nodes_.push_back({id, rule, tag, start, end, text, std::move(children), is_partial, std::move(invalid_utf8)});
|
||||
return id;
|
||||
}
|
||||
|
||||
@@ -127,6 +140,9 @@ struct common_peg_parse_result {
|
||||
|
||||
std::vector<common_peg_ast_id> nodes;
|
||||
|
||||
// Invalid UTF-8 consumed by this result, carried up to the enclosing AST nodes
|
||||
std::vector<common_peg_invalid_utf8> invalid_utf8;
|
||||
|
||||
common_peg_parse_result() = default;
|
||||
|
||||
common_peg_parse_result(common_peg_parse_result_type type, size_t start)
|
||||
@@ -135,8 +151,8 @@ struct common_peg_parse_result {
|
||||
common_peg_parse_result(common_peg_parse_result_type type, size_t start, size_t end)
|
||||
: type(type), start(start), end(end) {}
|
||||
|
||||
common_peg_parse_result(common_peg_parse_result_type type, size_t start, size_t end, std::vector<common_peg_ast_id> nodes)
|
||||
: type(type), start(start), end(end), nodes(std::move(nodes)) {}
|
||||
common_peg_parse_result(common_peg_parse_result_type type, size_t start, size_t end, std::vector<common_peg_ast_id> nodes, std::vector<common_peg_invalid_utf8> invalid_utf8 = {})
|
||||
: type(type), start(start), end(end), nodes(std::move(nodes)), invalid_utf8(std::move(invalid_utf8)) {}
|
||||
|
||||
bool fail() const { return type == COMMON_PEG_PARSE_RESULT_FAIL; }
|
||||
bool need_more_input() const { return type == COMMON_PEG_PARSE_RESULT_NEED_MORE_INPUT; }
|
||||
@@ -430,6 +446,7 @@ class common_peg_parser_builder {
|
||||
common_peg_parser space() { return add(common_peg_space_parser{}); }
|
||||
|
||||
// Matches all characters until a delimiter is found (delimiter not consumed).
|
||||
// Invalid UTF-8 is consumed and recorded on the AST nodes.
|
||||
// S -> (!delim .)*
|
||||
common_peg_parser until(const std::string & delimiter) { return add(common_peg_until_parser{{delimiter}}); }
|
||||
|
||||
|
||||
+19
-14
@@ -26,16 +26,16 @@ utf8_parse_result common_parse_utf8_codepoint(std::string_view input, size_t off
|
||||
|
||||
// Invalid: continuation byte as first byte
|
||||
if (!(input[offset] & 0x40)) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID);
|
||||
return utf8_parse_result(utf8_parse_result::INVALID, 0, 1);
|
||||
}
|
||||
|
||||
// 2-byte sequence
|
||||
if (!(input[offset] & 0x20)) {
|
||||
if (offset + 1 >= input.size()) {
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE);
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE, 0, 1);
|
||||
}
|
||||
if ((input[offset + 1] & 0xc0) != 0x80) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID);
|
||||
return utf8_parse_result(utf8_parse_result::INVALID, 0, 1);
|
||||
}
|
||||
auto result = ((input[offset] & 0x1f) << 6) | (input[offset + 1] & 0x3f);
|
||||
return utf8_parse_result(utf8_parse_result::SUCCESS, result, 2);
|
||||
@@ -43,11 +43,14 @@ utf8_parse_result common_parse_utf8_codepoint(std::string_view input, size_t off
|
||||
|
||||
// 3-byte sequence
|
||||
if (!(input[offset] & 0x10)) {
|
||||
if (offset + 2 >= input.size()) {
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE);
|
||||
}
|
||||
if ((input[offset + 1] & 0xc0) != 0x80 || (input[offset + 2] & 0xc0) != 0x80) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID);
|
||||
// Check one byte at a time so a bad byte is reported before a short input
|
||||
for (size_t i = 1; i < 3; i++) {
|
||||
if (offset + i >= input.size()) {
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE, 0, i);
|
||||
}
|
||||
if ((input[offset + i] & 0xc0) != 0x80) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID, 0, i);
|
||||
}
|
||||
}
|
||||
auto result = ((input[offset] & 0x0f) << 12) | ((input[offset + 1] & 0x3f) << 6) | (input[offset + 2] & 0x3f);
|
||||
return utf8_parse_result(utf8_parse_result::SUCCESS, result, 3);
|
||||
@@ -55,18 +58,20 @@ utf8_parse_result common_parse_utf8_codepoint(std::string_view input, size_t off
|
||||
|
||||
// 4-byte sequence
|
||||
if (!(input[offset] & 0x08)) {
|
||||
if (offset + 3 >= input.size()) {
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE);
|
||||
}
|
||||
if ((input[offset + 1] & 0xc0) != 0x80 || (input[offset + 2] & 0xc0) != 0x80 || (input[offset + 3] & 0xc0) != 0x80) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID);
|
||||
for (size_t i = 1; i < 4; i++) {
|
||||
if (offset + i >= input.size()) {
|
||||
return utf8_parse_result(utf8_parse_result::INCOMPLETE, 0, i);
|
||||
}
|
||||
if ((input[offset + i] & 0xc0) != 0x80) {
|
||||
return utf8_parse_result(utf8_parse_result::INVALID, 0, i);
|
||||
}
|
||||
}
|
||||
auto result = ((input[offset] & 0x07) << 18) | ((input[offset + 1] & 0x3f) << 12) | ((input[offset + 2] & 0x3f) << 6) | (input[offset + 3] & 0x3f);
|
||||
return utf8_parse_result(utf8_parse_result::SUCCESS, result, 4);
|
||||
}
|
||||
|
||||
// Invalid first byte
|
||||
return utf8_parse_result(utf8_parse_result::INVALID);
|
||||
return utf8_parse_result(utf8_parse_result::INVALID, 0, 1);
|
||||
}
|
||||
|
||||
bool common_utf8_is_complete(const std::string & s) {
|
||||
|
||||
+1
-1
@@ -9,7 +9,7 @@
|
||||
|
||||
struct utf8_parse_result {
|
||||
uint32_t codepoint; // Decoded codepoint (only valid if status == SUCCESS)
|
||||
size_t bytes_consumed; // How many bytes this codepoint uses (1-4)
|
||||
size_t bytes_consumed; // How many bytes this codepoint uses (1-4), or the length of the valid prefix if status != SUCCESS
|
||||
enum status { SUCCESS, INCOMPLETE, INVALID } status;
|
||||
|
||||
utf8_parse_result(enum status s, uint32_t cp = 0, size_t bytes = 0)
|
||||
|
||||
@@ -273,19 +273,35 @@ void test_unicode(testing &t) {
|
||||
});
|
||||
|
||||
t.test("malformed UTF-8", [](testing &t) {
|
||||
std::vector<test_case> test_cases {
|
||||
struct passthrough_case {
|
||||
std::string input;
|
||||
std::string expected_text;
|
||||
std::string expected_sanitized;
|
||||
};
|
||||
|
||||
std::vector<passthrough_case> test_cases {
|
||||
// Invalid UTF-8 bytes
|
||||
{std::string("Hello\xFF\xFE"), "", COMMON_PEG_PARSE_RESULT_FAIL},
|
||||
{std::string("Hello\xFF\xFE</tag>"), std::string("Hello\xFF\xFE"), "Hello\xEF\xBF\xBD\xEF\xBF\xBD"},
|
||||
|
||||
// Continuation byte without lead byte
|
||||
{std::string("Hello\x80World"), "", COMMON_PEG_PARSE_RESULT_FAIL},
|
||||
{std::string("Hello\x80World</tag>"), std::string("Hello\x80World"), "Hello\xEF\xBF\xBDWorld"},
|
||||
|
||||
// Invalid continuation byte
|
||||
{std::string("\xC3\x28"), "", COMMON_PEG_PARSE_RESULT_FAIL},
|
||||
// Invalid continuation byte, the lead byte is dropped and '(' survives
|
||||
{std::string("\xC3\x28</tag>"), std::string("\xC3\x28"), "\xEF\xBF\xBD("},
|
||||
|
||||
// Two good bytes of a 3-byte sequence then a bad third byte, the prefix is replaced once and the third byte is kept
|
||||
{std::string("\xE4\xB8" "A</tag>"), std::string("\xE4\xB8" "A"), "\xEF\xBF\xBD" "A"},
|
||||
{std::string("\xE4\xB8</tag>"), std::string("\xE4\xB8"), "\xEF\xBF\xBD"},
|
||||
|
||||
// Truncated sequence in a complete input, the leftover prefix is replaced once
|
||||
{std::string("Hello\xE4\xB8"), std::string("Hello\xE4\xB8"), "Hello\xEF\xBF\xBD"},
|
||||
|
||||
// Valid multi-byte content around the bad byte is left alone
|
||||
{std::string("\xE4\xBD\xA0\xFF\xE5\xA5\xBD</tag>"), std::string("\xE4\xBD\xA0\xFF\xE5\xA5\xBD"), "\xE4\xBD\xA0\xEF\xBF\xBD\xE5\xA5\xBD"},
|
||||
};
|
||||
|
||||
auto parser = build_peg_parser([](common_peg_parser_builder& p) {
|
||||
return p.until("</tag>");
|
||||
return p.tag("body", p.until("</tag>")) + p.optional(p.literal("</tag>"));
|
||||
});
|
||||
|
||||
for (size_t i = 0; i < test_cases.size(); i++) {
|
||||
@@ -296,10 +312,28 @@ void test_unicode(testing &t) {
|
||||
common_peg_parse_context ctx(tc.input);
|
||||
auto result = parser.parse(ctx);
|
||||
|
||||
assert_result_equal(t, tc.expected_result, result.type);
|
||||
assert_result_equal(t, COMMON_PEG_PARSE_RESULT_SUCCESS, result.type);
|
||||
const auto & node = ctx.ast.get(result.nodes[0]);
|
||||
t.assert_equal("raw text", tc.expected_text, std::string(node.text));
|
||||
t.assert_equal("sanitized text", tc.expected_sanitized, node.sanitized_text());
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
t.test("malformed UTF-8 rescanned by backtracking", [](testing &t) {
|
||||
// The failed alternative and the lookahead scan the same bad byte, it must only be recorded once
|
||||
auto parser = build_peg_parser([](common_peg_parser_builder& p) {
|
||||
return (p.until("<a>") + p.literal("<a>")) | (p.peek(p.until("<b>")) + p.until("<b>") + p.literal("<b>"));
|
||||
});
|
||||
|
||||
std::string input("x\xFFy<b>");
|
||||
common_peg_parse_context ctx(input);
|
||||
auto result = parser.parse(ctx);
|
||||
|
||||
assert_result_equal(t, COMMON_PEG_PARSE_RESULT_SUCCESS, result.type);
|
||||
t.assert_equal("invalid count", 1u, result.invalid_utf8.size());
|
||||
t.assert_equal("invalid offset", 1u, result.invalid_utf8[0].pos);
|
||||
});
|
||||
});
|
||||
|
||||
t.test("json_string parser", [](testing &t) {
|
||||
|
||||
@@ -23,6 +23,7 @@ static void test_command7_parser_compare(testing & t);
|
||||
static void test_prefix_tool_names(testing & t);
|
||||
static void test_tagged_peg_parser(testing & t);
|
||||
static void test_permute(testing & t);
|
||||
static void test_invalid_utf8(testing & t);
|
||||
|
||||
int main(int argc, char * argv[]) {
|
||||
testing t(std::cout);
|
||||
@@ -42,6 +43,7 @@ int main(int argc, char * argv[]) {
|
||||
t.test("prefix tool names", test_prefix_tool_names);
|
||||
t.test("tagged peg parser", test_tagged_peg_parser);
|
||||
t.test("permute", test_permute);
|
||||
t.test("invalid utf8", test_invalid_utf8);
|
||||
|
||||
return t.summary();
|
||||
}
|
||||
@@ -1069,3 +1071,36 @@ static void test_permute(testing & t) {
|
||||
)""", gbnf_of(parser));
|
||||
});
|
||||
}
|
||||
|
||||
static void test_invalid_utf8(testing & t) {
|
||||
auto parser = build_chat_peg_parser([](common_chat_peg_builder & p) {
|
||||
return "<think>" + p.reasoning(p.until("</think>")) + "</think>" + p.content(p.rest()) + p.end();
|
||||
});
|
||||
|
||||
t.test("replaced in reasoning and content", [&](testing & t) {
|
||||
std::string input("<think>plan\xFF\xFE</think>caf\xC3\xA9 \x80 done");
|
||||
common_peg_parse_context ctx(input);
|
||||
auto result = parser.parse(ctx);
|
||||
t.assert_true("success", result.success());
|
||||
|
||||
common_chat_msg msg;
|
||||
auto mapper = common_chat_peg_mapper(msg);
|
||||
mapper.from_ast(ctx.ast, result);
|
||||
|
||||
t.assert_equal("reasoning", "plan\xEF\xBF\xBD\xEF\xBF\xBD", msg.reasoning_content);
|
||||
t.assert_equal("content", "caf\xC3\xA9 \xEF\xBF\xBD done", msg.content);
|
||||
});
|
||||
|
||||
t.test("partial input keeps trailing incomplete sequence out", [&](testing & t) {
|
||||
std::string input("<think>x</think>a\x80" "b\xE4\xB8");
|
||||
common_peg_parse_context ctx(input, COMMON_PEG_PARSE_FLAG_LENIENT);
|
||||
auto result = parser.parse(ctx);
|
||||
t.assert_true("not fail", !result.fail());
|
||||
|
||||
common_chat_msg msg;
|
||||
auto mapper = common_chat_peg_mapper(msg);
|
||||
mapper.from_ast(ctx.ast, result);
|
||||
|
||||
t.assert_equal("content", "a\xEF\xBF\xBD" "b", msg.content);
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user