mirror of
https://github.com/fastapiadmin/FastapiAdmin.git
synced 2026-10-01 16:21:20 +00:00
refactor(backend): 调整模块结构、调度器逻辑与健康检查路由
This commit is contained in:
@@ -0,0 +1,10 @@
|
||||
"""工作流模块(存储与文件流转):
|
||||
|
||||
- ``source``: 存储源管理(OSS/COS/OBS/S3/SFTP/FTP 等连接配置)
|
||||
- ``core``: 对象存储协议适配器与工厂
|
||||
- ``storage``: 存储文件浏览
|
||||
- ``transfer``: 传输任务引擎
|
||||
- ``flow``: 工作流定义(画布 CRUD、发布、执行 API)
|
||||
|
||||
路由统一挂在 ``/workflow`` 下(见 ``app/api/v1/workflow.py``)。
|
||||
"""
|
||||
@@ -0,0 +1,667 @@
|
||||
"""存储协议适配器抽象基类
|
||||
|
||||
子类只需实现同步的 _sync_* 方法(协议相关逻辑),异步公开接口统一在
|
||||
基类中经 asyncio.to_thread 包装提供,避免阻塞事件循环。
|
||||
可选钩子:_sync_get_url(预签名 URL,不支持则省略,get_url 返回 None)、
|
||||
_sync_close(释放连接资源,无持久连接则省略,close 为空操作)。
|
||||
适配器实例按请求创建(无连接池复用),用完由调用方关闭。
|
||||
"""
|
||||
import asyncio
|
||||
import os
|
||||
import types
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import Any, Literal, get_args, get_origin
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.utils.crypto_util import CryptoUtil
|
||||
|
||||
|
||||
class StorageProtocolDefSchema(BaseModel):
|
||||
"""存储协议定义(/protocols 接口返回)"""
|
||||
|
||||
protocol: str = Field(..., description="协议标识")
|
||||
name: str = Field(..., description="协议名称")
|
||||
default_port: int = Field(..., description="默认端口")
|
||||
|
||||
|
||||
class AdvancedFieldDefSchema(BaseModel):
|
||||
"""高级配置字段定义元数据(/advanced-fields 接口返回,前端按类型动态渲染)"""
|
||||
|
||||
key: str = Field(..., description="字段名")
|
||||
label: str = Field(..., description="字段显示名")
|
||||
default: Any | None = Field(default=None, description="默认值")
|
||||
type: Literal["boolean", "number", "text", "select"] = Field("text", description="组件类型")
|
||||
options: list[str] | None = Field(default=None, description="下拉选项")
|
||||
|
||||
|
||||
class StorageProtocol(str, Enum):
|
||||
"""存储协议枚举"""
|
||||
|
||||
FTP = "ftp"
|
||||
FTPS = "ftps"
|
||||
SFTP = "sftp"
|
||||
S3 = "s3"
|
||||
OBS = "obs"
|
||||
OSS = "oss"
|
||||
COS = "cos"
|
||||
LOCAL = "local"
|
||||
|
||||
|
||||
# 各协议默认端口
|
||||
DEFAULT_PORTS: dict[StorageProtocol, int] = {
|
||||
StorageProtocol.FTP: 21,
|
||||
StorageProtocol.FTPS: 990,
|
||||
StorageProtocol.SFTP: 22,
|
||||
StorageProtocol.S3: 443,
|
||||
StorageProtocol.OBS: 443,
|
||||
StorageProtocol.OSS: 443,
|
||||
StorageProtocol.COS: 443,
|
||||
StorageProtocol.LOCAL: 0,
|
||||
}
|
||||
|
||||
|
||||
def encrypt_password(plain: str | None) -> str:
|
||||
"""明文密码 → 密文。空值原样返回空串。
|
||||
|
||||
统一走 CryptoUtil(独立数据加密密钥 + 支持轮换),不再自行从
|
||||
SECRET_KEY 派生,避免 JWT 签名密钥与落库敏感数据的加密密钥同源。
|
||||
"""
|
||||
return CryptoUtil.encrypt(plain)
|
||||
|
||||
|
||||
def decrypt_password(cipher: str | None) -> str:
|
||||
"""密文 → 明文密码。空值原样返回空串,解不开时抛业务异常。"""
|
||||
try:
|
||||
return CryptoUtil.decrypt(cipher)
|
||||
except CustomException as e:
|
||||
raise CustomException(msg=f"存储源密码解密失败:{e!s}")
|
||||
|
||||
|
||||
# 文件系统类协议与对象存储类协议的划分:仅用于文档说明与配置元数据,
|
||||
# 浏览/操作统一以"存储真根"为基准,path_prefix 不参与任何 key 计算。
|
||||
_FILE_SYSTEM_PROTOCOLS: frozenset[StorageProtocol] = frozenset(
|
||||
{StorageProtocol.FTP, StorageProtocol.FTPS, StorageProtocol.SFTP, StorageProtocol.LOCAL}
|
||||
)
|
||||
_OBJECT_STORE_PROTOCOLS: frozenset[StorageProtocol] = frozenset(
|
||||
{StorageProtocol.S3, StorageProtocol.OSS, StorageProtocol.COS, StorageProtocol.OBS}
|
||||
)
|
||||
|
||||
|
||||
# 流式传输阈值:SDK 高层上传超过该大小才触发分片(100GB,等效单次上传)
|
||||
_STREAM_THRESHOLD = 100 * 1024 * 1024 * 1024
|
||||
|
||||
|
||||
class StorageAdapterConfig(BaseModel):
|
||||
"""存储适配器配置(从 StorageSourceModel 剥离加密字段后注入,解耦 ORM 与协议层)"""
|
||||
|
||||
protocol: StorageProtocol = Field(description="存储协议")
|
||||
host: str | None = Field(default=None, description="主机地址(对象存储可不填)")
|
||||
port: int = Field(description="端口")
|
||||
username: str | None = Field(default=None, description="用户名/AccessKey")
|
||||
password: str | None = Field(default=None, description="密码/SecretKey(已解密)")
|
||||
bucket: str | None = Field(default=None, description="桶名(对象存储专用;FTP/SFTP/LOCAL 不使用)")
|
||||
endpoint: str | None = Field(default=None, description="接入点(对象存储)")
|
||||
scheme: str = Field(default="https", description="对象存储访问协议(http/https)")
|
||||
region: str | None = Field(default=None, description="区域(对象存储)")
|
||||
path_prefix: str | None = Field(default=None, description="路径前缀(仅作节点元数据/传输默认路径,不参与浏览与 key 计算)")
|
||||
is_secure: bool = Field(default=False, description="是否启用 TLS(FTPS 显示加密 / COS 访问协议)")
|
||||
implicit_tls: bool = Field(default=False, description="FTPS 是否隐式 TLS(默认显式)")
|
||||
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 等)")
|
||||
# 分片传输参数(对象存储分片上传,MB 单位;每个端点独立配置;单片上限 5GB 对齐各对象存储 SDK 上限)
|
||||
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)")
|
||||
# 传输方式(连线/任务级覆盖):stream 不触发分片(等效单次上传),multipart 使用分片参数;空=按分片参数走 SDK 高层上传
|
||||
transfer_mode: Literal["stream", "multipart"] | None = Field(default=None, description="传输方式")
|
||||
# SDK 高级配置(原始 JSON,由各适配器按协议解析合并默认值)
|
||||
advanced_config: dict = Field(default_factory=dict, description="SDK高级配置")
|
||||
|
||||
|
||||
def normalize_endpoint(endpoint: str | None, scheme: str = "https") -> str | None:
|
||||
"""规范化对象存储接入点:SDK 要求完整 URL(带协议),缺失时按 scheme 自动补充。
|
||||
|
||||
例如 ``s3.example.com`` -> ``https://s3.example.com``。
|
||||
"""
|
||||
if not endpoint:
|
||||
return endpoint
|
||||
return endpoint if "://" in endpoint else f"{scheme or 'https'}://{endpoint}"
|
||||
|
||||
|
||||
class StorageObject(BaseModel):
|
||||
"""远端文件对象信息"""
|
||||
|
||||
name: str = Field(description="文件/目录名")
|
||||
key: str = Field(description="相对存储真根的路径(path_prefix 仅作为普通目录层,不做剥离)")
|
||||
is_dir: bool = Field(default=False, description="是否目录")
|
||||
size: int | None = Field(default=None, description="大小(字节)")
|
||||
modified_time: datetime | None = Field(default=None, description="修改时间")
|
||||
|
||||
|
||||
class StoragePage(BaseModel):
|
||||
"""游标分页结果(对象存储 SDK 原生游标分页,无总条数,只能顺序翻页)"""
|
||||
|
||||
items: list[StorageObject] = Field(description="当前页条目")
|
||||
has_next: bool = Field(default=False, description="是否还有下一页")
|
||||
next_cursor: str | None = Field(default=None, description="下一页游标(无下一页为 None)")
|
||||
|
||||
|
||||
class BaseStorageAdapter(ABC):
|
||||
"""存储协议适配器抽象基类
|
||||
|
||||
子类实现同步 _sync_* 方法即可,基类统一经 asyncio.to_thread 包装为异步接口。
|
||||
可选钩子:_sync_get_url(生成预签名 URL,不支持则省略,get_url 返回 None)、
|
||||
_sync_close(释放连接资源,无持久连接则省略,close 为空操作)。
|
||||
适配器实例按请求创建(无连接池复用),用完由调用方关闭。
|
||||
"""
|
||||
|
||||
# 并发安全:持共享连接的协议(如 SFTP)不支持多线程并发读写同一连接
|
||||
concurrency_safe: bool = True
|
||||
|
||||
def __init__(self, config: StorageAdapterConfig) -> None:
|
||||
self.config = config
|
||||
# 当前操作桶(对象存储):默认取节点配置,多桶浏览时经 set_bucket 覆盖
|
||||
self.bucket_name = config.bucket or ""
|
||||
|
||||
def set_bucket(self, name: str) -> None:
|
||||
"""切换当前操作桶(对象存储多桶浏览用)。"""
|
||||
self.bucket_name = name
|
||||
|
||||
def _require_bucket(self) -> str:
|
||||
"""获取当前操作桶,未配置时抛出异常。"""
|
||||
if not self.bucket_name:
|
||||
raise CustomException(msg="存储源未配置 bucket")
|
||||
return self.bucket_name
|
||||
|
||||
def _sync_list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶(对象存储专用)。FTP/SFTP/LOCAL 等协议无需实现。"""
|
||||
raise NotImplementedError(f"{type(self).__name__} 不支持桶列表")
|
||||
|
||||
def _multipart_settings(self) -> tuple[int, int, bool]:
|
||||
"""返回 (分片大小, 并发数, 是否分片),分片大小单位字节;并发按内存预算护栏收敛。
|
||||
|
||||
源项目运行设置语义:分片大小默认 50MB、并发默认 6、内存预算默认 512MB。
|
||||
仅对象存储分片上传使用,FTP/SFTP/LOCAL 等协议忽略。
|
||||
传输方式为 stream 时返回超大阈值并关闭分片(等效单次上传)。
|
||||
"""
|
||||
if self.config.transfer_mode == "stream":
|
||||
return _STREAM_THRESHOLD, 1, False
|
||||
part_size = max(self.config.multipart_part_size, 5) * 1024 * 1024
|
||||
concurrency = max(self.config.multipart_concurrency, 1)
|
||||
budget = self.config.multipart_memory_budget * 1024 * 1024
|
||||
if budget > 0:
|
||||
max_by_budget = max(1, budget // part_size)
|
||||
concurrency = min(concurrency, max_by_budget)
|
||||
return part_size, concurrency, True
|
||||
|
||||
# ── 同步协议操作(子类实现,经 asyncio.to_thread 包装对外)─────────
|
||||
|
||||
@abstractmethod
|
||||
def _sync_test_connection(self) -> bool:
|
||||
"""测试连接是否可用。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
"""上传本地文件到远端,返回远端完整 key。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
"""下载远端文件到本地,返回本地路径。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
"""删除远端文件(目录递归删除由 delete_dir 处理)。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
"""判断远端文件是否存在。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
"""列出远端目录下的文件与目录(不含前缀)。"""
|
||||
|
||||
# ── 目录操作(子类实现,经 asyncio.to_thread 包装对外)─────────────
|
||||
|
||||
@abstractmethod
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
"""创建目录(remote_dir 为相对存储真根的目录路径)。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
"""删除空目录(对象存储删除占位标记)。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动(src/dst 为相对存储真根的路径,支持文件与目录)。"""
|
||||
|
||||
@abstractmethod
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制文件或目录(src/dst 为相对存储真根的路径)。"""
|
||||
|
||||
# ── 默认实现(子类可按协议优化覆盖)───────────────────────────────
|
||||
|
||||
def _sync_list_recursive(self, prefix: str) -> list[StorageObject]:
|
||||
"""递归列出(prefix 为相对存储真根的目录路径)。默认深度优先遍历;对象存储可覆盖为单次 API。"""
|
||||
result: list[StorageObject] = []
|
||||
stack = [prefix]
|
||||
while stack:
|
||||
current = stack.pop()
|
||||
for entry in self._sync_list(current):
|
||||
if entry.is_dir:
|
||||
stack.append(f"{current}/{entry.name}" if current else entry.name)
|
||||
result.append(entry)
|
||||
return result
|
||||
|
||||
def _sync_delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除目录(含目录本身)。默认:先删文件,再自底向上删空目录。"""
|
||||
for entry in reversed(self._sync_list_recursive(remote_dir)):
|
||||
if entry.is_dir:
|
||||
try:
|
||||
self._sync_rmdir(entry.key)
|
||||
except Exception as e:
|
||||
logger.warning("删除远端目录失败(将残留空目录): {}: {}", entry.key, e)
|
||||
else:
|
||||
self._sync_delete(entry.key)
|
||||
if remote_dir:
|
||||
try:
|
||||
self._sync_rmdir(remote_dir)
|
||||
except Exception as e:
|
||||
logger.warning("删除远端目录失败(将残留空目录): {}: {}", remote_dir, e)
|
||||
|
||||
def _sync_copy_dir(self, src: str, dst: str) -> None:
|
||||
"""递归复制目录(src/dst 为相对存储真根的路径):遍历后先建目录再逐文件复制。"""
|
||||
src = src.rstrip("/")
|
||||
for entry in self._sync_list_recursive(src):
|
||||
full = entry.key
|
||||
rel = full[len(src) + 1 :] if full.startswith(src + "/") else entry.key
|
||||
target = f"{dst}/{rel}".strip("/")
|
||||
if entry.is_dir:
|
||||
self._sync_mkdir(target)
|
||||
else:
|
||||
self._sync_copy(full, target)
|
||||
|
||||
def _sync_move_dir(self, src: str, dst: str) -> None:
|
||||
"""递归移动目录(对象存储无原生目录重命名):先复制后删除源。"""
|
||||
self._sync_copy_dir(src, dst)
|
||||
self._sync_delete_dir(src)
|
||||
|
||||
def _sync_upload_dir(self, local_dir: str, remote_dir: str, concurrency: int) -> int:
|
||||
"""批量上传目录:os.walk 收集文件后按 concurrency 并发上传,返回上传文件数。"""
|
||||
if not os.path.isdir(local_dir):
|
||||
raise CustomException(msg=f"本地目录不存在: {local_dir}")
|
||||
base = os.path.normpath(local_dir)
|
||||
tasks: list[tuple[str, str]] = []
|
||||
for root, _dirs, names in os.walk(base):
|
||||
for name in names:
|
||||
local_path = os.path.join(root, name)
|
||||
rel = os.path.relpath(local_path, base).replace(os.sep, "/")
|
||||
remote_path = f"{remote_dir}/{rel}".strip("/") if remote_dir else rel
|
||||
tasks.append((local_path, remote_path))
|
||||
if not tasks:
|
||||
return 0
|
||||
errors: list[Exception] = []
|
||||
|
||||
def _run(task: tuple[str, str]) -> None:
|
||||
try:
|
||||
self._sync_upload(task[0], task[1])
|
||||
except Exception as e:
|
||||
errors.append(e)
|
||||
|
||||
workers = max(1, concurrency) if self.concurrency_safe else 1
|
||||
if workers == 1:
|
||||
for task in tasks:
|
||||
_run(task)
|
||||
else:
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
list(pool.map(_run, tasks))
|
||||
if errors:
|
||||
raise CustomException(msg=f"批量上传失败 {len(errors)}/{len(tasks)} 个文件,首个错误: {errors[0]!s}")
|
||||
return len(tasks)
|
||||
|
||||
def _sync_download_dir(self, remote_dir: str, local_dir: str, concurrency: int) -> int:
|
||||
"""批量下载目录:递归列出后按 concurrency 并发下载,返回下载文件数。"""
|
||||
entries = [e for e in self._sync_list_recursive(remote_dir) if not e.is_dir]
|
||||
base = remote_dir.strip("/")
|
||||
tasks: list[tuple[str, str]] = []
|
||||
for e in entries:
|
||||
rel = e.key[len(base) + 1 :] if base and e.key.startswith(base + "/") else e.key
|
||||
tasks.append((e.key, os.path.join(local_dir, rel.replace("/", os.sep))))
|
||||
if not tasks:
|
||||
return 0
|
||||
errors: list[Exception] = []
|
||||
|
||||
def _run(task: tuple[str, str]) -> None:
|
||||
try:
|
||||
os.makedirs(os.path.dirname(task[1]), exist_ok=True)
|
||||
self._sync_download(task[0], task[1])
|
||||
except Exception as e:
|
||||
errors.append(e)
|
||||
|
||||
workers = max(1, concurrency) if self.concurrency_safe else 1
|
||||
if workers == 1:
|
||||
for task in tasks:
|
||||
_run(task)
|
||||
else:
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
list(pool.map(_run, tasks))
|
||||
if errors:
|
||||
raise CustomException(msg=f"批量下载失败 {len(errors)}/{len(tasks)} 个文件,首个错误: {errors[0]!s}")
|
||||
return len(tasks)
|
||||
|
||||
@staticmethod
|
||||
def _entries_from_keys(keys: list[str]) -> list[StorageObject]:
|
||||
"""由扁平 key 列表构造含隐含目录的条目列表(对象存储递归列举用,key 已剥离前缀)。"""
|
||||
dirs: set[str] = set()
|
||||
files: set[str] = set()
|
||||
for raw in keys:
|
||||
key = raw.rstrip("/")
|
||||
if not key:
|
||||
continue
|
||||
(dirs if raw.endswith("/") else files).add(key)
|
||||
parts = key.split("/")
|
||||
for i in range(1, len(parts)):
|
||||
dirs.add("/".join(parts[:i]))
|
||||
result: list[StorageObject] = []
|
||||
for k in sorted(dirs):
|
||||
result.append(StorageObject(name=k.rsplit("/", 1)[-1], key=k, is_dir=True))
|
||||
for k in sorted(files):
|
||||
result.append(StorageObject(name=k.rsplit("/", 1)[-1], key=k, is_dir=False))
|
||||
return result
|
||||
|
||||
# ── 异步公开接口(经 asyncio.to_thread 包装)─────────────────────
|
||||
|
||||
async def test_connection(self) -> bool:
|
||||
return await asyncio.to_thread(self._sync_test_connection)
|
||||
|
||||
async def upload(self, local_path: str, remote_path: str) -> str:
|
||||
return await asyncio.to_thread(self._sync_upload, local_path, remote_path)
|
||||
|
||||
async def download(self, remote_path: str, local_path: str) -> str:
|
||||
return await asyncio.to_thread(self._sync_download, remote_path, local_path)
|
||||
|
||||
async def delete(self, remote_path: str) -> None:
|
||||
await asyncio.to_thread(self._sync_delete, remote_path)
|
||||
|
||||
async def exists(self, remote_path: str) -> bool:
|
||||
return await asyncio.to_thread(self._sync_exists, remote_path)
|
||||
|
||||
async def list_files(
|
||||
self, prefix: str = "", page_size: int | None = None, cursor: str | None = None
|
||||
) -> list[StorageObject] | StoragePage:
|
||||
"""列出目录条目。
|
||||
|
||||
- 不传 page_size:全量拉取(目录选择器/搜索等场景)。
|
||||
- 传 page_size:游标分页。对象存储子类重写 _sync_list_page 走 SDK 原生游标
|
||||
(每页只拉一页数据);FTP/SFTP/LOCAL 等无服务器游标的协议用基类默认实现
|
||||
(全量列举 + 内存切片,游标即偏移量)。
|
||||
"""
|
||||
if page_size is not None:
|
||||
return await asyncio.to_thread(self._sync_list_page, prefix, page_size, cursor)
|
||||
return await asyncio.to_thread(self._sync_list, prefix)
|
||||
|
||||
def _sync_list_page(self, prefix: str, page_size: int, cursor: str | None) -> StoragePage:
|
||||
"""游标分页默认实现:全量列举后内存切片(游标为字符串偏移量,FTP/SFTP/LOCAL 使用)。"""
|
||||
items = self._sync_list(prefix)
|
||||
start = 0
|
||||
if cursor:
|
||||
try:
|
||||
start = int(cursor)
|
||||
except ValueError:
|
||||
start = 0
|
||||
page = items[start : start + page_size]
|
||||
has_next = start + page_size < len(items)
|
||||
return StoragePage(
|
||||
items=page,
|
||||
has_next=has_next,
|
||||
next_cursor=str(start + page_size) if has_next else None,
|
||||
)
|
||||
|
||||
async def list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶(对象存储专用)。"""
|
||||
return await asyncio.to_thread(self._sync_list_buckets)
|
||||
|
||||
async def list_recursive(self, prefix: str = "") -> list[StorageObject]:
|
||||
"""递归列出目录树(含隐含目录条目)。"""
|
||||
return await asyncio.to_thread(self._sync_list_recursive, prefix)
|
||||
|
||||
async def mkdir(self, remote_dir: str) -> None:
|
||||
"""创建目录。"""
|
||||
await asyncio.to_thread(self._sync_mkdir, remote_dir)
|
||||
|
||||
async def rmdir(self, remote_dir: str) -> None:
|
||||
"""删除空目录。"""
|
||||
await asyncio.to_thread(self._sync_rmdir, remote_dir)
|
||||
|
||||
async def delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除目录(含目录本身)。"""
|
||||
await asyncio.to_thread(self._sync_delete_dir, remote_dir)
|
||||
|
||||
async def rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动(文件或目录)。"""
|
||||
await asyncio.to_thread(self._sync_rename, src, dst)
|
||||
|
||||
async def copy(self, src: str, dst: str) -> None:
|
||||
"""复制(文件或目录)。"""
|
||||
await asyncio.to_thread(self._sync_copy, src, dst)
|
||||
|
||||
async def upload_dir(self, local_dir: str, remote_dir: str = "", concurrency: int = 1) -> int:
|
||||
"""批量上传本地目录到远端,返回上传文件数。"""
|
||||
return await asyncio.to_thread(self._sync_upload_dir, local_dir, remote_dir, concurrency)
|
||||
|
||||
async def download_dir(self, remote_dir: str, local_dir: str, concurrency: int = 1) -> int:
|
||||
"""批量下载远端目录到本地,返回下载文件数。"""
|
||||
return await asyncio.to_thread(self._sync_download_dir, remote_dir, local_dir, concurrency)
|
||||
|
||||
async def get_url(self, remote_path: str, expire: int = 3600) -> str | None:
|
||||
"""获取访问 URL(对象存储返回预签名 URL;FTP/SFTP/LOCAL 不支持返回 None)。"""
|
||||
sync = getattr(self, "_sync_get_url", None)
|
||||
if sync is None:
|
||||
return None
|
||||
return await asyncio.to_thread(sync, remote_path, expire)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""释放连接资源(无持久连接的协议实现为空操作)。"""
|
||||
sync = getattr(self, "_sync_close", None)
|
||||
if sync is None:
|
||||
return
|
||||
await asyncio.to_thread(sync)
|
||||
|
||||
|
||||
class S3AdvancedConfig(BaseModel):
|
||||
"""S3 兼容对象存储(boto3/botocore Config)"""
|
||||
|
||||
connect_timeout: int = Field(default=60, ge=1, le=300, description="连接超时(秒)")
|
||||
read_timeout: int = Field(default=60, ge=1, le=300, description="读取超时(秒)")
|
||||
max_attempts: int = Field(default=3, ge=1, le=10, description="最大重试次数")
|
||||
retries_mode: str = Field(default="standard", description="重试模式")
|
||||
max_pool_connections: int = Field(default=10, ge=1, le=100, description="连接池大小")
|
||||
tcp_keepalive: bool = Field(default=False, description="TCP 长连接保活")
|
||||
use_dualstack_endpoint: bool = Field(default=False, description="使用双栈端点(IPv4/IPv6)")
|
||||
signature_version: str = Field(default="s3v4", description="签名版本")
|
||||
addressing_style: str = Field(default="auto", description="寻址风格")
|
||||
proxies: str | None = Field(default=None, max_length=255, description="代理地址(http://host:port,HTTP/HTTPS 通用)")
|
||||
use_accelerate_endpoint: bool = Field(default=False, description="使用传输加速端点")
|
||||
parameter_validation: bool = Field(default=True, description="启用请求参数校验(关闭可提升性能)")
|
||||
us_east_1_regional_endpoint: str | None = Field(default=None, description="us-east-1 区域端点策略(regional/legacy)")
|
||||
request_checksum_calculation: str = Field(default="when_supported", description="请求校验和计算时机(when_supported/when_required)")
|
||||
response_checksum_validation: str = Field(default="when_supported", description="响应校验和验证时机(when_supported/when_required)")
|
||||
|
||||
|
||||
class OssAdvancedConfig(BaseModel):
|
||||
"""阿里云 OSS(alibabacloud_oss_v2 Config)"""
|
||||
|
||||
connect_timeout: int = Field(default=10, ge=1, le=300, description="连接超时(秒)")
|
||||
readwrite_timeout: int = Field(default=20, ge=1, le=300, description="读写超时(秒)")
|
||||
retry_max_attempts: int = Field(default=3, ge=1, le=10, description="最大重试次数")
|
||||
use_cname: bool = Field(default=False, description="使用自定义域名(CNAME)访问")
|
||||
use_path_style: bool = Field(default=False, description="路径风格寻址(兼容自建服务)")
|
||||
use_internal_endpoint: bool = Field(default=False, description="使用内网访问")
|
||||
use_accelerate_endpoint: bool = Field(default=False, description="使用传输加速端点")
|
||||
use_dualstack_endpoint: bool = Field(default=False, description="使用双栈端点(IPv4/IPv6)")
|
||||
insecure_skip_verify: bool = Field(default=False, description="跳过服务端证书校验")
|
||||
proxy_host: str | None = Field(default=None, max_length=255, description="代理服务器地址")
|
||||
signature_version: str = Field(default="v4", description="签名版本")
|
||||
disable_upload_crc64_check: bool = Field(default=False, description="关闭上传 CRC64 校验(提升性能)")
|
||||
disable_download_crc64_check: bool = Field(default=False, description="关闭下载 CRC64 校验(提升性能)")
|
||||
enabled_redirect: bool = Field(default=False, description="启用 HTTP 重定向")
|
||||
|
||||
|
||||
class CosAdvancedConfig(BaseModel):
|
||||
"""腾讯云 COS(cos-python-sdk-v5:CosConfig + CosS3Client retry)"""
|
||||
|
||||
appid: str | None = Field(default=None, max_length=64, description="账号 Appid(部分场景必需)")
|
||||
token: str | None = Field(default=None, max_length=1024, description="临时密钥 Token(临时密钥认证时填写)")
|
||||
endpoint: str | None = Field(default=None, max_length=255, description="自定义接入域名(默认按 region 解析)")
|
||||
timeout: int = Field(default=60, ge=1, le=300, description="请求超时(秒)")
|
||||
retry: int = Field(default=3, ge=0, le=10, description="最大重试次数")
|
||||
enable_md5: bool = Field(default=False, description="分片上传启用 MD5 校验")
|
||||
verify_ssl: bool = Field(default=True, description="验证服务端证书")
|
||||
auto_switch_domain_on_retry: bool = Field(default=False, description="重试时自动切换域名")
|
||||
pool_connections: int = Field(default=10, ge=1, le=100, description="连接池大小")
|
||||
pool_max_size: int = Field(default=100, ge=10, le=1000, description="连接池最大容量")
|
||||
keep_alive: bool = Field(default=True, description="HTTP 连接保活")
|
||||
allow_redirects: bool = Field(default=True, description="允许 HTTP 重定向")
|
||||
http_proxy: str | None = Field(default=None, max_length=255, description="HTTP 代理")
|
||||
https_proxy: str | None = Field(default=None, max_length=255, description="HTTPS 代理")
|
||||
|
||||
|
||||
class ObsAdvancedConfig(BaseModel):
|
||||
"""华为云 OBS(esdk-obs-python ObsClient,3.26.x 无独立连接超时参数)"""
|
||||
|
||||
timeout: int = Field(default=60, ge=10, le=600, description="请求超时(秒)")
|
||||
max_retry_count: int = Field(default=3, ge=1, le=5, description="最大重试次数")
|
||||
ssl_verify: bool = Field(default=False, description="验证服务端证书")
|
||||
is_cname: bool = Field(default=False, description="使用自定义域名访问")
|
||||
path_style: bool = Field(default=False, description="路径风格寻址(兼容自建服务)")
|
||||
pool_size: int = Field(default=10, ge=1, le=200, description="连接池大小")
|
||||
signature: str = Field(default="v4", description="签名类型(v2/v4/obs)")
|
||||
region: str | None = Field(default=None, max_length=64, description="区域(部分场景必需)")
|
||||
security_token: str | None = Field(default=None, max_length=1024, description="临时密钥 Token(临时密钥认证时填写)")
|
||||
|
||||
|
||||
class SftpAdvancedConfig(BaseModel):
|
||||
"""SFTP(paramiko SSHClient.connect)"""
|
||||
|
||||
connect_timeout: int = Field(default=30, ge=1, le=300, description="连接超时(秒)")
|
||||
banner_timeout: int = Field(default=30, ge=1, le=300, description="横幅超时(秒)")
|
||||
auth_timeout: int = Field(default=30, ge=1, le=300, description="认证超时(秒)")
|
||||
channel_timeout: int = Field(default=30, ge=1, le=300, description="通道操作超时(秒)")
|
||||
compress: bool = Field(default=False, description="启用传输压缩")
|
||||
allow_agent: bool = Field(default=True, description="允许使用 ssh-agent 认证")
|
||||
look_for_keys: bool = Field(default=True, description="查找本地密钥文件认证")
|
||||
keepalive_interval: int = Field(default=0, ge=0, le=3600, description="心跳间隔(秒,0 关闭)")
|
||||
|
||||
|
||||
class FtpAdvancedConfig(BaseModel):
|
||||
"""FTP / FTPS(ftplib)"""
|
||||
|
||||
timeout: int = Field(default=30, ge=1, le=300, description="连接超时(秒)")
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 各协议数据完整性校验能力对照(通常用于传输后的数据完整性保障,平时保持默认即可)
|
||||
#
|
||||
# S3 : 请求/响应 Checksum(CRC32/CRC32C 等算法)
|
||||
# 由 botocore 控制计算与验证时机:
|
||||
# - request_checksum_calculation: when_supported / when_required
|
||||
# - response_checksum_validation: when_supported / when_required
|
||||
# (注意:当前 botocore 版本运行时不接受文档中的 "never" 值)
|
||||
# OSS : CRC64(默认开启):
|
||||
# - disable_upload_crc64_check : 关闭上传校验(性能敏感时可开)
|
||||
# - disable_download_crc64_check : 关闭下载校验(性能敏感时可开)
|
||||
# COS : MD5(分片上传,默认关闭):
|
||||
# - enable_md5: 开启后分片上传附带 MD5 校验,可靠性提升、略有性能开销
|
||||
# OBS : SDK 自动校验(esdk-obs-python 无用户可控开关,内部计算并在失败时报错)
|
||||
#
|
||||
# 提示:以上校验字段均为"隐藏级"配置(不在前端高级配置面板展示),
|
||||
# 确有需要时通过接口在 advanced_config 中传入。
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 各协议传输加速能力对照(用于跨地域大文件传输加速,需先在云控制台开通)
|
||||
#
|
||||
# S3 : use_accelerate_endpoint(botocore 原生开关,advanced_config 隐藏字段)
|
||||
# OSS : use_accelerate_endpoint(oss_v2 原生开关,advanced_config 隐藏字段)
|
||||
# COS : 无 SDK 开关,改用加速域名接入:
|
||||
# endpoint 填 {bucket}.cos.accelerate.myqcloud.com
|
||||
# OBS : 无 SDK 开关,改用 CDN/自定义域名接入:
|
||||
# server 填 CDN 加速域名 + is_cname=true(advanced_config 隐藏字段)
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 协议 → 高级配置模型(local 无高级参数)
|
||||
ADVANCED_CONFIG_MODELS: dict[StorageProtocol, type[BaseModel]] = {
|
||||
StorageProtocol.S3: S3AdvancedConfig,
|
||||
StorageProtocol.OSS: OssAdvancedConfig,
|
||||
StorageProtocol.COS: CosAdvancedConfig,
|
||||
StorageProtocol.OBS: ObsAdvancedConfig,
|
||||
StorageProtocol.SFTP: SftpAdvancedConfig,
|
||||
StorageProtocol.FTP: FtpAdvancedConfig,
|
||||
StorageProtocol.FTPS: FtpAdvancedConfig,
|
||||
}
|
||||
|
||||
|
||||
def parse_advanced_config(protocol: StorageProtocol, raw: dict[str, Any] | None) -> BaseModel | None:
|
||||
"""解析端点高级配置 JSON 为对应协议模型;local 或无配置返回 None。"""
|
||||
model_cls = ADVANCED_CONFIG_MODELS.get(protocol)
|
||||
if model_cls is None:
|
||||
return None
|
||||
return model_cls.model_validate(raw or {})
|
||||
|
||||
|
||||
# 前端渲染元数据:协议 → 字段 key → 可选值(按协议区分同名枚举,如 signature_version)
|
||||
_SELECT_OPTIONS: dict[str, dict[str, list[str]]] = {
|
||||
"s3": {
|
||||
"signature_version": ["s3v4", "s3", "v2"],
|
||||
"addressing_style": ["auto", "path", "virtual"],
|
||||
"retries_mode": ["standard", "adaptive", "legacy"],
|
||||
"us_east_1_regional_endpoint": ["regional", "legacy"],
|
||||
"request_checksum_calculation": ["when_supported", "when_required"],
|
||||
"response_checksum_validation": ["when_supported", "when_required"],
|
||||
},
|
||||
"oss": {"signature_version": ["v4", "v1"]},
|
||||
"obs": {"signature": ["v2", "v4", "obs"]},
|
||||
}
|
||||
|
||||
|
||||
def _field_defs(model_cls: type[BaseModel], protocol_key: str) -> list[AdvancedFieldDefSchema]:
|
||||
defs: list[AdvancedFieldDefSchema] = []
|
||||
protocol_options = _SELECT_OPTIONS.get(protocol_key, {})
|
||||
for name, field_info in model_cls.model_fields.items():
|
||||
default = field_info.default
|
||||
if isinstance(default, BaseModel):
|
||||
default = None
|
||||
annotation = field_info.annotation
|
||||
if annotation is bool:
|
||||
field_type = "boolean"
|
||||
elif annotation is int:
|
||||
field_type = "number"
|
||||
elif annotation is str:
|
||||
field_type = "text"
|
||||
elif get_origin(annotation) is types.UnionType:
|
||||
# Optional[X]:取非 None 的分支推断组件类型
|
||||
inner = next((a for a in get_args(annotation) if a is not type(None)), str)
|
||||
field_type = "number" if inner is int else "text"
|
||||
else:
|
||||
field_type = "text"
|
||||
options = protocol_options.get(name)
|
||||
if options is not None:
|
||||
field_type = "select"
|
||||
defs.append(
|
||||
AdvancedFieldDefSchema(
|
||||
key=name,
|
||||
label=field_info.description or name,
|
||||
default=default,
|
||||
type=field_type,
|
||||
options=options,
|
||||
)
|
||||
)
|
||||
return defs
|
||||
|
||||
|
||||
ADVANCED_FIELD_DEFS: dict[str, list[AdvancedFieldDefSchema]] = {
|
||||
protocol.value: _field_defs(model_cls, protocol.value) for protocol, model_cls in ADVANCED_CONFIG_MODELS.items()
|
||||
}
|
||||
@@ -0,0 +1,364 @@
|
||||
from typing import Any
|
||||
|
||||
from qcloud_cos import CosConfig, CosS3Client
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import (
|
||||
BaseStorageAdapter,
|
||||
CosAdvancedConfig,
|
||||
StorageObject,
|
||||
StoragePage,
|
||||
StorageProtocol,
|
||||
)
|
||||
|
||||
|
||||
class CosStorageAdapter(BaseStorageAdapter):
|
||||
"""腾讯云 COS 存储适配器(cos-python-sdk-v5,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.COS
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
if not self.config.region:
|
||||
raise CustomException(msg="COS 存储源必须配置 region")
|
||||
if not self.config.username or not self.config.password:
|
||||
raise CustomException(msg="COS 存储源必须配置 SecretId/SecretKey")
|
||||
adv = CosAdvancedConfig(**self.config.advanced_config)
|
||||
proxies: dict[str, str] = {}
|
||||
if adv.http_proxy:
|
||||
proxies["http"] = adv.http_proxy
|
||||
if adv.https_proxy:
|
||||
proxies["https"] = adv.https_proxy
|
||||
cos_config = CosConfig(
|
||||
Region=self.config.region,
|
||||
SecretId=self.config.username or "",
|
||||
SecretKey=self.config.password or "",
|
||||
Scheme="https" if self.config.is_secure else "http",
|
||||
Token=adv.token,
|
||||
Appid=adv.appid,
|
||||
Endpoint=adv.endpoint,
|
||||
Timeout=adv.timeout,
|
||||
Proxies=proxies or None,
|
||||
VerifySSL=adv.verify_ssl,
|
||||
AutoSwitchDomainOnRetry=adv.auto_switch_domain_on_retry,
|
||||
PoolConnections=adv.pool_connections,
|
||||
PoolMaxSize=adv.pool_max_size,
|
||||
KeepAlive=adv.keep_alive,
|
||||
AllowRedirects=adv.allow_redirects,
|
||||
)
|
||||
self.client = CosS3Client(cos_config, retry=adv.retry)
|
||||
self._adv = adv
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
"""
|
||||
列出符合条件的bucket
|
||||
:param Bucket(string): 存储桶名称
|
||||
:param TagKey(string): 标签键
|
||||
:param TagValue(string): 标签值
|
||||
:param Region(string): 地域名称
|
||||
:param CreateTime(Timestamp): GMT时间戳, 和 Range 参数一起使用, 支持根据创建时间过滤存储桶
|
||||
:param Range(string): 和 CreateTime 参数一起使用, 支持根据创建时间过滤存储桶,支持枚举值 lt(创建时间早于 create-time)、gt(创建时间晚于 create-time)、lte(创建时间早于或等于 create-time)、gte(创建时间晚于或等于create-time)
|
||||
:param Marker(string): 起始标记, 从该标记之后(不含)按照 UTF-8 字典序返回存储桶条目
|
||||
:param MaxKeys(int): 单次返回最大的条目数量,默认值为2000,最大为2000
|
||||
|
||||
:return(dict): 账号下bucket相关信息.
|
||||
"""
|
||||
try:
|
||||
self.client.list_buckets()
|
||||
if self.client.head_bucket(Bucket=self._require_bucket()):
|
||||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.warning(f"COS 连接测试失败: {e}")
|
||||
return False
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
"""
|
||||
:param Bucket(string): 存储桶名称.
|
||||
:param key(string): 分块上传路径名.
|
||||
:param LocalFilePath(string): 本地文件路径名.
|
||||
:param PartSize(int): 分块的大小设置,单位为MB.
|
||||
:param MAXThread(int): 并发上传的最大线程数.
|
||||
:param EnableMD5(bool): 是否打开MD5校验.
|
||||
:param kwargs(dict): 设置请求headers.
|
||||
:return(dict): 成功上传文件的元信息.
|
||||
"""
|
||||
try:
|
||||
# 分片大小/并发取端点配置;文件小于分片大小自动单次上传(流式传输时该值为超大值,等效单次)
|
||||
part_size, concurrency, _ = self._multipart_settings()
|
||||
self.client.upload_file(
|
||||
Bucket=self._require_bucket(),
|
||||
Key=remote_path,
|
||||
LocalFilePath=local_path,
|
||||
PartSize=part_size // (1024 * 1024),
|
||||
MAXThread=concurrency,
|
||||
EnableMD5=self._adv.enable_md5,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
try:
|
||||
self.client.download_file(
|
||||
Bucket=self._require_bucket(),
|
||||
Key=remote_path,
|
||||
DestFilePath=local_path,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
"""
|
||||
单文件删除接口
|
||||
|
||||
:param Bucket(string): 存储桶名称.
|
||||
:param Key(string): COS路径.
|
||||
:param kwargs(dict): 设置请求headers.
|
||||
:return: dict.
|
||||
"""
|
||||
try:
|
||||
self.client.delete_object(Bucket=self._require_bucket(), Key=remote_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
try:
|
||||
return self.client.object_exists(Bucket=self._require_bucket(), Key=remote_path)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
"""列举目录层条目。COS 单次最多返回 1000 条,经 Marker 翻页拉全量。"""
|
||||
try:
|
||||
result: list[StorageObject] = []
|
||||
seen_dirs: set[str] = set()
|
||||
seen_files: set[str] = set()
|
||||
marker: str | None = None
|
||||
while True:
|
||||
kwargs: dict[str, Any] = {
|
||||
"Bucket": self._require_bucket(),
|
||||
"Prefix": prefix,
|
||||
"Delimiter": "/",
|
||||
}
|
||||
if marker:
|
||||
kwargs["Marker"] = marker
|
||||
resp = self.client.list_objects(**kwargs)
|
||||
for common in resp.get("CommonPrefixes", []):
|
||||
raw_key = common.get("Prefix", "").rstrip("/")
|
||||
if raw_key in seen_dirs:
|
||||
continue
|
||||
seen_dirs.add(raw_key)
|
||||
result.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in resp.get("Contents", []):
|
||||
raw_key = obj.get("Key", "")
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
if raw_key in seen_files:
|
||||
continue
|
||||
seen_files.add(raw_key)
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.get("Size"),
|
||||
modified_time=obj.get("LastModified"),
|
||||
)
|
||||
)
|
||||
if resp.get("IsTruncated"):
|
||||
marker = resp.get("NextMarker") or ""
|
||||
if not marker and resp.get("Contents"):
|
||||
marker = resp["Contents"][-1].get("Key", "") # 部分实现不返回 NextMarker 时回退最后 key
|
||||
if not marker:
|
||||
break
|
||||
else:
|
||||
break
|
||||
return result
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_page(self, prefix: str, page_size: int, cursor: str | None) -> StoragePage:
|
||||
"""游标分页:单次 SDK 请求只拉一页(COS Marker),翻页经前端回传游标。"""
|
||||
try:
|
||||
kwargs: dict[str, Any] = {
|
||||
"Bucket": self._require_bucket(),
|
||||
"Prefix": prefix,
|
||||
"Delimiter": "/",
|
||||
"MaxKeys": page_size,
|
||||
}
|
||||
if cursor:
|
||||
kwargs["Marker"] = cursor
|
||||
resp = self.client.list_objects(**kwargs)
|
||||
contents = resp.get("Contents", [])
|
||||
prefixes = resp.get("CommonPrefixes", [])
|
||||
truncated = bool(resp.get("IsTruncated"))
|
||||
# 条目数(含子目录前缀)少于请求页大小 → 必然已枚举完,忽略服务端边界误报
|
||||
if len(contents) + len(prefixes) < page_size:
|
||||
truncated = False
|
||||
next_cursor = resp.get("NextMarker") or ""
|
||||
if not next_cursor and truncated:
|
||||
# COS 带 Delimiter 时 NextMarker 可能缺失(如目录下只有子目录前缀、Contents 为空):
|
||||
# 取全部条目(子目录前缀 + 文件)字典序最后的 key 作为游标,避免翻页失效
|
||||
keys = [p.get("Prefix", "").rstrip("/") for p in prefixes if p.get("Prefix")] + [
|
||||
o.get("Key", "") for o in contents if o.get("Key")
|
||||
]
|
||||
if keys:
|
||||
next_cursor = max(keys)
|
||||
# COS 服务端在条目数恰好等于 MaxKeys 时可能保守返回 IsTruncated=true(实际已枚举完),
|
||||
# 用 MaxKeys=1 探测确认是否存在真正的下一页,避免出现空的下一页
|
||||
if truncated and next_cursor:
|
||||
try:
|
||||
probe = self.client.list_objects(
|
||||
Bucket=self._require_bucket(),
|
||||
Prefix=prefix,
|
||||
Delimiter="/",
|
||||
Marker=next_cursor,
|
||||
MaxKeys=1,
|
||||
)
|
||||
truncated = bool(probe.get("Contents")) or bool(probe.get("CommonPrefixes"))
|
||||
except Exception:
|
||||
pass # 探测失败时保持原判定
|
||||
items: list[StorageObject] = []
|
||||
for common in resp.get("CommonPrefixes", []):
|
||||
raw_key = common.get("Prefix", "").rstrip("/")
|
||||
items.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in resp.get("Contents", []):
|
||||
raw_key = obj.get("Key", "")
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
items.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.get("Size"),
|
||||
modified_time=obj.get("LastModified"),
|
||||
)
|
||||
)
|
||||
return StoragePage(
|
||||
items=items,
|
||||
has_next=truncated,
|
||||
next_cursor=next_cursor or None,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶。"""
|
||||
try:
|
||||
resp = self.client.list_buckets()
|
||||
buckets = (resp.get("Buckets") or {}).get("Bucket") or []
|
||||
return [b.get("Name", "") for b in buckets if b.get("Name")]
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 桶列表失败: {e!s}")
|
||||
|
||||
def _sync_get_url(self, remote_path: str, expire: int) -> str:
|
||||
"""生成预签名 URL
|
||||
:param Bucket(string): 存储桶名称.
|
||||
:param Key(string): COS路径.
|
||||
:param Method(string): HTTP请求的方法, 'PUT'|'POST'|'GET'|'DELETE'|'HEAD'
|
||||
:param Expired(int): 签名过期时间.
|
||||
:param Params(dict): 签入签名的参数
|
||||
:param Headers(dict): 签入签名的头部
|
||||
:param SignHost(bool): 是否将host算入签名.
|
||||
:return(string): 预先签名的URL.
|
||||
"""
|
||||
try:
|
||||
return self.client.get_presigned_url(
|
||||
Method="GET",
|
||||
Bucket=self._require_bucket(),
|
||||
Key=remote_path,
|
||||
Expired=expire,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 生成预签名 URL 失败: {e!s}")
|
||||
|
||||
# ── 目录操作(对象存储以 key/ 占位对象模拟目录)──────────────────
|
||||
|
||||
def _list_all_keys(self, prefix: str) -> list[str]:
|
||||
"""分页列举前缀下的全部对象 key。"""
|
||||
keys: list[str] = []
|
||||
marker = ""
|
||||
while True:
|
||||
resp = self.client.list_objects(Bucket=self._require_bucket(), Prefix=prefix, Marker=marker)
|
||||
keys.extend(obj["Key"] for obj in resp.get("Contents", []))
|
||||
if not resp.get("IsTruncated"):
|
||||
break
|
||||
# 未指定 delimiter 时服务端可能不返回 NextMarker,退化为当前页最后一个 key,避免死循环
|
||||
marker = resp.get("NextMarker") or (resp["Contents"][-1]["Key"] if resp.get("Contents") else "")
|
||||
if not marker:
|
||||
break
|
||||
return keys
|
||||
|
||||
def _is_dir(self, key: str) -> bool:
|
||||
"""判断 key 是否代表目录(存在占位对象或前缀下有任何对象)。"""
|
||||
try:
|
||||
resp = self.client.list_objects(Bucket=self._require_bucket(), Prefix=key.rstrip("/") + "/", MaxKeys=1)
|
||||
return bool(resp.get("Contents") or resp.get("CommonPrefixes"))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.put_object(Bucket=self._require_bucket(), Key=remote_dir.rstrip("/") + "/", Body=b"")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.delete_object(Bucket=self._require_bucket(), Key=remote_dir.rstrip("/") + "/")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 删除目录失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件用服务端 copy_object。"""
|
||||
try:
|
||||
if self._is_dir(src):
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
self.client.copy_object(
|
||||
Bucket=self._require_bucket(),
|
||||
Key=dst,
|
||||
CopySource={"Bucket": self._require_bucket(), "Key": src},
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 复制失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动:目录先复制后删除;文件 copy_object 后删源。"""
|
||||
if self._is_dir(src):
|
||||
self._sync_move_dir(src, dst)
|
||||
else:
|
||||
self._sync_copy(src, dst)
|
||||
self._sync_delete(src)
|
||||
|
||||
def _sync_list_recursive(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
keys = self._list_all_keys(prefix)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 递归列表失败: {e!s}")
|
||||
return self._entries_from_keys(keys)
|
||||
|
||||
def _sync_delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除:一次列举全部对象并分批批量删除(含占位目录对象)。"""
|
||||
try:
|
||||
keys = self._list_all_keys(remote_dir)
|
||||
marker = remote_dir.rstrip("/") + "/"
|
||||
if marker not in keys:
|
||||
keys.append(marker)
|
||||
for i in range(0, len(keys), 1000):
|
||||
self.client.delete_objects(
|
||||
Bucket=self._require_bucket(),
|
||||
Delete={"Objects": [{"Key": k} for k in keys[i : i + 1000]], "Quiet": True},
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"COS 递归删除失败: {e!s}")
|
||||
@@ -0,0 +1,32 @@
|
||||
from app.core.exceptions import CustomException
|
||||
from app.modules.workflow.core.base import BaseStorageAdapter, StorageAdapterConfig, StorageProtocol
|
||||
from app.modules.workflow.core.cos_adapter import CosStorageAdapter
|
||||
from app.modules.workflow.core.ftp_ftps_adapter import FtpStorageAdapter
|
||||
from app.modules.workflow.core.local_adapter import LocalStorageAdapter
|
||||
from app.modules.workflow.core.obs_adapter import ObsStorageAdapter
|
||||
from app.modules.workflow.core.oss_adapter import OssStorageAdapter
|
||||
from app.modules.workflow.core.s3_adapter import S3StorageAdapter
|
||||
from app.modules.workflow.core.sftp_adapter import SftpStorageAdapter
|
||||
|
||||
# 协议 → 适配器类映射(FTPS 复用 FTP 适配器,由配置区分显式/隐式 TLS)
|
||||
_STORAGE_ADAPTERS: dict[str, type[BaseStorageAdapter]] = {
|
||||
StorageProtocol.FTP.value: FtpStorageAdapter,
|
||||
StorageProtocol.FTPS.value: FtpStorageAdapter,
|
||||
StorageProtocol.SFTP.value: SftpStorageAdapter,
|
||||
StorageProtocol.S3.value: S3StorageAdapter,
|
||||
StorageProtocol.OBS.value: ObsStorageAdapter,
|
||||
StorageProtocol.OSS.value: OssStorageAdapter,
|
||||
StorageProtocol.COS.value: CosStorageAdapter,
|
||||
StorageProtocol.LOCAL.value: LocalStorageAdapter,
|
||||
}
|
||||
|
||||
|
||||
class StorageAdapterFactory:
|
||||
"""存储适配器工厂:根据协议创建对应适配器实例。"""
|
||||
|
||||
@staticmethod
|
||||
def create(config: StorageAdapterConfig) -> BaseStorageAdapter:
|
||||
adapter_cls = _STORAGE_ADAPTERS.get(config.protocol.value)
|
||||
if adapter_cls is None:
|
||||
raise CustomException(msg=f"不支持的存储协议: {config.protocol}")
|
||||
return adapter_cls(config)
|
||||
@@ -0,0 +1,253 @@
|
||||
import ftplib
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import tempfile
|
||||
from datetime import datetime
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import BaseStorageAdapter, FtpAdvancedConfig, StorageObject, StorageProtocol
|
||||
|
||||
|
||||
class _ImplicitFTP_TLS(ftplib.FTP_TLS):
|
||||
"""隐式 FTPS(默认端口 990)连接子类:socket 直连即套 TLS。"""
|
||||
|
||||
def connect(self, host: str = "", port: int = 0, timeout: int = -999, source_address=None) -> str:
|
||||
if host != "":
|
||||
self.host = host
|
||||
if port > 0:
|
||||
self.port = port
|
||||
if timeout != -999:
|
||||
self.timeout = timeout
|
||||
if source_address is not None:
|
||||
self.source_address = source_address
|
||||
context = self.context
|
||||
if context is None:
|
||||
# FTP_TLS 默认 context 未设置时,构造宽松校验的客户端上下文(兼容自签名证书)
|
||||
context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
|
||||
context.check_hostname = False
|
||||
context.verify_mode = ssl.CERT_NONE
|
||||
self.sock = context.wrap_socket(
|
||||
socket.create_connection((self.host, self.port), self.timeout, self.source_address),
|
||||
server_hostname=self.host,
|
||||
)
|
||||
self.file = self.sock.makefile("r", encoding=self.encoding)
|
||||
self.welcome = self.getresp()
|
||||
return self.welcome
|
||||
|
||||
|
||||
class FtpStorageAdapter(BaseStorageAdapter):
|
||||
"""FTP / FTPS 存储适配器(ftplib 标准库,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.FTP
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
self._adv = FtpAdvancedConfig(**self.config.advanced_config)
|
||||
|
||||
def _new_client(self) -> ftplib.FTP_TLS | ftplib.FTP:
|
||||
"""建立连接并登录。FTP 明文 / FTPS(显式或隐式 TLS)按配置选择。"""
|
||||
implicit = False
|
||||
if self.config.protocol == StorageProtocol.FTPS:
|
||||
# 隐式 TLS:implicit_tls 开关或 encrypt_type>=3(兼容两种配置写法)
|
||||
implicit = self.config.implicit_tls or (self.config.encrypt_type or 0) >= 3
|
||||
client = _ImplicitFTP_TLS() if implicit else ftplib.FTP_TLS()
|
||||
else:
|
||||
client = ftplib.FTP()
|
||||
client.encoding = self.config.encoding or "utf-8"
|
||||
client.connect(host=self.config.host, port=self.config.port, timeout=self._adv.timeout)
|
||||
# 传输模式:0=默认(被动) 1=主动 2=被动
|
||||
if self.config.connection_mode == 1:
|
||||
client.set_pasv(False)
|
||||
elif self.config.connection_mode == 2:
|
||||
client.set_pasv(True)
|
||||
if isinstance(client, ftplib.FTP_TLS):
|
||||
if not implicit:
|
||||
client.auth() # 显式 FTPS:升级 TLS 通道(隐式 TLS 连接时已是加密通道,重复升级会报错)
|
||||
client.login(user=self.config.username or "", passwd=self.config.password or "")
|
||||
if isinstance(client, ftplib.FTP_TLS):
|
||||
client.prot_p() # 数据通道加密
|
||||
return client
|
||||
|
||||
@staticmethod
|
||||
def _close_client(client: ftplib.FTP) -> None:
|
||||
try:
|
||||
client.quit()
|
||||
except Exception:
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
client = self._new_client()
|
||||
try:
|
||||
client.pwd()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"FTP 连接测试失败: {e}")
|
||||
return False
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
client = self._new_client()
|
||||
try:
|
||||
# 目标目录不存在时逐级创建(与 SFTP 上传行为对齐,否则 STOR 直接失败)
|
||||
dir_part = remote_path.rsplit("/", 1)[0] if "/" in remote_path else ""
|
||||
if dir_part:
|
||||
self._ensure_remote_dir(client, dir_part)
|
||||
with open(local_path, "rb") as f:
|
||||
client.storbinary(f"STOR {remote_path}", f)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 上传失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
return remote_path
|
||||
|
||||
@staticmethod
|
||||
def _ensure_remote_dir(client: ftplib.FTP, remote_dir: str) -> None:
|
||||
"""逐级创建远端目录(mkd -p,目录已存在的 550 错误静默跳过)。"""
|
||||
current = ""
|
||||
for part in [p for p in remote_dir.split("/") if p]:
|
||||
current = f"{current}/{part}"
|
||||
try:
|
||||
client.mkd(current)
|
||||
except ftplib.error_perm:
|
||||
pass
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
client = self._new_client()
|
||||
try:
|
||||
with open(local_path, "wb") as f:
|
||||
client.retrbinary(f"RETR {remote_path}", f.write)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 下载失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
client = self._new_client()
|
||||
try:
|
||||
client.delete(remote_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 删除失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
client = self._new_client()
|
||||
try:
|
||||
try:
|
||||
client.size(remote_path)
|
||||
return True
|
||||
except ftplib.error_perm:
|
||||
# 部分服务器不支持 SIZE,退化为 NLST 判断
|
||||
try:
|
||||
client.nlst(remote_path)
|
||||
return True
|
||||
except ftplib.error_perm:
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
client = self._new_client()
|
||||
try:
|
||||
# 空 prefix 表示浏览根目录:FTP 用 "." 表示登录后的当前目录
|
||||
entries = list(client.mlsd(prefix or "."))
|
||||
result: list[StorageObject] = []
|
||||
for name, facts in entries:
|
||||
if name in (".", ".."):
|
||||
continue
|
||||
modified_time = None
|
||||
raw_mtime = facts.get("modify")
|
||||
if raw_mtime:
|
||||
try:
|
||||
modified_time = datetime.strptime(raw_mtime, "%Y%m%d%H%M%S")
|
||||
except ValueError:
|
||||
modified_time = None
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=f"{prefix}/{name}" if prefix else name,
|
||||
is_dir=facts.get("type") == "dir",
|
||||
size=int(facts["size"]) if facts.get("size") else None,
|
||||
modified_time=modified_time,
|
||||
)
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 列表失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
# ── 目录操作(FTP mkd 不自动创建父目录,需逐级创建)────────────────
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
client = self._new_client()
|
||||
try:
|
||||
parts = [p for p in remote_dir.split("/") if p]
|
||||
current = ""
|
||||
for part in parts:
|
||||
current = f"{current}/{part}" if current else part
|
||||
try:
|
||||
client.mkd(current)
|
||||
except ftplib.error_perm:
|
||||
pass # 目录已存在时跳过(部分服务器返回 550)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 创建目录失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
client = self._new_client()
|
||||
try:
|
||||
client.rmd(remote_dir)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 删除目录失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
client = self._new_client()
|
||||
try:
|
||||
client.rename(src, dst)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 重命名失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
|
||||
@staticmethod
|
||||
def _is_dir(client: ftplib.FTP, path: str) -> bool:
|
||||
"""通过 MLSD 类型判断路径是否为目录。"""
|
||||
try:
|
||||
return any(facts.get("type") == "dir" for _name, facts in client.mlsd(path))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件经本地临时文件中转(FTP 无服务端复制)。"""
|
||||
client = self._new_client()
|
||||
try:
|
||||
if self._is_dir(client, src):
|
||||
self._close_client(client)
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
with tempfile.NamedTemporaryFile(delete=False) as tmp:
|
||||
tmp_path = tmp.name
|
||||
try:
|
||||
with open(tmp_path, "wb") as f:
|
||||
client.retrbinary(f"RETR {src}", f.write)
|
||||
with open(tmp_path, "rb") as f:
|
||||
client.storbinary(f"STOR {dst}", f)
|
||||
finally:
|
||||
os.remove(tmp_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"FTP 复制失败: {e!s}")
|
||||
finally:
|
||||
self._close_client(client)
|
||||
@@ -0,0 +1,120 @@
|
||||
import os
|
||||
import shutil
|
||||
from datetime import datetime
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import BaseStorageAdapter, StorageObject, StorageProtocol
|
||||
|
||||
|
||||
class LocalStorageAdapter(BaseStorageAdapter):
|
||||
"""本地磁盘/挂载目录存储适配器(零 SDK,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.LOCAL
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
self.root = config.host or ""
|
||||
|
||||
def _abs_path(self, remote_path: str) -> str:
|
||||
"""将远端相对路径映射为本地绝对路径,并防御路径穿越。
|
||||
|
||||
remote_path 为相对存储真根(host 根目录)的路径,直接 join 到 root 下。
|
||||
"""
|
||||
if ".." in remote_path:
|
||||
raise CustomException(msg="本地存储不允许路径穿越(..)")
|
||||
return os.path.join(self.root, remote_path.lstrip("/"))
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
if not self.root or not os.path.isdir(self.root):
|
||||
logger.warning(f"本地存储根目录不可用: {self.root}")
|
||||
return False
|
||||
return os.access(self.root, os.R_OK | os.W_OK)
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
dest = self._abs_path(remote_path)
|
||||
try:
|
||||
os.makedirs(os.path.dirname(dest), exist_ok=True)
|
||||
shutil.copy2(local_path, dest)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
src = self._abs_path(remote_path)
|
||||
try:
|
||||
shutil.copy2(src, local_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
target = self._abs_path(remote_path)
|
||||
try:
|
||||
if os.path.isdir(target):
|
||||
shutil.rmtree(target)
|
||||
elif os.path.exists(target):
|
||||
os.remove(target)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
return os.path.exists(self._abs_path(remote_path))
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
base_dir = self._abs_path(prefix) if prefix else self.root
|
||||
try:
|
||||
entries = os.scandir(base_dir)
|
||||
except OSError as e:
|
||||
raise CustomException(msg=f"本地存储列表失败: {e!s}")
|
||||
|
||||
result: list[StorageObject] = []
|
||||
for entry in entries:
|
||||
try:
|
||||
is_dir = entry.is_dir()
|
||||
stat = entry.stat()
|
||||
except OSError:
|
||||
continue
|
||||
key = f"{prefix}/{entry.name}" if prefix else entry.name
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=entry.name,
|
||||
key=key,
|
||||
is_dir=is_dir,
|
||||
size=None if is_dir else stat.st_size,
|
||||
modified_time=datetime.fromtimestamp(stat.st_mtime),
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
# ── 目录操作(本地文件系统原生支持,零额外开销)──────────────────
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
os.makedirs(self._abs_path(remote_dir), exist_ok=True)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
os.rmdir(self._abs_path(remote_dir))
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储删除目录失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
try:
|
||||
os.makedirs(os.path.dirname(self._abs_path(dst)), exist_ok=True)
|
||||
shutil.move(self._abs_path(src), self._abs_path(dst))
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储重命名失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
src_abs, dst_abs = self._abs_path(src), self._abs_path(dst)
|
||||
try:
|
||||
if os.path.isdir(src_abs):
|
||||
shutil.copytree(src_abs, dst_abs, dirs_exist_ok=True)
|
||||
else:
|
||||
os.makedirs(os.path.dirname(dst_abs), exist_ok=True)
|
||||
shutil.copy2(src_abs, dst_abs)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"本地存储复制失败: {e!s}")
|
||||
@@ -0,0 +1,487 @@
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from obs import ObsClient
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import (
|
||||
BaseStorageAdapter,
|
||||
ObsAdvancedConfig,
|
||||
StorageObject,
|
||||
StoragePage,
|
||||
StorageProtocol,
|
||||
normalize_endpoint,
|
||||
)
|
||||
|
||||
|
||||
class ObsStorageAdapter(BaseStorageAdapter):
|
||||
"""华为云 OBS 存储适配器(esdk-obs-python,同步调用经 asyncio.to_thread 包装)。
|
||||
SDK API 概览
|
||||
------------
|
||||
|
||||
一、桶管理(Bucket)
|
||||
|
||||
| 接口名 | 方法 | 功能描述 |
|
||||
| --- | --- | --- |
|
||||
| 创建桶 | ObsClient.createBucket | 创建桶。 |
|
||||
| 获取桶列表 | ObsClient.listBuckets | 查询桶列表,返回结果按照桶名字典序排列。 |
|
||||
| 判断桶是否存在 | ObsClient.headBucket | 判断桶是否存在。 |
|
||||
| 删除桶 | ObsClient.deleteBucket | 删除桶,待删除的桶必须为空。 |
|
||||
| 列举桶内对象 | ObsClient.listObjects | 列举桶内对象,默认返回最大1000个对象。 |
|
||||
| 列举桶内多版本对象 | ObsClient.listVersions | 列举桶内多版本对象,默认返回最大1000个多版本对象。 |
|
||||
| 列举分段上传任务 | ObsClient.listMultipartUploads | 列举指定桶中所有的初始化后还未合并或还未取消的分段上传任务。 |
|
||||
| 获取桶元数据 | ObsClient.getBucketMetadata | 对桶发送HEAD请求,获取桶的存储类型、CORS规则(如果已设置)等信息。 |
|
||||
| 获取桶区域位置 | ObsClient.getBucketLocation | 获取桶所在的区域位置。 |
|
||||
| 获取桶存量信息 | ObsClient.getBucketStorageInfo | 获取桶的存量信息,包含桶的空间大小以及对象个数。 |
|
||||
| 设置桶配额 | ObsClient.setBucketQuota | 设置桶的配额值,单位为字节,支持的最大值为2^63-1,配额值设为0表示桶的配额没有上限。 |
|
||||
| 获取桶配额 | ObsClient.getBucketQuota | 获取桶的配额值,0代表配额没有上限。 |
|
||||
| 设置桶存储类型 | ObsClient.setBucketStoragePolicy | 设置桶的存储类型,桶中对象的存储类型默认将与桶的存储类型保持一致。 |
|
||||
| 获取桶存储类型 | ObsClient.getBucketStoragePolicy | 获取桶的存储类型。 |
|
||||
| 设置桶ACL | ObsClient.setBucketAcl | 设置桶ACL。 |
|
||||
| 获取桶ACL | ObsClient.getBucketAcl | 获取桶ACL。 |
|
||||
| 设置桶日志管理配置 | ObsClient.setBucketLogging | 设置桶的访问日志配置。 |
|
||||
| 获取桶日志管理配置 | ObsClient.getBucketLogging | 获取桶的访问日志配置。 |
|
||||
| 设置桶策略 | ObsClient.setBucketPolicy | 配置桶的策略,如果桶已经存在一个策略,当前请求中的策略将完全覆盖桶中现存的策略。 |
|
||||
| 获取桶策略 | ObsClient.getBucketPolicy | 获取桶的策略配置。 |
|
||||
| 删除桶策略 | ObsClient.deleteBucketPolicy | 删除桶的策略配置。 |
|
||||
| 设置桶生命周期配置 | ObsClient.setBucketLifecycle | 配置桶的生命周期规则,实现定时转换桶中对象的存储类型,以及定时删除桶中对象的功能。 |
|
||||
| 获取桶生命周期配置 | ObsClient.getBucketLifecycle | 获取桶的生命周期规则。 |
|
||||
| 删除桶生命周期配置 | ObsClient.deleteBucketLifecycle | 删除桶所有的生命周期规则。 |
|
||||
| 设置桶Website配置 | ObsClient.setBucketWebsite | 设置桶的Website配置。 |
|
||||
| 获取桶Website配置 | ObsClient.getBucketWebsite | 获取桶的Website配置。 |
|
||||
| 删除桶Website配置 | ObsClient.deleteBucketWebsite | 删除指定桶的Website配置。 |
|
||||
| 设置桶多版本状态 | ObsClient.setBucketVersioning | 设置桶的多版本状态。 |
|
||||
| 获取桶多版本状态 | ObsClient.getBucketVersioning | 获取桶的多版本状态。 |
|
||||
| 设置桶CORS配置 | ObsClient.setBucketCors | 设置桶的跨域资源共享规则,以允许客户端浏览器进行跨域请求。 |
|
||||
| 获取桶CORS配置 | ObsClient.getBucketCors | 获取指定桶的跨域资源共享规则。 |
|
||||
| 删除桶CORS配置 | ObsClient.deleteBucketCors | 删除指定桶的跨域资源共享规则。 |
|
||||
| 设置桶标签 | ObsClient.setBucketTagging | 设置桶的标签。 |
|
||||
| 获取桶标签 | ObsClient.getBucketTagging | 获取指定桶的标签。 |
|
||||
| 删除桶标签 | ObsClient.deleteBucketTagging | 删除指定桶的标签。 |
|
||||
|
||||
二、对象管理(Object)
|
||||
|
||||
| 接口名 | 方法 | 功能描述 |
|
||||
| --- | --- | --- |
|
||||
| 上传对象 | ObsClient.putContent | 上传对象到指定桶中。 |
|
||||
| 上传文件 | ObsClient.putFile | 上传文件/文件夹到指定桶中。 |
|
||||
| 追加上传 | ObsClient.appendObject | 对同一个对象追加数据内容。 |
|
||||
| 下载对象 | ObsClient.getObject | 下载指定桶中的对象。 |
|
||||
| 复制对象 | ObsClient.copyObject | 为指定桶中的对象创建一个副本。 |
|
||||
| 删除对象 | ObsClient.deleteObject | 删除指定桶中的对象。 |
|
||||
| 批量删除对象 | ObsClient.deleteObjects | 批量删除指定桶中的多个对象。 |
|
||||
| 获取对象元数据 | ObsClient.getObjectMetadata | 对指定桶中的对象发送HEAD请求,获取对象的元数据信息。 |
|
||||
| 修改对象元数据 | ObsClient.setObjectMetadata | 修改指定桶中对象的元数据信息。 |
|
||||
| 设置对象ACL | ObsClient.setObjectAcl | 设置指定桶中对象ACL。 |
|
||||
| 获取对象ACL | ObsClient.getObjectAcl | 获取指定桶中对象ACL。 |
|
||||
|
||||
三、分段上传(Multipart Upload)
|
||||
|
||||
| 接口名 | 方法 | 功能描述 |
|
||||
| --- | --- | --- |
|
||||
| 初始化分段上传任务 | ObsClient.initiateMultipartUpload | 在指定桶中初始化分段上传任务。 |
|
||||
| 上传段 | ObsClient.uploadPart | 初始化分段上传任务后,通过分段上传任务的ID,上传段到指定桶中。 |
|
||||
| 复制段 | ObsClient.copyPart | 初始化分段上传任务后,通过分段上传任务的ID,复制段到指定桶中。 |
|
||||
| 列举已上传的段 | ObsClient.listParts | 通过分段上传任务的ID,列举指定桶中已上传的段。 |
|
||||
| 合并段 | ObsClient.completeMultipartUpload | 通过分段上传任务的ID,合并指定桶中已上传的段。 |
|
||||
| 取消分段上传任务 | ObsClient.abortMultipartUpload | 通过分段上传任务的ID,取消指定桶中的分段上传任务。 |
|
||||
|
||||
四、高级功能
|
||||
|
||||
| 接口名 | 方法 | 功能描述 |
|
||||
| --- | --- | --- |
|
||||
| 恢复归档存储对象 | ObsClient.restoreObject | 恢复指定桶中的归档存储对象。 |
|
||||
| 生成带授权信息的URL | ObsClient.createSignedUrl | 通过访问密钥、请求方法类型、请求参数等信息生成一个在Query参数中携带鉴权信息的URL,以对OBS服务进行特定操作。 |
|
||||
| 生成带授权信息的表单上传参数 | ObsClient.createPostSignature | 生成用于鉴权的请求参数,以进行基于浏览器的POST表单上传。 |
|
||||
| 断点续传上传 | ObsClient.uploadFile | 对分段上传的封装和加强,解决上传大文件时由于网络不稳定或程序崩溃导致上传失败的问题。 |
|
||||
| 断点续传下载 | ObsClient.downloadFile | 对范围下载的封装和加强,解决下载大对象到本地时由于网络不稳定或程序崩溃导致下载失败的问题。 |
|
||||
|
||||
五、工作流(WorkflowClient)
|
||||
|
||||
| 接口名 | 方法 | 功能描述 |
|
||||
| --- | --- | --- |
|
||||
| 创建工作流 | WorkflowClient.createWorkflow | 根据模板创建工作流。 |
|
||||
| 查询工作流 | WorkflowClient.getWorkflow | 按名称查询工作流。 |
|
||||
| 删除工作流 | WorkflowClient.deleteWorkflow | 删除存在的工作流。 |
|
||||
| 更新工作流 | WorkflowClient.updateWorkflow | 更新工作流。 |
|
||||
| 查询工作流列表 | WorkflowClient.listWorkflow | 查询工作流列表。 |
|
||||
| API触发启动工作流 | WorkflowClient.asyncAPIStartWorkflow | API触发启动工作流。 |
|
||||
| 查询工作流实例列表 | WorkflowClient.listWorkflowExecution | 查询工作流实例列表。 |
|
||||
| 查询工作流实例 | WorkflowClient.getWorkflowExecution | 查询工作流实例详细。 |
|
||||
| 恢复失败状态的工作流实例 | WorkflowClient.restoreFailedWorkflowExecution | 当且仅当工作流实例处于执行失败状态才能执行恢复操作。恢复后,工作流实例将从上次失败的状态处继续执行,已执行过的状态不会再执行。 |
|
||||
| 配置桶触发器 | WorkflowClient.putTriggerPolicy | 在桶上绑定工作流触发器。 |
|
||||
| 查询桶触发器 | WorkflowClient.getTriggerPolicy | 查询桶上绑定工作流触发器。 |
|
||||
| 删除桶触发器 | WorkflowClient.deleteTriggerPolicy | 删除在桶上绑定工作流触发器。 |
|
||||
|
||||
"""
|
||||
|
||||
protocol = StorageProtocol.OBS
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
if not self.config.endpoint:
|
||||
raise CustomException(msg="OBS 存储源必须配置 endpoint")
|
||||
adv = ObsAdvancedConfig(**self.config.advanced_config)
|
||||
self.client = ObsClient(
|
||||
access_key_id=self.config.username or "",
|
||||
secret_access_key=self.config.password or "",
|
||||
server=normalize_endpoint(self.config.endpoint, self.config.scheme),
|
||||
is_secure=self.config.is_secure,
|
||||
max_retry_count=adv.max_retry_count,
|
||||
timeout=adv.timeout,
|
||||
ssl_verify=adv.ssl_verify,
|
||||
is_cname=adv.is_cname,
|
||||
path_style=adv.path_style,
|
||||
pool_size=adv.pool_size,
|
||||
signature=adv.signature,
|
||||
region=adv.region,
|
||||
security_token=adv.security_token,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_ok(resp: Any) -> bool:
|
||||
"""判断 OBS 响应是否成功(status < 300)。SDK 未提供类型存根,故用 getattr 访问动态属性。"""
|
||||
status = getattr(resp, "status", None)
|
||||
return status is not None and status < 300
|
||||
|
||||
@staticmethod
|
||||
def _error_desc(resp: Any) -> str:
|
||||
"""提取 OBS 响应中的错误描述。"""
|
||||
code = getattr(resp, "errorCode", "") or ""
|
||||
message = getattr(resp, "errorMessage", "") or ""
|
||||
return f"{code} {message}".strip()
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
try:
|
||||
resp = self.client.listBuckets()
|
||||
if self._is_ok(resp):
|
||||
resp = self.client.headBucket(bucketName=self._require_bucket())
|
||||
if self._is_ok(resp):
|
||||
return True
|
||||
logger.warning(f"OBS 连接测试失败: {self._error_desc(resp)}")
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.warning(f"OBS 连接测试失败: {e}")
|
||||
return False
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
"""上传:默认 uploadFile(断点续传 + 分片并发),流式传输(stream)走 putContent 单次上传。"""
|
||||
try:
|
||||
if self.config.transfer_mode == "stream":
|
||||
with open(local_path, "rb") as f:
|
||||
resp = self.client.putContent(
|
||||
bucketName=self._require_bucket(),
|
||||
objectKey=remote_path,
|
||||
content=f,
|
||||
)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 上传失败: {self._error_desc(resp)}")
|
||||
return remote_path
|
||||
part_size, concurrency, _ = self._multipart_settings()
|
||||
resp = self.client.uploadFile(
|
||||
bucketName=self._require_bucket(),
|
||||
objectKey=remote_path,
|
||||
uploadFile=local_path,
|
||||
partSize=part_size,
|
||||
taskNum=concurrency,
|
||||
enableCheckpoint=True,
|
||||
)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 上传失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
"""下载(downloadFile 断点续传 + 分片并发)。"""
|
||||
try:
|
||||
resp = self.client.downloadFile(
|
||||
bucketName=self._require_bucket(),
|
||||
objectKey=remote_path,
|
||||
downloadFile=local_path,
|
||||
partSize=5 * 1024 * 1024,
|
||||
taskNum=3,
|
||||
enableCheckpoint=True,
|
||||
)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 下载失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
try:
|
||||
resp = self.client.deleteObject(bucketName=self._require_bucket(), objectKey=remote_path)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 删除失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
try:
|
||||
resp = self.client.getObjectMetadata(bucketName=self._require_bucket(), objectKey=remote_path)
|
||||
return self._is_ok(resp)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _to_utc_dt(value: str | None) -> datetime | None:
|
||||
"""OBS SDK 的 lastModified 为 UTC 字符串(如 '2026/08/25 09:25:27'),转为 aware datetime。"""
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return datetime.strptime(value, "%Y/%m/%d %H:%M:%S").replace(tzinfo=UTC)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
"""列举桶内“目录层”条目(Delimiter=/ 模拟文件夹,返回名称/是否目录/大小/修改时间)。
|
||||
OBS 单次最多返回 1000 条,经 marker 翻页拉全量。"""
|
||||
try:
|
||||
result: list[StorageObject] = []
|
||||
seen_dirs: set[str] = set()
|
||||
seen_files: set[str] = set()
|
||||
marker: str | None = None
|
||||
while True:
|
||||
kwargs: dict[str, Any] = {"bucketName": self._require_bucket(), "prefix": prefix, "delimiter": "/"}
|
||||
if marker:
|
||||
kwargs["marker"] = marker
|
||||
resp = self.client.listObjects(**kwargs)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 列表失败: {self._error_desc(resp)}")
|
||||
body = getattr(resp, "body", None)
|
||||
for common in getattr(body, "commonPrefixs", None) or []:
|
||||
raw_key = (getattr(common, "prefix", "") or "").rstrip("/")
|
||||
if raw_key in seen_dirs:
|
||||
continue
|
||||
seen_dirs.add(raw_key)
|
||||
result.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in getattr(body, "contents", None) or []:
|
||||
raw_key = getattr(obj, "key", "") or ""
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
if raw_key in seen_files:
|
||||
continue
|
||||
seen_files.add(raw_key)
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=getattr(obj, "size", None),
|
||||
modified_time=self._to_utc_dt(getattr(obj, "lastModified", None)),
|
||||
)
|
||||
)
|
||||
if getattr(body, "is_truncated", None):
|
||||
marker = getattr(body, "next_marker", None) or ""
|
||||
if not marker:
|
||||
break
|
||||
else:
|
||||
break
|
||||
return result
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_page(self, prefix: str, page_size: int, cursor: str | None) -> StoragePage:
|
||||
"""游标分页:单次 SDK 请求只拉一页(OBS marker),翻页经前端回传游标。"""
|
||||
try:
|
||||
kwargs: dict[str, Any] = {
|
||||
"bucketName": self._require_bucket(),
|
||||
"prefix": prefix,
|
||||
"delimiter": "/",
|
||||
"max_keys": page_size,
|
||||
}
|
||||
if cursor:
|
||||
kwargs["marker"] = cursor
|
||||
resp = self.client.listObjects(**kwargs)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 列表失败: {self._error_desc(resp)}")
|
||||
body = getattr(resp, "body", None)
|
||||
truncated = bool(getattr(body, "is_truncated", None))
|
||||
next_cursor = getattr(body, "next_marker", None) or None
|
||||
items: list[StorageObject] = []
|
||||
for common in getattr(body, "commonPrefixs", None) or []:
|
||||
raw_key = (getattr(common, "prefix", "") or "").rstrip("/")
|
||||
items.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in getattr(body, "contents", None) or []:
|
||||
raw_key = getattr(obj, "key", "") or ""
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
items.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=getattr(obj, "size", None),
|
||||
modified_time=self._to_utc_dt(getattr(obj, "lastModified", None)),
|
||||
)
|
||||
)
|
||||
return StoragePage(
|
||||
items=items,
|
||||
has_next=truncated,
|
||||
next_cursor=next_cursor,
|
||||
)
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶。"""
|
||||
try:
|
||||
resp = self.client.listBuckets()
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 桶列表失败: {self._error_desc(resp)}")
|
||||
body = getattr(resp, "body", None)
|
||||
buckets = getattr(body, "buckets", None) or []
|
||||
return [getattr(b, "name", "") or "" for b in buckets if getattr(b, "name", "")]
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 桶列表失败: {e!s}")
|
||||
|
||||
def _sync_get_url(self, remote_path: str, expire: int) -> str:
|
||||
try:
|
||||
resp = self.client.createSignedUrl("GET", bucketName=self._require_bucket(), objectKey=remote_path, expires=expire)
|
||||
return getattr(resp, "signedUrl", "")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 生成预签名 URL 失败: {e!s}")
|
||||
|
||||
# ── 目录操作(对象存储以 key/ 占位对象模拟目录)──────────────────
|
||||
|
||||
def _list_all_keys(self, prefix: str) -> list[str]:
|
||||
"""分页列举前缀下的全部对象 key。"""
|
||||
keys: list[str] = []
|
||||
marker: str | None = None
|
||||
while True:
|
||||
kwargs: dict[str, Any] = {"bucketName": self._require_bucket(), "prefix": prefix}
|
||||
if marker:
|
||||
kwargs["marker"] = marker
|
||||
resp = self.client.listObjects(**kwargs)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 列举对象失败: {self._error_desc(resp)}")
|
||||
body = getattr(resp, "body", None)
|
||||
for obj in getattr(body, "contents", None) or []:
|
||||
key = getattr(obj, "key", "")
|
||||
if key:
|
||||
keys.append(key)
|
||||
if not getattr(body, "is_truncated", False):
|
||||
break
|
||||
# 未指定 delimiter 时服务端可能不返回 next_marker,退化为当前页最后一个 key,避免死循环
|
||||
marker = getattr(body, "next_marker", None) or (keys[-1] if keys else "")
|
||||
if not marker:
|
||||
break
|
||||
return keys
|
||||
|
||||
def _is_dir(self, key: str) -> bool:
|
||||
"""判断 key 是否代表目录(存在占位对象或前缀下有任何对象)。"""
|
||||
try:
|
||||
resp = self.client.listObjects(bucketName=self._require_bucket(), prefix=key.rstrip("/") + "/", max_keys=1)
|
||||
if not self._is_ok(resp):
|
||||
return False
|
||||
body = getattr(resp, "body", None)
|
||||
return bool(getattr(body, "contents", None) or getattr(body, "commonPrefixs", None))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
resp = self.client.putContent(bucketName=self._require_bucket(), objectKey=remote_dir.rstrip("/") + "/", content=b"")
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 创建目录失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
resp = self.client.deleteObject(bucketName=self._require_bucket(), objectKey=remote_dir.rstrip("/") + "/")
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 删除目录失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 删除目录失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件用服务端 copyObject。"""
|
||||
try:
|
||||
if self._is_dir(src):
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
resp = self.client.copyObject(
|
||||
sourceBucketName=self._require_bucket(),
|
||||
sourceObjectKey=src,
|
||||
destBucketName=self._require_bucket(),
|
||||
destObjectKey=dst,
|
||||
)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 复制失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 复制失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动:优先 OBS 原生 renameFile(支持目录级),失败回退先复制后删除。"""
|
||||
try:
|
||||
resp = self.client.renameFile(bucketName=self._require_bucket(), objectKey=src, newObjectKey=dst)
|
||||
if self._is_ok(resp):
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
if self._is_dir(src):
|
||||
self._sync_move_dir(src, dst)
|
||||
else:
|
||||
self._sync_copy(src, dst)
|
||||
self._sync_delete(src)
|
||||
|
||||
def _sync_list_recursive(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
keys = self._list_all_keys(prefix)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 递归列表失败: {e!s}")
|
||||
return self._entries_from_keys(keys)
|
||||
|
||||
def _sync_delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除:一次列举全部对象并分批批量删除(含占位目录对象)。"""
|
||||
try:
|
||||
keys = self._list_all_keys(remote_dir)
|
||||
marker = remote_dir.rstrip("/") + "/"
|
||||
if marker not in keys:
|
||||
keys.append(marker)
|
||||
for i in range(0, len(keys), 1000):
|
||||
batch = keys[i : i + 1000]
|
||||
resp = self.client.deleteObjects(
|
||||
bucketName=self._require_bucket(),
|
||||
deleteObjectsRequest={"quiet": True, "objects": [{"key": k} for k in batch]},
|
||||
)
|
||||
if not self._is_ok(resp):
|
||||
raise CustomException(msg=f"OBS 递归删除失败: {self._error_desc(resp)}")
|
||||
except CustomException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OBS 递归删除失败: {e!s}")
|
||||
|
||||
def _sync_close(self) -> None:
|
||||
"""关闭 OBS 客户端连接。"""
|
||||
try:
|
||||
self.client.close()
|
||||
except Exception as e:
|
||||
logger.warning(f"OBS 关闭连接失败: {e}")
|
||||
@@ -0,0 +1,285 @@
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import alibabacloud_oss_v2 as oss
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import (
|
||||
BaseStorageAdapter,
|
||||
OssAdvancedConfig,
|
||||
StorageObject,
|
||||
StoragePage,
|
||||
StorageProtocol,
|
||||
normalize_endpoint,
|
||||
)
|
||||
|
||||
|
||||
class OssStorageAdapter(BaseStorageAdapter):
|
||||
"""阿里云 OSS 存储适配器(alibabacloud_oss_v2 SDK,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.OSS
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
if not self.config.endpoint:
|
||||
raise CustomException(msg="OSS 存储源必须配置 endpoint")
|
||||
if not self.config.region:
|
||||
raise CustomException(msg="OSS 存储源必须配置 region(V4 签名要求,如 cn-hangzhou)")
|
||||
if not self.config.username or not self.config.password:
|
||||
raise CustomException(msg="OSS 存储源必须配置 AccessKeyId 与 AccessKeySecret")
|
||||
cfg = oss.config.load_default()
|
||||
cfg.credentials_provider = oss.credentials.StaticCredentialsProvider(
|
||||
self.config.username,
|
||||
self.config.password
|
||||
)
|
||||
cfg.region = self.config.region
|
||||
cfg.endpoint = normalize_endpoint(self.config.endpoint, self.config.scheme)
|
||||
adv = OssAdvancedConfig(**self.config.advanced_config)
|
||||
cfg.connect_timeout = adv.connect_timeout
|
||||
cfg.readwrite_timeout = adv.readwrite_timeout
|
||||
cfg.retry_max_attempts = adv.retry_max_attempts
|
||||
cfg.use_cname = adv.use_cname
|
||||
cfg.use_path_style = adv.use_path_style
|
||||
cfg.insecure_skip_verify = adv.insecure_skip_verify
|
||||
cfg.use_internal_endpoint = adv.use_internal_endpoint
|
||||
cfg.use_accelerate_endpoint = adv.use_accelerate_endpoint
|
||||
cfg.proxy_host = adv.proxy_host
|
||||
cfg.signature_version = adv.signature_version
|
||||
cfg.disable_upload_crc64_check = adv.disable_upload_crc64_check
|
||||
cfg.disable_download_crc64_check = adv.disable_download_crc64_check
|
||||
cfg.enabled_redirect = adv.enabled_redirect
|
||||
cfg.use_dualstack_endpoint = adv.use_dualstack_endpoint
|
||||
self.client = oss.Client(cfg)
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
try:
|
||||
self.client.get_bucket_info(oss.GetBucketInfoRequest(bucket=self.bucket_name))
|
||||
return True
|
||||
except oss.exceptions.ServiceError as e:
|
||||
logger.warning(f"OSS 连接测试失败: {e.code} {e.message}")
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.warning(f"OSS 连接测试失败: {e}")
|
||||
return False
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
try:
|
||||
# 走 SDK 高层分片上传:普通小文件自动单次上传,大文件按端点分片参数并发分片;
|
||||
# 流式传输(stream)时 part_size 为超大值,等效单次上传
|
||||
part_size, concurrency, _ = self._multipart_settings()
|
||||
self.client.uploader.upload_file(
|
||||
oss.PutObjectRequest(bucket=self.bucket_name, key=remote_path),
|
||||
local_path,
|
||||
part_size=part_size,
|
||||
parallel_num=concurrency,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
try:
|
||||
self.client.get_object_to_file(
|
||||
oss.GetObjectRequest(bucket=self.bucket_name, key=remote_path),
|
||||
local_path,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
try:
|
||||
self.client.delete_object(oss.DeleteObjectRequest(bucket=self.bucket_name, key=remote_path))
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
try:
|
||||
self.client.head_object(oss.HeadObjectRequest(bucket=self.bucket_name, key=remote_path))
|
||||
return True
|
||||
except oss.exceptions.ServiceError as e:
|
||||
if e.status_code == 404:
|
||||
return False
|
||||
logger.warning(f"OSS 判断文件存在失败: {e.code} {e.message}")
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _to_utc_dt(value: int | float | datetime | None) -> datetime | None:
|
||||
"""兼容 SDK 返回的时间戳(int/float)与 datetime 两种类型。"""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
return value.astimezone(UTC) if value.tzinfo else value.replace(tzinfo=UTC)
|
||||
return datetime.fromtimestamp(value, tz=UTC)
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
paginator = self.client.list_objects_v2_paginator()
|
||||
result: list[StorageObject] = []
|
||||
for page in paginator.iter_page(oss.ListObjectsV2Request(bucket=self.bucket_name, prefix=prefix, delimiter="/")):
|
||||
for common in page.common_prefixes or []:
|
||||
raw_key = (common.prefix or "").rstrip("/")
|
||||
result.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in page.contents or []:
|
||||
raw_key = obj.key or ""
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.size,
|
||||
modified_time=self._to_utc_dt(obj.last_modified),
|
||||
)
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_page(self, prefix: str, page_size: int, cursor: str | None) -> StoragePage:
|
||||
"""游标分页:单次 SDK 请求只拉一页(OSS continuation_token),翻页经前端回传游标。"""
|
||||
try:
|
||||
result = self.client.list_objects_v2(
|
||||
oss.ListObjectsV2Request(
|
||||
bucket=self.bucket_name,
|
||||
prefix=prefix,
|
||||
delimiter="/",
|
||||
max_keys=page_size,
|
||||
continuation_token=cursor or None,
|
||||
)
|
||||
)
|
||||
truncated = bool(getattr(result, "is_truncated", False))
|
||||
next_cursor = getattr(result, "next_continuation_token", None) or None
|
||||
items: list[StorageObject] = []
|
||||
for common in result.common_prefixes or []:
|
||||
raw_key = (common.prefix or "").rstrip("/")
|
||||
items.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in result.contents or []:
|
||||
raw_key = obj.key or ""
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
items.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.size,
|
||||
modified_time=self._to_utc_dt(obj.last_modified),
|
||||
)
|
||||
)
|
||||
return StoragePage(items=items, has_next=truncated, next_cursor=next_cursor)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶。"""
|
||||
try:
|
||||
paginator = self.client.list_buckets_paginator()
|
||||
result: list[str] = []
|
||||
for page in paginator.iter_page(oss.ListBucketsRequest()):
|
||||
for b in page.buckets or []:
|
||||
if b.name:
|
||||
result.append(b.name)
|
||||
return result
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 桶列表失败: {e!s}")
|
||||
|
||||
def _sync_get_url(self, remote_path: str, expire: int) -> str | None:
|
||||
try:
|
||||
result = self.client.presign(
|
||||
oss.GetObjectRequest(bucket=self.bucket_name, key=remote_path),
|
||||
expires=timedelta(seconds=expire),
|
||||
)
|
||||
return result.url
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 生成预签名 URL 失败: {e!s}")
|
||||
|
||||
# ── 目录操作(对象存储以 key/ 占位对象模拟目录)──────────────────
|
||||
|
||||
def _list_all_keys(self, prefix: str) -> list[str]:
|
||||
"""分页列举前缀下的全部对象 key。"""
|
||||
keys: list[str] = []
|
||||
paginator = self.client.list_objects_v2_paginator()
|
||||
for page in paginator.iter_page(oss.ListObjectsV2Request(bucket=self.bucket_name, prefix=prefix)):
|
||||
for obj in page.contents or []:
|
||||
if obj.key:
|
||||
keys.append(obj.key)
|
||||
return keys
|
||||
|
||||
def _is_dir(self, key: str) -> bool:
|
||||
"""判断 key 是否代表目录(存在占位对象或前缀下有任何对象)。"""
|
||||
try:
|
||||
page = next(
|
||||
self.client.list_objects_v2_paginator().iter_page(
|
||||
oss.ListObjectsV2Request(bucket=self.bucket_name, prefix=key.rstrip("/") + "/", max_keys=1)
|
||||
)
|
||||
)
|
||||
return bool(page.contents or page.common_prefixes)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.put_object(oss.PutObjectRequest(bucket=self.bucket_name, key=remote_dir.rstrip("/") + "/", body=b""))
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.delete_object(oss.DeleteObjectRequest(bucket=self.bucket_name, key=remote_dir.rstrip("/") + "/"))
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 删除目录失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件用服务端 copy_object。"""
|
||||
try:
|
||||
if self._is_dir(src):
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
self.client.copy_object(
|
||||
oss.CopyObjectRequest(bucket=self.bucket_name, key=dst, source_bucket=self.bucket_name, source_key=src)
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 复制失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动:目录先复制后删除;文件 copy_object 后删源。"""
|
||||
if self._is_dir(src):
|
||||
self._sync_move_dir(src, dst)
|
||||
else:
|
||||
self._sync_copy(src, dst)
|
||||
self._sync_delete(src)
|
||||
|
||||
def _sync_list_recursive(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
keys = self._list_all_keys(prefix)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 递归列表失败: {e!s}")
|
||||
return self._entries_from_keys(keys)
|
||||
|
||||
def _sync_delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除:一次列举全部对象并分批批量删除(含占位目录对象)。"""
|
||||
try:
|
||||
keys = self._list_all_keys(remote_dir)
|
||||
marker = remote_dir.rstrip("/") + "/"
|
||||
if marker not in keys:
|
||||
keys.append(marker)
|
||||
for i in range(0, len(keys), 1000):
|
||||
self.client.delete_multiple_objects(
|
||||
oss.DeleteMultipleObjectsRequest(
|
||||
bucket=self.bucket_name,
|
||||
objects=[oss.DeleteObject(key=k) for k in keys[i : i + 1000]],
|
||||
quiet=True,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"OSS 递归删除失败: {e!s}")
|
||||
@@ -0,0 +1,306 @@
|
||||
from typing import Any
|
||||
|
||||
import boto3
|
||||
from boto3.s3.transfer import TransferConfig
|
||||
from botocore.config import Config
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import (
|
||||
BaseStorageAdapter,
|
||||
S3AdvancedConfig,
|
||||
StorageObject,
|
||||
StoragePage,
|
||||
StorageProtocol,
|
||||
normalize_endpoint,
|
||||
)
|
||||
|
||||
|
||||
class S3StorageAdapter(BaseStorageAdapter):
|
||||
"""S3 兼容对象存储适配器(boto3,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.S3
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
# 凭据为空时传空串而非 None,避免 boto3 走 EC2 实例元数据(IMDS)探测导致超时
|
||||
if not self.config.username or not self.config.password:
|
||||
raise CustomException(msg="S3 存储源必须配置 AccessKeyId 与 AccessKeySecret")
|
||||
if not self.config.bucket:
|
||||
raise CustomException(msg="S3 存储源必须配置 bucket")
|
||||
adv = S3AdvancedConfig(**self.config.advanced_config)
|
||||
proxies: dict[str, str] | None = None
|
||||
if adv.proxies:
|
||||
proxies = {"http": adv.proxies, "https": adv.proxies}
|
||||
s3_options: dict[str, Any] = {"addressing_style": adv.addressing_style}
|
||||
if adv.use_accelerate_endpoint:
|
||||
s3_options["use_accelerate_endpoint"] = True
|
||||
if adv.us_east_1_regional_endpoint:
|
||||
s3_options["us_east_1_regional_endpoint"] = adv.us_east_1_regional_endpoint
|
||||
config = Config(
|
||||
signature_version=adv.signature_version,
|
||||
retries={"max_attempts": adv.max_attempts, "mode": adv.retries_mode},
|
||||
connect_timeout=adv.connect_timeout,
|
||||
read_timeout=adv.read_timeout,
|
||||
max_pool_connections=adv.max_pool_connections,
|
||||
tcp_keepalive=adv.tcp_keepalive,
|
||||
use_dualstack_endpoint=adv.use_dualstack_endpoint,
|
||||
proxies=proxies,
|
||||
parameter_validation=adv.parameter_validation,
|
||||
request_checksum_calculation=adv.request_checksum_calculation,
|
||||
response_checksum_validation=adv.response_checksum_validation,
|
||||
s3=s3_options,
|
||||
)
|
||||
self.client = boto3.client(
|
||||
"s3",
|
||||
endpoint_url=normalize_endpoint(self.config.endpoint, self.config.scheme),
|
||||
region_name=self.config.region,
|
||||
aws_access_key_id=self.config.username or "",
|
||||
aws_secret_access_key=self.config.password or "",
|
||||
config=config,
|
||||
)
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
try:
|
||||
self.client.head_bucket(Bucket=self._require_bucket())
|
||||
return True
|
||||
except ClientError as e:
|
||||
code = e.response.get("Error", {}).get("Code", "")
|
||||
# 403 表示凭据有效但无权限查看桶,连接本身是通的
|
||||
if code == "403":
|
||||
return True
|
||||
logger.warning(f"S3 连接测试失败: {code} {e}")
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.warning(f"S3 连接测试失败: {e}")
|
||||
return False
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
try:
|
||||
# 按端点分片参数配置:超过分片大小阈值即分片并发上传(并发受内存预算护栏约束);
|
||||
# 流式传输(stream)时阈值为超大值,等效单次上传
|
||||
part_size, concurrency, _ = self._multipart_settings()
|
||||
transfer_config = TransferConfig(
|
||||
multipart_threshold=part_size,
|
||||
multipart_chunksize=part_size,
|
||||
max_concurrency=concurrency,
|
||||
)
|
||||
self.client.upload_file(local_path, self._require_bucket(), remote_path, Config=transfer_config)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
try:
|
||||
self.client.download_file(self._require_bucket(), remote_path, local_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
try:
|
||||
self.client.delete_object(Bucket=self._require_bucket(), Key=remote_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
try:
|
||||
self.client.head_object(Bucket=self._require_bucket(), Key=remote_path)
|
||||
return True
|
||||
except ClientError as e:
|
||||
if e.response.get("ResponseMetadata", {}).get("HTTPStatusCode") == 404:
|
||||
return False
|
||||
logger.warning(f"S3 head_object 失败: {e}")
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
"""列举目录层条目。S3 单次最多返回 1000 条,经 ContinuationToken 翻页拉全量。"""
|
||||
try:
|
||||
result: list[StorageObject] = []
|
||||
seen_dirs: set[str] = set()
|
||||
seen_files: set[str] = set()
|
||||
token: str | None = None
|
||||
while True:
|
||||
kwargs: dict[str, Any] = {
|
||||
"Bucket": self._require_bucket(),
|
||||
"Prefix": prefix,
|
||||
"Delimiter": "/",
|
||||
}
|
||||
if token:
|
||||
kwargs["ContinuationToken"] = token
|
||||
resp = self.client.list_objects_v2(**kwargs)
|
||||
for cp in resp.get("CommonPrefixes", []):
|
||||
raw_key = cp.get("Prefix", "").rstrip("/")
|
||||
if raw_key in seen_dirs:
|
||||
continue
|
||||
seen_dirs.add(raw_key)
|
||||
result.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in resp.get("Contents", []):
|
||||
raw_key = obj.get("Key", "")
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
if raw_key in seen_files:
|
||||
continue
|
||||
seen_files.add(raw_key)
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.get("Size"),
|
||||
modified_time=obj.get("LastModified"),
|
||||
)
|
||||
)
|
||||
if resp.get("IsTruncated") and resp.get("NextContinuationToken"):
|
||||
token = resp["NextContinuationToken"]
|
||||
else:
|
||||
break
|
||||
return result
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_page(self, prefix: str, page_size: int, cursor: str | None) -> StoragePage:
|
||||
"""游标分页:单次 SDK 请求只拉一页(S3 ContinuationToken),翻页经前端回传游标。"""
|
||||
try:
|
||||
kwargs: dict[str, Any] = {
|
||||
"Bucket": self._require_bucket(),
|
||||
"Prefix": prefix,
|
||||
"Delimiter": "/",
|
||||
"MaxKeys": page_size,
|
||||
}
|
||||
if cursor:
|
||||
kwargs["ContinuationToken"] = cursor
|
||||
resp = self.client.list_objects_v2(**kwargs)
|
||||
items: list[StorageObject] = []
|
||||
for cp in resp.get("CommonPrefixes", []):
|
||||
raw_key = cp.get("Prefix", "").rstrip("/")
|
||||
items.append(StorageObject(name=raw_key.rsplit("/", 1)[-1], key=raw_key, is_dir=True))
|
||||
for obj in resp.get("Contents", []):
|
||||
raw_key = obj.get("Key", "")
|
||||
if raw_key == prefix:
|
||||
continue
|
||||
name = raw_key.rsplit("/", 1)[-1]
|
||||
if not name: # key 以 "/" 结尾的目录占位对象,不按文件展示
|
||||
continue
|
||||
items.append(
|
||||
StorageObject(
|
||||
name=name,
|
||||
key=raw_key,
|
||||
is_dir=False,
|
||||
size=obj.get("Size"),
|
||||
modified_time=obj.get("LastModified"),
|
||||
)
|
||||
)
|
||||
return StoragePage(
|
||||
items=items,
|
||||
has_next=bool(resp.get("IsTruncated")),
|
||||
next_cursor=resp.get("NextContinuationToken"),
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 列表失败: {e!s}")
|
||||
|
||||
def _sync_list_buckets(self) -> list[str]:
|
||||
"""列出账号下全部存储桶。"""
|
||||
try:
|
||||
resp = self.client.list_buckets()
|
||||
return [b.get("Name", "") for b in resp.get("Buckets", []) if b.get("Name")]
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 桶列表失败: {e!s}")
|
||||
|
||||
def _sync_get_url(self, remote_path: str, expire: int) -> str:
|
||||
try:
|
||||
return self.client.generate_presigned_url(
|
||||
"get_object",
|
||||
Params={"Bucket": self._require_bucket(), "Key": remote_path},
|
||||
ExpiresIn=expire,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 生成预签名 URL 失败: {e!s}")
|
||||
|
||||
def _sync_close(self) -> None:
|
||||
"""关闭 boto3 客户端连接。"""
|
||||
close = getattr(self.client, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
# ── 目录操作(对象存储以 key/ 占位对象模拟目录)──────────────────
|
||||
|
||||
def _list_all_keys(self, prefix: str) -> list[str]:
|
||||
"""分页列举前缀下的全部对象 key。"""
|
||||
keys: list[str] = []
|
||||
paginator = self.client.get_paginator("list_objects_v2")
|
||||
for page in paginator.paginate(Bucket=self._require_bucket(), Prefix=prefix):
|
||||
keys.extend(obj["Key"] for obj in page.get("Contents", []))
|
||||
return keys
|
||||
|
||||
def _is_dir(self, key: str) -> bool:
|
||||
"""判断 key 是否代表目录(存在占位对象或前缀下有任何对象)。"""
|
||||
try:
|
||||
resp = self.client.list_objects_v2(
|
||||
Bucket=self._require_bucket(), Prefix=key.rstrip("/") + "/", MaxKeys=1
|
||||
)
|
||||
return bool(resp.get("Contents") or resp.get("CommonPrefixes"))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.put_object(Bucket=self._require_bucket(), Key=remote_dir.rstrip("/") + "/", Body=b"")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self.client.delete_object(Bucket=self._require_bucket(), Key=remote_dir.rstrip("/") + "/")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 删除目录失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件用服务端 copy_object。"""
|
||||
try:
|
||||
if self._is_dir(src):
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
self.client.copy_object(
|
||||
Bucket=self._require_bucket(),
|
||||
CopySource={"Bucket": self._require_bucket(), "Key": src},
|
||||
Key=dst,
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 复制失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
"""重命名/移动:目录先复制后删除;文件 copy_object 后删源。"""
|
||||
if self._is_dir(src):
|
||||
self._sync_move_dir(src, dst)
|
||||
else:
|
||||
self._sync_copy(src, dst)
|
||||
self._sync_delete(src)
|
||||
|
||||
def _sync_list_recursive(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
keys = self._list_all_keys(prefix)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 递归列表失败: {e!s}")
|
||||
return self._entries_from_keys(keys)
|
||||
|
||||
def _sync_delete_dir(self, remote_dir: str) -> None:
|
||||
"""递归删除:一次列举全部对象并分批批量删除(含占位目录对象)。"""
|
||||
try:
|
||||
keys = self._list_all_keys(remote_dir)
|
||||
marker = remote_dir.rstrip("/") + "/"
|
||||
if marker not in keys:
|
||||
keys.append(marker)
|
||||
for i in range(0, len(keys), 1000):
|
||||
self.client.delete_objects(
|
||||
Bucket=self._require_bucket(),
|
||||
Delete={"Objects": [{"Key": k} for k in keys[i : i + 1000]], "Quiet": True},
|
||||
)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"S3 递归删除失败: {e!s}")
|
||||
@@ -0,0 +1,217 @@
|
||||
"""SFTP 存储适配器
|
||||
|
||||
基于 paramiko。连接在首次使用时建立,并在适配器实例生命周期内复用
|
||||
(避免每次操作都重新进行 SSH 握手,与 FTP/FTPS 客户端保持一致);
|
||||
同步操作经基类 asyncio.to_thread 包装,不阻塞事件循环。
|
||||
用完由调用方调用 close() 释放连接。
|
||||
"""
|
||||
import os
|
||||
import stat
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import paramiko
|
||||
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import BaseStorageAdapter, SftpAdvancedConfig, StorageObject, StorageProtocol
|
||||
|
||||
|
||||
class SftpStorageAdapter(BaseStorageAdapter):
|
||||
"""SFTP 存储适配器(paramiko,同步调用经 asyncio.to_thread 包装)。"""
|
||||
|
||||
protocol = StorageProtocol.SFTP
|
||||
|
||||
# SFTP 通道复用同一连接,不支持多线程并发读写
|
||||
concurrency_safe: bool = False
|
||||
|
||||
def __init__(self, config) -> None:
|
||||
super().__init__(config)
|
||||
self._ssh: paramiko.SSHClient | None = None
|
||||
self._sftp: paramiko.SFTPClient | None = None
|
||||
self._adv = SftpAdvancedConfig(**self.config.advanced_config)
|
||||
|
||||
# ── 连接管理 ──────────────────────────────────────────────────
|
||||
|
||||
def _connect(self) -> paramiko.SFTPClient:
|
||||
"""建立并复用 SFTP 连接:首次使用时连接,之后直接复用。"""
|
||||
if self._sftp is not None:
|
||||
return self._sftp
|
||||
if not self.config.username or not self.config.password:
|
||||
raise CustomException(msg="SFTP 存储源必须配置用户名与密码")
|
||||
ssh = paramiko.SSHClient()
|
||||
ssh.set_missing_host_key_policy(paramiko.AutoAddPolicy())
|
||||
try:
|
||||
ssh.connect(
|
||||
hostname=self.config.host,
|
||||
port=self.config.port,
|
||||
username=self.config.username or "",
|
||||
password=self.config.password or "",
|
||||
timeout=self._adv.connect_timeout,
|
||||
banner_timeout=self._adv.banner_timeout,
|
||||
auth_timeout=self._adv.auth_timeout,
|
||||
channel_timeout=self._adv.channel_timeout,
|
||||
allow_agent=self._adv.allow_agent,
|
||||
look_for_keys=self._adv.look_for_keys,
|
||||
compress=self._adv.compress,
|
||||
)
|
||||
transport = ssh.get_transport()
|
||||
if transport is None or not transport.is_active():
|
||||
raise RuntimeError("SSH 连接已断开")
|
||||
if self._adv.keepalive_interval > 0:
|
||||
transport.set_keepalive(self._adv.keepalive_interval)
|
||||
sftp = ssh.open_sftp()
|
||||
# 给 SFTP 通道设置超时,防止服务器不响应时 listdir/stat 等操作无限阻塞
|
||||
channel = sftp.get_channel()
|
||||
if channel:
|
||||
channel.settimeout(self._adv.connect_timeout)
|
||||
except Exception:
|
||||
ssh.close()
|
||||
raise
|
||||
self._ssh = ssh
|
||||
self._sftp = sftp
|
||||
return sftp
|
||||
|
||||
@staticmethod
|
||||
def _ensure_remote_dir(client: paramiko.SFTPClient, remote_dir: str) -> None:
|
||||
"""递归创建远端目录(mkdir -p)。"""
|
||||
parts = [p for p in remote_dir.split("/") if p]
|
||||
current = ""
|
||||
for part in parts:
|
||||
current = f"{current}/{part}" if current else part
|
||||
try:
|
||||
client.stat(current)
|
||||
except FileNotFoundError:
|
||||
client.mkdir(current)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def _sync_close(self) -> None:
|
||||
"""关闭 SFTP 通道与 SSH 连接。"""
|
||||
if self._sftp:
|
||||
try:
|
||||
self._sftp.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._sftp = None
|
||||
if self._ssh:
|
||||
try:
|
||||
self._ssh.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._ssh = None
|
||||
|
||||
# ── 同步协议操作 ──────────────────────────────────────────────
|
||||
|
||||
def _sync_test_connection(self) -> bool:
|
||||
try:
|
||||
self._connect().listdir(".")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"SFTP 连接测试失败: {e}")
|
||||
return False
|
||||
|
||||
def _sync_upload(self, local_path: str, remote_path: str) -> str:
|
||||
client = self._connect()
|
||||
try:
|
||||
dir_part, _ = remote_path.rsplit("/", 1) if "/" in remote_path else ("", remote_path)
|
||||
if dir_part:
|
||||
self._ensure_remote_dir(client, dir_part)
|
||||
client.put(local_path, remote_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 上传失败: {e!s}")
|
||||
return remote_path
|
||||
|
||||
def _sync_download(self, remote_path: str, local_path: str) -> str:
|
||||
"""下载,支持断点续传:本地已有部分文件时从偏移处继续。"""
|
||||
client = self._connect()
|
||||
try:
|
||||
remote_size = client.stat(remote_path).st_size or 0
|
||||
local_size = os.path.getsize(local_path) if os.path.exists(local_path) else 0
|
||||
if local_size > remote_size:
|
||||
# 本地文件比远端还大(可能内容不一致),整文件重新下载
|
||||
os.remove(local_path)
|
||||
local_size = 0
|
||||
if local_size == remote_size:
|
||||
return local_path # 已完整下载,跳过
|
||||
if local_size:
|
||||
with client.open(remote_path, "rb") as rf, open(local_path, "ab") as lf:
|
||||
rf.seek(local_size)
|
||||
while chunk := rf.read(1024 * 1024):
|
||||
lf.write(chunk)
|
||||
else:
|
||||
client.get(remote_path, local_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 下载失败: {e!s}")
|
||||
return local_path
|
||||
|
||||
def _sync_delete(self, remote_path: str) -> None:
|
||||
try:
|
||||
self._connect().remove(remote_path)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 删除失败: {e!s}")
|
||||
|
||||
def _sync_exists(self, remote_path: str) -> bool:
|
||||
try:
|
||||
self._connect().stat(remote_path)
|
||||
return True
|
||||
except (FileNotFoundError, OSError):
|
||||
return False
|
||||
|
||||
def _sync_list(self, prefix: str) -> list[StorageObject]:
|
||||
try:
|
||||
# 空 prefix 表示浏览根目录:SFTP 用 "." 表示登录后的当前目录
|
||||
attrs = self._connect().listdir_attr(prefix or ".")
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 列表失败: {e!s}")
|
||||
result: list[StorageObject] = []
|
||||
for attr in attrs:
|
||||
result.append(
|
||||
StorageObject(
|
||||
name=attr.filename,
|
||||
key=f"{prefix}/{attr.filename}" if prefix else attr.filename,
|
||||
is_dir=bool(attr.st_mode and stat.S_ISDIR(attr.st_mode)),
|
||||
size=attr.st_size,
|
||||
modified_time=datetime.fromtimestamp(attr.st_mtime, tz=UTC) if attr.st_mtime else None,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
# ── 目录操作(SFTP 原生支持目录重命名,复制则流式传输)────────────
|
||||
|
||||
def _is_dir(self, path: str) -> bool:
|
||||
try:
|
||||
attr = self._connect().stat(path)
|
||||
return bool(attr.st_mode and stat.S_ISDIR(attr.st_mode))
|
||||
except (FileNotFoundError, OSError):
|
||||
return False
|
||||
|
||||
def _sync_mkdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self._ensure_remote_dir(self._connect(), remote_dir)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 创建目录失败: {e!s}")
|
||||
|
||||
def _sync_rmdir(self, remote_dir: str) -> None:
|
||||
try:
|
||||
self._connect().rmdir(remote_dir)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 删除目录失败: {e!s}")
|
||||
|
||||
def _sync_rename(self, src: str, dst: str) -> None:
|
||||
try:
|
||||
self._connect().posix_rename(src, dst)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 重命名失败: {e!s}")
|
||||
|
||||
def _sync_copy(self, src: str, dst: str) -> None:
|
||||
"""复制:目录走基类递归实现;文件用通道流式复制(SFTP 无服务端复制)。"""
|
||||
client = self._connect()
|
||||
try:
|
||||
if self._is_dir(src):
|
||||
self._sync_copy_dir(src, dst)
|
||||
return
|
||||
with client.open(src, "rb") as rf, client.open(dst, "wb") as wf:
|
||||
while chunk := rf.read(1024 * 1024):
|
||||
wf.write(chunk)
|
||||
except Exception as e:
|
||||
raise CustomException(msg=f"SFTP 复制失败: {e!s}")
|
||||
@@ -0,0 +1,94 @@
|
||||
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.flow.schema import WorkflowFlowCreateSchema, WorkflowFlowExecuteSchema, WorkflowFlowOutSchema, WorkflowFlowQueryParam, WorkflowFlowUpdateSchema
|
||||
from app.modules.workflow.flow.service import WorkflowFlowService
|
||||
|
||||
WorkflowFlowRouter = APIRouter(route_class=OperationLogRoute, prefix="/flow", tags=["传输流程"])
|
||||
|
||||
|
||||
@WorkflowFlowRouter.get("/page", summary="分页查询传输流程", response_model=ResponseSchema[PageResultSchema[WorkflowFlowOutSchema]])
|
||||
async def get_flow_page_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
page: Annotated[PaginationQueryParam, Depends()],
|
||||
search: Annotated[WorkflowFlowQueryParam, Query()],
|
||||
) -> JSONResponse:
|
||||
result: PageResultSchema[WorkflowFlowOutSchema] = await WorkflowFlowService(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="查询传输流程分页成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.get("/list", summary="查询传输流程列表", response_model=ResponseSchema[list[WorkflowFlowOutSchema]])
|
||||
async def get_flow_list_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
search: Annotated[WorkflowFlowQueryParam, Query()],
|
||||
) -> JSONResponse:
|
||||
result: list[WorkflowFlowOutSchema] = await WorkflowFlowService(auth, db).get_list(search=search)
|
||||
return SuccessResponse(data=result, msg="查询传输流程列表成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.get("/detail/{id}", summary="查询传输流程详情", response_model=ResponseSchema[WorkflowFlowOutSchema])
|
||||
async def get_flow_detail_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="流程ID", ge=1)],
|
||||
) -> JSONResponse:
|
||||
result: WorkflowFlowOutSchema = await WorkflowFlowService(auth, db).detail(id=id)
|
||||
return SuccessResponse(data=result, msg="查询传输流程详情成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.post("/create", status_code=status.HTTP_201_CREATED, summary="创建传输流程", response_model=ResponseSchema[WorkflowFlowOutSchema])
|
||||
async def create_flow_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:create"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
data: Annotated[WorkflowFlowCreateSchema, Body(description="流程创建参数")],
|
||||
) -> JSONResponse:
|
||||
result: WorkflowFlowOutSchema = await WorkflowFlowService(auth, db).create(data=data)
|
||||
return SuccessResponse(data=result, msg="创建传输流程成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.put("/update/{id}", summary="修改传输流程", response_model=ResponseSchema[WorkflowFlowOutSchema])
|
||||
async def update_flow_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="流程ID", ge=1)],
|
||||
data: Annotated[WorkflowFlowUpdateSchema, Body(description="流程修改参数")],
|
||||
) -> JSONResponse:
|
||||
result: WorkflowFlowOutSchema = await WorkflowFlowService(auth, db).update(id=id, data=data)
|
||||
return SuccessResponse(data=result, msg="修改传输流程成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.delete("/delete", summary="删除传输流程", response_model=ResponseSchema[None])
|
||||
async def delete_flow_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:flow:delete"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
ids: Annotated[list[int], Body(description="流程ID列表")],
|
||||
) -> JSONResponse:
|
||||
await WorkflowFlowService(auth, db).delete(ids=ids)
|
||||
return SuccessResponse(msg="删除传输流程成功")
|
||||
|
||||
|
||||
@WorkflowFlowRouter.post("/execute/{id}", summary="执行传输流程", response_model=ResponseSchema[list[int]])
|
||||
async def execute_flow_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:create"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="流程ID", ge=1)],
|
||||
data: Annotated[WorkflowFlowExecuteSchema | None, Body(description="执行参数(源文件/目录路径映射,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
task_ids: list[int] = await WorkflowFlowService(auth, db).execute(
|
||||
id=id, source_paths=data.source_paths if data else None
|
||||
)
|
||||
return SuccessResponse(data=task_ids, msg=f"执行传输流程成功,已生成 {len(task_ids)} 个传输任务")
|
||||
@@ -0,0 +1,48 @@
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy import delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_crud import CRUDBase
|
||||
from app.core.base_schema import AuthSchema
|
||||
|
||||
from .model import WorkflowFlowEdgeModel, WorkflowFlowModel, WorkflowFlowNodeModel
|
||||
from .schema import WorkflowFlowCreateSchema, WorkflowFlowUpdateSchema
|
||||
|
||||
|
||||
class WorkflowFlowCRUD(CRUDBase[WorkflowFlowModel, WorkflowFlowCreateSchema, WorkflowFlowUpdateSchema]):
|
||||
"""传输流程数据层"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
super().__init__(model=WorkflowFlowModel, auth=auth, db=db)
|
||||
|
||||
|
||||
class WorkflowFlowNodeCRUD(CRUDBase[WorkflowFlowNodeModel, object, object]):
|
||||
"""流程节点明细数据层
|
||||
|
||||
明细是 flow 的派生从属数据:随父流程全量覆写/删除,无独立数据权限主体,
|
||||
也没有回收站/软删消费场景。因此不沿用基类软删 delete(否则每次保存画布都会
|
||||
残留一套 is_deleted=1 的旧明细,无限累积),删除统一走下方物理删除方法。
|
||||
"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
super().__init__(model=WorkflowFlowNodeModel, auth=auth, db=db)
|
||||
|
||||
async def hard_delete_by_flow_ids(self, flow_ids: Sequence[int]) -> None:
|
||||
"""按 flow 物理删除明细(父流程行已过数据权限校验)。"""
|
||||
if not flow_ids:
|
||||
return
|
||||
_ = await self.db.execute(delete(WorkflowFlowNodeModel).where(WorkflowFlowNodeModel.flow_id.in_(flow_ids)))
|
||||
|
||||
|
||||
class WorkflowFlowEdgeCRUD(CRUDBase[WorkflowFlowEdgeModel, object, object]):
|
||||
"""流程连线明细数据层,删除语义同 WorkflowFlowNodeCRUD"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
super().__init__(model=WorkflowFlowEdgeModel, auth=auth, db=db)
|
||||
|
||||
async def hard_delete_by_flow_ids(self, flow_ids: Sequence[int]) -> None:
|
||||
"""按 flow 物理删除明细(父流程行已过数据权限校验)。"""
|
||||
if not flow_ids:
|
||||
return
|
||||
_ = await self.db.execute(delete(WorkflowFlowEdgeModel).where(WorkflowFlowEdgeModel.flow_id.in_(flow_ids)))
|
||||
@@ -0,0 +1,56 @@
|
||||
from sqlalchemy import JSON, Boolean, Integer, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.core.base_model import ModelMixin, UserMixin
|
||||
|
||||
|
||||
class WorkflowFlowModel(ModelMixin, UserMixin):
|
||||
"""传输流程:定义源节点 → 目标节点列表(parallel 多目标 / chain 链式)
|
||||
|
||||
graph 仅存画布布局与展示字段(节点位置、连线样式),业务配置落于
|
||||
flow_node / flow_edge 表,回显时由 service 组装,避免双份数据不一致。
|
||||
"""
|
||||
|
||||
__tablename__: str = "task_workflow_flow"
|
||||
__table_args__: dict[str, str] = {"comment": "传输流程定义表"}
|
||||
|
||||
name: Mapped[str] = mapped_column(String(64), nullable=False, index=True, comment="流程名称")
|
||||
task_type: Mapped[str] = mapped_column(String(16), nullable=False, default="parallel", comment="类型(parallel:多目标 chain:链式)")
|
||||
graph: Mapped[dict | None] = mapped_column(JSON, nullable=True, comment="VueFlow画布布局数据 {nodes:[{id,type,position,label}],edges:[{id,source,target,type,animated,style,label}]}")
|
||||
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="备注")
|
||||
|
||||
|
||||
class WorkflowFlowNodeModel(ModelMixin, UserMixin):
|
||||
"""流程画布节点(业务配置):节点关联的存储源与默认源目录。
|
||||
|
||||
节点在画布上的位置/名称等布局字段存 flow.graph 的 nodes 项(key=node_key)。
|
||||
"""
|
||||
|
||||
__tablename__: str = "task_workflow_flow_node"
|
||||
__table_args__: dict[str, str] = {"comment": "流程画布节点表"}
|
||||
|
||||
flow_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True, comment="流程ID")
|
||||
node_key: Mapped[str] = mapped_column(String(64), nullable=False, comment="画布节点ID")
|
||||
source_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True, comment="存储源ID")
|
||||
source_path: Mapped[str | None] = mapped_column(String(1024), default=None, nullable=True, comment="默认源目录")
|
||||
|
||||
|
||||
class WorkflowFlowEdgeModel(ModelMixin, UserMixin):
|
||||
"""流程画布连线(业务配置):传输方式与分片参数。
|
||||
|
||||
连线的目标目录由目标节点的默认源目录决定(节点 source_path),连线不再配置路径。
|
||||
连线的样式/动画等展示字段存 flow.graph 的 edges 项(key=edge_key)。
|
||||
"""
|
||||
|
||||
__tablename__: str = "task_workflow_flow_edge"
|
||||
__table_args__: dict[str, str] = {"comment": "流程画布连线表"}
|
||||
|
||||
flow_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True, comment="流程ID")
|
||||
edge_key: Mapped[str] = mapped_column(String(64), nullable=False, comment="画布连线ID")
|
||||
source_node_key: Mapped[str] = mapped_column(String(64), nullable=False, comment="源画布节点ID")
|
||||
target_node_key: Mapped[str] = mapped_column(String(64), nullable=False, comment="目标画布节点ID")
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False, comment="是否启用(禁用则不执行)")
|
||||
transfer_mode: Mapped[str | None] = mapped_column(String(16), default=None, nullable=True, comment="传输方式(stream/multipart,空用存储源默认)")
|
||||
multipart_part_size: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="分片大小(MB)")
|
||||
multipart_concurrency: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="分片并发数")
|
||||
@@ -0,0 +1,178 @@
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
from app.core.base_schema import BaseQueryParam, BaseSchema, UserByQueryParam, UserBySchema
|
||||
|
||||
FlowTaskType = Literal["parallel", "chain"]
|
||||
|
||||
|
||||
class FlowTargetSchema(BaseModel):
|
||||
"""流程目标节点配置(目标目录由目标节点默认源目录决定)"""
|
||||
|
||||
target_id: int = Field(..., ge=1, description="目标节点ID(存储源)")
|
||||
target_path: str = Field(..., max_length=1024, description="目标路径")
|
||||
|
||||
|
||||
class FlowSourceSchema(BaseModel):
|
||||
"""流程源节点概览(由画布连线实时派生)"""
|
||||
|
||||
source_id: int = Field(..., ge=1, description="源节点ID(存储源)")
|
||||
source_name: str | None = Field(default=None, description="源存储源名称")
|
||||
|
||||
|
||||
class FlowNodeSchema(BaseModel):
|
||||
"""画布节点业务配置(由画布拆分,落 flow_node 表)"""
|
||||
|
||||
node_key: str = Field(..., description="画布节点ID")
|
||||
source_id: int = Field(..., ge=1, description="存储源ID")
|
||||
source_path: str | None = Field(default=None, max_length=1024, description="默认源目录")
|
||||
|
||||
|
||||
class FlowEdgeSchema(BaseModel):
|
||||
"""画布连线业务配置(由画布拆分,落 flow_edge 表)
|
||||
|
||||
连线只定义传输方式;目标目录由目标节点的默认源目录(source_path)决定。
|
||||
"""
|
||||
|
||||
edge_key: str = Field(..., description="画布连线ID")
|
||||
source_node_key: str = Field(..., description="源画布节点ID")
|
||||
target_node_key: str = Field(..., description="目标画布节点ID")
|
||||
enabled: bool = Field(default=True, description="是否启用(禁用则不执行)")
|
||||
transfer_mode: str | None = Field(default=None, max_length=16, description="传输方式(stream/multipart,空用存储源默认)")
|
||||
multipart_part_size: int | None = Field(default=None, ge=1, description="分片大小(MB)")
|
||||
multipart_concurrency: int | None = Field(default=None, ge=1, description="分片并发数")
|
||||
|
||||
|
||||
class FlowLayoutNodeSchema(BaseModel):
|
||||
"""画布节点布局展示字段(存 flow.graph,业务配置在 flow_node 表)"""
|
||||
|
||||
id: str = Field(..., description="画布节点ID")
|
||||
type: str = Field("storage", description="节点类型")
|
||||
position: dict = Field(default_factory=lambda: {"x": 0, "y": 0}, description="节点位置")
|
||||
label: str | None = Field(default=None, description="节点显示名")
|
||||
style: dict | None = Field(default=None, description="节点样式")
|
||||
|
||||
|
||||
class FlowLayoutEdgeSchema(BaseModel):
|
||||
"""画布连线布局展示字段(存 flow.graph,业务配置在 flow_edge 表)"""
|
||||
|
||||
id: str = Field(..., description="画布连线ID")
|
||||
source: str = Field(..., description="源画布节点ID")
|
||||
target: str = Field(..., description="目标画布节点ID")
|
||||
type: str = Field("smoothstep", description="连线类型")
|
||||
animated: bool | None = Field(default=None, description="连线动画")
|
||||
style: dict | None = Field(default=None, description="连线样式")
|
||||
label: str | None = Field(default=None, description="连线显示名")
|
||||
|
||||
|
||||
class FlowGraphNodeDataSchema(BaseModel):
|
||||
"""画布节点回显数据(组装完整画布时写入 node.data)"""
|
||||
|
||||
source_id: int = Field(..., ge=1, description="存储源ID")
|
||||
source_path: str | None = Field(default=None, description="默认源目录")
|
||||
label: str | None = Field(default=None, description="节点显示名")
|
||||
protocol: str | None = Field(default=None, description="存储源协议")
|
||||
host: str | None = Field(default=None, description="主机地址")
|
||||
bucket: str | None = Field(default=None, description="桶名(对象存储)")
|
||||
endpoint: str | None = Field(default=None, description="对象存储地址")
|
||||
region: str | None = Field(default=None, description="区域")
|
||||
path_prefix: str | None = Field(default=None, description="路径前缀")
|
||||
|
||||
|
||||
class FlowGraphEdgeDataSchema(BaseModel):
|
||||
"""画布连线回显数据(组装完整画布时写入 edge.data)"""
|
||||
|
||||
enabled: bool = Field(default=True, description="是否启用(禁用则不执行)")
|
||||
transfer_mode: str | None = Field(default=None, description="传输方式")
|
||||
multipart_part_size: int | None = Field(default=None, description="分片大小(MB)")
|
||||
multipart_concurrency: int | None = Field(default=None, description="分片并发数")
|
||||
source_label: str | None = Field(default=None, description="源存储源名称")
|
||||
target_label: str | None = Field(default=None, description="目标存储源名称")
|
||||
source_protocol: str | None = Field(default=None, description="源存储源协议")
|
||||
target_protocol: str | None = Field(default=None, description="目标存储源协议")
|
||||
source_storage_id: int | None = Field(default=None, description="源存储源ID")
|
||||
target_storage_id: int | None = Field(default=None, description="目标存储源ID")
|
||||
|
||||
|
||||
class FlowSplitResultSchema(BaseModel):
|
||||
"""画布拆分结果:业务明细(node/edge)+ 精简布局 + 派生概览"""
|
||||
|
||||
layout: dict = Field(..., description="精简后的画布布局 {nodes:[{id,type,position,label}],edges:[{id,source,target,type,animated,style,label}]}")
|
||||
nodes: list[FlowNodeSchema] = Field(default_factory=list, description="节点业务明细")
|
||||
edges: list[FlowEdgeSchema] = Field(default_factory=list, description="连线业务明细")
|
||||
sources: list[FlowSourceSchema] = Field(default_factory=list, description="源节点列表(去重)")
|
||||
targets: list[FlowTargetSchema] = Field(default_factory=list, description="目标列表(去重)")
|
||||
|
||||
|
||||
class FlowTransferPlanSchema(BaseModel):
|
||||
"""执行计划:单条连线生成的传输任务参数"""
|
||||
|
||||
src_id: int = Field(..., ge=1, description="源存储源ID")
|
||||
tgt_id: int = Field(..., ge=1, description="目标存储源ID")
|
||||
src_path: str = Field(..., description="源文件/目录路径")
|
||||
tgt_path: str = Field(..., description="目标路径")
|
||||
transfer_mode: str = Field("stream", description="传输方式(stream/multipart)")
|
||||
multipart_part_size: int | None = Field(default=None, description="分片大小(MB)")
|
||||
multipart_concurrency: int | None = Field(default=None, description="分片并发数")
|
||||
|
||||
|
||||
class WorkflowFlowCreateSchema(BaseModel):
|
||||
"""创建传输流程(画布驱动:业务配置解析后落 flow_node/flow_edge 表)"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=64, description="流程名称")
|
||||
task_type: FlowTaskType = Field("parallel", description="类型(parallel:多目标 chain:链式)")
|
||||
graph: dict = Field(..., description="VueFlow画布数据 {nodes:[{id,type,position,data}],edges:[{id,source,target,data}]}")
|
||||
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
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_graph(self):
|
||||
"""画布必须包含传输连线(源节点 → 目标节点)。"""
|
||||
if not self.graph or not self.graph.get("edges"):
|
||||
raise ValueError("请至少添加一条传输连线(源节点 → 目标节点)")
|
||||
return self
|
||||
|
||||
|
||||
class WorkflowFlowUpdateSchema(WorkflowFlowCreateSchema):
|
||||
"""更新传输流程"""
|
||||
|
||||
|
||||
class WorkflowFlowExecuteSchema(BaseModel):
|
||||
"""执行传输流程参数"""
|
||||
|
||||
source_paths: dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description="源文件/目录路径映射 {源存储源ID: 路径},执行时必须为每条连线的源存储源指定",
|
||||
)
|
||||
|
||||
|
||||
class WorkflowFlowOutSchema(BaseSchema, UserBySchema):
|
||||
"""传输流程详情响应模型(sources/targets 由 service 从明细表派生)"""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
name: str | None = None
|
||||
task_type: FlowTaskType = "parallel"
|
||||
sources: list[FlowSourceSchema] = Field(default_factory=list, description="源节点列表(由画布派生)")
|
||||
targets: list[FlowTargetSchema] = Field(default_factory=list)
|
||||
graph: dict | None = None
|
||||
graph_stats: dict | None = None
|
||||
status: int = 0
|
||||
description: str | None = None
|
||||
|
||||
|
||||
class WorkflowFlowQueryParam(BaseQueryParam, UserByQueryParam):
|
||||
"""传输流程查询参数"""
|
||||
|
||||
name: str | None = Field(None, description="流程名称", json_schema_extra={"q": "like"})
|
||||
task_type: FlowTaskType | None = Field(None, description="类型(parallel/chain)", 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,489 @@
|
||||
from collections import defaultdict
|
||||
from collections.abc import Sequence
|
||||
from typing import cast
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_schema import AuthSchema, PageResultSchema
|
||||
from app.core.exceptions import CustomException
|
||||
from app.modules.workflow.source.crud import StorageSourceCRUD
|
||||
from app.modules.workflow.source.model import StorageSourceModel
|
||||
from app.modules.workflow.source.service import StorageSourceService
|
||||
from app.modules.workflow.transfer.schema import TransferMode, TransferTargetSchema, TransferTaskCreateSchema
|
||||
from app.modules.workflow.transfer.service import StorageTransferService
|
||||
from app.utils.common_util import search_to_dict
|
||||
|
||||
from .crud import WorkflowFlowCRUD, WorkflowFlowEdgeCRUD, WorkflowFlowNodeCRUD
|
||||
from .model import WorkflowFlowEdgeModel, WorkflowFlowModel, WorkflowFlowNodeModel
|
||||
from .schema import (
|
||||
FlowEdgeSchema,
|
||||
FlowGraphEdgeDataSchema,
|
||||
FlowGraphNodeDataSchema,
|
||||
FlowLayoutEdgeSchema,
|
||||
FlowLayoutNodeSchema,
|
||||
FlowNodeSchema,
|
||||
FlowSourceSchema,
|
||||
FlowSplitResultSchema,
|
||||
FlowTargetSchema,
|
||||
FlowTransferPlanSchema,
|
||||
WorkflowFlowCreateSchema,
|
||||
WorkflowFlowOutSchema,
|
||||
WorkflowFlowQueryParam,
|
||||
WorkflowFlowUpdateSchema,
|
||||
)
|
||||
|
||||
|
||||
class WorkflowFlowService:
|
||||
"""传输流程服务(源节点 → 目标节点,支持 1对多 / 多对1)
|
||||
|
||||
数据职责划分:
|
||||
- flow.graph:仅存画布布局与展示字段(节点位置、连线样式/动画),不做业务解析
|
||||
- flow_node / flow_edge 表:节点存储源与默认源目录、连线启用与传输方式(业务配置)
|
||||
- sources / targets:由连线明细实时派生,仅用于列表展示/校验
|
||||
|
||||
保存时拆分画布、回显时组装画布、执行直接读明细表,避免对 graph 的重复解析。
|
||||
"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
self.auth = auth
|
||||
self.db = db
|
||||
|
||||
# ── 内部工具 ────────────────────────────────────────────────────
|
||||
|
||||
def _crud(self) -> WorkflowFlowCRUD:
|
||||
return WorkflowFlowCRUD(self.auth, self.db)
|
||||
|
||||
async def _validate_nodes(
|
||||
self, sources: list[FlowSourceSchema], targets: list[FlowTargetSchema]
|
||||
) -> None:
|
||||
"""校验源节点与目标节点均存在且启用。"""
|
||||
node_service = StorageSourceService(self.auth, self.db)
|
||||
ids = [s.source_id for s in sources] + [t.target_id for t in targets]
|
||||
if ids:
|
||||
await node_service.get_active_sources(ids)
|
||||
|
||||
@staticmethod
|
||||
def _sources_from(
|
||||
edge_rows: Sequence[WorkflowFlowEdgeModel], node_map: dict[str, WorkflowFlowNodeModel]
|
||||
) -> list[FlowSourceSchema]:
|
||||
"""由连线明细派生源节点列表(去重)。"""
|
||||
out: list[FlowSourceSchema] = []
|
||||
seen: set[int] = set()
|
||||
for e in edge_rows:
|
||||
src_node = node_map.get(e.source_node_key)
|
||||
if not src_node:
|
||||
continue
|
||||
if src_node.source_id in seen:
|
||||
continue
|
||||
seen.add(src_node.source_id)
|
||||
out.append(FlowSourceSchema(source_id=src_node.source_id))
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def _targets_from(
|
||||
edge_rows: Sequence[WorkflowFlowEdgeModel], node_map: dict[str, WorkflowFlowNodeModel]
|
||||
) -> list[FlowTargetSchema]:
|
||||
"""由连线明细派生目标列表(去重),目标目录取目标节点的默认源目录。"""
|
||||
out: list[FlowTargetSchema] = []
|
||||
seen: set[int] = set()
|
||||
for e in edge_rows:
|
||||
tgt_node = node_map.get(e.target_node_key)
|
||||
if not tgt_node:
|
||||
continue
|
||||
if tgt_node.source_id in seen:
|
||||
continue
|
||||
seen.add(tgt_node.source_id)
|
||||
out.append(
|
||||
FlowTargetSchema(target_id=tgt_node.source_id, target_path=tgt_node.source_path or "")
|
||||
)
|
||||
return out
|
||||
|
||||
# ── 画布拆分/组装 ───────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _split_graph(graph: dict) -> FlowSplitResultSchema:
|
||||
"""校验并拆分提交的画布。
|
||||
|
||||
- 业务配置(节点 source_id/source_path、连线启用与传输方式)→ FlowNode/FlowEdge 明细
|
||||
- 布局展示(位置、连线样式/动画)→ 精简后的 layout
|
||||
同时完成画布完整性校验(存在连线、源≠目标),支持 1对多 / 多对1 拓扑。
|
||||
"""
|
||||
nodes = {n["id"]: n for n in (graph.get("nodes") or [])}
|
||||
edges = graph.get("edges") or []
|
||||
if not edges:
|
||||
raise CustomException(msg="画布中未找到有效的传输连线(源节点 → 目标节点)")
|
||||
|
||||
layout_nodes: list[dict] = []
|
||||
node_items: list[FlowNodeSchema] = []
|
||||
for nid, n in nodes.items():
|
||||
d = n.get("data") or {}
|
||||
layout_nodes.append(
|
||||
FlowLayoutNodeSchema(
|
||||
id=nid,
|
||||
type=n.get("type") or "storage",
|
||||
position=n.get("position") or {"x": 0, "y": 0},
|
||||
label=n.get("label") or d.get("label"),
|
||||
style=n.get("style"),
|
||||
).model_dump(exclude_none=True)
|
||||
)
|
||||
if d.get("source_id") is not None:
|
||||
node_items.append(
|
||||
FlowNodeSchema(node_key=nid, source_id=d["source_id"], source_path=d.get("source_path"))
|
||||
)
|
||||
node_map = {it.node_key: it for it in node_items}
|
||||
|
||||
sources: list[FlowSourceSchema] = []
|
||||
targets: list[FlowTargetSchema] = []
|
||||
edge_items: list[FlowEdgeSchema] = []
|
||||
layout_edges: list[dict] = []
|
||||
seen_edges: set[tuple[int, int]] = set()
|
||||
seen_sources: set[int] = set()
|
||||
seen_targets: set[int] = set()
|
||||
for e in edges:
|
||||
src_n = node_map.get(e.get("source"))
|
||||
tgt_n = node_map.get(e.get("target"))
|
||||
if not src_n or not tgt_n:
|
||||
continue
|
||||
sid, tid = src_n.source_id, tgt_n.source_id
|
||||
src_label = (nodes.get(e.get("source")) or {}).get("data", {}).get("label") or sid
|
||||
tgt_label = (nodes.get(e.get("target")) or {}).get("data", {}).get("label") or tid
|
||||
if sid == tid:
|
||||
raise CustomException(msg=f"流程保存失败,连线「{src_label} → {tgt_label}」源与目标不能相同")
|
||||
pair = (sid, tid)
|
||||
if pair in seen_edges:
|
||||
continue
|
||||
seen_edges.add(pair)
|
||||
if sid not in seen_sources:
|
||||
seen_sources.add(sid)
|
||||
sources.append(FlowSourceSchema(source_id=sid))
|
||||
if tid not in seen_targets:
|
||||
seen_targets.add(tid)
|
||||
# 目标目录由目标节点默认源目录决定
|
||||
targets.append(FlowTargetSchema(target_id=tid, target_path=tgt_n.source_path or ""))
|
||||
ed = e.get("data") or {}
|
||||
edge_items.append(
|
||||
FlowEdgeSchema(
|
||||
edge_key=e["id"],
|
||||
source_node_key=e.get("source"),
|
||||
target_node_key=e.get("target"),
|
||||
enabled=ed.get("enabled", True),
|
||||
transfer_mode=ed.get("transfer_mode"),
|
||||
multipart_part_size=ed.get("multipart_part_size"),
|
||||
multipart_concurrency=ed.get("multipart_concurrency"),
|
||||
)
|
||||
)
|
||||
layout_edges.append(
|
||||
FlowLayoutEdgeSchema(
|
||||
id=e["id"],
|
||||
source=e.get("source"),
|
||||
target=e.get("target"),
|
||||
type=e.get("type") or "smoothstep",
|
||||
animated=e.get("animated"),
|
||||
style=e.get("style"),
|
||||
label=e.get("label"),
|
||||
).model_dump(exclude_none=True)
|
||||
)
|
||||
|
||||
if not edge_items:
|
||||
raise CustomException(msg="画布中未找到有效的传输连线(源节点 → 目标节点)")
|
||||
return FlowSplitResultSchema(
|
||||
layout={"nodes": layout_nodes, "edges": layout_edges},
|
||||
nodes=node_items,
|
||||
edges=edge_items,
|
||||
sources=sources,
|
||||
targets=targets,
|
||||
)
|
||||
|
||||
async def _save_graph(self, flow_id: int, nodes: list[FlowNodeSchema], edges: list[FlowEdgeSchema]) -> None:
|
||||
"""覆写流程的业务明细(附属表物理删除后重建,不做逻辑删除)。
|
||||
|
||||
父流程行已在调用方经 WorkflowFlowCRUD 做过数据权限校验,明细随父全量覆写,
|
||||
删除统一走子表 CRUD 的物理清理方法(基类软删 delete 不适用于逐次覆写场景)。
|
||||
"""
|
||||
await WorkflowFlowNodeCRUD(self.auth, self.db).hard_delete_by_flow_ids([flow_id])
|
||||
await WorkflowFlowEdgeCRUD(self.auth, self.db).hard_delete_by_flow_ids([flow_id])
|
||||
user_id = self.auth.user.id
|
||||
for it in nodes:
|
||||
self.db.add(WorkflowFlowNodeModel(flow_id=flow_id, created_id=user_id, updated_id=user_id, **it.model_dump()))
|
||||
for it in edges:
|
||||
self.db.add(WorkflowFlowEdgeModel(flow_id=flow_id, created_id=user_id, updated_id=user_id, **it.model_dump()))
|
||||
await self.db.flush()
|
||||
|
||||
async def _load_flow_graph(self, flow_id: int) -> tuple[list[WorkflowFlowNodeModel], list[WorkflowFlowEdgeModel]]:
|
||||
result = await self.db.execute(
|
||||
select(WorkflowFlowNodeModel)
|
||||
.where(WorkflowFlowNodeModel.flow_id == flow_id)
|
||||
.order_by(WorkflowFlowNodeModel.id)
|
||||
)
|
||||
node_rows = list(result.scalars().all())
|
||||
result = await self.db.execute(
|
||||
select(WorkflowFlowEdgeModel)
|
||||
.where(WorkflowFlowEdgeModel.flow_id == flow_id)
|
||||
.order_by(WorkflowFlowEdgeModel.id)
|
||||
)
|
||||
edge_rows = list(result.scalars().all())
|
||||
return node_rows, edge_rows
|
||||
|
||||
async def _build_graph(
|
||||
self,
|
||||
flow: WorkflowFlowModel,
|
||||
node_rows: Sequence[WorkflowFlowNodeModel],
|
||||
edge_rows: Sequence[WorkflowFlowEdgeModel],
|
||||
) -> dict:
|
||||
"""由布局(flow.graph)+ 业务明细(node/edge 表)+ 存储源组装完整画布,供前端直接回显。"""
|
||||
layout = flow.graph or {}
|
||||
layout_nodes = {n["id"]: n for n in (layout.get("nodes") or [])}
|
||||
layout_edges = {e["id"]: e for e in (layout.get("edges") or [])}
|
||||
nodes_by_key = {n.node_key: n for n in node_rows}
|
||||
edges_by_key = {e.edge_key: e for e in edge_rows}
|
||||
|
||||
ids = {n.source_id for n in node_rows}
|
||||
src_map: dict[int, StorageSourceModel] = {}
|
||||
if ids:
|
||||
sources = await StorageSourceCRUD(self.auth, self.db).get_list(search={"id": ("in", sorted(ids))})
|
||||
src_map = {s.id: s for s in sources}
|
||||
|
||||
out_nodes: list[dict] = []
|
||||
for key, ln in layout_nodes.items():
|
||||
node_row = nodes_by_key.get(key)
|
||||
if not node_row:
|
||||
continue
|
||||
src = src_map.get(node_row.source_id)
|
||||
node = dict(ln)
|
||||
node["data"] = FlowGraphNodeDataSchema(
|
||||
source_id=node_row.source_id,
|
||||
source_path=node_row.source_path,
|
||||
label=ln.get("label") or (src.name if src else None),
|
||||
protocol=src.protocol if src else None,
|
||||
host=src.host if src else None,
|
||||
bucket=src.bucket if src else None,
|
||||
endpoint=src.endpoint if src else None,
|
||||
region=src.region if src else None,
|
||||
path_prefix=src.path_prefix if src else None,
|
||||
).model_dump(exclude_none=True)
|
||||
out_nodes.append(node)
|
||||
|
||||
out_edges: list[dict] = []
|
||||
for key, le in layout_edges.items():
|
||||
edge_row = edges_by_key.get(key)
|
||||
if not edge_row:
|
||||
continue
|
||||
src_node = nodes_by_key.get(edge_row.source_node_key)
|
||||
tgt_node = nodes_by_key.get(edge_row.target_node_key)
|
||||
src = src_map.get(src_node.source_id) if src_node else None
|
||||
tgt = src_map.get(tgt_node.source_id) if tgt_node else None
|
||||
edge = dict(le)
|
||||
edge["data"] = FlowGraphEdgeDataSchema(
|
||||
enabled=edge_row.enabled,
|
||||
transfer_mode=edge_row.transfer_mode,
|
||||
multipart_part_size=edge_row.multipart_part_size,
|
||||
multipart_concurrency=edge_row.multipart_concurrency,
|
||||
source_label=src.name if src else (str(src_node.source_id) if src_node else None),
|
||||
target_label=tgt.name if tgt else (str(tgt_node.source_id) if tgt_node else None),
|
||||
source_protocol=src.protocol if src else None,
|
||||
target_protocol=tgt.protocol if tgt else None,
|
||||
source_storage_id=src_node.source_id if src_node else None,
|
||||
target_storage_id=tgt_node.source_id if tgt_node else None,
|
||||
).model_dump(exclude_none=True)
|
||||
out_edges.append(edge)
|
||||
|
||||
return {"nodes": out_nodes, "edges": out_edges}
|
||||
|
||||
# ── 查询 ────────────────────────────────────────────────────────
|
||||
|
||||
async def _to_out(self, obj: WorkflowFlowModel) -> WorkflowFlowOutSchema:
|
||||
out = WorkflowFlowOutSchema.model_validate(obj)
|
||||
node_rows, edge_rows = await self._load_flow_graph(obj.id)
|
||||
out.graph_stats = {"node_count": len(node_rows), "edge_count": len(edge_rows)}
|
||||
if edge_rows:
|
||||
node_map = {n.node_key: n for n in node_rows}
|
||||
out.sources = self._sources_from(edge_rows, node_map)
|
||||
out.targets = self._targets_from(edge_rows, node_map)
|
||||
if obj.graph:
|
||||
out.graph = await self._build_graph(obj, node_rows, edge_rows)
|
||||
return out
|
||||
|
||||
async def _to_out_list(self, objs: Sequence[WorkflowFlowModel]) -> list[WorkflowFlowOutSchema]:
|
||||
"""批量组装列表概览:源/目标列表、画布统计,一次查询避免 N+1。"""
|
||||
flow_ids = [o.id for o in objs]
|
||||
nodes_by_flow: dict[int, list[WorkflowFlowNodeModel]] = defaultdict(list)
|
||||
edges_by_flow: dict[int, list[WorkflowFlowEdgeModel]] = defaultdict(list)
|
||||
if flow_ids:
|
||||
result = await self.db.execute(
|
||||
select(WorkflowFlowNodeModel).where(WorkflowFlowNodeModel.flow_id.in_(flow_ids))
|
||||
)
|
||||
for n in result.scalars().all():
|
||||
nodes_by_flow[n.flow_id].append(n)
|
||||
result = await self.db.execute(
|
||||
select(WorkflowFlowEdgeModel).where(WorkflowFlowEdgeModel.flow_id.in_(flow_ids))
|
||||
)
|
||||
for e in result.scalars().all():
|
||||
edges_by_flow[e.flow_id].append(e)
|
||||
|
||||
outs = [WorkflowFlowOutSchema.model_validate(o) for o in objs]
|
||||
for obj, out in zip(objs, outs, strict=False):
|
||||
out.graph = None
|
||||
ns = nodes_by_flow.get(obj.id, [])
|
||||
es = edges_by_flow.get(obj.id, [])
|
||||
out.graph_stats = {"node_count": len(ns), "edge_count": len(es)}
|
||||
if es:
|
||||
node_map = {n.node_key: n for n in ns}
|
||||
out.sources = self._sources_from(es, node_map)
|
||||
out.targets = self._targets_from(es, node_map)
|
||||
return outs
|
||||
|
||||
async def detail(self, id: int) -> WorkflowFlowOutSchema:
|
||||
obj = await self._crud().get_or_404(id=id)
|
||||
return await self._to_out(obj)
|
||||
|
||||
async def page(
|
||||
self,
|
||||
search: WorkflowFlowQueryParam | None,
|
||||
page_no: int,
|
||||
page_size: int,
|
||||
order_by: list[dict] | None = None,
|
||||
) -> PageResultSchema[WorkflowFlowOutSchema]:
|
||||
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),
|
||||
)
|
||||
items = await self._to_out_list(result.items)
|
||||
return PageResultSchema[WorkflowFlowOutSchema](
|
||||
page_no=result.page_no,
|
||||
page_size=result.page_size,
|
||||
total=result.total,
|
||||
has_next=result.has_next,
|
||||
items=items,
|
||||
)
|
||||
|
||||
async def get_list(self, search: WorkflowFlowQueryParam | None = None) -> list[WorkflowFlowOutSchema]:
|
||||
objs = await self._crud().get_list(search=search_to_dict(search), order_by=[{"id": "asc"}])
|
||||
return await self._to_out_list(objs)
|
||||
|
||||
# ── 写入 ────────────────────────────────────────────────────────
|
||||
|
||||
async def _prepare_and_save(
|
||||
self, flow_id: int | None, data: WorkflowFlowCreateSchema | WorkflowFlowUpdateSchema
|
||||
) -> tuple[dict, list[FlowNodeSchema], list[FlowEdgeSchema]]:
|
||||
"""拆分画布:业务配置写入明细表,布局存入 flow.graph。"""
|
||||
data_dict = data.model_dump(exclude_unset=True, exclude_none=True)
|
||||
graph = data_dict.pop("graph", None)
|
||||
if not graph:
|
||||
raise CustomException(msg="创建失败,请至少添加一条传输连线")
|
||||
split = self._split_graph(graph)
|
||||
await self._validate_nodes(split.sources, split.targets)
|
||||
data_dict["graph"] = split.layout
|
||||
if flow_id is not None:
|
||||
await self._save_graph(flow_id, split.nodes, split.edges)
|
||||
return data_dict, split.nodes, split.edges
|
||||
|
||||
async def create(self, data: WorkflowFlowCreateSchema) -> WorkflowFlowOutSchema:
|
||||
exist = await self._crud().get(name=data.name)
|
||||
if exist:
|
||||
raise CustomException(msg="创建失败,流程名称已存在")
|
||||
data_dict, node_items, edge_items = await self._prepare_and_save(None, data)
|
||||
obj = await self._crud().create(data=data_dict)
|
||||
await self._save_graph(obj.id, node_items, edge_items)
|
||||
return await self._to_out(obj)
|
||||
|
||||
async def update(self, id: int, data: WorkflowFlowUpdateSchema) -> WorkflowFlowOutSchema:
|
||||
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="更新失败,流程名称已存在")
|
||||
# _prepare_and_save 在 flow_id 非空时已覆写明细,无需再次 _save_graph
|
||||
data_dict, _, _ = await self._prepare_and_save(id, data)
|
||||
await self._crud().update(id=id, data=data_dict)
|
||||
obj = await self._crud().get_or_404(id=id)
|
||||
return await self._to_out(obj)
|
||||
|
||||
async def delete(self, ids: list[int]) -> None:
|
||||
if not ids:
|
||||
raise CustomException(msg="删除失败,删除对象不能为空")
|
||||
await self._crud().delete(ids=ids)
|
||||
# 子表随父流程一并物理清理(无软删消费场景)
|
||||
await WorkflowFlowNodeCRUD(self.auth, self.db).hard_delete_by_flow_ids(ids)
|
||||
await WorkflowFlowEdgeCRUD(self.auth, self.db).hard_delete_by_flow_ids(ids)
|
||||
|
||||
# ── 执行 ────────────────────────────────────────────────────────
|
||||
|
||||
async def execute(self, id: int, source_paths: dict[str, str] | None = None) -> list[int]:
|
||||
"""执行传输流程:直接读取业务明细表,每条启用的连线生成一个传输任务。
|
||||
|
||||
支持 1对多 / 多对1 拓扑;源文件/目录优先取执行时传入的 source_paths
|
||||
(按源存储源ID映射),未传入时回退使用节点配置的默认源目录;
|
||||
传输方式未配置时默认流式传输;禁用的连线不参与执行。
|
||||
"""
|
||||
obj = await self._crud().get_or_404(id=id, msg="执行失败,该流程不存在")
|
||||
node_rows, edge_rows = await self._load_flow_graph(id)
|
||||
# 只执行启用的连线(禁用的连线不生成传输任务)
|
||||
enabled_edges = [e for e in edge_rows if e.enabled]
|
||||
if not enabled_edges:
|
||||
raise CustomException(msg="执行失败,该流程画布没有启用的传输连线")
|
||||
nodes_by_key = {n.node_key: n for n in node_rows}
|
||||
source_paths = source_paths or {}
|
||||
|
||||
# 阶段一:解析并校验每条连线
|
||||
plans: list[FlowTransferPlanSchema] = []
|
||||
for e in enabled_edges:
|
||||
src_node = nodes_by_key.get(e.source_node_key)
|
||||
tgt_node = nodes_by_key.get(e.target_node_key)
|
||||
if not src_node or not tgt_node:
|
||||
raise CustomException(msg=f"执行失败,连线 {e.edge_key} 对应的节点不存在")
|
||||
src_id, tgt_id = src_node.source_id, tgt_node.source_id
|
||||
src_path = (source_paths.get(str(src_id)) or (src_node.source_path or "")).strip()
|
||||
# 目标目录由目标节点默认源目录决定
|
||||
tgt_path = (tgt_node.source_path or "").strip()
|
||||
edge_label = f"「存储源{src_id} → 存储源{tgt_id}」"
|
||||
if not src_id or not tgt_id:
|
||||
raise CustomException(msg=f"执行失败,连线 {edge_label} 存在无效的存储源节点")
|
||||
if src_id == tgt_id:
|
||||
raise CustomException(msg=f"执行失败,连线 {edge_label} 源与目标不能相同")
|
||||
if not src_path or not tgt_path:
|
||||
missing = "未选择源文件/目录" if not src_path else "目标节点未配置默认目录"
|
||||
raise CustomException(msg=f"执行失败,连线 {edge_label} {missing}")
|
||||
plans.append(
|
||||
FlowTransferPlanSchema(
|
||||
src_id=src_id,
|
||||
tgt_id=tgt_id,
|
||||
src_path=src_path,
|
||||
tgt_path=tgt_path,
|
||||
transfer_mode=e.transfer_mode or "stream",
|
||||
multipart_part_size=e.multipart_part_size,
|
||||
multipart_concurrency=e.multipart_concurrency,
|
||||
)
|
||||
)
|
||||
|
||||
# 阶段二:统一校验所有涉及的存储源可用,再逐个生成传输任务
|
||||
node_service = StorageSourceService(self.auth, self.db)
|
||||
sources = await node_service.get_active_sources(
|
||||
[p.src_id for p in plans] + [p.tgt_id for p in plans]
|
||||
)
|
||||
name_map = {s.id: s.name for s in sources}
|
||||
|
||||
transfer_service = StorageTransferService(self.auth, self.db)
|
||||
task_ids: list[int] = []
|
||||
for p in plans:
|
||||
src_label = name_map.get(p.src_id) or f"存储源{p.src_id}"
|
||||
tgt_label = name_map.get(p.tgt_id) or f"存储源{p.tgt_id}"
|
||||
# 任务名最大 128 字符,超长时截断避免 422
|
||||
task_name = f"{obj.name}-{src_label}→{tgt_label}"[:128]
|
||||
task_id = await transfer_service.create(
|
||||
TransferTaskCreateSchema(
|
||||
name=task_name,
|
||||
task_type="parallel",
|
||||
source_type="remote",
|
||||
source_id=p.src_id,
|
||||
source_path=p.src_path,
|
||||
targets=[TransferTargetSchema(target_id=p.tgt_id, target_path=p.tgt_path)],
|
||||
# flow 侧以 str 保存传输方式,接口侧字面量校验由 pydantic 兜底,此处仅收窄静态类型
|
||||
transfer_mode=cast(TransferMode | None, p.transfer_mode),
|
||||
multipart_part_size=p.multipart_part_size,
|
||||
multipart_concurrency=p.multipart_concurrency,
|
||||
)
|
||||
)
|
||||
task_ids.append(task_id)
|
||||
return task_ids
|
||||
@@ -0,0 +1,8 @@
|
||||
# 见 docs/PLUGIN_ARCHITECTURE.md
|
||||
|
||||
name = "workflow"
|
||||
title = "工作流"
|
||||
version = "1.0.0"
|
||||
description = "存储源管理、对象存储适配(OSS/COS/OBS/S3/SFTP/FTP)、文件浏览、传输任务与流程编排。"
|
||||
optional = true
|
||||
tags = ["workflow", "storage", "transfer"]
|
||||
@@ -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
|
||||
@@ -0,0 +1,174 @@
|
||||
import os
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Body, Depends, File, Form, Query, Security, UploadFile
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.common.response import ResponseSchema, SuccessResponse, UploadFileResponse
|
||||
from app.core.base_schema import AuthSchema
|
||||
from app.core.dependencies import AuthPermission, db_getter
|
||||
from app.core.router_class import OperationLogRoute
|
||||
from app.modules.workflow.core.base import StorageObject, StoragePage
|
||||
|
||||
from .schema import StoragePathCreateSchema, StoragePathResultSchema, StorageUploadResultSchema
|
||||
from .service import StorageFileService
|
||||
|
||||
StorageFileRouter = APIRouter(route_class=OperationLogRoute, prefix="/storage", tags=["存储管理"])
|
||||
|
||||
|
||||
def _delete_temp_file(path: str) -> None:
|
||||
"""响应发送后清理临时下载文件。"""
|
||||
try:
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
@StorageFileRouter.post("/upload", summary="上传文件到存储源", response_model=ResponseSchema[StorageUploadResultSchema])
|
||||
async def upload_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:upload"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
file: Annotated[UploadFile, File(description="上传文件")],
|
||||
source_id: Annotated[int | None, Form(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
remote_path: Annotated[str | None, Form(description="远端目录路径(不传自动生成文件名)")] = None,
|
||||
bucket: Annotated[str | None, Form(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: StorageUploadResultSchema = await StorageFileService(auth, db).upload(
|
||||
source_id=source_id, file=file, remote_path=remote_path, bucket=bucket
|
||||
)
|
||||
return SuccessResponse(data=result, msg="上传文件成功")
|
||||
|
||||
|
||||
@StorageFileRouter.post("/download", summary="下载存储源文件", response_model=None)
|
||||
async def download_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:download"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
background_tasks: BackgroundTasks,
|
||||
remote_path: Annotated[str, Body(description="远端文件路径")],
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> UploadFileResponse:
|
||||
local_path, file_name = await StorageFileService(auth, db).download(
|
||||
source_id=source_id, remote_path=remote_path, bucket=bucket
|
||||
)
|
||||
background_tasks.add_task(_delete_temp_file, local_path)
|
||||
return UploadFileResponse(file_path=local_path, filename=file_name)
|
||||
|
||||
|
||||
@StorageFileRouter.post("/download_dir", summary="下载存储源目录(递归打包ZIP)", response_model=None)
|
||||
async def download_dir_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:download"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
background_tasks: BackgroundTasks,
|
||||
remote_path: Annotated[str, Body(description="远端目录路径")],
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> UploadFileResponse:
|
||||
local_path, file_name = await StorageFileService(auth, db).download_dir(
|
||||
source_id=source_id, remote_path=remote_path, bucket=bucket
|
||||
)
|
||||
background_tasks.add_task(_delete_temp_file, local_path)
|
||||
return UploadFileResponse(file_path=local_path, filename=file_name)
|
||||
|
||||
|
||||
@StorageFileRouter.delete("/delete", summary="删除存储源文件", response_model=ResponseSchema[None])
|
||||
async def delete_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:delete"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
remote_path: Annotated[str, Body(description="远端文件路径")],
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
await StorageFileService(auth, db).delete(source_id=source_id, remote_path=remote_path, bucket=bucket)
|
||||
return SuccessResponse(msg="删除文件成功")
|
||||
|
||||
|
||||
@StorageFileRouter.get("/list", summary="查询存储源文件列表", response_model=None)
|
||||
async def list_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
source_id: Annotated[int | None, Query(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
prefix: Annotated[str | None, Query(description="目录前缀(可选)")] = None,
|
||||
bucket: Annotated[str | None, Query(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
page_size: Annotated[int | None, Query(ge=1, le=500, description="每页数量(传了才走游标分页,不传返回全量)")] = None,
|
||||
cursor: Annotated[str | None, Query(description="下一页游标(上一页返回的 next_cursor,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: StoragePage | list[StorageObject] = await StorageFileService(auth, db).list_files(
|
||||
source_id=source_id, prefix=prefix or "", bucket=bucket, page_size=page_size, cursor=cursor
|
||||
)
|
||||
return SuccessResponse(data=result, msg="查询文件列表成功")
|
||||
|
||||
|
||||
@StorageFileRouter.get("/buckets", summary="查询存储源桶列表", response_model=ResponseSchema[list[str]])
|
||||
async def list_storage_buckets_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
source_id: Annotated[int | None, Query(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: list[str] = await StorageFileService(auth, db).list_buckets(source_id=source_id)
|
||||
return SuccessResponse(data=result, msg="查询桶列表成功")
|
||||
|
||||
|
||||
@StorageFileRouter.post("/copy", summary="复制/移动文件", response_model=ResponseSchema[StoragePathResultSchema])
|
||||
async def copy_or_move_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
source_id: Annotated[int | None, Body(description="源存储源ID(不传使用默认存储源)")] = None,
|
||||
source_path: Annotated[str, Body(description="源文件路径")] = "",
|
||||
target_id: Annotated[int, Body(description="目标存储源ID")] = 0,
|
||||
target_path: Annotated[str, Body(description="目标路径")] = "",
|
||||
move: Annotated[bool, Body(description="是否为移动(true 移动/重命名,false 复制)")] = False,
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: StoragePathResultSchema = await StorageFileService(auth, db).copy_or_move(
|
||||
source_id=source_id,
|
||||
source_path=source_path,
|
||||
target_id=target_id,
|
||||
target_path=target_path,
|
||||
move=move,
|
||||
bucket=bucket,
|
||||
)
|
||||
return SuccessResponse(data=result, msg="操作文件成功")
|
||||
|
||||
|
||||
@StorageFileRouter.put("/rename", summary="重命名/移动文件", response_model=ResponseSchema[StoragePathResultSchema])
|
||||
async def rename_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
source_path: Annotated[str, Body(description="原路径")] = "",
|
||||
target_path: Annotated[str, Body(description="新路径")] = "",
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: StoragePathResultSchema = await StorageFileService(auth, db).rename(
|
||||
source_id=source_id, src_path=source_path, dst_path=target_path, bucket=bucket
|
||||
)
|
||||
return SuccessResponse(data=result, msg="重命名成功")
|
||||
|
||||
|
||||
@StorageFileRouter.post("/mkdir", summary="新建目录", response_model=ResponseSchema[StoragePathCreateSchema])
|
||||
async def mkdir_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
remote_dir: Annotated[str, Body(description="目录路径")] = "",
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: StoragePathCreateSchema = await StorageFileService(auth, db).mkdir(source_id=source_id, remote_dir=remote_dir, bucket=bucket)
|
||||
return SuccessResponse(data=result, msg="新建目录成功")
|
||||
|
||||
|
||||
@StorageFileRouter.post("/share", summary="生成分享链接", response_model=ResponseSchema[str | None])
|
||||
async def share_storage_file_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:storage:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
remote_path: Annotated[str, Body(description="远端文件路径")] = "",
|
||||
source_id: Annotated[int | None, Body(description="存储源ID(不传使用默认存储源)")] = None,
|
||||
expire: Annotated[int, Body(description="有效期(秒)", ge=60, le=604800)] = 3600,
|
||||
bucket: Annotated[str | None, Body(description="存储桶(对象存储多桶浏览用,可选)")] = None,
|
||||
) -> JSONResponse:
|
||||
result: str | None = await StorageFileService(auth, db).share(
|
||||
source_id=source_id, remote_path=remote_path, expire=expire, bucket=bucket
|
||||
)
|
||||
return SuccessResponse(data=result, msg="生成分享链接成功")
|
||||
@@ -0,0 +1,23 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class StorageUploadResultSchema(BaseModel):
|
||||
"""上传结果"""
|
||||
|
||||
file_path: str = Field(..., description="远端路径")
|
||||
file_name: str = Field(..., description="远端文件名")
|
||||
origin_name: str = Field(..., description="原始文件名")
|
||||
file_url: str | None = Field(default=None, description="访问链接")
|
||||
|
||||
|
||||
class StoragePathResultSchema(BaseModel):
|
||||
"""路径操作结果(复制/移动/重命名)"""
|
||||
|
||||
source_path: str = Field(..., description="源路径")
|
||||
target_path: str = Field(..., description="目标路径")
|
||||
|
||||
|
||||
class StoragePathCreateSchema(BaseModel):
|
||||
"""新建目录结果"""
|
||||
|
||||
path: str = Field(..., description="目录路径")
|
||||
@@ -0,0 +1,323 @@
|
||||
import asyncio
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
import zipfile
|
||||
from pathlib import Path
|
||||
|
||||
import aiofiles
|
||||
from fastapi import UploadFile
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_schema import AuthSchema
|
||||
from app.core.exceptions import CustomException
|
||||
from app.modules.workflow.core.base import (
|
||||
_OBJECT_STORE_PROTOCOLS,
|
||||
BaseStorageAdapter,
|
||||
StorageAdapterConfig,
|
||||
StorageObject,
|
||||
StoragePage,
|
||||
)
|
||||
from app.modules.workflow.core.factory import StorageAdapterFactory
|
||||
from app.modules.workflow.source.service import StorageSourceService
|
||||
from app.utils.upload_util import UploadUtil
|
||||
|
||||
from .schema import StoragePathCreateSchema, StoragePathResultSchema, StorageUploadResultSchema
|
||||
|
||||
|
||||
class StorageFileService:
|
||||
"""存储文件操作服务(上传/下载/删除/列表/预签名URL)"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
self.auth = auth
|
||||
self.db = db
|
||||
|
||||
# ── 内部工具 ────────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _validate_remote_path(remote_path: str) -> str:
|
||||
"""规范化并校验远端相对路径(禁止路径穿越)。"""
|
||||
if not remote_path or not remote_path.strip():
|
||||
raise CustomException(msg="请提供文件路径")
|
||||
parts = [p for p in remote_path.replace("\\", "/").split("/") if p not in ("", ".")]
|
||||
if any(p == ".." for p in parts) or "\x00" in remote_path:
|
||||
raise CustomException(msg="非法的文件路径")
|
||||
return "/".join(parts)
|
||||
|
||||
async def _get_source(self, source_id: int | None) -> StorageAdapterConfig:
|
||||
"""获取存储源并构造适配器配置(密码已解密,含 SDK 高级配置)。"""
|
||||
source = await StorageSourceService(self.auth, self.db).get_active_source(source_id)
|
||||
return StorageSourceService._build_config(source)
|
||||
|
||||
async def _get_adapter(self, source_id: int | None, bucket: str | None = None) -> BaseStorageAdapter:
|
||||
"""构造适配器并切换当前操作桶(对象存储多桶浏览:桶随请求传入,适配器按请求创建)。"""
|
||||
config = await self._get_source(source_id)
|
||||
adapter = StorageAdapterFactory.create(config)
|
||||
if bucket:
|
||||
adapter.set_bucket(bucket)
|
||||
return adapter
|
||||
|
||||
@staticmethod
|
||||
def _entries(result: list[StorageObject] | StoragePage) -> list[StorageObject]:
|
||||
"""全量列举结果规整为条目列表(探测/目录守卫等需要遍历条目的场景)。"""
|
||||
return result.items if isinstance(result, StoragePage) else result
|
||||
|
||||
@staticmethod
|
||||
async def _save_to_temp(file: UploadFile, suffix: str = "") -> str:
|
||||
"""将上传文件内容落盘到系统临时目录,返回临时路径。"""
|
||||
fd, path = tempfile.mkstemp(suffix=suffix)
|
||||
os.close(fd)
|
||||
try:
|
||||
async with aiofiles.open(path, "wb") as f:
|
||||
while chunk := await file.read(1024 * 1024):
|
||||
await f.write(chunk)
|
||||
except Exception:
|
||||
os.unlink(path)
|
||||
raise
|
||||
finally:
|
||||
await file.seek(0)
|
||||
return path
|
||||
|
||||
@staticmethod
|
||||
def _zip_directory(src_root: Path, zip_path: str, arc_root: Path) -> None:
|
||||
"""把 src_root 下所有文件打包到 zip_path(相对于 arc_root 计算归档名)。
|
||||
|
||||
ZIP_DEFLATED 是纯 CPU + 磁盘 IO 的同步操作,目录较大时会在事件循环里
|
||||
独占数秒,因此统一由调用方通过 asyncio.to_thread 在线程中执行。
|
||||
"""
|
||||
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
for f in src_root.rglob("*"):
|
||||
if f.is_file():
|
||||
zf.write(f, f.relative_to(arc_root).as_posix())
|
||||
|
||||
# ── 业务方法 ────────────────────────────────────────────────────
|
||||
|
||||
async def upload(
|
||||
self,
|
||||
source_id: int | None,
|
||||
file: UploadFile,
|
||||
remote_path: str | None = None,
|
||||
bucket: str | None = None,
|
||||
) -> StorageUploadResultSchema:
|
||||
"""上传文件到远端存储。remote_path 为空时自动生成安全文件名。"""
|
||||
if not file or not file.filename:
|
||||
raise CustomException(msg="请选择要上传的文件")
|
||||
|
||||
if not UploadUtil.check_path_traversal(file.filename):
|
||||
raise CustomException(msg="文件名包含非法字符")
|
||||
extension = UploadUtil.get_extension_from_filename(file.filename)
|
||||
if not extension:
|
||||
raise CustomException(msg="无法识别文件类型")
|
||||
if UploadUtil.is_dangerous_extension(extension):
|
||||
raise CustomException(msg=f"不允许上传此类型的文件: {extension}")
|
||||
UploadUtil.check_file_size(file)
|
||||
|
||||
# 确定远端路径
|
||||
if remote_path:
|
||||
if remote_path.endswith("/"):
|
||||
# 以 / 结尾视为目录:保留原文件名,拼接到目录下
|
||||
dir_path = self._validate_remote_path(remote_path)
|
||||
target = f"{dir_path}/{file.filename}"
|
||||
else:
|
||||
target = self._validate_remote_path(remote_path)
|
||||
# 大小写不敏感比较扩展名,避免 photo.JPG 被追加成 photo.JPG.jpg
|
||||
if not target.lower().endswith(extension.lower()):
|
||||
target = f"{target}{extension}"
|
||||
else:
|
||||
target = UploadUtil.generate_safe_filename(file.filename, extension)
|
||||
|
||||
temp_path = await self._save_to_temp(file, suffix=extension)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
await adapter.upload(temp_path, target)
|
||||
file_url = await adapter.get_url(target)
|
||||
finally:
|
||||
await adapter.close()
|
||||
os.unlink(temp_path)
|
||||
|
||||
return StorageUploadResultSchema(
|
||||
file_path=target,
|
||||
file_name=Path(target).name,
|
||||
origin_name=file.filename,
|
||||
file_url=file_url,
|
||||
)
|
||||
|
||||
async def download(self, source_id: int | None, remote_path: str, bucket: str | None = None) -> tuple[str, str]:
|
||||
"""下载远端文件到临时目录,返回 (本地临时路径, 文件名)。"""
|
||||
target = self._validate_remote_path(remote_path)
|
||||
extension = Path(target).suffix
|
||||
fd, temp_path = tempfile.mkstemp(suffix=extension)
|
||||
os.close(fd)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
local_path = await adapter.download(target, temp_path)
|
||||
except Exception:
|
||||
os.unlink(temp_path)
|
||||
raise
|
||||
finally:
|
||||
await adapter.close()
|
||||
return local_path, Path(target).name
|
||||
|
||||
async def download_dir(self, source_id: int | None, remote_path: str, bucket: str | None = None) -> tuple[str, str]:
|
||||
"""递归下载目录并打包 ZIP,返回 (zip临时路径, zip文件名)。"""
|
||||
target = self._validate_remote_path(remote_path)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
tmp_root = tempfile.mkdtemp(prefix="stor_dir_")
|
||||
zip_path = ""
|
||||
try:
|
||||
dir_name = Path(target).name or "download"
|
||||
local_dir = Path(tmp_root) / dir_name
|
||||
local_dir.mkdir(parents=True, exist_ok=True)
|
||||
await adapter.download_dir(target, str(local_dir), concurrency=3)
|
||||
fd, zip_path = tempfile.mkstemp(suffix=".zip")
|
||||
os.close(fd)
|
||||
await asyncio.to_thread(self._zip_directory, local_dir, zip_path, Path(tmp_root))
|
||||
except Exception:
|
||||
if zip_path:
|
||||
os.unlink(zip_path)
|
||||
raise
|
||||
finally:
|
||||
await adapter.close()
|
||||
await asyncio.to_thread(shutil.rmtree, tmp_root, ignore_errors=True)
|
||||
return zip_path, f"{dir_name}.zip"
|
||||
|
||||
async def delete(self, source_id: int | None, remote_path: str, bucket: str | None = None) -> None:
|
||||
target = self._validate_remote_path(remote_path)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
# 目录探测:对象存储对不存在的目录 key 执行 delete 会幂等成功(不抛异常),
|
||||
# 无法触发 delete_dir 回退,导致子文件残留;先探测再决定删除策略。
|
||||
is_dir = False
|
||||
try:
|
||||
objects = self._entries(await adapter.list_files(target))
|
||||
t = target.rstrip("/")
|
||||
is_dir = any(o.is_dir or o.key.startswith(f"{t}/") for o in objects)
|
||||
except Exception:
|
||||
is_dir = False
|
||||
if is_dir:
|
||||
await adapter.delete_dir(target)
|
||||
else:
|
||||
try:
|
||||
await adapter.delete(target)
|
||||
except Exception:
|
||||
# 文件删除失败时回退为递归删除:兼容 FTP/SFTP 等协议无原生目录删除
|
||||
await adapter.delete_dir(target)
|
||||
finally:
|
||||
await adapter.close()
|
||||
|
||||
async def exists(self, source_id: int | None, remote_path: str, bucket: str | None = None) -> bool:
|
||||
target = self._validate_remote_path(remote_path)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
return await adapter.exists(target)
|
||||
finally:
|
||||
await adapter.close()
|
||||
|
||||
async def list_files(
|
||||
self,
|
||||
source_id: int | None,
|
||||
prefix: str = "",
|
||||
bucket: str | None = None,
|
||||
page_size: int | None = None,
|
||||
cursor: str | None = None,
|
||||
) -> list[StorageObject] | StoragePage:
|
||||
"""列出目录条目。
|
||||
|
||||
- 不传 page_size:全量拉取(目录选择器/搜索等场景)。
|
||||
- 传 page_size:游标分页(对象存储走 SDK 原生游标,FTP/SFTP/LOCAL 走内存切片)。
|
||||
"""
|
||||
safe_prefix = self._validate_remote_path(prefix) if prefix else ""
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
if page_size is not None:
|
||||
return await adapter.list_files(safe_prefix, page_size=page_size, cursor=cursor)
|
||||
return await adapter.list_files(safe_prefix)
|
||||
finally:
|
||||
await adapter.close()
|
||||
|
||||
async def list_buckets(self, source_id: int | None) -> list[str]:
|
||||
"""列出账号下全部存储桶(仅对象存储协议;FTP/SFTP/LOCAL 无桶概念返回空列表)。"""
|
||||
config = await self._get_source(source_id)
|
||||
if config.protocol not in _OBJECT_STORE_PROTOCOLS:
|
||||
return []
|
||||
adapter = StorageAdapterFactory.create(config)
|
||||
try:
|
||||
return await adapter.list_buckets()
|
||||
finally:
|
||||
await adapter.close()
|
||||
|
||||
async def copy_or_move(
|
||||
self,
|
||||
source_id: int | None,
|
||||
source_path: str,
|
||||
target_id: int,
|
||||
target_path: str,
|
||||
move: bool = False,
|
||||
bucket: str | None = None,
|
||||
) -> StoragePathResultSchema:
|
||||
"""复制/移动文件:跨端点时下载到临时再上传;同端点 move 即重命名。"""
|
||||
src = self._validate_remote_path(source_path)
|
||||
dst = self._validate_remote_path(target_path)
|
||||
if move and source_id == target_id and src == dst:
|
||||
raise CustomException(msg="源路径与目标路径相同")
|
||||
source_config = await self._get_source(source_id)
|
||||
target_config = await self._get_source(target_id)
|
||||
src_adapter = StorageAdapterFactory.create(source_config)
|
||||
dst_adapter = StorageAdapterFactory.create(target_config)
|
||||
# 多桶浏览:操作在指定桶内进行;同端点复制/移动目标桶与源桶一致
|
||||
if bucket:
|
||||
src_adapter.set_bucket(bucket)
|
||||
if target_id == source_id:
|
||||
dst_adapter.set_bucket(bucket)
|
||||
# 目录守卫:本接口为文件级实现(临时文件 download/upload),目录会走进
|
||||
# download_dir 的目录逻辑而报错,且可能移入自身子树导致数据丢失,显式拦截。
|
||||
try:
|
||||
entries = self._entries(await src_adapter.list_files(src))
|
||||
except Exception:
|
||||
entries = []
|
||||
if any((e.is_dir and e.key.rstrip("/") == src) or e.key.startswith(f"{src}/") for e in entries):
|
||||
raise CustomException(msg="暂不支持目录复制/移动,请使用传输任务")
|
||||
fd, temp_path = tempfile.mkstemp(suffix=Path(dst).suffix)
|
||||
os.close(fd)
|
||||
try:
|
||||
await src_adapter.download(src, temp_path)
|
||||
await dst_adapter.upload(temp_path, dst)
|
||||
if move:
|
||||
await src_adapter.delete(src)
|
||||
finally:
|
||||
await src_adapter.close()
|
||||
await dst_adapter.close()
|
||||
os.unlink(temp_path)
|
||||
return StoragePathResultSchema(source_path=src, target_path=dst)
|
||||
|
||||
async def rename(self, source_id: int | None, src_path: str, dst_path: str, bucket: str | None = None) -> StoragePathResultSchema:
|
||||
"""重命名/移动(同一存储源内)。"""
|
||||
src = self._validate_remote_path(src_path)
|
||||
dst = self._validate_remote_path(dst_path)
|
||||
if src == dst:
|
||||
raise CustomException(msg="源路径与目标路径相同")
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
await adapter.rename(src, dst)
|
||||
finally:
|
||||
await adapter.close()
|
||||
return StoragePathResultSchema(source_path=src, target_path=dst)
|
||||
|
||||
async def mkdir(self, source_id: int | None, remote_dir: str, bucket: str | None = None) -> StoragePathCreateSchema:
|
||||
"""新建目录。"""
|
||||
path = self._validate_remote_path(remote_dir)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
await adapter.mkdir(path)
|
||||
finally:
|
||||
await adapter.close()
|
||||
return StoragePathCreateSchema(path=path)
|
||||
|
||||
async def share(self, source_id: int | None, remote_path: str, expire: int = 3600, bucket: str | None = None) -> str | None:
|
||||
"""生成分享链接(对象存储为预签名 URL,FTP/SFTP/LOCAL 返回 None)。"""
|
||||
target = self._validate_remote_path(remote_path)
|
||||
adapter = await self._get_adapter(source_id, bucket)
|
||||
try:
|
||||
return await adapter.get_url(target, expire=expire)
|
||||
finally:
|
||||
await adapter.close()
|
||||
@@ -0,0 +1 @@
|
||||
"""文件传输与工作流模块"""
|
||||
@@ -0,0 +1,124 @@
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Annotated, cast
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, File, Form, Path, Query, Request, Security, UploadFile
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.sse import EventSourceResponse, ServerSentEvent
|
||||
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.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.core.router_class import OperationLogRoute
|
||||
from app.core.sse_manager import SSE_QUEUE_MAX_SIZE
|
||||
from app.modules.workflow.transfer.schema import (
|
||||
TransferTaskCreateResultSchema,
|
||||
TransferTaskCreateSchema,
|
||||
TransferTaskOutSchema,
|
||||
TransferTaskQueryParam,
|
||||
TransferTaskType,
|
||||
)
|
||||
from app.modules.workflow.transfer.service import StorageTransferService
|
||||
from app.modules.workflow.transfer.sse_manager import transfer_stream_manager
|
||||
|
||||
StorageTransferRouter = APIRouter(route_class=OperationLogRoute, prefix="/transfer", tags=["文件传输"])
|
||||
|
||||
|
||||
@StorageTransferRouter.post("/task", summary="创建传输任务(远端源)", response_model=ResponseSchema[TransferTaskCreateResultSchema])
|
||||
async def create_transfer_task_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:create"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
data: Annotated[TransferTaskCreateSchema, Body(description="任务参数(远端源)")],
|
||||
) -> JSONResponse:
|
||||
task_id: int = await StorageTransferService(auth, db).create(data=data)
|
||||
return SuccessResponse(data=TransferTaskCreateResultSchema(id=task_id), msg="创建传输任务成功")
|
||||
|
||||
|
||||
@StorageTransferRouter.post("/task/upload", summary="创建传输任务(本地上传源)", response_model=ResponseSchema[TransferTaskCreateResultSchema])
|
||||
async def create_local_transfer_task_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:create"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
file: Annotated[UploadFile, File(description="本地源文件")],
|
||||
name: Annotated[str, Form(description="任务名称")],
|
||||
task_type: Annotated[str, Form(description="任务类型(parallel:多目标 chain:链式)")],
|
||||
targets: Annotated[str, Form(description="目标列表JSON,如 [{\"target_id\":1,\"target_path\":\"a.txt\"}]")],
|
||||
) -> JSONResponse:
|
||||
try:
|
||||
targets_data = json.loads(targets)
|
||||
except (json.JSONDecodeError, TypeError) as e:
|
||||
raise CustomException(msg=f"targets 参数格式错误: {e!s}") from e
|
||||
data = TransferTaskCreateSchema(name=name, task_type=cast("TransferTaskType", task_type), source_type="local", targets=targets_data)
|
||||
task_id: int = await StorageTransferService(auth, db).create_local(data=data, file=file)
|
||||
return SuccessResponse(data=TransferTaskCreateResultSchema(id=task_id), msg="创建传输任务成功")
|
||||
|
||||
|
||||
@StorageTransferRouter.get("/task/page", summary="分页查询传输任务", response_model=ResponseSchema[PageResultSchema[TransferTaskOutSchema]])
|
||||
async def get_transfer_task_page_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
page: Annotated[PaginationQueryParam, Depends()],
|
||||
search: Annotated[TransferTaskQueryParam, Query()],
|
||||
) -> JSONResponse:
|
||||
result: PageResultSchema[TransferTaskOutSchema] = await StorageTransferService(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="查询传输任务分页成功")
|
||||
|
||||
|
||||
@StorageTransferRouter.get("/task/{id}", summary="查询传输任务详情", response_model=ResponseSchema[TransferTaskOutSchema])
|
||||
async def get_transfer_task_detail_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:query"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="任务ID", ge=1)],
|
||||
) -> JSONResponse:
|
||||
result: TransferTaskOutSchema = await StorageTransferService(auth, db).detail(task_id=id)
|
||||
return SuccessResponse(data=result, msg="查询传输任务详情成功")
|
||||
|
||||
|
||||
@StorageTransferRouter.post("/task/{id}/cancel", summary="取消传输任务", response_model=ResponseSchema[None])
|
||||
async def cancel_transfer_task_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:update"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
id: Annotated[int, Path(description="任务ID", ge=1)],
|
||||
) -> JSONResponse:
|
||||
await StorageTransferService(auth, db).cancel(task_id=id)
|
||||
return SuccessResponse(msg="已请求取消传输任务")
|
||||
|
||||
|
||||
@StorageTransferRouter.delete("/task", summary="删除传输任务", response_model=ResponseSchema[None])
|
||||
async def delete_transfer_task_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission(["module_task:workflow:transfer:delete"]))],
|
||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||
ids: Annotated[list[int], Body(description="任务ID列表")],
|
||||
) -> JSONResponse:
|
||||
await StorageTransferService(auth, db).delete(ids=ids)
|
||||
return SuccessResponse(msg="删除传输任务成功")
|
||||
|
||||
|
||||
@StorageTransferRouter.get("/stream", summary="传输任务进度实时流(SSE)", response_class=EventSourceResponse)
|
||||
async def transfer_stream_controller(
|
||||
auth: Annotated[AuthSchema, Security(AuthPermission())],
|
||||
request: Request,
|
||||
) -> AsyncGenerator[ServerSentEvent, None]:
|
||||
"""传输任务实时进度通道(SSE):任务进度按创建者推送,客户端断开自动重连。
|
||||
|
||||
令牌经 Authorization 头携带(复用 HTTP 认证链,不进 URL);事件名 task_update,
|
||||
载荷为任务完整状态(含步骤)。空闲保活由框架内置 ping(15s 注释行)处理。
|
||||
"""
|
||||
queue: asyncio.Queue = asyncio.Queue(maxsize=SSE_QUEUE_MAX_SIZE)
|
||||
transfer_stream_manager.connect(auth.user.id, queue, redis=getattr(request.app.state, "redis", None))
|
||||
logger.info("传输进度 SSE 已连接: user={}", auth.user.id)
|
||||
try:
|
||||
while True:
|
||||
message = await queue.get()
|
||||
yield ServerSentEvent(event=str(message.get("type", "message")), data=message.get("data", message))
|
||||
finally:
|
||||
transfer_stream_manager.disconnect(auth.user.id, queue)
|
||||
logger.info("传输进度 SSE 已断开: user={}", auth.user.id)
|
||||
@@ -0,0 +1,13 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_crud import CRUDBase
|
||||
from app.core.base_schema import AuthSchema
|
||||
|
||||
from .model import StorageTransferTaskModel
|
||||
|
||||
|
||||
class StorageTransferTaskCRUD(CRUDBase[StorageTransferTaskModel, object, object]):
|
||||
"""文件传输任务 CRUD"""
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
super().__init__(StorageTransferTaskModel, auth, db)
|
||||
@@ -0,0 +1,382 @@
|
||||
"""文件传输任务执行引擎
|
||||
|
||||
- parallel(多目标):单源依次输出到多个目标端点
|
||||
- chain(链式):步骤串联,上一步目标端点即下一步源,链条长度不限
|
||||
- 进度按步骤粒度统计(SDK 无逐字节回调),实时写入 DB 并经 SSE 推送;广播仅由
|
||||
任务生命周期事件驱动(终态只推一次收尾帧),前端按"有无进行中任务"维持连接
|
||||
- 后台任务在独立 DB 会话中运行,不阻塞请求;取消采用内存标志(当前步骤执行完毕后生效)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
from datetime import UTC, datetime
|
||||
from typing import cast
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.database import async_db_session
|
||||
from app.core.logger import logger
|
||||
from app.modules.workflow.core.base import StorageAdapterConfig
|
||||
from app.modules.workflow.core.factory import StorageAdapterFactory
|
||||
from app.modules.workflow.source.model import StorageSourceModel
|
||||
from app.modules.workflow.source.service import StorageSourceService
|
||||
from app.modules.workflow.transfer.registry import transfer_task_registry
|
||||
from app.modules.workflow.transfer.sse_manager import transfer_stream_manager
|
||||
|
||||
from .model import StorageTransferStepModel, StorageTransferTaskModel
|
||||
from .schema import TransferMode, TransferStepOutSchema, TransferTaskOutSchema
|
||||
|
||||
# 单步最大可显示进度(步骤执行中为流动状态,完成后置 100)
|
||||
_STEP_RUNNING_PROGRESS = 50
|
||||
|
||||
|
||||
def _step_payload(step: StorageTransferStepModel) -> dict:
|
||||
return TransferStepOutSchema.model_validate(step).model_dump(mode="json")
|
||||
|
||||
|
||||
def _task_payload(task: StorageTransferTaskModel, steps: list[StorageTransferStepModel]) -> dict:
|
||||
out = TransferTaskOutSchema.model_validate(task)
|
||||
out.steps = [TransferStepOutSchema.model_validate(s) for s in steps]
|
||||
return out.model_dump(mode="json")
|
||||
|
||||
|
||||
async def _broadcast(task: StorageTransferTaskModel, steps: list[StorageTransferStepModel]) -> None:
|
||||
# 进度广播是尽力而为:SSE 抖动/客户端断连不得中断传输流水线本身
|
||||
try:
|
||||
await transfer_stream_manager.send_to_user(
|
||||
task.created_id,
|
||||
{"type": "task_update", "data": _task_payload(task, steps)},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("传输任务 {} 进度广播失败(不影响传输): {}", task.id, e)
|
||||
|
||||
|
||||
async def _build_config(db: AsyncSession, source_id: int) -> StorageAdapterConfig | None:
|
||||
"""获取存储源并构造适配器配置(复用存储源服务,含 SDK 高级配置);不存在或停用返回 None。"""
|
||||
source = await db.get(StorageSourceModel, source_id)
|
||||
if source is None or source.status == 1:
|
||||
return None
|
||||
return StorageSourceService._build_config(source)
|
||||
|
||||
|
||||
def _remove_local_temp(task: StorageTransferTaskModel) -> None:
|
||||
"""清理本地源任务的临时文件(幂等:文件不存在时静默忽略)。目录源为本地真实目录,不删除。"""
|
||||
if task.source_type == "local" and task.source_path and os.path.isfile(task.source_path):
|
||||
try:
|
||||
os.unlink(task.source_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _local_path_size(path: str) -> int:
|
||||
"""计算本地文件/目录总大小;目录需递归遍历,属阻塞操作,调用方须经 to_thread 执行。"""
|
||||
if os.path.isfile(path):
|
||||
return os.path.getsize(path)
|
||||
if os.path.isdir(path):
|
||||
return sum(os.path.getsize(os.path.join(r, f)) for r, _, files in os.walk(path) for f in files)
|
||||
return 0
|
||||
|
||||
|
||||
async def _remove_local_temp_dir(temp_dir: str) -> None:
|
||||
"""清理下载到临时目录的远端目录(递归删除属阻塞操作)。"""
|
||||
if temp_dir and os.path.isdir(temp_dir):
|
||||
await asyncio.to_thread(shutil.rmtree, temp_dir, ignore_errors=True)
|
||||
|
||||
|
||||
async def _resolve_source_size(adapter, source_path: str) -> int:
|
||||
"""尽力获取远端源大小:文件精确匹配,目录递归求和(各协议通用)。"""
|
||||
try:
|
||||
objects = await adapter.list_files(source_path)
|
||||
target = source_path.strip("/")
|
||||
for obj in objects:
|
||||
if obj.is_dir or obj.key.startswith(f"{target}/"):
|
||||
total = 0
|
||||
for e in await adapter.list_recursive(source_path):
|
||||
if not e.is_dir and e.size:
|
||||
total += e.size
|
||||
return total
|
||||
for obj in objects:
|
||||
if not obj.is_dir and obj.key == target and obj.size:
|
||||
return obj.size
|
||||
except Exception as e:
|
||||
logger.warning("获取远端源大小失败,进度不显示总大小: {}: {}", source_path, e)
|
||||
return 0
|
||||
|
||||
|
||||
async def _is_dir_path(adapter, path: str) -> bool:
|
||||
"""判断远端路径是否为目录。
|
||||
|
||||
列目录失败直接抛出:若吞掉异常返回 False,目录源会被误判为单文件
|
||||
走进错误分支,网络抖动场景下造成静默的数据丢失/漏传。
|
||||
"""
|
||||
objects = await adapter.list_files(path)
|
||||
target = path.strip("/")
|
||||
return any(o.is_dir or o.key.startswith(f"{target}/") for o in objects)
|
||||
|
||||
|
||||
def _list_local_files(root: str) -> list[tuple[str, str, int]]:
|
||||
"""递归列出目录下全部文件,返回 (绝对路径, 相对路径, 字节数)。
|
||||
|
||||
整棵树的一次性遍历属阻塞操作,调用方须经 to_thread 执行,避免大目录
|
||||
在事件循环里长时间占位。
|
||||
"""
|
||||
files: list[tuple[str, str, int]] = []
|
||||
for dirpath, _, names in os.walk(root):
|
||||
for name in names:
|
||||
local_file = os.path.join(dirpath, name)
|
||||
try:
|
||||
size = os.path.getsize(local_file)
|
||||
except OSError:
|
||||
size = 0
|
||||
files.append((local_file, os.path.relpath(local_file, root).replace(os.sep, "/"), size))
|
||||
return files
|
||||
|
||||
|
||||
def _prepare_local_dirs(paths: list[str]) -> None:
|
||||
"""批量创建目标文件的父目录(去重,避免逐文件 mkdir 的重复系统调用)。"""
|
||||
for parent in {os.path.dirname(p) for p in paths if os.path.dirname(p)}:
|
||||
os.makedirs(parent, exist_ok=True)
|
||||
|
||||
|
||||
async def _download_dir(adapter, source_path: str, local_dir: str) -> None:
|
||||
"""递归下载远端目录到本地临时目录(保留相对结构)。"""
|
||||
base = source_path.strip("/")
|
||||
targets = [
|
||||
(
|
||||
e.key,
|
||||
os.path.join(
|
||||
local_dir,
|
||||
(e.key[len(base) + 1 :] if base and e.key.startswith(base + "/") else e.key).replace("/", os.sep),
|
||||
),
|
||||
)
|
||||
for e in await adapter.list_recursive(source_path)
|
||||
if not e.is_dir
|
||||
]
|
||||
await asyncio.to_thread(_prepare_local_dirs, [local_file for _, local_file in targets])
|
||||
for key, local_file in targets:
|
||||
await adapter.download(key, local_file)
|
||||
|
||||
|
||||
async def _run_step(db: AsyncSession, task: StorageTransferTaskModel, step: StorageTransferStepModel) -> bool:
|
||||
"""执行单个传输步骤,成功返回 True。"""
|
||||
started_at = datetime.now(UTC)
|
||||
step.status = "running"
|
||||
step.started_at = started_at
|
||||
step.progress = _STEP_RUNNING_PROGRESS
|
||||
await db.commit()
|
||||
await _broadcast(task, await _load_steps(db, task.id))
|
||||
|
||||
temp_path: str | None = None
|
||||
temp_dir: str | None = None
|
||||
src_adapter = None
|
||||
dst_adapter = None
|
||||
try:
|
||||
# 解析源:目录(远端递归下载/本地直接引用)与单文件分别处理
|
||||
is_dir_source = False
|
||||
if step.source_id is not None:
|
||||
src_config = await _build_config(db, step.source_id)
|
||||
if src_config is None:
|
||||
raise RuntimeError(f"源存储源 {step.source_id} 不存在或已停用")
|
||||
src_adapter = StorageAdapterFactory.create(src_config)
|
||||
is_dir_source = await _is_dir_path(src_adapter, step.source_path or "")
|
||||
if is_dir_source:
|
||||
temp_dir = tempfile.mkdtemp(prefix="transfer_dir_")
|
||||
await _download_dir(src_adapter, step.source_path or "", temp_dir)
|
||||
else:
|
||||
fd, temp_path = tempfile.mkstemp(prefix="transfer_", suffix=os.path.splitext(step.target_path)[1])
|
||||
os.close(fd)
|
||||
await src_adapter.download(step.source_path or "", temp_path)
|
||||
elif os.path.isdir(step.source_path or ""):
|
||||
is_dir_source = True
|
||||
temp_dir = step.source_path or ""
|
||||
else:
|
||||
temp_path = step.source_path or ""
|
||||
|
||||
if temp_path and not os.path.exists(temp_path):
|
||||
raise RuntimeError("源文件不存在")
|
||||
if temp_dir and not os.path.isdir(temp_dir):
|
||||
raise RuntimeError("源目录不存在")
|
||||
if not temp_path and not temp_dir:
|
||||
raise RuntimeError("源文件不存在")
|
||||
|
||||
dst_config = await _build_config(db, step.target_id)
|
||||
if dst_config is None:
|
||||
raise RuntimeError(f"目标存储源 {step.target_id} 不存在或已停用")
|
||||
# 连线/任务级传输参数覆盖目标端点配置(分片上传发生在目标端点);未指定则用存储源默认
|
||||
# DB 列以 str 保存传输方式(上游 schema 已按 Literal 校验),此处仅收窄静态类型
|
||||
if step.transfer_mode:
|
||||
dst_config.transfer_mode = cast(TransferMode | None, step.transfer_mode)
|
||||
if step.multipart_part_size:
|
||||
dst_config.multipart_part_size = step.multipart_part_size
|
||||
if step.multipart_concurrency:
|
||||
dst_config.multipart_concurrency = step.multipart_concurrency
|
||||
dst_adapter = StorageAdapterFactory.create(dst_config)
|
||||
|
||||
size = 0
|
||||
if is_dir_source:
|
||||
# 目录上传:一次性列出本地临时目录结构(遍历在线程中完成),再按相对结构写回目标路径
|
||||
assert temp_dir is not None # 目录源已落盘到临时目录或引用本地真实目录
|
||||
for local_file, rel, file_size in await asyncio.to_thread(_list_local_files, temp_dir):
|
||||
remote = f"{step.target_path}/{rel}".strip("/")
|
||||
await dst_adapter.upload(local_file, remote)
|
||||
size += file_size
|
||||
else:
|
||||
assert temp_path is not None # 单文件源:已落盘或指向本地文件,且上面已校验存在
|
||||
size = os.path.getsize(temp_path)
|
||||
await dst_adapter.upload(temp_path, step.target_path)
|
||||
|
||||
elapsed = (datetime.now(UTC) - started_at).total_seconds() or 0.01
|
||||
speed = size / elapsed
|
||||
step.total_size = size
|
||||
step.transferred_size = size
|
||||
step.speed = speed
|
||||
step.status = "success"
|
||||
step.progress = 100
|
||||
step.finished_at = datetime.now(UTC)
|
||||
task.transferred_size += size
|
||||
task.speed = speed
|
||||
if task.total_size > 0:
|
||||
task.progress = min(99, int(task.transferred_size * 100 / task.total_size))
|
||||
await db.commit()
|
||||
await _broadcast(task, await _load_steps(db, task.id))
|
||||
return True
|
||||
except Exception as e:
|
||||
msg = str(e) or e.__class__.__name__
|
||||
step.status = "failed"
|
||||
step.error_msg = msg
|
||||
step.finished_at = datetime.now(UTC)
|
||||
task.status = "failed"
|
||||
task.error_msg = msg
|
||||
task.finished_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
await _broadcast(task, await _load_steps(db, task.id))
|
||||
logger.warning("传输任务 {}(步骤 {}) 失败: {}", task.id, step.step_order, msg)
|
||||
return False
|
||||
finally:
|
||||
if src_adapter is not None:
|
||||
await src_adapter.close()
|
||||
if dst_adapter is not None:
|
||||
await dst_adapter.close()
|
||||
if temp_path and step.source_id is not None and os.path.exists(temp_path):
|
||||
os.unlink(temp_path)
|
||||
# 远端目录源下载到临时目录,必须清理;本地目录源为真实目录,不删除
|
||||
if step.source_id is not None:
|
||||
await _remove_local_temp_dir(temp_dir or "")
|
||||
|
||||
|
||||
async def _load_steps(db: AsyncSession, task_id: int) -> list[StorageTransferStepModel]:
|
||||
result = await db.execute(
|
||||
select(StorageTransferStepModel)
|
||||
.where(
|
||||
StorageTransferStepModel.task_id == task_id,
|
||||
StorageTransferStepModel.is_deleted.is_(False),
|
||||
)
|
||||
.order_by(StorageTransferStepModel.step_order)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def execute_transfer_task(task_id: int) -> None:
|
||||
"""后台执行传输任务(由创建接口在事务提交后启动)。
|
||||
|
||||
外壳兜底:任何未预期异常(前置解析、db.refresh 等)都必须把任务从
|
||||
pending/running 落到 failed,否则任务永久卡在进行中状态无法收敛。
|
||||
"""
|
||||
try:
|
||||
await _execute_transfer_task(task_id)
|
||||
except Exception as e:
|
||||
logger.exception("传输任务 {} 执行异常中断", task_id)
|
||||
try:
|
||||
async with async_db_session() as db:
|
||||
task = await db.get(StorageTransferTaskModel, task_id)
|
||||
if task is not None and task.status in ("pending", "running"):
|
||||
task.status = "failed"
|
||||
task.error_msg = f"任务执行异常中断: {e}"[:500]
|
||||
task.finished_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
await _broadcast(task, await _load_steps(db, task_id))
|
||||
except Exception:
|
||||
logger.exception("传输任务 {} 失败状态回写异常", task_id)
|
||||
|
||||
|
||||
async def _execute_transfer_task(task_id: int) -> None:
|
||||
"""传输任务执行主体。"""
|
||||
async with async_db_session() as db:
|
||||
task = await db.get(StorageTransferTaskModel, task_id)
|
||||
if task is None or task.status != "pending":
|
||||
return
|
||||
steps = await _load_steps(db, task_id)
|
||||
if not steps:
|
||||
task.status = "failed"
|
||||
task.error_msg = "任务没有可执行的步骤"
|
||||
task.finished_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
return
|
||||
|
||||
# 解析源文件大小,用于总进度估算
|
||||
if task.source_type == "local" and task.source_path:
|
||||
task.source_size = await asyncio.to_thread(_local_path_size, task.source_path)
|
||||
elif task.source_type == "remote" and task.source_id:
|
||||
config = await _build_config(db, task.source_id)
|
||||
if config is None:
|
||||
task.status = "failed"
|
||||
task.error_msg = f"源存储源 {task.source_id} 不存在或已停用"
|
||||
task.finished_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
await _broadcast(task, steps)
|
||||
return
|
||||
adapter = StorageAdapterFactory.create(config)
|
||||
try:
|
||||
task.source_size = await _resolve_source_size(adapter, task.source_path or "")
|
||||
finally:
|
||||
await adapter.close()
|
||||
# 总字节 = 源大小 × 步骤数(每步传输一次源文件,parallel 与 chain 相同)
|
||||
task.total_size = (task.source_size or 0) * len(steps)
|
||||
task.status = "running"
|
||||
task.started_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
await _broadcast(task, steps)
|
||||
|
||||
completed = 0
|
||||
canceled = False
|
||||
deleted = False
|
||||
for step in steps:
|
||||
# 任务被软删除后中止执行(防竞态:删除请求已标记取消并软删记录)
|
||||
await db.refresh(task)
|
||||
if task.is_deleted:
|
||||
deleted = True
|
||||
break
|
||||
if transfer_task_registry.is_canceled(task_id):
|
||||
canceled = True
|
||||
break
|
||||
if await _run_step(db, task, step):
|
||||
completed += 1
|
||||
else:
|
||||
break
|
||||
|
||||
transfer_task_registry.clear(task_id)
|
||||
if deleted:
|
||||
# 任务已删除:不再修改其状态,仅记录日志
|
||||
logger.info("传输任务 {} 已删除,中止执行", task_id)
|
||||
_remove_local_temp(task)
|
||||
return
|
||||
if canceled:
|
||||
task.status = "canceled"
|
||||
task.error_msg = None
|
||||
for step in steps:
|
||||
if step.status == "pending":
|
||||
step.status = "canceled"
|
||||
step.finished_at = datetime.now(UTC)
|
||||
elif completed == len(steps):
|
||||
task.status = "success"
|
||||
task.progress = 100
|
||||
task.finished_at = datetime.now(UTC)
|
||||
await db.commit()
|
||||
await _broadcast(task, steps)
|
||||
logger.info("传输任务 {} 结束: {}", task_id, task.status)
|
||||
|
||||
# 清理本地源临时文件
|
||||
_remove_local_temp(task)
|
||||
@@ -0,0 +1,59 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import BigInteger, DateTime, Float, ForeignKey, Integer, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.core.base_model import ModelMixin, UserMixin
|
||||
|
||||
|
||||
class StorageTransferTaskModel(ModelMixin, UserMixin):
|
||||
"""文件传输任务模型(多目标 / 链式)"""
|
||||
|
||||
__tablename__: str = "task_workflow_transfer_task"
|
||||
__table_args__: dict[str, str] = {"comment": "文件传输任务表"}
|
||||
|
||||
name: Mapped[str] = mapped_column(String(128), nullable=False, comment="任务名称")
|
||||
task_type: Mapped[str] = mapped_column(String(16), nullable=False, comment="任务类型(parallel:多目标 chain:链式)")
|
||||
source_type: Mapped[str] = mapped_column(String(16), nullable=False, comment="源类型(local:本地 remote:远端)")
|
||||
source_id: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="源存储源ID(本地源为空)")
|
||||
source_path: Mapped[str | None] = mapped_column(String(1024), default=None, nullable=True, comment="源远端路径(本地源为服务端临时文件)")
|
||||
source_name: Mapped[str | None] = mapped_column(String(512), default=None, nullable=True, comment="源文件名")
|
||||
source_size: Mapped[int | None] = mapped_column(BigInteger, default=None, nullable=True, comment="源文件大小(字节)")
|
||||
status: Mapped[str] = mapped_column(String(16), default="pending", nullable=False, index=True, comment="状态(pending/running/success/failed/canceled)")
|
||||
total_size: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False, comment="总字节")
|
||||
transferred_size: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False, comment="已传输字节")
|
||||
progress: Mapped[int] = mapped_column(Integer, default=0, nullable=False, comment="进度(0-100)")
|
||||
speed: Mapped[float] = mapped_column(Float, default=0.0, nullable=False, comment="实时速度(B/s)")
|
||||
error_msg: Mapped[str | None] = mapped_column(Text, default=None, nullable=True, comment="错误信息")
|
||||
started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), default=None, nullable=True, comment="开始时间")
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), default=None, nullable=True, comment="结束时间")
|
||||
|
||||
|
||||
class StorageTransferStepModel(ModelMixin):
|
||||
"""文件传输步骤模型(一个任务展开为多个步骤)"""
|
||||
|
||||
__tablename__: str = "task_workflow_transfer_step"
|
||||
__table_args__: dict[str, str] = {"comment": "文件传输步骤表"}
|
||||
|
||||
task_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("task_workflow_transfer_task.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
index=True,
|
||||
comment="任务ID",
|
||||
)
|
||||
step_order: Mapped[int] = mapped_column(Integer, nullable=False, comment="步骤序号(从0开始)")
|
||||
source_id: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="源存储源ID(首步本地源为空)")
|
||||
source_path: Mapped[str | None] = mapped_column(String(1024), default=None, nullable=True, comment="源路径(本地源为服务端临时文件)")
|
||||
target_id: Mapped[int] = mapped_column(Integer, nullable=False, comment="目标存储源ID")
|
||||
target_path: Mapped[str] = mapped_column(String(1024), nullable=False, comment="目标路径")
|
||||
transfer_mode: Mapped[str | None] = mapped_column(String(16), default=None, nullable=True, comment="传输方式(stream:流式 multipart:分片;空=用存储源默认)")
|
||||
multipart_part_size: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="分片大小(MB,分片传输时覆盖存储源配置)")
|
||||
multipart_concurrency: Mapped[int | None] = mapped_column(Integer, default=None, nullable=True, comment="分片上传并发路数(分片传输时覆盖存储源配置)")
|
||||
status: Mapped[str] = mapped_column(String(16), default="pending", nullable=False, comment="状态(pending/running/success/failed/canceled)")
|
||||
progress: Mapped[int] = mapped_column(Integer, default=0, nullable=False, comment="进度(0-100)")
|
||||
speed: Mapped[float] = mapped_column(Float, default=0.0, nullable=False, comment="实时速度(B/s)")
|
||||
total_size: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False, comment="本步总字节")
|
||||
transferred_size: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False, comment="本步已传输字节")
|
||||
error_msg: Mapped[str | None] = mapped_column(Text, default=None, nullable=True, comment="错误信息")
|
||||
started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), default=None, nullable=True, comment="开始时间")
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), default=None, nullable=True, comment="结束时间")
|
||||
@@ -0,0 +1,20 @@
|
||||
"""传输任务运行时注册表(单实例部署)"""
|
||||
|
||||
|
||||
class TransferTaskRegistry:
|
||||
"""维护任务取消标志等运行时状态(不持久化,重启后任务按状态恢复)"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._cancel_flags: dict[int, bool] = {}
|
||||
|
||||
def mark_cancel(self, task_id: int) -> None:
|
||||
self._cancel_flags[task_id] = True
|
||||
|
||||
def is_canceled(self, task_id: int) -> bool:
|
||||
return self._cancel_flags.get(task_id, False)
|
||||
|
||||
def clear(self, task_id: int) -> None:
|
||||
self._cancel_flags.pop(task_id, None)
|
||||
|
||||
|
||||
transfer_task_registry = TransferTaskRegistry()
|
||||
@@ -0,0 +1,135 @@
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from app.core.base_schema import BaseQueryParam, BaseSchema
|
||||
|
||||
TransferStatus = Literal["pending", "running", "success", "failed", "canceled"]
|
||||
TransferTaskType = Literal["parallel", "chain"]
|
||||
TransferSourceType = Literal["local", "remote"]
|
||||
TransferMode = Literal["stream", "multipart"]
|
||||
|
||||
|
||||
class TransferTargetSchema(BaseModel):
|
||||
"""传输目标配置"""
|
||||
|
||||
target_id: int = Field(..., ge=1, description="目标存储源ID")
|
||||
target_path: str = Field(..., min_length=1, max_length=1024, description="目标路径")
|
||||
|
||||
|
||||
class TransferTaskCreateSchema(BaseModel):
|
||||
"""创建传输任务(远端源,JSON 提交)"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=128, description="任务名称")
|
||||
task_type: TransferTaskType = Field(..., description="任务类型(parallel:多目标 chain:链式)")
|
||||
source_type: TransferSourceType = Field(default="remote", description="源类型(local:本地 remote:远端)")
|
||||
source_id: int | None = Field(default=None, ge=1, description="源存储源ID(remote 必填)")
|
||||
source_path: str | None = Field(default=None, min_length=1, max_length=1024, description="源远端路径(remote 必填)")
|
||||
targets: list[TransferTargetSchema] = Field(..., min_length=1, description="目标列表(顺序即链式执行顺序)")
|
||||
transfer_mode: TransferMode | None = Field(default=None, description="传输方式(stream:流式 multipart:分片;空=用存储源默认)")
|
||||
multipart_part_size: int | None = Field(default=None, ge=5, le=5000, description="分片大小(MB,分片传输时覆盖存储源配置)")
|
||||
multipart_concurrency: int | None = Field(default=None, ge=1, le=64, description="分片上传并发路数(分片传输时覆盖存储源配置)")
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_source(self):
|
||||
if self.source_type == "remote" and not self.source_id:
|
||||
raise ValueError("远端源必须指定源存储源 source_id")
|
||||
if self.source_type == "remote" and not self.source_path:
|
||||
raise ValueError("远端源必须指定源路径 source_path")
|
||||
return self
|
||||
|
||||
|
||||
class LocalUploadInfoSchema(BaseModel):
|
||||
"""本地源上传文件信息(服务端临时文件)"""
|
||||
|
||||
source_path: str = Field(..., description="服务端临时文件路径")
|
||||
source_name: str = Field(..., description="原始文件名")
|
||||
source_size: int = Field(..., ge=0, description="文件大小(字节)")
|
||||
|
||||
|
||||
class TransferTaskStoreSchema(BaseModel):
|
||||
"""传输任务落库模型(create 展开明细后持久化)"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=128, description="任务名称")
|
||||
task_type: TransferTaskType = Field(..., description="任务类型")
|
||||
source_type: TransferSourceType = Field(default="remote", description="源类型")
|
||||
source_id: int | None = Field(default=None, ge=1, description="源存储源ID")
|
||||
source_path: str | None = Field(default=None, max_length=1024, description="源路径")
|
||||
source_name: str | None = Field(default=None, max_length=255, description="源文件名")
|
||||
source_size: int | None = Field(default=None, ge=0, description="源文件大小")
|
||||
status: TransferStatus = Field(default="pending", description="状态")
|
||||
|
||||
|
||||
class TransferStepCreateSchema(BaseModel):
|
||||
"""传输步骤落库模型(由目标列表展开)"""
|
||||
|
||||
step_order: int = Field(..., ge=0, description="步骤序号")
|
||||
source_id: int | None = Field(default=None, ge=1, description="源存储源ID")
|
||||
source_path: str | None = Field(default=None, max_length=1024, description="源路径")
|
||||
target_id: int = Field(..., ge=1, description="目标存储源ID")
|
||||
target_path: str = Field(..., min_length=1, max_length=1024, description="目标路径")
|
||||
transfer_mode: TransferMode | None = Field(default=None, description="传输方式")
|
||||
multipart_part_size: int | None = Field(default=None, ge=1, description="分片大小(MB)")
|
||||
multipart_concurrency: int | None = Field(default=None, ge=1, description="分片并发数")
|
||||
|
||||
|
||||
class TransferStepOutSchema(BaseSchema):
|
||||
"""传输步骤详情"""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
task_id: int = Field(description="任务ID")
|
||||
step_order: int = Field(description="步骤序号")
|
||||
source_id: int | None = Field(default=None, description="源存储源ID")
|
||||
source_path: str | None = Field(default=None, description="源路径")
|
||||
target_id: int = Field(description="目标存储源ID")
|
||||
target_path: str = Field(description="目标路径")
|
||||
transfer_mode: TransferMode | None = Field(default=None, description="传输方式(空=用存储源默认)")
|
||||
multipart_part_size: int | None = Field(default=None, description="分片大小(MB)")
|
||||
multipart_concurrency: int | None = Field(default=None, description="分片上传并发路数")
|
||||
status: TransferStatus = Field(description="状态")
|
||||
progress: int = Field(default=0, description="进度(0-100)")
|
||||
speed: float = Field(default=0.0, description="速度(B/s)")
|
||||
total_size: int = Field(default=0, description="本步总字节")
|
||||
transferred_size: int = Field(default=0, description="本步已传输字节")
|
||||
error_msg: str | None = Field(default=None, description="错误信息")
|
||||
started_at: datetime | None = Field(default=None, description="开始时间")
|
||||
finished_at: datetime | None = Field(default=None, description="结束时间")
|
||||
|
||||
|
||||
class TransferTaskCreateResultSchema(BaseModel):
|
||||
"""创建传输任务结果"""
|
||||
|
||||
id: int = Field(..., ge=1, description="任务ID")
|
||||
|
||||
|
||||
class TransferTaskOutSchema(BaseSchema):
|
||||
"""传输任务详情"""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
name: str = Field(description="任务名称")
|
||||
task_type: TransferTaskType = Field(description="任务类型")
|
||||
source_type: TransferSourceType = Field(description="源类型")
|
||||
source_id: int | None = Field(default=None, description="源存储源ID")
|
||||
source_path: str | None = Field(default=None, description="源远端路径")
|
||||
source_name: str | None = Field(default=None, description="源文件名")
|
||||
source_size: int | None = Field(default=None, description="源文件大小")
|
||||
status: TransferStatus = Field(description="状态")
|
||||
total_size: int = Field(default=0, description="总字节")
|
||||
transferred_size: int = Field(default=0, description="已传输字节")
|
||||
progress: int = Field(default=0, description="进度(0-100)")
|
||||
speed: float = Field(default=0.0, description="实时速度(B/s)")
|
||||
error_msg: str | None = Field(default=None, description="错误信息")
|
||||
started_at: datetime | None = Field(default=None, description="开始时间")
|
||||
finished_at: datetime | None = Field(default=None, description="结束时间")
|
||||
steps: list[TransferStepOutSchema] = Field(default_factory=list, description="传输步骤")
|
||||
|
||||
|
||||
class TransferTaskQueryParam(BaseQueryParam):
|
||||
"""传输任务查询参数"""
|
||||
|
||||
name: str | None = Field(None, description="任务名称", json_schema_extra={"q": "like"})
|
||||
task_type: TransferTaskType | None = Field(None, description="任务类型", json_schema_extra={"q": "eq"})
|
||||
status: TransferStatus | None = Field(None, description="状态", json_schema_extra={"q": "eq"})
|
||||
@@ -0,0 +1,255 @@
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import aiofiles
|
||||
from fastapi import UploadFile
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.base_schema import AuthSchema, PageResultSchema
|
||||
from app.core.exceptions import CustomException
|
||||
from app.modules.workflow.source.service import StorageSourceService
|
||||
from app.modules.workflow.transfer.engine import _broadcast, execute_transfer_task
|
||||
from app.modules.workflow.transfer.registry import transfer_task_registry
|
||||
from app.utils.common_util import search_to_dict
|
||||
|
||||
from .crud import StorageTransferTaskCRUD
|
||||
from .model import StorageTransferStepModel, StorageTransferTaskModel
|
||||
from .schema import (
|
||||
LocalUploadInfoSchema,
|
||||
TransferStepCreateSchema,
|
||||
TransferStepOutSchema,
|
||||
TransferTargetSchema,
|
||||
TransferTaskCreateSchema,
|
||||
TransferTaskOutSchema,
|
||||
TransferTaskQueryParam,
|
||||
TransferTaskStoreSchema,
|
||||
)
|
||||
|
||||
|
||||
class StorageTransferService:
|
||||
"""文件传输任务服务(创建 / 查询 / 取消 / 删除)"""
|
||||
|
||||
# 后台传输任务强引用集合:asyncio 对未保存引用的 Task 可能随时 GC(官方文档警告),
|
||||
# 完成后通过 done callback 移除,避免泄漏
|
||||
_BG_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
def __init__(self, auth: AuthSchema, db: AsyncSession) -> None:
|
||||
self.auth = auth
|
||||
self.db = db
|
||||
|
||||
# ── 内部工具 ────────────────────────────────────────────────────
|
||||
|
||||
def _crud(self) -> StorageTransferTaskCRUD:
|
||||
return StorageTransferTaskCRUD(self.auth, self.db)
|
||||
|
||||
async def _validate_targets(self, targets: list[TransferTargetSchema]) -> None:
|
||||
"""校验目标存储源均存在且启用(批量一次查询)。"""
|
||||
source_service = StorageSourceService(self.auth, self.db)
|
||||
await source_service.get_active_sources([t.target_id for t in targets])
|
||||
|
||||
@staticmethod
|
||||
def _build_steps(data: TransferTaskCreateSchema, local_source_path: str | None = None) -> list[TransferStepCreateSchema]:
|
||||
"""展开步骤:chain 下每步源继承上一步的目标;parallel 下每步源均为任务源。
|
||||
|
||||
local 源时任务源为服务端临时文件,需显式传入 local_source_path。
|
||||
"""
|
||||
steps: list[TransferStepCreateSchema] = []
|
||||
prev_id, prev_path = data.source_id, local_source_path or data.source_path
|
||||
for order, target in enumerate(data.targets):
|
||||
if data.task_type == "chain":
|
||||
source_id, source_path = prev_id, prev_path
|
||||
else:
|
||||
source_id, source_path = data.source_id, local_source_path or data.source_path
|
||||
steps.append(
|
||||
TransferStepCreateSchema(
|
||||
step_order=order,
|
||||
source_id=source_id,
|
||||
source_path=source_path,
|
||||
target_id=target.target_id,
|
||||
target_path=target.target_path,
|
||||
transfer_mode=data.transfer_mode,
|
||||
multipart_part_size=data.multipart_part_size,
|
||||
multipart_concurrency=data.multipart_concurrency,
|
||||
)
|
||||
)
|
||||
if data.task_type == "chain":
|
||||
prev_id, prev_path = target.target_id, target.target_path
|
||||
return steps
|
||||
|
||||
async def _persist(self, data: TransferTaskCreateSchema, local_info: LocalUploadInfoSchema | None = None) -> int:
|
||||
"""落库任务与步骤(pending),随后启动后台执行。"""
|
||||
task = await self._crud().create(
|
||||
TransferTaskStoreSchema(
|
||||
name=data.name,
|
||||
task_type=data.task_type,
|
||||
source_type=data.source_type,
|
||||
source_id=data.source_id,
|
||||
source_path=local_info.source_path if local_info else data.source_path,
|
||||
source_name=local_info.source_name if local_info else ((data.source_path or "").rsplit("/", 1)[-1] or None),
|
||||
source_size=local_info.source_size if local_info else None,
|
||||
).model_dump()
|
||||
)
|
||||
local_source_path = local_info.source_path if local_info else None
|
||||
for step_data in self._build_steps(data, local_source_path=local_source_path):
|
||||
self.db.add(StorageTransferStepModel(task_id=task.id, **step_data.model_dump()))
|
||||
# 事务边界在 HTTP 层(db_getter 的 session.begin()),此处只 flush 不 commit
|
||||
await self.db.flush()
|
||||
self._launch_after_commit(task.id)
|
||||
return task.id
|
||||
|
||||
def _launch_after_commit(self, task_id: int) -> None:
|
||||
"""请求事务提交后再启动后台传输。
|
||||
|
||||
后台任务使用独立会话,若提前启动会读不到未提交的任务行,
|
||||
execute_transfer_task 将静默返回,任务永久停留在 pending。
|
||||
事务回滚时 after_commit 不触发,任务既未落库也不会启动。
|
||||
"""
|
||||
|
||||
@event.listens_for(self.db.sync_session, "after_commit", once=True)
|
||||
def _launch_on_commit(_session) -> None:
|
||||
bg_task = asyncio.create_task(execute_transfer_task(task_id))
|
||||
StorageTransferService._BG_TASKS.add(bg_task)
|
||||
bg_task.add_done_callback(StorageTransferService._BG_TASKS.discard)
|
||||
|
||||
# ── 创建 ────────────────────────────────────────────────────────
|
||||
|
||||
async def create(self, data: TransferTaskCreateSchema) -> int:
|
||||
"""创建远端源传输任务。"""
|
||||
source_service = StorageSourceService(self.auth, self.db)
|
||||
if data.source_type == "remote":
|
||||
await source_service.get_active_source(data.source_id)
|
||||
await self._validate_targets(data.targets)
|
||||
return await self._persist(data)
|
||||
|
||||
async def create_local(self, data: TransferTaskCreateSchema, file: UploadFile) -> int:
|
||||
"""创建本地源传输任务:文件保存到服务端临时目录,执行完毕后自动清理。"""
|
||||
if not file or not file.filename:
|
||||
raise CustomException(msg="请选择要上传的文件")
|
||||
await self._validate_targets(data.targets)
|
||||
fd, temp_path = tempfile.mkstemp(prefix="transfer_upload_", suffix=os.path.splitext(file.filename)[1])
|
||||
os.close(fd)
|
||||
try:
|
||||
async with aiofiles.open(temp_path, "wb") as f:
|
||||
while chunk := await file.read(1024 * 1024):
|
||||
await f.write(chunk)
|
||||
except Exception:
|
||||
os.unlink(temp_path)
|
||||
raise
|
||||
finally:
|
||||
await file.seek(0)
|
||||
return await self._persist(
|
||||
data,
|
||||
local_info=LocalUploadInfoSchema(
|
||||
source_path=temp_path,
|
||||
source_name=file.filename,
|
||||
source_size=os.path.getsize(temp_path),
|
||||
),
|
||||
)
|
||||
|
||||
# ── 查询 ────────────────────────────────────────────────────────
|
||||
|
||||
async def page(
|
||||
self,
|
||||
search: TransferTaskQueryParam | None,
|
||||
page_no: int,
|
||||
page_size: int,
|
||||
order_by: list[dict] | None = None,
|
||||
) -> PageResultSchema[TransferTaskOutSchema]:
|
||||
result = await self._crud().page(
|
||||
offset=(page_no - 1) * page_size,
|
||||
limit=page_size,
|
||||
order_by=order_by or [{"id": "desc"}],
|
||||
search=search_to_dict(search),
|
||||
)
|
||||
items = [TransferTaskOutSchema.model_validate(obj) for obj in result.items]
|
||||
# 批量加载当前页任务的步骤(前端列表依赖 steps 展示目标/信息列)
|
||||
if items:
|
||||
task_ids = [item.id for item in items]
|
||||
step_result = await self.db.execute(
|
||||
select(StorageTransferStepModel)
|
||||
.where(
|
||||
StorageTransferStepModel.task_id.in_(task_ids),
|
||||
StorageTransferStepModel.is_deleted.is_(False),
|
||||
)
|
||||
.order_by(StorageTransferStepModel.step_order)
|
||||
)
|
||||
steps_map: dict[int, list[TransferStepOutSchema]] = {}
|
||||
for step in step_result.scalars().all():
|
||||
steps_map.setdefault(step.task_id, []).append(TransferStepOutSchema.model_validate(step))
|
||||
for item in items:
|
||||
if item.id is not None:
|
||||
item.steps = steps_map.get(item.id, [])
|
||||
return PageResultSchema[TransferTaskOutSchema](
|
||||
page_no=result.page_no,
|
||||
page_size=result.page_size,
|
||||
total=result.total,
|
||||
has_next=result.has_next,
|
||||
items=items,
|
||||
)
|
||||
|
||||
async def detail(self, task_id: int) -> TransferTaskOutSchema:
|
||||
task = await self._crud().get_or_404(id=task_id)
|
||||
out = TransferTaskOutSchema.model_validate(task)
|
||||
result = await self.db.execute(
|
||||
select(StorageTransferStepModel)
|
||||
.where(
|
||||
StorageTransferStepModel.task_id == task_id,
|
||||
StorageTransferStepModel.is_deleted.is_(False),
|
||||
)
|
||||
.order_by(StorageTransferStepModel.step_order)
|
||||
)
|
||||
out.steps = [TransferStepOutSchema.model_validate(step) for step in result.scalars().all()]
|
||||
return out
|
||||
|
||||
# ── 操作 ────────────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _remove_local_temp(task: StorageTransferTaskModel) -> None:
|
||||
"""清理本地源任务的临时文件(幂等:文件不存在时静默忽略)。"""
|
||||
if task.source_type == "local" and task.source_path:
|
||||
try:
|
||||
os.unlink(task.source_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
async def _push_task(self, task: StorageTransferTaskModel) -> None:
|
||||
"""将任务最新状态推送到其创建者的 WebSocket(复用引擎广播逻辑)。"""
|
||||
result = await self.db.execute(
|
||||
select(StorageTransferStepModel)
|
||||
.where(
|
||||
StorageTransferStepModel.task_id == task.id,
|
||||
StorageTransferStepModel.is_deleted.is_(False),
|
||||
)
|
||||
.order_by(StorageTransferStepModel.step_order)
|
||||
)
|
||||
await _broadcast(task, list(result.scalars().all()))
|
||||
|
||||
async def cancel(self, task_id: int) -> None:
|
||||
task = await self._crud().get_or_404(id=task_id)
|
||||
if task.status == "pending":
|
||||
# 走 CRUDBase.update:自动补 updated_id 审计字段(flush 由 HTTP 层统一提交)
|
||||
await self._crud().update(id=task_id, data={"status": "canceled", "finished_at": datetime.now(UTC)})
|
||||
# pending 任务未启动引擎,需在此清理本地源临时文件并即时推送状态
|
||||
self._remove_local_temp(task)
|
||||
await self._push_task(task)
|
||||
elif task.status == "running":
|
||||
transfer_task_registry.mark_cancel(task_id)
|
||||
|
||||
async def delete(self, ids: list[int]) -> None:
|
||||
for task_id in ids:
|
||||
transfer_task_registry.mark_cancel(task_id)
|
||||
# 清理本地源任务的临时文件(pending 任务引擎不会执行,需兜底清理)
|
||||
result = await self.db.execute(
|
||||
select(StorageTransferTaskModel)
|
||||
.where(
|
||||
StorageTransferTaskModel.id.in_(ids),
|
||||
StorageTransferTaskModel.source_type == "local",
|
||||
StorageTransferTaskModel.is_deleted.is_(False),
|
||||
)
|
||||
)
|
||||
for task in result.scalars().all():
|
||||
self._remove_local_temp(task)
|
||||
await self._crud().delete(ids=ids)
|
||||
@@ -0,0 +1,5 @@
|
||||
"""传输任务 SSE 连接管理(按用户推送任务进度)"""
|
||||
|
||||
from app.core.sse_manager import SSEConnectionManager
|
||||
|
||||
transfer_stream_manager = SSEConnectionManager(channel="transfer")
|
||||
Reference in New Issue
Block a user