mirror of
https://github.com/fastapiadmin/FastapiAdmin.git
synced 2026-09-30 07:46:13 +00:00
refactor(backend): 调整模块结构、调度器逻辑与健康检查路由
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, Path, Query, Security, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.common.response import ResponseSchema, SuccessResponse
|
||||
from app.core.base_schema import AuthSchema, PageResultSchema, PaginationQueryParam
|
||||
from app.core.dependencies import AuthPermission, db_getter
|
||||
from app.core.router_class import OperationLogRoute
|
||||
from app.modules.workflow.core.base import (
|
||||
ADVANCED_FIELD_DEFS,
|
||||
DEFAULT_PORTS,
|
||||
AdvancedFieldDefSchema,
|
||||
StorageProtocol,
|
||||
StorageProtocolDefSchema,
|
||||
)
|
||||
|
||||
from .schema import StorageSourceCreateSchema, StorageSourceOutSchema, StorageSourceQueryParam, StorageSourceTestSchema, StorageSourceUpdateSchema
|
||||
from .service import StorageSourceService
|
||||
|
||||
StorageSourceRouter = APIRouter(route_class=OperationLogRoute, prefix="/source", tags=["存储源管理"])
|
||||
|
||||
|
||||
@StorageSourceRouter.get("/protocols", summary="查询支持的存储协议", response_model=ResponseSchema[list[StorageProtocolDefSchema]])
|
||||
async def get_storage_protocols_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
) -> JSONResponse:
|
||||
result: list[StorageProtocolDefSchema] = [
|
||||
StorageProtocolDefSchema(protocol=p.value, name=p.name, default_port=DEFAULT_PORTS[p])
|
||||
for p in StorageProtocol
|
||||
]
|
||||
return SuccessResponse(data=result, msg="查询存储协议成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.get("/advanced-fields", summary="查询存储 SDK 高级配置字段定义", response_model=ResponseSchema[dict[str, list[AdvancedFieldDefSchema]]])
|
||||
async def get_advanced_fields_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
) -> JSONResponse:
|
||||
"""返回按协议分组的 SDK 高级配置字段元数据,供前端「高级设置」面板按协议动态渲染。"""
|
||||
return SuccessResponse(data=ADVANCED_FIELD_DEFS, msg="查询高级配置字段成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.get("/page", summary="分页查询存储源", response_model=ResponseSchema[PageResultSchema[StorageSourceOutSchema]])
|
||||
async def get_storage_source_page_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
page: Annotated[PaginationQueryParam, Depends()],
|
||||
search: Annotated[StorageSourceQueryParam, Query()],
|
||||
) -> JSONResponse:
|
||||
result: PageResultSchema[StorageSourceOutSchema] = await StorageSourceService(auth, db).page(
|
||||
search=search,
|
||||
page_no=page.page_no,
|
||||
page_size=page.page_size,
|
||||
order_by=page.order_by,
|
||||
)
|
||||
return SuccessResponse(data=result, msg="查询存储源分页成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.get("/list", summary="查询存储源列表", response_model=ResponseSchema[list[StorageSourceOutSchema]])
|
||||
async def get_storage_source_list_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
search: Annotated[StorageSourceQueryParam, Query()],
|
||||
) -> JSONResponse:
|
||||
result: list[StorageSourceOutSchema] = await StorageSourceService(auth, db).get_list(search=search)
|
||||
return SuccessResponse(data=result, msg="查询存储源列表成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.get("/detail/{id}", summary="查询存储源详情", response_model=ResponseSchema[StorageSourceOutSchema])
|
||||
async def get_storage_source_detail_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="存储源ID", ge=1)],
|
||||
) -> JSONResponse:
|
||||
result: StorageSourceOutSchema = await StorageSourceService(auth, db).detail(id=id)
|
||||
return SuccessResponse(data=result, msg="查询存储源详情成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.post("/create", status_code=status.HTTP_201_CREATED, summary="创建存储源", response_model=ResponseSchema[StorageSourceOutSchema])
|
||||
async def create_storage_source_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:create"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
data: Annotated[StorageSourceCreateSchema, Body(description="存储源创建参数")],
|
||||
) -> JSONResponse:
|
||||
result: StorageSourceOutSchema = await StorageSourceService(auth, db).create(data=data)
|
||||
return SuccessResponse(data=result, msg="创建存储源成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.put("/update/{id}", summary="修改存储源", response_model=ResponseSchema[StorageSourceOutSchema])
|
||||
async def update_storage_source_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="存储源ID", ge=1)],
|
||||
data: Annotated[StorageSourceUpdateSchema, Body(description="存储源修改参数")],
|
||||
) -> JSONResponse:
|
||||
result: StorageSourceOutSchema = await StorageSourceService(auth, db).update(id=id, data=data)
|
||||
return SuccessResponse(data=result, msg="修改存储源成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.delete("/delete", summary="删除存储源", response_model=ResponseSchema[None])
|
||||
async def delete_storage_source_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:delete"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
ids: Annotated[list[int], Body(description="存储源ID列表")],
|
||||
) -> JSONResponse:
|
||||
await StorageSourceService(auth, db).delete(ids=ids)
|
||||
return SuccessResponse(msg="删除存储源成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.post("/test/{id}", summary="测试存储源连接", response_model=ResponseSchema[bool])
|
||||
async def test_storage_source_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="存储源ID", ge=1)],
|
||||
) -> JSONResponse:
|
||||
result: bool = await StorageSourceService(auth, db).test_connection(id=id)
|
||||
return SuccessResponse(data=result, msg="连接成功")
|
||||
|
||||
|
||||
@StorageSourceRouter.post("/test", summary="测试存储源连接(配置)", response_model=ResponseSchema[bool])
|
||||
async def test_storage_source_config_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:node:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
data: Annotated[StorageSourceTestSchema, Body(description="存储源连接配置")],
|
||||
) -> JSONResponse:
|
||||
result: bool = await StorageSourceService(auth, db).test_config(data=data)
|
||||
return SuccessResponse(data=result, msg="连接成功")
|
||||
@@ -0,0 +1,14 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_crud import CRUDBase
|
||||
from app.core.base_schema import AuthSchema
|
||||
|
||||
from .model import StorageSourceModel
|
||||
from .schema import StorageSourceCreateSchema, StorageSourceUpdateSchema
|
||||
|
||||
|
||||
class StorageSourceCRUD(CRUDBase[StorageSourceModel, StorageSourceCreateSchema, StorageSourceUpdateSchema]):
|
||||
"""存储源模块数据层"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
super().__init__(model=StorageSourceModel, auth=auth, db=db)
|
||||
@@ -0,0 +1,35 @@
|
||||
from sqlalchemy import JSON, Boolean, Integer, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.core.base_model import ModelMixin, UserMixin
|
||||
|
||||
|
||||
class StorageSourceModel(ModelMixin, UserMixin):
|
||||
"""存储源配置模型"""
|
||||
|
||||
__tablename__: str = "task_workflow_storage_source"
|
||||
__table_args__: dict[str, str] = {"comment": "存储源配置表"}
|
||||
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True, nullable=False, index=True, comment="存储源名称")
|
||||
protocol: Mapped[str] = mapped_column(String(16), nullable=False, index=True, comment="协议(ftp/ftps/sftp/s3/obs/oss/cos/local)")
|
||||
host: Mapped[str | None] = mapped_column(String(255), default=None, nullable=True, comment="主机地址(对象存储可不填,用 endpoint)")
|
||||
port: Mapped[int] = mapped_column(Integer, nullable=False, comment="端口号(FTP/FTPS/SFTP 用)")
|
||||
username: Mapped[str | None] = mapped_column(String(255), default=None, nullable=True, comment="用户名/AccessKey")
|
||||
password: Mapped[str | None] = mapped_column(Text, default=None, nullable=True, comment="密码/SecretKey(Fernet加密)")
|
||||
bucket: Mapped[str | None] = mapped_column(String(255), default=None, nullable=True, comment="桶名(对象存储专用)")
|
||||
endpoint: Mapped[str | None] = mapped_column(String(255), default=None, nullable=True, comment="接入点(对象存储)")
|
||||
region: Mapped[str | None] = mapped_column(String(64), default=None, nullable=True, comment="区域(对象存储)")
|
||||
path_prefix: Mapped[str | None] = mapped_column(String(255), default=None, nullable=True, comment="统一路径前缀")
|
||||
is_secure: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False, comment="是否启用TLS(FTPS)")
|
||||
implicit_tls: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False, comment="FTPS是否隐式TLS(默认显式)")
|
||||
scheme: Mapped[str] = mapped_column(String(16), default="https", nullable=False, comment="对象存储访问协议(http/https)")
|
||||
encrypt_type: Mapped[int] = mapped_column(Integer, default=1, nullable=False, comment="FTPS加密类型(0=明文 1=显式TLS可用时 2=要求显式TLS 3=隐式TLS)")
|
||||
connection_mode: Mapped[int] = mapped_column(Integer, default=0, nullable=False, comment="FTP/FTPS传输模式(0=默认 1=主动 2=被动)")
|
||||
encoding: Mapped[str] = mapped_column(String(16), default="UTF-8", nullable=False, comment="FTP/FTPS/SFTP编码(utf-8/gbk 等)")
|
||||
multipart_part_size: Mapped[int] = mapped_column(Integer, default=50, nullable=False, comment="分片大小(MB,对象存储分片上传)")
|
||||
multipart_concurrency: Mapped[int] = mapped_column(Integer, default=6, nullable=False, comment="分片上传并发路数")
|
||||
multipart_memory_budget: Mapped[int] = mapped_column(Integer, default=512, nullable=False, comment="分片上传内存预算(MB)")
|
||||
advanced_config: Mapped[dict | None] = mapped_column(JSON, default=None, nullable=True, comment="SDK高级配置(JSON,按协议解析)")
|
||||
is_default: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False, comment="是否默认存储源")
|
||||
status: Mapped[int] = mapped_column(Integer, default=0, nullable=False, comment="状态(0:启用 1:停用)")
|
||||
description: Mapped[str | None] = mapped_column(Text, default=None, nullable=True, comment="备注")
|
||||
@@ -0,0 +1,157 @@
|
||||
import codecs
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
from app.core.base_schema import BaseQueryParam, BaseSchema, UserByQueryParam, UserBySchema
|
||||
from app.modules.workflow.core.base import DEFAULT_PORTS, StorageProtocol
|
||||
|
||||
|
||||
class StorageSourceConfigSchema(BaseModel):
|
||||
"""存储源连接配置模型(创建/测试共用)"""
|
||||
|
||||
protocol: StorageProtocol = Field(..., description="协议(ftp/ftps/sftp/s3/obs/oss/cos/local)")
|
||||
host: str | None = Field(default=None, max_length=255, description="主机地址/根目录(对象存储可不填,用 endpoint)")
|
||||
port: int | None = Field(default=None, ge=0, le=65535, description="端口(local 协议为0;不传则使用协议默认端口)")
|
||||
username: str | None = Field(default=None, max_length=255, description="用户名/AccessKey")
|
||||
password: str | None = Field(default=None, max_length=512, description="密码/SecretKey")
|
||||
bucket: str | None = Field(default=None, max_length=255, description="桶名/根目录")
|
||||
endpoint: str | None = Field(default=None, max_length=255, description="接入点(对象存储)")
|
||||
region: str | None = Field(default=None, max_length=64, description="区域(对象存储)")
|
||||
path_prefix: str | None = Field(default=None, max_length=255, description="统一路径前缀")
|
||||
is_secure: bool = Field(default=False, description="是否启用TLS(FTPS)")
|
||||
implicit_tls: bool = Field(default=False, description="FTPS是否隐式TLS(默认显式)")
|
||||
scheme: str = Field(default="https", description="对象存储访问协议(http/https)")
|
||||
encrypt_type: int = Field(default=1, description="FTPS加密类型(0=明文 1=显式TLS可用时 2=要求显式TLS 3=隐式TLS)")
|
||||
connection_mode: int = Field(default=0, description="FTP/FTPS传输模式(0=默认 1=主动 2=被动)")
|
||||
encoding: str = Field(default="UTF-8", description="FTP/FTPS/SFTP编码(utf-8/gbk 等)")
|
||||
# 分片传输参数(对象存储分片上传,每个端点独立配置;FTP/SFTP/LOCAL 等协议忽略)
|
||||
multipart_part_size: int = Field(default=50, ge=5, le=5000, description="分片大小(MB)")
|
||||
multipart_concurrency: int = Field(default=6, ge=1, le=64, description="分片上传并发路数")
|
||||
multipart_memory_budget: int = Field(default=512, ge=8, le=10240, description="分片上传内存预算(MB)")
|
||||
# SDK 高级配置(按协议解析,见 core/base.py;留空使用默认值)
|
||||
advanced_config: dict | None = Field(default=None, description="SDK高级配置(JSON)")
|
||||
|
||||
@field_validator("path_prefix")
|
||||
@classmethod
|
||||
def validate_path_prefix(cls, value: str | None) -> str | None:
|
||||
if value:
|
||||
value = value.strip()
|
||||
if ".." in value or "\x00" in value:
|
||||
raise ValueError("路径前缀包含非法字符")
|
||||
return value
|
||||
|
||||
@field_validator("encoding")
|
||||
@classmethod
|
||||
def validate_encoding(cls, value: str | None) -> str | None:
|
||||
if value:
|
||||
try:
|
||||
codecs.lookup(value)
|
||||
except LookupError:
|
||||
raise ValueError(f"未知的编码: {value}(如 utf-8/gbk 等)")
|
||||
return value
|
||||
|
||||
@field_validator("advanced_config")
|
||||
@classmethod
|
||||
def validate_advanced_config(cls, value: dict | None, info) -> dict | None:
|
||||
"""校验协议 SDK 高级配置取值,避免非法值直达 SDK 报错。"""
|
||||
if not value:
|
||||
return value
|
||||
if info.data.get("protocol") == StorageProtocol.S3:
|
||||
s3_enum = {
|
||||
"signature_version": {"s3v4", "s3", "v2"},
|
||||
"retries_mode": {"standard", "adaptive", "legacy"},
|
||||
"addressing_style": {"auto", "path", "virtual"},
|
||||
}
|
||||
for key, allowed in s3_enum.items():
|
||||
if key in value and value[key] not in allowed:
|
||||
raise ValueError(f"S3 {key} 取值必须为 {'/'.join(sorted(allowed))}")
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_protocol_fields(self):
|
||||
"""按协议校验必填字段并填充默认端口。"""
|
||||
if self.port is None:
|
||||
self.port = DEFAULT_PORTS[self.protocol]
|
||||
# 对象存储类协议必须配置桶/空间名
|
||||
obj_store_protocols = (
|
||||
StorageProtocol.S3,
|
||||
StorageProtocol.OBS,
|
||||
StorageProtocol.OSS,
|
||||
StorageProtocol.COS,
|
||||
)
|
||||
if self.protocol in obj_store_protocols:
|
||||
# 对象存储协议由接入点 URL 决定端口,统一置 0(避免无意义的端口输入)
|
||||
self.port = 0
|
||||
if self.protocol in obj_store_protocols and not self.bucket:
|
||||
raise ValueError(f"{self.protocol.value} 协议必须配置 bucket")
|
||||
# 非对象存储协议必须配置主机地址
|
||||
if self.protocol not in obj_store_protocols and not self.host:
|
||||
raise ValueError(f"{self.protocol.value} 协议必须配置 host")
|
||||
# endpoint 必须配置的协议(接入点须为完整 URL,如 https://xxx)
|
||||
if self.protocol in (StorageProtocol.S3, StorageProtocol.OBS, StorageProtocol.OSS) and not self.endpoint:
|
||||
raise ValueError(f"{self.protocol.value} 协议必须配置 endpoint")
|
||||
if self.protocol == StorageProtocol.COS and not self.region:
|
||||
raise ValueError("cos 协议必须配置 region")
|
||||
if self.protocol == StorageProtocol.COS:
|
||||
# COS 默认走 https(适配器以 is_secure 决定 Scheme)
|
||||
self.is_secure = True
|
||||
appid = self.advanced_config.get("appid") if self.advanced_config else None
|
||||
if appid and not str(appid).isdigit():
|
||||
raise ValueError("COS appid 必须为纯数字")
|
||||
if self.protocol == StorageProtocol.FTPS:
|
||||
if self.encrypt_type < 1:
|
||||
raise ValueError("FTPS 加密方式不能为 0(明文),明文请使用 FTP 协议")
|
||||
# 显式 TLS(1/2) 默认端口 21,990 是隐式 TLS 端口
|
||||
if self.encrypt_type < 3 and self.port == DEFAULT_PORTS[StorageProtocol.FTPS]:
|
||||
self.port = 21
|
||||
self.is_secure = True
|
||||
return self
|
||||
|
||||
|
||||
class StorageSourceCreateSchema(StorageSourceConfigSchema):
|
||||
"""存储源创建模型"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=64, description="存储源名称")
|
||||
is_default: bool = Field(default=False, description="是否默认存储源")
|
||||
status: int = Field(default=0, ge=0, le=1, description="状态(0:启用 1:停用)")
|
||||
description: str | None = Field(default=None, max_length=255, description="备注")
|
||||
|
||||
@field_validator("name")
|
||||
@classmethod
|
||||
def validate_name(cls, value: str) -> str:
|
||||
value = value.strip()
|
||||
if not value:
|
||||
raise ValueError("存储源名称不能为空")
|
||||
return value
|
||||
|
||||
|
||||
class StorageSourceUpdateSchema(StorageSourceCreateSchema):
|
||||
"""存储源更新模型(password 为空表示不修改原密码)"""
|
||||
|
||||
|
||||
class StorageSourceTestSchema(StorageSourceConfigSchema):
|
||||
"""存储源连接测试模型(仅校验连接配置,不落库;密码留空且传 source_id 时回退已保存密码)"""
|
||||
|
||||
source_id: int | None = Field(default=None, ge=1, description="已保存的存储源ID(编辑态测试时使用)")
|
||||
|
||||
|
||||
class StorageSourceOutSchema(StorageSourceCreateSchema, BaseSchema, UserBySchema):
|
||||
"""存储源详情响应模型(密码永不明文返回)"""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
password: None = Field(default=None, exclude=True, repr=False, description="密码(不返回)")
|
||||
has_password: bool = Field(default=False, description="是否已配置密码")
|
||||
|
||||
@field_validator("password", mode="before")
|
||||
@classmethod
|
||||
def mask_password(cls, value) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class StorageSourceQueryParam(BaseQueryParam, UserByQueryParam):
|
||||
"""存储源管理查询参数"""
|
||||
|
||||
name: str | None = Field(None, description="存储源名称", json_schema_extra={"q": "like"})
|
||||
protocol: StorageProtocol | None = Field(None, description="协议", json_schema_extra={"q": "eq"})
|
||||
status: int | None = Field(None, ge=0, le=1, description="状态(0:启用 1:停用)", json_schema_extra={"q": "eq"})
|
||||
@@ -0,0 +1,204 @@
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import ColumnElement, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_schema import AuthSchema, PageResultSchema
|
||||
from app.core.exceptions import CustomException
|
||||
from app.modules.workflow.core.base import (
|
||||
DEFAULT_PORTS,
|
||||
StorageAdapterConfig,
|
||||
StorageProtocol,
|
||||
decrypt_password,
|
||||
encrypt_password,
|
||||
)
|
||||
from app.modules.workflow.core.factory import StorageAdapterFactory
|
||||
from app.utils.common_util import search_to_dict
|
||||
|
||||
from .crud import StorageSourceCRUD
|
||||
from .model import StorageSourceModel
|
||||
from .schema import StorageSourceCreateSchema, StorageSourceOutSchema, StorageSourceQueryParam, StorageSourceTestSchema, StorageSourceUpdateSchema
|
||||
|
||||
|
||||
class StorageSourceService:
|
||||
"""存储源管理服务"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
self.auth = auth
|
||||
self.db = db
|
||||
|
||||
# ── 内部工具 ────────────────────────────────────────────────────
|
||||
|
||||
def _crud(self) -> StorageSourceCRUD:
|
||||
return StorageSourceCRUD(self.auth, self.db)
|
||||
|
||||
@staticmethod
|
||||
def _to_out(obj: StorageSourceModel) -> StorageSourceOutSchema:
|
||||
out = StorageSourceOutSchema.model_validate(obj)
|
||||
out.has_password = bool(obj.password)
|
||||
return out
|
||||
|
||||
async def _clear_other_default(self, keep_id: int | None = None) -> None:
|
||||
"""取消其他存储源的默认标记,保证同时只有一个默认源。"""
|
||||
conditions: list[ColumnElement[bool]] = [StorageSourceModel.is_default.is_(True)]
|
||||
if keep_id is not None:
|
||||
conditions.append(StorageSourceModel.id != keep_id)
|
||||
await self.db.execute(update(StorageSourceModel).where(*conditions).values(is_default=False))
|
||||
|
||||
@staticmethod
|
||||
def _build_config_from(source: Any, password: str) -> StorageAdapterConfig:
|
||||
"""从 ORM 或表单对象构造适配器配置(password 为明文,统一处理端口/高级配置)。"""
|
||||
return StorageAdapterConfig(
|
||||
protocol=StorageProtocol(source.protocol),
|
||||
host=source.host,
|
||||
port=source.port or DEFAULT_PORTS[StorageProtocol(source.protocol)],
|
||||
username=source.username,
|
||||
password=password,
|
||||
bucket=source.bucket,
|
||||
endpoint=source.endpoint,
|
||||
scheme=source.scheme,
|
||||
region=source.region,
|
||||
path_prefix=source.path_prefix,
|
||||
is_secure=source.is_secure,
|
||||
implicit_tls=source.implicit_tls,
|
||||
encrypt_type=source.encrypt_type,
|
||||
connection_mode=source.connection_mode,
|
||||
encoding=source.encoding,
|
||||
multipart_part_size=source.multipart_part_size,
|
||||
multipart_concurrency=source.multipart_concurrency,
|
||||
multipart_memory_budget=source.multipart_memory_budget,
|
||||
advanced_config=source.advanced_config or {},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_config(obj: StorageSourceModel) -> StorageAdapterConfig:
|
||||
return StorageSourceService._build_config_from(obj, decrypt_password(obj.password))
|
||||
|
||||
# ── 查询 ────────────────────────────────────────────────────────
|
||||
|
||||
async def detail(self, id: int) -> StorageSourceOutSchema:
|
||||
obj = await self._crud().get_or_404(id=id)
|
||||
return self._to_out(obj)
|
||||
|
||||
async def page(
|
||||
self,
|
||||
search: StorageSourceQueryParam | None,
|
||||
page_no: int,
|
||||
page_size: int,
|
||||
order_by: list[dict] | None = None,
|
||||
) -> PageResultSchema[StorageSourceOutSchema]:
|
||||
result = await self._crud().page(
|
||||
offset=(page_no - 1) * page_size,
|
||||
limit=page_size,
|
||||
order_by=order_by or [{"id": "asc"}],
|
||||
search=search_to_dict(search),
|
||||
)
|
||||
return PageResultSchema[StorageSourceOutSchema](
|
||||
page_no=result.page_no,
|
||||
page_size=result.page_size,
|
||||
total=result.total,
|
||||
has_next=result.has_next,
|
||||
items=[self._to_out(obj) for obj in result.items],
|
||||
)
|
||||
|
||||
async def get_list(self, search: StorageSourceQueryParam | None = None) -> list[StorageSourceOutSchema]:
|
||||
objs = await self._crud().get_list(search=search_to_dict(search), order_by=[{"id": "asc"}])
|
||||
return [self._to_out(obj) for obj in objs]
|
||||
|
||||
# ── 写入 ────────────────────────────────────────────────────────
|
||||
|
||||
async def create(self, data: StorageSourceCreateSchema) -> StorageSourceOutSchema:
|
||||
exist = await self._crud().get(name=data.name)
|
||||
if exist:
|
||||
raise CustomException(msg="创建失败,存储源名称已存在")
|
||||
|
||||
payload = data.model_dump(exclude_none=True)
|
||||
if payload.get("password"):
|
||||
payload["password"] = encrypt_password(payload["password"])
|
||||
|
||||
obj = await self._crud().create(data=payload)
|
||||
if data.is_default:
|
||||
await self._clear_other_default(keep_id=obj.id)
|
||||
return self._to_out(obj)
|
||||
|
||||
async def update(self, id: int, data: StorageSourceUpdateSchema) -> StorageSourceOutSchema:
|
||||
await self._crud().get_or_404(id=id, msg="更新失败,该存储源不存在")
|
||||
exist = await self._crud().get(name=data.name)
|
||||
if exist and exist.id != id:
|
||||
raise CustomException(msg="更新失败,存储源名称已存在")
|
||||
|
||||
# 不用 exclude_none:None 表示用户清空可选字段(description/path_prefix/region 等),需落库为 NULL
|
||||
payload = data.model_dump(exclude_unset=True)
|
||||
if payload.get("password"):
|
||||
payload["password"] = encrypt_password(payload["password"])
|
||||
else:
|
||||
payload.pop("password", None) # 未传新密码则不修改
|
||||
|
||||
await self._crud().update(id=id, data=payload)
|
||||
obj = await self._crud().get_or_404(id=id)
|
||||
if data.is_default:
|
||||
await self._clear_other_default(keep_id=id)
|
||||
return self._to_out(obj)
|
||||
|
||||
async def delete(self, ids: list[int]) -> None:
|
||||
if not ids:
|
||||
raise CustomException(msg="删除失败,删除对象不能为空")
|
||||
await self._crud().delete(ids=ids)
|
||||
|
||||
# ── 连接测试 ────────────────────────────────────────────────────
|
||||
|
||||
async def test_connection(self, id: int) -> bool:
|
||||
obj = await self._crud().get_or_404(id=id, msg="该存储源不存在")
|
||||
adapter = StorageAdapterFactory.create(self._build_config(obj))
|
||||
try:
|
||||
ok = await adapter.test_connection()
|
||||
finally:
|
||||
await adapter.close()
|
||||
if not ok:
|
||||
raise CustomException(msg="连接失败,请检查存储源配置")
|
||||
return True
|
||||
|
||||
async def test_config(self, data: StorageSourceTestSchema) -> bool:
|
||||
"""使用表单提交的配置直接测试连接(不落库),密码留空且传 source_id 时回退已保存密码。"""
|
||||
password = data.password or ""
|
||||
if not password and data.source_id:
|
||||
obj = await self._crud().get_or_404(id=data.source_id, msg="该存储源不存在")
|
||||
password = decrypt_password(obj.password)
|
||||
adapter = StorageAdapterFactory.create(self._build_config_from(data, password))
|
||||
try:
|
||||
ok = await adapter.test_connection()
|
||||
finally:
|
||||
await adapter.close()
|
||||
if not ok:
|
||||
raise CustomException(msg="连接失败,请检查存储源配置")
|
||||
return True
|
||||
|
||||
# ── 供文件模块复用 ──────────────────────────────────────────────
|
||||
|
||||
async def get_active_source(self, source_id: int | None = None) -> StorageSourceModel:
|
||||
"""获取可用存储源:优先指定 id;否则默认源;再退化为任一启用源。"""
|
||||
if source_id:
|
||||
source = await self._crud().get_or_404(id=source_id)
|
||||
if source.status == 1:
|
||||
raise CustomException(msg="该存储源已停用")
|
||||
return source
|
||||
source = await self._crud().get(status=0, is_default=True)
|
||||
if source:
|
||||
return source
|
||||
source = await self._crud().get(status=0)
|
||||
if source:
|
||||
return source
|
||||
raise CustomException(msg="未配置可用的存储源,请先在存储源管理中创建")
|
||||
|
||||
async def get_active_sources(self, ids: list[int]) -> list[StorageSourceModel]:
|
||||
"""批量获取可用存储源:全部存在且启用,任一无效即抛错(一次查询,避免 N+1)。"""
|
||||
if not ids:
|
||||
return []
|
||||
objs = list(await self._crud().get_list(search={"id": ("in", ids)}))
|
||||
missing = [i for i in ids if i not in {o.id for o in objs}]
|
||||
if missing:
|
||||
raise CustomException(msg=f"存储源不存在: {missing}")
|
||||
disabled = [o.name for o in objs if o.status == 1]
|
||||
if disabled:
|
||||
raise CustomException(msg=f"存储源已停用: {', '.join(disabled)}")
|
||||
return objs
|
||||
Reference in New Issue
Block a user