refactor(task): 重构定时任务与存储路由

- 移除各路由文件顶部冗余注释
- 将 JobRouter/NodeRouter 重命名为 CornJobRouter/CornJobNodeRouter
- 新增存储浏览、节点、传输、工作流路由注册
- 调整路由导入来源与任务表名
- 优化 main.py 启动方式及环境配置加载
This commit is contained in:
zhangtao
2026-09-06 04:14:17 +08:00
parent d470c7eb1f
commit 9e80f970d0
78 changed files with 4383 additions and 4383 deletions
@@ -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.task.storage.transfer.schema import (
TransferTaskCreateResultSchema,
TransferTaskCreateSchema,
TransferTaskOutSchema,
TransferTaskQueryParam,
TransferTaskType,
)
from app.modules.task.storage.transfer.service import StorageTransferService
from app.modules.task.storage.transfer.sse_manager import transfer_stream_manager
StorageTransferRouter = APIRouter(route_class=OperationLogRoute, prefix="/storage/transfer", tags=["文件传输"])
@StorageTransferRouter.post("/task", summary="创建传输任务(远端源)", response_model=ResponseSchema[TransferTaskCreateResultSchema])
async def create_transfer_task_controller(
auth: Annotated[AuthSchema, Security(AuthPermission(["module_storage: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_storage: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_storage: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_storage: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_storage: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_storage: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.task.storage.core.base import StorageAdapterConfig
from app.modules.task.storage.core.factory import StorageAdapterFactory
from app.modules.task.storage.node.model import StorageNodeModel
from app.modules.task.storage.node.service import StorageNodeService
from app.modules.task.storage.transfer.registry import transfer_task_registry
from app.modules.task.storage.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(StorageNodeModel, source_id)
if source is None or source.status == 1:
return None
return StorageNodeService._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_storage_transfer"
__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_storage_transfer_step"
__table_args__: dict[str, str] = {"comment": "文件传输步骤表"}
task_id: Mapped[int] = mapped_column(
ForeignKey("task_storage_transfer.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.task.storage.node.service import StorageNodeService
from app.modules.task.storage.transfer.engine import _broadcast, execute_transfer_task
from app.modules.task.storage.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 = StorageNodeService(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 = StorageNodeService(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")