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
Expand Up @@ -16,6 +16,11 @@ Documentation
* Unset `PYTHONPATH` in `AGENTS.md` command suggestions.


Internal
--------
* Major dependency update to `PyMySQL` v1.2.3, changing SSL and ping defaults.


v2.25.3 (2026/09/19)
==============

Expand Down
32 changes: 7 additions & 25 deletions mycli/client_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,38 +520,20 @@ def reconnect(self, database: str = "") -> bool:
assert self.sqlexecute is not None
assert self.sqlexecute.conn is not None

# First pass with ping(reconnect=False) and minimal feedback levels. This definitely
# works as expected, and is a good idea especially when "connect" was used as a
# synonym for "use".
# First pass with ping() and minimal feedback levels. This definitely works as
# expected, and is a good idea especially when "connect" was used as a synonym
# for "use". Note that the default behavior of ping() changed in PyMySQL 1.2.x:
# it no longer reconnects.
try:
self.sqlexecute.conn.ping(reconnect=False)
self.sqlexecute.conn.ping()
if not database:
self.echo("Already connected.", fg="yellow")
return True
except pymysql.err.Error:
pass

# Second pass with ping(reconnect=True). It is not demonstrated that this pass ever
# gives the benefit it is looking for, _ie_ preserves session state. We need to test
# this with connection pooling.
try:
old_connection_id = self.sqlexecute.connection_id
self.logger.debug("Attempting to reconnect.")
self.echo("Reconnecting...", fg="yellow")
self.sqlexecute.conn.ping(reconnect=True)
# if a database is currently selected, set it on the conn again
if self.sqlexecute.dbname:
self.sqlexecute.conn.select_db(self.sqlexecute.dbname)
self.logger.debug("Reconnected successfully.")
self.echo("Reconnected successfully.", fg="yellow")
self.sqlexecute.reset_connection_id()
if old_connection_id != self.sqlexecute.connection_id:
self.echo("Any session state was reset.", fg="red")
return True
except pymysql.err.Error:
pass

