"""Test MySQL Adapter.""" import re from dataclasses import dataclass from typing import Any, Sequence import pytest from sqlalchemy import column from sqlalchemy.dialects import mysql as mysql_dialect from sqlalchemy.engine import URL import abacus_common_logic.db.adapters.mysql_adapter as mod from abacus_common_logic.db.adapters.mysql_adapter import ( MySQLAdapter, _build_url, _get_server_engine, _normalize_create_trigger, ) @dataclass class Cfg: """Config dataclass.""" MYSQL_DB_USER: str = 'u' MYSQL_DB_PASS: str = 'p' MYSQL_DB_HOST: str = 'h' MYSQL_DB_NAME: str = 'db' MYSQL_DB_PORT: int | str = 3306 # Fakes class _Preparer: def quote(self, ident: str) -> str: return f'`{ident}`' def quote_schema(self, ident: str) -> str: return f'`{ident}`' class _FakeDialect: def __init__(self): self.identifier_preparer = _Preparer() class _FakeResult: def __init__( self, rows: list[tuple] | None = None, scalar_value: Any | None = None ): self._rows = rows or [] self._scalar = scalar_value def first(self): return self._rows[0] if self._rows else None def scalar(self): return self._scalar # .scalars().all() class _Scalars: def __init__(self, seq: Sequence[Any]): self._seq = seq def all(self): # noqa: A003 return list(self._seq) def scalars(self): if self._rows: return _FakeResult._Scalars([r[0] for r in self._rows]) return _FakeResult._Scalars([]) class _FakeConn: def __init__(self, scripted: dict[str, Any] | None = None): self.dialect = _FakeDialect() self.sql = [] # record of exec_driver_sql strings self._scripted = scripted or {} # context-manager protocol (for connect/begin) def __enter__(self): return self def __exit__(self, *exc): return False def execution_options(self, **opts): return self def exec_driver_sql(self, sql: str, *args, **kwargs): self.sql.append(sql) if 'SELECT @@FOREIGN_KEY_CHECKS' in sql: return _FakeResult(scalar_value=self._scripted.get('fk_checks', 1)) if sql.startswith('SHOW CREATE TABLE'): # MySQL returns (table, create_stmt) create_stmt = self._scripted.get('show_create_table') return _FakeResult(rows=[('t', create_stmt)]) if sql.startswith('SHOW CREATE TRIGGER'): # MySQL returns (_, _, create_stmt) create_stmt = self._scripted.get('show_create_trigger') return _FakeResult(rows=[('name', 'timing', create_stmt)]) return _FakeResult() def execute(self, statement, params=None): txt = str(statement) if 'INFORMATION_SCHEMA.SCHEMATA' in txt: exists = self._scripted.get('db_exists', 0) return _FakeResult(scalar_value=exists) if 'INFORMATION_SCHEMA.TABLES' in txt and "TABLE_TYPE='BASE TABLE'" in txt: # Return rows as (TABLE_NAME,) names = self._scripted.get('tables', []) return _FakeResult(rows=[(n,) for n in names]) if 'INFORMATION_SCHEMA.TRIGGERS' in txt: names = self._scripted.get('triggers', []) return _FakeResult(rows=[(n,) for n in names]) return _FakeResult() class _FakeEngine: def __init__(self, database: str, scripted: dict[str, Any] | None = None): self.url = type('U', (), {'database': database}) self._scripted = scripted or {} self.disposed = False def connect(self): return _FakeConn(self._scripted) def begin(self): return _FakeConn(self._scripted) def dispose(self): self.disposed = True # Tests def test_build_url_components(): """Test build_url.""" cfg = Cfg() url = _build_url(cfg) # type: ignore assert isinstance(url, URL) assert url.drivername == 'mysql+pymysql' assert url.username == 'u' assert url.password == 'p' assert url.host == 'h' assert url.port == 3306 assert url.database == 'db' assert url.query['charset'] == 'utf8mb4' assert url.query['connect_timeout'] == '10' def test_get_server_engine_requires_db_name(monkeypatch): """Test _get_server_engine.""" cfg = Cfg(MYSQL_DB_NAME='') with pytest.raises(ValueError, match='MYSQL_DB_NAME'): _get_server_engine(cfg) # type: ignore def test_normalize_create_trigger_strips_definer_and_db_qualifier(): """Test _normalize_create_trigger.""" sql_in = ( 'CREATE DEFINER=`root`@`%` TRIGGER `db`.`my_trg` BEFORE INSERT ON `db`.`t`\n' 'FOR EACH ROW SET NEW.c = NOW()' ) out = _normalize_create_trigger('db', sql_in) assert out.startswith('CREATE TRIGGER IF NOT EXISTS') assert 'DEFINER' not in out assert '`db`.' not in out assert '`t`' in out def test_json_contains_compiles_mysql(): """Test json_contains.""" expr = MySQLAdapter().json_contains(column('c'), 'v', path='$.x') compiled = expr.compile(dialect=mysql_dialect.dialect()) # SQL shape sql_text = str(compiled).upper() assert 'JSON_CONTAINS' in sql_text assert 'JSON_EXTRACT' in sql_text # Bound parameters (names are dialect-dependent, so just check values) params = compiled.params.values() assert '$.x' in params assert 'v' in params def test_set_fk_executes_expected_sql(): """Test set_fk.""" conn = _FakeConn() a = MySQLAdapter() a.set_fk(conn, True) # type: ignore a.set_fk(conn, False) # type: ignore assert conn.sql[0].startswith('SET FOREIGN_KEY_CHECKS = 1') assert conn.sql[1].startswith('SET FOREIGN_KEY_CHECKS = 0') def test_truncate_tables_empty_is_noop(): """Test truncate zero tables.""" conn = _FakeConn() MySQLAdapter().truncate_tables(conn, []) # type: ignore assert conn.sql == [] # Confirm nothing ran def test_truncate_tables_toggles_fk_and_truncates_in_order(): """Test truncate tables.""" # Start with foreign keys on conn = _FakeConn(scripted={'fk_checks': 1}) # Truncate two tables a = MySQLAdapter() a.truncate_tables(conn, ['alpha', 'beta']) # type: ignore joined = ' ; '.join(conn.sql) assert re.search(r'SELECT @@FOREIGN_KEY_CHECKS', joined) # Foreign keys were disabled assert 'SET FOREIGN_KEY_CHECKS = 0' in joined # Tables were truncated assert 'TRUNCATE TABLE `alpha`' in joined assert 'TRUNCATE TABLE `beta`' in joined # Foreign keys were enabled assert 'SET FOREIGN_KEY_CHECKS = 1' in joined def test_create_db_executes_create_database_and_disposes(monkeypatch): """Test create_db creates DB.""" fake = _FakeEngine('server', scripted={}) def fake_get_server_engine(cfg): return (fake, 'newdb') monkeypatch.setattr(mod, '_get_server_engine', fake_get_server_engine) a = MySQLAdapter() a.create_db(Cfg()) # type: ignore # Make sure the engine was disposed assert fake.disposed is True def test_db_exists_true_false(monkeypatch): """Test db_exists.""" fake_true = _FakeEngine('server', scripted={'db_exists': 1}) fake_false = _FakeEngine('server', scripted={'db_exists': 0}) monkeypatch.setattr(mod, '_get_server_engine', lambda _: (fake_true, 'db')) assert MySQLAdapter().db_exists(Cfg()) is True # type: ignore assert fake_true.disposed is True monkeypatch.setattr(mod, '_get_server_engine', lambda _: (fake_false, 'db')) assert MySQLAdapter().db_exists(Cfg()) is False # type: ignore assert fake_false.disposed is True def test_create_engine_calls_sa_create_engine_with_url_and_pool_params(monkeypatch): """Test create_engine uses SQLAlchemy.""" captured = {} def fake_create_engine(url, **kw): captured['url'] = url captured['kw'] = kw return object() monkeypatch.setattr(mod, 'sa_create_engine', fake_create_engine) cfg = Cfg() eng = MySQLAdapter().create_engine(cfg) # type: ignore assert eng is not None assert isinstance(captured['url'], URL) assert captured['url'].drivername == 'mysql+pymysql' assert captured['kw']['pool_pre_ping'] is True assert captured['kw']['pool_recycle'] == 1800 # Clone helpers def test_clone_db_schema_skips_when_target_has_tables(monkeypatch): """Test _clone_db_schema skips when DB is not empty.""" dst = _FakeEngine('dst', scripted={'db_exists': 1}) src = _FakeEngine('src', scripted={}) # Patch connection.execute for dst to return a row def fake_connect_with_row(engine): class _C(_FakeConn): def execute(self, statement, params=None): txt = str(statement) if 'INFORMATION_SCHEMA.TABLES' in txt and 'LIMIT 1' in txt: return _FakeResult(rows=[(1,)]) return super().execute(statement, params) return _C(engine._scripted) monkeypatch.setattr(dst, 'connect', lambda: fake_connect_with_row(dst)) a = MySQLAdapter() a._clone_db_schema(src, dst) # type: ignore # No queries should have been executed assert dst.connect().sql == [] def test_clone_db_schema_creates_tables_when_empty(monkeypatch): """Test _clone_db_schema creates tables when DB is empty.""" src = _FakeEngine( 'sdb', scripted={ 'tables': ['t1'], 'show_create_table': 'CREATE TABLE `sdb`.`t1` (id int) ENGINE=InnoDB', }, ) dst = _FakeEngine('ddb', scripted={}) class _DstConn(_FakeConn): def execute(self, statement, params=None): txt = str(statement) if 'INFORMATION_SCHEMA.TABLES' in txt and 'LIMIT 1' in txt: return _FakeResult(rows=[]) return super().execute(statement, params) dst_conn = _DstConn(dst._scripted) monkeypatch.setattr(dst, 'begin', lambda: dst_conn) # Clone the tables a = MySQLAdapter() a._clone_db_schema(src, dst) # type: ignore s = ' ; '.join(dst_conn.sql) # Foreign keys were disabled assert 'SET FOREIGN_KEY_CHECKS = 0' in s # Tables were created assert 'CREATE TABLE IF NOT EXISTS `t1`' in s # Foreign keys were enabled assert 'SET FOREIGN_KEY_CHECKS = 1' in s def test_clone_db_triggers_no_triggers_short_circuits(): """Test _clone_db_triggers with zero triggers.""" src = _FakeEngine('sdb', scripted={'triggers': []}) dst = _FakeEngine('ddb', scripted={}) MySQLAdapter()._clone_db_triggers(src, dst) # type: ignore assert dst.begin().sql == [] def test_clone_db_triggers_creates_triggers_and_optionally_drops(monkeypatch): """Test _clone_db_triggers creates triggers.""" src = _FakeEngine( 'sdb', scripted={ 'triggers': ['trg1'], 'show_create_trigger': ( 'CREATE DEFINER=`r`@`%` TRIGGER `sdb`.`trg1` ' 'BEFORE INSERT ON `sdb`.`t` FOR EACH ROW SET NEW.c=1' ), }, ) dst = _FakeEngine('ddb', scripted={}) dst_tx_conn = _FakeConn({}) monkeypatch.setattr(dst, 'begin', lambda: dst_tx_conn) # Clone the triggers a = MySQLAdapter() a._clone_db_triggers(src, dst, drop_existing=True) # type: ignore s = ' ; '.join(dst_tx_conn.sql) # Drops the trigger if exists assert 'DROP TRIGGER IF EXISTS `trg1`' in s # Creates the trigger assert 'CREATE TRIGGER IF NOT EXISTS' in s