diff --git a/changelog.md b/changelog.md index 5bfd6789..b9b99642 100644 --- a/changelog.md +++ b/changelog.md @@ -6,6 +6,7 @@ Features * Remove support for Python 3.10. * Exit faster when using a Boundary tunnel. * Add C-o u keybindings to insert literal Unix timestamps. +* Ability to complete partial enum values. Documentation diff --git a/mycli/constants.py b/mycli/constants.py index 9ab216b3..404bec6d 100644 --- a/mycli/constants.py +++ b/mycli/constants.py @@ -42,3 +42,12 @@ EMPTY_PASSWORD_FLAG_SENTINEL = -1 DEFAULT_PROMPT = "\\t \\u@\\h:\\d> " + +MYSQL_ESCAPES = { + '0': '\0', + 'b': '\b', + 'n': '\n', + 'r': '\r', + 't': '\t', + 'Z': '\x1a', +} diff --git a/mycli/packages/completion_engine.py b/mycli/packages/completion_engine.py index 45ca7d38..961fd1dc 100644 --- a/mycli/packages/completion_engine.py +++ b/mycli/packages/completion_engine.py @@ -5,8 +5,9 @@ from typing import Any, Callable, Literal import sqlparse -from sqlparse.sql import Comparison, Identifier, Token, Where +from sqlparse.sql import Comparison, Having, Identifier, Token, Where +from mycli.constants import MYSQL_ESCAPES from mycli.packages.special.dsn_aliases import DSN_SUBCOMMANDS from mycli.packages.special.favoritequeries import FAVORITE_SUBCOMMANDS from mycli.packages.special.main import COMMANDS as SPECIAL_COMMANDS @@ -21,6 +22,7 @@ r"(?P(?:`[^`]+`|[\w$]+)(?:\.(?:`[^`]+`|[\w$]+))?)\s*=\s*$", re.IGNORECASE, ) +_QUOTED_ENUM_TOKENS = re.compile(r"""--(?=\s|$)[^\n]*|\#[^\n]*|/\*[\s\S]*?(?:\*/|$)|`(?:``|[^`])*`|['"]""") # missing because not binary # BETWEEN @@ -391,6 +393,54 @@ def _emit_binary_or_comma(ctx: SuggestContext) -> list[Suggestion]: return fallback +def _emit_quoted_enum_value_or_nothing(ctx: SuggestContext) -> list[Suggestion]: + # Skip comments and identifiers before looking for an unfinished string. + offset = 0 + while match := _QUOTED_ENUM_TOKENS.search(ctx.text_before_cursor, offset): + quote = match.group() + offset = match.end() + if quote not in ("'", '"'): + continue + start = match.start() + value: list[str] = [] + closed = False + while offset < len(ctx.text_before_cursor): + char = ctx.text_before_cursor[offset] + offset += 1 + if char == '\\' and offset < len(ctx.text_before_cursor): + escaped = ctx.text_before_cursor[offset] + value.append(MYSQL_ESCAPES.get(escaped, escaped)) + offset += 1 + elif char == quote: + if ctx.text_before_cursor[offset : offset + 1] == quote: + value.append(quote) + offset += 1 + else: + closed = True + break + else: + value.append(char) + if closed: + continue + prefix = ctx.text_before_cursor[:start] + if not _enum_value_suggestion(prefix, ctx.full_text): + return [] + # Reuse the normal clause and table resolution with the value removed. + context_text = prefix + ' ' + ctx.full_text[start:] + for suggestion in suggest_type(context_text, prefix + ' '): + if suggestion['type'] == 'enum_value': + return [ + { + **suggestion, + 'value_prefix': ''.join(value), + 'quote': quote, + 'replacement_length': len(ctx.text_before_cursor) - start, + } + ] + return [] + return [] + + def _word_starts_with_digit_or_dot(ctx: SuggestContext) -> bool: return bool(ctx.word_before_cursor and re.match(r'^[\d\.]', ctx.word_before_cursor[0])) @@ -403,6 +453,18 @@ def _word_inside_single_or_double_quotes(ctx: SuggestContext) -> bool: return bool(ctx.word_before_cursor and _is_single_or_double_quoted(ctx)) +def _word_inside_or_starts_with_quote(ctx: SuggestContext) -> bool: + return _word_starts_with_quote(ctx) or _word_inside_single_or_double_quotes(ctx) + + +def _word_could_be_enum_value(ctx: SuggestContext) -> bool: + # the check for the "=" or "having" values seem to be needed in case of failing to see a Having + # token, which seems to be a sqlparse bug + return _word_inside_or_starts_with_quote(ctx) and ( + isinstance(ctx.token, (Having, Where)) or (isinstance(ctx.token, Token) and ctx.token.value.lower() in ['=', 'having']) + ) + + def _token_is_none(ctx: SuggestContext) -> bool: return ctx.token is None @@ -432,6 +494,11 @@ def _token_is_binary_or_comma(ctx: SuggestContext) -> bool: SUGGEST_BASED_ON_LAST_TOKEN_RULES = [ + SuggestRule( + 'quoted_enum_value', + _word_could_be_enum_value, + _emit_quoted_enum_value_or_nothing, + ), SuggestRule( 'guard_number_or_dot', _word_starts_with_digit_or_dot, diff --git a/mycli/sqlcompleter.py b/mycli/sqlcompleter.py index 218694c2..24f49f99 100644 --- a/mycli/sqlcompleter.py +++ b/mycli/sqlcompleter.py @@ -1536,6 +1536,8 @@ def tiebreaker_key(candidate: str) -> tuple[float, str]: length_based_on_path = False source_file_completion_length: int | None = None config_property_length: int | None = None + enum_value_length: int | None = None + enum_quote: str | None = None completion_filter_text = text_for_len rank = 0 @@ -1875,6 +1877,18 @@ def tiebreaker_key(candidate: str) -> tuple[float, str]: suggestion["column"], suggestion.get("parent"), ) + if 'value_prefix' in suggestion: + prefix = suggestion['value_prefix'] + enum_quote = suggestion['quote'] + suffix = document.text_after_cursor + if not suffix.startswith(enum_quote) and enum_quote in suffix: + return [] + if suffix and not (suffix.startswith(enum_quote) or suffix[0].isspace() or suffix[0] in ';,)'): + return [] + enum_value_length = suggestion['replacement_length'] + completion_filter_text = prefix.lower() + completions = [(*item, rank) for item in self.find_fuzzy_matches(prefix, prefix.lower(), enum_values)] + break if enum_values: quoted_values = [self._quote_sql_string(value) for value in enum_values] completions = [ @@ -1904,7 +1918,16 @@ def completion_sort_key(item: tuple[str, int, int], text_for_len: str) -> tuple[ sorted_completions = sorted(completions, key=lambda item: completion_sort_key(item, completion_filter_text)) uniq_completions_str = dict.fromkeys(x[0] for x in sorted_completions) - if config_property_length is not None: + if enum_value_length is not None and enum_quote is not None: + closing_quote = '' if document.text_after_cursor.startswith(enum_quote) else enum_quote + return ( + Completion( + enum_quote + x.replace('\\', '\\\\').replace(enum_quote, enum_quote * 2) + closing_quote, + -enum_value_length, + ) + for x in uniq_completions_str + ) + elif config_property_length is not None: return (Completion(x, -config_property_length) for x in uniq_completions_str) elif source_file_completion_length is not None: return (Completion(x, -source_file_completion_length) for x in uniq_completions_str) diff --git a/test/pytests/test_completion_engine.py b/test/pytests/test_completion_engine.py index 3d50abf7..c85614cb 100644 --- a/test/pytests/test_completion_engine.py +++ b/test/pytests/test_completion_engine.py @@ -25,6 +25,7 @@ _emit_nothing, _emit_on, _emit_procedure, + _emit_quoted_enum_value_or_nothing, _emit_relation_like, _emit_relation_name, _emit_select_like, @@ -140,6 +141,94 @@ def test_where_equals_suggests_enum_values_first(): ]) +@pytest.mark.parametrize('quote', ["'", '"']) +@pytest.mark.parametrize('clause', ['WHERE', 'HAVING']) +@pytest.mark.parametrize('separator', ['=', ' = ']) +def test_quoted_enum_prefix_context(quote: str, clause: str, separator: str) -> None: + expression = f'SELECT * FROM tabl t {clause} `t`.`foo`{separator}{quote}in_pro' + assert suggest_type(expression, expression) == [ + { + 'type': 'enum_value', + 'tables': [(None, 'tabl', 't')], + 'column': '`foo`', + 'parent': '`t`', + 'value_prefix': 'in_pro', + 'quote': quote, + 'replacement_length': 7, + } + ] + + +@pytest.mark.parametrize( + 'expression', + [ + "SELECT * FROM tabl WHERE foo > 'pen", + "SELECT * FROM tabl WHERE foo = 'pending'", + "SELECT * FROM tabl WHERE 'foo = pen", + "SELECT * FROM tabl -- foo = 'pen", + "SELECT * FROM tabl # foo = 'pen", + "SELECT * FROM tabl /* foo = 'pen", + ], +) +def test_quoted_enum_context_excludes_other_strings_and_comments(expression: str) -> None: + assert not any(item['type'] == 'enum_value' for item in suggest_type(expression, expression)) + + +@pytest.mark.parametrize( + ('quote', 'prefix', 'decoded'), + [ + ("'", "O''Br", "O'Br"), + ('"', 'say ""he', 'say "he'), + ("'", r"O\'Br", "O'Br"), + ('"', r'say \"he', 'say "he'), + ("'", r'a\\b', 'a\\b'), + ("'", r'a\nb', 'a\nb'), + ("'", r'a\tb', 'a\tb'), + ("'", r'a\0b', 'a\0b'), + ("'", r'a\qb', 'aqb'), + ("'", 'unfinished\\', 'unfinished\\'), + ], +) +def test_quoted_enum_prefix_decodes_sql_escapes(quote: str, prefix: str, decoded: str) -> None: + expression = f'SELECT * FROM tabl WHERE foo = {quote}{prefix}' + + assert suggest_type(expression, expression) == [ + { + 'type': 'enum_value', + 'tables': [(None, 'tabl', None)], + 'column': 'foo', + 'parent': None, + 'value_prefix': decoded, + 'quote': quote, + 'replacement_length': len(prefix) + 1, + } + ] + + +@pytest.mark.parametrize('has_enum', [False, True]) +def test_quoted_enum_emitter_selects_only_enum_suggestions(monkeypatch: pytest.MonkeyPatch, has_enum: bool) -> None: + expression = "SELECT * FROM tabl WHERE foo = 'pen" + context = _build_suggest_context(None, expression, "'pen", expression, empty_identifier()) + enum_suggestion = {'type': 'enum_value', 'tables': [(None, 'tabl', None)], 'column': 'foo', 'parent': None} + suggestions = [{'type': 'keyword'}] + if has_enum: + suggestions.append(enum_suggestion) + monkeypatch.setattr(completion_engine, 'suggest_type', lambda *_args: suggestions) + + result = _emit_quoted_enum_value_or_nothing(context) + + expected = [{**enum_suggestion, 'value_prefix': 'pen', 'quote': "'", 'replacement_length': 4}] if has_enum else [] + assert result == expected + + +def test_quoted_enum_emitter_returns_nothing_when_context_has_no_suggestions(monkeypatch: pytest.MonkeyPatch) -> None: + expression = "SELECT * FROM tabl WHERE foo = 'pen" + context = _build_suggest_context(None, expression, "'pen", expression, empty_identifier()) + monkeypatch.setattr(completion_engine, 'suggest_type', lambda *_args: []) + + assert _emit_quoted_enum_value_or_nothing(context) == [] + + def test_enum_value_suggestion_returns_none_without_equals_context(): expression = 'SELECT * FROM tabl WHERE foo' suggestion = _enum_value_suggestion(expression, expression) diff --git a/test/pytests/test_smart_completion_public_schema_only.py b/test/pytests/test_smart_completion_public_schema_only.py index ce44f8c7..a2b29f55 100644 --- a/test/pytests/test_smart_completion_public_schema_only.py +++ b/test/pytests/test_smart_completion_public_schema_only.py @@ -355,6 +355,91 @@ def test_enum_value_completion(completer, complete_event): ] +@pytest.mark.parametrize('quote', ["'", '"']) +@pytest.mark.parametrize('prefix', ['', 'pen']) +def test_quoted_enum_completion(completer, complete_event, quote: str, prefix: str) -> None: + completer.completion_match_order = ('perfect',) + text = 'SELECT * FROM orders WHERE status = ' + quote + prefix + result = list(completer.get_completions(Document(text), complete_event)) + expected = ['pending', 'shipped'] if not prefix else ['pending'] + assert [(item.text, item.start_position) for item in result] == [(quote + value + quote, -len(prefix) - 1) for value in expected] + + +@pytest.mark.parametrize('quote', ["'", '"']) +@pytest.mark.parametrize('clause', ['WHERE', 'HAVING']) +def test_quoted_enum_completion_preserves_closing_quote(completer, complete_event, quote: str, clause: str) -> None: + text = f'SELECT * FROM orders o {clause} `o`.`status` = {quote}pen' + document = Document(text + quote + ';', cursor_position=len(text)) + result = list(completer.get_completions(document, complete_event)) + pending = next(item for item in result if item.text == quote + 'pending') + updated = text[: len(text) + pending.start_position] + pending.text + document.text_after_cursor + assert updated == f'SELECT * FROM orders o {clause} `o`.`status` = {quote}pending{quote};' + + +@pytest.mark.parametrize( + ('prefix', 'value'), + [('in pro', 'in progress'), ('a.b', 'a.b-c'), ("O''B", "O'Brien"), (r"O\'B", "O'Brien"), (r'a\\b', r'a\bc')], +) +def test_quoted_enum_completion_replaces_entire_prefix(completer, complete_event, prefix: str, value: str) -> None: + completer.extend_enum_values([('orders', 'status', [value])]) + text = "SELECT * FROM orders WHERE status = '" + prefix + result = list(completer.get_completions(Document(text), complete_event)) + assert len(result) == 1 + completion = result[0] + updated = text[: len(text) + completion.start_position] + completion.text + assert updated == "SELECT * FROM orders WHERE status = '" + value.replace('\\', '\\\\').replace("'", "''") + "'" + + +@pytest.mark.parametrize( + 'text', + [ + 'SELECT * FROM orders WHERE status = pen', + "SELECT * FROM orders WHERE status = 'zzzzzzzzzz", + "SELECT * FROM orders WHERE ordered_date = 'pen", + "SELECT 'status = pen", + "SELECT * FROM orders WHERE status = 'pending'", + "SELECT * FROM orders -- status = 'pen", + "SELECT * FROM orders /* status = 'pen", + ], +) +def test_quoted_enum_completion_does_not_leak_values(completer, complete_event, text: str) -> None: + result = list(completer.get_completions(Document(text), complete_event)) + assert not any(item.text in ("'pending'", "'shipped'") for item in result) + + +@pytest.mark.parametrize('suffix', ["ding'", " ding'"]) +def test_quoted_enum_completion_skips_existing_value_suffix(completer, complete_event, suffix: str) -> None: + text = "SELECT * FROM orders WHERE status = 'pen" + document = Document(text + suffix, cursor_position=len(text)) + assert list(completer.get_completions(document, complete_event)) == [] + + +@pytest.mark.parametrize('method, expected', [('perfect', []), ('regex', ["'pending'"])]) +def test_quoted_enum_completion_uses_configured_matching(completer, complete_event, method: str, expected: list[str]) -> None: + completer.completion_match_order = (method,) + text = "SELECT * FROM orders WHERE status = 'pnd" + assert [item.text for item in completer.get_completions(Document(text), complete_event)] == expected + + +@pytest.mark.parametrize('column, prefix', [('ordered_date', 'pen'), ('status', 'zzzzzzzzzz')]) +def test_quoted_enum_completion_has_no_sql_fallback(completer, complete_event, column: str, prefix: str) -> None: + text = f"SELECT * FROM orders WHERE {column} = '{prefix}" + assert list(completer.get_completions(Document(text), complete_event)) == [] + + +def test_quoted_enum_completion_after_previous_statement_and_literal(completer, complete_event) -> None: + text = "SELECT 'other'; SELECT * FROM orders WHERE status = 'shipped' OR status='pen" + result = list(completer.get_completions(Document(text), complete_event)) + assert [(item.text, item.start_position) for item in result] == [("'pending'", -4)] + + +def test_double_quoted_enum_completion_escapes_embedded_quotes(completer, complete_event) -> None: + completer.extend_enum_values([('orders', 'status', ['say "hello"'])]) + text = 'SELECT * FROM orders WHERE status = "say ""he' + result = list(completer.get_completions(Document(text), complete_event)) + assert [(item.text, item.start_position) for item in result] == [('"say ""hello"""', -9)] + + def test_function_name_completion(completer, complete_event): text = "SELECT MA" position = len("SELECT MA") diff --git a/test/pytests/test_sqlcompleter.py b/test/pytests/test_sqlcompleter.py index 0fb6e6e3..b8a1d844 100644 --- a/test/pytests/test_sqlcompleter.py +++ b/test/pytests/test_sqlcompleter.py @@ -8,6 +8,7 @@ from prompt_toolkit.document import Document import pytest +from mycli.packages.polars_completion import PolarsCompletion import mycli.sqlcompleter from mycli.sqlcompleter import Fuzziness, SQLCompleter @@ -41,6 +42,102 @@ def make_completer(**kwargs) -> SQLCompleter: return comp +def test_invalid_completion_tiebreaker_falls_back_to_frecency() -> None: + completer = make_completer(completion_tiebreaker='unknown') + + assert completer.completion_tiebreaker == 'frecency' + assert completer.completion_config_errors == ['Invalid completion_tiebreaker; using frecency.'] + + +def test_extend_builtin_functions_ignores_generator() -> None: + completer = make_completer() + original_functions = completer.functions.copy() + + completer.extend_functions((item for item in [('test', 'custom_function')]), builtin=True) + + assert completer.functions == original_functions + + +def test_polars_completions_preserve_display_and_replacement(monkeypatch: pytest.MonkeyPatch) -> None: + candidate = PolarsCompletion(text='select(', display='select', display_meta='DataFrame method', start_position=-3) + transform = Mock(return_value=[candidate]) + monkeypatch.setattr(mycli.sqlcompleter, 'complete_polars_transform', transform) + completer = make_completer() + text = 'SELECT 1 .| df.sel' + + result = list(completer.get_completions(Document(text), None)) + + transform.assert_called_once_with(text) + assert len(result) == 1 + assert result[0].text == 'select(' + assert result[0].start_position == -3 + assert result[0].display_text == 'select' + assert result[0].display_meta_text == 'DataFrame method' + + +def test_get_completions_can_override_smart_mode(monkeypatch: pytest.MonkeyPatch) -> None: + completer = make_completer(smart_completion=True) + matches = Mock(return_value=[('select', Fuzziness.PERFECT)]) + monkeypatch.setattr(completer, 'find_matches', matches) + suggestions = Mock(side_effect=AssertionError('smart completion must not run')) + monkeypatch.setattr(mycli.sqlcompleter, 'suggest_type', suggestions) + + result = list(completer.get_completions(Document('sel'), None, smart_completion=False)) + + assert [(item.text, item.start_position) for item in result] == [('select', -3)] + assert matches.call_args.kwargs['start_only'] is True + assert matches.call_args.kwargs['fuzzy'] is False + suggestions.assert_not_called() + assert completer.smart_completion is True + + +@pytest.mark.parametrize('suggestion_type', ['favoritequery', 'favoritequery_template_key']) +@pytest.mark.parametrize('instance_exists', [False, True]) +def test_favorite_completions_without_registry_methods( + monkeypatch: pytest.MonkeyPatch, suggestion_type: str, instance_exists: bool +) -> None: + if instance_exists: + monkeypatch.setattr(mycli.sqlcompleter.FavoriteQueries, 'instance', SimpleNamespace(), raising=False) + else: + monkeypatch.delattr(mycli.sqlcompleter.FavoriteQueries, 'instance', raising=False) + monkeypatch.setattr(mycli.sqlcompleter, 'suggest_type', lambda *args: [{'type': suggestion_type}]) + + assert list(make_completer().get_completions(Document('/f '), None)) == [] + + +def test_dsn_completions_without_registry(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delattr(mycli.sqlcompleter.DsnAliases, 'instance', raising=False) + monkeypatch.setattr(mycli.sqlcompleter, 'suggest_type', lambda *args: [{'type': 'dsn_alias'}]) + + assert list(make_completer().get_completions(Document('/dsn delete '), None)) == [] + + +def test_unknown_suggestion_type_is_ignored(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(mycli.sqlcompleter, 'suggest_type', lambda *args: [{'type': 'unknown'}]) + + assert list(make_completer().get_completions(Document(''), None)) == [] + + +def test_enum_without_metadata_preserves_other_suggestions(monkeypatch: pytest.MonkeyPatch) -> None: + completer = make_completer() + completer.keywords = ['select'] + monkeypatch.setattr( + mycli.sqlcompleter, + 'suggest_type', + lambda *args: [{'type': 'enum_value', 'tables': [], 'column': 'status'}, {'type': 'keyword'}], + ) + + assert [item.text for item in completer.get_completions(Document(''), None)] == ['SELECT'] + + +def test_quoted_enum_completion_rejects_nonempty_unclosed_suffix() -> None: + completer = make_completer() + prefix = "SELECT * FROM orders WHERE status = 'pen" + document = Document(prefix + 'ding', cursor_position=len(prefix)) + + assert list(completer.get_completions(document, None)) == [] + + @pytest.mark.parametrize( ('item', 'expected'), [