#!/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