Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,274 @@
|
||||
#!/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
|
||||
Reference in New Issue
Block a user