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
1 change: 1 addition & 0 deletions changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ Upcoming (TBD)
Bug Fixes
--------
* Let the `beep_after_seconds` slow query alert happen before paged output.
* Allow comma-containing favorite queries in `~/.myclirc`.


Internal
Expand Down
74 changes: 67 additions & 7 deletions mycli/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,24 +14,84 @@
logger = logging.getLogger(__name__)


class FavoriteQueryPreservingConfigObj(ConfigObj):
"""When reading, quote SQL text on the fly which ConfigObj would otherwise interpret as a list."""

# Buglet: the use of _get_triple_quote() does not allow values which
# contain _both_ possible triple-quote delimiters to be read.
def _parse(self, infile: list[str]) -> None:
if not self.list_values or self.unrepr:
super()._parse(infile)
return

lines = infile.copy()
in_favorites = False
index = 0
while index < len(lines):
line = lines[index]
if not line.strip() or line.lstrip().startswith('#'):
index += 1
continue
section = self._sectionmarker.match(line)
if section is not None:
_, opening, name, closing, _ = section.groups()
in_favorites = opening.count('[') == closing.count(']') == 1 and self._unquote(name) == 'favorite_queries'
else:
entry = self._keyword.match(line)
if entry is not None:
indent, key, value = entry.groups()
if value.startswith(('"""', "'''")):
# Skip complete multiline values, including section-like SQL.
try:
_, _, index = self._multiline(value, lines, index, len(lines) - 1)
except SyntaxError:
break
elif in_favorites and not value.startswith(('"', "'")):
match = self._nolistvalue.match(value)
if match is not None:
sql, comment = match.groups()
lines[index] = f'{indent}{key} = {self._get_triple_quote(sql) % sql} {comment or ""}'
index += 1

super()._parse(lines)


class TripleQuotedConfigValue(str):
"""A value that must retain triple quoting when written."""

quote: str | None

def __new__(cls, value: str, quote: str | None = None) -> 'TripleQuotedConfigValue':
instance = super().__new__(cls, value)
instance.quote = quote
return instance


class LimiitedQuotePreservingConfigObj(ConfigObj):
"""Useful for saving individual items without modifying the whole file.

Triplequotes must be manually added for multiline values, and despite the
name of the class, could change from double to single in style. If we
don't do this, multiline triplequoted strings lose their quotes entirely,
resulting in unreadable files.
Preserve existing quoting and add triple quotes for new multiline values.
"""

# Buglet: the use of _get_triple_quote() does not allow values which
# contain _both_ possible triple-quote delimiters to be saved.
def __init__(self, *args, **kwargs):
ConfigObj.__init__(self, *args, **kwargs)

def _unquote(self, value):
return value

def _quote(self, value, multiline=True):
def _multiline(self, value: str, infile: list[str], cur_index: int, maxline: int) -> tuple[str, str | None, int]:
parsed, comment, end_index = super()._multiline(value, infile, cur_index, maxline)
return TripleQuotedConfigValue(parsed, value[:3]), comment, end_index

def _quote(self, value: str, multiline: bool = True) -> str:
if isinstance(value, TripleQuotedConfigValue):
if value.quote is not None:
return f'{value.quote}{value}{value.quote}'
return self._get_triple_quote(value) % value
if '\n' in value:
return f"'''{value}'''"
return self._get_triple_quote(value) % value
return value


