Files
2026-06-08 18:14:59 +08:00

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",
]