435 lines
21 KiB
Python
435 lines
21 KiB
Python
#!/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()
|
|
|
|
|