Build lightweight AI agent admin
This commit is contained in:
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user