Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 50 additions & 23 deletions src/prism.c
Original file line number Diff line number Diff line change
Expand Up @@ -6703,11 +6703,11 @@ parse_symbol_encoding_validate_other(pm_parser_t *parser, const pm_token_t *loca
* to ensure that all characters are valid in the encoding.
*/
static PRISM_INLINE pm_node_flags_t
parse_symbol_encoding(pm_parser_t *parser, const pm_token_t *location, const pm_string_t *contents, bool validate) {
if (parser->explicit_encoding != NULL) {
parse_symbol_encoding(pm_parser_t *parser, const pm_encoding_t *explicit_encoding, const pm_token_t *location, const pm_string_t *contents, bool validate) {
if (explicit_encoding != NULL) {
// A Symbol may optionally have its encoding explicitly set. This will
// happen if an escape sequence results in a non-ASCII code point.
if (parser->explicit_encoding == PM_ENCODING_UTF_8_ENTRY) {
if (explicit_encoding == PM_ENCODING_UTF_8_ENTRY) {
if (validate) parse_symbol_encoding_validate_utf8(parser, location, contents);
return PM_SYMBOL_FLAGS_FORCED_UTF8_ENCODING;
} else if (parser->encoding == PM_ENCODING_US_ASCII_ENTRY) {
Expand Down Expand Up @@ -6761,7 +6761,7 @@ pm_symbol_node_create(pm_parser_t *parser, const pm_token_t *opening, const pm_t
*/
static pm_symbol_node_t *
pm_symbol_node_create_current_string(pm_parser_t *parser, const pm_token_t *opening, const pm_token_t *value, const pm_token_t *closing) {
pm_symbol_node_t *node = pm_symbol_node_create_unescaped(parser, opening, value, closing, &parser->current_string, parse_symbol_encoding(parser, value, &parser->current_string, false));
pm_symbol_node_t *node = pm_symbol_node_create_unescaped(parser, opening, value, closing, &parser->current_string, parse_symbol_encoding(parser, parser->explicit_encoding, value, &parser->current_string, false));
parser->current_string = PM_STRING_EMPTY;
return node;
}
Expand All @@ -6779,7 +6779,7 @@ pm_symbol_node_label_create(pm_parser_t *parser, const pm_token_t *token) {

assert((label.end - label.start) >= 0);
pm_string_shared_init(&node->unescaped, label.start, label.end);
pm_node_flag_set(UP(node), parse_symbol_encoding(parser, &label, &node->unescaped, false));
pm_node_flag_set(UP(node), parse_symbol_encoding(parser, parser->explicit_encoding, &label, &node->unescaped, false));

return node;
}
Expand Down Expand Up @@ -6855,7 +6855,7 @@ pm_string_node_to_symbol_node(pm_parser_t *parser, pm_string_node_t *node, const
.end = parser->start + node->content_loc.start + node->content_loc.length
};

pm_node_flag_set(UP(new_node), parse_symbol_encoding(parser, &content, &node->unescaped, true));
pm_node_flag_set(UP(new_node), parse_symbol_encoding(parser, parser->explicit_encoding, &content, &node->unescaped, true));

/* The old node is arena-allocated so no explicit free is needed. */
return new_node;
Expand Down Expand Up @@ -8582,6 +8582,10 @@ lex_identifier(pm_parser_t *parser, bool previous_command_start) {
// identifier, then we'll optionally accept it.
lex_state_set(parser, PM_LEX_STATE_ARG | PM_LEX_STATE_LABELED);
(void) match(parser, ':');

/* A label is a symbol lexed inline rather than through a lex
* mode, so it clears the encoding here. */
parser->explicit_encoding = NULL;
return PM_TOKEN_LABEL;
}

Expand All @@ -8608,6 +8612,10 @@ lex_identifier(pm_parser_t *parser, bool previous_command_start) {
// identifier, then we'll optionally accept it.
lex_state_set(parser, PM_LEX_STATE_ARG | PM_LEX_STATE_LABELED);
(void) match(parser, ':');

/* A label is a symbol lexed inline rather than through a lex
* mode, so it clears the encoding here. */
parser->explicit_encoding = NULL;
return PM_TOKEN_LABEL;
}
}
Expand Down Expand Up @@ -9533,6 +9541,13 @@ lex_question_mark(pm_parser_t *parser) {
return PM_TOKEN_QUESTION_MARK;
}

/*
* A literal takes its encoding from its own contents. Literals that push a
* lex mode clear this in lex_mode_push_*; a character literal is lexed
* inline, so it clears the encoding here.
*/
parser->explicit_encoding = NULL;

if (parser->current.end >= parser->end) {
pm_parser_err_current(parser, PM_ERR_INCOMPLETE_QUESTION_MARK);
pm_string_shared_init(&parser->current_string, parser->current.start + 1, parser->current.end);
Expand Down Expand Up @@ -11101,6 +11116,13 @@ parser_lex(pm_parser_t *parser) {
if (peek(parser) == '"' || peek(parser) == '\'') {
lex_mode_push_string(parser, peek(parser) == '"', false, '\0', *parser->current.end);
parser->current.end++;
} else {
/*
* A quoted symbol clears its encoding by pushing a lex
* mode above. A bare symbol is lexed inline, so it
* clears the encoding here.
*/
parser->explicit_encoding = NULL;
}

lex_state_set(parser, PM_LEX_STATE_FNAME);
Expand Down Expand Up @@ -16014,9 +16036,9 @@ PM_STATIC_ASSERT(__LINE__, ((int) PM_STRING_FLAGS_FORCED_UTF8_ENCODING) == ((int
* to potentially mark the string's flags to indicate how to encode it.
*/
static PRISM_INLINE pm_node_flags_t
parse_unescaped_encoding(const pm_parser_t *parser) {
if (parser->explicit_encoding != NULL) {
if (parser->explicit_encoding == PM_ENCODING_UTF_8_ENTRY) {
parse_unescaped_encoding(const pm_parser_t *parser, const pm_encoding_t *explicit_encoding) {
if (explicit_encoding != NULL) {
if (explicit_encoding == PM_ENCODING_UTF_8_ENTRY) {
// If the there's an explicit encoding and it's using a UTF-8 escape
// sequence, then mark the string as UTF-8.
return PM_STRING_FLAGS_FORCED_UTF8_ENCODING;
Expand Down Expand Up @@ -16046,7 +16068,7 @@ parse_string_part(pm_parser_t *parser, uint16_t depth) {
// ^^^^ ^ ^^^^
case PM_TOKEN_STRING_CONTENT: {
pm_node_t *node = UP(pm_string_node_create_current_string(parser, NULL, &parser->current, NULL));
pm_node_flag_set(node, parse_unescaped_encoding(parser));
pm_node_flag_set(node, parse_unescaped_encoding(parser, parser->explicit_encoding));

parser_lex(parser);
return node;
Expand Down Expand Up @@ -16220,7 +16242,7 @@ parse_symbol(pm_parser_t *parser, pm_lex_mode_t *lex_mode, pm_lex_state_t next_s

pm_symbol_node_t *symbol = pm_symbol_node_create(parser, &opening, &parser->previous, NULL);
pm_string_shared_init(&symbol->unescaped, parser->previous.start, parser->previous.end);
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, &parser->previous, &symbol->unescaped, false));
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, parser->explicit_encoding, &parser->previous, &symbol->unescaped, false));

return UP(symbol);
}
Expand Down Expand Up @@ -16321,7 +16343,7 @@ parse_symbol(pm_parser_t *parser, pm_lex_mode_t *lex_mode, pm_lex_state_t next_s
expect1(parser, PM_TOKEN_STRING_END, PM_ERR_SYMBOL_TERM_DYNAMIC);
}

return UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, &content, &unescaped, false)));
return UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, parser->explicit_encoding, &content, &unescaped, false)));
}

/**
Expand All @@ -16341,7 +16363,7 @@ parse_undef_argument(pm_parser_t *parser, uint16_t depth) {

pm_symbol_node_t *symbol = pm_symbol_node_create(parser, NULL, &parser->previous, NULL);
pm_string_shared_init(&symbol->unescaped, parser->previous.start, parser->previous.end);
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, &parser->previous, &symbol->unescaped, false));
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, parser->explicit_encoding, &parser->previous, &symbol->unescaped, false));

return UP(symbol);
}
Expand Down Expand Up @@ -16377,7 +16399,7 @@ parse_alias_argument(pm_parser_t *parser, bool first, uint16_t depth) {

pm_symbol_node_t *symbol = pm_symbol_node_create(parser, NULL, &parser->previous, NULL);
pm_string_shared_init(&symbol->unescaped, parser->previous.start, parser->previous.end);
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, &parser->previous, &symbol->unescaped, false));
pm_node_flag_set(UP(symbol), parse_symbol_encoding(parser, parser->explicit_encoding, &parser->previous, &symbol->unescaped, false));

return UP(symbol);
}
Expand Down Expand Up @@ -16682,7 +16704,7 @@ parse_strings(pm_parser_t *parser, pm_node_t *current, bool accepts_label, uint1
expect1(parser, PM_TOKEN_STRING_END, PM_ERR_STRING_LITERAL_EOF);
node = UP(pm_interpolated_string_node_create(parser, &opening, &parts, &parser->previous));
} else if (accept1(parser, PM_TOKEN_LABEL_END)) {
node = UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, &content, &unescaped, true)));
node = UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, parser->explicit_encoding, &content, &unescaped, true)));
if (!label_allowed) pm_parser_err_node(parser, node, PM_ERR_UNEXPECTED_LABEL);
} else if (match1(parser, PM_TOKEN_EOF)) {
pm_parser_err_token(parser, &opening, PM_ERR_STRING_LITERAL_EOF);
Expand All @@ -16702,11 +16724,12 @@ parse_strings(pm_parser_t *parser, pm_node_t *current, bool accepts_label, uint1
// plain string) or if it's not then it has interpolation.
pm_token_t content = parser->current;
pm_string_t unescaped = parser->current_string;
const pm_encoding_t *explicit_encoding = parser->explicit_encoding;
parser_lex(parser);

if (match2(parser, PM_TOKEN_STRING_END, PM_TOKEN_EOF)) {
node = UP(pm_string_node_create_unescaped(parser, &opening, &content, &parser->current, &unescaped));
pm_node_flag_set(node, parse_unescaped_encoding(parser));
pm_node_flag_set(node, parse_unescaped_encoding(parser, explicit_encoding));

// Kind of odd behavior, but basically if we have an
// unterminated string and it ends in a newline, we back up one
Expand All @@ -16721,14 +16744,14 @@ parse_strings(pm_parser_t *parser, pm_node_t *current, bool accepts_label, uint1
parser->previous.type = 0;
}
} else if (accept1(parser, PM_TOKEN_LABEL_END)) {
node = UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, &content, &unescaped, true)));
node = UP(pm_symbol_node_create_unescaped(parser, &opening, &content, &parser->previous, &unescaped, parse_symbol_encoding(parser, explicit_encoding, &content, &unescaped, true)));
if (!label_allowed) pm_parser_err_node(parser, node, PM_ERR_UNEXPECTED_LABEL);
} else {
// If we get here, then we have interpolation so we'll need
// to create a string or symbol node with interpolation.
pm_node_list_t parts = { 0 };
pm_node_t *part = UP(pm_string_node_create_unescaped(parser, NULL, &parser->previous, NULL, &unescaped));
pm_node_flag_set(part, parse_unescaped_encoding(parser));
pm_node_flag_set(part, parse_unescaped_encoding(parser, explicit_encoding));
pm_node_list_append(parser->arena, &parts, part);

while (!match3(parser, PM_TOKEN_STRING_END, PM_TOKEN_LABEL_END, PM_TOKEN_EOF)) {
Expand Down Expand Up @@ -18851,7 +18874,7 @@ parse_string_array(pm_parser_t *parser, uint16_t depth) {
}
case PM_TOKEN_STRING_CONTENT: {
pm_node_t *string = UP(pm_string_node_create_current_string(parser, NULL, &parser->current, NULL));
pm_node_flag_set(string, parse_unescaped_encoding(parser));
pm_node_flag_set(string, parse_unescaped_encoding(parser, parser->explicit_encoding));
parser_lex(parser);

if (current == NULL) {
Expand Down Expand Up @@ -18969,6 +18992,10 @@ parse_symbol_array(pm_parser_t *parser, uint16_t depth) {
while (!match2(parser, PM_TOKEN_STRING_END, PM_TOKEN_EOF)) {
switch (parser->current.type) {
case PM_TOKEN_WORDS_SEP: {
/* Reset the explicit encoding if we hit a separator since each
* element can have its own encoding. */
parser->explicit_encoding = NULL;

if (current == NULL) {
/* If we hit a separator before we have any content, then we
* don't need to do anything. */
Expand Down Expand Up @@ -19511,7 +19538,7 @@ parse_expression_prefix(pm_parser_t *parser, pm_binding_power_t binding_power, u
NULL
));

pm_node_flag_set(node, parse_unescaped_encoding(parser));
pm_node_flag_set(node, parse_unescaped_encoding(parser, parser->explicit_encoding));

// Skip past the character literal here, since now we have handled
// parser->explicit_encoding correctly.
Expand Down Expand Up @@ -19749,7 +19776,7 @@ parse_expression_prefix(pm_parser_t *parser, pm_binding_power_t binding_power, u
// content and we're at the end of the heredoc, so we can return
// just a string node with the heredoc opening and closing as
// its opening and closing.
pm_node_flag_set(part, parse_unescaped_encoding(parser));
pm_node_flag_set(part, parse_unescaped_encoding(parser, parser->explicit_encoding));
pm_string_node_t *cast = (pm_string_node_t *) part;

cast->opening_loc = TOK2LOC(parser, &opening);
Expand Down Expand Up @@ -20608,7 +20635,7 @@ parse_expression_prefix(pm_parser_t *parser, pm_binding_power_t binding_power, u

if (match1(parser, PM_TOKEN_STRING_END)) {
pm_node_t *node = UP(pm_xstring_node_create_unescaped(parser, &opening, &content, &parser->current, &unescaped));
pm_node_flag_set(node, parse_unescaped_encoding(parser));
pm_node_flag_set(node, parse_unescaped_encoding(parser, parser->explicit_encoding));
parser_lex(parser);
return node;
}
Expand All @@ -20618,7 +20645,7 @@ parse_expression_prefix(pm_parser_t *parser, pm_binding_power_t binding_power, u
node = pm_interpolated_xstring_node_create(parser, &opening, &opening);

pm_node_t *part = UP(pm_string_node_create_unescaped(parser, NULL, &parser->previous, NULL, &unescaped));
pm_node_flag_set(part, parse_unescaped_encoding(parser));
pm_node_flag_set(part, parse_unescaped_encoding(parser, parser->explicit_encoding));

pm_interpolated_xstring_node_append(parser->arena, node, part);
} else {
Expand Down
56 changes: 56 additions & 0 deletions test/prism/encoding/string_encoding_test.rb
Original file line number Diff line number Diff line change
Expand Up @@ -86,8 +86,64 @@ def test_multibyte_escapes
end
end

def test_encodings_are_not_leaked_across_array_elements
{
"[?\\u00E9, ?a, ?\\x61]" => [Encoding::UTF_8, Encoding::ASCII_8BIT, Encoding::ASCII_8BIT],
"[\"\\u00E9\", ?a]" => [Encoding::UTF_8, Encoding::ASCII_8BIT],
"[?\\u00E9, \"a\"]" => [Encoding::UTF_8, Encoding::ASCII_8BIT],
"[?\\u00E9, :a]" => [Encoding::UTF_8, Encoding::US_ASCII],
"[\"\\u00E9\", :a]" => [Encoding::UTF_8, Encoding::US_ASCII],
"[:\"\\u00E9\", :a]" => [Encoding::UTF_8, Encoding::US_ASCII],
"[:\"\\u00E9\", :+]" => [Encoding::UTF_8, Encoding::US_ASCII],
"%I[\\u00E9 a]" => [Encoding::UTF_8, Encoding::US_ASCII],
"%W[\\u00E9 a]" => [Encoding::UTF_8, Encoding::ASCII_8BIT]
}.each do |expression, expected|
result = Prism.parse("# encoding: ascii-8bit\n#{expression}")
assert_predicate result, :success?, "failed to parse: #{expression}"

actual = result.statement.elements.map { |node| binary_node_encoding(node) }
assert_equal expected, actual, expression
end
end

def test_encodings_are_not_leaked_across_hash_elements
{
"{\"\\u00E9\" => 1, a: 2}" => [Encoding::UTF_8, Encoding::US_ASCII],
"{\"\\u00E9\": :a}" => [Encoding::UTF_8],
"{\"\\u00E9\": ?a}" => [Encoding::UTF_8],
"{a: 1, \"\\u00E9\": 2}" => [Encoding::US_ASCII, Encoding::UTF_8]
}.each do |expression, expected|
["ascii-8bit", "us-ascii"].each do |encoding|
result = Prism.parse("# encoding: #{encoding}\n#{expression}")
assert_predicate result, :success?, "failed to parse: #{expression}"

actual = result.statement.elements.map { |assoc| binary_node_encoding(assoc.key) }
assert_equal expected, actual, "#{encoding} #{expression}"
end
end
end

private

def binary_node_encoding(node)
case node
when StringNode
node.forced_utf8_encoding? ? Encoding::UTF_8 : Encoding::ASCII_8BIT
when SymbolNode
if node.forced_utf8_encoding?
Encoding::UTF_8
elsif node.forced_binary_encoding?
Encoding::ASCII_8BIT
elsif node.forced_us_ascii_encoding?
Encoding::US_ASCII
else
Encoding::ASCII_8BIT
end
else
flunk "unexpected node type: #{node.class}"
end
end

def assert_encoding(encoding)
escapes = ["\\x00", "\\x7F", "\\x80", "\\xFF", "\\u{00}", "\\u{7F}", "\\u{80}", "\\M-\\C-?"]
escapes = escapes.concat(escapes.product(escapes).map(&:join))
Expand Down