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
5 changes: 5 additions & 0 deletions changelog.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,11 @@
Upcoming (TBD)
==============

Bug Fixes
--------
* Allow shell-style redirects with `/source` when the filename is unquoted.


Documentation
--------
* Badge color nit in `README.md`.
Expand Down
140 changes: 8 additions & 132 deletions mycli/client_commands.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,23 +4,23 @@
import logging
import os
import re
import shlex
from typing import TYPE_CHECKING, Any, cast

import click
import sqlparse

from mycli.compat import WIN
from mycli.config import write_default_config
from mycli.main_modes.repl import set_all_external_titles
from mycli.packages import special
from mycli.packages.batch_utils import statements_from_filehandle
from mycli.packages.filepaths import dir_path_exists
from mycli.packages.interactive_utils import confirm_destructive_query
from mycli.packages.ptoolkit.history import FileHistoryWithTimestamp
from mycli.packages.special import main as special_main
from mycli.packages.special.iocommands import expand_favorite_query
from mycli.packages.special.main import ArgType, SpecialCommandAlias
from mycli.packages.special.source import (
parse_source_arguments,
parse_source_filename,
source_special_command_is_safe,
)
from mycli.packages.sqlresult import SQLResult
from mycli.sqlexecute import SQLExecute

Expand All @@ -36,30 +36,6 @@
DSN_CONFIG_VALUE = object()
FAVORITES_CONFIG_VALUE = object()
HIDDEN_CONFIG_SECTIONS = frozenset({'alias_dsn', 'favorite_queries'})
INVALID_SOURCE_FILENAME = 'Source accepts exactly one filename; filenames containing spaces must be quoted.'
SOURCE_SAFE_SPECIAL_COMMANDS = frozenset({
'connect',
'fd',
'fs',
'help',
'l',
'nowarnings',
'ping',
'prompt',
'redirectformat',
'rehash',
'status',
'tableformat',
'timing',
'use',
'warnings',
'dt',
})
SOURCE_SAFE_SUBCOMMANDS = {
'config': frozenset({'help', 'get', 'search'}),
'dsn': frozenset({'help', 'list', 'show', 'save', 'delete'}),
'favorite': frozenset({'help', 'list', 'reload', 'run', 'save', 'delete'}),
}


def _render_config_value(value: Any) -> str:
Expand All @@ -68,106 +44,6 @@ def _render_config_value(value: Any) -> str:
return str(value)


def _parse_source_arguments(arg: str) -> tuple[str, bool, bool, bool]:
allow_special = False
show_queries = False
page_output = False
filename = arg
while arguments := filename.split(maxsplit=1):
if arguments[0] == '--special':
allow_special = True
elif arguments[0] == '--show':
show_queries = True
elif arguments[0] == '--page':
page_output = True
else:
break
filename = arguments[1] if len(arguments) == 2 else ''
return filename, allow_special, show_queries, page_output


def _has_unquoted_whitespace(value: str) -> bool:
quote: str | None = None
escaped = False
for character in value:
if escaped:
if quote is None and character.isspace():
return True
escaped = False
continue
if not WIN and character == '\\' and quote != "'":
escaped = True
continue
if character in ("'", '"'):
if quote is None:
quote = character
elif quote == character:
quote = None
elif quote is None and character.isspace():
return True
return False


def _parse_source_filename(filename: str) -> str:
if not filename:
return ''
if _has_unquoted_whitespace(filename):
raise ValueError(INVALID_SOURCE_FILENAME)
try:
arguments = shlex.split(filename, posix=not WIN)
except ValueError as error:
raise ValueError(f'Invalid source filename: {error}.') from None
if len(arguments) != 1:
raise ValueError(INVALID_SOURCE_FILENAME)
parsed_filename = arguments[0]
if WIN and len(parsed_filename) >= 2 and parsed_filename[0] == parsed_filename[-1] and parsed_filename[0] in ("'", '"'):
parsed_filename = parsed_filename[1:-1]
return parsed_filename


def _registered_special_command(query: str) -> tuple[str, str] | None:
command, _verbosity, arg = special.parse_special_command(query)
registered = special_main.COMMANDS.get(command)
if registered is None:
registered = special_main.COMMANDS.get(command.lower())
if registered is None:
return None
return registered.command.removeprefix('\\').removeprefix('/').lower(), arg


def _favorite_source_command_is_safe(arg: str) -> bool:
query, _error = expand_favorite_query(arg)
if query is None:
return True
return not any(special.is_special_command(statement.rstrip(';')) for statement in sqlparse.split(query))


def _source_special_command_is_safe(query: str) -> bool:
parsed = _registered_special_command(query)
if parsed is None:
return False

command, arg = parsed
if command == 'f':
return not arg or _favorite_source_command_is_safe(arg)
if command in ('fd', 'fs'):
return True
if command in SOURCE_SAFE_SPECIAL_COMMANDS:
return True

subcommands = SOURCE_SAFE_SUBCOMMANDS.get(command)
if subcommands is None:
return False
arguments = arg.split(maxsplit=1)
subcommand = arguments[0].lower() if arguments else 'help'
if subcommand not in subcommands:
return False
if command == 'favorite' and subcommand == 'run':
run_arg = arguments[1] if len(arguments) == 2 else ''
return not run_arg or _favorite_source_command_is_safe(run_arg)
return True


