Files
ai-agent-admin/backend-fastapi/core/database_connection/service.py
T
2026-06-08 18:14:59 +08:00

275 lines
10 KiB
Python

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""数据库连接服务"""
import logging
import re
from typing import Any, Dict, List, Optional, Tuple
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.timezone import format_datetime
from core.database_connection.model import DatabaseConnection
from core.database_connection.resolver import ConnectionResolver
from core.database_connection.types import ConnectionInfo
from core.database_manager.service import AsyncDatabaseManagerService, parse_database_url
from utils.secret_crypto import encrypt_secret
logger = logging.getLogger(__name__)
RESOURCE_TYPE = "database_connection"
RESOURCE_DISPLAY_NAME = "数据库连接管理"
_CODE_PATTERN = re.compile(r"^[a-zA-Z][a-zA-Z0-9_-]{0,49}$")
_ALLOWED_DB_TYPES = {"postgresql", "mysql", "sqlserver", "oracle"}
_DEFAULT_PORTS = {
"postgresql": 5432,
"mysql": 3306,
"sqlserver": 1433,
"oracle": 1521,
}
def _default_port(db_type: str) -> int:
return _DEFAULT_PORTS.get(db_type, 5432)
class DatabaseConnectionService:
@staticmethod
def validate_code(code: str) -> None:
if code == "default":
raise ValueError("code 'default' is reserved for system connection")
if not _CODE_PATTERN.match(code):
raise ValueError("Invalid connection code format")
@staticmethod
def validate_db_type(db_type: str) -> None:
if db_type not in _ALLOWED_DB_TYPES:
raise ValueError(f"Unsupported db_type: {db_type}")
@classmethod
async def check_code_exists(cls, db: AsyncSession, code: str, exclude_id: str = None) -> bool:
stmt = select(DatabaseConnection.id).where(
DatabaseConnection.code == code,
DatabaseConnection.is_deleted == False,
)
if exclude_id:
stmt = stmt.where(DatabaseConnection.id != exclude_id)
result = await db.execute(stmt)
return result.scalar_one_or_none() is not None
@classmethod
async def get_list(
cls,
db: AsyncSession,
page: int = 1,
page_size: int = 20,
application_id: str = None,
name: str = None,
code: str = None,
status: bool = None,
) -> Tuple[List[DatabaseConnection], int]:
stmt = select(DatabaseConnection).where(DatabaseConnection.is_deleted == False)
if application_id:
stmt = stmt.where(DatabaseConnection.application_id == application_id)
else:
stmt = stmt.where(DatabaseConnection.application_id.is_(None))
if name:
stmt = stmt.where(DatabaseConnection.name.contains(name))
if code:
stmt = stmt.where(DatabaseConnection.code.contains(code))
if status is not None:
stmt = stmt.where(DatabaseConnection.status == status)
count_stmt = select(func.count()).select_from(stmt.subquery())
total = (await db.execute(count_stmt)).scalar() or 0
stmt = stmt.order_by(
DatabaseConnection.sort.desc(),
DatabaseConnection.sys_create_datetime.desc(),
)
stmt = stmt.offset((page - 1) * page_size).limit(page_size)
items = list((await db.execute(stmt)).scalars().all())
return items, total
@classmethod
async def get_all_enabled(
cls,
db: AsyncSession,
application_id: str = None,
) -> List[DatabaseConnection]:
stmt = select(DatabaseConnection).where(
DatabaseConnection.is_deleted == False,
DatabaseConnection.status == True,
DatabaseConnection.is_system == False,
)
if application_id:
stmt = stmt.where(DatabaseConnection.application_id == application_id)
else:
stmt = stmt.where(DatabaseConnection.application_id.is_(None))
stmt = stmt.order_by(DatabaseConnection.sort.desc(), DatabaseConnection.name.asc())
return list((await db.execute(stmt)).scalars().all())
@classmethod
async def get_by_id(cls, db: AsyncSession, conn_id: str) -> Optional[DatabaseConnection]:
result = await db.execute(
select(DatabaseConnection).where(
DatabaseConnection.id == conn_id,
DatabaseConnection.is_deleted == False,
)
)
return result.scalar_one_or_none()
@classmethod
async def get_by_code(cls, db: AsyncSession, code: str) -> Optional[DatabaseConnection]:
result = await db.execute(
select(DatabaseConnection).where(
DatabaseConnection.code == code,
DatabaseConnection.is_deleted == False,
)
)
return result.scalar_one_or_none()
@classmethod
async def create(cls, db: AsyncSession, data: Dict[str, Any]) -> DatabaseConnection:
cls.validate_code(data["code"])
cls.validate_db_type(data["db_type"])
row = DatabaseConnection(
application_id=data.get("application_id"),
code=data["code"],
name=data["name"],
db_type=data["db_type"],
host=data["host"],
port=data.get("port") or _default_port(data["db_type"]),
user=data.get("user") or "",
password_enc=encrypt_secret((data.get("password") or "").strip()),
default_database=data.get("default_database") or "",
description=data.get("description") or "",
status=data.get("status", True),
is_system=False,
extra_options=data.get("extra_options") or {},
)
db.add(row)
await db.flush()
await db.refresh(row)
return row
@classmethod
async def update(cls, db: AsyncSession, row: DatabaseConnection, data: Dict[str, Any]) -> DatabaseConnection:
if row.is_system:
raise ValueError("System connection cannot be modified")
if data.get("db_type"):
cls.validate_db_type(data["db_type"])
for field in ("name", "db_type", "host", "port", "user", "default_database", "description", "status"):
if field in data and data[field] is not None:
setattr(row, field, data[field])
if data.get("extra_options") is not None:
row.extra_options = data["extra_options"]
if data.get("password"):
row.password_enc = encrypt_secret(str(data["password"]).strip())
await db.flush()
await db.refresh(row)
ConnectionResolver.invalidate_cache(row.code)
return row
@classmethod
async def delete(cls, db: AsyncSession, row: DatabaseConnection) -> None:
if row.is_system or row.code == "default":
raise ValueError("System connection cannot be deleted")
row.is_deleted = True
await db.flush()
ConnectionResolver.invalidate_cache(row.code)
@classmethod
async def test_connection_info(cls, info: ConnectionInfo) -> Dict[str, Any]:
service = AsyncDatabaseManagerService.from_connection_info(info)
result = await service.test_connection()
if not result.get("success"):
logger.error(
"Database connection registry test failed: code=%s message=%s",
info.code,
result.get("message"),
)
return result
@classmethod
async def test_config(cls, config: Dict[str, Any]) -> Dict[str, Any]:
db_type = config["db_type"]
cls.validate_db_type(db_type)
info = ConnectionInfo(
code="__test__",
db_type=db_type,
host=config["host"],
port=config.get("port") or _default_port(db_type),
user=config.get("user") or "",
password=config.get("password") or "",
database=config.get("default_database") or "",
)
return await cls.test_connection_info(info)
@classmethod
async def test_saved(cls, db: AsyncSession, row: DatabaseConnection) -> Dict[str, Any]:
info = await ConnectionResolver.resolve(row.code, db)
return await cls.test_connection_info(info)
@classmethod
def to_response_dict(cls, row: DatabaseConnection) -> Dict[str, Any]:
return {
"id": row.id,
"code": row.code,
"name": row.name,
"db_type": row.db_type,
"host": row.host,
"port": row.port,
"user": row.user or "",
"default_database": row.default_database or "",
"description": row.description or "",
"status": row.status,
"is_system": row.is_system,
"has_password": bool(row.password_enc),
"application_id": row.application_id,
"extra_options": row.extra_options or {},
"sort": row.sort or 0,
"sys_create_datetime": format_datetime(row.sys_create_datetime) if row.sys_create_datetime else None,
"sys_update_datetime": format_datetime(row.sys_update_datetime) if row.sys_update_datetime else None,
}
@classmethod
async def get_manager_configs(
cls,
db: AsyncSession,
application_id: str = None,
) -> List[Dict[str, Any]]:
configs: List[Dict[str, Any]] = []
default_info = ConnectionResolver.default_connection_info()
configs.append({
"db_name": "default",
"name": default_info.display_name or default_info.database or "default",
"display_name": default_info.display_name or default_info.database or "default",
"db_type": default_info.db_type,
"host": default_info.host,
"port": default_info.port,
"database": default_info.database,
"user": default_info.user,
"has_password": bool(default_info.password),
"is_system": True,
})
custom_rows = await cls.get_all_enabled(db, application_id)
for row in custom_rows:
configs.append({
"db_name": row.code,
"name": row.name,
"display_name": row.name,
"db_type": row.db_type,
"host": row.host,
"port": row.port,
"database": row.default_database or "",
"user": row.user or "",
"has_password": bool(row.password_enc),
"is_system": False,
})
return configs