From 437c439bab516401f9ea13c03645c716021d2106 Mon Sep 17 00:00:00 2001 From: Roland Walker Date: Sat, 26 Sep 2026 09:29:42 -0400 Subject: [PATCH] major dependency upgrade to PyMySQL v1.2.3 PyMySQL v1.2.0 was a breaking update which was avoided for some time. The defaults for ping() have changed such that reconnect is False, and reconnect=True is also deprecated with a user-visible warning. The defaults for SSL support on connect() have also changed: SSL mode is the default, and ssl_disabled must be set if SSL is not wanted. Changes * remove all reconnect parameters to ping() * remove reconnect=True pass for mycli's reconnect(), which may never have been useful anyway * add support for the ssl_disabled kwarg to PyMySQL's connect() * add a retry for PyMySQL's connect() on SSL failure with ssl_mode of "auto" * import "ssl" as "ssllib" in sqlexecute.py, since there is also a variable with the name "ssl" Each value of --ssl-mode on the client has been tested using a MySQL server both with and without SSL support. For the second case, a MySQL 8 server was run with --tls-version=''. On the client, the ability to connect under each combination was checked, along with consulting the output of /status. The behavior of the /connect command in various situations is more difficult to verify, but since we only removed the second of three passes to the reconnect behavior, we can be reasonably confident that the overall behavior is still good. --- changelog.md | 5 + mycli/client_connection.py | 32 ++----- mycli/main_modes/repl.py | 4 +- mycli/packages/special/dbcommands.py | 2 +- mycli/sqlexecute.py | 43 +++++---- pyproject.toml | 2 +- test/pytests/test_client_connection.py | 10 -- test/pytests/test_main_modes_repl.py | 12 +-- test/pytests/test_special_dbcommands.py | 5 - test/pytests/test_sqlexecute.py | 119 +++++++++++++++++++++--- 10 files changed, 151 insertions(+), 83 deletions(-) diff --git a/changelog.md b/changelog.md index b5c3b4b69..5bfd6789a 100644 --- a/changelog.md +++ b/changelog.md @@ -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) ============== diff --git a/mycli/client_connection.py b/mycli/client_connection.py index df9c5d5ba..d00f8e4dd 100644 --- a/mycli/client_connection.py +++ b/mycli/client_connection.py @@ -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") diff --git a/mycli/main_modes/repl.py b/mycli/main_modes/repl.py index 5aa94c992..cc6d404a4 100644 --- a/mycli/main_modes/repl.py +++ b/mycli/main_modes/repl.py @@ -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] @@ -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) diff --git a/mycli/packages/special/dbcommands.py b/mycli/packages/special/dbcommands.py index e61f83827..3c055f939 100644 --- a/mycli/packages/special/dbcommands.py +++ b/mycli/packages/special/dbcommands.py @@ -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')] diff --git a/mycli/sqlexecute.py b/mycli/sqlexecute.py index 847bb79a1..5a5468074 100644 --- a/mycli/sqlexecute.py +++ b/mycli/sqlexecute.py @@ -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 @@ -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, @@ -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) @@ -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) diff --git a/pyproject.toml b/pyproject.toml index c1efce1c3..541e8b0af 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/test/pytests/test_client_connection.py b/test/pytests/test_client_connection.py index a1483fef9..8811286e6 100644 --- a/test/pytests/test_client_connection.py +++ b/test/pytests/test_client_connection.py @@ -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]) diff --git a/test/pytests/test_main_modes_repl.py b/test/pytests/test_main_modes_repl.py index 841bb13ed..0a1ab07b2 100644 --- a/test/pytests/test_main_modes_repl.py +++ b/test/pytests/test_main_modes_repl.py @@ -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 @@ -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 @@ -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: @@ -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: @@ -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) diff --git a/test/pytests/test_special_dbcommands.py b/test/pytests/test_special_dbcommands.py index 8eb981521..1f348a5c6 100644 --- a/test/pytests/test_special_dbcommands.py +++ b/test/pytests/test_special_dbcommands.py @@ -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 @@ -193,7 +191,6 @@ 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: @@ -201,7 +198,6 @@ def test_ping_reports_not_connected_on_pymysql_error() -> None: 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: @@ -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: diff --git a/test/pytests/test_sqlexecute.py b/test/pytests/test_sqlexecute.py index 4436193d5..d13881aa7 100644 --- a/test/pytests/test_sqlexecute.py +++ b/test/pytests/test_sqlexecute.py @@ -3,6 +3,7 @@ from datetime import time import os from types import SimpleNamespace +from unittest.mock import Mock from prompt_toolkit.formatted_text import FormattedText import pymysql @@ -804,6 +805,98 @@ def fake_connect_sandbox(self, conn): assert executor.connection_id is None +def test_connect_enters_sandbox_after_ssl_fallback(monkeypatch: pytest.MonkeyPatch) -> None: + executor = make_executor_for_connect_tests() + executor.ssl = {'mode': 'auto'} + new_conn = DummyConnection(server_version='8.0.36') + connect = Mock( + side_effect=[ + pymysql.OperationalError(sqlexecute.CR_SSL_CONNECTION_ERROR, 'SSL unsupported'), + pymysql.OperationalError(sqlexecute.ER_MUST_CHANGE_PASSWORD, 'must change password'), + new_conn, + ] + ) + sandbox = Mock() + monkeypatch.setattr(sqlexecute.pymysql, 'connect', connect) + monkeypatch.setattr(executor, '_connect_sandbox', sandbox) + + executor.connect() + + assert connect.call_count == 3 + initial, plaintext, raw_handshake = [call.kwargs for call in connect.call_args_list] + assert 'ssl' in initial + for kwargs in (plaintext, raw_handshake): + assert 'ssl' not in kwargs + assert kwargs['ssl_disabled'] is True + assert plaintext['defer_connect'] is False + assert raw_handshake['defer_connect'] is True + assert raw_handshake['autocommit'] is None + assert raw_handshake['init_command'] is None + sandbox.assert_called_once_with(new_conn) + assert executor.conn is new_conn + assert executor.sandbox_mode is True + assert executor.server_info is None + assert executor.connection_id is None + + +def test_connect_ssl_fallback_succeeds_without_sandbox(monkeypatch: pytest.MonkeyPatch) -> None: + executor = make_executor_for_connect_tests() + executor.ssl = {'mode': 'auto'} + new_conn = DummyConnection(server_version='8.0.36') + connect = Mock( + side_effect=[ + pymysql.OperationalError(sqlexecute.CR_SSL_CONNECTION_ERROR, 'SSL unsupported'), + new_conn, + ] + ) + monkeypatch.setattr(sqlexecute.pymysql, 'connect', connect) + monkeypatch.setattr(executor, 'reset_connection_id', lambda: None) + monkeypatch.setattr(executor, '_probe_doris_version', lambda: None) + + executor.connect() + + assert connect.call_count == 2 + assert connect.call_args.kwargs['ssl_disabled'] is True + assert 'ssl' not in connect.call_args.kwargs + assert executor.conn is new_conn + assert executor.sandbox_mode is False + + +@pytest.mark.parametrize('error_code', [1045, sqlexecute.CR_SSL_CONNECTION_ERROR]) +def test_connect_ssl_fallback_propagates_other_errors(monkeypatch: pytest.MonkeyPatch, error_code: int) -> None: + executor = make_executor_for_connect_tests() + executor.ssl = {'mode': 'auto'} + retry_error = pymysql.OperationalError(error_code, 'retry failed') + connect = Mock( + side_effect=[ + pymysql.OperationalError(sqlexecute.CR_SSL_CONNECTION_ERROR, 'SSL unsupported'), + retry_error, + ] + ) + monkeypatch.setattr(sqlexecute.pymysql, 'connect', connect) + + with pytest.raises(pymysql.OperationalError) as exc_info: + executor.connect() + + assert exc_info.value is retry_error + assert connect.call_count == 2 + assert executor.conn is None + + +def test_connect_required_ssl_does_not_fall_back(monkeypatch: pytest.MonkeyPatch) -> None: + executor = make_executor_for_connect_tests() + executor.ssl = {'mode': 'on'} + error = pymysql.OperationalError(sqlexecute.CR_SSL_CONNECTION_ERROR, 'SSL unsupported') + connect = Mock(side_effect=error) + monkeypatch.setattr(sqlexecute.pymysql, 'connect', connect) + + with pytest.raises(pymysql.OperationalError) as exc_info: + executor.connect() + + assert exc_info.value is error + assert connect.call_count == 1 + + def test_connect_reraises_non_sandbox_operational_error(monkeypatch) -> None: executor = make_executor_for_connect_tests() executor.ssl = None @@ -1605,15 +1698,15 @@ def fake_create_default_context(cafile: str | None = None, capath: str | None = create_default_context_calls.append((cafile, capath)) return ctx - monkeypatch.setattr(sqlexecute.ssl, 'create_default_context', fake_create_default_context) + monkeypatch.setattr(sqlexecute.ssllib, 'create_default_context', fake_create_default_context) result = executor._create_ssl_ctx({}) assert result is ctx assert create_default_context_calls == [(None, None)] assert ctx.check_hostname is False - assert ctx.verify_mode == sqlexecute.ssl.CERT_NONE - assert ctx.minimum_version == sqlexecute.ssl.TLSVersion.TLSv1_2 + assert ctx.verify_mode == sqlexecute.ssllib.CERT_NONE + assert ctx.minimum_version == sqlexecute.ssllib.TLSVersion.TLSv1_2 assert ctx.maximum_version is None assert ctx.loaded_cert_chain is None assert ctx.cipher_string is None @@ -1629,7 +1722,7 @@ def fake_create_default_context(cafile: str | None = None, capath: str | None = return ctx monkeypatch.setattr( - sqlexecute.ssl, + sqlexecute.ssllib, 'create_default_context', fake_create_default_context, ) @@ -1646,26 +1739,26 @@ def fake_create_default_context(cafile: str | None = None, capath: str | None = assert result is ctx assert create_default_context_calls == [('/tmp/ca.pem', None)] assert ctx.check_hostname is False - assert ctx.verify_mode == sqlexecute.ssl.CERT_REQUIRED + assert ctx.verify_mode == sqlexecute.ssllib.CERT_REQUIRED assert ctx.loaded_cert_chain == ('/tmp/client-cert.pem', '/tmp/client-key.pem') assert ctx.cipher_string == 'ECDHE-RSA-AES256-GCM-SHA384' - assert ctx.minimum_version == sqlexecute.ssl.TLSVersion.TLSv1_3 - assert ctx.maximum_version == sqlexecute.ssl.TLSVersion.TLSv1_3 + assert ctx.minimum_version == sqlexecute.ssllib.TLSVersion.TLSv1_3 + assert ctx.maximum_version == sqlexecute.ssllib.TLSVersion.TLSv1_3 @pytest.mark.parametrize( ('tls_version', 'expected_version'), ( - ('TLSv1', sqlexecute.ssl.TLSVersion.TLSv1), - ('TLSv1.1', sqlexecute.ssl.TLSVersion.TLSv1_1), - ('TLSv1.2', sqlexecute.ssl.TLSVersion.TLSv1_2), + ('TLSv1', sqlexecute.ssllib.TLSVersion.TLSv1), + ('TLSv1.1', sqlexecute.ssllib.TLSVersion.TLSv1_1), + ('TLSv1.2', sqlexecute.ssllib.TLSVersion.TLSv1_2), ), ) def test_create_ssl_ctx_supports_legacy_tls_version_overrides(monkeypatch, tls_version: str, expected_version) -> None: executor = make_executor_for_run_tests() ctx = FakeSSLContext() - monkeypatch.setattr(sqlexecute.ssl, 'create_default_context', lambda **_kwargs: ctx) + monkeypatch.setattr(sqlexecute.ssllib, 'create_default_context', lambda **_kwargs: ctx) result = executor._create_ssl_ctx({'tls_version': tls_version}) @@ -1678,13 +1771,13 @@ def test_create_ssl_ctx_logs_invalid_tls_version_and_keeps_default_minimum(monke executor = make_executor_for_run_tests() ctx = FakeSSLContext() - monkeypatch.setattr(sqlexecute.ssl, 'create_default_context', lambda **_kwargs: ctx) + monkeypatch.setattr(sqlexecute.ssllib, 'create_default_context', lambda **_kwargs: ctx) with caplog.at_level('ERROR', logger='mycli.sqlexecute'): result = executor._create_ssl_ctx({'tls_version': 'SSLv3'}) assert result is ctx - assert ctx.minimum_version == sqlexecute.ssl.TLSVersion.TLSv1_2 + assert ctx.minimum_version == sqlexecute.ssllib.TLSVersion.TLSv1_2 assert ctx.maximum_version is None assert 'Invalid tls version: SSLv3' in caplog.text