def _iter_config_values(
config: Mapping[str, Any],
prefix: str = '',
Expand Down Expand Up @@ -412,11 +288,11 @@ def change_db(self, arg: str, **_) -> Generator[SQLResult, None, None]:
yield SQLResult(status=msg)

def execute_from_file(self, arg: str, **_) -> Generator[SQLResult, None, None]:
filename, allow_special, show_queries, page_output = _parse_source_arguments(arg)
filename, allow_special, show_queries, page_output = parse_source_arguments(arg)
if page_output:
yield SQLResult(command={'name': 'source_page'})
try:
filename = _parse_source_filename(filename)
filename = parse_source_filename(filename)
except ValueError as error:
yield SQLResult(status=str(error), is_error=True)
return
Expand Down Expand Up @@ -450,7 +326,7 @@ def execute_from_file(self, arg: str, **_) -> Generator[SQLResult, None, None]:
is_error=True,
)
return
if not _source_special_command_is_safe(special_query):
if not source_special_command_is_safe(special_query):
command, _verbosity, _arg = special.parse_special_command(special_query)
yield SQLResult(
status=f'Special command is never permitted in source files: {command}.',
Expand Down
32 changes: 31 additions & 1 deletion mycli/packages/hybrid_redirection.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,34 @@
import functools
import logging
import re
import shlex

import sqlglot

from mycli.compat import WIN
from mycli.packages.special.delimitercommand import DelimiterCommand
from mycli.packages.special.source import (
parse_source_arguments,
parse_source_filename,
)

logger = logging.getLogger(__name__)
delimiter_command = DelimiterCommand()
SOURCE_COMMAND_PATTERN = re.compile(r'^([/]?source|[/\\]\.)\s+', re.IGNORECASE)
SOURCE_OPTIONS_PATTERN = re.compile(
r'^([/]?source|[/\\]\.)\s+(?P<options>(?:(?:--special|--show|--page)\s+)*)',
re.IGNORECASE,
)


def tokenize_command(command: str) -> list[sqlglot.Token]:
"""Tokenize a command without treating source options as SQL comments."""
source_match = SOURCE_OPTIONS_PATTERN.match(command)
if source_match:
options_start, options_end = source_match.span('options')
options = command[options_start:options_end].replace('-', '_')
command = command[:options_start] + options + command[options_end:]
return sqlglot.tokenize(command)


def find_token_indices(tokens: list[sqlglot.Token]) -> dict[str, list[int]]:
Expand Down Expand Up @@ -44,6 +64,16 @@ def find_sql_part(
):
leftmost_dollar_pos = tokens[true_dollar_indices[0]].start
sql_part = command[0:leftmost_dollar_pos].strip().removesuffix(delimiter_command.current).rstrip()
if SOURCE_COMMAND_PATTERN.match(sql_part):
source_arg_str = SOURCE_COMMAND_PATTERN.sub('', sql_part)
try:
filename, _allow_special, _show_queries, _page_output = parse_source_arguments(source_arg_str)
filename = parse_source_filename(filename)
except ValueError:
return ''
if not filename:
return ''
return sql_part
try:
statements = sqlglot.parse(sql_part, read='mysql')
except sqlglot.errors.ParseError:
Expand Down Expand Up @@ -142,7 +172,7 @@ def get_redirect_components(command: str) -> tuple[str | None, str | None, str |
"""Get the parts of a hybrid shell-style redirect command."""

try:
tokens = sqlglot.tokenize(command)
tokens = tokenize_command(command)
except sqlglot.errors.TokenError:
return None, None, None, None

Expand Down
13 changes: 10 additions & 3 deletions mycli/packages/special/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
'mycli.packages.special.iocommands',
'mycli.packages.special.llm',
'mycli.packages.special.main',
'mycli.packages.special.source',
]

import os
Expand Down Expand Up @@ -54,6 +55,10 @@
write_pipe_once,
write_tee,
)
from mycli.packages.special.source import (
parse_source_arguments,
parse_source_filename,
)

if not os.environ.get('MYCLI_LLM_OFF'):
from mycli.packages.special.llm import (
Expand Down Expand Up @@ -112,12 +117,15 @@ def sql_using_llm(*args, **kwargs): # type: ignore[no-redef, misc]
'is_llm_command',
'is_pager_enabled',
'is_redirected',
'is_special_command',
'is_show_favorite_query',
'is_show_warnings_enabled',
'is_special_command',
'is_timing_enabled',
'list_databases',
'list_tables',
'open_external_editor',
'parse_source_arguments',
'parse_source_filename',
'parse_special_command',
'ping',
'register_special_command',
Expand All @@ -131,10 +139,9 @@ def sql_using_llm(*args, **kwargs): # type: ignore[no-redef, misc]
'set_pager',
'set_pager_enabled',
'set_redirect',
'set_show_favorite_query',
'set_show_warnings_enabled',
'set_timing_enabled',
'set_show_favorite_query',
'is_show_favorite_query',
'special_command',
'split_queries',
'sql_using_llm',
Expand Down
Loading
Loading