#!/usr/bin/env python # -*- coding: utf-8 -*- """MySQL 异步处理器""" import logging import time from typing import Any, Dict, List, Optional from core.database_manager.handlers.common import ( format_connection_error, format_size, log_database_connect_failure, serialize_row, ) from core.database_manager.handlers.pools import get_aiomysql_pool from core.database_manager.handlers.transaction_mixin import HandlerTransactionMixin from core.database_manager.sql_utils import quote_identifier, split_sql_statements logger = logging.getLogger(__name__) class MySQLHandler(HandlerTransactionMixin): """MySQL异步处理器""" def __init__(self, host: str, port: int, user: str, password: str, database: str): self.host = host self.port = port self.user = user self.password = password self.database = database self.conn = None self._pool = None self._last_connect_error: Optional[str] = None self._in_transaction = False async def _release_connection(self) -> None: if self.conn and self._pool: try: self._pool.release(self.conn) except Exception as exc: logger.warning("Error releasing MySQL connection: %s", exc) self.conn = None self._pool = None async def connect(self, database: str = None) -> bool: db = database or self.database try: await self._release_connection() pool = await get_aiomysql_pool( self.host, self.port, self.user, self.password, db, ) self._pool = pool self.conn = await pool.acquire() self._last_connect_error = None return True except Exception as e: self._last_connect_error = format_connection_error(e) log_database_connect_failure( db_type="mysql", host=self.host, port=self.port, user=self.user, database=db, detail=self._last_connect_error, ) self.conn = None self._pool = None return False async def close(self): await self._release_connection() async def _execute_query(self, query: str, params: tuple = None) -> List[Dict[str, Any]]: import aiomysql async with self.conn.cursor(aiomysql.DictCursor) as cursor: await cursor.execute(query, params or ()) rows = await cursor.fetchall() return [dict(row) for row in rows] async def _execute_command(self, command: str, params: tuple = None) -> int: async with self.conn.cursor() as cursor: await cursor.execute(command, params or ()) return cursor.rowcount async def get_databases(self) -> List[Dict[str, Any]]: try: if not await self.connect(): return [] query = """ SELECT SCHEMA_NAME as name, DEFAULT_CHARACTER_SET_NAME as encoding, DEFAULT_COLLATION_NAME as collation, (SELECT COUNT(*) FROM information_schema.TABLES WHERE TABLE_SCHEMA = SCHEMA_NAME AND TABLE_TYPE = 'BASE TABLE') as tables_count FROM information_schema.SCHEMATA WHERE SCHEMA_NAME NOT IN ('information_schema', 'mysql', 'performance_schema', 'sys') ORDER BY SCHEMA_NAME """ databases = await self._execute_query(query) for db in databases: size_result = await self._execute_query("SELECT ROUND(SUM(data_length + index_length), 2) as size_bytes FROM information_schema.TABLES WHERE table_schema = %s", (db['name'],)) size_bytes = size_result[0]['size_bytes'] if size_result and size_result[0]['size_bytes'] else 0 db['size_bytes'] = int(size_bytes) if size_bytes else 0 db['size'] = format_size(db['size_bytes']) return databases except Exception as e: logger.error(f"Failed to get databases: {e}") return [] finally: await self.close() async def create_database(self, name: str, charset: str = "utf8mb4", collation: str = "utf8mb4_unicode_ci", **kwargs) -> bool: try: if not await self.connect(): return False qname = quote_identifier(name, "mysql") await self._execute_command( f"CREATE DATABASE {qname} CHARACTER SET {charset} COLLATE {collation}" ) return True except Exception as e: logger.error(f"Failed to create database {name}: {e}") raise finally: await self.close() async def drop_database(self, name: str) -> bool: try: if not await self.connect(): return False qname = quote_identifier(name, "mysql") await self._execute_command(f"DROP DATABASE {qname}") return True except Exception as e: logger.error(f"Failed to drop database {name}: {e}") raise finally: await self.close() async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]: databases = await self.get_databases() return [{"name": db["name"], "owner": None, "tables_count": db.get("tables_count", 0)} for db in databases] async def get_tables(self, database: str = None, schema_name: str = None) -> List[Dict[str, Any]]: try: db_name = database or schema_name or self.database if not await self.connect(db_name): return [] query = """ SELECT TABLE_SCHEMA as schema_name, TABLE_NAME as table_name, TABLE_TYPE as table_type, TABLE_ROWS as row_count, DATA_LENGTH as data_length, INDEX_LENGTH as index_length, (DATA_LENGTH + INDEX_LENGTH) as total_size_bytes, TABLE_COMMENT as description FROM information_schema.TABLES WHERE TABLE_SCHEMA = %s AND TABLE_TYPE = 'BASE TABLE' ORDER BY TABLE_NAME """ tables = await self._execute_query(query, (db_name,)) for table in tables: table['total_size'] = format_size(table.get('total_size_bytes', 0) or 0) return tables except Exception as e: logger.error(f"Failed to get tables: {e}") return [] finally: await self.close() async def get_table_columns(self, table_name: str, schema_name: str = None) -> List[Dict[str, Any]]: try: db_name = schema_name or self.database if not await self.connect(db_name): return [] query = """ SELECT COLUMN_NAME as column_name, DATA_TYPE as data_type, IS_NULLABLE = 'YES' as is_nullable, COLUMN_DEFAULT as column_default, CHARACTER_MAXIMUM_LENGTH as character_maximum_length, NUMERIC_PRECISION as numeric_precision, NUMERIC_SCALE as numeric_scale, ORDINAL_POSITION as ordinal_position, COLUMN_KEY = 'PRI' as is_primary_key, COLUMN_KEY = 'UNI' as is_unique, COLUMN_COMMENT as description FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s ORDER BY ORDINAL_POSITION """ rows = await self._execute_query(query, (db_name, table_name)) return rows except Exception as e: logger.error(f"Failed to get table columns: {e}") return [] finally: await self.close() async def get_table_indexes(self, table_name: str, schema_name: str = None) -> List[Dict[str, Any]]: try: db_name = schema_name or self.database if not await self.connect(db_name): return [] query = """ SELECT INDEX_NAME as index_name, INDEX_TYPE as index_type, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) as columns, NON_UNIQUE = 0 as is_unique, INDEX_NAME = 'PRIMARY' as is_primary FROM information_schema.STATISTICS WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s GROUP BY INDEX_NAME, INDEX_TYPE, NON_UNIQUE ORDER BY INDEX_NAME """ indexes = await self._execute_query(query, (db_name, table_name)) for idx in indexes: idx['definition'] = f"INDEX {idx['index_name']} ({idx['columns']})" return indexes except Exception as e: logger.error(f"Failed to get table indexes: {e}") return [] finally: await self.close() async def get_table_constraints(self, table_name: str, schema_name: str = None) -> List[Dict[str, Any]]: try: db_name = schema_name or self.database if not await self.connect(db_name): return [] query = """ SELECT tc.CONSTRAINT_NAME as constraint_name, tc.CONSTRAINT_TYPE as constraint_type, GROUP_CONCAT(kcu.COLUMN_NAME) as columns, kcu.REFERENCED_TABLE_NAME as referenced_table, NULL as referenced_columns FROM information_schema.TABLE_CONSTRAINTS tc JOIN information_schema.KEY_COLUMN_USAGE kcu ON tc.CONSTRAINT_NAME = kcu.CONSTRAINT_NAME AND tc.TABLE_SCHEMA = kcu.TABLE_SCHEMA AND tc.TABLE_NAME = kcu.TABLE_NAME WHERE tc.TABLE_SCHEMA = %s AND tc.TABLE_NAME = %s GROUP BY tc.CONSTRAINT_NAME, tc.CONSTRAINT_TYPE, kcu.REFERENCED_TABLE_NAME ORDER BY tc.CONSTRAINT_TYPE, tc.CONSTRAINT_NAME """ constraints = await self._execute_query(query, (db_name, table_name)) for const in constraints: const['definition'] = f"{const['constraint_type']} ({const['columns']})" return constraints except Exception as e: logger.error(f"Failed to get table constraints: {e}") return [] finally: await self.close() async def get_table_structure(self, table_name: str, database: str = None, schema_name: str = None) -> Dict[str, Any]: db_name = database or schema_name or self.database tables = await self.get_tables(database=db_name, schema_name=db_name) table_info = next((t for t in tables if t['table_name'] == table_name), None) if not table_info: raise ValueError(f"Table {db_name}.{table_name} not found") columns = await self.get_table_columns(table_name, db_name) indexes = await self.get_table_indexes(table_name, db_name) constraints = await self.get_table_constraints(table_name, db_name) return {"table_info": table_info, "columns": columns, "indexes": indexes, "constraints": constraints} async def get_table_ddl(self, table_name: str, schema_name: str = None) -> str: try: db_name = schema_name or self.database if not await self.connect(db_name): return "-- 无法连接数据库" result = await self._execute_query(f"SHOW CREATE TABLE `{table_name}`") if result: return result[0].get('Create Table', f"-- 无法获取表 {table_name} 的DDL") return f"-- 无法获取表 {table_name} 的DDL" except Exception as e: return f"-- 获取DDL失败: {str(e)}" finally: await self.close() async def get_views(self, database: str = None, schema_name: str = None) -> List[Dict[str, Any]]: try: db_name = database or schema_name or self.database if not await self.connect(db_name): return [] query = """ SELECT TABLE_NAME as view_name, TABLE_SCHEMA as schema_name, VIEW_DEFINITION as view_definition, IS_UPDATABLE as is_updatable, CHECK_OPTION as check_option, 'VIEW' as view_type FROM information_schema.VIEWS WHERE TABLE_SCHEMA = %s ORDER BY TABLE_NAME """ result = await self._execute_query(query, (db_name,)) for row in result: row['is_updatable'] = row.get('is_updatable') == 'YES' return result except Exception as e: logger.error(f"Failed to get views: {e}") return [] finally: await self.close() async def get_view_structure(self, view_name: str, schema_name: str = None) -> Dict[str, Any]: try: db_name = schema_name or self.database if not await self.connect(db_name): raise ValueError("Failed to connect") view_query = "SELECT TABLE_NAME as view_name, TABLE_SCHEMA as schema_name, VIEW_DEFINITION as view_definition, IS_UPDATABLE as is_updatable, CHECK_OPTION as check_option, 'VIEW' as view_type FROM information_schema.VIEWS WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s" view_info = await self._execute_query(view_query, (db_name, view_name)) if not view_info: raise ValueError(f"View {db_name}.{view_name} not found") view_info = view_info[0] view_info['is_updatable'] = view_info.get('is_updatable') == 'YES' columns_query = "SELECT COLUMN_NAME as column_name, DATA_TYPE as data_type, IS_NULLABLE = 'YES' as is_nullable, ORDINAL_POSITION as ordinal_position, COLUMN_COMMENT as description FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s ORDER BY ORDINAL_POSITION" columns = await self._execute_query(columns_query, (db_name, view_name)) definition_sql = await self._get_view_definition_internal(view_name) dependencies = await self._get_view_dependencies_internal(view_name, db_name) return {"view_info": view_info, "columns": columns, "dependencies": dependencies, "definition_sql": definition_sql} except Exception as e: logger.error(f"Failed to get view structure: {e}") raise finally: await self.close() async def _get_view_definition_internal(self, view_name: str) -> str: try: result = await self._execute_query(f"SHOW CREATE VIEW `{view_name}`") if result: return result[0].get('Create View', f"-- 无法获取视图 {view_name} 的定义") except Exception as e: logger.error(f"Failed to get view definition: {e}") return f"-- 无法获取视图 {view_name} 的定义" async def _get_view_dependencies_internal(self, view_name: str, schema_name: str) -> List[str]: try: query = "SELECT DISTINCT REFERENCED_TABLE_NAME as table_name FROM information_schema.VIEW_TABLE_USAGE WHERE VIEW_SCHEMA = %s AND VIEW_NAME = %s AND REFERENCED_TABLE_NAME IS NOT NULL ORDER BY REFERENCED_TABLE_NAME" result = await self._execute_query(query, (schema_name, view_name)) return [row['table_name'] for row in result] except Exception as e: logger.error(f"Failed to get view dependencies: {e}") return [] async def get_view_definition(self, view_name: str, schema_name: str = None) -> str: try: db_name = schema_name or self.database if not await self.connect(db_name): return "-- 无法连接数据库" result = await self._get_view_definition_internal(view_name) return result except Exception as e: return f"-- 获取视图定义失败: {str(e)}" finally: await self.close() async def get_view_dependencies(self, view_name: str, schema_name: str = None) -> List[str]: try: db_name = schema_name or self.database if not await self.connect(db_name): return [] result = await self._get_view_dependencies_internal(view_name, db_name) return result except Exception as e: return [] finally: await self.close() async def query_data(self, table_name: str, schema_name: str = None, page: int = 1, page_size: int = 20, where: str = None, order_by: str = None, database: str = None) -> Dict[str, Any]: try: db_name = database or schema_name or self.database if not await self.connect(db_name): return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size} count_q = f"SELECT COUNT(*) as total FROM `{table_name}`" + (f" WHERE {where}" if where else "") count_result = await self._execute_query(count_q) total = count_result[0]['total'] if count_result else 0 offset = (page - 1) * page_size data_q = f"SELECT * FROM `{table_name}`" + (f" WHERE {where}" if where else "") + (f" ORDER BY {order_by}" if order_by else "") + f" LIMIT {page_size} OFFSET {offset}" rows = await self._execute_query(data_q) rows_list = [serialize_row(row) for row in rows] columns = list(rows_list[0].keys()) if rows_list else [] return {"columns": columns, "rows": rows_list, "total": total, "page": page, "page_size": page_size} except Exception as e: logger.error(f"Failed to query data: {e}") return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size} finally: await self.close() async def _run_sql_query(self, sql: str) -> list: return await self._execute_query(sql) async def _run_sql_command(self, sql: str) -> int: return await self._execute_command(sql) async def insert_data(self, table_name: str, data: Dict[str, Any], schema_name: str = None) -> Dict[str, Any]: try: db_name = schema_name or self.database if not await self.connect(db_name): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} columns = list(data.keys()) values = list(data.values()) placeholders = ', '.join(['%s'] * len(values)) query = f"INSERT INTO `{table_name}` ({', '.join(columns)}) VALUES ({placeholders})" affected_rows = await self._execute_command(query, tuple(values)) return {"success": True, "message": "插入成功", "affected_rows": affected_rows} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def update_data(self, table_name: str, data: Dict[str, Any], where: str, schema_name: str = None) -> Dict[str, Any]: try: db_name = schema_name or self.database if not await self.connect(db_name): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} set_clause = ', '.join([f'{k} = %s' for k in data.keys()]) query = f"UPDATE `{table_name}` SET {set_clause} WHERE {where}" affected_rows = await self._execute_command(query, tuple(data.values())) return {"success": True, "message": f"更新成功,影响 {affected_rows} 行", "affected_rows": affected_rows} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def delete_data(self, table_name: str, where: str, schema_name: str = None) -> Dict[str, Any]: try: db_name = schema_name or self.database if not await self.connect(db_name): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} query = f"DELETE FROM `{table_name}` WHERE {where}" affected_rows = await self._execute_command(query) return {"success": True, "message": f"删除成功,影响 {affected_rows} 行", "affected_rows": affected_rows} except Exception as e: return {"success": False, "message": str(e), "affected_rows": 0} finally: await self.close() async def execute_ddl(self, sql: str, database: str = None, schema_name: str = None) -> Dict[str, Any]: try: db_name = database or schema_name or self.database if not await self.connect(db_name): return {"success": False, "message": "数据库连接失败", "affected_rows": 0} statements = split_sql_statements(sql) if not statements: return {"success": False, "message": "无有效SQL语句", "affected_rows": 0} await self.conn.begin() try: for index, statement in enumerate(statements, start=1): try: await self._execute_command(statement) except Exception as exc: await self.conn.rollback() return { "success": False, "message": f"DDL执行失败 (第{index}条): {exc}", "affected_rows": 0, } await self.conn.commit() except Exception: await self.conn.rollback() raise return {"success": True, "message": "DDL执行成功", "affected_rows": 0} except Exception as e: return {"success": False, "message": f"DDL执行失败: {str(e)}", "affected_rows": 0} finally: await self.close()