Compare commits

..

2 Commits

Author SHA1 Message Date
xiaoxia 4cf2981a25 chore: retrigger CI
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m0s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m23s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m28s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m43s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 2m25s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 4m20s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 4m54s
AI Code Review / AI Code Review (pull_request) Successful in 7m0s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m38s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 12m43s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 43s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 55s
2026-09-25 23:19:02 +08:00
xiaoxia a2ea6cc51f fix(cover): 批量生成场景补封面模板选择;批量生成应用所选模板
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 11s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 57s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m37s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m51s
AI Code Review / AI Code Review (pull_request) Successful in 6m36s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m40s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Failing after 14m40s
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 1m6s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 3m46s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 4m1s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
问题:批量生成(count>1)时,Step6CoverSettings 走 isBatch 分支直接 return 了卡片列表,
没有渲染封面模板选择入口;useBatchCovers.generateOne 硬编码 generateCover('default', ...),
即使用户想选自定义模板也不生效。

修改:
1. Step6CoverSettings 批量分支:在『一键全部自动生成』按钮旁增加『⚙️ 封面模板』和『➕ 新建模板』按钮,
   并在批量分支底部渲染 CoverSettingsModal/CoverEditorModal,与单视频分支一致
2. 批量分支的提示条显示当前选中的模板名称,与单视频一致
3. 批量分支补 Modal loading 覆盖层(应用模板生成时的提示)
4. useBatchCovers:去掉 selectedTemplate 参数名的下划线前缀(之前故意不使用),
   generateOne 里 generateCover 第一个参数改用 selectedTemplate(默认 'default'),
   用户选了非默认模板时一键生成/单个生成都会应用该模板
5. generateOne 的依赖数组补 selectedTemplate
6. Step6CoverSettings 增加 onTemplateChange 回调 prop(用于父组件如需监听模板变化),
   同时 shared 的 canGenerate 在批量场景下设为 false,避免误触单视频生成逻辑
7. 批量上传 change 事件从直接拼接 blob 改为调用 batchCovers.uploadOne,走素材库上传流程
   (之前直接把 blob URL 塞 previewCovers,刷新后丢失且不入库)
