mirror of
https://github.com/fastapiadmin/FastapiAdmin.git
synced 2026-09-20 20:39:55 +00:00
refactor(task): 重构定时任务与存储路由
- 移除各路由文件顶部冗余注释 - 将 JobRouter/NodeRouter 重命名为 CornJobRouter/CornJobNodeRouter - 新增存储浏览、节点、传输、工作流路由注册 - 调整路由导入来源与任务表名 - 优化 main.py 启动方式及环境配置加载
This commit is contained in:
@@ -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")
|
||||
Reference in New Issue
Block a user