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,5 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
数据库管理模块
"""
@@ -0,0 +1,481 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
数据库管理API(异步版本)
"""
import logging
from typing import List
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from core.database_manager.schema import (
ColumnInfo,
ConstraintInfo,
DatabaseConfig,
DatabaseCreateIn,
DatabaseInfo,
DatabaseOperationOut,
DatabaseRenameIn,
DataOperationOut,
DeleteDataIn,
ExecuteDDLIn,
ExecuteDDLOut,
ExecuteSQLIn,
ExecuteSQLOut,
IndexInfo,
InsertDataIn,
QueryDataIn,
QueryDataOut,
SchemaCreateIn,
SchemaInfo,
SchemaOperationOut,
SchemaRenameIn,
TableInfo,
TableStructure,
UpdateDataIn,
ViewInfo,
ViewStructure,
)
from core.database_manager.service import (
AsyncDatabaseManagerService,
DatabaseManagerForbidden,
DatabaseManagerValidationError,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/database_manager", tags=["数据库管理"])
async def _get_service(db_name: str, db: AsyncSession) -> AsyncDatabaseManagerService:
try:
return await AsyncDatabaseManagerService.create(db_name, db)
except ValueError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
def _handle_manager_error(exc: Exception) -> None:
if isinstance(exc, DatabaseManagerForbidden):
raise HTTPException(status_code=403, detail=str(exc)) from exc
if isinstance(exc, DatabaseManagerValidationError):
raise HTTPException(status_code=400, detail=str(exc)) from exc
raise exc
# ============ 数据库配置 ============
@router.get("/configs", response_model=List[DatabaseConfig], summary="获取数据库配置列表")
async def get_database_configs(db: AsyncSession = Depends(get_db)):
"""获取所有配置的数据库连接信息"""
return await AsyncDatabaseManagerService.get_database_configs_async(db)
@router.post("/{db_name}/test", summary="测试数据库连接")
async def test_database_connection(db_name: str, db: AsyncSession = Depends(get_db)):
"""测试指定数据库的连接"""
service = await _get_service(db_name, db)
return await service.test_connection()
# ============ 数据库管理 ============
@router.get("/{db_name}/databases", response_model=List[DatabaseInfo], summary="获取数据库列表")
async def get_databases(db_name: str, db: AsyncSession = Depends(get_db)):
"""获取所有数据库"""
service = await _get_service(db_name, db)
return await service.get_databases()
@router.post("/{db_name}/databases", response_model=DatabaseOperationOut, summary="创建数据库")
async def create_database(
db_name: str,
data: DatabaseCreateIn,
db: AsyncSession = Depends(get_db),
):
"""创建新数据库"""
try:
service = await _get_service(db_name, db)
success = await service.create_database(
name=data.name,
owner=data.owner,
encoding=data.encoding,
template=data.template,
charset=data.charset,
collation=data.collation,
)
return DatabaseOperationOut(
success=success,
message="数据库创建成功" if success else "数据库创建失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to create database: {e}")
return DatabaseOperationOut(success=False, message=str(e))
@router.patch(
"/{db_name}/databases/{database_name}",
response_model=DatabaseOperationOut,
summary="重命名数据库",
)
async def rename_database(
db_name: str,
database_name: str,
data: DatabaseRenameIn,
db: AsyncSession = Depends(get_db),
):
"""重命名数据库(PostgreSQL / SQL Server"""
try:
service = await _get_service(db_name, db)
success = await service.rename_database(database_name, data.new_name)
return DatabaseOperationOut(
success=success,
message="数据库重命名成功" if success else "数据库重命名失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to rename database: {e}")
return DatabaseOperationOut(success=False, message=str(e))
@router.delete("/{db_name}/databases/{database_name}", response_model=DatabaseOperationOut, summary="删除数据库")
async def drop_database(
db_name: str,
database_name: str,
db: AsyncSession = Depends(get_db),
):
"""删除数据库"""
try:
service = await _get_service(db_name, db)
success = await service.drop_database(database_name)
return DatabaseOperationOut(
success=success,
message="数据库删除成功" if success else "数据库删除失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to drop database: {e}")
return DatabaseOperationOut(success=False, message=str(e))
# ============ Schema管理(PostgreSQL ============
@router.get("/{db_name}/schemas", response_model=List[SchemaInfo], summary="获取Schema列表")
async def get_schemas(
db_name: str,
database: str = None,
db: AsyncSession = Depends(get_db),
):
"""获取所有Schema(仅PostgreSQL"""
service = await _get_service(db_name, db)
return await service.get_schemas(database)
@router.post("/{db_name}/schemas", response_model=SchemaOperationOut, summary="创建Schema")
async def create_schema(
db_name: str,
data: SchemaCreateIn,
db: AsyncSession = Depends(get_db),
):
"""创建 SchemaPostgreSQL / SQL Server"""
try:
service = await _get_service(db_name, db)
success = await service.create_schema(data.name, data.database, data.owner)
return SchemaOperationOut(
success=success,
message="Schema 创建成功" if success else "Schema 创建失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to create schema: {e}")
return SchemaOperationOut(success=False, message=str(e))
@router.delete(
"/{db_name}/schemas/{schema_name}",
response_model=SchemaOperationOut,
summary="删除Schema",
)
async def drop_schema(
db_name: str,
schema_name: str,
database: str = Query(None, description="数据库名"),
cascade: bool = Query(True, description="是否 CASCADE"),
db: AsyncSession = Depends(get_db),
):
"""删除 SchemaPostgreSQL / SQL Server"""
try:
service = await _get_service(db_name, db)
success = await service.drop_schema(schema_name, database, cascade)
return SchemaOperationOut(
success=success,
message="Schema 删除成功" if success else "Schema 删除失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to drop schema: {e}")
return SchemaOperationOut(success=False, message=str(e))
@router.patch(
"/{db_name}/schemas/{schema_name}",
response_model=SchemaOperationOut,
summary="重命名Schema",
)
async def rename_schema(
db_name: str,
schema_name: str,
data: SchemaRenameIn,
db: AsyncSession = Depends(get_db),
):
"""重命名 Schema(仅 PostgreSQL"""
try:
service = await _get_service(db_name, db)
success = await service.rename_schema(schema_name, data.new_name, data.database)
return SchemaOperationOut(
success=success,
message="Schema 重命名成功" if success else "Schema 重命名失败",
)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
except Exception as e:
logger.error(f"Failed to rename schema: {e}")
return SchemaOperationOut(success=False, message=str(e))
# ============ 表管理 ============
@router.get("/{db_name}/tables", response_model=List[TableInfo], summary="获取表列表")
async def get_tables(
db_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取指定schema/database的所有表"""
service = await _get_service(db_name, db)
return await service.get_tables(database=database, schema_name=schema_name)
@router.get("/{db_name}/tables/{table_name}/structure", response_model=TableStructure, summary="获取表结构")
async def get_table_structure(
db_name: str,
table_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取表的详细结构"""
service = await _get_service(db_name, db)
try:
return await service.get_table_structure(
table_name, database=database, schema_name=schema_name
)
except ValueError as e:
raise HTTPException(status_code=404, detail=str(e)) from e
@router.get("/{db_name}/tables/{table_name}/ddl", summary="获取表DDL")
async def get_table_ddl(
db_name: str,
table_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取表的DDL语句"""
service = await _get_service(db_name, db)
ddl = await service.get_table_ddl(table_name, schema_name, database)
return {"ddl": ddl}
@router.get("/{db_name}/tables/{table_name}/columns", response_model=List[ColumnInfo], summary="获取表字段")
async def get_table_columns(
db_name: str,
table_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取表的字段信息"""
service = await _get_service(db_name, db)
return await service.get_table_columns(table_name, schema_name, database)
@router.get("/{db_name}/tables/{table_name}/indexes", response_model=List[IndexInfo], summary="获取表索引")
async def get_table_indexes(
db_name: str,
table_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取表的索引信息"""
service = await _get_service(db_name, db)
return await service.get_table_indexes(table_name, schema_name, database)
@router.get("/{db_name}/tables/{table_name}/constraints", response_model=List[ConstraintInfo], summary="获取表约束")
async def get_table_constraints(
db_name: str,
table_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取表的约束信息"""
service = await _get_service(db_name, db)
return await service.get_table_constraints(table_name, schema_name, database)
# ============ 视图管理 ============
@router.get("/{db_name}/views", response_model=List[ViewInfo], summary="获取视图列表")
async def get_views(
db_name: str,
database: str = Query(None, description="数据库名"),
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取指定schema/database的所有视图"""
service = await _get_service(db_name, db)
return await service.get_views(database=database, schema_name=schema_name)
@router.get("/{db_name}/views/{view_name}/structure", response_model=ViewStructure, summary="获取视图结构")
async def get_view_structure(
db_name: str,
view_name: str,
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取视图的详细结构"""
service = await _get_service(db_name, db)
try:
return await service.get_view_structure(view_name, schema_name)
except ValueError as e:
raise HTTPException(status_code=404, detail=str(e)) from e
@router.get("/{db_name}/views/{view_name}/definition", summary="获取视图定义")
async def get_view_definition(
db_name: str,
view_name: str,
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取视图的定义SQL"""
service = await _get_service(db_name, db)
definition = await service.get_view_definition(view_name, schema_name)
return {"definition": definition}
@router.get("/{db_name}/views/{view_name}/dependencies", response_model=List[str], summary="获取视图依赖")
async def get_view_dependencies(
db_name: str,
view_name: str,
schema_name: str = Query(None, description="Schema名"),
db: AsyncSession = Depends(get_db),
):
"""获取视图依赖的表列表"""
service = await _get_service(db_name, db)
return await service.get_view_dependencies(view_name, schema_name)
# ============ 数据查询 ============
@router.post("/{db_name}/query", response_model=QueryDataOut, summary="查询表数据")
async def query_data(
db_name: str,
data: QueryDataIn,
db: AsyncSession = Depends(get_db),
):
"""分页查询表数据"""
service = await _get_service(db_name, db)
return await service.query_data(
table_name=data.table_name,
schema_name=data.schema_name,
database=data.database,
page=data.page,
page_size=data.page_size,
where=data.where,
order_by=data.order_by,
)
# ============ SQL执行 ============
@router.post("/{db_name}/execute", response_model=ExecuteSQLOut, summary="执行SQL")
async def execute_sql(
db_name: str,
data: ExecuteSQLIn,
db: AsyncSession = Depends(get_db),
):
"""执行自定义SQL语句"""
service = await _get_service(db_name, db)
return await service.execute_sql(data.sql, data.is_query)
# ============ 数据操作 ============
@router.post("/{db_name}/data/insert", response_model=DataOperationOut, summary="插入数据")
async def insert_data(
db_name: str,
data: InsertDataIn,
db: AsyncSession = Depends(get_db),
):
"""向表中插入数据"""
service = await _get_service(db_name, db)
return await service.insert_data(data.table_name, data.data, data.schema_name)
@router.post("/{db_name}/data/update", response_model=DataOperationOut, summary="更新数据")
async def update_data(
db_name: str,
data: UpdateDataIn,
db: AsyncSession = Depends(get_db),
):
"""更新表中的数据"""
service = await _get_service(db_name, db)
return await service.update_data(
data.table_name, data.data, data.where, data.schema_name
)
@router.post("/{db_name}/data/delete", response_model=DataOperationOut, summary="删除数据")
async def delete_data(
db_name: str,
data: DeleteDataIn,
db: AsyncSession = Depends(get_db),
):
"""删除表中的数据"""
service = await _get_service(db_name, db)
return await service.delete_data(data.table_name, data.where, data.schema_name)
# ============ DDL操作 ============
@router.post("/{db_name}/execute/ddl", response_model=ExecuteDDLOut, summary="执行DDL语句")
async def execute_ddl(
db_name: str,
data: ExecuteDDLIn,
db: AsyncSession = Depends(get_db),
):
"""执行DDL语句(CREATE TABLE, ALTER TABLE, DROP TABLE等)"""
try:
service = await _get_service(db_name, db)
result = await service.execute_ddl(data.sql, data.database, data.schema_name)
return ExecuteDDLOut(**result)
except HTTPException:
raise
except (DatabaseManagerForbidden, DatabaseManagerValidationError) as e:
_handle_manager_error(e)
@@ -0,0 +1,491 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
DDL 生成器:canonical 字段类型 → 多方言 CREATE TABLE / SCHEMA / INDEX / COMMENT
与前端 web/apps/web-ele/src/utils/database-types.ts 的 mapToDbType 规则对齐。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
from core.database_manager.sql_utils import quote_identifier, quote_table
# 表单工作流系统字段(设计节点与 DDL 生成共用)
SYSTEM_FIELDS: List[Dict[str, Any]] = [
{
'name': 'id',
'type': 'varchar',
'maxLength': 36,
'comment': '主键ID',
'nullable': False,
'isPrimaryKey': True,
},
{
'name': 'sys_create_datetime',
'type': 'datetime',
'comment': '创建时间',
'nullable': True,
'isPrimaryKey': False,
},
{
'name': 'sys_update_datetime',
'type': 'datetime',
'comment': '更新时间',
'nullable': True,
'isPrimaryKey': False,
},
{
'name': 'sys_creator_id',
'type': 'varchar',
'maxLength': 36,
'comment': '创建人ID',
'nullable': True,
'isPrimaryKey': False,
},
{
'name': 'sys_modifier_id',
'type': 'varchar',
'maxLength': 36,
'comment': '修改人ID',
'nullable': True,
'isPrimaryKey': False,
},
{
'name': 'sys_dept_id',
'type': 'varchar',
'maxLength': 36,
'comment': '部门ID',
'nullable': True,
'isPrimaryKey': False,
},
{
'name': 'is_deleted',
'type': 'boolean',
'comment': '是否删除',
'nullable': False,
'isPrimaryKey': False,
},
{
'name': 'sort',
'type': 'int',
'comment': '排序',
'nullable': False,
'isPrimaryKey': False,
},
]
SYSTEM_INDEX_FIELDS = frozenset({
'sys_update_datetime',
'sys_create_datetime',
'sys_creator_id',
'sys_dept_id',
'is_deleted',
})
CANONICAL_TYPE_MAPPING: Dict[str, str] = {
'string': 'varchar',
'str': 'varchar',
'varchar': 'varchar',
'char': 'char',
'text': 'text',
'int': 'int',
'integer': 'int',
'bigint': 'bigint',
'smallint': 'smallint',
'decimal': 'decimal',
'numeric': 'decimal',
'float': 'float',
'double': 'double',
'datetime': 'datetime',
'timestamp': 'datetime',
'date': 'date',
'time': 'time',
'boolean': 'boolean',
'bool': 'boolean',
'json': 'json',
'jsonb': 'jsonb',
}
COMMON_TYPE_TO_DB_TYPE: Dict[str, Dict[str, str]] = {
'postgresql': {
'int': 'INTEGER',
'bigint': 'BIGINT',
'smallint': 'SMALLINT',
'float': 'REAL',
'double': 'DOUBLE PRECISION',
'datetime': 'TIMESTAMP',
'boolean': 'BOOLEAN',
'json': 'JSON',
'jsonb': 'JSONB',
'varchar': 'VARCHAR',
'char': 'CHAR',
'text': 'TEXT',
'decimal': 'DECIMAL',
'numeric': 'NUMERIC',
'date': 'DATE',
'time': 'TIME',
},
'mysql': {
'int': 'INT',
'bigint': 'BIGINT',
'smallint': 'SMALLINT',
'float': 'FLOAT',
'double': 'DOUBLE',
'datetime': 'DATETIME',
'boolean': 'TINYINT(1)',
'json': 'JSON',
'jsonb': 'JSON',
'varchar': 'VARCHAR',
'char': 'CHAR',
'text': 'TEXT',
'decimal': 'DECIMAL',
'numeric': 'DECIMAL',
'date': 'DATE',
'time': 'TIME',
},
'sqlserver': {
'int': 'INT',
'bigint': 'BIGINT',
'smallint': 'SMALLINT',
'float': 'FLOAT',
'double': 'FLOAT',
'datetime': 'DATETIME2',
'boolean': 'BIT',
'json': 'NVARCHAR(MAX)',
'jsonb': 'NVARCHAR(MAX)',
'varchar': 'NVARCHAR',
'char': 'NCHAR',
'text': 'NVARCHAR(MAX)',
'decimal': 'DECIMAL',
'numeric': 'NUMERIC',
'date': 'DATE',
'time': 'TIME',
},
'oracle': {
'int': 'NUMBER',
'bigint': 'NUMBER',
'smallint': 'NUMBER',
'float': 'BINARY_FLOAT',
'double': 'BINARY_DOUBLE',
'datetime': 'TIMESTAMP',
'boolean': 'NUMBER(1)',
'json': 'CLOB',
'jsonb': 'CLOB',
'varchar': 'VARCHAR2',
'char': 'CHAR',
'text': 'CLOB',
'decimal': 'NUMBER',
'numeric': 'NUMBER',
'date': 'DATE',
'time': 'TIMESTAMP',
},
}
def normalize_db_type(db_type: str) -> str:
db = (db_type or 'postgresql').lower().strip()
if db in ('mssql', 'sql server'):
return 'sqlserver'
if db in ('postgres', 'psql'):
return 'postgresql'
return db
def normalize_canonical_type(field_type: str) -> str:
if not field_type:
return 'varchar'
return CANONICAL_TYPE_MAPPING.get(field_type.lower().strip(), field_type.lower().strip())
def process_canonical_field(field: dict) -> dict:
"""标准化单个 canonical 字段配置(设计节点使用)。"""
field_name = field.get('name', '')
standard_type = normalize_canonical_type(field.get('type', 'varchar'))
processed: Dict[str, Any] = {
'name': field_name,
'type': standard_type,
'comment': field.get('comment', field_name),
'nullable': field.get('nullable', True),
'isPrimaryKey': field.get('isPrimaryKey', False),
}
if standard_type in ('varchar', 'char'):
processed['maxLength'] = field.get('maxLength', 255)
elif standard_type in ('decimal', 'numeric'):
processed['precision'] = field.get('precision', 10)
processed['scale'] = field.get('scale', 2)
return processed
def map_to_dialect_type(field: dict, db_type: str) -> str:
"""将 canonical 字段映射为带长度/精度的方言类型 SQL 片段。"""
db = normalize_db_type(db_type)
field_type = normalize_canonical_type(field.get('type', 'varchar'))
type_map = COMMON_TYPE_TO_DB_TYPE.get(db, COMMON_TYPE_TO_DB_TYPE['postgresql'])
base_type = type_map.get(field_type, field_type.upper())
if field_type == 'varchar' and base_type in ('VARCHAR', 'NVARCHAR', 'VARCHAR2'):
length = field.get('maxLength', 255)
return f'{base_type}({length})'
if field_type == 'char' and base_type in ('CHAR', 'NCHAR'):
length = field.get('maxLength', 10)
return f'{base_type}({length})'
if field_type in ('decimal', 'numeric') and base_type in ('DECIMAL', 'NUMERIC', 'NUMBER'):
precision = field.get('precision', 10)
scale = field.get('scale', 2)
if db == 'oracle':
return f'NUMBER({precision}, {scale})'
return f'{base_type}({precision}, {scale})'
return base_type
def quote_table_name(table_name: str, schema: str, db_type: str) -> str:
db = normalize_db_type(db_type)
if schema and db in ('postgresql', 'sqlserver', 'oracle'):
return quote_table(schema, table_name, db)
return quote_identifier(table_name, db)
def _escape_sql_literal(value: str) -> str:
return (value or '').replace("'", "''")
@dataclass
class CreateTableDdlResult:
create_sql: str = ''
comment_sqls: List[str] = field(default_factory=list)
skipped: bool = False
def build_create_schema_sql(schema: str, db_type: str) -> str:
db = normalize_db_type(db_type)
if not schema:
return ''
if db == 'postgresql':
quoted = quote_identifier(schema, db)
return f'CREATE SCHEMA IF NOT EXISTS {quoted};'
if db == 'sqlserver':
safe_schema = schema.replace("'", "''")
return f"""
IF NOT EXISTS (SELECT * FROM sys.schemas WHERE name = '{safe_schema}')
BEGIN
EXEC('CREATE SCHEMA [{schema}]')
END;
""".strip()
return ''
def _generate_field_definition(field: dict, db_type: str) -> str:
field_name = field.get('name', '')
if not field_name:
return ''
db = normalize_db_type(db_type)
field_type = normalize_canonical_type(field.get('type', 'varchar'))
nullable = field.get('nullable', True)
is_primary = field.get('isPrimaryKey', False)
quoted_name = quote_identifier(field_name, db)
sql_type = map_to_dialect_type(field, db)
parts = [quoted_name, sql_type]
if is_primary:
parts.append('PRIMARY KEY')
if not nullable and not is_primary:
parts.append('NOT NULL')
if field_type == 'boolean' and not nullable:
if db == 'postgresql':
parts.append('DEFAULT FALSE')
else:
parts.append('DEFAULT 0')
if field_name == 'sort' and field_type in ('int', 'integer', 'bigint', 'smallint'):
parts.append('DEFAULT 0')
return ' '.join(parts)
def _generate_create_index_sql(
idx_name: str,
full_table_name: str,
field_name: str,
table_name: str,
db_type: str,
) -> str:
db = normalize_db_type(db_type)
quoted_idx = quote_identifier(idx_name, db)
quoted_field = quote_identifier(field_name, db)
if db == 'postgresql':
return f'CREATE INDEX IF NOT EXISTS {quoted_idx} ON {full_table_name} ({quoted_field});'
if db == 'mysql':
safe_table = table_name.replace("'", "''")
safe_idx = idx_name.replace("'", "''")
return (
f"SET @exist := (SELECT COUNT(*) FROM information_schema.statistics "
f"WHERE table_schema=DATABASE() AND table_name='{safe_table}' "
f"AND index_name='{safe_idx}');\n"
f"SET @sqlstmt := IF(@exist > 0, 'SELECT ''index exists''', "
f"'CREATE INDEX `{idx_name}` ON {full_table_name} (`{field_name}`)');\n"
f'PREPARE stmt FROM @sqlstmt;\nEXECUTE stmt;\nDEALLOCATE PREPARE stmt;'
)
if db == 'oracle':
return f'CREATE INDEX {quoted_idx} ON {full_table_name} ({quoted_field});'
if db == 'sqlserver':
safe_idx = idx_name.replace("'", "''")
return (
f"IF NOT EXISTS (SELECT 1 FROM sys.indexes WHERE name = '{safe_idx}')\n"
f'CREATE INDEX {quoted_idx} ON {full_table_name} ({quoted_field});'
)
return ''
def _build_column_comments_sql(
fields: List[dict],
full_table_name: str,
db_type: str,
) -> List[str]:
db = normalize_db_type(db_type)
if db not in ('postgresql', 'oracle'):
return []
comment_sqls: List[str] = []
for fld in fields:
comment = fld.get('comment', '')
field_name = fld.get('name', '')
if not comment or not field_name:
continue
safe_comment = _escape_sql_literal(comment)
quoted_col = quote_identifier(field_name, db)
comment_sqls.append(
f"COMMENT ON COLUMN {full_table_name}.{quoted_col} IS '{safe_comment}';"
)
return comment_sqls
def _wrap_sqlserver_create_if_not_exists(
schema: str,
table_name: str,
inner_sql: str,
) -> str:
schema_part = schema.replace("'", "''")
table_part = table_name.replace("'", "''")
return f"""
IF NOT EXISTS (
SELECT 1 FROM sys.tables t
INNER JOIN sys.schemas s ON t.schema_id = s.schema_id
WHERE s.name = '{schema_part}' AND t.name = '{table_part}'
)
BEGIN
{inner_sql}
END;
""".strip()
def _wrap_oracle_create_if_not_exists(table_name: str, inner_sql: str) -> str:
safe_table = table_name.upper().replace("'", "''")
single_line = ' '.join(inner_sql.split())
escaped_sql = single_line.replace("'", "''")
return f"""
DECLARE
v_count NUMBER;
BEGIN
SELECT COUNT(*) INTO v_count FROM user_tables WHERE table_name = '{safe_table}';
IF v_count = 0 THEN
EXECUTE IMMEDIATE '{escaped_sql}';
END IF;
END;
""".strip()
def build_create_table_ddl(
table: dict,
*,
db_type: str,
if_exists: str = 'skip',
effective_schema: str = '',
) -> CreateTableDdlResult:
"""
根据 canonical 表定义生成 CREATE TABLE DDL。
Args:
table: 含 tableName、fields、meta
db_type: postgresql / mysql / sqlserver / oracle
if_exists: skip | error | replace
effective_schema: 优先使用的 schema
"""
db = normalize_db_type(db_type)
table_name = table.get('tableName', '')
fields = table.get('fields', [])
meta = table.get('meta', {})
schema = effective_schema or meta.get('schema', '')
if not table_name or not fields:
raise ValueError(f'{table_name or "(unknown)"} 缺少必要配置')
full_table_name = quote_table_name(table_name, schema, db)
sql_parts: List[str] = []
if if_exists == 'replace':
if db == 'postgresql':
sql_parts.append(f'DROP TABLE IF EXISTS {full_table_name} CASCADE;')
elif db == 'mysql':
sql_parts.append(f'DROP TABLE IF EXISTS {full_table_name};')
elif db == 'sqlserver':
sql_parts.append(
f'IF OBJECT_ID(N\'{full_table_name.replace("[", "").replace("]", "")}\', N\'U\') IS NOT NULL '
f'DROP TABLE {full_table_name};'
)
elif db == 'oracle':
sql_parts.append(f'BEGIN EXECUTE IMMEDIATE \'DROP TABLE {full_table_name}\'; EXCEPTION WHEN OTHERS THEN NULL; END;')
if db in ('postgresql', 'mysql'):
create_clause = 'CREATE TABLE IF NOT EXISTS' if if_exists == 'skip' else 'CREATE TABLE'
else:
create_clause = 'CREATE TABLE'
field_defs = []
for fld in fields:
field_def = _generate_field_definition(fld, db)
if field_def:
field_defs.append(f' {field_def}')
create_body = f'{create_clause} {full_table_name} (\n' + ',\n'.join(field_defs) + '\n);'
if db == 'sqlserver' and if_exists == 'skip' and schema:
create_body = _wrap_sqlserver_create_if_not_exists(schema, table_name, create_body)
elif db == 'oracle' and if_exists == 'skip':
create_body = _wrap_oracle_create_if_not_exists(table_name, create_body)
sql_parts.append(create_body)
field_names = {f.get('name', '') for f in fields}
for field_name in sorted(field_names & SYSTEM_INDEX_FIELDS):
idx_name = f'idx_{table_name}_{field_name}'
index_sql = _generate_create_index_sql(
idx_name, full_table_name, field_name, table_name, db,
)
if index_sql:
sql_parts.append(index_sql)
comment_sqls = _build_column_comments_sql(fields, full_table_name, db)
return CreateTableDdlResult(
create_sql='\n'.join(sql_parts),
comment_sqls=comment_sqls,
skipped=False,
)
@@ -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
@@ -0,0 +1,258 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
数据库管理Schema
"""
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
# ============ 数据库相关 ============
class DatabaseInfo(BaseModel):
"""数据库信息"""
name: str
owner: Optional[str] = None
encoding: Optional[str] = None
collation: Optional[str] = None
size: Optional[str] = None
size_bytes: Optional[int] = None
tables_count: Optional[int] = None
description: Optional[str] = None
class DatabaseCreateIn(BaseModel):
"""创建数据库输入"""
name: str
owner: Optional[str] = None
encoding: Optional[str] = "UTF8"
charset: Optional[str] = "utf8mb4" # MySQL
collation: Optional[str] = None
template: Optional[str] = "template0" # PostgreSQL
class DatabaseRenameIn(BaseModel):
"""重命名数据库输入(PostgreSQL)"""
new_name: str
class DatabaseOperationOut(BaseModel):
"""数据库操作输出"""
success: bool
message: str
# ============ Schema相关(PostgreSQL ============
class SchemaInfo(BaseModel):
"""Schema信息"""
name: str
owner: Optional[str] = None
tables_count: Optional[int] = None
class SchemaCreateIn(BaseModel):
"""创建Schema输入(PostgreSQL"""
name: str
database: Optional[str] = None
owner: Optional[str] = None
class SchemaRenameIn(BaseModel):
"""重命名Schema输入(PostgreSQL"""
new_name: str
database: Optional[str] = None
class SchemaOperationOut(BaseModel):
"""Schema操作输出"""
success: bool
message: str
# ============ 表相关 ============
class TableInfo(BaseModel):
"""表信息"""
schema_name: Optional[str] = None
table_name: str
table_type: Optional[str] = "BASE TABLE"
row_count: Optional[int] = None
total_size: Optional[str] = None
total_size_bytes: Optional[int] = None
table_size: Optional[str] = None
table_size_bytes: Optional[int] = None
indexes_size: Optional[str] = None
indexes_size_bytes: Optional[int] = None
data_length: Optional[int] = None # MySQL
index_length: Optional[int] = None # MySQL
description: Optional[str] = None
# ============ 视图相关 ============
class ViewInfo(BaseModel):
"""视图信息"""
schema_name: Optional[str] = None
view_name: str
view_definition: Optional[str] = None
is_updatable: Optional[bool] = None
check_option: Optional[str] = None
view_type: Optional[str] = "VIEW"
description: Optional[str] = None
class ViewColumn(BaseModel):
"""视图列信息"""
column_name: str
data_type: str
is_nullable: Optional[bool] = None
ordinal_position: Optional[int] = None
description: Optional[str] = None
class ViewStructure(BaseModel):
"""视图结构详情"""
view_info: ViewInfo
columns: List[ViewColumn]
dependencies: List[str]
definition_sql: str
# ============ 列信息 ============
class ColumnInfo(BaseModel):
"""字段信息"""
column_name: str
data_type: str
is_nullable: bool
column_default: Optional[str] = None
character_maximum_length: Optional[int] = None
numeric_precision: Optional[int] = None
numeric_scale: Optional[int] = None
ordinal_position: Optional[int] = None
is_primary_key: bool = False
is_unique: bool = False
description: Optional[str] = None
class IndexInfo(BaseModel):
"""索引信息"""
index_name: str
index_type: Optional[str] = None
columns: str
is_unique: bool = False
is_primary: bool = False
definition: Optional[str] = None
class ConstraintInfo(BaseModel):
"""约束信息"""
constraint_name: str
constraint_type: str
columns: Optional[str] = None
definition: Optional[str] = None
referenced_table: Optional[str] = None
referenced_columns: Optional[str] = None
class TableStructure(BaseModel):
"""表结构详情"""
table_info: TableInfo
columns: List[ColumnInfo]
indexes: List[IndexInfo]
constraints: List[ConstraintInfo]
# ============ 数据查询相关 ============
class QueryDataIn(BaseModel):
"""查询数据输入"""
table_name: str
schema_name: Optional[str] = None
database: Optional[str] = None
page: int = 1
page_size: int = 20
where: Optional[str] = None
order_by: Optional[str] = None
class QueryDataOut(BaseModel):
"""查询数据输出"""
columns: List[str]
rows: List[Dict[str, Any]]
total: int
page: int
page_size: int
# ============ SQL执行相关 ============
class ExecuteSQLIn(BaseModel):
"""执行SQL输入"""
sql: str
is_query: bool = True
class ExecuteSQLOut(BaseModel):
"""执行SQL输出"""
success: bool
message: str
columns: Optional[List[str]] = None
rows: Optional[List[Dict[str, Any]]] = None
affected_rows: Optional[int] = None
execution_time: float
# ============ 数据操作相关 ============
class InsertDataIn(BaseModel):
"""插入数据"""
table_name: str
schema_name: Optional[str] = None
data: Dict[str, Any]
class UpdateDataIn(BaseModel):
"""更新数据"""
table_name: str
schema_name: Optional[str] = None
data: Dict[str, Any]
where: str
class DeleteDataIn(BaseModel):
"""删除数据"""
table_name: str
schema_name: Optional[str] = None
where: str
class DataOperationOut(BaseModel):
"""数据操作输出"""
success: bool
message: str
affected_rows: int
# ============ DDL操作相关 ============
class ExecuteDDLIn(BaseModel):
"""执行DDL输入"""
sql: str = Field(..., description="DDL SQL语句")
database: Optional[str] = Field(None, description="数据库名")
schema_name: Optional[str] = Field(None, description="Schema名")
class ExecuteDDLOut(BaseModel):
"""执行DDL输出"""
success: bool = Field(..., description="是否成功")
message: str = Field(..., description="执行消息")
affected_rows: int = Field(default=0, description="影响行数")
# ============ 数据库配置相关 ============
class DatabaseConfig(BaseModel):
"""数据库配置信息"""
db_name: str
name: str
db_type: str
host: str
port: int
database: str
user: str
has_password: bool
is_system: bool = False
display_name: Optional[str] = None
@@ -0,0 +1,513 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
数据库管理服务(异步版本)
支持 PostgreSQL、MySQL、SQL Server、Oracle
"""
import inspect
import logging
from typing import Any, Dict, List, Optional
from urllib.parse import urlparse
from app.config import settings
from core.database_manager.handlers.common import (
format_connection_error,
format_size,
log_database_connect_failure,
serialize_row,
)
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.pools import close_all_manager_pools
from core.database_manager.handlers.sqlserver import SQLServerHandler
from core.database_manager.sql_utils import is_protected_database, is_protected_schema
logger = logging.getLogger(__name__)
class DatabaseManagerForbidden(Exception):
"""写操作被禁止(系统连接或受保护对象)"""
class DatabaseManagerValidationError(Exception):
"""参数或数据库类型不支持"""
_SCHEMA_DB_TYPES = frozenset({"postgresql", "sqlserver", "oracle"})
_DEFAULT_PORTS = {
"postgresql": 5432,
"mysql": 3306,
"sqlserver": 1433,
"oracle": 1521,
}
def parse_database_url(database_url: str) -> dict:
"""解析数据库URL"""
parsed = urlparse(database_url)
scheme = parsed.scheme.lower()
if "postgresql" in scheme or "postgres" in scheme:
db_type = "postgresql"
default_port = 5432
elif "mysql" in scheme:
db_type = "mysql"
default_port = 3306
elif "mssql" in scheme or "sqlserver" in scheme:
db_type = "sqlserver"
default_port = 1433
elif "oracle" in scheme:
db_type = "oracle"
default_port = 1521
else:
return None
return {
"db_type": db_type,
"host": parsed.hostname or "localhost",
"port": parsed.port or default_port,
"user": parsed.username or "",
"password": parsed.password or "",
"database": parsed.path.lstrip("/") if parsed.path else "",
}
async def aiomysql_connect(**kwargs):
"""创建 aiomysql 连接,自动过滤当前版本不支持的参数"""
import aiomysql
supported = inspect.signature(aiomysql.connect).parameters
filtered = {key: value for key, value in kwargs.items() if key in supported}
return await aiomysql.connect(**filtered)
class AsyncDatabaseManagerService:
"""异步数据库管理服务(工厂类)"""
def __init__(
self,
db_name: str = "default",
connection: Optional[Dict[str, Any]] = None,
):
self.db_name = db_name
if connection is None:
connection = parse_database_url(settings.DATABASE_URL)
if not connection:
raise ValueError("Invalid DATABASE_URL")
self.db_type = connection["db_type"]
self.host = connection["host"]
self.port = connection["port"]
self.user = connection["user"]
self.password = connection.get("password", "")
self.database = connection["database"]
self.is_system = connection.get("is_system", db_name == "default")
extra_options = connection.get("extra_options") or {}
if self.db_type == "postgresql":
self._handler = PostgreSQLHandler(
self.host, self.port, self.user, self.password, self.database
)
elif self.db_type == "mysql":
self._handler = MySQLHandler(
self.host, self.port, self.user, self.password, self.database
)
elif self.db_type == "sqlserver":
self._handler = SQLServerHandler(
self.host,
self.port,
self.user,
self.password,
self.database,
extra_options,
)
elif self.db_type == "oracle":
self._handler = OracleHandler(
self.host,
self.port,
self.user,
self.password,
self.database,
extra_options,
)
else:
raise ValueError(f"Unsupported database type: {self.db_type}")
@classmethod
def from_connection_info(cls, info) -> "AsyncDatabaseManagerService":
"""从 ConnectionInfo 创建服务实例"""
conn = info.to_handler_kwargs()
conn["db_type"] = info.db_type
conn["is_system"] = info.is_system
conn["extra_options"] = getattr(info, "extra_options", None) or {}
return cls(info.code, conn)
@classmethod
async def create(
cls,
db_name: str = "default",
db=None,
) -> "AsyncDatabaseManagerService":
"""按连接 code 解析并创建服务实例"""
from core.database_connection.resolver import ConnectionResolver
info = await ConnectionResolver.resolve(db_name, db)
return cls.from_connection_info(info)
@staticmethod
def get_database_configs() -> List[Dict[str, Any]]:
"""同步获取配置(仅 default,兼容旧调用)"""
configs = []
db_info = parse_database_url(settings.DATABASE_URL)
if db_info:
configs.append(
{
"db_name": "default",
"name": db_info["database"],
"display_name": db_info["database"],
"db_type": db_info["db_type"],
"host": db_info["host"],
"port": db_info["port"],
"database": db_info["database"],
"user": db_info["user"],
"has_password": bool(db_info["password"]),
"is_system": True,
}
)
return configs
@staticmethod
async def get_database_configs_async(db) -> List[Dict[str, Any]]:
"""获取全部连接配置(default + 自定义)"""
from core.database_connection.service import DatabaseConnectionService
return await DatabaseConnectionService.get_manager_configs(db)
def _uses_schema_layer(self) -> bool:
return self.db_type in _SCHEMA_DB_TYPES
def _default_schema(self) -> str:
if self.db_type == "postgresql":
return "public"
if self.db_type == "sqlserver":
return "dbo"
if self.db_type == "oracle":
return getattr(self._handler, "default_schema", (self.user or "").upper())
return ""
async def test_connection(self) -> Dict[str, Any]:
try:
if await self._handler.connect():
logger.info(
"Database connection test succeeded: db_name=%s db_type=%s target=%s",
self.db_name,
self.db_type,
self._build_connection_target(),
)
return {
"success": True,
"message": "数据库连接成功",
"db_name": self.db_name,
"db_type": self.db_type,
}
detail = getattr(self._handler, "_last_connect_error", None) or "未知错误"
target = self._build_connection_target()
message = f"数据库连接失败 ({target}): {detail}"
log_database_connect_failure(
db_type=self.db_type,
host=self.host,
port=self.port,
user=self.user,
database=self.database,
db_name=self.db_name,
detail=detail,
action="test",
)
return {
"success": False,
"message": message,
"db_name": self.db_name,
"db_type": self.db_type,
}
except Exception as e:
detail = format_connection_error(e)
target = self._build_connection_target()
message = f"数据库连接失败 ({target}): {detail}"
log_database_connect_failure(
db_type=self.db_type,
host=self.host,
port=self.port,
user=self.user,
database=self.database,
db_name=self.db_name,
detail=detail,
action="test",
)
return {
"success": False,
"message": message,
"db_name": self.db_name,
"db_type": self.db_type,
}
finally:
await self._handler.close()
def _build_connection_target(self) -> str:
"""构建连接目标描述,便于定位失败原因"""
parts = [f"{self.host}:{self.port}"]
if self.database:
parts.append(f"db={self.database}")
if self.user:
parts.append(f"user={self.user}")
return ", ".join(parts)
def ensure_write_allowed(self) -> None:
# if self.db_name == "default" or self.is_system:
# raise DatabaseManagerForbidden("系统连接不允许执行写操作")
pass
def ensure_database_name_allowed(self, name: str) -> None:
if is_protected_database(name, self.db_type):
raise DatabaseManagerForbidden(f"系统数据库 '{name}' 不允许此操作")
def ensure_schema_name_allowed(self, name: str) -> None:
if is_protected_schema(name, self.db_type):
raise DatabaseManagerForbidden(f"系统 Schema '{name}' 不允许此操作")
async def get_databases(self) -> List[Dict[str, Any]]:
return await self._handler.get_databases()
async def create_database(self, name: str, **kwargs) -> bool:
self.ensure_write_allowed()
self.ensure_database_name_allowed(name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 不支持创建 Database")
return await self._handler.create_database(name, **kwargs)
async def drop_database(self, name: str) -> bool:
self.ensure_write_allowed()
self.ensure_database_name_allowed(name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 不支持删除 Database")
return await self._handler.drop_database(name)
async def rename_database(self, name: str, new_name: str) -> bool:
self.ensure_write_allowed()
self.ensure_database_name_allowed(name)
self.ensure_database_name_allowed(new_name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 不支持重命名 Database")
if self.db_type not in ("postgresql", "sqlserver"):
raise DatabaseManagerValidationError("当前数据库类型不支持重命名 Database")
return await self._handler.rename_database(name, new_name)
async def create_schema(self, name: str, database: str = None, owner: str = None) -> bool:
self.ensure_write_allowed()
self.ensure_schema_name_allowed(name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 暂不支持 Schema 创建")
if self.db_type not in ("postgresql", "sqlserver"):
raise DatabaseManagerValidationError("当前数据库类型不支持 Schema 操作")
return await self._handler.create_schema(name, database, owner)
async def drop_schema(self, name: str, database: str = None, cascade: bool = True) -> bool:
self.ensure_write_allowed()
self.ensure_schema_name_allowed(name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 暂不支持 Schema 删除")
if self.db_type not in ("postgresql", "sqlserver"):
raise DatabaseManagerValidationError("当前数据库类型不支持 Schema 操作")
return await self._handler.drop_schema(name, database, cascade)
async def rename_schema(self, name: str, new_name: str, database: str = None) -> bool:
self.ensure_write_allowed()
self.ensure_schema_name_allowed(name)
self.ensure_schema_name_allowed(new_name)
if self.db_type == "oracle":
raise DatabaseManagerValidationError("Oracle 不支持重命名 Schema")
if self.db_type == "sqlserver":
raise DatabaseManagerValidationError("SQL Server 不支持重命名 Schema")
if self.db_type != "postgresql":
raise DatabaseManagerValidationError("当前数据库类型不支持 Schema 重命名")
return await self._handler.rename_schema(name, new_name, database)
async def get_schemas(self, database: str = None) -> List[Dict[str, Any]]:
return await self._handler.get_schemas(database)
async def get_tables(self, database: str = None, schema_name: str = None) -> List[Dict[str, Any]]:
schema = schema_name or (self._default_schema() if self._uses_schema_layer() else None)
return await self._handler.get_tables(database, schema)
async def get_table_columns(
self, table_name: str, schema_name: str = None, database: str = None
) -> List[Dict[str, Any]]:
if self._uses_schema_layer():
return await self._handler.get_table_columns(
table_name, schema_name or self._default_schema(), database
)
db = database or schema_name or self.database
return await self._handler.get_table_columns(table_name, db)
async def get_table_indexes(
self, table_name: str, schema_name: str = None, database: str = None
) -> List[Dict[str, Any]]:
if self._uses_schema_layer():
return await self._handler.get_table_indexes(
table_name, schema_name or self._default_schema(), database
)
db = database or schema_name or self.database
return await self._handler.get_table_indexes(table_name, db)
async def get_table_constraints(
self, table_name: str, schema_name: str = None, database: str = None
) -> List[Dict[str, Any]]:
if self._uses_schema_layer():
return await self._handler.get_table_constraints(
table_name, schema_name or self._default_schema(), database
)
db = database or schema_name or self.database
return await self._handler.get_table_constraints(table_name, db)
async def get_table_structure(
self, table_name: str, database: str = None, schema_name: str = None
) -> Dict[str, Any]:
if self._uses_schema_layer() and schema_name is None:
schema_name = self._default_schema()
return await self._handler.get_table_structure(table_name, database, schema_name)
async def get_table_ddl(
self, table_name: str, schema_name: str = None, database: str = None
) -> str:
if self._uses_schema_layer():
return await self._handler.get_table_ddl(
table_name, schema_name or self._default_schema(), database
)
db = database or schema_name or self.database
return await self._handler.get_table_ddl(table_name, db)
async def get_views(
self, database: str = None, schema_name: str = None
) -> List[Dict[str, Any]]:
if self._uses_schema_layer() and schema_name is None:
schema_name = self._default_schema()
return await self._handler.get_views(database, schema_name)
async def get_view_structure(self, view_name: str, schema_name: str = None) -> Dict[str, Any]:
if self._uses_schema_layer() and schema_name is None:
schema_name = self._default_schema()
return await self._handler.get_view_structure(view_name, schema_name)
async def get_view_definition(self, view_name: str, schema_name: str = None) -> str:
if self._uses_schema_layer() and schema_name is None:
schema_name = self._default_schema()
return await self._handler.get_view_definition(view_name, schema_name)
async def get_view_dependencies(self, view_name: str, schema_name: str = None) -> List[str]:
if self._uses_schema_layer() and schema_name is None:
schema_name = self._default_schema()
return await self._handler.get_view_dependencies(view_name, schema_name)
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]:
return await self._handler.query_data(
table_name, schema_name, page, page_size, where, order_by, database
)
async def _connect_for_execute(self, database: Optional[str] = None) -> bool:
"""执行 SQL 前连接目标库(PostgreSQL 支持按表单配置的 database 切换)。"""
if self.db_type == "postgresql" and database and str(database).strip():
return await self._handler.connect(str(database).strip())
return await self._handler.connect()
async def execute_sql(
self,
sql: str,
is_query: bool = True,
database: Optional[str] = None,
) -> Dict[str, Any]:
if not is_query:
self.ensure_write_allowed()
in_tx = self.in_transaction
if not in_tx and not await self._connect_for_execute(database):
return {
"success": False,
"message": "数据库连接失败",
"columns": None,
"rows": None,
"affected_rows": None,
"execution_time": 0,
}
return await self._handler.execute_sql(sql, is_query)
@property
def in_transaction(self) -> bool:
return bool(getattr(self._handler, "_in_transaction", False))
async def begin_transaction(self, database: Optional[str] = None) -> bool:
if not await self._connect_for_execute(database):
return False
return await self._handler.begin_transaction()
async def commit_transaction(self) -> None:
await self._handler.commit_transaction()
async def rollback_transaction(self) -> None:
await self._handler.rollback_transaction()
async def create_savepoint(self, name: str) -> None:
await self._handler.create_savepoint(name)
async def release_savepoint(self, name: str) -> None:
await self._handler.release_savepoint(name)
async def rollback_to_savepoint(self, name: str) -> None:
await self._handler.rollback_to_savepoint(name)
async def insert_data(
self, table_name: str, data: Dict[str, Any], schema_name: str = None
) -> Dict[str, Any]:
self.ensure_write_allowed()
return await self._handler.insert_data(table_name, data, schema_name)
async def update_data(
self, table_name: str, data: Dict[str, Any], where: str, schema_name: str = None
) -> Dict[str, Any]:
self.ensure_write_allowed()
return await self._handler.update_data(table_name, data, where, schema_name)
async def delete_data(
self, table_name: str, where: str, schema_name: str = None
) -> Dict[str, Any]:
self.ensure_write_allowed()
return await self._handler.delete_data(table_name, where, schema_name)
async def execute_ddl(
self, sql: str, database: str = None, schema_name: str = None
) -> Dict[str, Any]:
self.ensure_write_allowed()
return await self._handler.execute_ddl(sql, database, schema_name)
__all__ = [
"AsyncDatabaseManagerService",
"DatabaseManagerForbidden",
"DatabaseManagerValidationError",
"parse_database_url",
"close_all_manager_pools",
"aiomysql_connect",
"format_size",
"serialize_row",
"format_connection_error",
"log_database_connect_failure",
"PostgreSQLHandler",
"MySQLHandler",
"SQLServerHandler",
"OracleHandler",
]
@@ -0,0 +1,195 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""SQL 工具:标识符转义、语句拆分、系统对象黑名单"""
from typing import List
PROTECTED_DATABASES = frozenset({
"postgres",
"template0",
"template1",
"mysql",
"information_schema",
"performance_schema",
"sys",
"master",
"tempdb",
"model",
"msdb",
})
PROTECTED_SCHEMAS = frozenset({
"pg_catalog",
"information_schema",
"pg_toast",
"sys",
"guest",
"db_owner",
"db_accessadmin",
"db_securityadmin",
"db_ddladmin",
"db_backupoperator",
"db_datareader",
"db_datawriter",
"db_denydatareader",
"db_denydatawriter",
})
ORACLE_PROTECTED_SCHEMAS = frozenset({
"sys",
"system",
"outln",
"xdb",
"ctxsys",
"mdsys",
"orddata",
"ordsys",
"lbacsys",
"dbsnmp",
"appqossys",
"audsys",
"gsmadmin_internal",
"oav",
"ordplugins",
"si_informtn_schema",
"wmsys",
})
def is_protected_database(name: str, db_type: str = "") -> bool:
lower = name.lower()
if lower in PROTECTED_DATABASES:
return True
if (db_type or "").lower() == "oracle":
return False
return False
def is_protected_schema(name: str, db_type: str = "") -> bool:
lower = name.lower()
if lower in PROTECTED_SCHEMAS:
return True
if lower.startswith("db_"):
return True
if (db_type or "").lower() == "oracle" and lower in ORACLE_PROTECTED_SCHEMAS:
return True
return False
def _normalize_db_type(db_type: str) -> str:
db = (db_type or "postgresql").lower()
if db in ("mssql", "sql server"):
return "sqlserver"
if db in ("postgres", "psql"):
return "postgresql"
return db
def quote_identifier(name: str, db_type: str) -> str:
"""按数据库类型转义标识符"""
db = _normalize_db_type(db_type)
if db == "mysql":
return f"`{name.replace('`', '``')}`"
if db in ("sqlserver", "mssql"):
return f"[{name.replace(']', ']]')}]"
if db == "oracle":
return f'"{name.replace(chr(34), chr(34) + chr(34))}"'
return f'"{name.replace(chr(34), chr(34) + chr(34))}"'
def quote_table(schema: str | None, table: str, db_type: str) -> str:
if schema:
return f"{quote_identifier(schema, db_type)}.{quote_identifier(table, db_type)}"
return quote_identifier(table, db_type)
def split_sql_statements(sql: str) -> List[str]:
"""按分号拆分 SQL,忽略字符串与 dollar-quote 内的分号"""
if not sql or not sql.strip():
return []
statements: List[str] = []
current: List[str] = []
i = 0
n = len(sql)
in_single = False
in_double = False
in_backtick = False
dollar_tag: str | None = None
while i < n:
if dollar_tag is None and not in_single and not in_double and not in_backtick:
if sql[i] == "$":
j = i + 1
while j < n and sql[j] != "$" and (sql[j].isalnum() or sql[j] == "_"):
j += 1
if j < n and sql[j] == "$":
dollar_tag = sql[i : j + 1]
current.append(dollar_tag)
i = j + 1
continue
if sql[i] == "'":
in_single = True
current.append(sql[i])
i += 1
continue
if sql[i] == '"':
in_double = True
current.append(sql[i])
i += 1
continue
if sql[i] == "`":
in_backtick = True
current.append(sql[i])
i += 1
continue
if sql[i] == ";":
stmt = "".join(current).strip()
if stmt:
statements.append(stmt)
current = []
i += 1
continue
elif dollar_tag is not None:
if sql.startswith(dollar_tag, i):
current.append(dollar_tag)
i += len(dollar_tag)
dollar_tag = None
continue
elif in_single:
if sql[i] == "'" and i + 1 < n and sql[i + 1] == "'":
current.append("''")
i += 2
continue
if sql[i] == "'":
in_single = False
current.append(sql[i])
i += 1
continue
elif in_double:
if sql[i] == '"' and i + 1 < n and sql[i + 1] == '"':
current.append('""')
i += 2
continue
if sql[i] == '"':
in_double = False
current.append(sql[i])
i += 1
continue
elif in_backtick:
if sql[i] == "`" and i + 1 < n and sql[i + 1] == "`":
current.append("``")
i += 2
continue
if sql[i] == "`":
in_backtick = False
current.append(sql[i])
i += 1
continue
current.append(sql[i])
i += 1
stmt = "".join(current).strip()
if stmt:
statements.append(stmt)
return statements