2026-09-25 22:52:36 +08:00
36 changed files with 583 additions and 2926 deletions
-5
View File
@@ -16,7 +16,6 @@ from app.api.routes.generation_preview import router as generation_preview_route
from app.api.routes.generation_tasks import router as generation_tasks_router
from app.api.routes.generation_variant_plans import router as generation_variant_plans_router
from app.api.routes.gpu_lipsync import router as gpu_lipsync_router
from app.api.routes.gpu_relay import router as gpu_relay_router
from app.api.routes.health import router as health_check_router
from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router
@@ -206,10 +205,6 @@ api_router.include_router(
internal_render_router,
tags=["Internal"],
)
api_router.include_router(
gpu_relay_router,
tags=["GpuRelay"],
)
api_router.include_router(
scripts_router,
prefix="/scripts",
-173
View File
@@ -1,173 +0,0 @@
"""GPU 编码回传 relay 端点。
P4000 编码完成后通过 HTTP PUT 把结果 mp4 写到这里;Worker 在发起 GPU 请求时携带
带签名(token + 随机 key)的 URL,等待 P4000 写入后用同 URL 把文件 GET 回本地。
安全:
- 生产环境必须配置 GPU_ENCODE_RELAY_SECRET;token=xxx 查询参数必须匹配。
- key 为随机 hex,无法被枚举。
- 写入/读取后 worker 会调用 DELETE 主动清理;文件落地在 generated-files/gpu_relay/,
跟 generated-files 同卷,nginx 已对 generated-files 做静态挂载,但 gpu_relay/ 子目录
通过本接口走鉴权,不直接暴露为静态目录(文件名随机 + token 保护双重保险)。
"""
from __future__ import annotations
import logging
import os
import secrets
import time
import uuid
from pathlib import Path
from typing import Optional
from fastapi import APIRouter, HTTPException, Query, Request
from fastapi.responses import FileResponse, Response
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/gpu-relay", tags=["Internal-GpuRelay"])
_DEFAULT_SECRET_LOGGED = False
def _relay_dir() -> Path:
base = os.getenv("GENERATED_FILES_DIR", "/app/generated")
sub = os.getenv("GPU_ENCODE_RELAY_DIR", "gpu_relay")
p = Path(base) / sub
p.mkdir(parents=True, exist_ok=True)
return p
def _secret() -> str:
global _DEFAULT_SECRET_LOGGED
secret = (os.getenv("GPU_ENCODE_RELAY_SECRET", "") or "").strip()
if not secret:
env = (os.getenv("APP_ENV", os.getenv("ENV", "development"))).lower()
if env in ("production", "prod"):
# Production: raise so deployment fails fast
raise RuntimeError("GPU_ENCODE_RELAY_SECRET must be set in production")
# Dev: ephemeral random secret, log once
secret = os.environ.setdefault("GPU_ENCODE_RELAY_SECRET", secrets.token_urlsafe(32))
if not _DEFAULT_SECRET_LOGGED:
logger.warning(
"[gpu-relay] GPU_ENCODE_RELAY_SECRET not set; using ephemeral dev token (%s...)",
secret[:8],
)
_DEFAULT_SECRET_LOGGED = True
return secret
def _safe_key(key: str) -> str:
"""只允许合法文件名字符,防 path traversal。"""
k = key.strip()
if not k or "/" in k or "\\" in k or k in (".", "..") or not all(
c.isalnum() or c in "-_" for c in k
):
raise HTTPException(status_code=400, detail="invalid key")
return k
def _check_token(tok: Optional[str]) -> None:
if not tok or tok != _secret():
raise HTTPException(status_code=401, detail="unauthorized")
# ── Worker 侧:生成一个一次性 PUT URL ───────────────────────────────────
def build_relay_put_url(base_url: str, key: str, secret: str) -> str:
"""给 P4000 用的 PUT URL(含 token)。"""
return f"{base_url.rstrip('/')}/api/v1/internal/gpu-relay/{key}?token={secret}"
def build_relay_get_url(base_url: str, key: str, secret: str) -> str:
"""Worker 取回结果用的 GET URL。"""
return build_relay_put_url(base_url, key, secret)
def generate_key() -> str:
return uuid.uuid4().hex
# ── HTTP endpoints ──────────────────────────────────────────────────────
@router.put("/{key}")
async def put_object(
key: str,
request: Request,
token: Optional[str] = Query(None),
):
_check_token(token)
safe = _safe_key(key)
dst = _relay_dir() / safe
tmp = dst.with_suffix(dst.suffix + ".part")
size = 0
t0 = time.time()
try:
with open(tmp, "wb") as f:
async for chunk in request.stream():
f.write(chunk)
size += len(chunk)
os.replace(tmp, dst)
except Exception as e: # noqa: BLE001
if tmp.exists():
try:
tmp.unlink()
except OSError:
pass
logger.exception("[gpu-relay] PUT failed key=%s", safe)
raise HTTPException(status_code=500, detail=f"write failed: {e}") from e
logger.info(
"[gpu-relay] PUT key=%s size=%d took=%.2fs",
safe, size, time.time() - t0,
)
return {"ok": True, "key": safe, "size": size}
@router.get("/{key}")
async def get_object(
key: str,
token: Optional[str] = Query(None),
):
_check_token(token)
safe = _safe_key(key)
path = _relay_dir() / safe
if not path.exists():
raise HTTPException(status_code=404, detail="not found")
return FileResponse(
path=path,
media_type="video/mp4",
filename=f"{safe}.mp4",
)
@router.head("/{key}")
async def head_object(
key: str,
token: Optional[str] = Query(None),
):
_check_token(token)
safe = _safe_key(key)
path = _relay_dir() / safe
if not path.exists():
return Response(status_code=404)
return Response(
status_code=200,
media_type="video/mp4",
headers={"Content-Length": str(path.stat().st_size)},
)
@router.delete("/{key}")
async def delete_object(
key: str,
token: Optional[str] = Query(None),
):
_check_token(token)
safe = _safe_key(key)
path = _relay_dir() / safe
try:
if path.exists():
path.unlink()
except OSError as e:
raise HTTPException(status_code=500, detail=f"delete failed: {e}") from e
return {"ok": True, "key": safe}
+20 -16
View File
@@ -160,14 +160,12 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
)
await page.goto("/app/generate")
// ── 页面标题 ─────────────────────────────────────────────────
// GenerateHeader: <h2><ThunderboltOutlined />智能剪辑</h2>
// SVG icon 可能干扰 role=heading 的 accessible name,用文本包含兜底
await expect(page.getByText("智能剪辑").first()).toBeVisible({ timeout: 30000 })
await expect(page.getByRole("heading", { name: "智能剪辑" })).toBeVisible({
timeout: 30000,
})
// ── Step 1:默认随机混剪选中,点下一步 ──────────────────────────
// h3 实际文案: "🎬 选择剪辑模式"(非 "选择模式"),用正则包含匹配
await expect(page.getByText(/选择剪辑模式/)).toBeVisible()
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
await expect(page.getByText("随机混剪")).toBeVisible()
await page.getByRole("button", { name: /下一步/ }).click()
@@ -182,8 +180,11 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
await page.getByTestId("material-card").first().click()
await page.getByRole("button", { name: /下一步/ }).click()
// ── 数量弹窗:默认 1 个 → 确认 ───────────────────────────────
await expect(page.getByText("要生成几个视频?")).toBeVisible({ timeout: 5000 })
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// ── Step 3:填写标题 ──────────────────────────────────────────
// (#2048: PreviewCountModal 已移除,生成数量在 Step1 内设置)
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
const titleInput = page.getByPlaceholder("输入或从标题库选择")
await expect(titleInput).toBeVisible({ timeout: 5000 })
@@ -191,10 +192,9 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
await page.getByRole("button", { name: /下一步/ }).click()
// ── Step 4:确认生成 ──────────────────────────────────────────
// (#2024: Step4 不再显示"📋 生成配置"卡片,内容区仅显示进度/错误)
// 等待底部操作栏的「✨ 确认生成视频」按钮可见即可
await expect(page.getByText("📋 生成配置")).toBeVisible({ timeout: 10000 })
await expect(page.getByText("随机混剪")).toBeVisible()
const confirmBtn = page.getByRole("button", { name: /确认生成视频/ })
await expect(confirmBtn).toBeVisible({ timeout: 10000 })
await expect(confirmBtn).toBeEnabled({ timeout: 5000 })
const createTask = page.waitForResponse(
@@ -324,11 +324,12 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
)
await page.goto("/app/generate")
// ── 页面标题 ─────────────────────────────────────────────────
await expect(page.getByText("智能剪辑").first()).toBeVisible({ timeout: 30000 })
await expect(page.getByRole("heading", { name: "智能剪辑" })).toBeVisible({
timeout: 30000,
})
// ── Step 1:切到叙事剪辑 → 下一步 ────────────────────────────
await expect(page.getByText(/选择剪辑模式/)).toBeVisible()
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
await page.getByText("叙事剪辑").click()
await page.getByRole("button", { name: /下一步/ }).click()
@@ -350,8 +351,11 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
await page.getByTestId("material-card").first().click()
await page.getByRole("button", { name: /下一步/ }).click()
// ── 数量弹窗 ─────────────────────────────────────────────────
await expect(page.getByText("要生成几个视频?")).toBeVisible({ timeout: 5000 })
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// ── Step 3:填写标题(handleScriptModalConfirm 已预填 script.title,但我们再覆盖一次) ─
// (#2048: PreviewCountModal 已移除)
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
const titleInput2 = page.getByPlaceholder("输入或从标题库选择")
await expect(titleInput2).toBeVisible({ timeout: 5000 })
@@ -359,9 +363,9 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
await page.getByRole("button", { name: /下一步/ }).click()
// ── Step 4:确认生成 ──────────────────────────────────────────
// (#2024: Step4 不再显示"📋 生成配置"卡片)
await expect(page.getByText("📋 生成配置")).toBeVisible({ timeout: 10000 })
await expect(page.getByText("叙事剪辑")).toBeVisible()
const confirmBtn2 = page.getByRole("button", { name: /确认生成视频/ })
await expect(confirmBtn2).toBeVisible({ timeout: 10000 })
await expect(confirmBtn2).toBeEnabled({ timeout: 5000 })
const createTask2 = page.waitForResponse(
+1 -1
View File
@@ -161,7 +161,7 @@ test.describe("Core media upload flow", () => {
const asset = data.items.find((item) => item.name === "e2e-sample.mp4")
return asset ? `${asset.mime_type || asset.file_type || ""}:${asset.status}` : "missing"
},
{ timeout: 90_000, intervals: [3_000, 5_000, 10_000] },
{ timeout: 30_000, intervals: [1_000, 2_000, 3_000] },
)
.toMatch(/^(video\/quicktime|video\/mp4|video)?:ready$/)
-7
View File
@@ -4,13 +4,6 @@
<meta charset="UTF-8" />
<link rel="icon" type="image/svg+xml" href="/vite.svg" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<!-- 标题字体(#2001 / #font-selection 修复):Google Fonts CDN 引入中文字体,保证优设标题黑/抖音美好体/阿里普惠体等fallback可用 -->
<link rel="preconnect" href="https://fonts.googleapis.com" />
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin />
<link
href="https://fonts.googleapis.com/css2?family=Noto+Sans+SC:wght@400;500;700;900&family=Noto+Serif+SC:wght@400;700;900&family=ZCOOL+KuaiLe&family=ZCOOL+XiaoWei&family=ZCOOL+QingKe+HuangYou&family=Ma+Shan+Zheng&family=Long+Cang&family=Liu+Jian+Mao+Cao&family=Zhi+Mang+Xing&display=swap"
rel="stylesheet"
/>
<title>小虾 SaaS - 自动化视频剪辑平台</title>
</head>
<body>
+2 -8
View File
@@ -51,18 +51,12 @@ export interface GenerateCoverResponse {
/** AI 生成封面 — 从最终成片中抽帧(MediaKit 选帧) */
export async function generateCover(
templateId: string | undefined | null,
templateId: string,
data: GenerateCoverRequest,
): Promise<GenerateCoverResponse> {
// templateId 为空时不传该参数,让后端使用默认模板配置
// (前端此前用 "default" 作为占位符,该 id 不存在于后端模板库会 404)
const params: Record<string, string> = {}
if (templateId && templateId !== "default") {
params.template_id = templateId
}
const response = await apiClient.post<GenerateCoverResponse>("/generation/generate-cover", data, {
timeout: 300000,
params,
params: { template_id: templateId },
})
return response.data
}
@@ -85,13 +85,6 @@ export function useSharedCover(opts: UseSharedCoverOptions): UseSharedCoverRetur
config: t.config,
}))
setTemplates(list)
// 若当前选中 "default"(初始占位),自动解析为第一个系统模板的真实 id
// ("default" 不是后端真实模板 id,传过去会 404)
setSelectedTemplateId((prev) => {
if (prev !== "default") return prev
const firstSys = list.find((t) => t.is_system)
return firstSys?.id || list[0]?.id || "default"
})
} catch (err) {
const axiosErr = err as {
response?: {
@@ -119,11 +112,6 @@ export function useSharedCover(opts: UseSharedCoverOptions): UseSharedCoverRetur
}
}, [])
useEffect(() => {
// 挂载时拉一次模板列表,用于把 "default" 占位符解析成真实模板 id
void reloadTemplates()
}, [reloadTemplates])
useEffect(() => {
if (showCoverSettings) {
void reloadTemplates()
@@ -205,12 +193,7 @@ export function useSharedCover(opts: UseSharedCoverOptions): UseSharedCoverRetur
await deleteCoverTemplate(id)
setTemplates((prev) => prev.filter((t) => t.id !== id))
if (selectedTemplateId === id) {
// 删除后选中第一个系统模板作为兜底,避免 magic string "default" 传后端 404
setTemplates((prevAfter) => {
const firstSys = prevAfter.find((t) => t.is_system)
setSelectedTemplateId(firstSys?.id || prevAfter[0]?.id || "")
return prevAfter
})
setSelectedTemplateId("default")
}
} catch (err) {
const axiosErr = err as {
@@ -248,8 +231,7 @@ export function useSharedCover(opts: UseSharedCoverOptions): UseSharedCoverRetur
}
setGenerating(true)
try {
const tplId = selectedTemplateId && selectedTemplateId !== "default" ? selectedTemplateId : ""
const url = await generateFn(tplId)
const url = await generateFn(selectedTemplateId || "default")
if (!url) {
message.warning("封面生成未返回图片,请重试")
}
@@ -284,7 +266,7 @@ export function useSharedCover(opts: UseSharedCoverOptions): UseSharedCoverRetur
const selectedTemplateName =
templates.find((t) => t.id === selectedTemplateId)?.name ||
(selectedTemplateId === "default" || !selectedTemplateId ? "默认模板" : "自定义")
(selectedTemplateId === "default" ? "默认模板" : "自定义")
return {
templates,
+9 -39
View File
@@ -20,84 +20,54 @@ export const FONT_OPTIONS: FontOption[] = [
{
value: "优设标题黑",
label: "优设标题黑",
// 原版"优设标题黑"为商用字体,非开源;这里优先使用本地已安装字体,兜底用 ZCOOL QingKe HuangYou(站酷庆科黄油体,同为厚重黑体/海报风格,可从 Google Fonts 加载)
family:
'"YouSheBiaoTiHei","YouShe Title Black","ZCOOL QingKe HuangYou","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
'"YouShe Title Black","YouSheBiaoTiHei","Source Han Sans SC Heavy","Noto Sans SC","PingFang SC",sans-serif',
tag: "hot",
},
{
value: "阿里普惠体Bold",
label: "阿里普惠体Bold",
// 阿里普惠体需要从阿里官网下载;兜底用 Noto Sans SC 900(同等字重)
family:
'"Alibaba PuHuiTi","Alibaba PuHuiTi Bold","Alibaba Sans","Noto Sans SC",system-ui,"PingFang SC","Microsoft YaHei",sans-serif',
'"Alibaba PuHuiTi Bold","Alibaba PuHuiTi","Source Han Sans SC","PingFang SC",sans-serif',
tag: "hot",
},
{
value: "抖音美好体",
label: "抖音美好体",
// 抖音美好体版权字体兜底 ZCOOL KuaiLe(站酷快乐体,圆润卡通风格近似)
family:
'"Douyin Sans","DouyinSans","ZCOOL KuaiLe","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
family: '"Douyin Sans","DouyinSans","Source Han Sans SC","PingFang SC",sans-serif',
tag: "hot",
},
{
value: "思源黑体Heavy",
label: "思源黑体Heavy",
family:
'"Noto Sans SC","Source Han Sans SC","Source Han Sans CN Heavy","PingFang SC","Microsoft YaHei",sans-serif',
'"Source Han Sans SC Heavy","Noto Sans SC","Source Han Sans CN Heavy","PingFang SC",sans-serif',
tag: "new",
},
{
value: "思源黑体",
label: "思源黑体",
family: '"Noto Sans SC","Source Han Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
family: '"Source Han Sans SC","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
},
{
value: "思源宋体",
label: "思源宋体",
family: '"Noto Serif SC","Source Han Serif SC","Songti SC","SimSun",serif',
family: '"Source Han Serif SC","Noto Serif SC","Songti SC","SimSun",serif',
},
{
value: "苹方",
label: "苹方",
family:
'"PingFang SC",-apple-system,blinkmacsystemfont,"Helvetica Neue","Noto Sans SC",sans-serif',
family: '"PingFang SC",-apple-system,"Helvetica Neue",sans-serif',
},
{
value: "微软雅黑",
label: "微软雅黑",
family: '"Microsoft YaHei","PingFang SC","Noto Sans SC",sans-serif',
family: '"Microsoft YaHei","PingFang SC",sans-serif',
},
{
value: "楷体",
label: "楷体",
family: '"KaiTi","STKaiti","DFKai-SB","Kaiti SC",serif',
},
{
value: "站酷小薇体",
label: "站酷小薇体",
family: '"ZCOOL XiaoWei","Noto Serif SC",serif',
},
{
value: "马善政毛笔",
label: "马善政毛笔",
family: '"Ma Shan Zheng","STXingkai","KaiTi",cursive',
},
{
value: "龙藏体",
label: "龙藏体",
family: '"Long Cang","STXingkai","KaiTi",cursive',
},
{
value: "流江毛笔草",
label: "流江毛笔草",
family: '"Liu Jian Mao Cao","STXingkai",cursive',
},
{
value: "志莽行书",
label: "志莽行书",
family: '"Zhi Mang Xing","STXingkai",cursive',
family: '"KaiTi","STKaiti","DFKai-SB",serif',
},
]
@@ -412,7 +412,7 @@ const PanelTitleConfig: React.FC<PanelTitleConfigProps> = ({ titleConfig, onUpda
onUpdateStyle={handleUpdateStyle}
showCoverToggle
previewWidth={280}
enableTemplates={false}
enableTemplates
selectedTemplateId={selectedTemplateId}
onApplyTemplate={handleApplyTemplate}
activePreset={activePreset}
+59 -58
View File
@@ -14,6 +14,7 @@ import VoiceSelectModal from "./components/VoiceSelectModal"
import ScriptSelectModal from "./components/ScriptSelectModal"
import TtsVoiceModal from "./components/TtsVoiceModal"
import GenerateHeader from "./components/GenerateHeader"
import PreviewCountModal from "./components/PreviewCountModal"
import GenerateStepsBar from "./components/GenerateStepsBar"
import GenerateStepContent from "./components/GenerateStepContent"
import GenerateStepActions from "./components/GenerateStepActions"
@@ -106,7 +107,6 @@ const GeneratePage: React.FC = () => {
setPreviewCovers,
selectedVariantIds,
setSelectedVariantIds,
setSelectedTemplate,
} = formState
const isBatch = previewCount > 1
@@ -137,8 +137,8 @@ const GeneratePage: React.FC = () => {
}
}, [selectedVoice, isBatch, voiceModePerVideo, setVoiceLibraryIds])
/* ── 标题面板模式:默认展示样式参数(false);需要大卡片模板网格时再切 true ── */
const enableTemplates = false
/* ── 数量选择弹窗 ── */
const [countModalOpen, setCountModalOpen] = useState(false)
/* ── Step5 保存中状态 ── */
const [finishing, setFinishing] = useState(false)
@@ -199,7 +199,6 @@ const GeneratePage: React.FC = () => {
generated,
generateError,
generatedVideos,
currentTaskId,
batchTasks,
generate: handleGenerate,
retry: handleRetryGenerate,
@@ -245,37 +244,38 @@ const GeneratePage: React.FC = () => {
},
})
/* ── 对齐批量数组长度到 previewCount(用于进入 Step3 时) ── */
const ensureArraysAligned = useCallback(() => {
setPreviewTitles((prev) => {
const list = prev || []
if (list.length === previewCount) return list
const base = list[0] || titleSettings.title || ""
return Array.from({ length: previewCount }, (_, i) => list[i] ?? (i === 0 ? base : ""))
})
setVoiceLibraryIds((prev) => {
const list = prev || []
if (list.length === previewCount) return list
return Array.from({ length: previewCount }, (_, i) => list[i] ?? selectedVoice ?? "")
})
setPreviewCovers((prev) => {
const list = prev || []
if (list.length === previewCount) return list
return Array.from({ length: previewCount }, (_, i) => list[i] ?? "")
})
setSelectedVariantIds((prev) => {
if (prev && prev.length === previewCount) return prev
return Array.from({ length: previewCount }, (_, i) => i)
})
}, [
previewCount,
setPreviewTitles,
setVoiceLibraryIds,
setPreviewCovers,
setSelectedVariantIds,
titleSettings.title,
selectedVoice,
])
/* ── 数量弹窗确认 ── */
const handleCountConfirm = useCallback(
(count: number) => {
setPreviewCount(count)
setCountModalOpen(false)
setPreviewTitles((prev) => {
const list = prev || []
const base = list[0] || titleSettings.title || ""
return Array.from({ length: count }, (_, i) => list[i] ?? (i === 0 ? base : ""))
})
setVoiceLibraryIds((prev) => {
const list = prev || []
return Array.from({ length: count }, (_, i) => list[i] ?? selectedVoice ?? "")
})
setPreviewCovers((prev) => {
const list = prev || []
return Array.from({ length: count }, (_, i) => list[i] ?? "")
})
setSelectedVariantIds(Array.from({ length: count }, (_, i) => i))
setCurrentStep(3)
},
[
setPreviewCount,
setPreviewTitles,
setVoiceLibraryIds,
setPreviewCovers,
setSelectedVariantIds,
setCurrentStep,
titleSettings.title,
selectedVoice,
],
)
/* ── #1970:Step1 弹窗回调 ── */
const handleVoiceModalConfirm = useCallback(
@@ -398,7 +398,7 @@ const GeneratePage: React.FC = () => {
smartSelectedIds,
titleSettings,
generated,
onBeforeEnterStep3: ensureArraysAligned,
onOpenCountModal: () => setCountModalOpen(true),
onOpenStep1Modal: () => {
if (editMode === "random") {
setVoiceModalOpen(true)
@@ -429,31 +429,30 @@ const GeneratePage: React.FC = () => {
setFinishing(true)
const hide = message.loading("正在保存到视频库...", 0)
try {
// 收集需要 finalize 的任务 ID:批量用 batchTasks;单视频优先用 finalVideo.generation_task_id,兜底 currentTaskId
const singleTaskId = finalVideo?.generation_task_id || currentTaskId || ""
const taskIds =
batchTasks && batchTasks.length > 0
? batchTasks.map((t) => t.taskId).filter(Boolean)
: finalVideo?.generation_task_id
? [finalVideo.generation_task_id]
: []
// 单视频/批量:为每个任务调用 finalize(入库 + 绑定封面 + 自定义标题)
// 批量时必须按 batchTasks[i].variantIndex 对齐 previewCovers/previewTitles(taskIds 顺序不一定按变体序号)
if (isBatch && batchTasks.length > 0) {
// 单视频/批量:为每个 awaiting_cover 任务调用 finalize(入库 + 绑定封面 + 自定义标题)
if (isBatch && previewCovers.length > 0) {
await Promise.all(
batchTasks.map(async (task) => {
const vi = task.variantIndex
const coverUrl = previewCovers[vi] || ""
const title = previewTitles[vi] || titleSettings.title || ""
return finalizeGeneration(task.taskId, {
taskIds.map(async (taskId, idx) => {
const coverUrl = previewCovers[idx] || ""
return finalizeGeneration(taskId, {
cover_url: coverUrl || undefined,
custom_title: title,
custom_title: previewTitles[idx] || titleSettings.title || "",
})
}),
)
} else if (singleTaskId) {
} else if (finalVideo?.generation_task_id) {
const coverUrl = coverSettings.thumbnail_url || coverSettings.upload_url || ""
await finalizeGeneration(singleTaskId, {
await finalizeGeneration(finalVideo.generation_task_id, {
cover_url: coverUrl || undefined,
custom_title: titleSettings.title || "",
})
} else {
console.warn("[handleFinish] 未找到任务 ID,跳过 finalize 直接跳转")
}
hide()
@@ -482,7 +481,6 @@ const GeneratePage: React.FC = () => {
previewTitles,
titleSettings.title,
coverSettings,
currentTaskId,
navigate,
])
@@ -541,7 +539,7 @@ const GeneratePage: React.FC = () => {
onUpdateStyle={styleUpdaters.updateStyle}
activePreset={styleUpdaters.activePreset}
titlePresets={styleUpdaters.titlePresets}
enableTemplates={enableTemplates}
enableTemplates
selectedTemplateId={selectedTitleTemplateId}
onApplyTemplate={(settings, tpl) => {
styleUpdaters.applyTemplate(settings)
@@ -565,8 +563,6 @@ const GeneratePage: React.FC = () => {
generateError={generateError}
progress={progress}
generatedVideos={generatedVideos}
currentTaskId={currentTaskId}
onRetry={handleRetryGenerate}
onRetryBatchTask={handleRetryBatchTask}
onDismissError={handleDismissError}
@@ -581,9 +577,6 @@ const GeneratePage: React.FC = () => {
previewCovers={previewCovers}
onPreviewCoversChange={setPreviewCovers}
selectedVariantIds={selectedVariantIds}
selectedCoverTemplate={selectedTemplate}
onSelectedCoverTemplateChange={setSelectedTemplate}
onConfirmGenerate={handleConfirmGenerate}
/>
{/* ════ 步骤4(单视频):成片播放器 ════ */}
@@ -671,6 +664,14 @@ const GeneratePage: React.FC = () => {
</div>
</div>
{/* 数量选择弹窗 */}
<PreviewCountModal
open={countModalOpen}
defaultCount={1}
onConfirm={handleCountConfirm}
onCancel={() => setCountModalOpen(false)}
/>
{/* 音色克隆弹窗 */}
<CloneModal
open={cloneModalOpen}
@@ -89,12 +89,6 @@ export interface GenerateStepContentProps {
previewCovers: string[]
onPreviewCoversChange: (urls: string[]) => void
selectedVariantIds?: number[]
selectedCoverTemplate?: string
onSelectedCoverTemplateChange?: (templateId: string) => void
/** 单视频任务 ID(兜底,awaiting_cover 状态下 results 接口未入库时用) */
currentTaskId?: string
/** Step3 右上角确认生成按钮 */
onConfirmGenerate?: () => void | Promise<void>
}
export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) => {
@@ -148,9 +142,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
previewCovers,
onPreviewCoversChange,
selectedVariantIds,
selectedCoverTemplate,
onSelectedCoverTemplateChange,
onConfirmGenerate,
} = props
switch (currentStep) {
@@ -204,11 +195,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
previewCount={previewCount}
previewTitles={previewTitles}
onPreviewTitlesChange={onPreviewTitlesChange}
onConfirmGenerate={onConfirmGenerate}
generating={props.generating}
selectedCount={
props.previewCount && props.previewCount > 1 ? props.selectedVariantIds?.length || 1 : 1
}
/>
)
case 4:
@@ -264,9 +250,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
previewCovers={previewCovers}
onPreviewCoversChange={onPreviewCoversChange}
selectedVariantIndexes={selectedVariantIds}
selectedTemplate={selectedCoverTemplate}
onTemplateChange={onSelectedCoverTemplateChange}
currentTaskId={props.currentTaskId}
/>
)
default:
@@ -0,0 +1,127 @@
/**
* 生成数量选择弹窗(Issue #1677)
* Step1 选完模板点「下一步」时弹出:要生成几个视频?(1~10)
* 默认 1,回车 = 1(零额外操作)
*/
import React, { useState, useEffect, useRef } from "react"
import { MAX_PREVIEW_COUNT } from "../constants"
interface PreviewCountModalProps {
open: boolean
/** 默认值(上次选择,默认1) */
defaultCount?: number
onConfirm: (count: number) => void
onCancel: () => void
}
const PreviewCountModal: React.FC<PreviewCountModalProps> = ({
open,
defaultCount = 1,
onConfirm,
onCancel,
}) => {
const [count, setCount] = useState(defaultCount)
const inputRef = useRef<HTMLInputElement>(null)
useEffect(() => {
if (open) {
setCount(defaultCount)
// 弹窗打开后聚焦并选中,方便直接回车=默认1
setTimeout(() => inputRef.current?.focus(), 50)
}
}, [open, defaultCount])
const clamp = (n: number) => Math.max(1, Math.min(MAX_PREVIEW_COUNT, n || 1))
const handleConfirm = () => {
onConfirm(clamp(count))
}
const handleKeyDown = (e: React.KeyboardEvent) => {
if (e.key === "Enter") {
e.preventDefault()
handleConfirm()
}
if (e.key === "Escape") {
onCancel()
}
}
if (!open) return null
return (
<div className="xx-modal-mask" onClick={onCancel}>
<div className="xx-modal-box xx-count-modal" onClick={(e) => e.stopPropagation()}>
<h3 style={{ margin: "0 0 8px", fontSize: 18 }}>要生成几个视频?</h3>
<p style={{ margin: "0 0 20px", fontSize: 13, color: "var(--text-secondary, #666)" }}>
素材共用,AI 随机剪辑出不同版本,每个视频可独立设置标题、配音和封面
</p>
<div className="xx-count-selector">
<button
type="button"
className="xx-count-btn"
onClick={() => setCount((c) => clamp(c - 1))}
disabled={count <= 1}
aria-label="减少"
>
−
</button>
<input
ref={inputRef}
type="number"
min={1}
max={MAX_PREVIEW_COUNT}
value={count}
onChange={(e) => setCount(clamp(parseInt(e.target.value, 10) || 1))}
onKeyDown={handleKeyDown}
className="xx-count-input"
/>
<button
type="button"
className="xx-count-btn"
onClick={() => setCount((c) => clamp(c + 1))}
disabled={count >= MAX_PREVIEW_COUNT}
aria-label="增加"
>
+
</button>
</div>
<div className="xx-count-quick">
{[1, 3, 5, 10].map((n) => (
<button
key={n}
type="button"
className={`xx-count-chip ${count === n ? "active" : ""}`}
onClick={() => setCount(n)}
>
{n} 个
</button>
))}
</div>
<div className="xx-count-actions">
<button type="button" className="xx-btn xx-btn-ghost" onClick={onCancel}>
取消
</button>
<button type="button" className="xx-btn xx-btn-primary" onClick={handleConfirm}>
{count === 1 ? "生成 1 个视频" : `生成 ${count} 个视频`}
</button>
</div>
<p
style={{
margin: "12px 0 0",
fontSize: 12,
color: "var(--text-tertiary, #999)",
textAlign: "center",
}}
>
直接按回车 = 生成 1 个
</p>
</div>
</div>
)
}
export default PreviewCountModal
@@ -47,12 +47,6 @@ interface Step4TitleSettingsProps {
enableTemplates?: boolean
selectedTemplateId?: string | null
onApplyTemplate?: (settings: TitleSettings, template: TitleTemplate) => void
/** Step3 右上角「🎬 确认生成」主按钮 */
onConfirmGenerate?: () => void | Promise<void>
/** 是否生成中 */
generating?: boolean
/** 批量模式下勾选数量 */
selectedCount?: number
}
const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
@@ -75,9 +69,6 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
enableTemplates,
selectedTemplateId,
onApplyTemplate,
onConfirmGenerate,
generating,
selectedCount = 1,
} = props
const isBatch = previewCount > 1
@@ -99,47 +90,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
)
return (
<div className="xx-form-section" style={{ position: "relative" }}>
{/* ── 右上角「🎬 确认生成」主按钮 ── */}
{onConfirmGenerate && (
<button
type="button"
onClick={() => {
if (generating) return
void onConfirmGenerate()
}}
disabled={generating}
style={{
position: "absolute",
top: 0,
right: 0,
background: generating ? "#a78bfa" : "#7c3aed",
color: "#fff",
border: "none",
borderRadius: 10,
padding: "12px 24px",
fontSize: 15,
fontWeight: 600,
cursor: generating ? "not-allowed" : "pointer",
boxShadow: "0 4px 14px rgba(124,58,237,0.4)",
transition: "all .2s",
zIndex: 5,
whiteSpace: "nowrap",
}}
onMouseEnter={(e) => {
if (!generating) (e.currentTarget as HTMLButtonElement).style.background = "#6d28d9"
}}
onMouseLeave={(e) => {
if (!generating) (e.currentTarget as HTMLButtonElement).style.background = "#7c3aed"
}}
>
{generating
? "⏳ 生成中..."
: selectedCount > 1
? `🎬 确认生成 ${selectedCount} 个视频`
: "🎬 确认生成"}
</button>
)}
<div className="xx-form-section">
<h3>📝 选择标题</h3>
{!isBatch ? (
@@ -30,8 +30,6 @@ interface Step6CoverSettingsProps {
onPreviewCoversChange?: (urls: string[]) => void
selectedVariantIndexes?: number[]
onTemplateChange?: (templateId: string) => void
/** 单视频任务 ID(awaiting_cover 阶段 results 接口可能返回 preview-xxx 合成对象,兜底用) */
currentTaskId?: string
}
const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
@@ -49,32 +47,6 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
props.generatedVideos.find((v) => v.status === "completed" || v.status === "awaiting_cover") ||
props.generatedVideos[0]
/**
* 兜底任务/视频 ID:awaiting_cover 阶段后端 /results 可能还没有入库 GeneratedVideo,
* 只返回合成的 preview-{taskId} 轻量对象;此时用 currentTaskId 兜底让后端能找到任务。
* 同时统一抽取 taskId(generation_task_id 优先)用于日志/错误提示。
*/
const effectiveTaskId =
(finalVideo as { generation_task_id?: string } | undefined)?.generation_task_id ||
props.currentTaskId ||
""
const _rawVideoId =
(finalVideo as { id?: string; video_id?: string } | undefined)?.id ||
(finalVideo as { video_id?: string } | undefined)?.video_id ||
""
// preview-{taskId} 是后端合成的临时 id,gv_repo.get 查不到 → 不传 generated_video_id,
// 让后端走 plan.config.generation_task_id / rendered_storage_key 兜底路径。
const effectiveVideoId = _rawVideoId && !_rawVideoId.startsWith("preview-") ? _rawVideoId : ""
const effectiveVideoUrl = finalVideo?.file_url || finalVideo?.download_url || ""
/** 按钮可用:非批量 且 (有 finalVideo 对象或兜底 taskId) 且 视频状态已完成/等待封面/未设置 */
const isVideoReady =
!finalVideo ||
finalVideo.status === "completed" ||
finalVideo.status === "awaiting_cover" ||
!finalVideo.status
const canGenerateCover = !isBatch && (!!finalVideo || !!effectiveTaskId) && isVideoReady
const completedVideos = useMemo(
() =>
props.generatedVideos.filter(
@@ -88,58 +60,37 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
* 批量场景 canGenerate=false,避免 shared.generateAutoCover 被误触发
*/
const shared = useSharedCover({
canGenerate: canGenerateCover,
disabledHint: isBatch
? "批量场景请在上方操作卡片"
: !finalVideo && !effectiveTaskId
? "请先生成视频再选择封面"
: "视频尚未就绪,请稍候",
initialTemplateId: "default", // 封面模板独立于编辑模板,默认用 default
canGenerate: !!finalVideo && !isBatch,
disabledHint: isBatch ? "批量场景请在上方操作卡片" : "请先生成视频再选择封面",
initialTemplateId: props.selectedTemplate || "default",
generateFn: async (tplId) => {
if (isBatch) return null
if (!finalVideo && !effectiveTaskId) {
console.warn("[Cover] generateAutoCover: no finalVideo and no taskId")
return null
}
// 请求体:generated_video_id 仅在后端已入库(非 preview-xxx 合成id)时传;
// video_url 兜底让后端能直接下载视频抽帧;generation_task_id 后端已从 plan.config 自动读取。
const requestBody: {
generated_video_id?: string
video_url?: string
cover_type: "ai_frame"
title_config?: Record<string, unknown>
} = {
if (!finalVideo || isBatch) return null
const response = await apiGenerateCover(tplId, {
generated_video_id: finalVideo.id,
video_url: finalVideo.file_url || finalVideo.download_url || "",
cover_type: "ai_frame",
}
if (effectiveVideoId) {
requestBody.generated_video_id = effectiveVideoId
}
if (effectiveVideoUrl) {
requestBody.video_url = effectiveVideoUrl
}
if (props.titleSettings?.title) {
requestBody.title_config = {
text: props.titleSettings.title,
font: props.titleSettings.font,
font_size: props.titleSettings.size,
font_color: props.titleSettings.color,
position: props.titleSettings.position,
bold: props.titleSettings.bold,
stroke: props.titleSettings.stroke,
shadow: props.titleSettings.shadow,
}
}
console.log("[Cover] auto-generate request:", { tplId, ...requestBody })
const response = await apiGenerateCover(tplId, requestBody)
const url = response.cover?.image_url || response.cover?.thumbnail_url || ""
...(props.titleSettings?.title
? {
title_config: {
text: props.titleSettings.title,
font: props.titleSettings.font,
font_size: props.titleSettings.size,
font_color: props.titleSettings.color,
position: props.titleSettings.position,
bold: props.titleSettings.bold,
stroke: props.titleSettings.stroke,
shadow: props.titleSettings.shadow,
},
}
: {}),
})
const url = response.cover?.image_url || ""
if (url) {
props.onCoverSettingsChange({
...props.coverSettings,
thumbnail_url: url,
ai_suggested_time: response.cover?.frame_time ?? null,
})
} else {
console.warn("[Cover] generate returned empty url:", response)
}
return url
},
@@ -147,12 +98,6 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
// 选中模板变化时通知父组件(用于批量生成时透传 template_id)
const { onTemplateChange, selectedTemplate: parentSelectedTemplate } = props
// 父组件 selectedTemplate 变化时同步到子(例如从 Step1/Step4 切换到 Step6 时)
useEffect(() => {
if (parentSelectedTemplate && parentSelectedTemplate !== shared.selectedTemplateId) {
shared.handleSelectTemplate(parentSelectedTemplate)
}
}, [parentSelectedTemplate]) // eslint-disable-line react-hooks/exhaustive-deps
useEffect(() => {
if (isBatch && onTemplateChange && shared.selectedTemplateId !== parentSelectedTemplate) {
onTemplateChange(shared.selectedTemplateId)
@@ -186,7 +131,7 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
* 透传给 useBatchCovers,由其在 generateOne/generateAll 中发给后端。
*/
const batchCovers = useBatchCovers({
selectedTemplate: shared.selectedTemplateId,
selectedTemplate: shared.selectedTemplateId || "default",
generatedVideos: props.generatedVideos,
titles: batchTitles,
titleStyle: {
@@ -364,7 +309,7 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
<div className="xx-form-section">
<h3>🖼️ 选择封面</h3>
{(finalVideo || effectiveTaskId) && (
{finalVideo && (
<div
style={{
padding: "10px 14px",
@@ -376,7 +321,7 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
color: "var(--text-secondary, #666)",
}}
>
🎬 封面将从最终成片{finalVideo?.name ? `「${finalVideo.name}」` : ""}中智能选帧
🎬 封面将从最终成片「{finalVideo.name}」中智能选帧
{shared.selectedTemplateId && shared.selectedTemplateId !== "default" && (
<>
{" "}
@@ -390,9 +335,8 @@ const Step6CoverSettings: React.FC<Step6CoverSettingsProps> = (props) => {
<Button
buttonType="primary"
onClick={() => void shared.generateAutoCover()}
disabled={!canGenerateCover || shared.generating}
disabled={!finalVideo || shared.generating}
loading={shared.generating}
title={!canGenerateCover ? "请先完成视频生成" : ""}
>
✨ 自动生成封面
</Button>
@@ -7,7 +7,7 @@ import type {
TextDirection,
StrokeStyle,
} from "../../types/cover"
import { DEFAULT_EDITOR_CONFIG, ALL_FONTS } from "../../types/cover"
import { DEFAULT_EDITOR_CONFIG, PRESET_FONTS, SYSTEM_FONTS, ALL_FONTS } from "../../types/cover"
import Modal from "@/components/ui/Modal"
import Button from "@/components/ui/Button"
import "@/components/cover/cover.css"
@@ -16,19 +16,6 @@ import "@/components/cover/cover.css"
const mergeEditorConfig = (partial?: Partial<CoverEditorConfig> | null): CoverEditorConfig => {
const def = DEFAULT_EDITOR_CONFIG
const src = partial || {}
/** 兼容老模板:老版本 background 有 posX/posY/rotation,新版改为 offsetY(相对文字位置偏移)。
* 老模板黑底默认 posY 通常是 50(与文字对齐)或 80(副标题偏下),统一归一为 offsetY=0,
* 因为新版背景位置已自动跟随文字位置,offsetY 仅做相对微调。 */
const normalizeBg = (bg: Record<string, unknown> | undefined) => {
if (!bg) return {}
// 兼容老模板字段:posX/posY/rotation 在新版中已改为 offsetY(背景位置自动跟随文字)
const normalized = { ...bg }
delete (normalized as Record<string, unknown>).posX
delete (normalized as Record<string, unknown>).posY
delete (normalized as Record<string, unknown>).rotation
if (normalized.offsetY == null) normalized.offsetY = 0
return normalized
}
const mergeText = (
base: TextStyleConfig,
patch?: Partial<TextStyleConfig> | null,
@@ -36,10 +23,7 @@ const mergeEditorConfig = (partial?: Partial<CoverEditorConfig> | null): CoverEd
...base,
...(patch || {}),
position: { ...base.position, ...(patch?.position || {}) },
background: {
...base.background,
...normalizeBg(patch?.background as Record<string, unknown> | undefined),
},
background: { ...base.background, ...(patch?.background || {}) },
shadows: Array.isArray(patch?.shadows) ? [...patch!.shadows] : [...base.shadows],
})
return {
@@ -109,23 +93,27 @@ const PositionPair: React.FC<{
</div>
)
/* ── Font Select options(合并预设+系统,用圆点颜色区分 tag) ── */
const fontDotClass = (tag?: string) => {
if (tag === "preset") return "xx-ce-font-dot xx-ce-font-dot--preset"
if (tag === "hand") return "xx-ce-font-dot xx-ce-font-dot--hand"
if (tag === "serif") return "xx-ce-font-dot xx-ce-font-dot--serif"
if (tag === "mono") return "xx-ce-font-dot xx-ce-font-dot--mono"
return "xx-ce-font-dot xx-ce-font-dot--system"
}
const fontOptions = ALL_FONTS.map((f) => ({
label: (
<span>
<span className={fontDotClass(f.tag)} />
<span style={{ fontFamily: f.family }}>{f.name}</span>
</span>
),
value: f.name,
}))
/* ── Font Select options ── */
const fontOptions = [
...PRESET_FONTS.map((f) => ({
label: (
<span>
<span className="xx-ce-font-dot xx-ce-font-dot--preset" />
<span style={{ fontFamily: f.family }}>{f.name}</span>
</span>
),
value: f.name,
})),
...SYSTEM_FONTS.map((f) => ({
label: (
<span>
<span className="xx-ce-font-dot xx-ce-font-dot--system" />
<span style={{ fontFamily: f.family }}>{f.name}</span>
</span>
),
value: f.name,
})),
]
/* ── Find font family string from name ── */
const getFontFamily = (name: string): string => {
@@ -306,122 +294,6 @@ const TextStylePanel: React.FC<{
onChange={(v) => upd("rotation", v)}
/>
</div>
{/* 文字背景(剪映样式) */}
<div className="xx-ce-switch-item">
<div className="xx-ce-switch-row">
<span>文字背景</span>
<Switch
size="small"
checked={config.background?.enabled ?? false}
onChange={(v) => upd("background", { ...config.background, enabled: v })}
/>
</div>
{config.background?.enabled && (
<div style={{ marginTop: 8, paddingLeft: 4 }}>
<div className="xx-ce-row">
<label className="xx-ce-label">形状</label>
<div className="xx-ce-radio-group">
<button
type="button"
className={`xx-ce-radio-btn ${config.background.shape === "rectangle" ? "active" : ""}`}
onClick={() =>
upd("background", {
...config.background,
shape: "rectangle" as const,
})
}
>
矩形
</button>
<button
type="button"
className={`xx-ce-radio-btn ${config.background.shape === "polygon" ? "active" : ""}`}
onClick={() =>
upd("background", {
...config.background,
shape: "polygon" as const,
})
}
>
圆角
</button>
</div>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">背景颜色</label>
<ColorPicker
value={config.background.color || "#000000"}
onChange={(v) => upd("background", { ...config.background, color: v })}
/>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">不透明度: {config.background.opacity}%</label>
<Slider
min={0}
max={100}
step={1}
value={config.background.opacity}
onChange={(v) => upd("background", { ...config.background, opacity: v })}
/>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">
圆角:{" "}
{config.background.shape === "rectangle"
? 0
: Math.round((config.background.height ?? 20) / 2)}
px
</label>
<Slider
min={0}
max={50}
step={1}
value={
config.background.shape === "rectangle"
? 0
: Math.round((config.background.height ?? 20) / 2)
}
disabled={config.background.shape === "rectangle"}
onChange={(v) => {
// 圆角近似:通过 height 控制
void v
}}
/>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">宽度: {config.background.width}%</label>
<Slider
min={20}
max={200}
step={2}
value={config.background.width}
onChange={(v) => upd("background", { ...config.background, width: v })}
/>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">高度: {config.background.height}%</label>
<Slider
min={5}
max={80}
step={1}
value={config.background.height}
onChange={(v) => upd("background", { ...config.background, height: v })}
/>
</div>
<div className="xx-ce-row">
<label className="xx-ce-label">上下偏移: {config.background.offsetY ?? 0}%</label>
<Slider
min={-30}
max={30}
step={1}
value={config.background.offsetY ?? 0}
onChange={(v) => upd("background", { ...config.background, offsetY: v })}
/>
</div>
</div>
)}
</div>
</div>
)
}
@@ -444,59 +316,6 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
const portraitFileRef = useRef<HTMLInputElement>(null)
const [bgImageUrl, setBgImageUrl] = useState<string>(initCfg.backgroundImage || "")
const [portraitImageUrl, setPortraitImageUrl] = useState<string>(initCfg.portraitImage || "")
/* ── 文字拖拽状态 ── */
const canvasRef = useRef<HTMLDivElement | null>(null)
const dragRef = useRef<null | {
target: "title" | "subtitle"
startX: number
startY: number
startPosX: number
startPosY: number
}>(null)
const handleTextMouseDown = (
e: React.MouseEvent<HTMLDivElement>,
target: "title" | "subtitle",
) => {
e.preventDefault()
e.stopPropagation()
const tc = target === "title" ? cfg.title : cfg.subtitle
if (!tc) return
dragRef.current = {
target,
startX: e.clientX,
startY: e.clientY,
startPosX: tc.position.x,
startPosY: tc.position.y,
}
const onMove = (ev: MouseEvent) => {
const d = dragRef.current
if (!d || !canvasRef.current) return
const rect = canvasRef.current.getBoundingClientRect()
const dx = ((ev.clientX - d.startX) / rect.width) * 100
const dy = ((ev.clientY - d.startY) / rect.height) * 100
const newX = Math.max(0, Math.min(100, d.startPosX + dx))
const newY = Math.max(0, Math.min(100, d.startPosY + dy))
if (d.target === "title") {
setCfg((prev) => ({
...prev,
title: { ...prev.title, position: { x: newX, y: newY } },
}))
} else {
setCfg((prev) => ({
...prev,
subtitle: { ...prev.subtitle, position: { x: newX, y: newY } },
}))
}
}
const onUp = () => {
dragRef.current = null
window.removeEventListener("mousemove", onMove)
window.removeEventListener("mouseup", onUp)
}
window.addEventListener("mousemove", onMove)
window.addEventListener("mouseup", onUp)
}
// reset state when modal opens with a new template
useEffect(() => {
@@ -575,6 +394,19 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
const renderTextStyle = (tc: TextStyleConfig | undefined | null): React.CSSProperties => {
if (!tc) return {}
// wrap text by charsPerLine
const rawText = tc.text || ""
const lines: string[] = []
if (tc.direction === "vertical") {
lines.push(rawText)
} else {
for (let i = 0; i < rawText.length; i += Math.max(1, tc.charsPerLine)) {
lines.push(rawText.slice(i, i + Math.max(1, tc.charsPerLine)))
}
}
// store rendered as lines via data attribute; for JSX we'll render outside style
void lines
const style: React.CSSProperties = {
position: "absolute",
left: `${tc.position.x}%`,
@@ -590,12 +422,9 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
tc.strokeWidth > 0 ? `${Math.max(0.5, s(tc.strokeWidth))}px ${tc.strokeColor}` : undefined,
whiteSpace: tc.direction === "vertical" ? "pre-wrap" : "pre",
writingMode: tc.direction === "vertical" ? "vertical-rl" : undefined,
zIndex: 4,
zIndex: 3,
textAlign: "center",
userSelect: "none",
cursor: "grab",
padding: 0,
pointerEvents: "auto",
}
if (tc.shadows.length > 0) {
style.textShadow = tc.shadows
@@ -605,39 +434,6 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
return style
}
const renderTextBgStyle = (tc: TextStyleConfig | undefined | null): React.CSSProperties => {
if (!tc?.background?.enabled) return { display: "none" }
const bg = tc.background
// 背景位置跟随文字:left/top 对齐文字中心,用 offsetY(-50~50% 相对文字位置)做上下微调
// 这样文字拖拽时背景会自动跟随,不需要独立的位置控制
const radius = bg.shape === "polygon" ? `${Math.max(4, Math.round(bg.height / 4))}px` : "0"
const alpha = Math.max(0, Math.min(1, bg.opacity / 100))
const hex = (bg.color || "#000000").replace("#", "")
let r = 0,
g = 0,
b = 0
if (hex.length === 6) {
r = parseInt(hex.substring(0, 2), 16)
g = parseInt(hex.substring(2, 4), 16)
b = parseInt(hex.substring(4, 6), 16)
}
const rgba = `rgba(${r}, ${g}, ${b}, ${alpha})`
// bg.offsetY 是相对文字位置的上下偏移(-50~50,单位%画布高度),默认 0 表示与文字中心对齐
const offsetY = typeof bg.offsetY === "number" ? bg.offsetY : 0
return {
position: "absolute",
left: `${tc.position.x}%`,
top: `calc(${tc.position.y}% + ${offsetY}%)`,
width: `${bg.width}%`,
height: `${bg.height}%`,
transform: "translate(-50%, -50%)",
background: rgba,
borderRadius: radius,
zIndex: 3,
pointerEvents: "none",
}
}
const renderTextLines = (tc: TextStyleConfig | undefined | null): string => {
if (!tc) return ""
const raw = tc.text || ""
@@ -1061,7 +857,7 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
<span className="xx-ce-anchor-dot" style={{ top: "50%", right: "-12px" }} />
<span className="xx-ce-anchor-dot" style={{ bottom: "10%", left: "-12px" }} />
<div className="xx-ce-canvas" ref={canvasRef}>
<div className="xx-ce-canvas">
{/* base background layer */}
<div
className="xx-ce-canvas-base"
@@ -1118,29 +914,14 @@ const CoverEditorModal: React.FC<CoverEditorModalProps> = ({ open, onClose, temp
/>
)}
{/* Title bg */}
{cfg.title?.background?.enabled && <div style={renderTextBgStyle(cfg.title)} />}
{/* Subtitle bg */}
{cfg.subtitle?.background?.enabled && <div style={renderTextBgStyle(cfg.subtitle)} />}
{/* Title (draggable) */}
{/* Title */}
{cfg.title && (
<div
style={renderTextStyle(cfg.title)}
onMouseDown={(e) => handleTextMouseDown(e, "title")}
>
{renderTextLines(cfg.title)}
</div>
<div style={renderTextStyle(cfg.title)}>{renderTextLines(cfg.title)}</div>
)}
{/* Subtitle (draggable) */}
{/* Subtitle */}
{cfg.subtitle && (
<div
style={renderTextStyle(cfg.subtitle)}
onMouseDown={(e) => handleTextMouseDown(e, "subtitle")}
>
{renderTextLines(cfg.subtitle)}
</div>
<div style={renderTextStyle(cfg.subtitle)}>{renderTextLines(cfg.subtitle)}</div>
)}
{/* Mask overlay */}
@@ -1,4 +1,4 @@
import React, { useMemo, useState } from "react"
import React from "react"
import type { CoverTemplate } from "../../types/cover"
import Modal from "@/components/ui/Modal"
import Button from "@/components/ui/Button"
@@ -17,58 +17,15 @@ interface CoverSettingsModalProps {
onCreateNew: () => void
}
/** 模板缩略图:优先渲染 thumbnail_url;加载失败/无图时展示占位 */
const TemplateThumb: React.FC<{ tpl: CoverTemplate; isSelected: boolean }> = ({
tpl,
isSelected,
}) => {
const [errored, setErrored] = useState(false)
const url = tpl.thumbnail_url && !errored ? tpl.thumbnail_url : ""
// 随机柔和渐变做占位,保证卡片不会灰成一片
const placeholderBg = useMemo(() => {
const palettes = [
["#e0e0e0", "#c0c0c0"],
["#ef4444", "#b91c1c"],
["#374151", "#111827"],
["#3b82f6", "#1d4ed8"],
["#8b5cf6", "#6d28d9"],
["#f97316", "#ea580c"],
["#22c55e", "#15803d"],
["#06b6d4", "#0e7490"],
]
let h = 0
for (const ch of tpl.id || tpl.name || "") h = (h * 31 + ch.charCodeAt(0)) >>> 0
const [a, b] = palettes[h % palettes.length]
return `linear-gradient(135deg, ${a}, ${b})`
}, [tpl.id, tpl.name])
return (
<div
className="xx-cover-template-thumb"
style={{
background: url ? "#000" : placeholderBg,
position: "relative",
overflow: "hidden",
}}
>
{isSelected && <span className="xx-cover-template-check">✓</span>}
{url ? (
<img
src={url}
alt={tpl.name}
onError={() => setErrored(true)}
style={{
width: "100%",
height: "100%",
objectFit: "cover",
display: "block",
}}
/>
) : (
<span style={{ fontSize: 28, opacity: 0.5 }}>🖼️</span>
)}
</div>
)
const GRADIENT_MAP: Record<string, string> = {
default: "linear-gradient(135deg, #e0e0e0, #c0c0c0)",
"bold-red": "linear-gradient(135deg, #ef4444, #b91c1c)",
"elegant-black": "linear-gradient(135deg, #374151, #111827)",
"gradient-blue": "linear-gradient(135deg, #3b82f6, #1d4ed8)",
"gradient-purple": "linear-gradient(135deg, #8b5cf6, #6d28d9)",
"warm-orange": "linear-gradient(135deg, #f97316, #ea580c)",
"fresh-green": "linear-gradient(135deg, #22c55e, #15803d)",
"tech-blue": "linear-gradient(135deg, #06b6d4, #0e7490)",
}
const CoverSettingsModal: React.FC<CoverSettingsModalProps> = ({
@@ -140,7 +97,13 @@ const CoverSettingsModal: React.FC<CoverSettingsModalProps> = ({
className={`xx-cover-template-card${isSelected ? " selected" : ""}`}
onClick={() => onSelectTemplate(tpl.id)}
>
<TemplateThumb tpl={tpl} isSelected={isSelected} />
<div
className="xx-cover-template-thumb"
style={{ background: GRADIENT_MAP[tpl.id] || GRADIENT_MAP.default }}
>
{isSelected && <span className="xx-cover-template-check">✓</span>}
🖼️
</div>
<div className="xx-cover-template-info">
<div className="xx-cover-template-name">
{tpl.name}
@@ -153,7 +116,7 @@ const CoverSettingsModal: React.FC<CoverSettingsModalProps> = ({
onClick={() => onEditTemplate(tpl)}
title={tpl.is_system ? "基于此模板新建自定义模板" : "编辑模板"}
>
编辑
{tpl.is_system ? "复制" : "编辑"}
</Button>
{!tpl.is_system && (
<Button
@@ -68,96 +68,112 @@ const TitleMiniPreview: React.FC<Props> = ({
const text = (sampleText || settings.title || "预览标题").trim() || "预览标题"
useEffect(() => {
let cancelled = false
const draw = () => {
if (cancelled) return
const cvs = canvasRef.current
if (!cvs) return
const dpr = window.devicePixelRatio || 1
cvs.width = width * dpr
cvs.height = h * dpr
cvs.style.width = `${width}px`
cvs.style.height = `${h}px`
const ctx = cvs.getContext("2d")
if (!ctx) return
ctx.scale(dpr, dpr)
ctx.clearRect(0, 0, width, h)
const cvs = canvasRef.current
if (!cvs) return
const dpr = window.devicePixelRatio || 1
cvs.width = width * dpr
cvs.height = h * dpr
cvs.style.width = `${width}px`
cvs.style.height = `${h}px`
const ctx = cvs.getContext("2d")
if (!ctx) return
ctx.scale(dpr, dpr)
ctx.clearRect(0, 0, width, h)
// 背景(transparent 时跳过,用于叠加在图片上)
if (!transparent) {
ctx.fillStyle = "#111827"
ctx.fillRect(0, 0, width, h)
}
// 背景(transparent 时跳过,用于叠加在图片上)
if (!transparent) {
ctx.fillStyle = "#111827"
ctx.fillRect(0, 0, width, h)
}
// 分辨率缩放:以 360 宽为基准(对应 720p 的一半)
const scale = width / 360
const r = (v: number) => Math.round(v * scale)
// 分辨率缩放:以 360 宽为基准(对应 720p 的一半)
const scale = width / 360
const r = (v: number) => Math.round(v * scale)
// 字体
const size = r(settings.size)
const ff = getFontFamily(settings.font)
const parts: string[] = []
if (settings.italic) parts.push("italic")
if (settings.bold) parts.push("bold")
parts.push(`${size}px`, ff)
ctx.font = parts.join(" ")
ctx.textAlign = "center"
ctx.textBaseline = "middle"
ctx.fillStyle = settings.color
ctx.lineJoin = "round"
// 字体
const size = r(settings.size)
const ff = getFontFamily(settings.font)
const parts: string[] = []
if (settings.italic) parts.push("italic")
if (settings.bold) parts.push("bold")
parts.push(`${size}px`, ff)
ctx.font = parts.join(" ")
ctx.textAlign = "center"
ctx.textBaseline = "middle"
ctx.fillStyle = settings.color
ctx.lineJoin = "round"
// 阴影
const shadowEnabled = !!settings.shadow
const prevShadow = {
c: ctx.shadowColor,
b: ctx.shadowBlur,
ox: ctx.shadowOffsetX,
oy: ctx.shadowOffsetY,
// 阴影
const shadowEnabled = !!settings.shadow
const prevShadow = {
c: ctx.shadowColor,
b: ctx.shadowBlur,
ox: ctx.shadowOffsetX,
oy: ctx.shadowOffsetY,
}
if (shadowEnabled) {
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
}
// 换行
const lines = wrapLines(text, settings.maxCharsPerLine ?? 0)
const lineH = size * (settings.lineHeight ?? 1.2)
const totalH = lines.length * lineH
let startY: number
if (settings.position === "top") {
startY = size / 2 + r(settings.marginTop ?? 24)
} else if (settings.position === "center") {
startY = h / 2 - totalH / 2 + size / 2
} else {
// bottom
const botMargin = portrait ? r(24) : r(16)
startY = h - totalH - botMargin + size / 2
}
let centerX = width / 2
if (settings.position === "custom" && settings.posX != null) {
centerX = (settings.posX / 100) * width
}
// 背景块
if (settings.bgEnabled) {
const pad = r(settings.bgPadding ?? 12)
const rad = r(settings.bgRadius ?? 8)
let maxLineW = 0
for (const l of lines) {
const m = ctx.measureText(l)
if (m.width > maxLineW) maxLineW = m.width
}
const bw = maxLineW + pad * 2
const bh = totalH + pad * 2
const bx = centerX - bw / 2
const by = startY - size / 2 - pad + (size - lineH) / 2
ctx.shadowColor = "rgba(0,0,0,0)"
ctx.shadowBlur = 0
ctx.fillStyle = settings.bgColor ?? "rgba(0,0,0,0.5)"
roundRect(ctx, bx, by, bw, bh, rad)
ctx.fill()
// 恢复阴影
if (shadowEnabled) {
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
}
}
// 换行
const lines = wrapLines(text, settings.maxCharsPerLine ?? 0)
const lineH = size * (settings.lineHeight ?? 1.2)
const totalH = lines.length * lineH
let startY: number
if (settings.position === "top") {
startY = size / 2 + r(settings.marginTop ?? 24)
} else if (settings.position === "center") {
startY = h / 2 - totalH / 2 + size / 2
} else {
// bottom
const botMargin = portrait ? r(24) : r(16)
startY = h - totalH - botMargin + size / 2
}
let centerX = width / 2
if (settings.position === "custom" && settings.posX != null) {
centerX = (settings.posX / 100) * width
}
// 背景块
if (settings.bgEnabled) {
const pad = r(settings.bgPadding ?? 12)
const rad = r(settings.bgRadius ?? 8)
let maxLineW = 0
for (const l of lines) {
const m = ctx.measureText(l)
if (m.width > maxLineW) maxLineW = m.width
}
const bw = maxLineW + pad * 2
const bh = totalH + pad * 2
const bx = centerX - bw / 2
const by = startY - size / 2 - pad + (size - lineH) / 2
// 描边(先画,再画填充)
const strokeEnabled = !!settings.stroke && (settings.strokeWidth ?? 0) > 0
lines.forEach((line, i) => {
const y = startY + i * lineH
if (strokeEnabled) {
ctx.shadowColor = "rgba(0,0,0,0)"
ctx.shadowBlur = 0
ctx.fillStyle = settings.bgColor ?? "rgba(0,0,0,0.5)"
roundRect(ctx, bx, by, bw, bh, rad)
ctx.fill()
ctx.lineWidth = r(settings.strokeWidth ?? 4)
ctx.strokeStyle = settings.strokeColor ?? "#000000"
ctx.strokeText(line, centerX, y)
// 恢复阴影
if (shadowEnabled) {
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
@@ -166,46 +182,14 @@ const TitleMiniPreview: React.FC<Props> = ({
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
}
}
ctx.fillText(line, centerX, y)
})
// 描边(先画,再画填充)
const strokeEnabled = !!settings.stroke && (settings.strokeWidth ?? 0) > 0
lines.forEach((line, i) => {
const y = startY + i * lineH
if (strokeEnabled) {
ctx.shadowColor = "rgba(0,0,0,0)"
ctx.shadowBlur = 0
ctx.lineWidth = r(settings.strokeWidth ?? 4)
ctx.strokeStyle = settings.strokeColor ?? "#000000"
ctx.strokeText(line, centerX, y)
// 恢复阴影
if (shadowEnabled) {
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
}
}
ctx.fillText(line, centerX, y)
})
// 恢复
ctx.shadowColor = prevShadow.c
ctx.shadowBlur = prevShadow.b
ctx.shadowOffsetX = prevShadow.ox
ctx.shadowOffsetY = prevShadow.oy
}
// Web fonts (Google Fonts 等) 加载需要时间,等 fonts.ready 后再画,
// 避免第一次渲染用 sans-serif 画完再切字体造成「字体选择没反应」的错觉
if (typeof document !== "undefined" && document.fonts && document.fonts.ready) {
document.fonts.ready.then(() => {
if (cancelled) return
draw()
})
}
draw()
return () => {
cancelled = true
}
// 恢复
ctx.shadowColor = prevShadow.c
ctx.shadowBlur = prevShadow.b
ctx.shadowOffsetX = prevShadow.ox
ctx.shadowOffsetY = prevShadow.oy
}, [settings, width, h, text, transparent, portrait, background])
return (
+4 -18
View File
@@ -3540,19 +3540,10 @@
vertical-align: middle;
}
.xx-ce-font-dot--preset {
background: #10b981; /* 绿:预置爆款中文字体 */
}
.xx-ce-font-dot--hand {
background: #f59e0b; /* 橙:手写/书法字体 */
}
.xx-ce-font-dot--serif {
background: #8b5cf6; /* 紫:衬线字体 */
}
.xx-ce-font-dot--mono {
background: #6b7280; /* 灰:等宽字体 */
background: #10b981;
}
.xx-ce-font-dot--system {
background: #3b82f6; /* 蓝:系统无衬线 */
background: #3b82f6;
}
/* Shadow actions */
@@ -3789,12 +3780,7 @@
align-items: center;
}
/* Canvas 装饰层(背景/装饰/遮罩/底色/人物/文字背景色块)不接收鼠标事件,
但拖拽的标题/副标题文字(内联 cursor:grab)需要接收 mousedown。
已通过 renderTextStyle 显式设 pointer-events 以外的样式,因此此处只关掉纯装饰层。 */
.xx-ce-canvas-base,
.xx-ce-el-bg,
.xx-ce-el-portrait,
.xx-ce-el-mask {
/* Canvas text elements — ensure proper stacking */
.xx-ce-canvas > div {
pointer-events: none;
}
@@ -107,57 +107,54 @@ export function useBatchCovers({
addBusy(index)
try {
const titleText = titles[index] || ""
const response = await generateCover(
selectedTemplate && selectedTemplate !== "default" ? selectedTemplate : undefined,
{
generated_video_id: target.id,
video_url: target.file_url || target.download_url || "",
cover_type: "ai_frame",
...(titleText
? {
title_config: {
text: titleText,
font: titleStyle.font,
font_size: titleStyle.size,
font_color: titleStyle.color,
position: titleStyle.position,
bold: titleStyle.bold,
italic: titleStyle.italic,
stroke: titleStyle.stroke
? {
enabled: true,
width: titleStyle.strokeWidth ?? 4,
color: titleStyle.strokeColor ?? "#000000",
}
: { enabled: false },
shadow: titleStyle.shadow
? {
enabled: true,
offset_x: titleStyle.shadowOffsetX ?? 2,
offset_y: titleStyle.shadowOffsetY ?? 2,
blur: titleStyle.shadowBlur ?? 4,
color: titleStyle.shadowColor ?? "rgba(0,0,0,0.8)",
}
: { enabled: false },
line_height: titleStyle.lineHeight ?? 1.2,
margin_top: titleStyle.marginTop ?? 24,
max_chars_per_line: titleStyle.maxCharsPerLine ?? 0,
background: titleStyle.bgEnabled
? {
enabled: true,
color: titleStyle.bgColor,
padding: titleStyle.bgPadding,
radius: titleStyle.bgRadius,
}
: { enabled: false },
line_overrides: (titleStyle.lineOverrides ?? []) as Array<
Record<string, unknown>
>,
},
}
: {}),
},
)
const response = await generateCover(selectedTemplate || "default", {
generated_video_id: target.id,
video_url: target.file_url || target.download_url || "",
cover_type: "ai_frame",
...(titleText
? {
title_config: {
text: titleText,
font: titleStyle.font,
font_size: titleStyle.size,
font_color: titleStyle.color,
position: titleStyle.position,
bold: titleStyle.bold,
italic: titleStyle.italic,
stroke: titleStyle.stroke
? {
enabled: true,
width: titleStyle.strokeWidth ?? 4,
color: titleStyle.strokeColor ?? "#000000",
}
: { enabled: false },
shadow: titleStyle.shadow
? {
enabled: true,
offset_x: titleStyle.shadowOffsetX ?? 2,
offset_y: titleStyle.shadowOffsetY ?? 2,
blur: titleStyle.shadowBlur ?? 4,
color: titleStyle.shadowColor ?? "rgba(0,0,0,0.8)",
}
: { enabled: false },
line_height: titleStyle.lineHeight ?? 1.2,
margin_top: titleStyle.marginTop ?? 24,
max_chars_per_line: titleStyle.maxCharsPerLine ?? 0,
background: titleStyle.bgEnabled
? {
enabled: true,
color: titleStyle.bgColor,
padding: titleStyle.bgPadding,
radius: titleStyle.bgRadius,
}
: { enabled: false },
line_overrides: (titleStyle.lineOverrides ?? []) as Array<
Record<string, unknown>
>,
},
}
: {}),
})
const url = response.cover?.image_url || response.cover?.thumbnail_url || ""
if (url) {
patchCover(index, url)
@@ -22,8 +22,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
const [generated, setGenerated] = useState(false)
const [generateError, setGenerateError] = useState<string | null>(null)
const [generatedVideos, setGeneratedVideos] = useState<GeneratedVideo[]>([])
/** 单视频模式:当前任务 ID(封面 finalize 需要) */
const [currentTaskId, setCurrentTaskId] = useState<string>("")
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
@@ -122,8 +120,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
setGenerated(false)
setGenerateError(null)
setBatchTasks([])
setGeneratedVideos([])
setCurrentTaskId("")
clearTimer()
try {
@@ -345,10 +341,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
}
if (taskIds.length > 1) {
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
setCurrentTaskId("")
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
} else {
setCurrentTaskId(taskIds[0])
startPolling(taskIds[0])
}
} catch (err) {
@@ -423,7 +417,6 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
generated,
generateError,
generatedVideos,
currentTaskId,
generate,
retry,
retryBatchTask,
@@ -4,7 +4,7 @@
*
* - 步骤1(选择模式):下一步分支由外层弹窗处理(VoiceSelectModal / ScriptSelectModal),
* 本 hook 的 goNext 仅在未选模式时拦截;外层 Modal onConfirm 里主动 setCurrentStep(2)。
* - 步骤2(选择素材):直接进入步骤3,数组长度对齐由 onBeforeEnterStep3 保证。
* - 步骤2(选择素材):弹数量选择弹窗(PreviewCountModal),确认后跳步骤3。
* - 步骤3 底部按钮是「确认生成视频」(由 GenerateStepActions 调 onConfirmGenerate),
* 创建成功后跳步骤4;本 hook 的 goNext 只负责 2→3 和 4→5 的「下一步」。
* - 步骤4(确认生成进度页):全部渲染完成后「下一步」解锁进封面。
@@ -23,10 +23,10 @@ export interface UseStepNavigationOptions {
titleSettings: TitleSettings
/** 是否已完成视频生成(步骤4全部渲染完成后才能进入封面) */
generated: boolean
/** 点素材下一步时弹出数量选择弹窗 */
onOpenCountModal: () => void
/** 步骤1下一步:根据 editMode 打开对应弹窗(随机→配音 / 叙事→文案) */
onOpenStep1Modal: () => void
/** 进入步骤3前自动对齐数组(previewTitles/voiceLibraryIds/previewCovers/selectedVariantIds)长度到 previewCount */
onBeforeEnterStep3?: () => void
}
export interface UseStepNavigationReturn {
@@ -42,8 +42,8 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
selectedMaterials,
smartSelectedIds,
generated,
onOpenCountModal,
onOpenStep1Modal,
onBeforeEnterStep3,
} = options
const goNext = () => {
@@ -62,9 +62,8 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
message.warning("请先进行智能匹配并选择素材")
return
}
// 直接进入步骤3(生成数量在 Step1 已设置);对齐数组长度
onBeforeEnterStep3?.()
setCurrentStep(3)
// 弹数量选择弹窗
onOpenCountModal()
return
}
// 步骤4(确认生成):全部渲染完成后才能下一步进封面
+47 -85
View File
@@ -63,8 +63,9 @@ export interface TextBackground {
shape: TextBgShape
width: number
height: number
/** 相对文字的上下偏移(百分比),背景自动跟随文字位置 */
offsetY: number
posX: number
posY: number
rotation: number
}
/** 文字样式配置(主标题/副标题共用) */
@@ -153,7 +154,9 @@ export const DEFAULT_TITLE_CONFIG: TextStyleConfig = {
shape: "polygon",
width: 30,
height: 10,
offsetY: 0,
posX: 50,
posY: 50,
rotation: 0,
},
}
@@ -181,7 +184,9 @@ export const DEFAULT_SUBTITLE_CONFIG: TextStyleConfig = {
shape: "rectangle",
width: 100,
height: 20,
offsetY: 8,
posX: 50,
posY: 80,
rotation: 0,
},
}
@@ -217,89 +222,46 @@ export const DEFAULT_EDITOR_CONFIG: CoverEditorConfig = {
maskShape: "矩形",
}
/** 预置字体(已与 @/components/title/constants 字体表保持一致;自定义商业字体兜底 Google Fonts 开源中文字体) */
// 封面编辑器预置字体:与标题样式字体列表保持一致(从 @/components/title/constants 同步),
// 并补全西文常用系统字体,保证在中英文环境下都有可用字体。
// 注:需要配合 index.html 引入的 Google Fonts(Noto Sans SC / ZCOOL / Ma Shan Zheng 等)。
export interface CoverFont {
name: string
family: string
tag?: "preset" | "hand" | "serif" | "sans" | "mono"
}
/** 预置中文字体(爆款/常用) */
export const PRESET_FONTS: CoverFont[] = [
{
name: "优设标题黑",
family:
'"YouSheBiaoTiHei","ZCOOL QingKe HuangYou","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
tag: "preset",
},
{
name: "阿里普惠体Bold",
family:
'"Alibaba PuHuiTi","Alibaba Sans","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
tag: "preset",
},
{
name: "抖音美好体",
family:
'"Douyin Sans","ZCOOL KuaiLe","Noto Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
tag: "preset",
},
{
name: "思源黑体Heavy",
family: '"Noto Sans SC","Source Han Sans SC Heavy","PingFang SC","Microsoft YaHei",sans-serif',
tag: "preset",
},
{
name: "思源黑体",
family: '"Noto Sans SC","Source Han Sans SC","PingFang SC","Microsoft YaHei",sans-serif',
tag: "preset",
},
{
name: "思源宋体",
family: '"Noto Serif SC","Source Han Serif SC","Songti SC","SimSun",serif',
tag: "serif",
},
{ name: "站酷小薇体", family: '"ZCOOL XiaoWei","Noto Serif SC",serif', tag: "preset" },
{ name: "马善政毛笔", family: '"Ma Shan Zheng","STXingkai","KaiTi",cursive', tag: "hand" },
{ name: "龙藏体", family: '"Long Cang","STXingkai",cursive', tag: "hand" },
{ name: "楷体", family: '"KaiTi","STKaiti","DFKai-SB",serif', tag: "serif" },
{
name: "苹方",
family: '"PingFang SC",-apple-system,"Helvetica Neue",sans-serif',
tag: "sans",
},
{
name: "微软雅黑",
family: '"Microsoft YaHei","PingFang SC","Noto Sans SC",sans-serif',
tag: "sans",
},
/** 预置字体 */
export const PRESET_FONTS = [
{ name: "思源黑体", family: "'Noto Sans SC', sans-serif" },
{ name: "斗鱼追光体2.0", family: "'DouYu ZhuangGuangTi', sans-serif" },
{ name: "抖音美好体", family: "'DouYin MeiHaoTi', sans-serif" },
]
/** 系统字体(西文 + 通用中文) */
export const SYSTEM_FONTS: CoverFont[] = [
{ name: "Arial", family: "Arial, Helvetica, sans-serif", tag: "sans" },
{ name: "Helvetica", family: "Helvetica, Arial, sans-serif", tag: "sans" },
{ name: "Times New Roman", family: '"Times New Roman", Times, serif', tag: "serif" },
{ name: "Georgia", family: "Georgia, serif", tag: "serif" },
{ name: "Verdana", family: "Verdana, Geneva, sans-serif", tag: "sans" },
{ name: "Tahoma", family: "Tahoma, Geneva, sans-serif", tag: "sans" },
{ name: "Impact", family: 'Impact, "Arial Black", sans-serif', tag: "sans" },
{ name: "Comic Sans MS", family: '"Comic Sans MS", cursive', tag: "hand" },
{ name: "Courier New", family: '"Courier New", Courier, monospace', tag: "mono" },
{ name: "宋体", family: "SimSun, 'Noto Serif SC', serif", tag: "serif" },
{ name: "黑体", family: "SimHei, 'Noto Sans SC', sans-serif", tag: "sans" },
{ name: "仿宋", family: "FangSong, 'Noto Serif SC', serif", tag: "serif" },
{ name: "Trebuchet MS", family: '"Trebuchet MS", sans-serif', tag: "sans" },
{ name: "Lucida Console", family: '"Lucida Console", Monaco, monospace', tag: "mono" },
{ name: "Palatino", family: 'Palatino, "Palatino Linotype", serif', tag: "serif" },
{ name: "Garamond", family: "Garamond, serif", tag: "serif" },
{ name: "Calibri", family: "Calibri, sans-serif", tag: "sans" },
{ name: "Cambria", family: "Cambria, serif", tag: "serif" },
{ name: "Candara", family: "Candara, sans-serif", tag: "sans" },
{ name: "Consolas", family: "Consolas, monospace", tag: "mono" },
/** 系统字体 */
export const SYSTEM_FONTS = [
{ name: "Arial", family: "Arial, sans-serif" },
{ name: "Helvetica", family: "Helvetica, sans-serif" },
{ name: "Times New Roman", family: "'Times New Roman', serif" },
{ name: "Georgia", family: "Georgia, serif" },
{ name: "Verdana", family: "Verdana, sans-serif" },
{ name: "Tahoma", family: "Tahoma, sans-serif" },
{ name: "Impact", family: "Impact, sans-serif" },
{ name: "Comic Sans MS", family: "'Comic Sans MS', cursive" },
{ name: "Courier New", family: "'Courier New', monospace" },
{ name: "微软雅黑", family: "'Microsoft YaHei', sans-serif" },
{ name: "宋体", family: "SimSun, serif" },
{ name: "黑体", family: "SimHei, sans-serif" },
{ name: "楷体", family: "KaiTi, serif" },
{ name: "仿宋", family: "FangSong, serif" },
{ name: "Trebuchet MS", family: "'Trebuchet MS', sans-serif" },
{ name: "Lucida Console", family: "'Lucida Console', monospace" },
{ name: "Palatino", family: "Palatino, serif" },
{ name: "Garamond", family: "Garamond, serif" },
{ name: "Bookman", family: "Bookman, serif" },
{ name: "Avant Garde", family: "'Avant Garde', sans-serif" },
{ name: "Calibri", family: "Calibri, sans-serif" },
{ name: "Cambria", family: "Cambria, serif" },
{ name: "Candara", family: "Candara, sans-serif" },
{ name: "Consolas", family: "Consolas, monospace" },
{ name: "Constantia", family: "Constantia, serif" },
{ name: "Corbel", family: "Corbel, sans-serif" },
{ name: "Franklin Gothic", family: "'Franklin Gothic', sans-serif" },
{ name: "Gill Sans", family: "'Gill Sans', sans-serif" },
{ name: "Optima", family: "Optima, sans-serif" },
{ name: "Futura", family: "Futura, sans-serif" },
{ name: "Rockwell", family: "Rockwell, serif" },
]
/** 所有字体列表 */
@@ -20,6 +20,7 @@ import "@/pages/generate/components/Step4TitleSettings"
import "@/pages/generate/components/Step5VoiceSelect"
import "@/pages/generate/components/Step3VoiceWithMode"
import "@/pages/generate/components/BatchGenerationGrid"
import "@/pages/generate/components/PreviewCountModal"
import "@/pages/generate/components/GenerateStepContent"
import "@/pages/generate/components/voice/VoiceRecommendSection"
import "@/pages/generate/components/voice/VoiceChoiceCard"
@@ -63,7 +63,6 @@ from packages.domain.render_layer_utils import clip_playback_speed as _clip_play
from packages.domain.render_layer_utils import estimate_total_duration as _estimate_total_duration_pure
from packages.domain.render_layer_utils import resolve_layer_role as _resolve_layer_role_pure
from packages.domain.tts_config import TtsConfig
from packages.shared.gpu_encoder import GpuEncodeError, get_gpu_encoder
logger = logging.getLogger(__name__)
@@ -1622,27 +1621,20 @@ class UnifiedRenderService:
effective_duration,
has_audio,
)
# 尝试 GPU NVENC 加速
gpu_ok = False
if self._gpu_encode_available():
mezz_path = output_path.parent / f".{output_path.stem}.mezz{output_path.suffix}"
gpu_ok = self._ffmpeg_output_to_mezzanine(command, mezz_path, output_path)
if not gpu_ok:
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
stderr_text = (e.stderr or "").strip()
stderr_tail = stderr_text[-1500:] if len(stderr_text) > 1500 else stderr_text
logger.error(
"直通渲染失败: plan_id=%s clip=%s exit_code=%d\nvf=%s\nstderr(last 1500):\n%s",
self.plan.id,
clip.clip_id,
e.returncode,
vf_str[:2000],
stderr_tail,
)
raise
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
stderr_text = (e.stderr or "").strip()
stderr_tail = stderr_text[-1500:] if len(stderr_text) > 1500 else stderr_text
logger.error(
"直通渲染失败: plan_id=%s clip=%s exit_code=%d\nvf=%s\nstderr(last 1500):\n%s",
self.plan.id,
clip.clip_id,
e.returncode,
vf_str[:2000],
stderr_tail,
)
raise
return has_audio
@@ -2178,124 +2170,6 @@ class UnifiedRenderService:
filter_complex = ";".join(filter_parts)
return filter_complex, input_args
# ── GPU NVENC 加速 ────────────────────────────────────────────────────
def _gpu_encode_available(self) -> bool:
"""GPU 编码客户端是否已配置且健康(缓存健康状态,单任务内只探测一次)。"""
if not getattr(self, "_gpu_health_ok", None):
client = get_gpu_encoder()
if client is None:
self._gpu_health_ok = False
return False
try:
health = client.check_health()
if health.ready:
logger.info(
"[gpu-encoder] healthy endpoint=%s gpu=%s",
client.endpoint,
health.gpu_name,
)
self._gpu_health_ok = True
else:
logger.warning(
"[gpu-encoder] not ready: %s (endpoint=%s)",
health.error,
client.endpoint,
)
self._gpu_health_ok = False
except Exception as e: # noqa: BLE001
logger.warning("[gpu-encoder] health probe error (CPU fallback): %s", e)
self._gpu_health_ok = False
return self._gpu_health_ok
def _ffmpeg_output_to_mezzanine(
self,
base_command: list[str],
mezzanine_path: Path,
output_path: Path,
) -> bool:
"""用 CPU ultrafast 把滤镜链输出到 mezzanine_path,然后调 GPU 做最终编码。
base_command: 原本要执行的完整 ffmpeg 命令(含 -c:v libx264 -crf X -preset Y ... output_path)
我们把最后一个参数(output_path)替换成 mezzanine_path,并把编码参数改成 ultrafast,
成功后调用 gpu_encoder 做 nvenc 编码到 output_path。
任何失败返回 False,调用方走原始 CPU 路径。
"""
client = get_gpu_encoder()
if client is None:
return False
# 构造 mezzanine 命令:替换编码参数和输出路径
mezz_cmd = list(base_command)
# 找到编码参数位置并替换
try:
i_crf = mezz_cmd.index("-crf")
mezz_cmd[i_crf + 1] = "20"
i_preset = mezz_cmd.index("-preset")
mezz_cmd[i_preset + 1] = "ultrafast"
except ValueError:
logger.warning("[gpu-encoder] could not find -crf/-preset in command, skip gpu")
return False
# 如果命令有音频编码 -c:a aac,我们保留音频让 GPU 侧不用单独处理
# (P4000 的 ffmpeg_args 可以直接 copy 音频?这里简单起见:把音频编码留在 mezzanine,
# 然后 GPU 侧直接 -c:a copy,避免重编码损失)
has_audio = "-c:a" in mezz_cmd
# 替换输出路径(最后一个参数)
mezz_cmd[-1] = str(mezzanine_path)
# 1) 跑 mezzanine
mezzanine_path.parent.mkdir(parents=True, exist_ok=True)
t0 = time.time()
try:
run_ffmpeg(mezz_cmd)
except subprocess.CalledProcessError as e:
logger.warning("[gpu-encoder] mezzanine encode failed (CPU fallback): %s", e)
return False
logger.info(
"[gpu-encoder] mezzanine ready: %s (%.1fs, %d bytes), dispatching to P4000 nvenc...",
mezzanine_path.name,
time.time() - t0,
mezzanine_path.stat().st_size if mezzanine_path.exists() else 0,
)
# 2) GPU nvenc encode(含上传 mezzanine → OSS → P4000 下载+编码 → relay 回传)
try:
# GPU 侧:-i in.mp4 -c:v h264_nvenc ... 音频 copy(mezzanine 里音频已是 aac)
audio_args = ["-c:a", "copy"] if has_audio else None
client.encode_mezzanine_to_output(
mezzanine_path,
output_path,
audio_args=audio_args,
)
logger.info(
"[gpu-encoder] GPU nvenc encode done: %s (total %.1fs)",
output_path.name,
time.time() - t0,
)
return True
except GpuEncodeError as e:
logger.warning("[gpu-encoder] GPU encode failed (CPU fallback): %s", e)
# 删除可能残留的不完整 output
try:
if output_path.exists():
output_path.unlink()
except OSError:
pass
return False
except Exception as e: # noqa: BLE001
logger.warning("[gpu-encoder] GPU encode unexpected error (CPU fallback): %s", e)
return False
finally:
# 清理 mezzanine
try:
if mezzanine_path.exists():
mezzanine_path.unlink()
except OSError:
pass
def _execute_ffmpeg(
self,
filter_complex: str,
@@ -2335,28 +2209,20 @@ class UnifiedRenderService:
input_args.count("-i"),
output_path,
)
# 尝试 GPU NVENC 加速:先出 ultrafast mezzanine,再交给 P4000 做最终编码
gpu_ok = False
if self._gpu_encode_available():
mezz_path = output_path.parent / f".{output_path.stem}.mezz{output_path.suffix}"
gpu_ok = self._ffmpeg_output_to_mezzanine(command, mezz_path, output_path)
if not gpu_ok:
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
# 额外记录 filter_complex + stderr,方便排查滤镜链构建问题
stderr_text = (e.stderr or "").strip()
stderr_tail = stderr_text[-1500:] if len(stderr_text) > 1500 else stderr_text
logger.error(
"渲染失败: plan_id=%s exit_code=%d\nfilter_complex:\n%s\nstderr(last 1500):\n%s",
self.plan.id,
e.returncode,
filter_complex[:5000],
stderr_tail,
)
raise
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
# 额外记录 filter_complex + stderr,方便排查滤镜链构建问题
stderr_text = (e.stderr or "").strip()
stderr_tail = stderr_text[-1500:] if len(stderr_text) > 1500 else stderr_text
logger.error(
"渲染失败: plan_id=%s exit_code=%d\nfilter_complex:\n%s\nstderr(last 1500):\n%s",
self.plan.id,
e.returncode,
filter_complex[:5000],
stderr_tail,
)
raise
def _build_sticker_filters(self, input_label: str, output_label: str) -> tuple[str, list[str]]:
"""构建贴纸叠加滤镜链.
-6
View File
@@ -76,10 +76,4 @@ celery_app.conf.beat_schedule = {
"schedule": 600.0, # 每 10 分钟(秒)
"options": {"expires": 540},
},
# 音色克隆卡死巡检:worker 重启/消息丢失后 processing 卡 10 分钟标 failed,用户可点重试
"cleanup-stale-voice-clones": {
"task": "worker.cleanup_stale_voice_clones",
"schedule": 300.0, # 每 5 分钟
"options": {"expires": 240},
},
}
-77
View File
@@ -287,80 +287,3 @@ def _recover_stuck_ingest_jobs_on_ready(sender, **kwargs): # pragma: no cover
logger.info("Worker 启动 ingest 恢复完成,共重新派单 %d 个卡死任务", recovered)
except Exception as e: # noqa: BLE001 — 启动恢复失败不能阻断 worker 起服
logger.error("启动 ingest 恢复扫描失败(beat 巡检仍会兜底标 failed): %s", e, exc_info=True)
def recover_stale_voice_clones_on_startup(timeout_minutes: int = 10) -> int:
"""Worker 启动时恢复卡死在 processing 的音色克隆任务。
容器重启/进程 OOM 时 worker 中正在轮询的克隆任务会丢失,
voice_clone_profiles 永久卡在 processing 无兜底。启动时扫描
updated_at 超过 timeout_minutes 的 processing 记录,直接标记
为 failed(错误信息指引用户重试)。选择标 failed 而非重新派单,
因为 CosyVoice 侧的 voice_id 无法在无上下文下恢复轮询,重试需
用户确认后显式触发。
Args:
timeout_minutes: 判定卡死的阈值,默认 10 分钟
Returns:
恢复的记录数
"""
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
)
try:
session = SessionLocal()
try:
repo = SQLAlchemyVoiceCloneProfileRepository(session)
count = repo.cleanup_stale_processing(timeout_minutes)
finally:
session.close()
if count > 0:
logger.warning("启动时恢复了 %d 个卡死在 processing 的音色克隆(超时 %d 分钟)", count, timeout_minutes)
else:
logger.info("无卡死 processing 音色克隆需要恢复")
return count
except Exception as e:
logger.error("启动时音色克隆恢复扫描失败(beat 巡检仍会兜底): %s", e, exc_info=True)
return 0
@worker_ready.connect
def _recover_stuck_voice_clones_on_ready(sender, **kwargs):
"""Worker 启动完成后恢复卡死的音色克隆任务。"""
try:
recovered = recover_stale_voice_clones_on_startup()
logger.info("Worker 启动音色克隆恢复完成,共标记 %d 个卡死任务为 failed", recovered)
except Exception as e:
logger.error("启动音色克隆恢复失败(beat 巡检仍会兜底标 failed): %s", e, exc_info=True)
@worker_ready.connect
def _probe_gpu_encoder_on_ready(sender, **kwargs):
"""Worker 启动完成后探测 P4000 GPU NVENC 节点状态,打日志。"""
try:
from packages.shared.gpu_encoder import get_gpu_encoder
client = get_gpu_encoder()
if client is None:
logger.info(
"[gpu-encoder] disabled (ENABLE_GPU_ENCODE=false or endpoint not configured), using CPU libx264"
)
return
health = client.check_health()
if health.ready:
logger.info(
"[gpu-encoder] NVENC enabled: endpoint=%s gpu=%s worker=%s",
client.endpoint,
health.gpu_name,
health.worker,
)
else:
logger.warning(
"[gpu-encoder] configured but NOT ready: %s (endpoint=%s) — falling back to CPU",
health.error,
client.endpoint,
)
except Exception as e: # noqa: BLE001
logger.warning("[gpu-encoder] startup probe error (will retry on first job, CPU fallback): %s", e)
-41
View File
@@ -22,10 +22,6 @@ from packages.application.ingest_orphan_cleanup import (
INGEST_PROCESSING_TIMEOUT_MINUTES,
)
# 音色克隆 processing 超时:正常克隆轮询最多 5 分钟,10 分钟无更新视为卡死
VOICE_CLONE_PROCESSING_TIMEOUT_MINUTES = 10
logger = logging.getLogger(__name__)
@@ -129,40 +125,3 @@ def scheduled_cleanup_stale_ingest_jobs(
purged,
)
return {"stale_jobs": total_jobs, "assets_to_error": total_assets, "purged_messages": purged}
@shared_task(name="worker.cleanup_stale_voice_clones")
def scheduled_cleanup_stale_voice_clones(
processing_timeout_minutes: int = VOICE_CLONE_PROCESSING_TIMEOUT_MINUTES,
) -> dict:
"""Celery Beat: 清理卡死在 processing 的音色克隆档案。
每 5 分钟执行一次。worker 重启/Celery 消息丢失/进程 OOM 时,
已 prefetch 的克隆任务消息丢失,voice_clone_profile 永久卡在 processing。
超过 processing_timeout_minutes 未更新的记录标记为 failed,
错误信息指引用户点击重试。
"""
from worker_app.db import SessionLocal
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
)
session = None
try:
session = SessionLocal()
repo = SQLAlchemyVoiceCloneProfileRepository(session)
count = repo.cleanup_stale_processing(processing_timeout_minutes)
if count > 0:
logger.warning(
"[Beat] 清理了 %d 个卡死 processing 的音色克隆(超时 %d 分钟)",
count,
processing_timeout_minutes,
)
return {"cleaned": count}
except Exception as e:
logger.error("[Beat] 清理卡死音色克隆失败: %s", e, exc_info=True)
return {"cleaned": 0, "error": str(e)}
finally:
if session is not None:
session.close()
@@ -44,7 +44,6 @@ def process_voice_clone(self: Task, profile_id: str) -> dict:
# P2-2 修复:session 初始化为 None,避免 SessionLocal() 抛异常时
# finally 块中 session.close() 触发 UnboundLocalError
session = None
logger.info(f"Voice clone task started: profile_id={profile_id}")
try:
session = SessionLocal()
repo = SQLAlchemyVoiceCloneProfileRepository(session)
-31
View File
@@ -1,31 +0,0 @@
# Staging GPU relay plain-HTTP vhost (P4000 NVENC 编码回传入口)
# - 监听 8092 端口纯 HTTP(绕开 HTTPS 证书与 P4000 httpx SSL 问题)
# - 代理到本机 staging API 的 /api/ 路径(127.0.0.1:8000 是 docker 映射端口)
# - P4000 通过 Tailscale 直连宿主机 100.69.73.60:8092 PUT 编码结果
# - Worker 通过 Docker DNS (xiaoxia-api-staging:8000) 直接 GET/DELETE,
# 不经宿主机 nginx,避免 UFW FORWARD DROP 阻断
#
# 部署:cp infra/nginx/gpu-relay-staging.conf /etc/nginx/conf.d/ && nginx -t && systemctl reload nginx
server {
listen 8092;
server_name _;
client_max_body_size 2048m;
location /api/ {
proxy_pass http://127.0.0.1:8000/api/;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_request_buffering off;
proxy_read_timeout 600s;
proxy_send_timeout 600s;
}
location = /health {
proxy_pass http://127.0.0.1:8000/health;
}
}
@@ -136,39 +136,6 @@ class SQLAlchemyVoiceCloneProfileRepository:
)
return {voice_id: profile_id for voice_id, profile_id in rows}
def cleanup_stale_processing(self, timeout_minutes: int = 10) -> int:
"""清理超时卡在 processing 的克隆档案。
worker 重启、Celery 任务丢失或 OOM 被杀时,processing 档案会永久卡住。
updated_at < NOW() - timeout_minutes 的 processing 记录,标记为 failed
并附带明确错误信息,用户可在前端点击「重试」。
Args:
timeout_minutes: 超时分钟数,默认 10 分钟(正常克隆 < 5 分钟)
Returns:
清理的记录数
"""
from datetime import UTC, datetime, timedelta
cutoff = datetime.now(UTC) - timedelta(minutes=timeout_minutes)
models = (
self.session.query(VoiceCloneProfileModel)
.filter(
VoiceCloneProfileModel.status == "processing",
VoiceCloneProfileModel.updated_at < cutoff,
)
.all()
)
count = 0
for model in models:
model.status = "failed"
model.error_message = f"克隆任务执行超时(超过 {timeout_minutes} 分钟未更新,可能因服务重启中断),请重试"
count += 1
if count > 0:
self.session.commit()
return count
@staticmethod
def _model_to_entity(model: VoiceCloneProfileModel) -> VoiceCloneProfile:
return VoiceCloneProfile(
-61
View File
@@ -129,67 +129,6 @@ class SharedSettings(BaseSettings):
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
gpu_worker_stale_seconds: int = 300
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
enable_gpu_encode: bool = Field(
default=False,
validation_alias=AliasChoices("ENABLE_GPU_ENCODE", "enable_gpu_encode"),
)
# P4000 编码节点地址(Tailscale 内网),例如 http://100.105.75.67:8900
gpu_encode_endpoint: str = Field(
default="",
validation_alias=AliasChoices("GPU_ENCODE_ENDPOINT", "gpu_encode_endpoint"),
)
# GPU 回传临时文件走公网/内网 nginx(/gpu-relay/ 已加 location);
# 形如 http://100.69.73.60/gpu-relay (不带尾斜杠)
gpu_encode_relay_base_url: str = Field(
default="",
validation_alias=AliasChoices("GPU_ENCODE_RELAY_BASE_URL", "gpu_encode_relay_base_url"),
description="P4000 回传结果用的外部 URL(worker 通过该 URL 提供给 P4000 PUT),如 http://100.69.73.60:8092",
)
# Worker→API 内网直连 URL(Docker DNS),用于 worker 自己下载/清理 relay 文件。
# 未配置时回退到 relay_base_url(本地开发/单节点)。
gpu_encode_relay_internal_base_url: str = Field(
default="",
validation_alias=AliasChoices("GPU_ENCODE_RELAY_INTERNAL_BASE_URL", "gpu_encode_relay_internal_base_url"),
)
# 同步调用超时(秒):含编码+上传回传,5 分钟足够短视频
gpu_encode_sync_timeout: int = 300
# 异步轮询总超时(秒):长视频走 async + 轮询
gpu_encode_async_timeout: int = 1800
# 轮询间隔(秒)
gpu_encode_poll_interval: float = 3.0
# 启动探测超时(秒)
gpu_encode_health_timeout: float = 3.0
# NVENC 默认编码参数(可被调用方覆盖)
gpu_encode_vcodec: str = "h264_nvenc"
gpu_encode_preset: str = "p4" # NVENC preset: p1(最快)~p7(最好),p4 为均衡
gpu_encode_crf: int = 23
gpu_encode_bitrate: str = "" # 空则用 crf;非空则用 -b:v 模式
# GPU 编码失败时是否自动降级到 CPU(默认 True);设为 False 可在 CI/测试中暴露错误
gpu_encode_fallback_cpu: bool = Field(
default=True,
validation_alias=AliasChoices("GPU_ENCODE_FALLBACK_CPU", "gpu_encode_fallback_cpu"),
)
# P4000 → relay 回传鉴权 token(query 参数 token=xxx)。
# 生产环境必须设置;未设置且非 production 时自动生成随机值(写日志方便排查)。
gpu_encode_relay_secret: str = Field(
default="",
validation_alias=AliasChoices("GPU_ENCODE_RELAY_SECRET", "gpu_encode_relay_secret"),
)
# GPU 中间片在 OSS 的临时前缀(worker 上传 mezzanine 供 P4000 下载)
gpu_encode_oss_tmp_prefix: str = Field(
default="tmp/gpu-mezzanine/",
validation_alias=AliasChoices("GPU_ENCODE_OSS_TMP_PREFIX", "gpu_encode_oss_tmp_prefix"),
)
# relay 写入目录(相对于 generated-files 根目录)
gpu_encode_relay_dir: str = Field(
default="gpu_relay",
validation_alias=AliasChoices("GPU_ENCODE_RELAY_DIR", "gpu_encode_relay_dir"),
)
# relay 文件保留时间(秒),worker 下载完成后会主动删除,此为兜底清理 TTL
gpu_encode_relay_ttl: int = 3600
@property
def effective_database_url(self) -> str:
"""返回实际使用的数据库 URL。
-400
View File
@@ -1,400 +0,0 @@
"""P4000 NVENC 远程编码客户端。
完整链路(encode_video_file):
1. CPU 滤镜已在本地生成 mezzanine 中间片(libx264 ultrafast)
2. 上传 mezzanine 到 OSS 临时前缀,拿到签名 GET URL
3. 生成 relay 一次性 key,构造两个带 token 的 URL:
- put_url:给 P4000 回传结果,走 relay_base_url(外部可达,通常是 host:port 经 nginx)
- get/del_url:worker 自己下载+清理用,走 relay_internal_base_url(Docker DNS 直连 API)
4. POST P4000 /api/render/sync:inputs={"in.mp4": "<oss-signed-url>"}, output_url="<put_url>"
ffmpeg_args: -i in.mp4 [-vf <vf>] -c:v h264_nvenc ... -an/-c:a aac -f mp4 pipe:1
5. P4000 编码完成后 PUT 最终 mp4 到 put_url,API 服务落盘到 /app/generated/gpu_relay/<key>
6. 本客户端通过 get_url(Docker 内网)下载最终文件到 output_path,然后 DELETE 清理
7. 删除 OSS 临时 mezzanine
任何环节失败抛 GpuEncodeError,调用方应 fallback 到 CPU libx264。
"""
from __future__ import annotations
import json
import logging
import os
import socket
import time
import urllib.error
import urllib.parse
import urllib.request
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional
logger = logging.getLogger(__name__)
class GpuEncodeError(RuntimeError):
"""GPU 编码失败(网络/超时/ffmpeg/upload/download 任一环节)。调用方应 fallback 到 CPU。"""
@dataclass
class GpuHealth:
healthy: bool
worker: str = ""
gpu_name: str = ""
nvenc_h264: bool = False
nvenc_hevc: bool = False
error: str = ""
@property
def ready(self) -> bool:
return self.healthy and self.nvenc_h264
class GpuEncoderClient:
def __init__(
self,
endpoint: str,
relay_base_url: str,
*,
relay_internal_base_url: str = "",
sync_timeout: int = 300,
health_timeout: float = 3.0,
vcodec: str = "h264_nvenc",
preset: str = "p4",
crf: int = 23,
bitrate: str = "",
relay_secret: str = "",
oss_tmp_prefix: str = "tmp/gpu-mezzanine/",
) -> None:
self.endpoint = endpoint.rstrip("/")
self.relay_base_url = relay_base_url.rstrip("/")
# Worker→API 内网访问地址(Docker DNS 直连,如 http://xiaoxia-api-staging:8000)。
# 未配置时回退到 relay_base_url(本地开发/单节点)。
self.relay_internal_base_url = (
relay_internal_base_url.rstrip("/") if relay_internal_base_url else self.relay_base_url
)
self.sync_timeout = sync_timeout
self.health_timeout = health_timeout
self.vcodec = vcodec
self.preset = preset
self.crf = crf
self.bitrate = bitrate
self._relay_secret = relay_secret
self.oss_tmp_prefix = oss_tmp_prefix.rstrip("/") + "/" if oss_tmp_prefix else "tmp/gpu-mezzanine/"
RELAY_PATH_PREFIX = "/api/v1/internal/gpu-relay"
# ------------------------------------------------------------------
# URL builders
# ------------------------------------------------------------------
def _relay_url_from_base(self, base_url: str, key: str, secret: str) -> str:
return f"{base_url}{self.RELAY_PATH_PREFIX}/{key}?token={urllib.parse.quote(secret, safe='')}"
def _relay_put_url(self, key: str, secret: str) -> str:
"""给 P4000 回传结果用的 URL(外部可达)。"""
return self._relay_url_from_base(self.relay_base_url, key, secret)
def _relay_internal_url(self, key: str, secret: str) -> str:
"""Worker 自己 GET/DELETE 用的 URL(Docker 内网)。"""
return self._relay_url_from_base(self.relay_internal_base_url, key, secret)
# ------------------------------------------------------------------
# Health
# ------------------------------------------------------------------
def check_health(self) -> GpuHealth:
url = f"{self.endpoint}/health"
try:
with urllib.request.urlopen(url, timeout=self.health_timeout) as resp:
data = json.loads(resp.read().decode("utf-8"))
except (urllib.error.URLError, socket.timeout, TimeoutError, json.JSONDecodeError, ConnectionError) as e:
return GpuHealth(healthy=False, error=f"health probe failed: {e}")
try:
return GpuHealth(
healthy=data.get("status") == "healthy",
worker=str(data.get("worker", "")),
gpu_name=(data.get("gpu") or {}).get("name", ""),
nvenc_h264=bool((data.get("nvenc") or {}).get("h264_nvenc")),
nvenc_hevc=bool((data.get("nvenc") or {}).get("hevc_nvenc")),
)
except Exception as e: # noqa: BLE001
return GpuHealth(healthy=False, error=f"malformed health response: {e}")
# ------------------------------------------------------------------
# High-level: encode a mezzanine file to final output
# ------------------------------------------------------------------
def encode_mezzanine_to_output(
self,
mezzanine_path: Path,
output_path: Path,
*,
extra_video_args: Optional[list[str]] = None,
audio_args: Optional[list[str]] = None,
timeout: Optional[int] = None,
) -> dict[str, Any]:
"""把 mezzanine(CPU 滤镜已完成)交给 P4000 NVENC 编码,结果写到 output_path。
extra_video_args: -i 之后、-c:v 之前插入的 ffmpeg 参数(如分辨率/帧率调整)。
audio_args: 音频编码参数(如 ["-c:a","aac","-b:a","128k"]);None 表示 -an 无音频。
"""
if not mezzanine_path.exists():
raise GpuEncodeError(f"mezzanine file not found: {mezzanine_path}")
if not self.relay_base_url:
raise GpuEncodeError("gpu_encode_relay_base_url not configured")
timeout = timeout or self.sync_timeout
t_total = time.time()
oss_key: Optional[str] = None
relay_key: Optional[str] = None
try:
# 1. upload mezzanine → OSS
input_url, oss_key = self._upload_mezzanine(mezzanine_path)
logger.debug("[gpu-encoder] mezzanine uploaded: oss_key=%s", oss_key)
# 2. prepare relay URLs (PUT 走外部 URL 给 P4000;GET/DELETE 走内部 Docker 网络)
relay_key = uuid.uuid4().hex
secret = self._get_relay_secret()
put_url = self._relay_put_url(relay_key, secret)
get_url = self._relay_internal_url(relay_key, secret)
del_url = get_url # 内部 URL,DELETE method
# 3. build ffmpeg args
ffmpeg_args = ["-y", "-i", "in.mp4"]
if extra_video_args:
ffmpeg_args.extend(extra_video_args)
ffmpeg_args.extend(["-c:v", self.vcodec, "-preset", self.preset])
if self.bitrate:
ffmpeg_args.extend(["-b:v", self.bitrate])
else:
ffmpeg_args.extend(["-cq", str(self.crf)])
ffmpeg_args.extend(["-pix_fmt", "yuv420p", "-movflags", "+faststart"])
if audio_args:
ffmpeg_args.extend(audio_args)
else:
ffmpeg_args.append("-an")
ffmpeg_args.extend(["-f", "mp4", "pipe:1"])
# 4. call P4000 sync render
body = {
"inputs": {"in.mp4": input_url},
"ffmpeg_args": ffmpeg_args,
"output_url": put_url,
"timeout": int(timeout),
}
job = self._post_sync(body, mezzanine_path=mezzanine_path)
logger.info(
"[gpu-encoder] P4000 done: job_id=%s rc=%s size=%s dur=%ss",
job.get("job_id"),
job.get("ffmpeg_rc"),
job.get("size"),
job.get("duration"),
)
# 5. download result from relay to output_path
output_path.parent.mkdir(parents=True, exist_ok=True)
size = self._download_to_file(get_url, output_path)
# 6. cleanup relay
self._relay_delete(del_url)
logger.info(
"[gpu-encoder] encode ok: %s → %s (%d bytes) total=%.2fs",
mezzanine_path.name,
output_path.name,
size,
time.time() - t_total,
)
return {"job": job, "output_size": size, "output_path": str(output_path)}
except GpuEncodeError:
raise
except Exception as e: # noqa: BLE001
raise GpuEncodeError(f"unexpected: {e}") from e
finally:
# cleanup OSS mezzanine (best-effort)
if oss_key:
try:
self._delete_oss(oss_key)
except Exception as e: # noqa: BLE001
logger.warning("[gpu-encoder] failed to delete OSS mezzanine %s: %s", oss_key, e)
# relay cleanup also best-effort (done above after download)
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _get_relay_secret(self) -> str:
if self._relay_secret:
return self._relay_secret
# read from env (same var API server uses)
env = (os.getenv("APP_ENV", os.getenv("ENV", "development"))).lower()
secret = (os.getenv("GPU_ENCODE_RELAY_SECRET", "") or "").strip()
if not secret:
if env in ("production", "prod"):
raise GpuEncodeError("GPU_ENCODE_RELAY_SECRET must be set in production")
# dev: fail - worker should always have a secret explicitly set (or same ephemeral won't match)
raise GpuEncodeError("GPU_ENCODE_RELAY_SECRET not set")
return secret
def _post_sync(self, body: dict[str, Any], *, mezzanine_path: Path) -> dict[str, Any]:
url = f"{self.endpoint}/api/render/sync"
req_timeout = body.get("timeout", self.sync_timeout) + 60
payload = json.dumps(body).encode("utf-8")
req = urllib.request.Request(
url,
data=payload,
headers={"Content-Type": "application/json"},
method="POST",
)
t0 = time.time()
try:
with urllib.request.urlopen(req, timeout=req_timeout) as resp:
raw = resp.read().decode("utf-8")
except urllib.error.HTTPError as e:
detail = e.read().decode("utf-8", errors="replace")[:1000]
raise GpuEncodeError(f"P4000 HTTP {e.code}: {detail}") from e
except (urllib.error.URLError, socket.timeout, TimeoutError, ConnectionError) as e:
raise GpuEncodeError(f"P4000 connection error: {e}") from e
try:
result = json.loads(raw)
except json.JSONDecodeError as e:
raise GpuEncodeError(f"P4000 bad JSON: {raw[:500]}") from e
dt = time.time() - t0
status = result.get("status")
ffmpeg_rc = result.get("ffmpeg_rc")
uploaded = result.get("uploaded")
if status != "completed" or ffmpeg_rc != 0:
err = result.get("message") or result.get("error") or "unknown"
raise GpuEncodeError(f"P4000 job failed: status={status} rc={ffmpeg_rc} err={err!s:.500}")
# P4000 has a known bug where uploaded=true even on PUT SSL failure;
# we will verify by downloading, so don't hard-fail here but log
if not uploaded:
logger.warning("[gpu-encoder] P4000 reports uploaded=false (will verify via download)")
result["_roundtrip"] = dt
return result
def _download_to_file(self, url: str, output_path: Path) -> int:
"""GET url → write to output_path. Returns bytes written."""
tmp = output_path.with_suffix(output_path.suffix + ".gpu_tmp")
size = 0
try:
with urllib.request.urlopen(url, timeout=self.sync_timeout) as resp:
if resp.status != 200:
raise GpuEncodeError(f"relay GET returned HTTP {resp.status}")
with open(tmp, "wb") as f:
while True:
chunk = resp.read(1024 * 256)
if not chunk:
break
f.write(chunk)
size += len(chunk)
if size == 0:
raise GpuEncodeError("relay returned empty file")
os.replace(tmp, output_path)
return size
except (urllib.error.URLError, socket.timeout, TimeoutError, ConnectionError) as e:
if tmp.exists():
try:
tmp.unlink()
except OSError:
pass
raise GpuEncodeError(f"failed to download from relay: {e}") from e
def _relay_delete(self, url: str) -> None:
try:
req = urllib.request.Request(url, method="DELETE")
with urllib.request.urlopen(req, timeout=10) as resp:
resp.read()
except Exception as e: # noqa: BLE001
logger.debug("[gpu-encoder] relay cleanup delete failed: %s", e)
# ------------------------------------------------------------------
# OSS helpers (optional - storage may not be available in all envs)
# ------------------------------------------------------------------
def _upload_mezzanine(self, path: Path) -> tuple[str, str]:
"""Upload mezzanine to OSS tmp prefix, return (signed_get_url, oss_key)."""
try:
from packages.shared.storage import get_storage_service
except ImportError as e:
raise GpuEncodeError(f"storage service unavailable: {e}") from e
storage = get_storage_service()
if storage is None or storage.bucket is None:
raise GpuEncodeError("OSS storage not configured; cannot upload mezzanine")
key = f"{self.oss_tmp_prefix}{uuid.uuid4().hex}.mp4"
try:
storage.upload_file(str(path), key, content_type="video/mp4")
except Exception as e: # noqa: BLE001
raise GpuEncodeError(f"failed to upload mezzanine to OSS: {e}") from e
# Generate signed GET URL (1h expiry)
signed = storage.get_download_url(key, expires_seconds=3600)
return signed, key
def _delete_oss(self, key: str) -> None:
try:
from packages.shared.storage import get_storage_service
storage = get_storage_service()
if storage is not None and storage.bucket is not None:
storage.delete_file(key)
except Exception as e: # noqa: BLE001
logger.debug("[gpu-encoder] OSS delete %s failed: %s", key, e)
# ── Singleton factory ────────────────────────────────────────────────────
_default_client: Optional[GpuEncoderClient] = None
_default_client_initialized: bool = False
def _build_client_from_settings() -> Optional[GpuEncoderClient]:
try:
from packages.config import get_shared_settings
settings = get_shared_settings()
except Exception: # noqa: BLE001
return None
if not getattr(settings, "enable_gpu_encode", False):
return None
endpoint = (getattr(settings, "gpu_encode_endpoint", "") or "").strip()
relay = (getattr(settings, "gpu_encode_relay_base_url", "") or "").strip()
relay_internal = (getattr(settings, "gpu_encode_relay_internal_base_url", "") or "").strip()
if not endpoint or not relay:
return None
return GpuEncoderClient(
endpoint=endpoint,
relay_base_url=relay,
relay_internal_base_url=relay_internal,
sync_timeout=getattr(settings, "gpu_encode_sync_timeout", 300),
health_timeout=getattr(settings, "gpu_encode_health_timeout", 3.0),
vcodec=getattr(settings, "gpu_encode_vcodec", "h264_nvenc"),
preset=getattr(settings, "gpu_encode_preset", "p4"),
crf=getattr(settings, "gpu_encode_crf", 23),
bitrate=getattr(settings, "gpu_encode_bitrate", "") or "",
relay_secret=getattr(settings, "gpu_encode_relay_secret", "") or "",
oss_tmp_prefix=getattr(settings, "gpu_encode_oss_tmp_prefix", "tmp/gpu-mezzanine/"),
)
def get_gpu_encoder() -> Optional[GpuEncoderClient]:
"""返回进程级单例;未启用或未配置返回 None。"""
global _default_client, _default_client_initialized
if not _default_client_initialized:
_default_client_initialized = True
try:
_default_client = _build_client_from_settings()
except Exception as e: # noqa: BLE001
logger.warning("[gpu-encoder] failed to init client (CPU fallback): %s", e)
_default_client = None
return _default_client
def reset_gpu_encoder_for_tests() -> None:
global _default_client, _default_client_initialized
_default_client = None
_default_client_initialized = False
# Convenience
def is_gpu_encode_enabled() -> bool:
return get_gpu_encoder() is not None
-590
View File
@@ -1,590 +0,0 @@
"""GpuEncoderClient 单元测试:mock HTTP,覆盖 health/sync/fallback/singleton 等完整路径。"""
from __future__ import annotations
import json
import socket
import sys
import urllib.error
import urllib.request
from http.client import HTTPResponse
from io import BytesIO
from pathlib import Path
from unittest import mock
import pytest
from packages.shared.gpu_encoder import (
GpuEncodeError,
GpuEncoderClient,
GpuHealth,
_build_client_from_settings,
get_gpu_encoder,
is_gpu_encode_enabled,
reset_gpu_encoder_for_tests,
)
@pytest.fixture(autouse=True)
def _reset_singleton():
reset_gpu_encoder_for_tests()
yield
reset_gpu_encoder_for_tests()
@pytest.fixture
def client():
return GpuEncoderClient(
endpoint="http://gpu.example.com:8900",
relay_base_url="http://api.example.com",
relay_internal_base_url="http://api-internal:8000",
sync_timeout=60,
health_timeout=2,
relay_secret="test-secret",
)
def _fake_response(status: int = 200, body: dict | bytes | None = None, headers=None):
if isinstance(body, dict):
data = json.dumps(body).encode("utf-8")
elif body is None:
data = b""
else:
data = body
bio = BytesIO(data)
resp = mock.MagicMock(spec=HTTPResponse)
resp.status = status
resp.read.side_effect = lambda n=-1: bio.read(n)
resp.__enter__ = mock.MagicMock(return_value=resp)
resp.__exit__ = mock.MagicMock(return_value=False)
return resp
# ── Health check ────────────────────────────────────────────────────
class TestHealthCheck:
def test_healthy_nvenc_available(self, client):
body = {
"status": "healthy",
"worker": "w1",
"gpu": {"name": "Quadro P4000"},
"nvenc": {"h264_nvenc": True, "hevc_nvenc": True},
}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=body)):
h = client.check_health()
assert h.healthy and h.nvenc_h264 and h.ready
assert h.gpu_name == "Quadro P4000"
def test_connection_error_returns_unhealthy(self, client):
with mock.patch("urllib.request.urlopen", side_effect=urllib.error.URLError("timeout")):
h = client.check_health()
assert not h.healthy
assert "health probe failed" in h.error
def test_bad_json_returns_unhealthy(self, client):
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=b"not json")):
h = client.check_health()
assert not h.healthy
def test_nvenc_unavailable(self, client):
body = {"status": "healthy", "gpu": {"name": "t"}, "nvenc": {"h264_nvenc": False}}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=body)):
h = client.check_health()
assert h.healthy and not h.ready
def test_malformed_response_inner_exception(self, client):
"""data 是合法 JSON 但 gpu 字段类型错(字符串)触发内部 except."""
body = {"status": "healthy", "gpu": "not-a-dict", "nvenc": {}}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=body)):
h = client.check_health()
assert not h.healthy
assert "malformed" in h.error
# ── _post_sync ──────────────────────────────────────────────────────
class TestPostSync:
def test_completed_job_returns_dict(self, client):
result_body = {
"job_id": "j1",
"status": "completed",
"ffmpeg_rc": 0,
"uploaded": True,
"duration": 5.1,
"size": 123456,
}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=result_body)) as m:
res = client._post_sync(
{
"inputs": {"in.mp4": "http://x"},
"ffmpeg_args": ["-i", "in.mp4"],
"output_url": "http://relay/k?token=s",
"timeout": 30,
},
mezzanine_path=Path("/tmp/fake.mp4"),
)
assert res["status"] == "completed" and res["ffmpeg_rc"] == 0
req = m.call_args[0][0]
assert req.full_url == "http://gpu.example.com:8900/api/render/sync"
def test_ffmpeg_failure_raises(self, client):
body = {"status": "failed", "ffmpeg_rc": 1, "message": "Invalid data"}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=body)):
with pytest.raises(GpuEncodeError, match="rc=1"):
client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
def test_http_4xx_raises(self, client):
err = urllib.error.HTTPError(
url="http://gpu/render/sync", code=422, msg="Unprocessable", hdrs={}, fp=BytesIO(b"bad request")
)
with mock.patch("urllib.request.urlopen", side_effect=err):
with pytest.raises(GpuEncodeError, match="HTTP 422"):
client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
def test_connection_error_raises(self, client):
with mock.patch("urllib.request.urlopen", side_effect=urllib.error.URLError("conn refused")):
with pytest.raises(GpuEncodeError, match="connection error"):
client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
def test_timeout_error_raises(self, client):
with mock.patch("urllib.request.urlopen", side_effect=socket.timeout("timed out")):
with pytest.raises(GpuEncodeError, match="connection error"):
client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
def test_bad_json_raises(self, client):
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=b"not-json")):
with pytest.raises(GpuEncodeError, match="bad JSON"):
client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
def test_uploaded_false_logs_warning_but_succeeds(self, client, caplog):
body = {"status": "completed", "ffmpeg_rc": 0, "uploaded": False, "job_id": "j"}
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=body)), caplog.at_level("WARNING"):
res = client._post_sync(
{"inputs": {}, "ffmpeg_args": [], "output_url": "", "timeout": 10}, mezzanine_path=Path("/tmp/x")
)
assert res["status"] == "completed"
assert "uploaded=false" in caplog.text
# ── Relay URL builders ──────────────────────────────────────────────
class TestRelayUrl:
def test_put_url_uses_external_base(self, client):
url = client._relay_put_url("abc123", "secret!")
assert "abc123" in url
assert "token=secret%21" in url
assert url.startswith("http://api.example.com/api/v1/internal/gpu-relay/")
def test_internal_url_uses_internal_base(self, client):
url = client._relay_internal_url("abc123", "s")
assert url.startswith("http://api-internal:8000/api/v1/internal/gpu-relay/abc123")
def test_internal_url_falls_back_to_external_when_not_set(self):
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api.example.com", relay_secret="s")
put = c._relay_put_url("k", "s")
internal = c._relay_internal_url("k", "s")
assert put.startswith("http://api.example.com/")
assert internal == put
def test_encode_uses_different_put_and_get_urls(self, client):
put_url = client._relay_put_url("k", "test-secret")
get_url = client._relay_internal_url("k", "test-secret")
assert "api.example.com" in put_url and "api-internal:8000" in get_url and put_url != get_url
# ── _get_relay_secret ──────────────────────────────────────────────
class TestGetRelaySecret:
def test_explicit_secret_used(self, client):
assert client._get_relay_secret() == "test-secret"
def test_env_secret_used_when_not_explicit(self, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "from-env")
monkeypatch.setenv("APP_ENV", "staging")
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api")
assert c._get_relay_secret() == "from-env"
def test_prod_without_secret_raises(self, monkeypatch):
monkeypatch.delenv("GPU_ENCODE_RELAY_SECRET", raising=False)
monkeypatch.setenv("APP_ENV", "production")
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api")
with pytest.raises(GpuEncodeError, match="GPU_ENCODE_RELAY_SECRET"):
c._get_relay_secret()
def test_dev_without_secret_raises(self, monkeypatch):
"""未设置 secret 且非 production 也 raise(worker 必须显式配置)。"""
monkeypatch.delenv("GPU_ENCODE_RELAY_SECRET", raising=False)
monkeypatch.setenv("APP_ENV", "development")
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api")
with pytest.raises(GpuEncodeError, match="GPU_ENCODE_RELAY_SECRET not set"):
c._get_relay_secret()
# ── _download_to_file ──────────────────────────────────────────────
class TestDownloadToFile:
def test_writes_file(self, client, tmp_path):
data = b"hello" * 1000
out = tmp_path / "out.mp4"
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=data)):
size = client._download_to_file("http://relay/k?token=s", out)
assert size == len(data) and out.read_bytes() == data
def test_empty_file_raises(self, client, tmp_path):
out = tmp_path / "out.mp4"
with mock.patch("urllib.request.urlopen", return_value=_fake_response(body=b"")):
with pytest.raises(GpuEncodeError, match="empty file"):
client._download_to_file("http://relay/k", out)
assert not out.exists()
def test_non_200_status_raises(self, client, tmp_path):
out = tmp_path / "o.mp4"
with mock.patch("urllib.request.urlopen", return_value=_fake_response(status=404, body=b"")):
with pytest.raises(GpuEncodeError, match="HTTP 404"):
client._download_to_file("http://relay/k", out)
def test_url_error_cleans_up_tmp(self, client, tmp_path):
out = tmp_path / "o.mp4"
tmp_file = out.with_suffix(out.suffix + ".gpu_tmp")
tmp_file.write_bytes(b"partial")
assert tmp_file.exists()
with mock.patch("urllib.request.urlopen", side_effect=urllib.error.URLError("net down")):
with pytest.raises(GpuEncodeError, match="failed to download"):
client._download_to_file("http://relay/k", out)
assert not tmp_file.exists()
# ── encode_mezzanine_to_output ─────────────────────────────────────
class TestEncodeMezzanine:
def test_happy_path_with_audio(self, client, tmp_path):
mezz = tmp_path / "mezz.mp4"
mezz.write_bytes(b"M" * 100)
out = tmp_path / "out" / "final.mp4"
with (
mock.patch.object(client, "_upload_mezzanine", return_value=("http://oss/signed", "osskey1")),
mock.patch.object(
client,
"_post_sync",
return_value={
"job_id": "j1",
"status": "completed",
"ffmpeg_rc": 0,
"uploaded": True,
"size": 5000,
"duration": 1.2,
},
) as m_post,
mock.patch.object(client, "_download_to_file", return_value=5000) as m_dl,
mock.patch.object(client, "_relay_delete") as m_del,
mock.patch.object(client, "_delete_oss") as m_ossdel,
):
result = client.encode_mezzanine_to_output(mezz, out, audio_args=["-c:a", "aac"])
assert result["output_size"] == 5000 and str(out) == result["output_path"]
body = m_post.call_args[0][0]
assert "-c:a" in body["ffmpeg_args"] and "aac" in body["ffmpeg_args"]
assert "-an" not in body["ffmpeg_args"]
assert body["output_url"].startswith("http://api.example.com/")
assert "api-internal:8000" in m_dl.call_args[0][0]
m_del.assert_called_once()
m_ossdel.assert_called_once_with("osskey1")
def test_happy_path_no_audio_uses_an_and_cq(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"M")
out = tmp_path / "o.mp4"
with (
mock.patch.object(client, "_upload_mezzanine", return_value=("http://oss/u", "k")),
mock.patch.object(
client,
"_post_sync",
return_value={"status": "completed", "ffmpeg_rc": 0, "uploaded": True, "job_id": "j"},
) as m_post,
mock.patch.object(client, "_download_to_file", return_value=100),
mock.patch.object(client, "_relay_delete"),
mock.patch.object(client, "_delete_oss"),
):
client.encode_mezzanine_to_output(mezz, out)
body = m_post.call_args[0][0]
assert "-an" in body["ffmpeg_args"] and "-cq" in body["ffmpeg_args"]
assert str(client.crf) in body["ffmpeg_args"]
def test_bitrate_set_uses_bv_instead_of_cq(self, tmp_path):
c = GpuEncoderClient(
endpoint="http://gpu",
relay_base_url="http://api",
relay_internal_base_url="http://api-int:8000",
relay_secret="s",
bitrate="2M",
)
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
out = tmp_path / "o.mp4"
with (
mock.patch.object(c, "_upload_mezzanine", return_value=("http://oss/u", "k")),
mock.patch.object(
c, "_post_sync", return_value={"status": "completed", "ffmpeg_rc": 0, "uploaded": True, "job_id": "j"}
) as m_post,
mock.patch.object(c, "_download_to_file", return_value=10),
mock.patch.object(c, "_relay_delete"),
mock.patch.object(c, "_delete_oss"),
):
c.encode_mezzanine_to_output(mezz, out, extra_video_args=["-vf", "scale=1280:-2"])
body = m_post.call_args[0][0]
assert "-b:v" in body["ffmpeg_args"] and "2M" in body["ffmpeg_args"]
assert "-cq" not in body["ffmpeg_args"]
assert "-vf" in body["ffmpeg_args"]
def test_mezzanine_not_found_raises(self, client, tmp_path):
with pytest.raises(GpuEncodeError, match="mezzanine file not found"):
client.encode_mezzanine_to_output(tmp_path / "nope.mp4", tmp_path / "o.mp4")
def test_relay_base_not_configured_raises(self, tmp_path):
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="", relay_secret="s")
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
with pytest.raises(GpuEncodeError, match="relay_base_url"):
c.encode_mezzanine_to_output(mezz, tmp_path / "o.mp4")
def test_unexpected_exception_is_wrapped(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
with (
mock.patch.object(client, "_upload_mezzanine", return_value=("http://oss/u", "k")),
mock.patch.object(client, "_post_sync", side_effect=RuntimeError("boom")),
mock.patch.object(client, "_delete_oss"),
):
with pytest.raises(GpuEncodeError, match="unexpected: boom"):
client.encode_mezzanine_to_output(mezz, tmp_path / "o.mp4")
def test_gpu_encode_error_re_raised_directly(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
with (
mock.patch.object(client, "_upload_mezzanine", return_value=("http://oss/u", "k")),
mock.patch.object(client, "_post_sync", side_effect=GpuEncodeError("direct fail")),
mock.patch.object(client, "_delete_oss"),
):
with pytest.raises(GpuEncodeError, match="direct fail"):
client.encode_mezzanine_to_output(mezz, tmp_path / "o.mp4")
def test_oss_cleanup_runs_on_failure(self, client, tmp_path, caplog):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
with (
mock.patch.object(client, "_upload_mezzanine", return_value=("http://oss/u", "ossk")),
mock.patch.object(client, "_post_sync", side_effect=GpuEncodeError("enc fail")),
mock.patch.object(client, "_delete_oss", side_effect=Exception("oss down")) as m_ossdel,
caplog.at_level("WARNING"),
):
with pytest.raises(GpuEncodeError):
client.encode_mezzanine_to_output(mezz, tmp_path / "o.mp4")
m_ossdel.assert_called_once_with("ossk")
# ── _relay_delete ──────────────────────────────────────────────────
class TestRelayDelete:
def test_exception_is_swallowed(self, client, caplog):
with mock.patch("urllib.request.urlopen", side_effect=RuntimeError("boom")), caplog.at_level("DEBUG"):
client._relay_delete("http://relay/k?token=s")
assert "cleanup delete failed" in caplog.text
def test_success_issues_delete(self, client):
with mock.patch("urllib.request.urlopen", return_value=_fake_response(status=204, body=b"")) as m:
client._relay_delete("http://relay/k?token=s")
assert m.call_args[0][0].get_method() == "DELETE"
# ── OSS helpers ────────────────────────────────────────────────────
class TestOssHelpers:
def test_upload_storage_import_error(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
# 删除 sys.modules 中 packages.shared.storage 使导入失败
saved = sys.modules.pop("packages.shared.storage", None)
try:
real_import = __builtins__.__import__ if hasattr(__builtins__, "__import__") else __import__
def fake_import(name, *a, **kw):
if name == "packages.shared.storage" or name.startswith("packages.shared.storage."):
raise ImportError("no storage")
return real_import(name, *a, **kw)
with mock.patch("builtins.__import__", side_effect=fake_import):
with pytest.raises(GpuEncodeError, match="storage service unavailable"):
client._upload_mezzanine(mezz)
finally:
if saved is not None:
sys.modules["packages.shared.storage"] = saved
def test_upload_storage_none(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = None
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
with pytest.raises(GpuEncodeError, match="OSS storage not configured"):
client._upload_mezzanine(mezz)
def test_upload_bucket_none(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
svc = mock.MagicMock()
svc.bucket = None
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = svc
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
with pytest.raises(GpuEncodeError, match="OSS storage not configured"):
client._upload_mezzanine(mezz)
def test_upload_failure_raises(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
svc = mock.MagicMock()
svc.bucket = object()
svc.upload_file.side_effect = RuntimeError("oss err")
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = svc
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
with pytest.raises(GpuEncodeError, match="failed to upload mezzanine"):
client._upload_mezzanine(mezz)
def test_upload_success(self, client, tmp_path):
mezz = tmp_path / "m.mp4"
mezz.write_bytes(b"x")
svc = mock.MagicMock()
svc.bucket = object()
svc.get_download_url.return_value = "https://oss/signed?sig=abc"
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = svc
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
url, key = client._upload_mezzanine(mezz)
assert url.startswith("https://oss/signed")
assert key.startswith("tmp/gpu-mezzanine/") and key.endswith(".mp4")
svc.upload_file.assert_called_once()
def test_delete_oss_exception_swallowed(self, client, caplog):
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.side_effect = RuntimeError("svc down")
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}), caplog.at_level("DEBUG"):
client._delete_oss("somekey")
assert "OSS delete" in caplog.text
def test_delete_oss_bucket_none_noop(self, client):
svc = mock.MagicMock()
svc.bucket = None
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = svc
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
client._delete_oss("k")
svc.delete_file.assert_not_called()
def test_delete_oss_success(self, client):
svc = mock.MagicMock()
svc.bucket = object()
fake_mod = mock.MagicMock()
fake_mod.get_storage_service.return_value = svc
with mock.patch.dict("sys.modules", {"packages.shared.storage": fake_mod}):
client._delete_oss("k")
svc.delete_file.assert_called_once_with("k")
# ── Singleton / factory ────────────────────────────────────────────
class TestSingletonFactory:
def test_build_client_import_error_returns_none(self):
saved = sys.modules.get("packages.config")
sys.modules["packages.config"] = None
try:
with mock.patch("builtins.__import__", side_effect=RuntimeError("no cfg")):
assert _build_client_from_settings() is None
finally:
if saved is not None:
sys.modules["packages.config"] = saved
def test_build_client_not_enabled_returns_none(self):
s = mock.MagicMock()
s.enable_gpu_encode = False
fake_mod = mock.MagicMock()
fake_mod.get_shared_settings.return_value = s
with mock.patch.dict("sys.modules", {"packages.config": fake_mod}):
assert _build_client_from_settings() is None
def test_build_client_missing_endpoint(self):
s = mock.MagicMock()
s.enable_gpu_encode = True
s.gpu_encode_endpoint = ""
s.gpu_encode_relay_base_url = "http://api"
s.gpu_encode_relay_internal_base_url = ""
fake_mod = mock.MagicMock()
fake_mod.get_shared_settings.return_value = s
with mock.patch.dict("sys.modules", {"packages.config": fake_mod}):
assert _build_client_from_settings() is None
def test_build_client_missing_relay(self):
s = mock.MagicMock()
s.enable_gpu_encode = True
s.gpu_encode_endpoint = "http://gpu"
s.gpu_encode_relay_base_url = ""
s.gpu_encode_relay_internal_base_url = ""
fake_mod = mock.MagicMock()
fake_mod.get_shared_settings.return_value = s
with mock.patch.dict("sys.modules", {"packages.config": fake_mod}):
assert _build_client_from_settings() is None
def test_build_client_success(self):
s = mock.MagicMock()
s.enable_gpu_encode = True
s.gpu_encode_endpoint = "http://gpu"
s.gpu_encode_relay_base_url = "http://api/"
s.gpu_encode_relay_internal_base_url = "http://api-int:8000/"
s.gpu_encode_sync_timeout = 120
s.gpu_encode_health_timeout = 1.0
s.gpu_encode_vcodec = "h264_nvenc"
s.gpu_encode_preset = "p7"
s.gpu_encode_crf = 20
s.gpu_encode_bitrate = ""
s.gpu_encode_relay_secret = "s"
s.gpu_encode_oss_tmp_prefix = "tmp/x/"
fake_mod = mock.MagicMock()
fake_mod.get_shared_settings.return_value = s
with mock.patch.dict("sys.modules", {"packages.config": fake_mod}):
c = _build_client_from_settings()
assert c is not None and c.endpoint == "http://gpu"
assert c.relay_base_url == "http://api"
assert c.relay_internal_base_url == "http://api-int:8000"
assert c.preset == "p7" and c.crf == 20
def test_get_gpu_encoder_init_failure_returns_none(self, caplog):
with (
mock.patch("packages.shared.gpu_encoder._build_client_from_settings", side_effect=RuntimeError("boom")),
caplog.at_level("WARNING"),
):
assert get_gpu_encoder() is None
assert "failed to init client" in caplog.text
def test_get_gpu_encoder_returns_singleton_and_enabled(self):
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api", relay_secret="s")
with mock.patch("packages.shared.gpu_encoder._build_client_from_settings", return_value=c):
assert get_gpu_encoder() is c and get_gpu_encoder() is c
assert is_gpu_encode_enabled() is True
def test_is_gpu_encode_enabled_when_none(self):
with mock.patch("packages.shared.gpu_encoder._build_client_from_settings", return_value=None):
assert is_gpu_encode_enabled() is False
# ── Constructor edge cases ────────────────────────────────────────
class TestConstructor:
def test_oss_tmp_prefix_empty_uses_default(self):
c = GpuEncoderClient(endpoint="http://gpu", relay_base_url="http://api", relay_secret="s", oss_tmp_prefix="")
assert c.oss_tmp_prefix == "tmp/gpu-mezzanine/"
def test_oss_tmp_prefix_strips_and_adds_slash(self):
c = GpuEncoderClient(
endpoint="http://gpu", relay_base_url="http://api", relay_secret="s", oss_tmp_prefix="tmp/foo"
)
assert c.oss_tmp_prefix == "tmp/foo/"
-240
View File
@@ -1,240 +0,0 @@
"""gpu_relay API 路由单元测试:覆盖 helper 函数 + PUT/GET/HEAD/DELETE handler。"""
from __future__ import annotations
import os
from pathlib import Path
from unittest import mock
import pytest
from fastapi import HTTPException
from apps.api.app.api.routes import gpu_relay
# ── _relay_dir ────────────────────────────────────────────────────────
class TestRelayDir:
def test_default_dir(self, tmp_path, monkeypatch):
monkeypatch.delenv("GENERATED_FILES_DIR", raising=False)
monkeypatch.delenv("GPU_ENCODE_RELAY_DIR", raising=False)
# 用 tmp_path 作 base
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
p = gpu_relay._relay_dir()
assert p == tmp_path / "gpu_relay"
assert p.exists()
def test_custom_subdir(self, tmp_path, monkeypatch):
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
monkeypatch.setenv("GPU_ENCODE_RELAY_DIR", "custom_relay")
p = gpu_relay._relay_dir()
assert p == tmp_path / "custom_relay"
assert p.exists()
# ── _secret ──────────────────────────────────────────────────────────
class TestSecret:
def test_explicit_secret_returned(self, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "topsecret")
gpu_relay._DEFAULT_SECRET_LOGGED = False
assert gpu_relay._secret() == "topsecret"
def test_prod_without_secret_raises(self, monkeypatch):
monkeypatch.delenv("GPU_ENCODE_RELAY_SECRET", raising=False)
monkeypatch.setenv("APP_ENV", "production")
with pytest.raises(RuntimeError, match="GPU_ENCODE_RELAY_SECRET must be set"):
gpu_relay._secret()
def test_dev_without_secret_generates_ephemeral(self, monkeypatch, caplog):
monkeypatch.delenv("GPU_ENCODE_RELAY_SECRET", raising=False)
monkeypatch.setenv("APP_ENV", "development")
gpu_relay._DEFAULT_SECRET_LOGGED = False
with caplog.at_level("WARNING"):
secret = gpu_relay._secret()
assert len(secret) > 16
assert "ephemeral dev token" in caplog.text
# 第二次调用不再 log(_DEFAULT_SECRET_LOGGED=True)
before = len(caplog.records)
secret2 = gpu_relay._secret()
assert secret2 == secret
assert len(caplog.records) == before
# 清理
monkeypatch.delenv("GPU_ENCODE_RELAY_SECRET", raising=False)
# ── _safe_key ────────────────────────────────────────────────────────
class TestSafeKey:
@pytest.mark.parametrize("bad", ["", "../etc", "a/b", "a\\b", ".", "..", "a b", "a%b"])
def test_invalid_keys_rejected(self, bad):
with pytest.raises(HTTPException) as ei:
gpu_relay._safe_key(bad)
assert ei.value.status_code == 400
@pytest.mark.parametrize("good", ["abc123", "ABC-Def_01", "a" * 32])
def test_valid_keys_accepted(self, good):
assert gpu_relay._safe_key(good) == good
def test_strips_whitespace(self):
assert gpu_relay._safe_key(" abc ") == "abc"
# ── _check_token ─────────────────────────────────────────────────────
class TestCheckToken:
def test_missing_token_401(self):
with pytest.raises(HTTPException) as ei:
gpu_relay._check_token(None)
assert ei.value.status_code == 401
def test_wrong_token_401(self, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "correct")
with pytest.raises(HTTPException) as ei:
gpu_relay._check_token("wrong")
assert ei.value.status_code == 401
def test_correct_token_passes(self, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "correct")
assert gpu_relay._check_token("correct") is None
# ── build_relay_* helpers ───────────────────────────────────────────
class TestBuildRelayUrls:
def test_put_url(self):
url = gpu_relay.build_relay_put_url("http://api.example.com/", "k1", "s")
assert url == "http://api.example.com/api/v1/internal/gpu-relay/k1?token=s"
def test_get_url_same_as_put(self):
assert gpu_relay.build_relay_get_url("http://api", "k", "s") == gpu_relay.build_relay_put_url(
"http://api", "k", "s"
)
def test_generate_key_is_hex(self):
k = gpu_relay.generate_key()
assert len(k) == 32
int(k, 16) # valid hex
# ── PUT endpoint ────────────────────────────────────────────────────
@pytest.mark.asyncio
class TestPutObject:
async def test_put_writes_file(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
# async request.stream 模拟
async def _stream():
yield b"chunk1"
yield b"chunk2"
req = mock.MagicMock()
req.stream = _stream
resp = await gpu_relay.put_object(key="abc123", request=req, token="s")
assert resp["ok"] is True
assert resp["size"] == len(b"chunk1") + len(b"chunk2")
p = tmp_path / "gpu_relay" / "abc123"
assert p.read_bytes() == b"chunk1chunk2"
# .part 临时文件应已 rename
assert not p.with_suffix(p.suffix + ".part").exists()
async def test_put_invalid_key_400(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
req = mock.MagicMock()
with pytest.raises(HTTPException) as ei:
await gpu_relay.put_object(key="../bad", request=req, token="s")
assert ei.value.status_code == 400
async def test_put_bad_token_401(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "correct")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
req = mock.MagicMock()
with pytest.raises(HTTPException) as ei:
await gpu_relay.put_object(key="abc", request=req, token="wrong")
assert ei.value.status_code == 401
async def test_put_write_error_cleans_tmp(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
async def _bad_stream():
yield b"x"
raise OSError("disk full")
req = mock.MagicMock()
req.stream = _bad_stream
with pytest.raises(HTTPException) as ei:
await gpu_relay.put_object(key="abc", request=req, token="s")
assert ei.value.status_code == 500
# tmp 文件被清理
part = tmp_path / "gpu_relay" / "abc.part"
assert not part.exists()
# ── GET endpoint ────────────────────────────────────────────────────
@pytest.mark.asyncio
class TestGetObject:
async def test_get_missing_404(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
with pytest.raises(HTTPException) as ei:
await gpu_relay.get_object(key="nope", token="s")
assert ei.value.status_code == 404
async def test_get_returns_file(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
p = tmp_path / "gpu_relay" / "exist"
p.parent.mkdir(parents=True, exist_ok=True)
p.write_bytes(b"viddata")
resp = await gpu_relay.get_object(key="exist", token="s")
assert resp.media_type == "video/mp4"
# ── HEAD endpoint ───────────────────────────────────────────────────
@pytest.mark.asyncio
class TestHeadObject:
async def test_head_missing_404(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
resp = await gpu_relay.head_object(key="nope", token="s")
assert resp.status_code == 404
async def test_head_returns_content_length(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
p = tmp_path / "gpu_relay" / "k"
p.parent.mkdir(parents=True, exist_ok=True)
p.write_bytes(b"12345")
resp = await gpu_relay.head_object(key="k", token="s")
assert resp.status_code == 200
assert resp.headers["Content-Length"] == "5"
# ── DELETE endpoint ──────────────────────────────────────────────────
@pytest.mark.asyncio
class TestDeleteObject:
async def test_delete_existing(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
p = tmp_path / "gpu_relay" / "k"
p.parent.mkdir(parents=True, exist_ok=True)
p.write_bytes(b"x")
resp = await gpu_relay.delete_object(key="k", token="s")
assert resp["ok"] is True
assert not p.exists()
async def test_delete_missing_is_noop(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
# 不存在时不应 404,返回 ok
resp = await gpu_relay.delete_object(key="nope", token="s")
assert resp["ok"] is True
async def test_delete_unlink_error_500(self, tmp_path, monkeypatch):
monkeypatch.setenv("GPU_ENCODE_RELAY_SECRET", "s")
monkeypatch.setenv("GENERATED_FILES_DIR", str(tmp_path))
p = tmp_path / "gpu_relay" / "k"
p.parent.mkdir(parents=True, exist_ok=True)
p.write_bytes(b"x")
with mock.patch.object(Path, "unlink", side_effect=OSError("perm denied")):
with pytest.raises(HTTPException) as ei:
await gpu_relay.delete_object(key="k", token="s")
assert ei.value.status_code == 500
-156
View File
@@ -1,156 +0,0 @@
"""SQLAlchemyVoiceCloneProfileRepository.cleanup_stale_processing 单元测试。
通过 monkeypatch sys.modules['packages.adapters.sqlalchemy_impl.models'],
注入一个具备 SQLAlchemy 列比较语义(== / < 返回可链式 .all() 的 mock)的假模型类,
不依赖真实 DB,也不会触发 SQLAlchemy 映射。
"""
from __future__ import annotations
import sys
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
class _Col:
"""模拟 SQLAlchemy Column:比较运算返回 MagicMock,可被 filter 链式调用。"""
def __init__(self, name: str):
self._name = name
def __eq__(self, other): # type: ignore[override]
return MagicMock(name=f"{self._name}=={other!r}")
def __ne__(self, other): # type: ignore[override]
return MagicMock(name=f"{self._name}!={other!r}")
def __lt__(self, other):
return MagicMock(name=f"{self._name}<{other!r}")
def __gt__(self, other):
return MagicMock(name=f"{self._name}>{other!r}")
def __le__(self, other):
return MagicMock(name=f"{self._name}<={other!r}")
def __ge__(self, other):
return MagicMock(name=f"{self._name}>={other!r}")
def __hash__(self):
return id(self)
class _FakeVoiceCloneProfileModel:
"""假模型:类属性是 _Col;实例上可读写 status/error_message/updated_at。"""
status = _Col("status")
updated_at = _Col("updated_at")
id = _Col("id")
error_message = _Col("error_message")
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
# ── 预注入 mock 模型模块,避免真实 import 拉起 DB / SQLAlchemy 映射 ──
_fake_models = SimpleNamespace(VoiceCloneProfileModel=_FakeVoiceCloneProfileModel)
sys.modules.setdefault("packages.adapters.sqlalchemy_impl.models", _fake_models)
if "packages.adapters.sqlalchemy_impl.voice_clone_profile_repository" in sys.modules:
mod = sys.modules["packages.adapters.sqlalchemy_impl.voice_clone_profile_repository"]
mod.VoiceCloneProfileModel = _FakeVoiceCloneProfileModel # type: ignore[attr-defined]
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
)
def _make_fake_row(
*,
status: str = "processing",
updated_at: datetime | None = None,
error_message: str = "",
) -> _FakeVoiceCloneProfileModel:
return _FakeVoiceCloneProfileModel(
status=status,
error_message=error_message,
updated_at=updated_at or datetime.now(UTC),
)
def _make_repo(fake_rows: list[_FakeVoiceCloneProfileModel]):
"""构造 repo + mock session。
生产代码使用 .query(Model).filter(A, B).all()(一次 filter,两个表达式参数)。
"""
session = MagicMock()
filtered = MagicMock()
filtered.all.return_value = list(fake_rows)
session.query.return_value.filter.return_value = filtered
repo = SQLAlchemyVoiceCloneProfileRepository.__new__(SQLAlchemyVoiceCloneProfileRepository)
repo.session = session
return repo, session
class TestCleanupStaleProcessing:
"""cleanup_stale_processing 行为测试。"""
def test_no_stale_records_returns_zero_and_no_commit(self):
"""无卡死记录时返回 0,不调用 commit。"""
repo, session = _make_repo([])
assert repo.cleanup_stale_processing() == 0
session.commit.assert_not_called()
def test_stale_record_marked_failed_with_timeout_message(self):
"""超时 processing 记录被标记为 failed,错误信息包含超时分钟数。"""
old = _make_fake_row(updated_at=datetime.now(UTC) - timedelta(minutes=15))
repo, session = _make_repo([old])
count = repo.cleanup_stale_processing(timeout_minutes=10)
assert count == 1
assert old.status == "failed"
assert "超时" in old.error_message
assert "10" in old.error_message
session.commit.assert_called_once()
def test_error_message_reflects_custom_timeout(self):
"""自定义 timeout_minutes 会反映在错误信息里。"""
old = _make_fake_row(updated_at=datetime.now(UTC) - timedelta(hours=1))
repo, _session = _make_repo([old])
repo.cleanup_stale_processing(timeout_minutes=5)
assert old.status == "failed"
assert "5" in old.error_message
def test_multiple_stale_records_all_cleaned_in_single_commit(self):
"""多条卡死记录都被清理,返回正确计数并只 commit 一次。"""
m1 = _make_fake_row(updated_at=datetime.now(UTC) - timedelta(minutes=20))
m2 = _make_fake_row(updated_at=datetime.now(UTC) - timedelta(minutes=11))
repo, session = _make_repo([m1, m2])
assert repo.cleanup_stale_processing(timeout_minutes=10) == 2
assert m1.status == "failed"
assert m2.status == "failed"
session.commit.assert_called_once()
def test_queries_model_with_status_and_updated_at_filters(self):
"""query 被调用,filter 同时传入 status=='processing' 与 updated_at<cutoff 两个条件。
注:不同测试加载顺序下 sys.modules['packages...models'] 可能是真模型类
(因为其他测试文件已先 import),所以这里不断言模型类身份,
只断言 query/filter 被正确调用。
"""
repo, session = _make_repo([])
repo.cleanup_stale_processing()
assert session.query.called, "session.query 应被调用"
# filter 被调用一次,且传入两个过滤表达式
q = session.query.return_value
assert q.filter.called, "query.filter 应被调用"
args_f, _kwargs = q.filter.call_args
assert len(args_f) == 2, f"filter 应接收 2 个位置参数(status + updated_at),实际 {len(args_f)}"
+40 -50
View File
@@ -8,14 +8,13 @@
Celery bind=True 任务的底层函数签名为 (self, profile_id),
CosyVoiceService 在 voice_clone.py 中被实例化传入 workflow,必须 mock 防止真实初始化。
跨环境兼容(_resolve_task):
不同 Celery 版本 / Python 版本 / 是否有 active Celery app,task 对象形态不同:
1) Celery Proxy(LocalProxy/LazyProxy):import 结果是代理对象,调用
_get_current_object() 可能抛 RuntimeError(无 active context),必须 try 保护。
成功取到真实 Task 实例后,使用 bound method .run。
2) Celery Task 实例(bind=True 时 @task 返回的典型形态):直接有 .run/.retry。
3) 原始函数(某些环境装饰器未生效或 patch 时序问题):需手动传 mock_self。
统一返回 (callable, mock_self, real_task),调用方不需要重复解析。
跨环境兼容:
Python 3.13 + Celery 5.4.0 → import 返回 Celery Proxy
→ _get_current_object() 返回 Task 实例 → .run 是 bound method(self 已绑定)
→ 调用方式:task.run(profile_id),retry mock 在 task.run.retry
Python 3.10 + Celery 5.4.0 → import 返回原始函数(装饰器未生效)
→ 签名 (self, profile_id),需手动传 mock_self
→ 调用方式:func(mock_self, profile_id),retry mock 在 mock_self.retry
"""
from __future__ import annotations
@@ -59,33 +58,24 @@ def _make_mock_profile(
def _resolve_task(task_obj):
"""解析 Celery 任务对象,兼容 Proxy / Task 实例 / 原始函数三种形态。
"""解析 Celery 任务对象,返回 (callable, mock_self_or_none)。
所有分支均做异常保护,避免因 Celery Proxy 在无 app context 时抛错导致测试挂掉。
跨环境兼容 Celery Proxy / Task 实例 / 原始函数三种情况。
Returns:
tuple: (callable, mock_self, real_task)
- callable: 最终执行用的可调用对象
- mock_self: 仅原始函数分支需要手动传入 mock self;其他分支为 None
- real_task: 真实 Task 实例(Proxy 分支为 _get_current_object() 结果;
Task 分支为 task_obj 本身;原始函数分支为 None)。用于 patch .retry。
tuple: (callable, mock_self)
- Proxy/Task: callable 是 bound method task.run,mock_self=None
- 原始函数: callable 是原始函数,mock_self 需由调用方提供
"""
# Case 1: Celery Proxy → 安全尝试 _get_current_object()
# Case 1: Celery Proxy → 提取 Task 实例的 .run(bound method)
if hasattr(task_obj, "_get_current_object"):
try:
real_task = task_obj._get_current_object()
if real_task is not None and hasattr(real_task, "run"):
return real_task.run, None, real_task
except Exception:
# 无 active app context 或 Proxy 未绑定,退化为其他分支处理
pass
real_task = task_obj._get_current_object()
return real_task.run, None
# Case 2: Celery Task 实例(非 Proxy)
if hasattr(task_obj, "run") and hasattr(task_obj, "retry"):
return task_obj.run, None, task_obj
# Case 3: 原始函数(装饰器未生效)
return task_obj, MagicMock(), None
return task_obj.run, None
# Case 3: 原始函数(CI 环境中装饰器未生效)
return task_obj, MagicMock()
# ── 成功场景 ──────────────────────────────────────────────
@@ -120,8 +110,8 @@ class TestProcessVoiceCloneSuccess:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, _ = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self is not None else ("profile-123",)
func, mock_self = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self else ("profile-123",)
result = func(*args)
assert result["ok"] is True
@@ -158,8 +148,8 @@ class TestProcessVoiceCloneSuccess:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, _ = _resolve_task(process_voice_clone)
args = (mock_self, "nonexistent") if mock_self is not None else ("nonexistent",)
func, mock_self = _resolve_task(process_voice_clone)
args = (mock_self, "nonexistent") if mock_self else ("nonexistent",)
result = func(*args)
assert result["ok"] is False
@@ -199,24 +189,24 @@ class TestProcessVoiceCloneTimeout:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, real_task = _resolve_task(process_voice_clone)
func, mock_self = _resolve_task(process_voice_clone)
if mock_self is not None:
# 设置 retry mock:根据环境不同,retry 在不同对象上
if mock_self is None:
# Proxy/Task 环境:retry 在 Task 实例上(func 是 bound method task.run)
real_task = process_voice_clone._get_current_object()
mock_retry = MagicMock()
mock_retry.side_effect = Retry("retrying")
with patch.object(real_task, "retry", mock_retry):
with pytest.raises(Retry):
func("profile-123")
mock_retry.assert_called_once()
else:
# 原始函数环境:retry 在 mock_self 上
mock_self.retry.side_effect = Retry("retrying")
with pytest.raises(Retry):
func(mock_self, "profile-123")
mock_self.retry.assert_called_once()
else:
# Proxy/Task 环境:retry 在 Task 实例上。用 _resolve_task 返回的 real_task,
# 避免再次 _get_current_object() 在无 context 时抛 AttributeError。
retry_target = real_task if real_task is not None else process_voice_clone
mock_retry = MagicMock()
mock_retry.side_effect = Retry("retrying")
with patch.object(retry_target, "retry", mock_retry):
with pytest.raises(Retry):
func("profile-123")
mock_retry.assert_called_once()
mock_session.rollback.assert_called_once()
mock_session.close.assert_called_once()
@@ -253,8 +243,8 @@ class TestProcessVoiceCloneFailure:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, _ = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self is not None else ("profile-123",)
func, mock_self = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self else ("profile-123",)
result = func(*args)
assert result["ok"] is False
@@ -287,8 +277,8 @@ class TestProcessVoiceCloneFailure:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, _ = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self is not None else ("profile-123",)
func, mock_self = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self else ("profile-123",)
result = func(*args)
assert result["ok"] is False
@@ -321,8 +311,8 @@ class TestProcessVoiceCloneFailure:
from worker_app.tasks.voice_clone import process_voice_clone
func, mock_self, _ = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self is not None else ("profile-123",)
func, mock_self = _resolve_task(process_voice_clone)
args = (mock_self, "profile-123") if mock_self else ("profile-123",)
result = func(*args)
assert result["ok"] is False