275 lines
10 KiB
Python
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
|