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