# Third pass with sqlexecute.connect() should always work, but always resets session state.
# Second pass with sqlexecute.connect() should always work, and always resets
# session state.
try:
self.logger.debug("Creating new connection")
self.echo("Creating new connection...", fg="yellow")
Expand Down
4 changes: 2 additions & 2 deletions mycli/main_modes/repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -347,7 +347,7 @@ def render_prompt_string(
if r'\b' in checker_string:
connection = getattr(sqlexecute, 'conn', None)
if connection:
connection.ping(reconnect=False)
connection.ping()
server_status = getattr(connection, 'server_status', 0) or 0
transaction_indicator = '[TX]' if server_status & SERVER_STATUS_IN_TRANS else ''
strings = [x.replace(r'\b', transaction_indicator) for x in strings]
Expand Down Expand Up @@ -709,7 +709,7 @@ def _keepalive_hook(
try:
assert mycli.sqlexecute is not None
assert mycli.sqlexecute.conn is not None
mycli.sqlexecute.conn.ping(reconnect=False)
mycli.sqlexecute.conn.ping()
except Exception as e:
mycli.logger.debug('keepalive ping error %r', e)

Expand Down
2 changes: 1 addition & 1 deletion mycli/packages/special/dbcommands.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ def ping(cur: Cursor, arg: str | None = None, **_) -> list[SQLResult]:
return [SQLResult(status='Syntax: /ping.')]

try:
cur.connection.ping(reconnect=False)
cur.connection.ping()
except Error:
return [SQLResult(status='Not connected')]
return [SQLResult(status='Connected')]
Expand Down
43 changes: 26 additions & 17 deletions mycli/sqlexecute.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,14 @@
import enum
import logging
import re
import ssl
import ssl as ssllib
from typing import Any, Generator, Iterable

from prompt_toolkit.formatted_text import FormattedText
import pymysql
from pymysql.connections import Connection
from pymysql.constants import FIELD_TYPE
from pymysql.constants.CR import CR_SSL_CONNECTION_ERROR
from pymysql.converters import conversions, convert_date, convert_datetime, convert_time, decoders
from pymysql.cursors import Cursor

Expand Down Expand Up @@ -257,9 +258,10 @@ def connect(
client_flag |= pymysql.constants.CLIENT.MULTI_STATEMENTS
client_flag |= pymysql.constants.CLIENT.HANDLE_EXPIRED_PASSWORDS

ssl_context = None
if ssl:
ssl_context = self._create_ssl_ctx(ssl)
ssl_kwargs: dict[str, Any] = {'ssl': self._create_ssl_ctx(ssl)}
else:
ssl_kwargs = {'ssl_disabled': True}

connect_kwargs: dict[str, Any] = {
"database": db,
Expand All @@ -274,16 +276,23 @@ def connect(
"client_flag": client_flag,
"local_infile": local_infile or False,
"conv": conv,
"ssl": ssl_context, # type: ignore[arg-type]
"program_name": "mycli",
"defer_connect": defer_connect,
"init_command": init_command or None,
"cursorclass": pymysql.cursors.SSCursor if unbuffered else pymysql.cursors.Cursor,
**ssl_kwargs,
}

self.sandbox_mode = False
try:
conn = pymysql.connect(**connect_kwargs) # type: ignore[misc]
try:
conn = pymysql.connect(**connect_kwargs) # type: ignore[misc]
except pymysql.OperationalError as e:
if e.args[0] != CR_SSL_CONNECTION_ERROR or not ssl or ssl.get('mode') != 'auto':
raise
del connect_kwargs['ssl']
connect_kwargs['ssl_disabled'] = True
conn = pymysql.connect(**connect_kwargs) # type: ignore[misc]
except pymysql.OperationalError as e:
if e.args[0] == ER_MUST_CHANGE_PASSWORD:
# Post-handshake queries (SET NAMES, SET AUTOCOMMIT, init_command)
Expand Down Expand Up @@ -624,35 +633,35 @@ def _connect_sandbox(conn: Connection) -> None:
finally:
conn.set_character_set = original_set_charset # type: ignore[assignment]

def _create_ssl_ctx(self, sslp: dict) -> ssl.SSLContext:
def _create_ssl_ctx(self, sslp: dict) -> ssllib.SSLContext:
ca = sslp.get("ca")
capath = sslp.get("capath")
hasnoca = ca is None and capath is None
ctx = ssl.create_default_context(cafile=ca, capath=capath)
ctx = ssllib.create_default_context(cafile=ca, capath=capath)
ctx.check_hostname = not hasnoca and sslp.get("check_hostname", True)
ctx.verify_mode = ssl.CERT_NONE if hasnoca else ssl.CERT_REQUIRED
ctx.verify_mode = ssllib.CERT_NONE if hasnoca else ssllib.CERT_REQUIRED
if "cert" in sslp:
ctx.load_cert_chain(sslp["cert"], keyfile=sslp.get("key"))
if "cipher" in sslp:
ctx.set_ciphers(sslp["cipher"])

ctx.minimum_version = ssl.TLSVersion.TLSv1_2
ctx.minimum_version = ssllib.TLSVersion.TLSv1_2

if "tls_version" in sslp:
tls_version = sslp["tls_version"]

if tls_version == "TLSv1":
ctx.minimum_version = ssl.TLSVersion.TLSv1
ctx.maximum_version = ssl.TLSVersion.TLSv1
ctx.minimum_version = ssllib.TLSVersion.TLSv1
ctx.maximum_version = ssllib.TLSVersion.TLSv1
elif tls_version == "TLSv1.1":
ctx.minimum_version = ssl.TLSVersion.TLSv1_1
ctx.maximum_version = ssl.TLSVersion.TLSv1_1
ctx.minimum_version = ssllib.TLSVersion.TLSv1_1
ctx.maximum_version = ssllib.TLSVersion.TLSv1_1
elif tls_version == "TLSv1.2":
ctx.minimum_version = ssl.TLSVersion.TLSv1_2
ctx.maximum_version = ssl.TLSVersion.TLSv1_2
ctx.minimum_version = ssllib.TLSVersion.TLSv1_2
ctx.maximum_version = ssllib.TLSVersion.TLSv1_2
elif tls_version == "TLSv1.3":
ctx.minimum_version = ssl.TLSVersion.TLSv1_3
ctx.maximum_version = ssl.TLSVersion.TLSv1_3
ctx.minimum_version = ssllib.TLSVersion.TLSv1_3
ctx.maximum_version = ssllib.TLSVersion.TLSv1_3
else:
_logger.error("Invalid tls version: %s", tls_version)

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ dependencies = [
"cryptography ~= 50.0.1",
"Pygments ~= 2.21.0",
"prompt_toolkit >= 3.0.41,<4.0.0",
"PyMySQL ~= 1.1.2",
"PyMySQL ~= 1.2.3",
"sqlparse ~= 0.6.0",
"sqlglot ~= 30.17.0",
"sqlglotc ~= 30.17.0",
Expand Down
10 changes: 0 additions & 10 deletions test/pytests/test_client_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -1471,16 +1471,6 @@ def test_reconnect_returns_true_when_ping_succeeds() -> None:
assert client.echo_calls == [(('Already connected.',), {'fg': 'yellow'})]


def test_reconnect_uses_ping_reconnect_and_selects_current_database() -> None:
client = DummyClient()
conn = FakeConn([pymysql.err.Error('stale'), None])
client.sqlexecute = FakeReconnectSQLExecute(conn, connection_id=10, dbname='selected')
client.sqlexecute.next_connection_id = 10

assert client.reconnect(database='newdb') is True
assert conn.select_db_calls == ['selected']


def test_reconnect_reports_session_reset_when_connection_id_changes() -> None:
client = DummyClient()
conn = FakeConn([pymysql.err.Error('stale'), None])
Expand Down
12 changes: 3 additions & 9 deletions test/pytests/test_main_modes_repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,10 +94,8 @@ class FakeConnection:
def __init__(self, ping_exc: Exception | None = None, cursor_value: Any = 'cursor') -> None:
self.ping_exc = ping_exc
self.cursor_value = cursor_value
self.ping_calls: list[bool] = []

def ping(self, reconnect: bool = False) -> None:
self.ping_calls.append(reconnect)
def ping(self) -> None:
if self.ping_exc is not None:
raise self.ping_exc

Expand Down Expand Up @@ -703,8 +701,7 @@ def make_transaction_prompt_cli(connection: Any) -> Any:
def test_transaction_prompt_reads_refreshed_flag(status: int | None, expected: str) -> None:
connection = SimpleNamespace(server_status=0 if expected else 1, cursor=pytest.fail)

def ping(*, reconnect: bool) -> None:
assert reconnect is False
def ping() -> None:
connection.server_status = status

connection.ping = ping
Expand Down Expand Up @@ -762,8 +759,6 @@ def test_transaction_prompt_pings_only_for_active_escape(format_string: str, exp
repl_mode.render_prompt_string(cli, format_string, 0)

assert ping.call_count == expected_calls
if expected_calls:
ping.assert_called_once_with(reconnect=False)


def test_transaction_prompt_reuses_cached_render_without_ping() -> None:
Expand All @@ -773,7 +768,7 @@ def test_transaction_prompt_reuses_cached_render_without_ping() -> None:
repl_mode.render_prompt_string(cli, r'\b', 0)
repl_mode.render_prompt_string(cli, r'\b', 0)

ping.assert_called_once_with(reconnect=False)
ping.assert_called_once_with()


def test_render_prompt_string_includes_current_edit_mode() -> None:
Expand Down Expand Up @@ -1330,7 +1325,6 @@ def test_keepalive_hook_covers_threshold_and_errors() -> None:
assert cli._keepalive_counter == 1
repl_mode._keepalive_hook(cli, None)
assert cli._keepalive_counter == 0
assert cli.sqlexecute.conn.ping_calls == [False]

cli.sqlexecute.conn = FakeConnection(ping_exc=RuntimeError('boom'))
repl_mode._keepalive_hook(cli, None)
Expand Down
5 changes: 0 additions & 5 deletions test/pytests/test_special_dbcommands.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,13 +30,11 @@ def __init__(
self.unix_socket = unix_socket
self._thread_id_value = thread_id_value
self.ping_error = ping_error
self.ping_calls: list[bool] = []

def thread_id(self) -> int:
return self._thread_id_value

def ping(self, reconnect: bool = True) -> None:
self.ping_calls.append(reconnect)
if self.ping_error is not None:
raise self.ping_error

Expand Down Expand Up @@ -193,15 +191,13 @@ def test_ping_reports_connected_without_reconnecting() -> None:
cursor = FakeCursor(query_results={}, connection=connection)

assert ping(cursor) == [SQLResult(status='Connected')]
assert connection.ping_calls == [False]


def test_ping_reports_not_connected_on_pymysql_error() -> None:
connection = FakeConnection(ping_error=Error('connection lost'))
cursor = FakeCursor(query_results={}, connection=connection)

assert ping(cursor) == [SQLResult(status='Not connected')]
assert connection.ping_calls == [False]


def test_ping_propagates_unrelated_errors() -> None:
Expand All @@ -217,7 +213,6 @@ def test_ping_rejects_arguments_without_contacting_server() -> None:
cursor = FakeCursor(query_results={}, connection=connection)

assert ping(cursor, arg='unexpected') == [SQLResult(status='Syntax: /ping.')]
assert connection.ping_calls == []


def test_ping_command_registration() -> None:
Expand Down
Loading
Loading