Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,14 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Handler 导出"""
|
||||
from core.database_manager.handlers.mysql import MySQLHandler
|
||||
from core.database_manager.handlers.oracle import OracleHandler
|
||||
from core.database_manager.handlers.postgresql import PostgreSQLHandler
|
||||
from core.database_manager.handlers.sqlserver import SQLServerHandler
|
||||
|
||||
__all__ = [
|
||||
"PostgreSQLHandler",
|
||||
"MySQLHandler",
|
||||
"SQLServerHandler",
|
||||
"OracleHandler",
|
||||
]
|
||||
@@ -0,0 +1,68 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Handler 共享工具"""
|
||||
import logging
|
||||
from typing import Any, Dict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def format_size(size_bytes: int) -> str:
|
||||
"""格式化字节大小"""
|
||||
if not size_bytes:
|
||||
return "0 bytes"
|
||||
if size_bytes >= 1073741824:
|
||||
return f"{size_bytes / 1073741824:.2f} GB"
|
||||
if size_bytes >= 1048576:
|
||||
return f"{size_bytes / 1048576:.2f} MB"
|
||||
if size_bytes >= 1024:
|
||||
return f"{size_bytes / 1024:.2f} KB"
|
||||
return f"{size_bytes} bytes"
|
||||
|
||||
|
||||
def serialize_row(row: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""序列化行数据"""
|
||||
for key, value in row.items():
|
||||
if hasattr(value, "isoformat"):
|
||||
row[key] = value.isoformat()
|
||||
elif isinstance(value, bytes):
|
||||
row[key] = value.decode("utf-8", errors="replace")
|
||||
elif isinstance(value, (set, frozenset)):
|
||||
row[key] = list(value)
|
||||
return row
|
||||
|
||||
|
||||
def format_connection_error(exc: Exception) -> str:
|
||||
"""将驱动异常转为可读的错误详情"""
|
||||
if exc is None:
|
||||
return "未知错误"
|
||||
msg = str(exc).strip()
|
||||
if msg:
|
||||
return msg
|
||||
return exc.__class__.__name__
|
||||
|
||||
|
||||
def log_database_connect_failure(
|
||||
*,
|
||||
db_type: str,
|
||||
host: str,
|
||||
port: int,
|
||||
user: str = "",
|
||||
database: str = "",
|
||||
db_name: str = "",
|
||||
detail: str,
|
||||
action: str = "connect",
|
||||
) -> None:
|
||||
"""记录数据库连接失败详情到后台日志"""
|
||||
logger.error(
|
||||
"Database connection failed [%s]: db_name=%s db_type=%s target=%s:%s "
|
||||
"database=%s user=%s error=%s",
|
||||
action,
|
||||
db_name or "-",
|
||||
db_type,
|
||||
host,
|
||||
port,
|
||||
database or "-",
|
||||
user or "-",
|
||||
detail,
|
||||
)
|
||||
@@ -0,0 +1,434 @@
|
||||
#!/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()
|
||||
|
||||
|
||||
@@ -0,0 +1,673 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Oracle 异步处理器(oracledb async pool)"""
|
||||
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_oracledb_pool
|
||||
from core.database_manager.handlers.transaction_mixin import HandlerTransactionMixin
|
||||
from core.database_manager.sql_utils import quote_identifier, quote_table, split_sql_statements
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ORACLE_SYSTEM_SCHEMAS = (
|
||||
"SYS",
|
||||
"SYSTEM",
|
||||
"OUTLN",
|
||||
"XDB",
|
||||
"CTXSYS",
|
||||
"MDSYS",
|
||||
"ORDDATA",
|
||||
"ORDSYS",
|
||||
"LBACSYS",
|
||||
"DBSNMP",
|
||||
"APPQOSSYS",
|
||||
"AUDSYS",
|
||||
"GSMADMIN_INTERNAL",
|
||||
"OJVMSYS",
|
||||
"ORDPLUGINS",
|
||||
"SI_INFORMTN_SCHEMA",
|
||||
"WMSYS",
|
||||
)
|
||||
|
||||
|
||||
class OracleHandler(HandlerTransactionMixin):
|
||||
"""Oracle 异步处理器"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.user = user
|
||||
self.password = password
|
||||
self.database = database
|
||||
self.extra_options = extra_options or {}
|
||||
self.conn = None
|
||||
self._pool = None
|
||||
self._last_connect_error: Optional[str] = None
|
||||
self._in_transaction = False
|
||||
|
||||
@property
|
||||
def default_schema(self) -> str:
|
||||
return (self.user or "").upper()
|
||||
|
||||
async def _release_connection(self) -> None:
|
||||
if self.conn and self._pool:
|
||||
try:
|
||||
await self._pool.release(self.conn)
|
||||
except Exception as exc:
|
||||
logger.warning("Error releasing Oracle connection: %s", exc)
|
||||
self.conn = None
|
||||
self._pool = None
|
||||
|
||||
async def connect(self, database: str = None) -> bool:
|
||||
_ = database # Oracle 使用 service name,不按 PG 方式切换 database
|
||||
try:
|
||||
await self._release_connection()
|
||||
pool = await get_oracledb_pool(
|
||||
self.host,
|
||||
self.port,
|
||||
self.user,
|
||||
self.password,
|
||||
self.database,
|
||||
self.extra_options,
|
||||
)
|
||||
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="oracle",
|
||||
host=self.host,
|
||||
port=self.port,
|
||||
user=self.user,
|
||||
database=self.database,
|
||||
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: dict = None) -> List[Dict[str, Any]]:
|
||||
async with self.conn.cursor() as cursor:
|
||||
await cursor.execute(query, params or {})
|
||||
if not cursor.description:
|
||||
return []
|
||||
columns = [col[0].lower() for col in cursor.description]
|
||||
rows = await cursor.fetchall()
|
||||
return [dict(zip(columns, row)) for row in rows]
|
||||
|
||||
async def _execute_command(self, command: str, params: dict = None) -> int:
|
||||
async with self.conn.cursor() as cursor:
|
||||
await cursor.execute(command, params or {})
|
||||
return cursor.rowcount
|
||||
|
||||
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 begin_transaction(self) -> bool:
|
||||
if self._in_transaction:
|
||||
return True
|
||||
if not await self.connect():
|
||||
return False
|
||||
self._in_transaction = True
|
||||
return True
|
||||
|
||||
async def commit_transaction(self) -> None:
|
||||
if not self._in_transaction:
|
||||
return
|
||||
try:
|
||||
await self.conn.commit()
|
||||
finally:
|
||||
self._in_transaction = False
|
||||
await self.close()
|
||||
|
||||
async def rollback_transaction(self) -> None:
|
||||
if not self._in_transaction:
|
||||
return
|
||||
try:
|
||||
await self.conn.rollback()
|
||||
finally:
|
||||
self._in_transaction = False
|
||||
await self.close()
|
||||
|
||||
async def get_databases(self) -> List[Dict[str, Any]]:
|
||||
"""Oracle 返回虚拟 database(service name)"""
|
||||
try:
|
||||
if not await self.connect():
|
||||
return []
|
||||
service_name = self.database or self.default_schema or "ORCL"
|
||||
count_rows = await self._execute_query(
|
||||
"SELECT COUNT(*) AS tables_count FROM user_tables"
|
||||
)
|
||||
tables_count = count_rows[0]["tables_count"] if count_rows else 0
|
||||
return [
|
||||
{
|
||||
"name": service_name,
|
||||
"owner": self.user,
|
||||
"encoding": "UTF-8",
|
||||
"collation": None,
|
||||
"size": "0 bytes",
|
||||
"size_bytes": 0,
|
||||
"description": "Oracle service instance",
|
||||
"tables_count": tables_count,
|
||||
}
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error("Failed to get databases: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def create_database(self, name: str, **kwargs) -> bool:
|
||||
raise NotImplementedError("Oracle 不支持创建 Database")
|
||||
|
||||
async def drop_database(self, name: str) -> bool:
|
||||
raise NotImplementedError("Oracle 不支持删除 Database")
|
||||
|
||||
async def rename_database(self, name: str, new_name: str) -> bool:
|
||||
raise NotImplementedError("Oracle 不支持重命名 Database")
|
||||
|
||||
async def create_schema(self, name: str, database: str = None, owner: str = None) -> bool:
|
||||
raise NotImplementedError("Oracle Schema 创建需要 DBA 权限,暂不支持")
|
||||
|
||||
async def drop_schema(self, name: str, database: str = None, cascade: bool = True) -> bool:
|
||||
raise NotImplementedError("Oracle Schema 删除需要 DBA 权限,暂不支持")
|
||||
|
||||
async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool:
|
||||
raise NotImplementedError("Oracle 不支持重命名 Schema")
|
||||
|
||||
async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
placeholders = ", ".join([f":s{i}" for i in range(len(_ORACLE_SYSTEM_SCHEMAS))])
|
||||
params = {f"s{i}": schema for i, schema in enumerate(_ORACLE_SYSTEM_SCHEMAS)}
|
||||
query = f"""
|
||||
SELECT username AS name,
|
||||
username AS owner,
|
||||
(SELECT COUNT(*)
|
||||
FROM all_tables t
|
||||
WHERE t.owner = u.username) AS tables_count
|
||||
FROM all_users u
|
||||
WHERE u.username NOT IN ({placeholders})
|
||||
ORDER BY u.username
|
||||
"""
|
||||
return await self._execute_query(query, params)
|
||||
except Exception as e:
|
||||
logger.error("Failed to get schemas: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_tables(
|
||||
self, database: str = None, schema_name: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
query = """
|
||||
SELECT owner AS schema_name,
|
||||
table_name,
|
||||
'BASE TABLE' AS table_type,
|
||||
num_rows AS row_count,
|
||||
0 AS total_size_bytes
|
||||
FROM all_tables
|
||||
WHERE owner = :owner
|
||||
ORDER BY table_name
|
||||
"""
|
||||
tables = await self._execute_query(query, {"owner": schema})
|
||||
for table in tables:
|
||||
table["total_size"] = format_size(0)
|
||||
return tables
|
||||
except Exception as e:
|
||||
logger.error("Failed to get tables: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_columns(
|
||||
self, table_name: str, schema_name: str = None, database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
query = """
|
||||
SELECT c.column_name,
|
||||
c.data_type,
|
||||
CASE WHEN c.nullable = 'Y' THEN 1 ELSE 0 END AS is_nullable,
|
||||
c.data_default AS column_default,
|
||||
c.char_length AS character_maximum_length,
|
||||
c.data_precision AS numeric_precision,
|
||||
c.data_scale AS numeric_scale,
|
||||
c.column_id AS ordinal_position,
|
||||
CASE WHEN pk.column_name IS NOT NULL THEN 1 ELSE 0 END AS is_primary_key,
|
||||
CASE WHEN uq.column_name IS NOT NULL THEN 1 ELSE 0 END AS is_unique,
|
||||
cc.comments AS description
|
||||
FROM all_tab_columns c
|
||||
LEFT JOIN (
|
||||
SELECT cols.column_name, cols.owner, cols.table_name
|
||||
FROM all_constraints cons
|
||||
JOIN all_cons_columns cols
|
||||
ON cons.constraint_name = cols.constraint_name
|
||||
AND cons.owner = cols.owner
|
||||
WHERE cons.constraint_type = 'P'
|
||||
) pk ON pk.owner = c.owner
|
||||
AND pk.table_name = c.table_name
|
||||
AND pk.column_name = c.column_name
|
||||
LEFT JOIN (
|
||||
SELECT cols.column_name, cols.owner, cols.table_name
|
||||
FROM all_constraints cons
|
||||
JOIN all_cons_columns cols
|
||||
ON cons.constraint_name = cols.constraint_name
|
||||
AND cons.owner = cols.owner
|
||||
WHERE cons.constraint_type = 'U'
|
||||
) uq ON uq.owner = c.owner
|
||||
AND uq.table_name = c.table_name
|
||||
AND uq.column_name = c.column_name
|
||||
LEFT JOIN all_col_comments cc
|
||||
ON cc.owner = c.owner
|
||||
AND cc.table_name = c.table_name
|
||||
AND cc.column_name = c.column_name
|
||||
WHERE c.owner = :owner AND c.table_name = :table_name
|
||||
ORDER BY c.column_id
|
||||
"""
|
||||
return await self._execute_query(
|
||||
query, {"owner": schema, "table_name": table_name.upper()}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table columns: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_indexes(
|
||||
self, table_name: str, schema_name: str = None, database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
query = """
|
||||
SELECT i.index_name,
|
||||
i.index_type,
|
||||
LISTAGG(c.column_name, ', ') WITHIN GROUP (ORDER BY c.column_position) AS columns,
|
||||
CASE WHEN i.uniqueness = 'UNIQUE' THEN 1 ELSE 0 END AS is_unique,
|
||||
CASE WHEN i.index_name IN (
|
||||
SELECT constraint_name FROM all_constraints
|
||||
WHERE owner = :owner AND table_name = :table_name
|
||||
AND constraint_type = 'P'
|
||||
) THEN 1 ELSE 0 END AS is_primary,
|
||||
'' AS definition
|
||||
FROM all_indexes i
|
||||
JOIN all_ind_columns c
|
||||
ON i.owner = c.index_owner
|
||||
AND i.index_name = c.index_name
|
||||
WHERE i.table_owner = :owner AND i.table_name = :table_name
|
||||
GROUP BY i.index_name, i.index_type, i.uniqueness
|
||||
ORDER BY i.index_name
|
||||
"""
|
||||
indexes = await self._execute_query(
|
||||
query, {"owner": schema, "table_name": table_name.upper()}
|
||||
)
|
||||
for idx in indexes:
|
||||
unique = "UNIQUE " if idx.get("is_unique") else ""
|
||||
idx["definition"] = (
|
||||
f"CREATE {unique}INDEX {idx['index_name']} "
|
||||
f"ON {schema}.{table_name.upper()} ({idx.get('columns', '')})"
|
||||
)
|
||||
return indexes
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table indexes: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_constraints(
|
||||
self, table_name: str, schema_name: str = None, database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
query = """
|
||||
SELECT c.constraint_name,
|
||||
c.constraint_type,
|
||||
LISTAGG(cc.column_name, ', ') WITHIN GROUP (ORDER BY cc.position) AS columns,
|
||||
ref.table_name AS referenced_table,
|
||||
LISTAGG(ref_cc.column_name, ', ') WITHIN GROUP (ORDER BY ref_cc.position)
|
||||
AS referenced_columns,
|
||||
c.search_condition AS definition
|
||||
FROM all_constraints c
|
||||
LEFT JOIN all_cons_columns cc
|
||||
ON c.owner = cc.owner
|
||||
AND c.constraint_name = cc.constraint_name
|
||||
LEFT JOIN all_constraints ref
|
||||
ON c.r_owner = ref.owner
|
||||
AND c.r_constraint_name = ref.constraint_name
|
||||
LEFT JOIN all_cons_columns ref_cc
|
||||
ON ref.owner = ref_cc.owner
|
||||
AND ref.constraint_name = ref_cc.constraint_name
|
||||
WHERE c.owner = :owner AND c.table_name = :table_name
|
||||
GROUP BY c.constraint_name, c.constraint_type, ref.table_name, c.search_condition
|
||||
ORDER BY c.constraint_type, c.constraint_name
|
||||
"""
|
||||
return await self._execute_query(
|
||||
query, {"owner": schema, "table_name": table_name.upper()}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table constraints: %s", 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]:
|
||||
schema = schema_name or self.default_schema
|
||||
tables = await self.get_tables(database=database, schema_name=schema)
|
||||
table_info = next(
|
||||
(t for t in tables if t["table_name"].upper() == table_name.upper()),
|
||||
None,
|
||||
)
|
||||
if not table_info:
|
||||
raise ValueError(f"Table {schema}.{table_name} not found")
|
||||
columns = await self.get_table_columns(table_name, schema, database)
|
||||
indexes = await self.get_table_indexes(table_name, schema, database)
|
||||
constraints = await self.get_table_constraints(table_name, schema, database)
|
||||
return {
|
||||
"table_info": table_info,
|
||||
"columns": columns,
|
||||
"indexes": indexes,
|
||||
"constraints": constraints,
|
||||
}
|
||||
|
||||
async def get_table_ddl(
|
||||
self, table_name: str, schema_name: str = None, database: str = None
|
||||
) -> str:
|
||||
try:
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
columns = await self.get_table_columns(table_name, schema, database)
|
||||
if not columns:
|
||||
return f"-- 无法获取表 {schema}.{table_name} 的DDL"
|
||||
full_name = quote_table(schema, table_name.upper(), "oracle")
|
||||
ddl_lines = [f"CREATE TABLE {full_name} ("]
|
||||
col_defs = []
|
||||
for col in columns:
|
||||
col_def = f" {quote_identifier(col['column_name'], 'oracle')} {col['data_type']}"
|
||||
if col.get("character_maximum_length"):
|
||||
col_def += f"({col['character_maximum_length']})"
|
||||
elif col.get("numeric_precision"):
|
||||
scale = col.get("numeric_scale")
|
||||
col_def += f"({col['numeric_precision']}"
|
||||
col_def += f",{scale})" if scale is not None else ")"
|
||||
if not col.get("is_nullable"):
|
||||
col_def += " NOT NULL"
|
||||
col_defs.append(col_def)
|
||||
ddl_lines.append(",\n".join(col_defs))
|
||||
ddl_lines.append(");")
|
||||
return "\n".join(ddl_lines)
|
||||
except Exception as e:
|
||||
return f"-- 获取DDL失败: {str(e)}"
|
||||
|
||||
async def get_views(
|
||||
self, database: str = None, schema_name: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
query = """
|
||||
SELECT view_name,
|
||||
owner AS schema_name,
|
||||
text AS view_definition,
|
||||
0 AS is_updatable,
|
||||
NULL AS check_option,
|
||||
'VIEW' AS view_type
|
||||
FROM all_views
|
||||
WHERE owner = :owner
|
||||
ORDER BY view_name
|
||||
"""
|
||||
return await self._execute_query(query, {"owner": schema})
|
||||
except Exception as e:
|
||||
logger.error("Failed to get views: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_structure(self, view_name: str, schema_name: str = None) -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
raise ValueError("Failed to connect")
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
view_query = """
|
||||
SELECT view_name, owner AS schema_name, text AS view_definition,
|
||||
0 AS is_updatable, NULL AS check_option, 'VIEW' AS view_type
|
||||
FROM all_views
|
||||
WHERE owner = :owner AND view_name = :view_name
|
||||
"""
|
||||
view_info = await self._execute_query(
|
||||
view_query, {"owner": schema, "view_name": view_name.upper()}
|
||||
)
|
||||
if not view_info:
|
||||
raise ValueError(f"View {schema}.{view_name} not found")
|
||||
view_info = view_info[0]
|
||||
columns = await self.get_table_columns(view_name, schema)
|
||||
definition_sql = await self.get_view_definition(view_name, schema)
|
||||
dependencies = await self.get_view_dependencies(view_name, schema)
|
||||
return {
|
||||
"view_info": view_info,
|
||||
"columns": columns,
|
||||
"dependencies": dependencies,
|
||||
"definition_sql": definition_sql,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error("Failed to get view structure: %s", e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_definition(self, view_name: str, schema_name: str = None) -> str:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return "-- 无法连接数据库"
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
rows = await self._execute_query(
|
||||
"""
|
||||
SELECT text AS definition
|
||||
FROM all_views
|
||||
WHERE owner = :owner AND view_name = :view_name
|
||||
""",
|
||||
{"owner": schema, "view_name": view_name.upper()},
|
||||
)
|
||||
if rows and rows[0].get("definition"):
|
||||
full = quote_table(schema, view_name.upper(), "oracle")
|
||||
return f"CREATE OR REPLACE VIEW {full} AS\n{rows[0]['definition']}"
|
||||
return f"-- 无法获取视图 {schema}.{view_name} 的定义"
|
||||
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:
|
||||
if not await self.connect():
|
||||
return []
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
rows = await self._execute_query(
|
||||
"""
|
||||
SELECT DISTINCT referenced_name AS table_name
|
||||
FROM all_dependencies
|
||||
WHERE owner = :owner
|
||||
AND name = :view_name
|
||||
AND type = 'VIEW'
|
||||
AND referenced_type IN ('TABLE', 'VIEW')
|
||||
ORDER BY referenced_name
|
||||
""",
|
||||
{"owner": schema, "view_name": view_name.upper()},
|
||||
)
|
||||
return [row["table_name"] for row in rows]
|
||||
except Exception as e:
|
||||
logger.error("Failed to get view dependencies: %s", 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:
|
||||
if not await self.connect(database):
|
||||
return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size}
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
full_table = quote_table(schema, table_name.upper(), "oracle")
|
||||
count_q = f"SELECT COUNT(*) AS total FROM {full_table}"
|
||||
if where:
|
||||
count_q += f" WHERE {where}"
|
||||
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 {full_table}"
|
||||
if where:
|
||||
data_q += f" WHERE {where}"
|
||||
if order_by:
|
||||
data_q += f" ORDER BY {order_by}"
|
||||
data_q += f" OFFSET {offset} ROWS FETCH NEXT {page_size} ROWS ONLY"
|
||||
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("Failed to query data: %s", e)
|
||||
return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size}
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def insert_data(
|
||||
self, table_name: str, data: Dict[str, Any], schema_name: str = None
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
full_table = quote_table(schema, table_name.upper(), "oracle")
|
||||
columns = list(data.keys())
|
||||
quoted_columns = ", ".join(
|
||||
quote_identifier(col, "oracle") for col in columns
|
||||
)
|
||||
placeholders = ", ".join([f":v{i}" for i in range(len(columns))])
|
||||
params = {f"v{i}": v for i, v in enumerate(data.values())}
|
||||
query = f"INSERT INTO {full_table} ({quoted_columns}) VALUES ({placeholders})"
|
||||
affected_rows = await self._execute_command(query, params)
|
||||
return {"success": True, "message": "插入成功", "affected_rows": affected_rows or 1}
|
||||
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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
full_table = quote_table(schema, table_name.upper(), "oracle")
|
||||
set_parts = []
|
||||
params = {}
|
||||
for i, (k, v) in enumerate(data.items()):
|
||||
key = f"s{i}"
|
||||
set_parts.append(f"{quote_identifier(k, 'oracle')} = :{key}")
|
||||
params[key] = v
|
||||
query = f"UPDATE {full_table} SET {', '.join(set_parts)} WHERE {where}"
|
||||
affected_rows = await self._execute_command(query, params)
|
||||
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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema = (schema_name or self.default_schema).upper()
|
||||
full_table = quote_table(schema, table_name.upper(), "oracle")
|
||||
query = f"DELETE FROM {full_table} 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:
|
||||
if not await self.connect(database):
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
statements = split_sql_statements(sql)
|
||||
if not statements:
|
||||
return {"success": False, "message": "无有效SQL语句", "affected_rows": 0}
|
||||
for index, statement in enumerate(statements, start=1):
|
||||
try:
|
||||
await self._execute_command(statement)
|
||||
except Exception as exc:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"DDL执行失败 (第{index}条): {exc}",
|
||||
"affected_rows": 0,
|
||||
}
|
||||
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()
|
||||
@@ -0,0 +1,297 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""外部数据库连接池(按连接目标复用)"""
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_POOL_MAX_SIZE = 5
|
||||
_pool_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def _pool_identity(user: str, password: str) -> str:
|
||||
"""生成连接池身份标识(含密码摘要,避免换密后复用旧池)"""
|
||||
digest = hashlib.sha256(f"{user or ''}\0{password or ''}".encode()).hexdigest()[:12]
|
||||
return f"{user or '-'}#{digest}"
|
||||
|
||||
|
||||
_pg_pool_cache: Dict[str, Any] = {}
|
||||
_mysql_pool_cache: Dict[str, Any] = {}
|
||||
_mssql_pool_cache: Dict[str, Any] = {}
|
||||
_oracle_pool_cache: Dict[str, Any] = {}
|
||||
|
||||
|
||||
async def get_asyncpg_pool(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
):
|
||||
import asyncpg
|
||||
|
||||
db_name = database or "postgres"
|
||||
key = f"pg://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
|
||||
async with _pool_lock:
|
||||
pool = _pg_pool_cache.get(key)
|
||||
if pool is None or getattr(pool, "_closed", False):
|
||||
pool = await asyncpg.create_pool(
|
||||
host=host,
|
||||
port=port,
|
||||
user=user,
|
||||
password=password,
|
||||
database=db_name,
|
||||
min_size=0,
|
||||
max_size=_POOL_MAX_SIZE,
|
||||
timeout=10,
|
||||
command_timeout=60,
|
||||
)
|
||||
_pg_pool_cache[key] = pool
|
||||
logger.debug("Created asyncpg pool: %s (max=%s)", key, _POOL_MAX_SIZE)
|
||||
return pool
|
||||
|
||||
|
||||
async def get_aiomysql_pool(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
):
|
||||
import aiomysql
|
||||
|
||||
db_name = database or ""
|
||||
key = f"mysql://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
|
||||
async with _pool_lock:
|
||||
pool = _mysql_pool_cache.get(key)
|
||||
if pool is None or pool.closed:
|
||||
pool = await aiomysql.create_pool(
|
||||
host=host,
|
||||
port=port,
|
||||
user=user,
|
||||
password=password,
|
||||
db=db_name,
|
||||
charset="utf8mb4",
|
||||
autocommit=True,
|
||||
minsize=0,
|
||||
maxsize=_POOL_MAX_SIZE,
|
||||
connect_timeout=10,
|
||||
)
|
||||
_mysql_pool_cache[key] = pool
|
||||
logger.debug("Created aiomysql pool: %s (max=%s)", key, _POOL_MAX_SIZE)
|
||||
return pool
|
||||
|
||||
|
||||
def build_mssql_dsn(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
) -> str:
|
||||
opts = extra_options or {}
|
||||
driver = opts.get("odbc_driver", "ODBC Driver 18 for SQL Server")
|
||||
db_name = database or "master"
|
||||
encrypt = opts.get("encrypt", "yes")
|
||||
trust = opts.get("trust_server_certificate", "yes")
|
||||
return (
|
||||
f"DRIVER={{{driver}}};"
|
||||
f"SERVER={host},{port};"
|
||||
f"DATABASE={db_name};"
|
||||
f"UID={user};"
|
||||
f"PWD={password};"
|
||||
f"Encrypt={encrypt};"
|
||||
f"TrustServerCertificate={trust};"
|
||||
)
|
||||
|
||||
|
||||
async def get_aioodbc_pool(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
import aioodbc
|
||||
|
||||
db_name = database or "master"
|
||||
key = f"mssql://{_pool_identity(user, password)}@{host}:{port}/{db_name}"
|
||||
async with _pool_lock:
|
||||
pool = _mssql_pool_cache.get(key)
|
||||
if pool is None or pool.closed:
|
||||
dsn = build_mssql_dsn(host, port, user, password, db_name, extra_options)
|
||||
pool = await aioodbc.create_pool(
|
||||
dsn=dsn,
|
||||
minsize=0,
|
||||
maxsize=_POOL_MAX_SIZE,
|
||||
autocommit=True,
|
||||
)
|
||||
_mssql_pool_cache[key] = pool
|
||||
logger.debug("Created aioodbc pool: %s (max=%s)", key, _POOL_MAX_SIZE)
|
||||
return pool
|
||||
|
||||
|
||||
def build_oracle_dsn(host: str, port: int, service_name: str) -> str:
|
||||
import oracledb
|
||||
|
||||
service = service_name or "ORCL"
|
||||
return oracledb.make_dsn(host, port, service_name=service)
|
||||
|
||||
|
||||
async def get_oracledb_pool(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
import oracledb
|
||||
|
||||
service = database or (extra_options or {}).get("service_name") or "ORCL"
|
||||
key = f"oracle://{_pool_identity(user, password)}@{host}:{port}/{service}"
|
||||
async with _pool_lock:
|
||||
pool = _oracle_pool_cache.get(key)
|
||||
if pool is None or (hasattr(pool, "opened") and not pool.opened):
|
||||
dsn = build_oracle_dsn(host, port, service)
|
||||
pool = oracledb.create_pool_async(
|
||||
user=user,
|
||||
password=password,
|
||||
dsn=dsn,
|
||||
min=0,
|
||||
max=_POOL_MAX_SIZE,
|
||||
)
|
||||
await pool.open()
|
||||
_oracle_pool_cache[key] = pool
|
||||
logger.debug("Created oracledb pool: %s (max=%s)", key, _POOL_MAX_SIZE)
|
||||
return pool
|
||||
|
||||
|
||||
async def probe_asyncpg_connection(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
) -> None:
|
||||
"""单次 PostgreSQL 探测(不走连接池,避免多 worker 下池缓存导致测试结果抖动)"""
|
||||
import asyncpg
|
||||
|
||||
conn = await asyncpg.connect(
|
||||
host=host,
|
||||
port=port,
|
||||
user=user,
|
||||
password=password,
|
||||
database=database or "postgres",
|
||||
timeout=10,
|
||||
)
|
||||
try:
|
||||
await conn.execute("SELECT 1")
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def probe_aiomysql_connection(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
) -> None:
|
||||
"""单次 MySQL 探测"""
|
||||
import aiomysql
|
||||
|
||||
conn = await aiomysql.connect(
|
||||
host=host,
|
||||
port=port,
|
||||
user=user,
|
||||
password=password,
|
||||
db=database or "",
|
||||
charset="utf8mb4",
|
||||
connect_timeout=10,
|
||||
)
|
||||
try:
|
||||
async with conn.cursor() as cursor:
|
||||
await cursor.execute("SELECT 1")
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
async def probe_aioodbc_connection(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""单次 SQL Server 探测"""
|
||||
import aioodbc
|
||||
|
||||
dsn = build_mssql_dsn(host, port, user, password, database or "master", extra_options)
|
||||
conn = await aioodbc.connect(dsn=dsn, timeout=10)
|
||||
try:
|
||||
async with conn.cursor() as cursor:
|
||||
await cursor.execute("SELECT 1")
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def probe_oracledb_connection(
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""单次 Oracle 探测"""
|
||||
import oracledb
|
||||
|
||||
service = database or (extra_options or {}).get("service_name") or "ORCL"
|
||||
dsn = build_oracle_dsn(host, port, service)
|
||||
conn = await oracledb.connect_async(user=user, password=password, dsn=dsn)
|
||||
try:
|
||||
async with conn.cursor() as cursor:
|
||||
await cursor.execute("SELECT 1 FROM DUAL")
|
||||
finally:
|
||||
await conn.close()
|
||||
|
||||
|
||||
async def close_all_manager_pools() -> None:
|
||||
"""关闭所有外部数据库连接池(应用 shutdown 时调用)"""
|
||||
async with _pool_lock:
|
||||
for key, pool in list(_pg_pool_cache.items()):
|
||||
try:
|
||||
await pool.close()
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to close asyncpg pool %s: %s", key, exc)
|
||||
_pg_pool_cache.clear()
|
||||
|
||||
for key, pool in list(_mysql_pool_cache.items()):
|
||||
try:
|
||||
pool.close()
|
||||
await pool.wait_closed()
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to close aiomysql pool %s: %s", key, exc)
|
||||
_mysql_pool_cache.clear()
|
||||
|
||||
for key, pool in list(_mssql_pool_cache.items()):
|
||||
try:
|
||||
pool.close()
|
||||
await pool.wait_closed()
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to close aioodbc pool %s: %s", key, exc)
|
||||
_mssql_pool_cache.clear()
|
||||
|
||||
for key, pool in list(_oracle_pool_cache.items()):
|
||||
try:
|
||||
await pool.close()
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to close oracledb pool %s: %s", key, exc)
|
||||
_oracle_pool_cache.clear()
|
||||
@@ -0,0 +1,546 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""PostgreSQL 异步处理器"""
|
||||
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_asyncpg_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 PostgreSQLHandler(HandlerTransactionMixin):
|
||||
"""PostgreSQL异步处理器"""
|
||||
|
||||
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:
|
||||
await self._pool.release(self.conn)
|
||||
except Exception as exc:
|
||||
logger.warning("Error releasing PostgreSQL connection: %s", exc)
|
||||
self.conn = None
|
||||
self._pool = None
|
||||
|
||||
async def connect(self, database: str = None) -> bool:
|
||||
if database:
|
||||
self.database = database
|
||||
db = (self.database or "").strip() or "postgres"
|
||||
try:
|
||||
await self._release_connection()
|
||||
|
||||
pool = await get_asyncpg_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="postgresql",
|
||||
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 get_databases(self) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return []
|
||||
query = """
|
||||
SELECT d.datname as name, pg_catalog.pg_get_userbyid(d.datdba) as owner,
|
||||
pg_catalog.pg_encoding_to_char(d.encoding) as encoding, d.datcollate as collation,
|
||||
pg_catalog.pg_size_pretty(pg_catalog.pg_database_size(d.datname)) as size,
|
||||
pg_catalog.pg_database_size(d.datname) as size_bytes,
|
||||
pg_catalog.shobj_description(d.oid, 'pg_database') as description
|
||||
FROM pg_catalog.pg_database d WHERE d.datistemplate = false ORDER BY d.datname
|
||||
"""
|
||||
rows = await self.conn.fetch(query)
|
||||
return [dict(row) for row in rows]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get databases: {e}")
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def create_database(self, name: str, owner: str = None, encoding: str = "UTF8", template: str = "template0", **kwargs) -> bool:
|
||||
try:
|
||||
if not await self.connect("postgres"):
|
||||
return False
|
||||
qname = quote_identifier(name, "postgresql")
|
||||
query = f"CREATE DATABASE {qname} ENCODING '{encoding}' TEMPLATE {template}"
|
||||
if owner:
|
||||
query += f" OWNER {quote_identifier(owner, 'postgresql')}"
|
||||
await self.conn.execute(query)
|
||||
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("postgres"):
|
||||
return False
|
||||
qname = quote_identifier(name, "postgresql")
|
||||
await self.conn.execute(
|
||||
f"SELECT pg_terminate_backend(pid) FROM pg_stat_activity "
|
||||
f"WHERE datname = '{name.replace(chr(39), chr(39) + chr(39))}' "
|
||||
f"AND pid <> pg_backend_pid()"
|
||||
)
|
||||
await self.conn.execute(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 rename_database(self, name: str, new_name: str) -> bool:
|
||||
try:
|
||||
if not await self.connect("postgres"):
|
||||
return False
|
||||
old_q = quote_identifier(name, "postgresql")
|
||||
new_q = quote_identifier(new_name, "postgresql")
|
||||
await self.conn.execute(f"ALTER DATABASE {old_q} RENAME TO {new_q}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to rename database {name} to {new_name}: {e}")
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def create_schema(self, name: str, database: str = None, owner: str = None) -> bool:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return False
|
||||
qname = quote_identifier(name, "postgresql")
|
||||
query = f"CREATE SCHEMA {qname}"
|
||||
if owner:
|
||||
query += f" AUTHORIZATION {quote_identifier(owner, 'postgresql')}"
|
||||
await self.conn.execute(query)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create schema {name}: {e}")
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def drop_schema(self, name: str, database: str = None, cascade: bool = True) -> bool:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return False
|
||||
qname = quote_identifier(name, "postgresql")
|
||||
cascade_clause = " CASCADE" if cascade else ""
|
||||
await self.conn.execute(f"DROP SCHEMA {qname}{cascade_clause}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to drop schema {name}: {e}")
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return False
|
||||
old_q = quote_identifier(name, "postgresql")
|
||||
new_q = quote_identifier(new_name, "postgresql")
|
||||
await self.conn.execute(f"ALTER SCHEMA {old_q} RENAME TO {new_q}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to rename schema {name} to {new_name}: {e}")
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT schema_name as name, schema_owner as owner,
|
||||
(SELECT count(*) FROM information_schema.tables WHERE table_schema = schema_name) as tables_count
|
||||
FROM information_schema.schemata
|
||||
WHERE schema_name NOT IN ('pg_catalog', 'information_schema', 'pg_toast') ORDER BY schema_name
|
||||
"""
|
||||
rows = await self.conn.fetch(query)
|
||||
return [dict(row) for row in rows]
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get schemas: {e}")
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_tables(self, database: str = None, schema_name: str = "public") -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT schemaname as schema_name, tablename as table_name, 'BASE TABLE' as table_type,
|
||||
pg_size_pretty(pg_total_relation_size(schemaname || '.' || tablename)) as total_size,
|
||||
pg_total_relation_size(schemaname || '.' || tablename) as total_size_bytes,
|
||||
pg_size_pretty(pg_relation_size(schemaname || '.' || tablename)) as table_size,
|
||||
pg_relation_size(schemaname || '.' || tablename) as table_size_bytes,
|
||||
pg_size_pretty(pg_total_relation_size(schemaname || '.' || tablename) - pg_relation_size(schemaname || '.' || tablename)) as indexes_size,
|
||||
(pg_total_relation_size(schemaname || '.' || tablename) - pg_relation_size(schemaname || '.' || tablename)) as indexes_size_bytes,
|
||||
obj_description((schemaname || '.' || tablename)::regclass) as description
|
||||
FROM pg_catalog.pg_tables WHERE schemaname = $1 ORDER BY tablename
|
||||
"""
|
||||
rows = await self.conn.fetch(query, schema_name)
|
||||
tables = [dict(row) for row in rows]
|
||||
for table in tables:
|
||||
try:
|
||||
result = await self.conn.fetchval(f'SELECT count(*) FROM "{schema_name}"."{table["table_name"]}"')
|
||||
table['row_count'] = result
|
||||
except:
|
||||
table['row_count'] = None
|
||||
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 = "public", database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT c.column_name, c.data_type, c.is_nullable = 'YES' as is_nullable, c.column_default,
|
||||
c.character_maximum_length, c.numeric_precision, c.numeric_scale, c.ordinal_position,
|
||||
EXISTS(SELECT 1 FROM information_schema.table_constraints tc JOIN information_schema.key_column_usage kcu ON tc.constraint_name = kcu.constraint_name WHERE tc.table_schema = c.table_schema AND tc.table_name = c.table_name AND kcu.column_name = c.column_name AND tc.constraint_type = 'PRIMARY KEY') as is_primary_key,
|
||||
EXISTS(SELECT 1 FROM information_schema.table_constraints tc JOIN information_schema.key_column_usage kcu ON tc.constraint_name = kcu.constraint_name WHERE tc.table_schema = c.table_schema AND tc.table_name = c.table_name AND kcu.column_name = c.column_name AND tc.constraint_type = 'UNIQUE') as is_unique,
|
||||
col_description((c.table_schema || '.' || c.table_name)::regclass, c.ordinal_position) as description
|
||||
FROM information_schema.columns c WHERE c.table_schema = $1 AND c.table_name = $2 ORDER BY c.ordinal_position
|
||||
"""
|
||||
rows = await self.conn.fetch(query, schema_name, table_name)
|
||||
return [dict(row) for row in 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 = "public", database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT i.indexname as index_name, am.amname as index_type,
|
||||
array_to_string(array_agg(a.attname ORDER BY k.ordinality), ', ') as columns,
|
||||
i.indexdef as definition, ix.indisunique as is_unique, ix.indisprimary as is_primary
|
||||
FROM pg_indexes i
|
||||
JOIN pg_class c ON c.relname = i.tablename AND c.relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname)
|
||||
JOIN pg_index ix ON ix.indexrelid = (SELECT oid FROM pg_class WHERE relname = i.indexname AND relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname))
|
||||
JOIN pg_am am ON am.oid = (SELECT relam FROM pg_class WHERE relname = i.indexname AND relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = i.schemaname))
|
||||
CROSS JOIN LATERAL unnest(ix.indkey) WITH ORDINALITY AS k(attnum, ordinality)
|
||||
JOIN pg_attribute a ON a.attrelid = c.oid AND a.attnum = k.attnum
|
||||
WHERE i.schemaname = $1 AND i.tablename = $2
|
||||
GROUP BY i.indexname, am.amname, i.indexdef, ix.indisunique, ix.indisprimary ORDER BY i.indexname
|
||||
"""
|
||||
rows = await self.conn.fetch(query, schema_name, table_name)
|
||||
return [dict(row) for row in rows]
|
||||
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 = "public", database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT
|
||||
c.conname AS constraint_name,
|
||||
CASE c.contype
|
||||
WHEN 'p' THEN 'PRIMARY KEY'
|
||||
WHEN 'f' THEN 'FOREIGN KEY'
|
||||
WHEN 'u' THEN 'UNIQUE'
|
||||
WHEN 'c' THEN 'CHECK'
|
||||
ELSE c.contype::text
|
||||
END AS constraint_type,
|
||||
(
|
||||
SELECT string_agg(a.attname, ', ' ORDER BY u.ord)
|
||||
FROM unnest(c.conkey) WITH ORDINALITY AS u(attnum, ord)
|
||||
JOIN pg_attribute a
|
||||
ON a.attrelid = c.conrelid
|
||||
AND a.attnum = u.attnum
|
||||
AND a.attnum > 0
|
||||
) AS columns,
|
||||
pg_get_constraintdef(c.oid) AS definition,
|
||||
ref_cls.relname AS referenced_table,
|
||||
(
|
||||
SELECT string_agg(a.attname, ', ' ORDER BY u.ord)
|
||||
FROM unnest(c.confkey) WITH ORDINALITY AS u(attnum, ord)
|
||||
JOIN pg_attribute a
|
||||
ON a.attrelid = c.confrelid
|
||||
AND a.attnum = u.attnum
|
||||
AND a.attnum > 0
|
||||
) AS referenced_columns
|
||||
FROM pg_constraint c
|
||||
JOIN pg_class cls ON cls.oid = c.conrelid
|
||||
JOIN pg_namespace n ON n.oid = cls.relnamespace
|
||||
LEFT JOIN pg_class ref_cls ON ref_cls.oid = c.confrelid
|
||||
WHERE n.nspname = $1
|
||||
AND cls.relname = $2
|
||||
AND c.contype IN ('p', 'f', 'u', 'c')
|
||||
ORDER BY c.contype, c.conname
|
||||
"""
|
||||
rows = await self.conn.fetch(query, schema_name, table_name)
|
||||
return [dict(row) for row in rows]
|
||||
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 = "public") -> Dict[str, Any]:
|
||||
tables = await self.get_tables(database=database, schema_name=schema_name)
|
||||
table_info = next((t for t in tables if t['table_name'] == table_name), None)
|
||||
if not table_info:
|
||||
raise ValueError(f"Table {schema_name}.{table_name} not found")
|
||||
columns = await self.get_table_columns(table_name, schema_name, database)
|
||||
indexes = await self.get_table_indexes(table_name, schema_name, database)
|
||||
constraints = await self.get_table_constraints(table_name, schema_name, database)
|
||||
return {"table_info": table_info, "columns": columns, "indexes": indexes, "constraints": constraints}
|
||||
|
||||
async def get_table_ddl(self, table_name: str, schema_name: str = "public", database: str = None) -> str:
|
||||
try:
|
||||
columns = await self.get_table_columns(table_name, schema_name, database)
|
||||
indexes = await self.get_table_indexes(table_name, schema_name, database)
|
||||
constraints = await self.get_table_constraints(table_name, schema_name, database)
|
||||
ddl_lines = [f'CREATE TABLE "{schema_name}"."{table_name}" (']
|
||||
col_defs = []
|
||||
for col in columns:
|
||||
col_def = f' "{col["column_name"]}" {col["data_type"]}'
|
||||
if col.get("character_maximum_length"):
|
||||
col_def = f' "{col["column_name"]}" {col["data_type"]}({col["character_maximum_length"]})'
|
||||
if not col["is_nullable"]:
|
||||
col_def += " NOT NULL"
|
||||
if col.get("column_default"):
|
||||
col_def += f' DEFAULT {col["column_default"]}'
|
||||
col_defs.append(col_def)
|
||||
for pk in [c for c in constraints if c["constraint_type"] == "PRIMARY KEY"]:
|
||||
col_defs.append(f' PRIMARY KEY ({pk["columns"]})')
|
||||
ddl_lines.append(',\n'.join(col_defs))
|
||||
ddl_lines.append(');')
|
||||
for idx in indexes:
|
||||
if not idx["is_primary"]:
|
||||
unique = "UNIQUE " if idx["is_unique"] else ""
|
||||
ddl_lines.append(f'\nCREATE {unique}INDEX "{idx["index_name"]}" ON "{schema_name}"."{table_name}" ({idx["columns"]});')
|
||||
return '\n'.join(ddl_lines)
|
||||
except Exception as e:
|
||||
return f"-- 获取DDL失败: {str(e)}"
|
||||
|
||||
async def get_views(self, database: str = None, schema_name: str = "public") -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT table_name as view_name, table_schema as schema_name, view_definition, is_updatable = 'YES' as is_updatable, check_option, 'VIEW' as view_type
|
||||
FROM information_schema.views WHERE table_schema = $1
|
||||
UNION ALL
|
||||
SELECT matviewname as view_name, schemaname as schema_name, definition as view_definition, false as is_updatable, NULL as check_option, 'MATERIALIZED VIEW' as view_type
|
||||
FROM pg_matviews WHERE schemaname = $1 ORDER BY view_name
|
||||
"""
|
||||
rows = await self.conn.fetch(query, schema_name)
|
||||
return [dict(row) for row in rows]
|
||||
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 = "public") -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
raise ValueError("Failed to connect")
|
||||
view_query = "SELECT table_name as view_name, table_schema as schema_name, view_definition, is_updatable = 'YES' as is_updatable, check_option, 'VIEW' as view_type FROM information_schema.views WHERE table_schema = $1 AND table_name = $2"
|
||||
view_info = await self.conn.fetchrow(view_query, schema_name, view_name)
|
||||
view_type = "VIEW"
|
||||
if not view_info:
|
||||
matview_query = "SELECT matviewname as view_name, schemaname as schema_name, definition as view_definition, false as is_updatable, NULL as check_option, 'MATERIALIZED VIEW' as view_type FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2"
|
||||
view_info = await self.conn.fetchrow(matview_query, schema_name, view_name)
|
||||
view_type = "MATERIALIZED VIEW"
|
||||
if not view_info:
|
||||
raise ValueError(f"View {schema_name}.{view_name} not found")
|
||||
view_info = dict(view_info)
|
||||
columns_query = "SELECT column_name, data_type, is_nullable = 'YES' as is_nullable, ordinal_position, col_description((table_schema || '.' || table_name)::regclass::oid, ordinal_position) as description FROM information_schema.columns WHERE table_schema = $1 AND table_name = $2 ORDER BY ordinal_position"
|
||||
columns = [dict(c) for c in await self.conn.fetch(columns_query, schema_name, view_name)]
|
||||
deps_query = "SELECT DISTINCT ref_nsp.nspname || '.' || ref_cl.relname as table_name FROM pg_depend d JOIN pg_rewrite r ON r.oid = d.objid JOIN pg_class v ON v.oid = r.ev_class JOIN pg_namespace v_nsp ON v_nsp.oid = v.relnamespace JOIN pg_class ref_cl ON ref_cl.oid = d.refobjid JOIN pg_namespace ref_nsp ON ref_nsp.oid = ref_cl.relnamespace WHERE v.relkind IN ('v', 'm') AND v_nsp.nspname = $1 AND v.relname = $2 AND ref_cl.relkind IN ('r', 'v', 'm') AND d.deptype = 'n' ORDER BY table_name"
|
||||
dependencies = [row['table_name'] for row in await self.conn.fetch(deps_query, schema_name, view_name)]
|
||||
if view_type == "MATERIALIZED VIEW":
|
||||
def_result = await self.conn.fetchval("SELECT definition FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2", schema_name, view_name)
|
||||
definition_sql = f'CREATE MATERIALIZED VIEW "{schema_name}"."{view_name}" AS\n{def_result}' if def_result else f"-- 无法获取视图定义"
|
||||
else:
|
||||
def_result = await self.conn.fetchval("SELECT pg_get_viewdef($1::regclass, true)", f"{schema_name}.{view_name}")
|
||||
definition_sql = f'CREATE OR REPLACE VIEW "{schema_name}"."{view_name}" AS\n{def_result}' if def_result else f"-- 无法获取视图定义"
|
||||
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(self, view_name: str, schema_name: str = "public") -> str:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return "-- 无法连接数据库"
|
||||
is_view = await self.conn.fetchval("SELECT 1 FROM information_schema.views WHERE table_schema = $1 AND table_name = $2", schema_name, view_name)
|
||||
if is_view:
|
||||
result = await self.conn.fetchval("SELECT pg_get_viewdef($1::regclass, true)", f"{schema_name}.{view_name}")
|
||||
definition_sql = f'CREATE OR REPLACE VIEW "{schema_name}"."{view_name}" AS\n{result}' if result else "-- 无法获取视图定义"
|
||||
else:
|
||||
result = await self.conn.fetchval("SELECT definition FROM pg_matviews WHERE schemaname = $1 AND matviewname = $2", schema_name, view_name)
|
||||
definition_sql = f'CREATE MATERIALIZED VIEW "{schema_name}"."{view_name}" AS\n{result}' if result else "-- 无法获取视图定义"
|
||||
return definition_sql
|
||||
except Exception as e:
|
||||
return f"-- 获取视图定义失败: {str(e)}"
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_dependencies(self, view_name: str, schema_name: str = "public") -> List[str]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return []
|
||||
query = "SELECT DISTINCT ref_nsp.nspname || '.' || ref_cl.relname as table_name FROM pg_depend d JOIN pg_rewrite r ON r.oid = d.objid JOIN pg_class v ON v.oid = r.ev_class JOIN pg_namespace v_nsp ON v_nsp.oid = v.relnamespace JOIN pg_class ref_cl ON ref_cl.oid = d.refobjid JOIN pg_namespace ref_nsp ON ref_nsp.oid = ref_cl.relnamespace WHERE v.relkind IN ('v', 'm') AND v_nsp.nspname = $1 AND v.relname = $2 AND ref_cl.relkind IN ('r', 'v', 'm') AND d.deptype = 'n' ORDER BY table_name"
|
||||
rows = await self.conn.fetch(query, schema_name, view_name)
|
||||
return [row['table_name'] for row in rows]
|
||||
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:
|
||||
if not await self.connect(database):
|
||||
return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size}
|
||||
schema_name = schema_name or "public"
|
||||
full_table = f'"{ schema_name}"."{table_name}"'
|
||||
count_q = f"SELECT COUNT(*) FROM {full_table}" + (f" WHERE {where}" if where else "")
|
||||
total = await self.conn.fetchval(count_q)
|
||||
offset = (page - 1) * page_size
|
||||
data_q = f"SELECT * FROM {full_table}" + (f" WHERE {where}" if where else "") + (f" ORDER BY {order_by}" if order_by else "") + f" LIMIT {page_size} OFFSET {offset}"
|
||||
rows = await self.conn.fetch(data_q)
|
||||
rows_list = [serialize_row(dict(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.conn.fetch(sql)
|
||||
|
||||
async def _run_sql_command(self, sql: str) -> int:
|
||||
result = await self.conn.execute(sql)
|
||||
if result and str(result).split()[-1].isdigit():
|
||||
return int(str(result).split()[-1])
|
||||
return 0
|
||||
|
||||
async def insert_data(self, table_name: str, data: Dict[str, Any], schema_name: str = None) -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "public"
|
||||
columns = list(data.keys())
|
||||
values = list(data.values())
|
||||
placeholders = ', '.join([f'${i+1}' for i in range(len(values))])
|
||||
query = f'INSERT INTO "{schema_name}"."{table_name}" ({", ".join(columns)}) VALUES ({placeholders})'
|
||||
await self.conn.execute(query, *values)
|
||||
return {"success": True, "message": "插入成功", "affected_rows": 1}
|
||||
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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "public"
|
||||
set_clause = ', '.join([f'{k} = ${i+1}' for i, k in enumerate(data.keys())])
|
||||
query = f'UPDATE "{schema_name}"."{table_name}" SET {set_clause} WHERE {where}'
|
||||
result = await self.conn.execute(query, *data.values())
|
||||
affected_rows = int(result.split()[-1]) if result and result.split()[-1].isdigit() else 0
|
||||
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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "public"
|
||||
query = f'DELETE FROM "{schema_name}"."{table_name}" WHERE {where}'
|
||||
result = await self.conn.execute(query)
|
||||
affected_rows = int(result.split()[-1]) if result and result.split()[-1].isdigit() else 0
|
||||
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:
|
||||
if not await self.connect(database):
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
if schema_name:
|
||||
await self.conn.execute(
|
||||
f"SET search_path TO {quote_identifier(schema_name, 'postgresql')}, public"
|
||||
)
|
||||
statements = split_sql_statements(sql)
|
||||
if not statements:
|
||||
return {"success": False, "message": "无有效SQL语句", "affected_rows": 0}
|
||||
async with self.conn.transaction():
|
||||
for index, statement in enumerate(statements, start=1):
|
||||
try:
|
||||
await self.conn.execute(statement)
|
||||
except Exception as exc:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"DDL执行失败 (第{index}条): {exc}",
|
||||
"affected_rows": 0,
|
||||
}
|
||||
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()
|
||||
|
||||
|
||||
@@ -0,0 +1,725 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""SQL Server 异步处理器(aioodbc)"""
|
||||
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_aioodbc_pool
|
||||
from core.database_manager.handlers.transaction_mixin import HandlerTransactionMixin
|
||||
from core.database_manager.sql_utils import quote_identifier, quote_table, split_sql_statements
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _normalize_mssql_default(value: Optional[str]) -> Optional[str]:
|
||||
"""Strip SQL Server INFORMATION_SCHEMA default wrappers, e.g. ((0)) -> 0."""
|
||||
if not value:
|
||||
return None
|
||||
normalized = value.strip()
|
||||
while normalized.startswith("(") and normalized.endswith(")"):
|
||||
normalized = normalized[1:-1].strip()
|
||||
if not normalized or normalized.upper() == "NULL":
|
||||
return None
|
||||
return normalized
|
||||
|
||||
|
||||
_MSSQL_SYSTEM_SCHEMAS = (
|
||||
"sys",
|
||||
"INFORMATION_SCHEMA",
|
||||
"guest",
|
||||
"db_owner",
|
||||
"db_accessadmin",
|
||||
"db_securityadmin",
|
||||
"db_ddladmin",
|
||||
"db_backupoperator",
|
||||
"db_datareader",
|
||||
"db_datawriter",
|
||||
"db_denydatareader",
|
||||
"db_denydatawriter",
|
||||
)
|
||||
|
||||
|
||||
class SQLServerHandler(HandlerTransactionMixin):
|
||||
"""SQL Server 异步处理器"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: str,
|
||||
database: str,
|
||||
extra_options: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.user = user
|
||||
self.password = password
|
||||
self.database = database or "master"
|
||||
self.extra_options = extra_options or {}
|
||||
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:
|
||||
await self._pool.release(self.conn)
|
||||
except Exception as exc:
|
||||
logger.warning("Error releasing SQL Server connection: %s", exc)
|
||||
self.conn = None
|
||||
self._pool = None
|
||||
|
||||
async def connect(self, database: str = None) -> bool:
|
||||
db = database or self.database or "master"
|
||||
try:
|
||||
await self._release_connection()
|
||||
pool = await get_aioodbc_pool(
|
||||
self.host,
|
||||
self.port,
|
||||
self.user,
|
||||
self.password,
|
||||
db,
|
||||
self.extra_options,
|
||||
)
|
||||
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="sqlserver",
|
||||
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]]:
|
||||
async with self.conn.cursor() as cursor:
|
||||
await cursor.execute(query, params or ())
|
||||
if not cursor.description:
|
||||
return []
|
||||
columns = [col[0] for col in cursor.description]
|
||||
rows = await cursor.fetchall()
|
||||
return [dict(zip(columns, 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 _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 _exec_tx_statement(self, sql: str) -> None:
|
||||
upper = sql.strip().upper()
|
||||
if upper == "BEGIN":
|
||||
await self._run_sql_command("BEGIN TRANSACTION")
|
||||
elif upper == "COMMIT":
|
||||
await self._run_sql_command("COMMIT TRANSACTION")
|
||||
elif upper == "ROLLBACK":
|
||||
await self._run_sql_command("ROLLBACK TRANSACTION")
|
||||
else:
|
||||
await self._run_sql_command(sql)
|
||||
|
||||
async def get_databases(self) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect("master"):
|
||||
return []
|
||||
query = """
|
||||
SELECT d.name,
|
||||
SUSER_SNAME(d.owner_sid) AS owner,
|
||||
d.collation_name AS collation,
|
||||
CAST(
|
||||
(SELECT SUM(CAST(mf.size AS BIGINT)) * 8192
|
||||
FROM sys.master_files mf
|
||||
WHERE mf.database_id = d.database_id) AS BIGINT
|
||||
) AS size_bytes
|
||||
FROM sys.databases d
|
||||
WHERE d.state = 0
|
||||
ORDER BY d.name
|
||||
"""
|
||||
databases = await self._execute_query(query)
|
||||
for db in databases:
|
||||
size_bytes = int(db.get("size_bytes") or 0)
|
||||
db["size"] = format_size(size_bytes)
|
||||
db["size_bytes"] = size_bytes
|
||||
return databases
|
||||
except Exception as e:
|
||||
logger.error("Failed to get databases: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def create_database(self, name: str, **kwargs) -> bool:
|
||||
try:
|
||||
if not await self.connect("master"):
|
||||
return False
|
||||
qname = quote_identifier(name, "sqlserver")
|
||||
await self._execute_command(f"CREATE DATABASE {qname}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to create database %s: %s", name, e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def drop_database(self, name: str) -> bool:
|
||||
try:
|
||||
if not await self.connect("master"):
|
||||
return False
|
||||
qname = quote_identifier(name, "sqlserver")
|
||||
await self._execute_command(
|
||||
f"ALTER DATABASE {qname} SET SINGLE_USER WITH ROLLBACK IMMEDIATE"
|
||||
)
|
||||
await self._execute_command(f"DROP DATABASE {qname}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to drop database %s: %s", name, e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def rename_database(self, name: str, new_name: str) -> bool:
|
||||
try:
|
||||
if not await self.connect("master"):
|
||||
return False
|
||||
old_q = quote_identifier(name, "sqlserver")
|
||||
new_q = quote_identifier(new_name, "sqlserver")
|
||||
await self._execute_command(f"ALTER DATABASE {old_q} MODIFY NAME = {new_q}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to rename database %s to %s: %s", name, new_name, e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def create_schema(self, name: str, database: str = None, owner: str = None) -> bool:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return False
|
||||
qname = quote_identifier(name, "sqlserver")
|
||||
query = f"CREATE SCHEMA {qname}"
|
||||
if owner:
|
||||
query += f" AUTHORIZATION {quote_identifier(owner, 'sqlserver')}"
|
||||
await self._execute_command(query)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to create schema %s: %s", name, e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def drop_schema(self, name: str, database: str = None, cascade: bool = True) -> bool:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return False
|
||||
qname = quote_identifier(name, "sqlserver")
|
||||
await self._execute_command(f"DROP SCHEMA {qname}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error("Failed to drop schema %s: %s", name, e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool:
|
||||
raise NotImplementedError("SQL Server 不支持直接重命名 Schema")
|
||||
|
||||
async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
placeholders = ", ".join(["?"] * len(_MSSQL_SYSTEM_SCHEMAS))
|
||||
query = f"""
|
||||
SELECT s.name,
|
||||
USER_NAME(s.principal_id) AS owner,
|
||||
(SELECT COUNT(*)
|
||||
FROM sys.tables t
|
||||
WHERE t.schema_id = s.schema_id) AS tables_count
|
||||
FROM sys.schemas s
|
||||
WHERE s.name NOT IN ({placeholders})
|
||||
AND s.name NOT LIKE 'db[_]%'
|
||||
ORDER BY s.name
|
||||
"""
|
||||
return await self._execute_query(query, _MSSQL_SYSTEM_SCHEMAS)
|
||||
except Exception as e:
|
||||
logger.error("Failed to get schemas: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_tables(
|
||||
self, database: str = None, schema_name: str = "dbo"
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT t.TABLE_SCHEMA AS schema_name,
|
||||
t.TABLE_NAME AS table_name,
|
||||
t.TABLE_TYPE AS table_type,
|
||||
CAST(p.rows AS BIGINT) AS row_count
|
||||
FROM INFORMATION_SCHEMA.TABLES t
|
||||
LEFT JOIN sys.tables st
|
||||
ON st.name = t.TABLE_NAME
|
||||
AND SCHEMA_NAME(st.schema_id) = t.TABLE_SCHEMA
|
||||
LEFT JOIN sys.partitions p
|
||||
ON st.object_id = p.object_id
|
||||
AND p.index_id IN (0, 1)
|
||||
WHERE t.TABLE_SCHEMA = ?
|
||||
AND t.TABLE_TYPE = 'BASE TABLE'
|
||||
ORDER BY t.TABLE_NAME
|
||||
"""
|
||||
return await self._execute_query(query, (schema_name,))
|
||||
except Exception as e:
|
||||
logger.error("Failed to get tables: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_columns(
|
||||
self, table_name: str, schema_name: str = "dbo", database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT c.COLUMN_NAME AS column_name,
|
||||
c.DATA_TYPE AS data_type,
|
||||
CASE WHEN c.IS_NULLABLE = 'YES' THEN 1 ELSE 0 END AS is_nullable,
|
||||
c.COLUMN_DEFAULT AS column_default,
|
||||
c.CHARACTER_MAXIMUM_LENGTH AS character_maximum_length,
|
||||
c.NUMERIC_PRECISION AS numeric_precision,
|
||||
c.NUMERIC_SCALE AS numeric_scale,
|
||||
c.ORDINAL_POSITION AS ordinal_position,
|
||||
CASE WHEN pk.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS is_primary_key,
|
||||
CASE WHEN uq.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS is_unique,
|
||||
ep.value AS description
|
||||
FROM INFORMATION_SCHEMA.COLUMNS c
|
||||
LEFT JOIN (
|
||||
SELECT ku.TABLE_SCHEMA, ku.TABLE_NAME, ku.COLUMN_NAME
|
||||
FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc
|
||||
JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE ku
|
||||
ON tc.CONSTRAINT_NAME = ku.CONSTRAINT_NAME
|
||||
AND tc.TABLE_SCHEMA = ku.TABLE_SCHEMA
|
||||
AND tc.TABLE_NAME = ku.TABLE_NAME
|
||||
WHERE tc.CONSTRAINT_TYPE = 'PRIMARY KEY'
|
||||
) pk ON pk.TABLE_SCHEMA = c.TABLE_SCHEMA
|
||||
AND pk.TABLE_NAME = c.TABLE_NAME
|
||||
AND pk.COLUMN_NAME = c.COLUMN_NAME
|
||||
LEFT JOIN (
|
||||
SELECT ku.TABLE_SCHEMA, ku.TABLE_NAME, ku.COLUMN_NAME
|
||||
FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc
|
||||
JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE ku
|
||||
ON tc.CONSTRAINT_NAME = ku.CONSTRAINT_NAME
|
||||
AND tc.TABLE_SCHEMA = ku.TABLE_SCHEMA
|
||||
AND tc.TABLE_NAME = ku.TABLE_NAME
|
||||
WHERE tc.CONSTRAINT_TYPE = 'UNIQUE'
|
||||
) uq ON uq.TABLE_SCHEMA = c.TABLE_SCHEMA
|
||||
AND uq.TABLE_NAME = c.TABLE_NAME
|
||||
AND uq.COLUMN_NAME = c.COLUMN_NAME
|
||||
LEFT JOIN sys.schemas ss ON ss.name = c.TABLE_SCHEMA
|
||||
LEFT JOIN sys.tables st ON st.schema_id = ss.schema_id AND st.name = c.TABLE_NAME
|
||||
LEFT JOIN sys.columns sc ON sc.object_id = st.object_id AND sc.name = c.COLUMN_NAME
|
||||
LEFT JOIN sys.extended_properties ep
|
||||
ON ep.major_id = sc.object_id
|
||||
AND ep.minor_id = sc.column_id
|
||||
AND ep.name = 'MS_Description'
|
||||
WHERE c.TABLE_SCHEMA = ? AND c.TABLE_NAME = ?
|
||||
ORDER BY c.ORDINAL_POSITION
|
||||
"""
|
||||
columns = await self._execute_query(query, (schema_name, table_name))
|
||||
for col in columns:
|
||||
col["column_default"] = _normalize_mssql_default(col.get("column_default"))
|
||||
return columns
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table columns: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_indexes(
|
||||
self, table_name: str, schema_name: str = "dbo", database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT i.name AS index_name,
|
||||
i.type_desc AS index_type,
|
||||
STUFF((
|
||||
SELECT ', ' + c.name
|
||||
FROM sys.index_columns ic
|
||||
JOIN sys.columns c
|
||||
ON ic.object_id = c.object_id
|
||||
AND ic.column_id = c.column_id
|
||||
WHERE ic.object_id = i.object_id
|
||||
AND ic.index_id = i.index_id
|
||||
ORDER BY ic.key_ordinal
|
||||
FOR XML PATH('')
|
||||
), 1, 2, '') AS columns,
|
||||
i.is_unique,
|
||||
i.is_primary_key AS is_primary,
|
||||
'' AS definition
|
||||
FROM sys.indexes i
|
||||
JOIN sys.tables t ON i.object_id = t.object_id
|
||||
JOIN sys.schemas s ON t.schema_id = s.schema_id
|
||||
WHERE s.name = ? AND t.name = ? AND i.name IS NOT NULL
|
||||
ORDER BY i.name
|
||||
"""
|
||||
indexes = await self._execute_query(query, (schema_name, table_name))
|
||||
for idx in indexes:
|
||||
unique = "UNIQUE " if idx.get("is_unique") else ""
|
||||
idx["definition"] = f"CREATE {unique}INDEX {idx['index_name']} ON {schema_name}.{table_name} ({idx.get('columns', '')})"
|
||||
return indexes
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table indexes: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_constraints(
|
||||
self, table_name: str, schema_name: str = "dbo", database: str = None
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT tc.CONSTRAINT_NAME AS constraint_name,
|
||||
tc.CONSTRAINT_TYPE AS constraint_type,
|
||||
STUFF((
|
||||
SELECT ', ' + kcu.COLUMN_NAME
|
||||
FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE kcu
|
||||
WHERE kcu.CONSTRAINT_NAME = tc.CONSTRAINT_NAME
|
||||
AND kcu.TABLE_SCHEMA = tc.TABLE_SCHEMA
|
||||
AND kcu.TABLE_NAME = tc.TABLE_NAME
|
||||
FOR XML PATH('')
|
||||
), 1, 2, '') AS columns,
|
||||
kcu2.TABLE_NAME AS referenced_table,
|
||||
STUFF((
|
||||
SELECT ', ' + fk.COLUMN_NAME
|
||||
FROM INFORMATION_SCHEMA.KEY_COLUMN_USAGE fk
|
||||
WHERE fk.CONSTRAINT_NAME = tc.CONSTRAINT_NAME
|
||||
AND fk.TABLE_SCHEMA = tc.TABLE_SCHEMA
|
||||
FOR XML PATH('')
|
||||
), 1, 2, '') AS referenced_columns,
|
||||
'' AS definition
|
||||
FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS tc
|
||||
LEFT JOIN INFORMATION_SCHEMA.REFERENTIAL_CONSTRAINTS rc
|
||||
ON tc.CONSTRAINT_NAME = rc.CONSTRAINT_NAME
|
||||
AND tc.TABLE_SCHEMA = rc.CONSTRAINT_SCHEMA
|
||||
LEFT JOIN INFORMATION_SCHEMA.KEY_COLUMN_USAGE kcu2
|
||||
ON rc.UNIQUE_CONSTRAINT_NAME = kcu2.CONSTRAINT_NAME
|
||||
WHERE tc.TABLE_SCHEMA = ? AND tc.TABLE_NAME = ?
|
||||
ORDER BY tc.CONSTRAINT_TYPE, tc.CONSTRAINT_NAME
|
||||
"""
|
||||
constraints = await self._execute_query(query, (schema_name, table_name))
|
||||
for const in constraints:
|
||||
const["definition"] = f"{const['constraint_type']} ({const.get('columns', '')})"
|
||||
return constraints
|
||||
except Exception as e:
|
||||
logger.error("Failed to get table constraints: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_table_structure(
|
||||
self, table_name: str, database: str = None, schema_name: str = "dbo"
|
||||
) -> Dict[str, Any]:
|
||||
tables = await self.get_tables(database=database, schema_name=schema_name)
|
||||
table_info = next((t for t in tables if t["table_name"] == table_name), None)
|
||||
if not table_info:
|
||||
raise ValueError(f"Table {schema_name}.{table_name} not found")
|
||||
columns = await self.get_table_columns(table_name, schema_name, database)
|
||||
indexes = await self.get_table_indexes(table_name, schema_name, database)
|
||||
constraints = await self.get_table_constraints(table_name, schema_name, database)
|
||||
return {
|
||||
"table_info": table_info,
|
||||
"columns": columns,
|
||||
"indexes": indexes,
|
||||
"constraints": constraints,
|
||||
}
|
||||
|
||||
async def get_table_ddl(
|
||||
self, table_name: str, schema_name: str = "dbo", database: str = None
|
||||
) -> str:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return "-- 无法连接数据库"
|
||||
full_name = quote_table(schema_name, table_name, "sqlserver")
|
||||
rows = await self._execute_query(
|
||||
"SELECT OBJECT_DEFINITION(OBJECT_ID(?)) AS ddl",
|
||||
(f"{schema_name}.{table_name}",),
|
||||
)
|
||||
if rows and rows[0].get("ddl"):
|
||||
return rows[0]["ddl"]
|
||||
columns = await self.get_table_columns(table_name, schema_name, database)
|
||||
if not columns:
|
||||
return f"-- 无法获取表 {full_name} 的DDL"
|
||||
ddl_lines = [f"CREATE TABLE {full_name} ("]
|
||||
col_defs = []
|
||||
for col in columns:
|
||||
col_def = f" {quote_identifier(col['column_name'], 'sqlserver')} {col['data_type']}"
|
||||
if col.get("character_maximum_length"):
|
||||
col_def += f"({col['character_maximum_length']})"
|
||||
if not col.get("is_nullable"):
|
||||
col_def += " NOT NULL"
|
||||
col_defs.append(col_def)
|
||||
ddl_lines.append(",\n".join(col_defs))
|
||||
ddl_lines.append(");")
|
||||
return "\n".join(ddl_lines)
|
||||
except Exception as e:
|
||||
return f"-- 获取DDL失败: {str(e)}"
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_views(
|
||||
self, database: str = None, schema_name: str = "dbo"
|
||||
) -> List[Dict[str, Any]]:
|
||||
try:
|
||||
if not await self.connect(database):
|
||||
return []
|
||||
query = """
|
||||
SELECT TABLE_NAME AS view_name,
|
||||
TABLE_SCHEMA AS schema_name,
|
||||
VIEW_DEFINITION AS view_definition,
|
||||
CASE WHEN IS_UPDATABLE = 'YES' THEN 1 ELSE 0 END AS is_updatable,
|
||||
CHECK_OPTION AS check_option,
|
||||
'VIEW' AS view_type
|
||||
FROM INFORMATION_SCHEMA.VIEWS
|
||||
WHERE TABLE_SCHEMA = ?
|
||||
ORDER BY TABLE_NAME
|
||||
"""
|
||||
return await self._execute_query(query, (schema_name,))
|
||||
except Exception as e:
|
||||
logger.error("Failed to get views: %s", e)
|
||||
return []
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_structure(self, view_name: str, schema_name: str = "dbo") -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
raise ValueError("Failed to connect")
|
||||
view_query = """
|
||||
SELECT TABLE_NAME AS view_name, TABLE_SCHEMA AS schema_name,
|
||||
VIEW_DEFINITION AS view_definition,
|
||||
CASE WHEN IS_UPDATABLE = 'YES' THEN 1 ELSE 0 END AS is_updatable,
|
||||
CHECK_OPTION AS check_option, 'VIEW' AS view_type
|
||||
FROM INFORMATION_SCHEMA.VIEWS
|
||||
WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ?
|
||||
"""
|
||||
view_info = await self._execute_query(view_query, (schema_name, view_name))
|
||||
if not view_info:
|
||||
raise ValueError(f"View {schema_name}.{view_name} not found")
|
||||
view_info = view_info[0]
|
||||
columns = await self.get_table_columns(view_name, schema_name)
|
||||
definition_sql = await self.get_view_definition(view_name, schema_name)
|
||||
dependencies = await self.get_view_dependencies(view_name, schema_name)
|
||||
return {
|
||||
"view_info": view_info,
|
||||
"columns": columns,
|
||||
"dependencies": dependencies,
|
||||
"definition_sql": definition_sql,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error("Failed to get view structure: %s", e)
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_definition(self, view_name: str, schema_name: str = "dbo") -> str:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return "-- 无法连接数据库"
|
||||
rows = await self._execute_query(
|
||||
"""
|
||||
SELECT VIEW_DEFINITION AS definition
|
||||
FROM INFORMATION_SCHEMA.VIEWS
|
||||
WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ?
|
||||
""",
|
||||
(schema_name, view_name),
|
||||
)
|
||||
if rows and rows[0].get("definition"):
|
||||
full = quote_table(schema_name, view_name, "sqlserver")
|
||||
return f"CREATE VIEW {full} AS\n{rows[0]['definition']}"
|
||||
return f"-- 无法获取视图 {schema_name}.{view_name} 的定义"
|
||||
except Exception as e:
|
||||
return f"-- 获取视图定义失败: {str(e)}"
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def get_view_dependencies(self, view_name: str, schema_name: str = "dbo") -> List[str]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return []
|
||||
rows = await self._execute_query(
|
||||
"""
|
||||
SELECT DISTINCT REFERENCED_ENTITY_NAME AS table_name
|
||||
FROM sys.sql_expression_dependencies d
|
||||
JOIN sys.views v ON d.referencing_id = v.object_id
|
||||
JOIN sys.schemas s ON v.schema_id = s.schema_id
|
||||
WHERE s.name = ? AND v.name = ?
|
||||
AND d.referenced_entity_name IS NOT NULL
|
||||
ORDER BY REFERENCED_ENTITY_NAME
|
||||
""",
|
||||
(schema_name, view_name),
|
||||
)
|
||||
return [row["table_name"] for row in rows]
|
||||
except Exception as e:
|
||||
logger.error("Failed to get view dependencies: %s", 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:
|
||||
if not await self.connect(database):
|
||||
return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size}
|
||||
schema_name = schema_name or "dbo"
|
||||
full_table = quote_table(schema_name, table_name, "sqlserver")
|
||||
count_q = f"SELECT COUNT(*) AS total FROM {full_table}"
|
||||
if where:
|
||||
count_q += f" WHERE {where}"
|
||||
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 {full_table}"
|
||||
if where:
|
||||
data_q += f" WHERE {where}"
|
||||
if order_by:
|
||||
data_q += f" ORDER BY {order_by}"
|
||||
data_q += f" OFFSET {offset} ROWS FETCH NEXT {page_size} ROWS ONLY"
|
||||
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("Failed to query data: %s", e)
|
||||
return {"columns": [], "rows": [], "total": 0, "page": page, "page_size": page_size}
|
||||
finally:
|
||||
await self.close()
|
||||
|
||||
async def insert_data(
|
||||
self, table_name: str, data: Dict[str, Any], schema_name: str = None
|
||||
) -> Dict[str, Any]:
|
||||
try:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "dbo"
|
||||
full_table = quote_table(schema_name, table_name, "sqlserver")
|
||||
columns = list(data.keys())
|
||||
quoted_columns = ", ".join(
|
||||
quote_identifier(col, "sqlserver") for col in columns
|
||||
)
|
||||
placeholders = ", ".join(["?"] * len(columns))
|
||||
query = f"INSERT INTO {full_table} ({quoted_columns}) VALUES ({placeholders})"
|
||||
affected_rows = await self._execute_command(query, tuple(data.values()))
|
||||
return {"success": True, "message": "插入成功", "affected_rows": affected_rows or 1}
|
||||
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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "dbo"
|
||||
full_table = quote_table(schema_name, table_name, "sqlserver")
|
||||
set_clause = ", ".join(
|
||||
f"{quote_identifier(k, 'sqlserver')} = ?" for k in data.keys()
|
||||
)
|
||||
query = f"UPDATE {full_table} 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:
|
||||
if not await self.connect():
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
schema_name = schema_name or "dbo"
|
||||
full_table = quote_table(schema_name, table_name, "sqlserver")
|
||||
query = f"DELETE FROM {full_table} 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:
|
||||
if not await self.connect(database):
|
||||
return {"success": False, "message": "数据库连接失败", "affected_rows": 0}
|
||||
statements = split_sql_statements(sql)
|
||||
if not statements:
|
||||
return {"success": False, "message": "无有效SQL语句", "affected_rows": 0}
|
||||
for index, statement in enumerate(statements, start=1):
|
||||
try:
|
||||
await self._execute_command(statement)
|
||||
except Exception as exc:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"DDL执行失败 (第{index}条): {exc}",
|
||||
"affected_rows": 0,
|
||||
}
|
||||
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()
|
||||
@@ -0,0 +1,121 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Handler 事务支持(表单业务库与 default 行为对齐)"""
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Dict
|
||||
|
||||
from core.database_manager.handlers.common import serialize_row
|
||||
|
||||
_SAVEPOINT_NAME_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
|
||||
|
||||
def validate_savepoint_name(name: str) -> str:
|
||||
if not _SAVEPOINT_NAME_RE.match(name or ""):
|
||||
raise ValueError(f"Invalid savepoint name: {name}")
|
||||
return name
|
||||
|
||||
|
||||
class HandlerTransactionMixin:
|
||||
"""在 execute_sql 上叠加请求级事务;子类需实现 _run_sql_command / _run_sql_query。"""
|
||||
|
||||
_in_transaction: bool = False
|
||||
|
||||
async def _run_sql_command(self, sql: str) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
async def _run_sql_query(self, sql: str) -> list:
|
||||
raise NotImplementedError
|
||||
|
||||
async def _exec_tx_statement(self, sql: str) -> None:
|
||||
await self._run_sql_command(sql)
|
||||
|
||||
async def begin_transaction(self) -> bool:
|
||||
if self._in_transaction:
|
||||
return True
|
||||
if not await self.connect():
|
||||
return False
|
||||
await self._exec_tx_statement("BEGIN")
|
||||
self._in_transaction = True
|
||||
return True
|
||||
|
||||
async def commit_transaction(self) -> None:
|
||||
if not self._in_transaction:
|
||||
return
|
||||
try:
|
||||
await self._exec_tx_statement("COMMIT")
|
||||
finally:
|
||||
self._in_transaction = False
|
||||
await self.close()
|
||||
|
||||
async def rollback_transaction(self) -> None:
|
||||
if not self._in_transaction:
|
||||
return
|
||||
try:
|
||||
await self._exec_tx_statement("ROLLBACK")
|
||||
finally:
|
||||
self._in_transaction = False
|
||||
await self.close()
|
||||
|
||||
async def create_savepoint(self, name: str) -> None:
|
||||
sp = validate_savepoint_name(name)
|
||||
await self._exec_tx_statement(f"SAVEPOINT {sp}")
|
||||
|
||||
async def release_savepoint(self, name: str) -> None:
|
||||
sp = validate_savepoint_name(name)
|
||||
await self._exec_tx_statement(f"RELEASE SAVEPOINT {sp}")
|
||||
|
||||
async def rollback_to_savepoint(self, name: str) -> None:
|
||||
sp = validate_savepoint_name(name)
|
||||
await self._exec_tx_statement(f"ROLLBACK TO SAVEPOINT {sp}")
|
||||
|
||||
async def _execute_sql_impl(self, sql: str, is_query: bool = True) -> Dict[str, Any]:
|
||||
start_time = time.time()
|
||||
try:
|
||||
if not self.conn and not await self.connect():
|
||||
return {
|
||||
"success": False,
|
||||
"message": "数据库连接失败",
|
||||
"columns": None,
|
||||
"rows": None,
|
||||
"affected_rows": None,
|
||||
"execution_time": 0,
|
||||
}
|
||||
if is_query:
|
||||
rows = await self._run_sql_query(sql)
|
||||
execution_time = time.time() - start_time
|
||||
rows_list = [serialize_row(dict(row) if not isinstance(row, dict) else row) for row in rows]
|
||||
columns = list(rows_list[0].keys()) if rows_list else []
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"查询成功,返回 {len(rows_list)} 行",
|
||||
"columns": columns,
|
||||
"rows": rows_list,
|
||||
"affected_rows": None,
|
||||
"execution_time": round(execution_time, 3),
|
||||
}
|
||||
affected_rows = await self._run_sql_command(sql)
|
||||
execution_time = time.time() - start_time
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"执行成功,影响 {affected_rows} 行",
|
||||
"columns": None,
|
||||
"rows": None,
|
||||
"affected_rows": affected_rows,
|
||||
"execution_time": round(execution_time, 3),
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"success": False,
|
||||
"message": str(e),
|
||||
"columns": None,
|
||||
"rows": None,
|
||||
"affected_rows": None,
|
||||
"execution_time": round(time.time() - start_time, 3),
|
||||
}
|
||||
|
||||
async def execute_sql(self, sql: str, is_query: bool = True) -> Dict[str, Any]:
|
||||
result = await self._execute_sql_impl(sql, is_query)
|
||||
if not self._in_transaction:
|
||||
await self.close()
|
||||
return result
|
||||
Reference in New Issue
Block a user