Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
@@ -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 返回虚拟 databaseservice 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