232 lines
8.8 KiB
Python
Executable File
232 lines
8.8 KiB
Python
Executable File
"""用户端应用路由。
|
||
|
||
面向普通用户暴露“应用中心”能力:列出已发布工作流、上传输入创建任务、
|
||
查询进度、重试失败任务以及下载产物。用户只看到输入 -> 进度 -> 结果。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import shutil
|
||
import uuid
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
|
||
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile
|
||
from fastapi.responses import FileResponse
|
||
|
||
from wov_app.db import Database
|
||
|
||
router = APIRouter(tags=["apps"])
|
||
|
||
|
||
def _now_iso() -> str:
|
||
"""返回当前 UTC 时间的 ISO 格式字符串。"""
|
||
return datetime.now(timezone.utc).isoformat()
|
||
|
||
|
||
def _get_db() -> Database:
|
||
"""从 FastAPI 应用状态中延迟获取数据库实例。"""
|
||
from wov_app.main import app
|
||
|
||
return app.state.db
|
||
|
||
|
||
@router.get("/api/apps")
|
||
def list_apps(db: Database = Depends(_get_db)) -> list[dict]:
|
||
"""返回全部已发布工作流及其最新版本定义。"""
|
||
apps = []
|
||
for workflow in db.list_workflows():
|
||
# 草稿工作流不对用户端可见。
|
||
if not workflow["published"]:
|
||
continue
|
||
latest = db.get_latest_workflow_version(workflow["id"])
|
||
apps.append(
|
||
{
|
||
"id": workflow["id"],
|
||
"name": workflow["name"],
|
||
"description": workflow["description"],
|
||
"version": workflow["latest_version"],
|
||
"definition": latest["definition"] if latest else None,
|
||
}
|
||
)
|
||
return apps
|
||
|
||
|
||
@router.post("/api/apps/{workflow_id}/runs")
|
||
async def create_run(
|
||
workflow_id: str,
|
||
file: UploadFile = File(...),
|
||
params: str = Form(default=""),
|
||
db: Database = Depends(_get_db),
|
||
) -> dict:
|
||
"""接收用户上传文件,创建排队中的工作流任务。"""
|
||
workflow = db.get_workflow(workflow_id)
|
||
# 只允许对已发布且存在版本的工作流发起任务。
|
||
if workflow is None or not workflow["published"]:
|
||
raise HTTPException(status_code=404, detail="published workflow not found")
|
||
|
||
latest = db.get_latest_workflow_version(workflow_id)
|
||
if latest is None:
|
||
raise HTTPException(status_code=422, detail="workflow has no version")
|
||
|
||
run_id = f"run_{uuid.uuid4().hex[:12]}"
|
||
# 使用安全文件名,避免路径穿越。
|
||
filename = Path(file.filename or "upload.bin").name
|
||
from wov_app.config import STORAGE_DIR
|
||
|
||
# 上传文件按 run 隔离存放,调度器通过 input_uri 引用。
|
||
input_dir = STORAGE_DIR / "uploads" / run_id
|
||
input_dir.mkdir(parents=True, exist_ok=True)
|
||
input_uri = input_dir / filename
|
||
content = await file.read()
|
||
input_uri.write_bytes(content)
|
||
|
||
# 可选参数覆盖(如前端框选的 crop):{节点ID: {参数: 值}},随任务持久化。
|
||
param_overrides = None
|
||
if params.strip():
|
||
try:
|
||
parsed = json.loads(params)
|
||
except json.JSONDecodeError as exc:
|
||
raise HTTPException(status_code=422, detail="params must be valid JSON") from exc
|
||
if not isinstance(parsed, dict):
|
||
raise HTTPException(status_code=422, detail="params must be a JSON object")
|
||
param_overrides = parsed
|
||
|
||
now = _now_iso()
|
||
db.create_run(
|
||
{
|
||
"id": run_id,
|
||
"workflow_id": workflow_id,
|
||
"workflow_version": latest["version"],
|
||
"status": "QUEUED",
|
||
"progress": 0,
|
||
"input_uri": str(input_uri),
|
||
"param_overrides": param_overrides,
|
||
"created_at": now,
|
||
"updated_at": now,
|
||
}
|
||
)
|
||
return {
|
||
"id": run_id,
|
||
"status": "QUEUED",
|
||
"progress": 0,
|
||
"artifacts": [],
|
||
}
|
||
|
||
|
||
@router.get("/api/runs")
|
||
def list_runs(db: Database = Depends(_get_db)) -> list[dict]:
|
||
"""返回最近的运行记录,供任务管理页展示。"""
|
||
return db.list_runs()
|
||
|
||
|
||
@router.get("/api/runs/{run_id}")
|
||
def get_run(run_id: str, db: Database = Depends(_get_db)) -> dict:
|
||
"""返回任务详情,并附带当前产物列表。"""
|
||
run = db.get_run(run_id)
|
||
if run is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
run["artifacts"] = db.list_artifacts(run_id)
|
||
return run
|
||
|
||
|
||
@router.post("/api/runs/{run_id}/retry")
|
||
def retry_run(run_id: str, db: Database = Depends(_get_db)) -> dict:
|
||
"""重置失败任务为排队状态,清空旧产物后重新执行。"""
|
||
run = db.get_run(run_id)
|
||
if run is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
if run["status"] != "FAILED":
|
||
raise HTTPException(status_code=422, detail="only failed runs can be retried")
|
||
# reset_run 会清空进度、错误和旧产物,确保从头开始。
|
||
db.reset_run(run_id, _now_iso())
|
||
# 重试前清除可能残留的暂停信号(任务失败时信号文件可能仍在)。
|
||
from wov_app.config import STORAGE_DIR
|
||
|
||
(STORAGE_DIR / "runs" / run_id / "paused.flag").unlink(missing_ok=True)
|
||
return {"id": run_id, "status": "QUEUED"}
|
||
|
||
|
||
@router.post("/api/runs/{run_id}/pause")
|
||
def pause_run(run_id: str, db: Database = Depends(_get_db)) -> dict:
|
||
"""暂停任务:排队中或运行中的任务可暂停,运行中的任务在节点边界停下。"""
|
||
run = db.get_run(run_id)
|
||
if run is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
if run["status"] not in ("QUEUED", "RUNNING"):
|
||
raise HTTPException(status_code=422, detail="only queued or running runs can be paused")
|
||
db.pause_run(run_id, _now_iso())
|
||
# 写入暂停信号文件:运行中的节点(如 OCR)逐帧检查到后立即中止,
|
||
# 由调度器保持 PAUSED;resume 时清除。
|
||
from wov_app.config import STORAGE_DIR
|
||
|
||
run_dir = STORAGE_DIR / "runs" / run_id
|
||
run_dir.mkdir(parents=True, exist_ok=True)
|
||
(run_dir / "paused.flag").write_text("", encoding="utf-8")
|
||
return {"id": run_id, "status": "PAUSED"}
|
||
|
||
|
||
@router.post("/api/runs/{run_id}/resume")
|
||
def resume_run(run_id: str, db: Database = Depends(_get_db)) -> dict:
|
||
"""继续任务:暂停的任务恢复排队,由调度器从断点继续执行。"""
|
||
run = db.get_run(run_id)
|
||
if run is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
if run["status"] != "PAUSED":
|
||
raise HTTPException(status_code=422, detail="only paused runs can be resumed")
|
||
db.resume_run(run_id, _now_iso())
|
||
# 清除暂停信号文件,避免节点误判仍处于暂停状态。
|
||
from wov_app.config import STORAGE_DIR
|
||
|
||
(STORAGE_DIR / "runs" / run_id / "paused.flag").unlink(missing_ok=True)
|
||
return {"id": run_id, "status": "QUEUED"}
|
||
|
||
@router.delete("/api/runs/{run_id}")
|
||
def delete_run(run_id: str, db: Database = Depends(_get_db)) -> dict:
|
||
"""删除上传任务及其私有文件;批量 run 必须通过批量任务入口删除。"""
|
||
run = db.get_run(run_id)
|
||
if run is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
# 批量 run 的输入是用户原始视频,且由 batch_videos 关联管理。
|
||
# 在任何数据库/文件删除之前拒绝,避免误删媒体库或留下悬空的批量明细。
|
||
if run.get("source") == "batch":
|
||
raise HTTPException(status_code=422, detail="请通过批量任务入口删除该任务")
|
||
from wov_app.config import STORAGE_DIR
|
||
|
||
# 先删数据库记录(含产物表),再清理磁盘上的上传与中间产物。
|
||
db.delete_run(run_id)
|
||
# 删除范围仅来自该任务的私有存储布局,绝不由 input_uri 推导:
|
||
# 即使历史上传记录引用外部路径,也必须保留源视频和同目录的用户文件。
|
||
shutil.rmtree(STORAGE_DIR / "uploads" / run_id, ignore_errors=True)
|
||
# 步骤产物位于 <storage>/runs/<run_id>/,整目录一并删除。
|
||
shutil.rmtree(STORAGE_DIR / "runs" / run_id, ignore_errors=True)
|
||
return {"deleted": run_id}
|
||
|
||
@router.get("/api/runs/{run_id}/artifacts")
|
||
def list_run_artifacts(run_id: str, db: Database = Depends(_get_db)) -> list[dict]:
|
||
"""返回任务全部产物记录。"""
|
||
if db.get_run(run_id) is None:
|
||
raise HTTPException(status_code=404, detail="run not found")
|
||
return db.list_artifacts(run_id)
|
||
|
||
|
||
@router.get("/api/runs/{run_id}/artifacts/{artifact_name}")
|
||
def download_artifact(
|
||
run_id: str,
|
||
artifact_name: str,
|
||
db: Database = Depends(_get_db),
|
||
) -> FileResponse:
|
||
"""按任务与产物名下载文件,文件缺失时返回 404。"""
|
||
artifact = db.get_artifact(run_id, artifact_name)
|
||
if artifact is None:
|
||
raise HTTPException(status_code=404, detail="artifact not found")
|
||
path = Path(artifact["uri"])
|
||
if not path.is_file():
|
||
raise HTTPException(status_code=404, detail="artifact file missing")
|
||
return FileResponse(
|
||
path,
|
||
media_type=artifact["mime_type"],
|
||
filename=path.name,
|
||
)
|