"""DB adapter for MySQL.""" import logging import re from flask import Flask from flask_sqlalchemy import SQLAlchemy from sqlalchemy import ( Boolean, cast as sa_cast, create_engine as sa_create_engine, func, literal, text, ) from sqlalchemy.dialects.mysql import JSON from sqlalchemy.engine import URL, Connection, Engine from sqlalchemy.exc import OperationalError from sqlalchemy.sql.elements import ClauseElement from .adapter import Adapter from .utils import DialectName log = logging.getLogger(__name__) class MySQLConfig: """MySQL Config.""" MYSQL_DB_USER: str MYSQL_DB_PASS: str MYSQL_DB_HOST: str MYSQL_DB_NAME: str MYSQL_DB_PORT: int | str class MySQLAdapter(Adapter): """DB adapter for MySQL.""" dialect = DialectName('mysql') def clone_db(self, src_engine: Engine, dst_engine: Engine) -> None: """Clone the entire MySQL database from src_engine into dst_engine.""" self._clone_db_schema(src_engine, dst_engine) self._clone_db_data(src_engine, dst_engine) self._clone_db_triggers(src_engine, dst_engine, False) def create_db(self, config: MySQLConfig) -> None: """Create the DB.""" engine, dbname = _get_server_engine(config) try: with engine.begin() as conn: conn.execute( text( f'CREATE DATABASE IF NOT EXISTS `{dbname}` ' 'CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci' ) ) finally: engine.dispose() def create_engine(self, config: MySQLConfig) -> Engine: """Create MySQL DB engine.""" url = _build_url(config) return sa_create_engine(url, pool_pre_ping=True, pool_recycle=1800) def db_exists(self, config: MySQLConfig) -> bool: """Check if a DB exists.""" engine, dbname = _get_server_engine(config) try: with engine.connect() as conn: exists = conn.execute( text(""" SELECT 1 FROM INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME = :db LIMIT 1 """), {'db': dbname}, ).scalar() return bool(exists) finally: engine.dispose() def drop_db(self, conn: Connection, db: str) -> None: """Drop the DB.""" prep = conn.dialect.identifier_preparer db_quoted = prep.quote_schema(db) conn.execute(text(f'DROP DATABASE IF EXISTS {db_quoted}')) def json_contains(self, col, value, *, path: str = '$'): """Get the JSON_CONTAINS equivalent function.""" if value is None: candidate = sa_cast(literal('null'), JSON) elif isinstance(value, ClauseElement): candidate = value else: # Numbers stay numbers, booleans stay booleans, strings stay strings, etc. candidate = literal(value, type_=JSON) # Build the target JSON document at path target = func.json_extract(col, literal(path)) return func.json_contains(target, candidate).cast(Boolean) def set_fk(self, conn: Connection, enable: bool = True) -> None: """Control foreign key constraints.""" conn.exec_driver_sql('SET FOREIGN_KEY_CHECKS = %s' % (1 if enable else 0)) def setup_db(self, db: SQLAlchemy, app: Flask, config: MySQLConfig) -> None: """Set up the DB with a Flask app.""" url = _build_url(config).render_as_string(hide_password=False) app.config.setdefault('SQLALCHEMY_TRACK_MODIFICATIONS', False) app.config['SQLALCHEMY_DATABASE_URI'] = url db.init_app(app) def truncate_tables(self, conn: Connection, tables: list[str]) -> None: """Hard reset tables (and AUTO_INCREMENT) before a test transaction starts.""" if not tables: return # Run DDL in autocommit clone conn_ac = conn.execution_options(isolation_level='AUTOCOMMIT') preparer = conn_ac.dialect.identifier_preparer # Store original FK setting current_fk = conn_ac.exec_driver_sql('SELECT @@FOREIGN_KEY_CHECKS').scalar() fk_enabled = bool(int(current_fk)) if current_fk is not None else True self.set_fk(conn_ac, False) try: for t in tables: name = preparer.quote(t) conn_ac.exec_driver_sql(f'TRUNCATE TABLE {name}') finally: self.set_fk(conn_ac, fk_enabled) def _clone_db_data( self, src_engine: Engine, dst_engine: Engine, tables: list[str] | None = None ) -> None: """Copy data from source_db -> target_db (same host).""" source_db = src_engine.url.database # Resolve table list if not tables: with src_engine.connect() as sconn: tables = ( sconn.execute( text(""" SELECT TABLE_NAME FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = :db AND TABLE_TYPE='BASE TABLE' ORDER BY TABLE_NAME """), {'db': source_db}, ) .scalars() .all() ) with dst_engine.begin() as dconn: self.set_fk(dconn, False) for t in tables: # Skip if target already has rows (idempotence guard) has_rows = dconn.exec_driver_sql(f'SELECT 1 FROM `{t}` LIMIT 1').first() if has_rows: continue # Copy rows dconn.exec_driver_sql( f'INSERT INTO `{t}` SELECT * FROM `{source_db}`.`{t}`' ) # Copy AUTO_INCREMENT from source with src_engine.connect() as sconn: src_ai = sconn.execute( text(""" SELECT AUTO_INCREMENT FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA=:db AND TABLE_NAME=:t """), {'db': source_db, 't': t}, ).scalar() if src_ai: # may be None for tables without AI dconn.exec_driver_sql( f'ALTER TABLE `{t}` AUTO_INCREMENT={int(src_ai)}' ) self.set_fk(dconn, True) def _clone_db_schema(self, src_engine: Engine, dst_engine: Engine) -> None: """Create all base tables in dest from source (structure only).""" # Check if target already has tables. with dst_engine.connect() as dconn: has_any = dconn.execute( text(""" SELECT 1 FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA=:db LIMIT 1 """), {'db': dst_engine.url.database}, ).first() # If yes, assume schema already exists. if has_any: return # Copy schema from source to destination with src_engine.connect() as sconn, dst_engine.begin() as dconn: # Disable FK checks so creation order doesn't matter self.set_fk(dconn, False) # Get table names source_db = src_engine.url.database table_names = ( sconn.execute( text(""" SELECT TABLE_NAME FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = :db AND TABLE_TYPE='BASE TABLE' ORDER BY TABLE_NAME """), {'db': source_db}, ) .scalars() .all() ) # Copy tables for t in table_names: row = sconn.exec_driver_sql( f'SHOW CREATE TABLE `{source_db}`.`{t}`' ).first() create_stmt = row[1] # Second column is the CREATE statement # Make idempotent create_stmt = _normalize_create_table(source_db, create_stmt) dconn.exec_driver_sql(create_stmt) # Re-enable FK checks self.set_fk(dconn, True) def _clone_db_triggers( self, src_engine: Engine, dst_engine: Engine, drop_existing: bool = False ) -> None: """Clone all triggers from source -> target (same host). Requires TRIGGER privilege on target. Performs the following: - Strips DEFINER to avoid privilege issues. - De-qualifies any `source_db`.`table` references. - Optionally drops existing triggers first for idempotence. """ src_db = src_engine.url.database dst_db = dst_engine.url.database with src_engine.connect() as s: names = ( s.execute( text(""" SELECT TRIGGER_NAME FROM INFORMATION_SCHEMA.TRIGGERS WHERE TRIGGER_SCHEMA=:db ORDER BY TRIGGER_NAME """), {'db': src_db}, ) .scalars() .all() ) if not names: return with dst_engine.begin() as d: d.exec_driver_sql(f'USE `{dst_db}`') try: d.exec_driver_sql('SET SESSION log_bin_trust_function_creators=1') except OperationalError: pass for name in names: with src_engine.connect() as s: row = s.exec_driver_sql(f'SHOW CREATE TRIGGER `{name}`').first() create_stmt = _normalize_create_trigger(src_db, row[2]) if drop_existing: d.exec_driver_sql(f'DROP TRIGGER IF EXISTS `{name}`') try: d.exec_driver_sql(create_stmt) except OperationalError as e: if getattr(getattr(e, 'orig', None), 'args', [None])[0] == 1419: log.warning( '1419: Enable log_bin_trust_function_creators ' 'or disable binlog for tests; skipping trigger %s', name, ) def _build_url(config: MySQLConfig) -> URL: return URL.create( 'mysql+pymysql', username=str(config.MYSQL_DB_USER), password=str(config.MYSQL_DB_PASS), host=str(config.MYSQL_DB_HOST), port=int(config.MYSQL_DB_PORT), database=str(config.MYSQL_DB_NAME), query={ 'charset': 'utf8mb4', 'connect_timeout': '10', 'read_timeout': '30', 'write_timeout': '30', }, ) def _get_server_engine(config: MySQLConfig) -> tuple[Engine, str]: full_url = _build_url(config) dbname = full_url.database or '' if not dbname: raise ValueError('Config must specify MYSQL_DB_NAME (database).') server_url = URL.create( drivername=full_url.drivername, username=full_url.username, password=full_url.password, host=full_url.host, port=full_url.port, database=None, query=full_url.query, ) engine = sa_create_engine(server_url, pool_pre_ping=True, pool_recycle=1800) return engine, dbname def _normalize_create_table(db: str, sql: str) -> str: # Make idempotent (only the first CREATE TABLE) sql = sql.replace('CREATE TABLE ', 'CREATE TABLE IF NOT EXISTS ', 1) # Remove backticked schema qualifiers: `src_db`. pattern = rf'`{re.escape(db)}`\.' sql = re.sub(pattern, '', sql) return sql def _normalize_create_trigger(db: str, sql: str) -> str: sql = re.sub( r'CREATE\s+DEFINER=`[^`]+`@`[^`]+`\s+TRIGGER', 'CREATE TRIGGER IF NOT EXISTS', sql, count=1, flags=re.I, ) return sql.replace(f'`{db}`.', '')