514 lines
20 KiB
Python
514 lines
20 KiB
Python
#!/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",
|
|
]
|