Expand Down Expand Up @@ -70,7 +130,7 @@ def read_config_file(
if preserve_quotes:
config = LimiitedQuotePreservingConfigObj(f, interpolation=False, encoding="utf8", list_values=False)
else:
config = ConfigObj(f, interpolation=False, encoding="utf8", list_values=list_values)
config = FavoriteQueryPreservingConfigObj(f, interpolation=False, encoding="utf8", list_values=list_values)
except ConfigObjError as e:
if raise_errors:
raise
Expand Down
5 changes: 3 additions & 2 deletions mycli/packages/special_commands/favorite_queries.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from jinja2 import meta, nodes
from jinja2.sandbox import SandboxedEnvironment

from mycli.config import log, read_config_file, read_config_files
from mycli.config import TripleQuotedConfigValue, log, read_config_file, read_config_files

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -283,6 +283,7 @@ def reload(self) -> None:
def _clean_query(self, query: str | None) -> str | None:
if not query:
return query

query = query.lstrip(' \t\n\r')
query = query.rstrip(' \t\n\r')
query = query.removesuffix(';')
Expand Down Expand Up @@ -317,7 +318,7 @@ def save(self, name: str, query: str) -> None:
config.encoding = "utf-8"
section_existed = self.section_name in config
previous_query = config.get(self.section_name, {}).get(name, MISSING)
self._set_query(config, name, query)
self._set_query(config, name, TripleQuotedConfigValue(query))
try:
config.write()
except Exception:
Expand Down
55 changes: 53 additions & 2 deletions mycli_test/pytests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from mycli import config as config_module
from mycli.config import (
LimiitedQuotePreservingConfigObj,
TripleQuotedConfigValue,
_remove_pad,
create_default_config,
get_mylogin_cnf_path,
Expand Down Expand Up @@ -181,6 +182,33 @@ def test_quote_preserving_config_retains_quotes_and_quotes_multiline_values() ->
assert config._quote('first line\nsecond line') == "'''first line\nsecond line'''"


@pytest.mark.parametrize('quote', ['"""', "'''"])
@pytest.mark.parametrize('value', ["SELECT 1, '#tag', 2", "SELECT 1,\n'#tag'"])
def test_quote_preserving_config_retains_triple_quoted_values(quote: str, value: str) -> None:
text = f'[favorite_queries]\nq = {quote}{value}{quote}\n'
config = read_config_file(StringIO(text), preserve_quotes=True)
assert isinstance(config, LimiitedQuotePreservingConfigObj)
output = BytesIO()
config.write(output)
assert output.getvalue().decode('utf-8') == text


def test_quote_preserving_config_quotes_new_marked_value() -> None:
config = read_config_file(StringIO('[favorite_queries]\n'), preserve_quotes=True)
assert config is not None
query = "SELECT 1, '#tag', 2"
config['favorite_queries']['q'] = TripleQuotedConfigValue(query)
output = BytesIO()

config.write(output)

text = output.getvalue().decode('utf-8')
assert text == f"[favorite_queries]\nq = '''{query}'''\n"
reloaded = read_config_file(StringIO(text), raise_errors=True)
assert reloaded is not None
assert reloaded['favorite_queries']['q'] == query


def test_read_config_files_merges_files_in_order(monkeypatch: pytest.MonkeyPatch) -> None:
defaults = config_module.ConfigObj({'main': {'default': 'yes', 'color': 'default'}})
first = config_module.ConfigObj({'main': {'color': 'blue'}})
Expand Down Expand Up @@ -266,7 +294,7 @@ def test_read_config_file_permission_error(monkeypatch, caplog) -> None:
def raise_oserror(*_args, **_kwargs):
raise OSError(13, 'denied', '/tmp/test.cnf')

monkeypatch.setattr(config_module, 'ConfigObj', raise_oserror)
monkeypatch.setattr(config_module, 'FavoriteQueryPreservingConfigObj', raise_oserror)

with caplog.at_level(logging.WARNING, logger='mycli.config'):
assert read_config_file('/tmp/test.cnf') is None
Expand All @@ -281,13 +309,36 @@ def test_read_config_file_can_raise_parse_errors(tmp_path) -> None:
read_config_file(str(invalid_path), raise_errors=True)


@pytest.mark.parametrize('value', ['"""unterminated', "'''closed''' trailing"])
def test_read_config_file_raises_for_invalid_multiline_value(value: str) -> None:
contents = f'[favorite_queries]\nvalid = SELECT 1\nbroken = {value}\n'

with pytest.raises(ConfigObjError) as exc_info:
read_config_file(StringIO(contents), raise_errors=True)

assert exc_info.value.line_number == 3
assert exc_info.value.config['favorite_queries']['valid'] == 'SELECT 1'


def test_read_config_file_recovers_from_invalid_multiline_value(caplog: pytest.LogCaptureFixture) -> None:
contents = '[favorite_queries]\nvalid = SELECT 1\nbroken = """unterminated\n'

with caplog.at_level(logging.WARNING, logger='mycli.config'):
config = read_config_file(StringIO(contents))

assert config is not None
assert config['favorite_queries'] == {'valid': 'SELECT 1'}
assert 'Unable to parse line 3' in caplog.text
assert 'Using successfully parsed config values.' in caplog.text


def test_read_config_file_can_raise_io_errors(monkeypatch) -> None:
error = OSError(13, 'denied', '/tmp/test.cnf')

def raise_oserror(*_args, **_kwargs):
raise error

monkeypatch.setattr(config_module, 'ConfigObj', raise_oserror)
monkeypatch.setattr(config_module, 'FavoriteQueryPreservingConfigObj', raise_oserror)

with pytest.raises(OSError) as exc_info:
read_config_file('/tmp/test.cnf', raise_errors=True)
Expand Down
Loading
Loading