refactor(backend): 调整模块结构、调度器逻辑与健康检查路由

This commit is contained in:
zhangtao
2026-09-06 00:54:24 +08:00
parent 92081815f4
commit d470c7eb1f
324 changed files with 3435 additions and 7404 deletions
@@ -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