#!/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", ]