Compare commits
91 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| db621b4fcb | |||
| 60cacdf280 | |||
| 08de0d9946 | |||
| b0b81a5d60 | |||
| d959dd874f | |||
| dcd0c56827 | |||
| 585bab9313 | |||
| 4a449ae496 | |||
| 112f0eb277 | |||
| 2e2d1cd73e | |||
| 32473485d7 | |||
| 1591259bb8 | |||
| a1f25a4426 | |||
| 65a77e3fb6 | |||
| 9a57b0d5b8 | |||
| 4d98e98b57 | |||
| 81e1eb47fb | |||
| d3e4d6a07d | |||
| 0d6ce433d0 | |||
| eb2b009b33 | |||
| fbd89b4089 | |||
| 9b50e0696e | |||
| 9af73dcd86 | |||
| 6002f7a5e4 | |||
| 7e88440ca9 | |||
| fbf8844f25 | |||
| 34ffe14aae | |||
| 4fa3e4eb92 | |||
| a59a6a588a | |||
| f1621ace9f | |||
| f9daa08b2e | |||
| 66409fde6f | |||
| 3c016af076 | |||
| 428b9eeb3b | |||
| 43dfd6d425 | |||
| 73b9f7f97c | |||
| ed208ba3f7 | |||
| 99d6a47ec5 | |||
| 9e0959cb85 | |||
| 8fb75162d4 | |||
| 01f3e7d4c1 | |||
| 283f48d3f3 | |||
| f262b0cdd1 | |||
| 08362f556b | |||
| c2d8ebab4b | |||
| 1c9d5f81bb | |||
| 62720a4c70 | |||
| abf9a7fbdf | |||
| 54a3c50035 | |||
| f2bc951903 | |||
| bd2ac2d8c6 | |||
| 3c13b3837d | |||
| 3a757afbd3 | |||
| 7f44e82a27 | |||
| a61bda738a | |||
| da0b310102 | |||
| ea4b74216f | |||
| f8cf74f42d | |||
| fac88318bd | |||
| 46bce0dee8 | |||
| 20751428a8 | |||
| f655462fd6 | |||
| 787bb0ee31 | |||
| 7a2d235685 | |||
| 78c7fcca5c | |||
| 14b9a9d0fc | |||
| 132ca6bb70 | |||
| b2ba78c16c | |||
| 605a3eb841 | |||
| bf6c66d71c | |||
| 30bf66b307 | |||
| b68c29c69b | |||
| 4da3eae11a | |||
| c7a34fb297 | |||
| 115b428cb3 | |||
| cd4274553c | |||
| d825756c67 | |||
| b4724a866f | |||
| 731d3297b3 | |||
| ef6766dd58 | |||
| 0c41816a6d | |||
| ebb3c79d63 | |||
| 4f5ae52a40 | |||
| a425103b4f | |||
| 018e1bcb9b | |||
| 221eed2a25 | |||
| 35c00ccbb7 | |||
| 67a1ed6430 | |||
| ae733312db | |||
| 43a584d041 | |||
| c7662f0515 |
+34
-4
@@ -198,8 +198,38 @@ DOUBAO_TIMEOUT=30
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分扣点总开关:默认 false(对现有用户零影响)。
|
||||
# P2 阶段各业务路由逐个接入 @points_gate 时,用
|
||||
# `if settings.points_enabled: ...`
|
||||
# 包裹扣点逻辑;所有路由接入完成并验证通过后再在 staging/prod 打开。
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
# - false:生成视频/口型同步/数字人/AI标题/TTS/克隆音色等所有功能对登录
|
||||
# 用户免费放行,不扣积分、不做余额拦截;积分余额/流水/会员状态查询接口
|
||||
# 保留可用,但数据不再变动。积分相关的表、代码、接口均保留不删除。
|
||||
# - 恢复积分:设置 ENABLE_CREDIT_SYSTEM=true 即可,无需改代码。
|
||||
ENABLE_CREDIT_SYSTEM=false
|
||||
# 旧开关名(兼容别名):与 ENABLE_CREDIT_SYSTEM 任一为 true 即启用。
|
||||
POINTS_ENABLED=false
|
||||
|
||||
# ==================== 抖音解析多源轮询 (#1963) ====================
|
||||
# 无需配置 Key 也可使用(P0 免费源可用),配置 Key 可增加兜底能力
|
||||
|
||||
# TikHub API Key (https://tikhub.io) — $0.001/次起,注册送$0.05
|
||||
TIKHUB_API_KEY=
|
||||
|
||||
# apizero.cn API Key (https://v1.apizero.cn) — 国内抖音解析服务
|
||||
APIZERO_API_KEY=
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询口型同步)====================
|
||||
# GPU Worker 长期鉴权 Token,Worker 端 .env 的 GPU_WORKER_TOKEN 必须与此一致
|
||||
# 留空时 development 环境允许匿名访问(仅本地调试),staging/production 必须配置
|
||||
GPU_WORKER_TOKEN=
|
||||
# 单任务超时(秒),processing 超过此时长无任务心跳才回退 pending 或标记 failed
|
||||
# #1970:RTX2060 6G 推理 720p 长视频需 5 分钟以上,默认 900
|
||||
GPU_TASK_TIMEOUT_SECONDS=900
|
||||
# 是否启用 GPU 口型同步(开关)。开启后需同时有 Worker 在心跳窗口内(5分钟)才会走 GPU 路径;
|
||||
# 开关关闭 / 无可用 Worker / GPU 任务失败或超时 → 自动回退现有 MediaKit 云端 lipsync
|
||||
USE_GPU_LIPSYNC=false
|
||||
# 业务侧轮询 GPU 任务结果的间隔(秒)
|
||||
GPU_LIPSYNC_POLL_INTERVAL=5
|
||||
# 业务侧等待 GPU 任务总超时(秒);超时回退 MediaKit
|
||||
GPU_LIPSYNC_WAIT_TIMEOUT=1200
|
||||
# Worker 心跳新鲜度窗口(秒),last_heartbeat_at 在此窗口内视为在线
|
||||
GPU_WORKER_STALE_SECONDS=300
|
||||
|
||||
|
||||
@@ -1186,8 +1186,12 @@ jobs:
|
||||
DOUBAO_API_KEY: "${{ secrets.DOUBAO_API_KEY }}"
|
||||
DOUBAO_MODEL: "${{ secrets.DOUBAO_MODEL }}"
|
||||
DOUBAO_BASE_URL: "${{ secrets.DOUBAO_BASE_URL }}"
|
||||
DOUBAO_VISION_MODEL: "${{ secrets.DOUBAO_VISION_MODEL }}"
|
||||
WECHAT_APP_ID: "${{ secrets.WECHAT_APP_ID }}"
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
APIZERO_API_KEY: "${{ secrets.APIZERO_API_KEY }}"
|
||||
GPU_WORKER_TOKEN: "${{ secrets.GPU_WORKER_TOKEN }}"
|
||||
run: |
|
||||
set -eu
|
||||
echo "Rendering .env from template + secrets..."
|
||||
@@ -1288,6 +1292,14 @@ jobs:
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/.env"
|
||||
echo "✅ .env uploaded to staging server"
|
||||
|
||||
# 上传抖音 cookies 文件到 staging host(供容器挂载)
|
||||
echo "Uploading Douyin cookies to staging server..."
|
||||
ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" \
|
||||
"mkdir -p /var/lib/xiaoxia-saas-staging/configs"
|
||||
scp -P "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no deploy/configs/douyin_cookies.txt \
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/configs/douyin_cookies.txt"
|
||||
echo "✅ Douyin cookies uploaded"
|
||||
|
||||
# 通过环境变量传递凭证,避免命令行引号转义问题
|
||||
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh"
|
||||
|
||||
@@ -1630,8 +1642,12 @@ jobs:
|
||||
DOUBAO_API_KEY: "${{ secrets.DOUBAO_API_KEY }}"
|
||||
DOUBAO_MODEL: "${{ secrets.DOUBAO_MODEL }}"
|
||||
DOUBAO_BASE_URL: "${{ secrets.DOUBAO_BASE_URL }}"
|
||||
DOUBAO_VISION_MODEL: "${{ secrets.DOUBAO_VISION_MODEL }}"
|
||||
WECHAT_APP_ID: "${{ secrets.WECHAT_APP_ID }}"
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
APIZERO_API_KEY: "${{ secrets.APIZERO_API_KEY }}"
|
||||
GPU_WORKER_TOKEN: "${{ secrets.GPU_WORKER_TOKEN }}"
|
||||
run: |
|
||||
set -eu
|
||||
echo "Rendering .env from template + secrets..."
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""#1894: drop obsolete script title fields (title_text/title_category/title_config)
|
||||
|
||||
Revision ID: 078_drop_script_title_fields
|
||||
Revises: 077_merge_title_libs
|
||||
Create Date: 2026-09-16
|
||||
|
||||
口播文案(scripts)不再自带配套标题、标题分类和标题样式字段。
|
||||
智能剪辑 / AI 数字人等生成场景各自通过入参配置标题,不再从文案读取。
|
||||
保留字段:title(名称)、content(正文)、segments(分段)、tags(标签)。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "078_drop_script_title_fields"
|
||||
down_revision = "077_merge_title_libs"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
with op.batch_alter_table("scripts") as batch:
|
||||
batch.drop_column("title_config")
|
||||
batch.drop_column("title_category")
|
||||
batch.drop_column("title_text")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("scripts") as batch:
|
||||
batch.add_column(sa.Column("title_text", sa.String(500), nullable=False, server_default=""))
|
||||
batch.add_column(sa.Column("title_category", sa.String(50), nullable=False, server_default=""))
|
||||
batch.add_column(sa.Column("title_config", sa.JSON, nullable=False, server_default="{}"))
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add asset_atom_clips table
|
||||
|
||||
Revision ID: 079_asset_atom_clips
|
||||
Revises: 078_drop_script_title_fields
|
||||
Create Date: 2026-09-17
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "079_asset_atom_clips"
|
||||
down_revision = "078_drop_script_title_fields"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"asset_atom_clips",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column(
|
||||
"asset_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("assets.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("start_time", sa.Float(), nullable=False),
|
||||
sa.Column("end_time", sa.Float(), nullable=False),
|
||||
sa.Column("duration", sa.Float(), nullable=False),
|
||||
sa.Column("clip_index", sa.Integer(), nullable=False),
|
||||
sa.Column("tags", sa.JSON(), nullable=False, server_default=sa.text("'[]'")),
|
||||
sa.Column("scene_change_at", sa.Float(), nullable=True),
|
||||
sa.Column(
|
||||
"is_fallback",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("NOW()"),
|
||||
),
|
||||
)
|
||||
# 按素材查片段并按索引排序(复合索引前缀可独立用于 asset_id 过滤)
|
||||
op.create_index(
|
||||
"ix_asset_atom_clips_asset_index",
|
||||
"asset_atom_clips",
|
||||
["asset_id", "clip_index"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_asset_atom_clips_asset_index", table_name="asset_atom_clips")
|
||||
op.drop_table("asset_atom_clips")
|
||||
@@ -0,0 +1,37 @@
|
||||
"""add edit_plan_clips.atom_clip_id for #1970
|
||||
|
||||
Revision ID: 080_edit_plan_clips_atom_clip_id
|
||||
Revises: 079_asset_atom_clips
|
||||
Create Date: 2026-09-17
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "080_edit_plan_clips_atom_clip_id"
|
||||
down_revision = "079_asset_atom_clips"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"edit_plan_clips",
|
||||
sa.Column(
|
||||
"atom_clip_id",
|
||||
sa.String(36),
|
||||
nullable=False,
|
||||
server_default=sa.text("''"),
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_edit_plan_clips_atom_clip_id",
|
||||
"edit_plan_clips",
|
||||
["atom_clip_id"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_edit_plan_clips_atom_clip_id", table_name="edit_plan_clips")
|
||||
op.drop_column("edit_plan_clips", "atom_clip_id")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add gpu_lipsync_tasks and gpu_workers tables for MuseTalk reverse-poll worker
|
||||
|
||||
Revision ID: 081_add_gpu_lipsync
|
||||
Revises: 080_edit_plan_clips_atom_clip_id
|
||||
Create Date: 2026-09-18
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "081_add_gpu_lipsync"
|
||||
down_revision = "080_edit_plan_clips_atom_clip_id"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# GPU Worker 注册表
|
||||
op.create_table(
|
||||
"gpu_workers",
|
||||
sa.Column("worker_id", sa.String(100), primary_key=True),
|
||||
sa.Column("hostname", sa.String(200), nullable=False, server_default=""),
|
||||
sa.Column("gpu_name", sa.String(200), nullable=False, server_default=""),
|
||||
sa.Column("free_vram_mb", sa.Integer(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("capabilities", sa.String(500), nullable=False, server_default=""),
|
||||
sa.Column("last_heartbeat_at", sa.DateTime(), nullable=True, index=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# GPU 口型同步任务表
|
||||
op.create_table(
|
||||
"gpu_lipsync_tasks",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("lipsync_job_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("project_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("video_url", sa.Text(), nullable=False),
|
||||
sa.Column("audio_url", sa.Text(), nullable=False),
|
||||
sa.Column("result_url", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("result_duration", sa.Float(), nullable=False, server_default=sa.text("0.0")),
|
||||
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
|
||||
sa.Column("worker_id", sa.String(100), nullable=False, server_default="", index=True),
|
||||
sa.Column("attempt", sa.Integer(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("error_msg", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("started_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("finished_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("last_heartbeat_at", sa.DateTime(), nullable=True),
|
||||
)
|
||||
op.create_index("ix_gpu_lipsync_status_created", "gpu_lipsync_tasks", ["status", "created_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_gpu_lipsync_status_created", table_name="gpu_lipsync_tasks")
|
||||
op.drop_table("gpu_lipsync_tasks")
|
||||
op.drop_table("gpu_workers")
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add ai_tags to asset_atom_clips for #1970 fragment-level AI tagging
|
||||
|
||||
Revision ID: 082_atom_clip_ai_tags
|
||||
Revises: 081_add_gpu_lipsync
|
||||
Create Date: 2026-09-18
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "082_atom_clip_ai_tags"
|
||||
down_revision = "081_add_gpu_lipsync"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"asset_atom_clips",
|
||||
sa.Column("ai_tags", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("asset_atom_clips", "ai_tags")
|
||||
@@ -14,6 +14,7 @@ from app.api.routes.generation_cover import router as generation_cover_router
|
||||
from app.api.routes.generation_preview import router as generation_preview_router
|
||||
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.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
|
||||
@@ -211,3 +212,8 @@ api_router.include_router(
|
||||
prefix="/usage",
|
||||
tags=["Usage"],
|
||||
)
|
||||
api_router.include_router(
|
||||
gpu_lipsync_router,
|
||||
prefix="/gpu",
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
|
||||
@@ -25,12 +25,27 @@ def check_project_access(project_id: str, user_id: str, project_repository) -> N
|
||||
raise HTTPException(status_code=403, detail="无权访问该项目")
|
||||
|
||||
|
||||
_LEGACY_PLANS = {"standard", "pro", "enterprise", "basic", "premium"}
|
||||
|
||||
|
||||
def get_user_plan(user_id: str, user_repository: UserRepository) -> str:
|
||||
"""获取用户的订阅计划名称。"""
|
||||
"""获取用户的会员类型,兼容旧档位值。
|
||||
|
||||
旧档位 standard/pro/enterprise/basic/premium 统一映射到当前体系:
|
||||
- standard/basic → monthly
|
||||
- pro/premium/enterprise → quarterly
|
||||
"""
|
||||
user = user_repository.find_by_id(user_id)
|
||||
if user is None:
|
||||
return "free"
|
||||
return getattr(user, "subscription_plan", "free") or "free"
|
||||
plan = getattr(user, "subscription_plan", "free") or "free"
|
||||
if plan in {"standard", "basic"}:
|
||||
return "monthly"
|
||||
if plan in {"pro", "premium", "enterprise"}:
|
||||
return "quarterly"
|
||||
if plan not in {"free", "monthly", "quarterly", "yearly"}:
|
||||
return "free"
|
||||
return plan
|
||||
|
||||
|
||||
def require_project_and_library(
|
||||
|
||||
@@ -16,10 +16,12 @@ from app.core.task_enqueue import (
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_asset_repository,
|
||||
get_cosyvoice_service,
|
||||
get_db_session,
|
||||
get_generated_video_repository,
|
||||
get_generation_task_repository,
|
||||
get_project_repository,
|
||||
get_voice_clone_profile_repository,
|
||||
)
|
||||
from app.schemas.generated_video import (
|
||||
GeneratedVideoResponse,
|
||||
@@ -132,6 +134,8 @@ def _select_assets_from_library(
|
||||
mode: str,
|
||||
count: int,
|
||||
rng=None,
|
||||
script_tags: list | None = None,
|
||||
tag_names_by_id: dict | None = None,
|
||||
) -> list[str]:
|
||||
"""根据选取模式从素材库中选取 ready 状态的视频素材 ID。
|
||||
|
||||
@@ -141,6 +145,8 @@ def _select_assets_from_library(
|
||||
count: 选取数量,0 表示全部(仅 smart 模式有效)
|
||||
rng: 可选随机源(smart 模式排序噪声用),生产环境不传则内部随机;
|
||||
测试可注入固定种子或零噪声随机源获得确定性结果。
|
||||
script_tags: #1970 叙事模式文案标签;非空时标签命中素材优先,不足再用其余素材兜底。
|
||||
tag_names_by_id: asset_id → 素材标签名列表(素材只存 tag_ids 时由调用方查名称注入)。
|
||||
|
||||
Returns:
|
||||
选中的素材 ID 列表
|
||||
@@ -150,6 +156,20 @@ def _select_assets_from_library(
|
||||
if not ready_video_assets:
|
||||
return []
|
||||
|
||||
# 叙事模式(#1970 PR3):文案标签命中池优先;无任何命中时完全降级为现有随机逻辑。
|
||||
if script_tags:
|
||||
from packages.domain.narrative_match import pick_narrative_assets
|
||||
|
||||
limit = count if count > 0 else None
|
||||
picked = pick_narrative_assets(
|
||||
ready_video_assets,
|
||||
script_tags=script_tags,
|
||||
tag_names_by_id=tag_names_by_id,
|
||||
limit=limit,
|
||||
rng=rng,
|
||||
)
|
||||
return [a.id for a in picked]
|
||||
|
||||
if mode == "smart":
|
||||
# 智能匹配:统一使用 packages/domain/smart_match.py 的多维评分+多样性选取
|
||||
# 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%)
|
||||
@@ -162,16 +182,78 @@ def _select_assets_from_library(
|
||||
return [a.id for a in ready_video_assets]
|
||||
|
||||
|
||||
# #1970 PR3:video_ratio → 默认输出分辨率(显式 output_width/output_height 优先)
|
||||
_VIDEO_RATIO_DIMENSIONS = {
|
||||
"9:16": (1080, 1920),
|
||||
"16:9": (1920, 1080),
|
||||
"1:1": (1080, 1080),
|
||||
"3:4": (1080, 1440),
|
||||
"4:3": (1440, 1080),
|
||||
}
|
||||
|
||||
|
||||
def _resolve_output_dimensions(request: CreateGenerationTaskRequest) -> tuple[int, int]:
|
||||
"""解析输出分辨率:显式 output_width/output_height 非旧默认值时优先,否则按 video_ratio。
|
||||
|
||||
前端 #1973 总是同时传 video_ratio 与具体分辨率,两者一致;此函数主要服务
|
||||
只传比例的调用方,并保证旧调用(不传比例)维持 1280x720 行为。
|
||||
"""
|
||||
width, height = request.output_width, request.output_height
|
||||
ratio = (request.video_ratio or "").strip()
|
||||
if ratio in _VIDEO_RATIO_DIMENSIONS and (width, height) == (1280, 720):
|
||||
return _VIDEO_RATIO_DIMENSIONS[ratio]
|
||||
return width, height
|
||||
|
||||
|
||||
def _load_asset_tag_names(db: Session, assets: list, user_id: str) -> dict[str, list[str]]:
|
||||
"""叙事模式:查 TagModel 名称,构造 asset_id → 标签名列表(失败返回空 dict 降级随机)。"""
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetTagModel, TagModel
|
||||
|
||||
tag_ids = {tid for a in assets for tid in (getattr(a, "tag_ids", None) or [])}
|
||||
if not tag_ids:
|
||||
return {}
|
||||
name_rows = (
|
||||
db.query(TagModel.id, TagModel.name).filter(TagModel.id.in_(tag_ids), TagModel.user_id == user_id).all()
|
||||
)
|
||||
name_by_id = {row.id: row.name for row in name_rows}
|
||||
links = db.query(AssetTagModel.asset_id, AssetTagModel.tag_id).filter(AssetTagModel.tag_id.in_(tag_ids)).all()
|
||||
index: dict[str, list[str]] = {}
|
||||
for asset_id, tag_id in links:
|
||||
name = name_by_id.get(tag_id)
|
||||
if name:
|
||||
index.setdefault(asset_id, []).append(name)
|
||||
return index
|
||||
except Exception: # noqa: BLE001 - 标签匹配是加分项,查询失败不阻断生成
|
||||
logger.warning("[叙事模式] 素材标签查询失败,降级随机选片", exc_info=True)
|
||||
return {}
|
||||
|
||||
|
||||
def _writeback_edit_plan_config(
|
||||
plan_id: str,
|
||||
task_id: str,
|
||||
title_config: dict | None,
|
||||
db: Session,
|
||||
dedup_enabled: bool | None = None,
|
||||
video_index: int | None = None,
|
||||
assembly_mode: str | None = None,
|
||||
script_id: str | None = None,
|
||||
video_ratio: str | None = None,
|
||||
) -> None:
|
||||
"""[已下沉] 路由层兼容别名 → app.services.generation_common.writeback_edit_plan_config。"""
|
||||
from app.services.generation_common import writeback_edit_plan_config
|
||||
|
||||
return writeback_edit_plan_config(plan_id, task_id, title_config, db)
|
||||
return writeback_edit_plan_config(
|
||||
plan_id,
|
||||
task_id,
|
||||
title_config,
|
||||
db,
|
||||
dedup_enabled=dedup_enabled,
|
||||
video_index=video_index,
|
||||
assembly_mode=assembly_mode,
|
||||
script_id=script_id,
|
||||
video_ratio=video_ratio,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_project_and_library(
|
||||
@@ -221,16 +303,63 @@ def create_generation_task(
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
db: Session = Depends(get_db_session),
|
||||
cosyvoice_service: Any = Depends(get_cosyvoice_service),
|
||||
voice_clone_repository: Any = Depends(get_voice_clone_profile_repository),
|
||||
) -> BatchGenerationTaskResponse:
|
||||
logger.info(
|
||||
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, count=%d",
|
||||
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, assembly=%s, count=%d",
|
||||
authenticated_user.user.id,
|
||||
request.template_id,
|
||||
len(request.asset_ids),
|
||||
request.asset_select_mode,
|
||||
request.assembly_mode,
|
||||
request.count,
|
||||
)
|
||||
|
||||
# video_ratio → 默认分辨率(显式分辨率优先)
|
||||
request.output_width, request.output_height = _resolve_output_dimensions(request)
|
||||
|
||||
# ── #1970 PR3 叙事模式:入队前同步合成配音并落为 audio asset ──
|
||||
# 合成结果覆盖 voice_library_id(下游按 audio asset id 消费),失败直接 4xx 不入队。
|
||||
narrative_script_tags: list = []
|
||||
if request.assembly_mode == "narrative":
|
||||
from app.config import settings as _settings
|
||||
from app.services.narrative_service import NarrativeError, prepare_narrative_voice
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.tts_job_repository import SQLAlchemyTTSJobRepository
|
||||
|
||||
try:
|
||||
narrative_ctx = prepare_narrative_voice(
|
||||
db=db,
|
||||
user_id=authenticated_user.user.id,
|
||||
script_id=request.script_id,
|
||||
tts_voice_id=request.tts_voice_id,
|
||||
tts_voice_source=request.tts_voice_source,
|
||||
tts_repository=SQLAlchemyTTSJobRepository(db),
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
voice_clone_repository=voice_clone_repository,
|
||||
asset_repository=asset_repository,
|
||||
asset_library_repository=asset_library_repository,
|
||||
project_repository=project_repository,
|
||||
storage_service=get_storage_service(),
|
||||
points_enabled=bool(getattr(_settings, "points_enabled", False)),
|
||||
is_member=bool(getattr(authenticated_user.user, "is_member", False)),
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
except NarrativeError as e:
|
||||
logger.warning("[叙事模式] 配音前置处理失败: %s", e.message)
|
||||
raise HTTPException(status_code=e.status_code, detail=e.message) from e
|
||||
|
||||
request.voice_library_id = narrative_ctx.voice_asset_id
|
||||
narrative_script_tags = list(getattr(narrative_ctx.script, "tags", None) or [])
|
||||
logger.info(
|
||||
"[叙事模式] 配音已就绪: script_id=%s, tts_job=%s, voice_asset=%s, duration=%.2f",
|
||||
request.script_id,
|
||||
narrative_ctx.tts_job_id,
|
||||
narrative_ctx.voice_asset_id,
|
||||
narrative_ctx.audio_duration,
|
||||
)
|
||||
|
||||
try:
|
||||
project_id, asset_library_id = _resolve_project_and_library(
|
||||
request, project_repository, asset_library_repository, asset_repository, authenticated_user
|
||||
@@ -256,19 +385,29 @@ def create_generation_task(
|
||||
|
||||
# 素材库自动匹配:当未显式指定 asset_ids 时,按模式自动选取
|
||||
if not resolved_asset_ids:
|
||||
_tag_index = (
|
||||
_load_asset_tag_names(db, assets, authenticated_user.user.id) if narrative_script_tags else None
|
||||
)
|
||||
resolved_asset_ids = _select_assets_from_library(
|
||||
assets,
|
||||
mode=request.asset_select_mode,
|
||||
count=request.asset_select_count,
|
||||
script_tags=narrative_script_tags or None,
|
||||
tag_names_by_id=_tag_index,
|
||||
)
|
||||
elif project_id and not resolved_asset_ids and request.asset_select_mode in ("smart",):
|
||||
# 项目级模式:未指定 asset_ids 且选择了 smart 模式时,也自动选取
|
||||
elif project_id and not resolved_asset_ids and (request.asset_select_mode in ("smart",) or narrative_script_tags):
|
||||
# 项目级模式:未指定 asset_ids 且选择了 smart 模式(或叙事模式按标签匹配)时自动选取
|
||||
assets = asset_repository.find_by_project(project_id)
|
||||
if assets:
|
||||
_tag_index = (
|
||||
_load_asset_tag_names(db, assets, authenticated_user.user.id) if narrative_script_tags else None
|
||||
)
|
||||
resolved_asset_ids = _select_assets_from_library(
|
||||
assets,
|
||||
mode=request.asset_select_mode,
|
||||
count=request.asset_select_count,
|
||||
script_tags=narrative_script_tags or None,
|
||||
tag_names_by_id=_tag_index,
|
||||
)
|
||||
if not resolved_asset_ids:
|
||||
raise HTTPException(
|
||||
@@ -332,6 +471,10 @@ def create_generation_task(
|
||||
task_id=preview_task.id,
|
||||
title_config=fallback_title_config,
|
||||
db=db,
|
||||
dedup_enabled=request.dedup_enabled,
|
||||
assembly_mode=request.assembly_mode,
|
||||
script_id=request.script_id or None,
|
||||
video_ratio=request.video_ratio or None,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
@@ -476,9 +619,12 @@ def create_generation_task(
|
||||
variant_plan_ids.append(_plan0.id)
|
||||
|
||||
# #1855 P0:批次区间避让表,从变体0实际clips构建初始值(公共函数)
|
||||
from app.services.generation_common import collect_plan_atom_clip_ids as _collect_atom_ids
|
||||
from app.services.generation_common import collect_plan_segments as _collect_segments
|
||||
|
||||
_batch_segments = _collect_segments(_plan0.id, _plan_svc._clip_repo)
|
||||
# #1970:批次内原子片段硬避让集合
|
||||
_batch_atom_ids: list[str] = _collect_atom_ids(_plan0.id, _plan_svc._clip_repo)
|
||||
|
||||
# 变体 1..N-1 独立选片(传入累积batch_segments做素材区间避让)
|
||||
for task_index in range(1, count):
|
||||
@@ -493,6 +639,7 @@ def create_generation_task(
|
||||
name_suffix=f"批量{task_index + 1}",
|
||||
voice_duration=voice_durations[task_index] if task_index < len(voice_durations) else 0.0,
|
||||
batch_segments=_batch_segments,
|
||||
batch_used_atom_ids=_batch_atom_ids,
|
||||
)
|
||||
break
|
||||
except ValueError as ve:
|
||||
@@ -529,6 +676,8 @@ def create_generation_task(
|
||||
_new_segs = _collect_segments(variant.id, _plan_svc._clip_repo)
|
||||
for _aid, _ivs in _new_segs.items():
|
||||
_batch_segments.setdefault(_aid, []).extend(_ivs)
|
||||
# #1970:同步累积原子片段ID
|
||||
_batch_atom_ids.extend(_collect_atom_ids(variant.id, _plan_svc._clip_repo))
|
||||
except Exception:
|
||||
logger.exception("[生成任务] 变体%d 区间收集失败(不阻断)", task_index)
|
||||
|
||||
@@ -672,6 +821,11 @@ def create_generation_task(
|
||||
task_id=task.id,
|
||||
title_config=variant_title_config,
|
||||
db=db,
|
||||
dedup_enabled=request.dedup_enabled,
|
||||
video_index=task_index,
|
||||
assembly_mode=request.assembly_mode,
|
||||
script_id=request.script_id or None,
|
||||
video_ratio=request.video_ratio or None,
|
||||
)
|
||||
|
||||
if safe_enqueue_generation_task(
|
||||
@@ -762,6 +916,7 @@ def confirm_generation(
|
||||
generation_task_repository.update(source_task)
|
||||
|
||||
# 同步标题到 EditPlan.config
|
||||
# #1970:确认生成复用预览计划,dedup_enabled 沿用计划已有值,不在此覆盖
|
||||
if confirmed_title_config and source_task.source_edit_plan_id:
|
||||
_writeback_edit_plan_config(
|
||||
plan_id=source_task.source_edit_plan_id,
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""GPU MuseTalk Worker 反向轮询路由 — /api/v1/gpu/lipsync/*.
|
||||
|
||||
仅面向部署在用户 RTX2060 本地的 GPU Worker 脚本,不面向前端用户。
|
||||
鉴权方式:长期 API Token(`Authorization: Bearer <GPU_WORKER_TOKEN>`),不走用户 JWT。
|
||||
|
||||
接口:
|
||||
POST /api/v1/gpu/register Worker 注册/心跳
|
||||
GET /api/v1/gpu/lipsync/poll Worker 轮询拉任务(无任务返回 204)
|
||||
POST /api/v1/gpu/lipsync/result Worker multipart 上传结果视频/上报失败
|
||||
GET /api/v1/gpu/lipsync/status/{id} 业务侧查询任务状态(内部接口,暂开放给登录用户)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tempfile
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.gpu_lipsync import (
|
||||
GpuLipsyncPollResponse,
|
||||
GpuLipsyncResultResponse,
|
||||
GpuLipsyncStatusResponse,
|
||||
GpuLipsyncTaskPayload,
|
||||
GpuWorkerRegisterRequest,
|
||||
GpuWorkerRegisterResponse,
|
||||
)
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
from fastapi import (
|
||||
APIRouter,
|
||||
Depends,
|
||||
File,
|
||||
Form,
|
||||
HTTPException,
|
||||
Query,
|
||||
Request,
|
||||
UploadFile,
|
||||
status,
|
||||
)
|
||||
from fastapi.responses import Response
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 复用 bearer scheme 抽 Token,但不校验用户 JWT
|
||||
_gpu_bearer = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
def _verify_gpu_token(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Depends(_gpu_bearer),
|
||||
) -> str:
|
||||
"""校验 GPU Worker Token,返回 worker 提供的 token 串(仅用于日志,不做身份识别).
|
||||
|
||||
- development 且未配置 token → 直接放行(方便本地调试)。
|
||||
- production/staging 未配置 token → 拒绝(避免裸奔)。
|
||||
- token 不匹配 → 401。
|
||||
"""
|
||||
settings = get_api_settings()
|
||||
expected = (settings.gpu_worker_token or "").strip()
|
||||
is_dev = settings.environment == "development"
|
||||
if not expected:
|
||||
if is_dev:
|
||||
return credentials.credentials if credentials else ""
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="GPU_WORKER_TOKEN not configured on server",
|
||||
)
|
||||
if credentials is None or credentials.scheme.lower() != "bearer":
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Missing bearer token")
|
||||
if credentials.credentials != expected:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid GPU worker token")
|
||||
return credentials.credentials
|
||||
|
||||
|
||||
def _get_svc(db=Depends(get_db_session)) -> GpuLipsyncService:
|
||||
return GpuLipsyncService(db)
|
||||
|
||||
|
||||
# ── POST /register — Worker 注册/心跳 ──────────────────────────────
|
||||
|
||||
|
||||
@router.post("/register", response_model=GpuWorkerRegisterResponse)
|
||||
def register_worker(
|
||||
body: GpuWorkerRegisterRequest,
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
svc.register_worker(
|
||||
worker_id=body.worker_id,
|
||||
hostname=body.hostname,
|
||||
gpu_name=body.gpu_name,
|
||||
free_vram_mb=body.free_vram_mb,
|
||||
capabilities=body.capabilities,
|
||||
task_id=body.task_id,
|
||||
)
|
||||
return GpuWorkerRegisterResponse(ok=True, server_time=datetime.now(UTC), message="ok")
|
||||
|
||||
|
||||
# ── GET /lipsync/poll — Worker 轮询拉任务 ─────────────────────────
|
||||
|
||||
|
||||
@router.get("/lipsync/poll")
|
||||
def poll_task(
|
||||
worker_id: str = Query(..., min_length=1, max_length=100, description="Worker 唯一 ID"),
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
task = svc.poll_task(worker_id=worker_id)
|
||||
if task is None:
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
payload = GpuLipsyncTaskPayload(
|
||||
task_id=task.id,
|
||||
video_url=getattr(task, "_signed_video_url", task.video_url),
|
||||
audio_url=getattr(task, "_signed_audio_url", task.audio_url),
|
||||
lipsync_job_id=task.lipsync_job_id or "",
|
||||
user_id=task.user_id or "",
|
||||
project_id=task.project_id or "",
|
||||
created_at=task.created_at,
|
||||
upload_url=getattr(task, "_signed_upload_url", ""),
|
||||
upload_method="PUT",
|
||||
expires_at=getattr(task, "_upload_expires_at", datetime.now(UTC)),
|
||||
)
|
||||
return GpuLipsyncPollResponse(task=payload)
|
||||
|
||||
|
||||
# ── POST /lipsync/result — Worker 上报结果(multipart) ─────────────
|
||||
|
||||
|
||||
@router.post("/lipsync/result", response_model=GpuLipsyncResultResponse)
|
||||
async def report_result(
|
||||
request: Request,
|
||||
task_id: str = Form(...),
|
||||
worker_id: str = Form(...),
|
||||
success: bool = Form(True),
|
||||
duration_seconds: float = Form(0.0),
|
||||
error_msg: str = Form(""),
|
||||
result: Optional[UploadFile] = File(None),
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
# 参数校验:
|
||||
# - success=true + result 文件 → API 代为上传到 OSS(方便 Worker 端实现)
|
||||
# - success=true + 无文件 → Worker 已经自己 PUT 到预签名 upload_url,直接确认
|
||||
# - success=false → 不上传文件,错误信息通过 error_msg 传递
|
||||
if success and result is not None:
|
||||
# 把文件落盘到临时目录,然后 PUT 到预签名 URL
|
||||
storage = get_storage_service()
|
||||
result_key = svc._result_key(task_id)
|
||||
upload_url = storage.get_upload_url(result_key, expires_seconds=3600, content_type="video/mp4")
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="gpu_result_") as tmpdir:
|
||||
tmp_path = Path(tmpdir) / "result.mp4"
|
||||
content = await result.read()
|
||||
if not content:
|
||||
raise HTTPException(status_code=400, detail="上传的 result 文件为空")
|
||||
tmp_path.write_bytes(content)
|
||||
headers = {"Content-Type": "video/mp4"}
|
||||
with open(tmp_path, "rb") as f:
|
||||
resp = requests.put(upload_url, data=f, headers=headers, timeout=300)
|
||||
if resp.status_code >= 400:
|
||||
logger.error(
|
||||
"上传 GPU 结果到 OSS 失败: status=%d body=%s",
|
||||
resp.status_code,
|
||||
resp.text[:500],
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail=f"上传结果视频到 OSS 失败 (HTTP {resp.status_code})",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.exception("上传 GPU 结果视频异常: %s", exc)
|
||||
raise HTTPException(status_code=500, detail=f"上传结果视频异常: {exc}") from exc
|
||||
elif not success:
|
||||
# 失败时忽略 result 文件(即便传了也没用)
|
||||
pass
|
||||
# 其他情况:success=true 且无文件 → Worker 已自行 PUT 到预签名 URL,直接标记完成
|
||||
|
||||
try:
|
||||
task = svc.report_result(
|
||||
task_id=task_id,
|
||||
worker_id=worker_id,
|
||||
success=success,
|
||||
duration_seconds=duration_seconds,
|
||||
error_msg=error_msg,
|
||||
)
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
return GpuLipsyncResultResponse(
|
||||
ok=True,
|
||||
task_id=task.id,
|
||||
status=task.status,
|
||||
message="ok",
|
||||
)
|
||||
|
||||
|
||||
# ── GET /lipsync/status/{task_id} — 业务侧查询状态 ─────────────────
|
||||
# 说明:此接口会被 lipsync_service 内部在业务流程里直接读 DB,不通过 HTTP。
|
||||
# 但仍暴露一个简单查询接口,方便调试和前端轮询(如后续需要)。暂不做用户权限校验,
|
||||
# task_id 本身是 UUID,不可枚举。
|
||||
|
||||
|
||||
@router.get("/lipsync/status/{task_id}", response_model=GpuLipsyncStatusResponse)
|
||||
def get_task_status(
|
||||
task_id: str,
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
):
|
||||
task = svc.get_task(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return GpuLipsyncStatusResponse(
|
||||
task_id=task.id,
|
||||
status=task.status,
|
||||
result_url=task.result_url,
|
||||
result_duration=task.result_duration,
|
||||
error_msg=task.error_msg,
|
||||
worker_id=task.worker_id,
|
||||
attempt=task.attempt,
|
||||
created_at=task.created_at,
|
||||
started_at=task.started_at,
|
||||
finished_at=task.finished_at,
|
||||
)
|
||||
@@ -295,7 +295,14 @@ def get_lipsync_job(
|
||||
from datetime import datetime as _dt
|
||||
|
||||
_now = _dt.now(UTC)
|
||||
_stale = job.updated_at is None or (_now - job.updated_at).total_seconds() > 30
|
||||
_upd = job.updated_at
|
||||
# DB 返回的 DateTime 列可能是 naive(取决于方言/驱动):代码写入统一用
|
||||
# datetime.now(UTC),经 SQLAlchemy 存入 TIMESTAMP WITHOUT TIMEZONE 后再
|
||||
# 读回就是 UTC wall clock 的 naive datetime,直接补 UTC tz 即可;避免
|
||||
# TypeError: can't subtract offset-naive and offset-aware datetimes。
|
||||
if _upd is not None and _upd.tzinfo is None:
|
||||
_upd = _upd.replace(tzinfo=UTC)
|
||||
_stale = _upd is None or (_now - _upd).total_seconds() > 30
|
||||
if _stale:
|
||||
try:
|
||||
refreshed = svc.refresh_job_status(job_id, current_user.user.id)
|
||||
|
||||
@@ -12,6 +12,7 @@ from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.points import (
|
||||
DailyUsageResponse,
|
||||
@@ -44,6 +45,12 @@ from packages.domain.points_service import PointsService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _credits_enabled() -> bool:
|
||||
"""积分系统总开关(ENABLE_CREDIT_SYSTEM),关闭时全部功能免费放行。"""
|
||||
return bool(getattr(settings, "points_enabled", False))
|
||||
|
||||
|
||||
# ── 两个 router ──
|
||||
points_router = APIRouter()
|
||||
usage_router = APIRouter()
|
||||
@@ -172,6 +179,19 @@ def check_points(
|
||||
"valid_scenes": sorted(POINTS_SCENES.keys()),
|
||||
},
|
||||
)
|
||||
|
||||
# 积分系统暂停(ENABLE_CREDIT_SYSTEM=false):所有场景直接放行,需 0 积分
|
||||
if not _credits_enabled():
|
||||
svc = _get_service()
|
||||
account = svc.get_or_create_account(current_user.user.id, db)
|
||||
return PointsCheckResponse(
|
||||
allowed=True,
|
||||
required_points=0,
|
||||
current_balance=account["balance"],
|
||||
remaining_after=account["balance"],
|
||||
is_free_quota=False,
|
||||
)
|
||||
|
||||
is_mem = _is_member(current_user)
|
||||
mt = _member_type(current_user)
|
||||
|
||||
@@ -209,8 +229,19 @@ def deduct_points(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""积分扣减(内部服务调用)。"""
|
||||
"""积分扣减(内部服务调用)。
|
||||
|
||||
积分系统暂停(ENABLE_CREDIT_SYSTEM=false)时为 no-op:不扣分、余额不变,
|
||||
直接返回成功,保证内部调用方拿到 success=True 继续业务流程。
|
||||
"""
|
||||
svc = _get_service()
|
||||
if not _credits_enabled():
|
||||
account = svc.get_or_create_account(current_user.user.id, db)
|
||||
return SimpleMessageResponse(
|
||||
success=True,
|
||||
message="积分系统已暂停,未扣减积分",
|
||||
data={"transaction_id": "", "balance": account["balance"]},
|
||||
)
|
||||
result = svc.deduct_points(
|
||||
user_id=current_user.user.id,
|
||||
amount=body.amount,
|
||||
@@ -243,11 +274,7 @@ def refund_points(
|
||||
"""积分退还(内部服务调用)。"""
|
||||
from packages.adapters.sqlalchemy_impl.models import PointsTransactionModel
|
||||
|
||||
txn = (
|
||||
db.query(PointsTransactionModel)
|
||||
.filter(PointsTransactionModel.id == body.transaction_id)
|
||||
.first()
|
||||
)
|
||||
txn = db.query(PointsTransactionModel).filter(PointsTransactionModel.id == body.transaction_id).first()
|
||||
if txn is None:
|
||||
raise HTTPException(status_code=404, detail="交易记录不存在")
|
||||
if txn.user_id != current_user.user.id:
|
||||
|
||||
@@ -36,9 +36,6 @@ def _to_response(script) -> ScriptResponse:
|
||||
for s in segments
|
||||
],
|
||||
tags=script.tags or [],
|
||||
title_text=getattr(script, "title_text", "") or "",
|
||||
title_category=getattr(script, "title_category", "") or "",
|
||||
title_config=getattr(script, "title_config", None) or {},
|
||||
created_at=script.created_at,
|
||||
updated_at=script.updated_at,
|
||||
)
|
||||
@@ -73,9 +70,6 @@ def create_script(
|
||||
content=request.content,
|
||||
segments=[s.model_dump() for s in request.segments],
|
||||
tags=request.tags,
|
||||
title_text=request.title_text or "",
|
||||
title_category=request.title_category or "",
|
||||
title_config=request.title_config or {},
|
||||
)
|
||||
return _to_response(script)
|
||||
|
||||
@@ -110,9 +104,6 @@ def update_script(
|
||||
content=request.content,
|
||||
segments=[s.model_dump() for s in request.segments] if request.segments is not None else None,
|
||||
tags=request.tags,
|
||||
title_text=request.title_text,
|
||||
title_category=request.title_category,
|
||||
title_config=request.title_config,
|
||||
)
|
||||
except ScriptNotFoundError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
|
||||
|
||||
@@ -1,7 +1,12 @@
|
||||
"""Scripts AI 能力路由 — Issue #1893.
|
||||
"""Scripts AI 能力路由 — Issue #1893/#1963.
|
||||
|
||||
三个 AI 工具接口(均挂载在 /api/v1/scripts 前缀下):
|
||||
- POST /extract-from-douyin 从抖音视频提取文案(yt-dlp 下载 + ASR 转写)
|
||||
- POST /extract-from-douyin 从抖音视频提取文案
|
||||
- 入口自动从分享文本中正则提取 http(s) URL,兼容 "复制链接" 粘贴场景
|
||||
- 多源轮询解析(douyin_resolver):App Feed API → TikHub → apizero
|
||||
- 拿到 MP4 直链后优先走火山 MediaKit ASR,失败回退下载+本地 ASR
|
||||
- ASR 空结果时使用 Feed desc 兜底,图文视频直接返回 desc
|
||||
- 所有源均失败时返回具体错误信息(不暴露内部细节)
|
||||
- POST /ai-rewrite AI 文案改写(复用豆包 LLM)
|
||||
- POST /ai-generate-titles AI 标题生成(复用 generate_smart_titles)
|
||||
"""
|
||||
@@ -12,6 +17,8 @@ import logging
|
||||
import os
|
||||
import re
|
||||
import tempfile
|
||||
import time
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session
|
||||
@@ -23,6 +30,12 @@ from app.schemas.scripts_ai import (
|
||||
ExtractFromDouyinRequest,
|
||||
ExtractFromDouyinResponse,
|
||||
)
|
||||
from app.services.douyin_resolver import available_providers, resolve_douyin_video
|
||||
from app.services.mediakit_client import (
|
||||
MediaKitClient,
|
||||
MediaKitError,
|
||||
get_mediakit_client,
|
||||
)
|
||||
from app.services.script_asr_service import (
|
||||
ASRNotConfiguredError,
|
||||
ASRTranscriptionError,
|
||||
@@ -38,253 +51,504 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 抖音 URL 校验:支持短链 v.douyin.com 和长链 www.douyin.com/video/
|
||||
_DOUYIN_URL_RE = re.compile(
|
||||
r"^(https?://)?(v\.douyin\.com/\S+|www\.douyin\.com/video/\S+)$",
|
||||
_DOUYIN_DEBUG_ERRORS = os.environ.get("DOUYIN_DEBUG_ERRORS", "").lower() in (
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
) or os.environ.get(
|
||||
"APP_ENV", ""
|
||||
).lower() in ("staging", "dev", "development", "test")
|
||||
|
||||
_TAIL_PUNCT = ".,;:!?,。;:!?))]》" + chr(34) + chr(39) + "<>"
|
||||
_URL_EXTRACT_RE = re.compile(r"https?://\S+", re.IGNORECASE)
|
||||
_DOUYIN_HOST_RE = re.compile(
|
||||
r"(^|\.)(douyin\.com|iesdouyin\.com|amemv\.com)$",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_ANY_SCHEME_RE = re.compile(r"^[a-z][a-z0-9+.-]*://\S+", re.IGNORECASE)
|
||||
|
||||
|
||||
def _validate_douyin_url(url: str) -> None:
|
||||
"""校验抖音 URL 格式,不合法时抛 HTTPException(400)."""
|
||||
if not url or not url.strip():
|
||||
def _dbg(key, val):
|
||||
logger.debug("douyin_extract %s=%s", key, str(val)[:200])
|
||||
|
||||
|
||||
def _extract_url_from_text(raw):
|
||||
if not raw:
|
||||
return None
|
||||
m = _URL_EXTRACT_RE.search(raw)
|
||||
if m:
|
||||
return m.group(0).rstrip(_TAIL_PUNCT)
|
||||
short = re.search(
|
||||
r"(?:^|(?<![a-z0-9/:]))((?:v|www)\.douyin\.com/\S+|douyin\.com/(?:video|note)/\S+)",
|
||||
raw,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
if short:
|
||||
return "https://" + short.group(1).rstrip(_TAIL_PUNCT)
|
||||
return None
|
||||
|
||||
|
||||
def _extract_and_validate_douyin_url(raw_input):
|
||||
raw = (raw_input or "").strip()
|
||||
if not raw:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="链接不能为空")
|
||||
|
||||
url = _extract_url_from_text(raw)
|
||||
|
||||
if not url:
|
||||
if _ANY_SCHEME_RE.search(raw):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无效的抖音链接,仅支持 http(s) 协议",
|
||||
)
|
||||
short = re.search(
|
||||
r"(?:^|(?<![a-z0-9]))((?:v|www)\.douyin\.com/\S+|douyin\.com/(?:video|note)/\S+)",
|
||||
raw,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
if short:
|
||||
url = "https://" + short.group(1).rstrip(_TAIL_PUNCT)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="未在输入中找到有效抖音链接,请粘贴包含 v.douyin.com 或 www.douyin.com 的分享文本",
|
||||
)
|
||||
|
||||
if not re.match(r"^https?://", url, re.IGNORECASE):
|
||||
url = "https://" + url
|
||||
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
host = parsed.hostname or ""
|
||||
scheme = (parsed.scheme or "").lower()
|
||||
except Exception:
|
||||
host = ""
|
||||
scheme = ""
|
||||
if scheme not in ("http", "https"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="链接不能为空",
|
||||
detail="无效的抖音链接,仅支持 http(s) 协议",
|
||||
)
|
||||
if not _DOUYIN_URL_RE.match(url.strip()):
|
||||
if not _DOUYIN_HOST_RE.search(host):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无效的抖音链接,仅支持 v.douyin.com 短链或 www.douyin.com/video/ 长链",
|
||||
detail="无效的抖音链接,仅支持 douyin.com 域名(v.douyin.com 短链或 www.douyin.com 长链)",
|
||||
)
|
||||
return url
|
||||
|
||||
|
||||
# ── 1. 从抖音视频提取文案 ─────────────────────────────────────────────────────
|
||||
# ── MediaKitClient ASR 扩展(monkey patch) ────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/extract-from-douyin",
|
||||
response_model=ExtractFromDouyinResponse,
|
||||
)
|
||||
def _mk_post_json(self, path, payload):
|
||||
import httpx
|
||||
|
||||
if not self.is_available:
|
||||
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
|
||||
url = self._base_url + path
|
||||
try:
|
||||
with httpx.Client(timeout=self._timeout) as http:
|
||||
resp = http.post(url, headers=self._headers(), json=payload)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except httpx.TimeoutException as exc:
|
||||
raise MediaKitError("MediaKit API 超时 (%ss)" % self._timeout, code="Timeout") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise MediaKitError(
|
||||
"MediaKit API HTTP %s: %s" % (exc.response.status_code, exc.response.text[:300]),
|
||||
code="HttpError",
|
||||
) from exc
|
||||
except httpx.RequestError as exc:
|
||||
raise MediaKitError("MediaKit API 网络错误: %s" % exc, code="NetworkError") from exc
|
||||
if data.get("success") is False and data.get("error"):
|
||||
err = data["error"] if isinstance(data["error"], dict) else {"message": str(data["error"])}
|
||||
raise MediaKitError(
|
||||
err.get("message", "请求失败"),
|
||||
code=err.get("code", "RequestFailed"),
|
||||
)
|
||||
return data
|
||||
|
||||
|
||||
def _mk_get_json(self, path):
|
||||
import httpx
|
||||
|
||||
if not self.is_available:
|
||||
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
|
||||
url = self._base_url + path
|
||||
try:
|
||||
with httpx.Client(timeout=self._timeout) as http:
|
||||
resp = http.get(url, headers=self._headers())
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
except httpx.TimeoutException as exc:
|
||||
raise MediaKitError("MediaKit API 超时 (%ss)" % self._timeout, code="Timeout") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise MediaKitError(
|
||||
"MediaKit API HTTP %s: %s" % (exc.response.status_code, exc.response.text[:300]),
|
||||
code="HttpError",
|
||||
) from exc
|
||||
except httpx.RequestError as exc:
|
||||
raise MediaKitError("MediaKit API 网络错误: %s" % exc, code="NetworkError") from exc
|
||||
|
||||
|
||||
def _mediakit_asr_submit(self, video_url):
|
||||
"""提交语音转字幕任务(POST /tools/asr-subtitles)。返回 task_id。"""
|
||||
data = self._post_json(
|
||||
"/tools/asr-subtitles",
|
||||
{"video_url": video_url, "language": "cmn-Hans-CN"},
|
||||
)
|
||||
task_id = data.get("task_id")
|
||||
if not task_id:
|
||||
raise MediaKitError("MediaKit ASR 提交响应缺少 task_id: %s" % str(data)[:200])
|
||||
return task_id
|
||||
|
||||
|
||||
def _mediakit_asr_poll(self, task_id, poll_interval=2.0, max_attempts=90):
|
||||
"""轮询 ASR 任务直到 completed/failed。返回 (text, duration)。"""
|
||||
for attempt in range(max_attempts):
|
||||
time.sleep(poll_interval)
|
||||
try:
|
||||
data = self._get_json("/tasks/" + task_id)
|
||||
except MediaKitError as exc:
|
||||
if attempt < max_attempts - 1 and getattr(exc, "code", "") in ("Timeout", "NetworkError"):
|
||||
logger.warning("MediaKit ASR 轮询异常(第%d次),将重试: %s", attempt + 1, exc)
|
||||
continue
|
||||
raise
|
||||
st = data.get("status")
|
||||
if st in ("completed", "success"):
|
||||
result = data.get("result") or {}
|
||||
subs = result.get("subtitles") or []
|
||||
text = "".join(s.get("subtitle_text", "") for s in subs if isinstance(s, dict))
|
||||
duration = float(result.get("duration") or 0.0)
|
||||
return text.strip(), duration
|
||||
if st == "failed":
|
||||
err = data.get("error")
|
||||
if isinstance(err, dict):
|
||||
msg = err.get("message") or "unknown"
|
||||
code = err.get("code") or "TaskFailed"
|
||||
elif isinstance(err, str):
|
||||
msg, code = err, "TaskFailed"
|
||||
else:
|
||||
msg, code = "unknown", "TaskFailed"
|
||||
raise MediaKitError("MediaKit ASR 任务失败: %s" % msg, code=code)
|
||||
raise MediaKitError(
|
||||
"MediaKit ASR 超时(%ss 未完成)" % int(poll_interval * max_attempts),
|
||||
code="Timeout",
|
||||
)
|
||||
|
||||
|
||||
# 绑定到类(零侵入)
|
||||
if not hasattr(MediaKitClient, "_post_json"):
|
||||
MediaKitClient._post_json = _mk_post_json
|
||||
if not hasattr(MediaKitClient, "_get_json"):
|
||||
MediaKitClient._get_json = _mk_get_json
|
||||
if not hasattr(MediaKitClient, "asr_submit"):
|
||||
MediaKitClient.asr_submit = _mediakit_asr_submit
|
||||
if not hasattr(MediaKitClient, "asr_poll"):
|
||||
MediaKitClient.asr_poll = _mediakit_asr_poll
|
||||
|
||||
|
||||
# ── 下载 + 本地 ASR 兜底 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
def _direct_url_download_and_local_asr(direct_url, page_url, temp_dir):
|
||||
"""通过直链下载 MP4,再做本地 ASR。返回 (text, duration)。"""
|
||||
import os
|
||||
|
||||
import httpx
|
||||
|
||||
video_path = os.path.join(temp_dir, "video.mp4")
|
||||
try:
|
||||
with httpx.Client(timeout=90, follow_redirects=True, verify=False) as http:
|
||||
with http.stream(
|
||||
"GET",
|
||||
direct_url,
|
||||
headers={
|
||||
"User-Agent": (
|
||||
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) "
|
||||
"AppleWebKit/537.36 (KHTML, like Gecko) "
|
||||
"Chrome/128.0.0.0 Safari/537.36"
|
||||
),
|
||||
"Referer": "https://www.douyin.com/",
|
||||
"Accept": "*/*",
|
||||
"Accept-Language": "zh-CN,zh;q=0.9",
|
||||
},
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
downloaded = 0
|
||||
with open(video_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=65536):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
downloaded += len(chunk)
|
||||
if downloaded == 0:
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="直链下载为空")
|
||||
except HTTPException:
|
||||
raise
|
||||
except httpx.TimeoutException:
|
||||
logger.warning("直链下载超时: %s", page_url)
|
||||
raise HTTPException(status_code=status.HTTP_504_GATEWAY_TIMEOUT, detail="视频下载超时,请稍后重试") from None
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("直链下载失败: url=%s err=%s", page_url, exc)
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="视频下载失败: " + str(exc)[:200]) from exc
|
||||
|
||||
try:
|
||||
text = transcribe_to_text(video_path)
|
||||
return text.strip(), 0.0
|
||||
except ASRNotConfiguredError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=str(exc)) from exc
|
||||
except ASRTranscriptionError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
logger.exception("直链下载后 ASR 转写异常: path=%s", video_path)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="语音识别失败: " + str(exc)[:200],
|
||||
) from exc
|
||||
|
||||
|
||||
# ── 1. 从抖音视频提取文案 ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/douyin/__debug_diag")
|
||||
def douyin_diag():
|
||||
"""[Staging/Dev only] 抖音解析源诊断。"""
|
||||
import time as _t
|
||||
|
||||
import httpx as _httpx
|
||||
from app.services.douyin_resolver import APIZERO_API_KEY as _api_key_apizero
|
||||
from app.services.douyin_resolver import TIKHUB_API_KEY as _api_key_tikhub
|
||||
|
||||
results = {
|
||||
"providers": available_providers(),
|
||||
"env": {
|
||||
"APP_ENV": os.environ.get("APP_ENV", ""),
|
||||
"MEDIAKIT_CONFIGURED": bool(os.environ.get("MEDIAKIT_API_KEY", "")),
|
||||
},
|
||||
}
|
||||
|
||||
test_url = "https://v.douyin.com/hb-giW8cC1Q/"
|
||||
|
||||
t0 = _t.time()
|
||||
try:
|
||||
r = resolve_douyin_video(test_url)
|
||||
results["resolver"] = {
|
||||
"ok": bool(r),
|
||||
"source": r.source if r else None,
|
||||
"desc_len": len(r.desc) if r else 0,
|
||||
"has_video_url": bool(r.video_url) if r else False,
|
||||
"url_domain": r.video_url.split("/")[2] if r and r.video_url and "/" in r.video_url else None,
|
||||
"time": round(_t.time() - t0, 2),
|
||||
}
|
||||
except Exception as e:
|
||||
results["resolver"] = {"ok": False, "error": str(e)[:200], "time": round(_t.time() - t0, 2)}
|
||||
|
||||
if _api_key_apizero:
|
||||
t0 = _t.time()
|
||||
try:
|
||||
with _httpx.Client(timeout=8, verify=False) as c:
|
||||
r = c.get(
|
||||
"https://v1.apizero.cn/api/video-parse",
|
||||
params={"url": test_url, "flat": 2},
|
||||
headers={"Authorization": f"Bearer {_api_key_apizero}"},
|
||||
)
|
||||
results["apizero"] = {"status": r.status_code, "prefix": r.text[:200], "time": round(_t.time() - t0, 2)}
|
||||
except Exception as e:
|
||||
results["apizero"] = {"error": str(e)[:200], "time": round(_t.time() - t0, 2)}
|
||||
|
||||
if _api_key_tikhub:
|
||||
t0 = _t.time()
|
||||
try:
|
||||
with _httpx.Client(timeout=8, verify=False) as c:
|
||||
r = c.get(
|
||||
"https://api.tikhub.io/api/v1/douyin/web/get_aweme_id",
|
||||
params={"url": test_url},
|
||||
headers={"Authorization": f"Bearer {_api_key_tikhub}"},
|
||||
)
|
||||
results["tikhub"] = {"status": r.status_code, "prefix": r.text[:200], "time": round(_t.time() - t0, 2)}
|
||||
except Exception as e:
|
||||
results["tikhub"] = {"error": str(e)[:200], "time": round(_t.time() - t0, 2)}
|
||||
|
||||
return results
|
||||
|
||||
|
||||
@router.post("/extract-from-douyin", response_model=ExtractFromDouyinResponse)
|
||||
@points_gate("douyin_extract")
|
||||
def extract_from_douyin(
|
||||
request: ExtractFromDouyinRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
) -> ExtractFromDouyinResponse:
|
||||
"""从抖音视频下载无水印视频并通过 ASR 提取文案."""
|
||||
source_url = request.url.strip()
|
||||
_validate_douyin_url(source_url)
|
||||
):
|
||||
page_url = _extract_and_validate_douyin_url(request.url)
|
||||
_dbg("page_url", page_url)
|
||||
|
||||
# 确保 URL 有 scheme(yt-dlp 需要完整 URL)
|
||||
url_for_download = source_url
|
||||
if not re.match(r"^https?://", url_for_download, re.IGNORECASE):
|
||||
url_for_download = "https://" + url_for_download
|
||||
# ── Phase A:多源轮询解析 MP4 直链 ──
|
||||
last_err_stage = "parse"
|
||||
t0 = time.time()
|
||||
result = resolve_douyin_video(page_url)
|
||||
resolve_elapsed = time.time() - t0
|
||||
logger.info("抖音解析耗时: %.2fs providers=%s", resolve_elapsed, available_providers())
|
||||
|
||||
text: str = ""
|
||||
duration: float = 0.0
|
||||
direct_url = result.video_url if result else None
|
||||
feed_desc = (result.desc or "").strip() if result else ""
|
||||
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="douyin_extract_") as temp_dir:
|
||||
# 延迟导入 yt-dlp,避免模块缺失时影响其他路由启动
|
||||
try:
|
||||
import yt_dlp
|
||||
except ImportError as exc:
|
||||
logger.error("yt-dlp 未安装,抖音提取功能不可用: %s", exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="抖音提取功能暂不可用(缺少依赖 yt-dlp)",
|
||||
) from exc
|
||||
# 图文视频(无 video_url 但有 desc)直接返回文案,跳过 ASR
|
||||
if result and not direct_url and feed_desc:
|
||||
logger.info("图文视频直接返回文案: source=%s desc_len=%d", result.source, len(feed_desc))
|
||||
return ExtractFromDouyinResponse(
|
||||
text=feed_desc,
|
||||
duration_seconds=0.0,
|
||||
source_url=page_url,
|
||||
)
|
||||
|
||||
ydl_opts = {
|
||||
"format": "best[ext=mp4]/best",
|
||||
"outtmpl": f"{temp_dir}/%(id)s.%(ext)s",
|
||||
"quiet": True,
|
||||
"no_warnings": True,
|
||||
"noplaylist": True,
|
||||
}
|
||||
if not direct_url:
|
||||
if _DOUYIN_DEBUG_ERRORS:
|
||||
detail = f"抖音视频链接解析失败,请检查链接是否正确或稍后重试 [debug: providers={available_providers()}]"
|
||||
else:
|
||||
detail = "抖音视频链接解析失败,请检查链接是否正确或稍后重试"
|
||||
logger.warning("抖音解析全部失败: url=%s providers=%s", page_url, available_providers())
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=detail)
|
||||
|
||||
try:
|
||||
ydl = yt_dlp.YoutubeDL(ydl_opts)
|
||||
info = ydl.extract_info(url_for_download, download=True)
|
||||
except yt_dlp.utils.DownloadError as exc:
|
||||
# yt-dlp 官方异常类型:HTTP 错误、短链失效、视频下架等
|
||||
msg = str(exc)
|
||||
logger.warning("抖音下载失败: url=%s error=%s", source_url, msg)
|
||||
# 404/视频不存在/不可下载 → 400;网络问题/上游异常 → 502
|
||||
is_bad_url = any(
|
||||
kw in msg.lower() for kw in ("404", "not found", "unable to download webpage", "unsupported url", "no video formats")
|
||||
# ── Phase B:ASR 转文字 ──
|
||||
mk_client = get_mediakit_client()
|
||||
text = ""
|
||||
duration = 0.0
|
||||
|
||||
# B1:MediaKit 云端 ASR(不下载视频,最快)
|
||||
if mk_client.is_available:
|
||||
last_err_stage = "asr"
|
||||
try:
|
||||
task_id = mk_client.asr_submit(direct_url)
|
||||
text, duration = mk_client.asr_poll(task_id)
|
||||
text = text.strip()
|
||||
if text:
|
||||
logger.info(
|
||||
"抖音 MediaKit ASR 成功: source=%s text_len=%d duration=%.1f total_time=%.1fs",
|
||||
result.source,
|
||||
len(text),
|
||||
duration,
|
||||
time.time() - t0,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST if is_bad_url else status.HTTP_502_BAD_GATEWAY,
|
||||
detail=("无法解析该抖音链接,请确认链接有效且视频未被下架" if is_bad_url else f"视频下载失败: {msg[:200]}"),
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.exception("抖音视频下载异常: url=%s", source_url)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"视频下载失败: {str(exc)[:200]}",
|
||||
) from exc
|
||||
else:
|
||||
logger.info("抖音 MediaKit ASR 返回空文本(无旁白/BGM视频)")
|
||||
except MediaKitError as exc:
|
||||
logger.warning("MediaKit ASR 失败,回退本地 ASR: %s", exc)
|
||||
text = ""
|
||||
|
||||
if info is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无法解析该抖音链接",
|
||||
)
|
||||
# B2:回退下载 + 本地 ASR
|
||||
if not text:
|
||||
last_err_stage = "download"
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="douyin_extract_") as temp_dir:
|
||||
text, dl_duration = _direct_url_download_and_local_asr(direct_url, page_url, temp_dir)
|
||||
text = (text or "").strip()
|
||||
if dl_duration and not duration:
|
||||
duration = dl_duration
|
||||
if text:
|
||||
logger.info(
|
||||
"抖音本地 ASR 成功: source=%s text_len=%d total_time=%.1fs",
|
||||
result.source,
|
||||
len(text),
|
||||
time.time() - t0,
|
||||
)
|
||||
last_err_stage = "asr"
|
||||
except HTTPException as exc:
|
||||
# 下载超时(504)是明确的网络错误,直接抛出
|
||||
if exc.status_code == status.HTTP_504_GATEWAY_TIMEOUT:
|
||||
raise
|
||||
# 本地 ASR 不可用/失败(502/503)时记录后继续走 desc 兜底,
|
||||
# 不直接抛 502,避免 API 镜像缺 worker 模块时整条链路挂掉
|
||||
logger.warning("本地 ASR 链路失败(status=%d): %s", exc.status_code, exc.detail)
|
||||
text = ""
|
||||
# 如果是下载失败(非ASR错误),保持stage为download
|
||||
if "语音识别" in str(exc.detail) or "ASR" in str(exc.detail):
|
||||
last_err_stage = "asr"
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("本地 ASR 链路异常: %s", exc)
|
||||
text = ""
|
||||
|
||||
video_path = ydl.prepare_filename(info)
|
||||
try:
|
||||
duration = float(info.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
duration = 0.0
|
||||
# ── Phase C:结果判定 & 兜底 ──
|
||||
|
||||
# 校验下载的文件是否真的存在(某些 yt-dlp 版本可能 info 成功但未下载到文件)
|
||||
if not os.path.isfile(video_path) or os.path.getsize(video_path) == 0:
|
||||
logger.error("yt-dlp 未产生有效视频文件: path=%s", video_path)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="视频下载异常:未获取到有效文件",
|
||||
)
|
||||
# ASR 空结果(无旁白视频)→ 使用解析源 desc 兜底
|
||||
if not text and feed_desc:
|
||||
text = feed_desc
|
||||
logger.info("抖音 ASR 空结果,使用解析源 desc 兜底: desc_len=%d", len(text))
|
||||
|
||||
# ASR 转写(兜底捕获所有异常,避免 500)
|
||||
try:
|
||||
text = transcribe_to_text(video_path)
|
||||
except ASRNotConfiguredError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
except ASRTranscriptionError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.exception("ASR 转写异常: path=%s", video_path)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"语音识别失败: {str(exc)[:200]}",
|
||||
) from exc
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
# 最后兜底:任何未捕获异常都转成 502/400,不允许冒泡成 500
|
||||
logger.exception("抖音文案提取未预期异常: url=%s", source_url)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"抖音文案提取失败: {str(exc)[:200]}",
|
||||
) from exc
|
||||
if not text:
|
||||
stage_msg = {
|
||||
"parse": "抖音视频链接解析失败,请检查链接是否正确或稍后重试",
|
||||
"download": "抖音视频下载失败,请检查网络或稍后重试",
|
||||
"asr": "抖音语音识别失败,请稍后重试或手动输入文案",
|
||||
}
|
||||
user_msg = stage_msg.get(last_err_stage, "抖音链接解析暂时不可用,请稍后重试或手动输入文案")
|
||||
if _DOUYIN_DEBUG_ERRORS:
|
||||
user_msg = user_msg + f" [debug: stage={last_err_stage} source={result.source}]"
|
||||
logger.warning("抖音文案提取失败: url=%s stage=%s source=%s", page_url, last_err_stage, result.source)
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=user_msg)
|
||||
|
||||
return ExtractFromDouyinResponse(
|
||||
text=text,
|
||||
duration_seconds=duration,
|
||||
source_url=source_url,
|
||||
source_url=page_url,
|
||||
)
|
||||
|
||||
|
||||
# ── 2. AI 文案改写 ───────────────────────────────────────────────────────────
|
||||
# ── 2. AI 文案改写 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/ai-rewrite",
|
||||
response_model=AiRewriteResponse,
|
||||
)
|
||||
@router.post("/ai-rewrite", response_model=AiRewriteResponse)
|
||||
@points_gate("ai_rewrite")
|
||||
def ai_rewrite(
|
||||
request: AiRewriteRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
) -> AiRewriteResponse:
|
||||
"""使用豆包大模型改写文案."""
|
||||
):
|
||||
content = (request.content or "").strip()
|
||||
if not content:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文案内容不能为空",
|
||||
)
|
||||
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="文案内容不能为空")
|
||||
style = request.style or "口语化"
|
||||
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="AI 服务不可用,请联系管理员配置豆包大模型 API Key",
|
||||
)
|
||||
|
||||
system_prompt = (
|
||||
"你是一个专业的短视频文案改写专家。请对以下文案进行改写,"
|
||||
"要求:保留原意、口语化、适合短视频口播、调整语序避免查重。"
|
||||
)
|
||||
if style:
|
||||
system_prompt += f"\n风格要求:{style}"
|
||||
|
||||
user_prompt = f"请改写以下文案:\n\n{content}"
|
||||
|
||||
system_prompt = system_prompt + "\n风格要求:" + style
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
{"role": "user", "content": "请改写以下文案:\n\n" + content},
|
||||
]
|
||||
|
||||
try:
|
||||
rewritten = client.chat_completion(
|
||||
messages=messages,
|
||||
temperature=0.8,
|
||||
max_tokens=2048,
|
||||
)
|
||||
rewritten = client.chat_completion(messages=messages, temperature=0.8, max_tokens=2048)
|
||||
except Exception as exc:
|
||||
logger.error("AI 改写调用失败: %s", exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"AI 改写失败: {exc}",
|
||||
) from exc
|
||||
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="AI 改写失败: " + str(exc)) from exc
|
||||
if not rewritten:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="AI 改写未返回有效结果",
|
||||
)
|
||||
|
||||
return AiRewriteResponse(
|
||||
original=content,
|
||||
rewritten=rewritten.strip(),
|
||||
style=style,
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="AI 改写未返回有效结果")
|
||||
return AiRewriteResponse(original=content, rewritten=rewritten.strip(), style=style)
|
||||
|
||||
|
||||
# ── 3. AI 标题生成 ───────────────────────────────────────────────────────────
|
||||
# ── 3. AI 标题生成 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/ai-generate-titles",
|
||||
response_model=AiGenerateTitlesResponse,
|
||||
)
|
||||
@router.post("/ai-generate-titles", response_model=AiGenerateTitlesResponse)
|
||||
@points_gate("ai_title")
|
||||
def ai_generate_titles(
|
||||
request: AiGenerateTitlesRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
) -> AiGenerateTitlesResponse:
|
||||
"""使用现有 generate_smart_titles 生成标题."""
|
||||
):
|
||||
content = (request.content or "").strip()
|
||||
if not content:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文案内容不能为空",
|
||||
)
|
||||
|
||||
# count 限制在 1-5(Pydantic ge=1 le=5 已校验),但为兼容直接调用场景截断
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="文案内容不能为空")
|
||||
count = max(1, min(5, request.count))
|
||||
|
||||
from app.services.ai_service import generate_smart_titles
|
||||
|
||||
result = generate_smart_titles(
|
||||
description=content,
|
||||
style="viral",
|
||||
count=count,
|
||||
)
|
||||
|
||||
result = generate_smart_titles(description=content, style="viral", count=count)
|
||||
titles = result.get("titles", [])[:count]
|
||||
|
||||
return AiGenerateTitlesResponse(titles=titles)
|
||||
|
||||
@@ -10,9 +10,11 @@ from typing import Any
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_user_repository
|
||||
from app.schemas.subscription import (
|
||||
BillingCycle,
|
||||
BillingRecord,
|
||||
ChangePlanRequest,
|
||||
ChangePlanResponse,
|
||||
MembershipType,
|
||||
SimpleResponse,
|
||||
SubscriptionInfo,
|
||||
ToggleAutoRenewRequest,
|
||||
@@ -26,43 +28,18 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# ============ 配额定义(硬编码,后续可迁移到配置中心) ============
|
||||
# ============ 会员展示名称(与 packages.domain.points_rules.MEMBERSHIP_PRICES 对应)============
|
||||
|
||||
PLAN_QUOTAS = {
|
||||
"free": {"max_projects": 3, "max_storage_gb": 10},
|
||||
"standard": {"max_projects": 10, "max_storage_gb": 50},
|
||||
"pro": {"max_projects": -1, "max_storage_gb": 100},
|
||||
"enterprise": {"max_projects": -1, "max_storage_gb": 1000},
|
||||
_PLAN_NAMES: dict[str, str] = {
|
||||
MembershipType.FREE: "免费用户",
|
||||
MembershipType.MONTHLY: "月卡会员",
|
||||
MembershipType.QUARTERLY: "季卡会员",
|
||||
MembershipType.YEARLY: "年卡会员",
|
||||
}
|
||||
|
||||
|
||||
# ============ Helper Functions ============
|
||||
|
||||
|
||||
def _get_plan_name(plan_id: str) -> str:
|
||||
"""获取套餐显示名称"""
|
||||
plan_names = {
|
||||
"free": "体验版",
|
||||
"standard": "标准版",
|
||||
"pro": "专业版",
|
||||
"enterprise": "企业版",
|
||||
}
|
||||
return plan_names.get(plan_id, "未知套餐")
|
||||
|
||||
|
||||
def _get_plan_price(plan_id: str, billing_cycle: str) -> float:
|
||||
"""获取套餐价格"""
|
||||
prices = {
|
||||
("free", "monthly"): 0,
|
||||
("free", "yearly"): 0,
|
||||
("standard", "monthly"): 99,
|
||||
("standard", "yearly"): 999,
|
||||
("pro", "monthly"): 299,
|
||||
("pro", "yearly"): 2999,
|
||||
("enterprise", "monthly"): 999,
|
||||
("enterprise", "yearly"): 9999,
|
||||
}
|
||||
return prices.get((plan_id, billing_cycle), 0)
|
||||
return _PLAN_NAMES.get(plan_id, "免费用户")
|
||||
|
||||
|
||||
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
|
||||
@@ -75,15 +52,20 @@ def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
|
||||
period_start = now.isoformat()
|
||||
period_end = now.isoformat()
|
||||
|
||||
plan_id = user.user.subscription_plan or MembershipType.FREE
|
||||
# 旧档位(standard/pro/enterprise)统一降级为 monthly,避免前端炸掉
|
||||
if plan_id in {"standard", "pro", "enterprise"}:
|
||||
plan_id = MembershipType.MONTHLY
|
||||
|
||||
return SubscriptionInfo(
|
||||
id=f"sub-{user.user.id[:8]}",
|
||||
plan_id=user.user.subscription_plan or "free",
|
||||
plan_name=_get_plan_name(user.user.subscription_plan or "free"),
|
||||
plan_id=plan_id,
|
||||
plan_name=_get_plan_name(plan_id),
|
||||
status=user.user.subscription_status or "active",
|
||||
billing_cycle="monthly",
|
||||
billing_cycle=plan_id if plan_id != MembershipType.FREE else BillingCycle.MONTHLY,
|
||||
current_period_start=period_start,
|
||||
current_period_end=period_end,
|
||||
amount=_get_plan_price(user.user.subscription_plan or "free", "monthly"),
|
||||
amount=0 if plan_id == MembershipType.FREE else 0, # 金额由前端 /plans 接口展示
|
||||
auto_renew=True,
|
||||
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
|
||||
)
|
||||
@@ -115,11 +97,11 @@ def list_membership_plans(
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == "monthly":
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == "quarterly":
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == "yearly":
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
@@ -151,7 +133,7 @@ async def get_billing_records(
|
||||
return [
|
||||
BillingRecord(
|
||||
id=r.id,
|
||||
plan_name=r.plan_name,
|
||||
plan_name=_get_plan_name(r.plan_name),
|
||||
amount=r.amount,
|
||||
billing_cycle=r.billing_cycle,
|
||||
status=r.status,
|
||||
@@ -165,6 +147,10 @@ async def get_billing_records(
|
||||
session.close()
|
||||
|
||||
|
||||
_VALID_PLANS = {MembershipType.MONTHLY, MembershipType.QUARTERLY, MembershipType.YEARLY}
|
||||
_VALID_CYCLES = {BillingCycle.MONTHLY, BillingCycle.QUARTERLY, BillingCycle.YEARLY}
|
||||
|
||||
|
||||
@router.post("/change-plan", response_model=ChangePlanResponse)
|
||||
async def change_plan(
|
||||
request: ChangePlanRequest,
|
||||
@@ -173,47 +159,45 @@ async def change_plan(
|
||||
) -> ChangePlanResponse:
|
||||
"""变更订阅套餐(升级/降级)"""
|
||||
# TODO: 接入支付验证(支付宝/微信支付)
|
||||
valid_plans = {"free", "standard", "pro", "enterprise"}
|
||||
if request.target_plan_id not in valid_plans:
|
||||
target_plan = request.target_plan_id
|
||||
if target_plan not in _VALID_PLANS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"无效的套餐ID。支持的套餐: {', '.join(valid_plans)}",
|
||||
detail=f"无效的会员类型。支持: {', '.join(sorted(_VALID_PLANS))}",
|
||||
)
|
||||
|
||||
valid_cycles = {"monthly", "yearly"}
|
||||
if request.billing_cycle not in valid_cycles:
|
||||
if request.billing_cycle not in _VALID_CYCLES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="无效的计费周期。支持: monthly, yearly",
|
||||
detail=f"无效的计费周期。支持: {', '.join(sorted(_VALID_CYCLES))}",
|
||||
)
|
||||
|
||||
user = current_user.user
|
||||
current_plan = user.subscription_plan or "free"
|
||||
target_plan = request.target_plan_id
|
||||
current_plan = user.subscription_plan or MembershipType.FREE
|
||||
# 旧档位归一化,避免永远显示"您已经是xxx"
|
||||
if current_plan in {"standard", "pro", "enterprise"}:
|
||||
current_plan = MembershipType.MONTHLY
|
||||
|
||||
if current_plan == target_plan:
|
||||
return ChangePlanResponse(
|
||||
success=False,
|
||||
message=f"您已经是 {_get_plan_name(target_plan)}",
|
||||
message=f"您已经是{_get_plan_name(target_plan)}",
|
||||
)
|
||||
|
||||
# 通过 dataclasses.replace 创建新实例(不直接修改 dataclass)
|
||||
quotas = PLAN_QUOTAS.get(target_plan, PLAN_QUOTAS["free"])
|
||||
updated_user = replace(
|
||||
user,
|
||||
subscription_plan=target_plan,
|
||||
subscription_status="active",
|
||||
max_projects=quotas["max_projects"],
|
||||
max_storage_gb=quotas["max_storage_gb"],
|
||||
max_projects=-1, # 付费会员不限项目数
|
||||
max_storage_gb=100,
|
||||
)
|
||||
user_repository.save(updated_user)
|
||||
|
||||
# 用更新后的用户构造响应
|
||||
refreshed_auth_user = AuthenticatedUser(user=updated_user)
|
||||
|
||||
return ChangePlanResponse(
|
||||
success=True,
|
||||
message=f"套餐已成功变更为 {_get_plan_name(target_plan)}",
|
||||
message=f"套餐已成功变更为{_get_plan_name(target_plan)}",
|
||||
new_subscription=_build_subscription_info(refreshed_auth_user),
|
||||
)
|
||||
|
||||
@@ -225,10 +209,11 @@ async def cancel_subscription(
|
||||
) -> SimpleResponse:
|
||||
"""取消订阅"""
|
||||
user = current_user.user
|
||||
if user.subscription_plan == "free":
|
||||
plan_id = user.subscription_plan or MembershipType.FREE
|
||||
if plan_id == MembershipType.FREE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="体验版无需取消",
|
||||
detail="免费用户无需取消订阅",
|
||||
)
|
||||
|
||||
updated_user = replace(user, subscription_status="cancelled")
|
||||
@@ -236,7 +221,7 @@ async def cancel_subscription(
|
||||
|
||||
return SimpleResponse(
|
||||
success=True,
|
||||
message="订阅已取消,当前周期结束后停止服务",
|
||||
message="订阅已取消,当前周期结束后将降级为免费用户",
|
||||
)
|
||||
|
||||
|
||||
@@ -262,11 +247,14 @@ async def payment_callback(
|
||||
if SessionLocal is None:
|
||||
raise HTTPException(status_code=500, detail="Database not available")
|
||||
|
||||
# 仅接受当前会员体系的 plan 值
|
||||
if plan not in _VALID_PLANS:
|
||||
raise HTTPException(status_code=400, detail=f"未知的会员类型: {plan}")
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
repo = SQLAlchemyBillingRepository(session)
|
||||
|
||||
# 创建账单记录
|
||||
record_id = uuid.uuid4().hex
|
||||
repo.create(
|
||||
{
|
||||
@@ -279,19 +267,20 @@ async def payment_callback(
|
||||
}
|
||||
)
|
||||
|
||||
# 在事务中标记支付成功并更新订阅
|
||||
repo.mark_paid(record_id, payment_method, payment_id)
|
||||
|
||||
# 计算到期时间
|
||||
days = 365 if billing_cycle == "yearly" else 30
|
||||
days_map = {BillingCycle.MONTHLY: 30, BillingCycle.QUARTERLY: 90, BillingCycle.YEARLY: 365}
|
||||
days = days_map.get(billing_cycle, 30)
|
||||
expires_at = datetime.now(UTC) + timedelta(days=days)
|
||||
repo.update_subscription_on_payment(user_id, plan, expires_at)
|
||||
|
||||
return {"success": True, "message": "支付成功", "record_id": record_id}
|
||||
except HTTPException:
|
||||
session.rollback()
|
||||
raise
|
||||
except Exception as e:
|
||||
session.rollback()
|
||||
logger.error(f"支付回调处理失败: user_id={user_id}, plan={plan}, error={e}")
|
||||
# 不返回原始异常信息,避免泄漏内部实现细节
|
||||
logger.error("支付回调处理失败: user_id=%s, plan=%s, error=%s", user_id, plan, e)
|
||||
raise HTTPException(status_code=500, detail="支付处理失败,请稍后重试") from e
|
||||
finally:
|
||||
session.close()
|
||||
@@ -303,10 +292,5 @@ async def toggle_auto_renew(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> SimpleResponse:
|
||||
"""切换自动续费"""
|
||||
# TODO: 实际需要在数据库中存储 auto_renew 字段
|
||||
status_text = "已开启自动续费" if request.enabled else "已关闭自动续费"
|
||||
|
||||
return SimpleResponse(
|
||||
success=True,
|
||||
message=status_text,
|
||||
)
|
||||
return SimpleResponse(success=True, message=status_text)
|
||||
|
||||
@@ -1,243 +1,35 @@
|
||||
"""Title library CRUD routes.
|
||||
"""Title library routes — DEPRECATED (#1894).
|
||||
|
||||
.. deprecated::
|
||||
标题库 API 已废弃(#1894),标题配置已整合到 scripts 模型。
|
||||
所有接口保留向后兼容,但返回 Warning header 并记录日志。
|
||||
独立标题库已废弃。前端应直接调用 GET /api/v1/scripts 获取文案列表,
|
||||
取每条文案的 `title` 字段作为标题候选。
|
||||
|
||||
所有 /api/v1/titles 端点统一返回 HTTP 410 Gone。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from app.api.routes._helpers import get_user_plan
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session, get_user_repository
|
||||
from app.schemas.title_library import (
|
||||
CreateTitleLibraryRequest,
|
||||
ListTitleLibraryResponse,
|
||||
TitleLibraryItemResponse,
|
||||
UpdateTitleLibraryRequest,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
|
||||
from packages.application.title_library.commands import (
|
||||
CreateTitleLibraryCommand,
|
||||
PickTitleCommand,
|
||||
UpdateTitleLibraryCommand,
|
||||
)
|
||||
from packages.application.title_library.use_cases import (
|
||||
CreateTitleLibraryUseCase,
|
||||
DeleteTitleLibraryUseCase,
|
||||
GetTitleLibraryUseCase,
|
||||
ListTitleLibraryUseCase,
|
||||
NotFoundError,
|
||||
PickTitleUseCase,
|
||||
QuotaExceededError,
|
||||
UpdateTitleLibraryUseCase,
|
||||
)
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from fastapi import APIRouter, Response, status
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEPRECATION_WARNING = (
|
||||
'299 - "Title library API is deprecated; migrate to scripts.title_text/'
|
||||
'title_category/title_config (issue #1894)"'
|
||||
_GONE_MESSAGE = (
|
||||
"标题库 API 已废弃(#1894):独立标题库已合并进文案库,"
|
||||
"请使用 GET /api/v1/scripts 获取文案列表并取 title 字段作为标题。"
|
||||
)
|
||||
|
||||
|
||||
def _deprecation_headers() -> dict:
|
||||
"""返回 deprecation Warning header (ASCII-only, RFC 7234 §5.5)."""
|
||||
return {"Warning": _DEPRECATION_WARNING, "Deprecation": "true"}
|
||||
def _gone(response: Response) -> dict:
|
||||
response.status_code = status.HTTP_410_GONE
|
||||
response.headers["Deprecation"] = "true"
|
||||
response.headers["Sunset"] = "Tue, 16 Sep 2026 00:00:00 GMT"
|
||||
return {"error": {"code": "GONE", "message": _GONE_MESSAGE}}
|
||||
|
||||
|
||||
def _log_deprecation(endpoint: str) -> None:
|
||||
logger.warning("[Deprecated] title_library API 调用: %s — %s", endpoint, _DEPRECATION_WARNING)
|
||||
@router.api_route("", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
|
||||
def titles_root_gone(response: Response) -> dict:
|
||||
return _gone(response)
|
||||
|
||||
|
||||
def _get_title_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTitleLibraryRepository:
|
||||
return SQLAlchemyTitleLibraryRepository(session)
|
||||
|
||||
|
||||
def _to_response(item) -> TitleLibraryItemResponse:
|
||||
return TitleLibraryItemResponse(
|
||||
id=item.id,
|
||||
user_id=item.user_id,
|
||||
name=item.name,
|
||||
text=item.text,
|
||||
category=item.category,
|
||||
description=item.description,
|
||||
tags=item.tags,
|
||||
usage_count=item.usage_count,
|
||||
is_active=item.is_active,
|
||||
created_at=item.created_at,
|
||||
updated_at=item.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=ListTitleLibraryResponse)
|
||||
def list_titles(
|
||||
response: Response,
|
||||
category: Optional[str] = Query(None),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> ListTitleLibraryResponse:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
|
||||
_log_deprecation("list_titles")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListTitleLibraryUseCase(title_repository)
|
||||
items = use_case.execute(user_id, category=category, skip=skip, limit=limit)
|
||||
total = title_repository.count_by_user(user_id)
|
||||
return ListTitleLibraryResponse(
|
||||
items=[_to_response(i) for i in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/pick", response_model=TitleLibraryItemResponse)
|
||||
def pick_title(
|
||||
response: Response,
|
||||
category: Optional[str] = Query(None, description="按分类筛选,不填则从全部标题中选"),
|
||||
exclude_ids: Optional[str] = Query(
|
||||
None,
|
||||
description="排除的标题ID(逗号分隔),用于批量生成时避免重复",
|
||||
),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代.
|
||||
|
||||
智能选择一个标题。
|
||||
|
||||
策略:优先使用次数少的,从最少的前5个中随机选一个,兼顾公平和多样性。
|
||||
"""
|
||||
_log_deprecation("pick_title")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
exclude_list: list[str] = []
|
||||
if exclude_ids:
|
||||
exclude_list = [t.strip() for t in exclude_ids.split(",") if t.strip()]
|
||||
|
||||
use_case = PickTitleUseCase(title_repository)
|
||||
item = use_case.execute(
|
||||
PickTitleCommand(
|
||||
user_id=user_id,
|
||||
category=category,
|
||||
exclude_ids=exclude_list,
|
||||
)
|
||||
)
|
||||
if item is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="标题库为空,请先添加标题",
|
||||
)
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.get("/{title_id}", response_model=TitleLibraryItemResponse)
|
||||
def get_title(
|
||||
title_id: str,
|
||||
response: Response,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
|
||||
_log_deprecation("get_title")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = GetTitleLibraryUseCase(title_repository)
|
||||
item = use_case.execute(title_id, user_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.post("", response_model=TitleLibraryItemResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_title(
|
||||
response: Response,
|
||||
request: CreateTitleLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
|
||||
_log_deprecation("create_title")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
plan_name = get_user_plan(user_id, user_repository)
|
||||
command = CreateTitleLibraryCommand(
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
category=request.category,
|
||||
description=request.description,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = CreateTitleLibraryUseCase(title_repository)
|
||||
try:
|
||||
item = use_case.execute(command, plan_name=plan_name)
|
||||
except QuotaExceededError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=f"标题库配额已满({exc.used}/{exc.limit}),请升级套餐",
|
||||
) from exc
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.put("/{title_id}", response_model=TitleLibraryItemResponse)
|
||||
def update_title(
|
||||
title_id: str,
|
||||
response: Response,
|
||||
request: UpdateTitleLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
|
||||
_log_deprecation("update_title")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id=title_id,
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
category=request.category,
|
||||
description=request.description,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = UpdateTitleLibraryUseCase(title_repository)
|
||||
try:
|
||||
item = use_case.execute(command)
|
||||
except NotFoundError as _e:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found") from _e
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.delete("/{title_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
|
||||
def delete_title(
|
||||
title_id: str,
|
||||
response: Response,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> Response:
|
||||
"""[Deprecated] 请使用 scripts API 的 title_text/title_category 字段替代."""
|
||||
_log_deprecation("delete_title")
|
||||
for k, v in _deprecation_headers().items():
|
||||
response.headers[k] = v
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = DeleteTitleLibraryUseCase(title_repository)
|
||||
deleted = use_case.execute(title_id, user_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
|
||||
return
|
||||
@router.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
|
||||
def titles_subpath_gone(response: Response, path: str) -> dict:
|
||||
return _gone(response)
|
||||
|
||||
@@ -98,6 +98,24 @@ class CreateGenerationTaskRequest(BaseModel):
|
||||
description="各变体独立标题文字数组:长度1=共用,长度=count=独立。为空时使用 title_config.text",
|
||||
)
|
||||
|
||||
# ── 智能降重开关(#1970)──
|
||||
# True(默认):edge_crop + 片段级微变换(hflip/变速/亮度/对比度/饱和度/BGM偏移)全部生效;
|
||||
# False:跳过 edge_crop、不注入微变换,渲染确定性(固定种子)。
|
||||
dedup_enabled: bool = Field(default=True, description="智能降重开关,默认开启;关闭后跳过边缘裁切与微变换")
|
||||
|
||||
# ── 剪辑组装模式(#1970 PR3)──
|
||||
# random(默认,完全兼容现有随机混剪)/ narrative(叙事剪辑:文案→TTS 配音→标签匹配画面)
|
||||
assembly_mode: str = Field(default="random", description="组装模式:random=随机混剪(默认),narrative=叙事剪辑")
|
||||
# 叙事模式必填:文案库 scripts.id(后端据此读取 content 合成 TTS)
|
||||
script_id: str = Field(default="", description="叙事模式必填:文案库 ID")
|
||||
# 叙事模式必填:TTS 音色 ID(preset 为 CosyVoice 音色 id;clone 为克隆档案 id)
|
||||
tts_voice_id: str = Field(default="", description="叙事模式必填:TTS 音色 ID(系统音色或克隆档案 ID)")
|
||||
tts_voice_source: str = Field(default="preset", description="TTS 音色来源:preset=系统预设(默认),clone=克隆音色")
|
||||
# 视频比例:当前前端 9:16/16:9;与 output_width/output_height 并存,传了具体分辨率时以分辨率为准
|
||||
video_ratio: str = Field(
|
||||
default="", description="视频比例,如 9:16(默认竖屏)/16:9;与显式分辨率冲突时以分辨率为准"
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_variant_arrays(self) -> "CreateGenerationTaskRequest":
|
||||
"""变体数组字段长度校验 + #1749 配音严格守卫。
|
||||
@@ -127,6 +145,26 @@ class CreateGenerationTaskRequest(BaseModel):
|
||||
raise ValueError(f"variant_plan_ids 长度({len(self.variant_plan_ids)})必须与 count({self.count})一致")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_assembly_mode(self) -> "CreateGenerationTaskRequest":
|
||||
"""#1970 组装模式与叙事模式入参校验。"""
|
||||
if self.assembly_mode not in ("random", "narrative"):
|
||||
raise ValueError("assembly_mode 仅支持 'random'(默认)或 'narrative'")
|
||||
if self.tts_voice_source not in ("preset", "clone"):
|
||||
raise ValueError("tts_voice_source 仅支持 'preset' 或 'clone'")
|
||||
if self.video_ratio:
|
||||
parts = self.video_ratio.split(":")
|
||||
if len(parts) != 2 or not all(p.isdigit() and int(p) > 0 for p in parts):
|
||||
raise ValueError("video_ratio 格式必须为 '宽:高',如 9:16 或 16:9")
|
||||
if self.video_ratio not in ("9:16", "16:9", "1:1", "3:4", "4:3"):
|
||||
raise ValueError("video_ratio 仅支持 9:16 / 16:9 / 1:1 / 3:4 / 4:3")
|
||||
if self.assembly_mode == "narrative":
|
||||
if not self.script_id.strip():
|
||||
raise ValueError("叙事模式(narrative)必须提供 script_id(文案库 ID)")
|
||||
if not self.tts_voice_id.strip():
|
||||
raise ValueError("叙事模式(narrative)必须提供 tts_voice_id(TTS 音色 ID)")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
|
||||
has_project = bool(self.project_id.strip())
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
"""GPU MuseTalk 反向轮询 API Schema 定义.
|
||||
|
||||
面向部署在用户 RTX2060 本地的 GPU Worker 脚本,不面向前端用户。
|
||||
Worker 用长期 GPU_WORKER_TOKEN 鉴权(不是用户 JWT)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ── Worker 注册/心跳 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuWorkerRegisterRequest(BaseModel):
|
||||
"""Worker 启动/心跳时上报自身信息."""
|
||||
|
||||
worker_id: str = Field(..., min_length=1, max_length=100, description="Worker 唯一 ID(机器名+UUID 等)")
|
||||
hostname: str = Field("", max_length=200, description="主机名,用于运维排查")
|
||||
gpu_name: str = Field("", max_length=200, description="GPU 型号,如 'NVIDIA GeForce RTX 2060'")
|
||||
free_vram_mb: int = Field(0, ge=0, description="当前空闲显存(MB)")
|
||||
capabilities: str = Field("musetalk", max_length=500, description="能力列表,逗号分隔,如 'musetalk'")
|
||||
task_id: Optional[str] = Field(
|
||||
None,
|
||||
max_length=64,
|
||||
description=(
|
||||
"当前正在处理的任务 ID。Worker 推理期间定期心跳时携带,"
|
||||
"服务端同步刷新该任务 last_heartbeat_at,防止长推理被误判超时;空闲时不传"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class GpuWorkerRegisterResponse(BaseModel):
|
||||
ok: bool = True
|
||||
server_time: datetime
|
||||
message: str = "ok"
|
||||
|
||||
|
||||
# ── 轮询任务 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncTaskPayload(BaseModel):
|
||||
"""下发给 Worker 的任务载荷(含预签名下载 URL)."""
|
||||
|
||||
task_id: str
|
||||
video_url: str = Field(..., description="人物视频预签名下载 URL(GET)")
|
||||
audio_url: str = Field(..., description="驱动音频预签名下载 URL(GET)")
|
||||
lipsync_job_id: str = ""
|
||||
user_id: str = ""
|
||||
project_id: str = ""
|
||||
created_at: datetime
|
||||
upload_url: str = Field(..., description="结果视频预签名上传 URL(PUT, video/mp4)")
|
||||
upload_method: str = Field("PUT", description="上传方式,目前只支持 PUT")
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class GpuLipsyncPollResponse(BaseModel):
|
||||
"""Worker poll 的返回:200 带任务,204 无任务."""
|
||||
|
||||
task: Optional[GpuLipsyncTaskPayload] = None
|
||||
|
||||
|
||||
# ── Worker 上报结果 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncResultRequest(BaseModel):
|
||||
"""Worker 通过 multipart 上传结果时携带的字段(非文件字段)."""
|
||||
|
||||
task_id: str = Field(..., min_length=1, max_length=64)
|
||||
worker_id: str = Field(..., min_length=1, max_length=100)
|
||||
success: bool = Field(True, description="true=成功(此时必须上传 result 视频文件);false=失败")
|
||||
duration_seconds: float = Field(0.0, ge=0, description="合成后视频时长(秒),成功时应填入")
|
||||
error_msg: str = Field("", max_length=2000, description="失败原因,success=false 时必填")
|
||||
|
||||
|
||||
class GpuLipsyncResultResponse(BaseModel):
|
||||
ok: bool = True
|
||||
task_id: str
|
||||
status: str # done / failed
|
||||
message: str = "ok"
|
||||
|
||||
|
||||
# ── 业务侧查询任务状态 ────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncStatusResponse(BaseModel):
|
||||
task_id: str
|
||||
status: str
|
||||
result_url: str = ""
|
||||
result_duration: float = 0.0
|
||||
error_msg: str = ""
|
||||
worker_id: str = ""
|
||||
attempt: int = 0
|
||||
created_at: datetime
|
||||
started_at: Optional[datetime] = None
|
||||
finished_at: Optional[datetime] = None
|
||||
|
||||
|
||||
# ── 创建任务(内部服务调用) ──────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncCreateRequest(BaseModel):
|
||||
"""服务层内部创建 GPU 任务用(不通过 HTTP 暴露给 Worker/前端)."""
|
||||
|
||||
video_url: str # 已可访问的 OSS key 或公网 URL(API 侧会转预签名)
|
||||
audio_url: str
|
||||
lipsync_job_id: str = ""
|
||||
user_id: str = ""
|
||||
project_id: str = ""
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -22,9 +22,6 @@ class ScriptResponse(BaseModel):
|
||||
content: str
|
||||
segments: list[ScriptSegment] = Field(default_factory=list)
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
title_text: str = ""
|
||||
title_category: str = ""
|
||||
title_config: Dict[str, Any] = Field(default_factory=dict)
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
@@ -39,9 +36,6 @@ class CreateScriptRequest(BaseModel):
|
||||
content: str = ""
|
||||
segments: list[ScriptSegment] = Field(default_factory=list)
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
title_text: str = ""
|
||||
title_category: str = ""
|
||||
title_config: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class UpdateScriptRequest(BaseModel):
|
||||
@@ -49,6 +43,3 @@ class UpdateScriptRequest(BaseModel):
|
||||
content: Optional[str] = None
|
||||
segments: Optional[list[ScriptSegment]] = None
|
||||
tags: Optional[list[str]] = None
|
||||
title_text: Optional[str] = None
|
||||
title_category: Optional[str] = None
|
||||
title_config: Optional[Dict[str, Any]] = None
|
||||
|
||||
@@ -7,15 +7,21 @@ from typing import Optional
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ============ Enums / Types ============
|
||||
# 会员体系(#1951/#1955 实装):
|
||||
# free — 免费用户
|
||||
# monthly — 月卡
|
||||
# quarterly — 季卡
|
||||
# yearly — 年卡
|
||||
# 已废弃档位:standard / pro / enterprise(保留常量名便于识别旧字段,但不在 API 中暴露)
|
||||
|
||||
|
||||
class PlanType(str):
|
||||
"""套餐类型"""
|
||||
class MembershipType(str):
|
||||
"""会员类型(与 packages.domain.points_rules.MEMBERSHIP_PRICES 一致)"""
|
||||
|
||||
FREE = "free"
|
||||
STANDARD = "standard"
|
||||
PRO = "pro"
|
||||
ENTERPRISE = "enterprise"
|
||||
MONTHLY = "monthly"
|
||||
QUARTERLY = "quarterly"
|
||||
YEARLY = "yearly"
|
||||
|
||||
|
||||
class SubscriptionStatus(str):
|
||||
@@ -40,6 +46,7 @@ class BillingCycle(str):
|
||||
"""计费周期"""
|
||||
|
||||
MONTHLY = "monthly"
|
||||
QUARTERLY = "quarterly"
|
||||
YEARLY = "yearly"
|
||||
|
||||
|
||||
@@ -95,8 +102,8 @@ class SimpleResponse(BaseModel):
|
||||
class ChangePlanRequest(BaseModel):
|
||||
"""升级/降级请求"""
|
||||
|
||||
target_plan_id: str = Field(..., description="目标套餐ID")
|
||||
billing_cycle: str = Field(..., description="计费周期: monthly/yearly")
|
||||
target_plan_id: str = Field(..., description="目标会员类型: monthly/quarterly/yearly")
|
||||
billing_cycle: str = Field(..., description="计费周期: monthly/quarterly/yearly")
|
||||
|
||||
|
||||
class ToggleAutoRenewRequest(BaseModel):
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
"""抖音视频解析多源轮询服务。
|
||||
|
||||
优先级(P0 最高):
|
||||
P0: App Feed API 直连(零成本,不用 API Key,当前最稳定)
|
||||
P1: TikHub API(付费 $0.001/次起,稳定)
|
||||
P2: apizero.cn 极数本源(按量付费,国内延迟低)
|
||||
|
||||
任一源成功即返回 MP4 直链 + 标题/文案;所有源均失败时返回 None。
|
||||
每个解析源独立超时(5-10s),总耗时不超过所有源超时之和(实际快速失败时远小于此)。
|
||||
未配置 API Key 的源自动跳过;无任何 Key 时 P0 仍可使用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── API Keys from env ──────────────────────────────────────────────────
|
||||
TIKHUB_API_KEY = os.environ.get("TIKHUB_API_KEY", "").strip()
|
||||
APIZERO_API_KEY = os.environ.get("APIZERO_API_KEY", "").strip()
|
||||
|
||||
# ── Timeouts (seconds) ────────────────────────────────────────────────
|
||||
_TIMEOUT_APP_FEED = 12
|
||||
_TIMEOUT_TIKHUB = 6
|
||||
_TIMEOUT_APIZERO = 6
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolveResult:
|
||||
video_url: str # MP4 直链;图文视频时为空字符串
|
||||
desc: str # 视频标题/描述文案
|
||||
source: str # 解析源名称,用于日志/metrics
|
||||
|
||||
|
||||
# ── URL preprocessing ─────────────────────────────────────────────────
|
||||
_AWEME_ID_RE = re.compile(
|
||||
r"(?:douyin\.com/(?:video|note)/|iesdouyin\.com/share/video/|aweme_id=)(\d{15,25})",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _extract_url_from_text(text: str) -> str:
|
||||
"""从任意分享文本中提取首个 http(s) URL。"""
|
||||
if not text:
|
||||
return ""
|
||||
m = re.search(r"https?://\S+", text)
|
||||
return m.group(0).rstrip("。,!?!?,,;;\"'))】") if m else "" # noqa: B005
|
||||
|
||||
|
||||
def _canonicalize_url(url: str, timeout: int = 8) -> str:
|
||||
"""跟随 v.douyin.com 短链 302 重定向,返回完整 URL。失败时返回原 URL。"""
|
||||
if "v.douyin.com" not in url and "iesdouyin.com" not in url:
|
||||
return url
|
||||
try:
|
||||
with httpx.Client(
|
||||
timeout=timeout, follow_redirects=True, verify=False, headers={"User-Agent": "Mozilla/5.0"}
|
||||
) as c:
|
||||
resp = c.get(url)
|
||||
return str(resp.url)
|
||||
except Exception as exc:
|
||||
logger.debug("短链解析失败: %s (%s)", url, exc)
|
||||
return url
|
||||
|
||||
|
||||
# ── Provider P0: App Feed API (零成本直连) ────────────────────────────
|
||||
def _resolve_app_feed(url: str, timeout: int = _TIMEOUT_APP_FEED) -> Optional[ResolveResult]:
|
||||
"""抖音 Android App Feed API 直连 — 零依赖、无需 Key、目前最稳定。"""
|
||||
from packages.douyin_parser import fetch_douyin_video_url
|
||||
|
||||
video_url, desc = fetch_douyin_video_url(url, timeout=timeout, max_retries=2)
|
||||
if video_url:
|
||||
return ResolveResult(video_url=video_url, desc=desc or "", source="app_feed")
|
||||
if desc:
|
||||
# 图文视频:video_url 为 None 但 desc 可用
|
||||
return ResolveResult(video_url="", desc=desc, source="app_feed_image")
|
||||
return None
|
||||
|
||||
|
||||
# ── Provider P1: TikHub ───────────────────────────────────────────────
|
||||
def _resolve_tikhub(url: str, api_key: str, timeout: int = _TIMEOUT_TIKHUB) -> Optional[ResolveResult]:
|
||||
"""TikHub API: https://api.tikhub.io/
|
||||
两步:get_aweme_id → fetch_one_video
|
||||
"""
|
||||
if not api_key:
|
||||
return None
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
aweme_id = _AWEME_ID_RE.search(url or "")
|
||||
aweme_id = aweme_id.group(1) if aweme_id else None
|
||||
|
||||
if not aweme_id:
|
||||
try:
|
||||
with httpx.Client(timeout=timeout, verify=False) as c:
|
||||
r = c.get(
|
||||
"https://api.tikhub.io/api/v1/douyin/web/get_aweme_id",
|
||||
headers=headers,
|
||||
params={"url": url},
|
||||
)
|
||||
data = r.json()
|
||||
aweme_id = (data.get("data") or {}).get("aweme_id")
|
||||
except Exception as exc:
|
||||
logger.warning("TikHub get_aweme_id 失败: %s", exc)
|
||||
return None
|
||||
if not aweme_id:
|
||||
return None
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=timeout, verify=False) as c:
|
||||
r = c.get(
|
||||
"https://api.tikhub.io/api/v1/douyin/app/v3/fetch_one_video",
|
||||
headers=headers,
|
||||
params={"aweme_id": aweme_id},
|
||||
)
|
||||
data = r.json()
|
||||
video = (data.get("data") or {}).get("video") or {}
|
||||
urls = []
|
||||
for k in ("download_addr", "play_addr_h264", "play_addr"):
|
||||
urls = (video.get(k) or {}).get("url_list") or []
|
||||
if urls:
|
||||
break
|
||||
if not urls:
|
||||
# bit_rate 兜底
|
||||
for br in video.get("bit_rate") or []:
|
||||
urls = (br.get("play_addr") or {}).get("url_list") or []
|
||||
if urls:
|
||||
break
|
||||
if not urls:
|
||||
return None
|
||||
# 优先 CDN 直链
|
||||
video_url = urls[0]
|
||||
for u in urls:
|
||||
if any(h in u for h in ("douyinvod.com", "bytecdn.com", "365yg.com")):
|
||||
video_url = u
|
||||
break
|
||||
desc = (data.get("data") or {}).get("desc", "")
|
||||
# 检测图文
|
||||
images = (data.get("data") or {}).get("images") or []
|
||||
if images and not any(h in video_url for h in ("douyinvod.com", "bytecdn.com", "amemv.com")):
|
||||
# 图文且无视频直链
|
||||
if desc:
|
||||
return ResolveResult(video_url="", desc=desc, source="tikhub_image")
|
||||
return None
|
||||
return ResolveResult(video_url=video_url, desc=desc or "", source="tikhub")
|
||||
except Exception as exc:
|
||||
logger.warning("TikHub fetch_one_video 失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── Provider P2: apizero.cn ──────────────────────────────────────────
|
||||
def _resolve_apizero(url: str, api_key: str, timeout: int = _TIMEOUT_APIZERO) -> Optional[ResolveResult]:
|
||||
"""apizero.cn 极数本源: https://v1.apizero.cn/api/video-parse?url=...&flat=2"""
|
||||
if not api_key:
|
||||
return None
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
try:
|
||||
with httpx.Client(timeout=timeout, verify=False) as c:
|
||||
r = c.get(
|
||||
"https://v1.apizero.cn/api/video-parse",
|
||||
headers=headers,
|
||||
params={"url": url, "flat": 2},
|
||||
)
|
||||
data = r.json()
|
||||
d = data.get("data") or {}
|
||||
video_list = d.get("video_list") or []
|
||||
if not video_list:
|
||||
return None
|
||||
video_url = video_list[0].get("url", "")
|
||||
desc = d.get("title", "") or d.get("desc", "") or d.get("author", "")
|
||||
if not video_url:
|
||||
return None
|
||||
return ResolveResult(video_url=video_url, desc=desc, source="apizero")
|
||||
except Exception as exc:
|
||||
logger.warning("apizero 解析失败: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
# ── Main API ──────────────────────────────────────────────────────────
|
||||
def resolve_douyin_video(page_url: str) -> Optional[ResolveResult]:
|
||||
"""按 P0→P1→P2 顺序轮询解析抖音视频。
|
||||
|
||||
Args:
|
||||
page_url: 抖音 URL 或含 URL 的分享文本。
|
||||
|
||||
Returns:
|
||||
ResolveResult 或 None(所有源均失败)。
|
||||
图文视频时 video_url 为空字符串、desc 为文案。
|
||||
"""
|
||||
url = _extract_url_from_text(page_url) or page_url
|
||||
url = _canonicalize_url(url)
|
||||
|
||||
providers = [
|
||||
("app_feed", lambda: _resolve_app_feed(url)),
|
||||
("tikhub", lambda: _resolve_tikhub(url, TIKHUB_API_KEY)),
|
||||
("apizero", lambda: _resolve_apizero(url, APIZERO_API_KEY)),
|
||||
]
|
||||
|
||||
enabled_count = 0
|
||||
for name, fn in providers:
|
||||
if name == "tikhub" and not TIKHUB_API_KEY:
|
||||
continue
|
||||
if name == "apizero" and not APIZERO_API_KEY:
|
||||
continue
|
||||
enabled_count += 1
|
||||
t0 = time.time()
|
||||
try:
|
||||
result = fn()
|
||||
elapsed = time.time() - t0
|
||||
if result:
|
||||
domain = result.video_url.split("/")[2] if result.video_url and "/" in result.video_url else "(image)"
|
||||
logger.info(
|
||||
"抖音解析成功: source=%s url_domain=%s desc_len=%d time=%.2fs",
|
||||
result.source,
|
||||
domain,
|
||||
len(result.desc),
|
||||
elapsed,
|
||||
)
|
||||
return result
|
||||
logger.debug("解析源 %s 返回空 (%.2fs)", name, elapsed)
|
||||
except Exception as exc:
|
||||
logger.warning("解析源 %s 异常 (%.2fs): %s", name, time.time() - t0, exc)
|
||||
|
||||
if enabled_count == 0:
|
||||
logger.error("无任何抖音解析源可用:请检查 App Feed API 网络连通性")
|
||||
else:
|
||||
logger.warning("所有 %d 个抖音解析源均失败: url=%s", enabled_count, url)
|
||||
return None
|
||||
|
||||
|
||||
def available_providers() -> list[str]:
|
||||
"""返回当前可用的解析源列表(用于诊断)。"""
|
||||
provs = ["app_feed"]
|
||||
if TIKHUB_API_KEY:
|
||||
provs.append("tikhub")
|
||||
if APIZERO_API_KEY:
|
||||
provs.append("apizero")
|
||||
return provs
|
||||
@@ -423,6 +423,7 @@ class EditPlanService:
|
||||
clip_type=clip.clip_type,
|
||||
order=clip.order,
|
||||
asset_id=clip.asset_id,
|
||||
atom_clip_id=clip_item.get("atom_clip_id", ""),
|
||||
text_content=clip.text_content,
|
||||
start_time=clip.start_time,
|
||||
duration=clip.duration,
|
||||
@@ -474,6 +475,7 @@ class EditPlanService:
|
||||
voice_duration: float = 0.0,
|
||||
rng=None,
|
||||
batch_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
batch_used_atom_ids: set[str] | list[str] | None = None,
|
||||
) -> EditPlan:
|
||||
"""为批量变体生成独立 plan:完整重跑单视频选片流程(#1743)。
|
||||
|
||||
@@ -608,18 +610,69 @@ class EditPlanService:
|
||||
st = float(c.start_time or 0.0)
|
||||
batch_segments_resolved.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
|
||||
|
||||
clips_data = reselect_clips_for_variant(
|
||||
source_clips_data,
|
||||
pool_ids,
|
||||
asset_durations=durations,
|
||||
asset_scene_points=scene_points,
|
||||
historical_used_segments=historical,
|
||||
batch_segments=batch_segments_resolved,
|
||||
target_durations=target_durations,
|
||||
rng=rng,
|
||||
)
|
||||
clips_data = None
|
||||
# #1970 原子片段级变体重选:候选素材已切片时优先按原子片段选片
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.domain.atom_clip_resolver import flatten_candidates, load_atom_clips_for_assets
|
||||
from packages.domain.atom_clip_selector import reselect_clips_from_atoms
|
||||
|
||||
# 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit)
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
|
||||
# 兜底切片只需要时长;本方法已查出 durations,封装一个只读假素材仓储
|
||||
class _DurationOnlyAssetRepo:
|
||||
def __init__(self, durations_map: dict[str, float]) -> None:
|
||||
self._durations = durations_map
|
||||
|
||||
def get(self, asset_id: str):
|
||||
if asset_id not in self._durations:
|
||||
return None
|
||||
|
||||
class _A:
|
||||
pass
|
||||
|
||||
a = _A()
|
||||
a.duration = self._durations[asset_id]
|
||||
return a
|
||||
|
||||
clips_by_asset = load_atom_clips_for_assets(
|
||||
pool_ids,
|
||||
atom_clip_repo=atom_repo,
|
||||
asset_repo=_DurationOnlyAssetRepo(durations),
|
||||
)
|
||||
atom_candidates = flatten_candidates(clips_by_asset)
|
||||
if atom_candidates:
|
||||
# 历史成片已用原子片段(降权);批次内前序变体已用(硬避让)
|
||||
historical_atom_ids = set(
|
||||
self._clip_repo.list_recent_atom_clip_ids_by_user(
|
||||
created_by_user_id or source.created_by_user_id or "",
|
||||
limit=200,
|
||||
)
|
||||
)
|
||||
clips_data = reselect_clips_from_atoms(
|
||||
source_clips_data,
|
||||
atom_candidates,
|
||||
historical_atom_ids=historical_atom_ids,
|
||||
batch_used_atom_ids=(set(batch_used_atom_ids) if batch_used_atom_ids else None),
|
||||
rng=rng,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("原子片段变体重选失败,回退整条素材选片", exc_info=True)
|
||||
clips_data = None
|
||||
|
||||
if clips_data is None:
|
||||
clips_data = reselect_clips_for_variant(
|
||||
source_clips_data,
|
||||
pool_ids,
|
||||
asset_durations=durations,
|
||||
asset_scene_points=scene_points,
|
||||
historical_used_segments=historical,
|
||||
batch_segments=batch_segments_resolved,
|
||||
target_durations=target_durations,
|
||||
rng=rng,
|
||||
) # 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit)
|
||||
for item in clips_data:
|
||||
aid = item.get("asset_id", "")
|
||||
if aid:
|
||||
|
||||
@@ -61,10 +61,17 @@ def writeback_edit_plan_config(
|
||||
task_id: str,
|
||||
title_config: dict | None,
|
||||
db: Session,
|
||||
dedup_enabled: bool | None = None,
|
||||
video_index: int | None = None,
|
||||
assembly_mode: str | None = None,
|
||||
script_id: str | None = None,
|
||||
video_ratio: str | None = None,
|
||||
) -> None:
|
||||
"""任务入队成功后,回写 EditPlan.config:generation_task_id + title_config。
|
||||
|
||||
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
|
||||
#1970:dedup_enabled 非 None 时一并写入,worker 据此决定 edge_crop/微变换;
|
||||
PR3 叙事模式再写 assembly_mode/script_id/video_ratio(可追溯,不影响渲染)。
|
||||
失败只记日志,不影响任务创建。
|
||||
"""
|
||||
if not plan_id:
|
||||
@@ -80,6 +87,16 @@ def writeback_edit_plan_config(
|
||||
current_config = plan_model.config if isinstance(plan_model.config, dict) else {}
|
||||
merged = dict(current_config)
|
||||
merged["generation_task_id"] = task_id
|
||||
if dedup_enabled is not None:
|
||||
merged["dedup_enabled"] = bool(dedup_enabled)
|
||||
if video_index is not None:
|
||||
merged["video_index"] = int(video_index)
|
||||
if assembly_mode:
|
||||
merged["assembly_mode"] = assembly_mode
|
||||
if script_id:
|
||||
merged["script_id"] = script_id
|
||||
if video_ratio:
|
||||
merged["video_ratio"] = video_ratio
|
||||
|
||||
if title_config:
|
||||
# #1901 统一字段名为 "title"(worker sync_configs_to_plan 写的是 "title")
|
||||
@@ -157,6 +174,33 @@ def collect_plan_segments(
|
||||
return segs
|
||||
|
||||
|
||||
def collect_plan_atom_clip_ids(
|
||||
plan_id: str,
|
||||
clip_repo: Any,
|
||||
*,
|
||||
page_size: int = 500,
|
||||
) -> list[str]:
|
||||
"""分页读取 plan 所有 clips,收集已选用的原子片段 ID(#1970)。
|
||||
|
||||
用于批量变体间原子片段级硬避让:同一原子片段在同批次内只用一次。
|
||||
旧路径 clips 的 atom_clip_id 为空串,自动忽略。
|
||||
"""
|
||||
ids: list[str] = []
|
||||
sk, pg = 0, page_size
|
||||
while True:
|
||||
batch = clip_repo.list_by_plan(plan_id, skip=sk, limit=pg)
|
||||
if not batch:
|
||||
break
|
||||
for c in batch:
|
||||
acid = getattr(c, "atom_clip_id", "") or ""
|
||||
if acid:
|
||||
ids.append(acid)
|
||||
if len(batch) < pg:
|
||||
break
|
||||
sk += pg
|
||||
return ids
|
||||
|
||||
|
||||
def resolve_latest_plan_by_template(
|
||||
db: Session,
|
||||
*,
|
||||
|
||||
@@ -0,0 +1,382 @@
|
||||
"""GPU MuseTalk 口型同步服务 — 反向轮询模式.
|
||||
|
||||
职责:
|
||||
1. 创建任务(由 lipsync 业务流程调用),为输入/输出生成预签名 URL,任务入队;
|
||||
2. Worker 心跳注册(register):登记/刷新 worker 状态;
|
||||
3. Worker 轮询拉任务(poll):原子地 CLAIM 一条 pending 任务,返回预签名 URL;
|
||||
4. Worker 上报结果(report_result):标记 done/failed,失败可重试;
|
||||
5. 业务侧查询状态(get_status)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Optional
|
||||
|
||||
from app.core.storage import get_storage_service
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel, GpuWorkerModel
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 任务在 processing 超过此时长仍未完成 → 超时回退 pending 或置 failed
|
||||
MAX_ATTEMPTS = 3
|
||||
|
||||
|
||||
class GpuLipsyncService:
|
||||
"""GPU 口型同步服务(无状态方法,每次调用从 DI 拿 db/storage)."""
|
||||
|
||||
RESULT_PREFIX = "gpu-lipsync/results/"
|
||||
INPUT_SIGN_EXPIRES_PAD = 600 # 输入预签名 URL 在任务超时基础上再加 10min 余量
|
||||
|
||||
# ── 公共入口 ────────────────────────────────────────────────────
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
self.settings = get_api_settings()
|
||||
self.storage = get_storage_service()
|
||||
|
||||
# ── Worker 注册/心跳 ────────────────────────────────────────────
|
||||
|
||||
def register_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
hostname: str = "",
|
||||
gpu_name: str = "",
|
||||
free_vram_mb: int = 0,
|
||||
capabilities: str = "musetalk",
|
||||
task_id: Optional[str] = None,
|
||||
) -> GpuWorkerModel:
|
||||
"""Worker 注册/心跳。
|
||||
|
||||
task_id 非空时(Worker 推理期间的任务级心跳),同步把对应 processing
|
||||
任务的 last_heartbeat_at 续到当前时间,使长推理不会被
|
||||
``_recover_timed_out_tasks`` 误回退。任务已结束 / 不属于该 worker
|
||||
(如已被超时回收重新派发)时忽略,不报错。
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none()
|
||||
if worker is None:
|
||||
worker = GpuWorkerModel(
|
||||
worker_id=worker_id,
|
||||
hostname=hostname,
|
||||
gpu_name=gpu_name,
|
||||
free_vram_mb=free_vram_mb,
|
||||
capabilities=capabilities,
|
||||
last_heartbeat_at=now,
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(worker)
|
||||
else:
|
||||
worker.hostname = hostname or worker.hostname
|
||||
worker.gpu_name = gpu_name or worker.gpu_name
|
||||
worker.free_vram_mb = free_vram_mb
|
||||
worker.capabilities = capabilities or worker.capabilities
|
||||
worker.last_heartbeat_at = now
|
||||
if task_id:
|
||||
self._touch_task_heartbeat(task_id, worker_id, now)
|
||||
self.db.commit()
|
||||
return worker
|
||||
|
||||
# ── 轮询拉任务(Worker 调用) ──────────────────────────────────
|
||||
|
||||
def poll_task(self, worker_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
"""原子地认领一条最早的 pending 任务,返回给 worker;无任务返回 None.
|
||||
|
||||
同时会:
|
||||
- 把 processing 状态且真正超时(任务心跳停滞超过
|
||||
gpu_task_timeout_seconds;Worker 推理期会通过 register(task_id=...)
|
||||
续心跳,长推理不会误判)的任务回退为 pending(attempt++,超过
|
||||
MAX_ATTEMPTS 置 failed),让其它 worker 认领。
|
||||
- 刷新 worker 心跳。
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
self._recover_timed_out_tasks(now)
|
||||
# 更新 worker 心跳
|
||||
self._touch_worker(worker_id, now)
|
||||
|
||||
# 选一条最早 pending 任务(FOR UPDATE SKIP LOCKED 语义:简单起见先查再锁状态)
|
||||
task = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(GpuLipsyncTaskModel.status == "pending")
|
||||
.order_by(GpuLipsyncTaskModel.created_at.asc())
|
||||
.first()
|
||||
)
|
||||
if task is None:
|
||||
self.db.commit()
|
||||
return None
|
||||
|
||||
# 原子 claim:用 UPDATE WHERE status=pending 避免并发
|
||||
upd_rows = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(
|
||||
GpuLipsyncTaskModel.id == task.id,
|
||||
GpuLipsyncTaskModel.status == "pending",
|
||||
)
|
||||
.update(
|
||||
{
|
||||
GpuLipsyncTaskModel.status: "processing",
|
||||
GpuLipsyncTaskModel.worker_id: worker_id,
|
||||
GpuLipsyncTaskModel.started_at: now,
|
||||
GpuLipsyncTaskModel.last_heartbeat_at: now,
|
||||
GpuLipsyncTaskModel.attempt: GpuLipsyncTaskModel.attempt + 1,
|
||||
GpuLipsyncTaskModel.updated_at: now,
|
||||
},
|
||||
synchronize_session=False,
|
||||
)
|
||||
)
|
||||
self.db.commit()
|
||||
if upd_rows == 0:
|
||||
# 被其它 worker 抢先了
|
||||
return None
|
||||
self.db.refresh(task)
|
||||
# 生成预签名输入/输出 URL(在 claim 时动态生成,避免长时间过期)
|
||||
expires = self.settings.gpu_task_timeout_seconds + self.INPUT_SIGN_EXPIRES_PAD
|
||||
task._signed_video_url = self.storage.get_download_url(task.video_url, expires_seconds=expires)
|
||||
task._signed_audio_url = self.storage.get_download_url(task.audio_url, expires_seconds=expires)
|
||||
task._signed_upload_url = self.storage.get_upload_url(
|
||||
self._result_key(task.id),
|
||||
expires_seconds=expires,
|
||||
content_type="video/mp4",
|
||||
)
|
||||
task._upload_expires_at = now + timedelta(seconds=expires)
|
||||
return task
|
||||
|
||||
# ── 上报结果 ──────────────────────────────────────────────────
|
||||
|
||||
def report_result(
|
||||
self,
|
||||
task_id: str,
|
||||
worker_id: str,
|
||||
success: bool,
|
||||
duration_seconds: float = 0.0,
|
||||
error_msg: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task = self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
if task is None:
|
||||
raise KeyError(f"task {task_id} not found")
|
||||
now = datetime.now(UTC)
|
||||
if success:
|
||||
task.status = "done"
|
||||
task.result_url = self._result_key(task_id)
|
||||
task.result_duration = duration_seconds or 0.0
|
||||
task.error_msg = ""
|
||||
task.finished_at = now
|
||||
else:
|
||||
# 失败:若仍可重试(已尝试次数 < MAX_ATTEMPTS)→ 回退 pending;否则 → failed
|
||||
if task.attempt < MAX_ATTEMPTS:
|
||||
task.status = "pending"
|
||||
task.worker_id = ""
|
||||
task.started_at = None
|
||||
task.error_msg = error_msg[:2000]
|
||||
logger.warning(
|
||||
"GPU 任务 %s 在 worker %s 上失败,回退 pending 等待重试(attempt=%d): %s",
|
||||
task_id,
|
||||
worker_id,
|
||||
task.attempt,
|
||||
error_msg[:200],
|
||||
)
|
||||
else:
|
||||
task.status = "failed"
|
||||
task.error_msg = error_msg[:2000]
|
||||
task.finished_at = now
|
||||
logger.error(
|
||||
"GPU 任务 %s 失败达到最大重试次数 %d,置为 failed: %s",
|
||||
task_id,
|
||||
MAX_ATTEMPTS,
|
||||
error_msg[:200],
|
||||
)
|
||||
task.updated_at = now
|
||||
task.last_heartbeat_at = now
|
||||
self._touch_worker(worker_id, now)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
return task
|
||||
|
||||
# ── 业务侧查询 ────────────────────────────────────────────────
|
||||
|
||||
def get_task(self, task_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
return self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
|
||||
def get_by_lipsync_job(self, lipsync_job_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
return (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(GpuLipsyncTaskModel.lipsync_job_id == lipsync_job_id)
|
||||
.order_by(GpuLipsyncTaskModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
# ── 创建任务(业务侧调用) ────────────────────────────────────
|
||||
|
||||
def create_task(
|
||||
self,
|
||||
video_url: str,
|
||||
audio_url: str,
|
||||
lipsync_job_id: str = "",
|
||||
user_id: str = "",
|
||||
project_id: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task_id = str(uuid.uuid4())
|
||||
now = datetime.now(UTC)
|
||||
task = GpuLipsyncTaskModel(
|
||||
id=task_id,
|
||||
lipsync_job_id=lipsync_job_id,
|
||||
user_id=user_id,
|
||||
project_id=project_id,
|
||||
video_url=video_url,
|
||||
audio_url=audio_url,
|
||||
status="pending",
|
||||
attempt=0,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
self.db.add(task)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
logger.info(
|
||||
"创建 GPU 口型任务 %s (lipsync_job=%s, user=%s)",
|
||||
task_id,
|
||||
lipsync_job_id,
|
||||
user_id,
|
||||
)
|
||||
return task
|
||||
|
||||
# ── 内部辅助 ──────────────────────────────────────────────────
|
||||
|
||||
def _result_key(self, task_id: str) -> str:
|
||||
return f"{self.RESULT_PREFIX}{task_id}.mp4"
|
||||
|
||||
def _touch_task_heartbeat(self, task_id: str, worker_id: str, now: datetime) -> None:
|
||||
"""Worker 推理期间的任务级心跳:只刷新属于该 worker 且仍在 processing 的任务。
|
||||
|
||||
任务不存在 / 已被超时回收重新派发 / 已完成 → 静默忽略(此时旧 worker 的
|
||||
结果上报会被结果接口按最终态处理)。
|
||||
"""
|
||||
task = self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
if task is None:
|
||||
return
|
||||
if task.status != "processing" or task.worker_id != worker_id:
|
||||
logger.info(
|
||||
"忽略过期任务心跳 task=%s worker=%s(status=%s owner=%s)",
|
||||
task_id,
|
||||
worker_id,
|
||||
task.status,
|
||||
task.worker_id,
|
||||
)
|
||||
return
|
||||
task.last_heartbeat_at = now
|
||||
task.updated_at = now
|
||||
self.db.flush()
|
||||
|
||||
def _touch_worker(self, worker_id: str, now: datetime) -> None:
|
||||
if not worker_id:
|
||||
return
|
||||
worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none()
|
||||
if worker is not None:
|
||||
worker.last_heartbeat_at = now
|
||||
self.db.flush()
|
||||
else:
|
||||
# 自注册(poll 时允许自动建一个空 worker 记录,运维可见)
|
||||
worker = GpuWorkerModel(
|
||||
worker_id=worker_id,
|
||||
hostname="",
|
||||
gpu_name="",
|
||||
free_vram_mb=0,
|
||||
capabilities="musetalk",
|
||||
last_heartbeat_at=now,
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(worker)
|
||||
self.db.flush()
|
||||
|
||||
def _recover_timed_out_tasks(self, now: datetime) -> None:
|
||||
"""扫描 processing 状态且真正超时的任务,回退 pending 或失败。
|
||||
|
||||
判定只看任务自身 last_heartbeat_at:claim 时写入,Worker 推理期间通过
|
||||
/gpu/register(task_id=...) 每 30s 续期。因此仅在 Worker 崩溃/断网
|
||||
(任务心跳停滞超过 gpu_task_timeout_seconds)时才回收,
|
||||
不会因 Worker 主循环忙于推理而误回退。
|
||||
"""
|
||||
timeout = self.settings.gpu_task_timeout_seconds
|
||||
cutoff = now - timedelta(seconds=timeout)
|
||||
stuck_tasks = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(
|
||||
GpuLipsyncTaskModel.status == "processing",
|
||||
GpuLipsyncTaskModel.last_heartbeat_at < cutoff,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for t in stuck_tasks:
|
||||
if t.attempt >= MAX_ATTEMPTS:
|
||||
t.status = "failed"
|
||||
t.error_msg = f"worker 心跳超时({timeout}s),重试次数已耗尽"
|
||||
t.finished_at = now
|
||||
else:
|
||||
t.status = "pending"
|
||||
t.worker_id = ""
|
||||
t.started_at = None
|
||||
t.error_msg = f"worker 心跳超时({timeout}s),等待重试"
|
||||
logger.warning("GPU 任务 %s 心跳超时,回退 pending(attempt=%d)", t.id, t.attempt)
|
||||
t.updated_at = now
|
||||
if stuck_tasks:
|
||||
self.db.flush()
|
||||
|
||||
# ── 业务侧辅助 ──────────────────────────────────────────────────
|
||||
|
||||
def has_available_worker(self) -> bool:
|
||||
"""判断是否有 Worker 在心跳新鲜窗口内可用."""
|
||||
stale_cutoff = datetime.now(UTC) - timedelta(seconds=self.settings.gpu_worker_stale_seconds)
|
||||
return (
|
||||
self.db.query(GpuWorkerModel).filter(GpuWorkerModel.last_heartbeat_at >= stale_cutoff).first() is not None
|
||||
)
|
||||
|
||||
def wait_for_result(
|
||||
self,
|
||||
task_id: str,
|
||||
timeout_seconds: Optional[int] = None,
|
||||
poll_interval: Optional[float] = None,
|
||||
) -> Optional[GpuLipsyncTaskModel]:
|
||||
"""同步轮询等待 GPU 任务完成。
|
||||
|
||||
Args:
|
||||
task_id: 任务 ID(由 create_task 返回)
|
||||
timeout_seconds: 总超时,默认取 settings.gpu_lipsync_wait_timeout
|
||||
poll_interval: 轮询间隔秒,默认取 settings.gpu_lipsync_poll_interval
|
||||
|
||||
Returns:
|
||||
终态 task(status=done/failed);超时返回 None(此时调用方应回退 MediaKit)。
|
||||
等待期间会自动调用 _recover_timed_out_tasks 做超时回收。
|
||||
"""
|
||||
import time
|
||||
|
||||
timeout = timeout_seconds if timeout_seconds is not None else self.settings.gpu_lipsync_wait_timeout
|
||||
interval = poll_interval if poll_interval is not None else self.settings.gpu_lipsync_poll_interval
|
||||
deadline = time.monotonic() + timeout
|
||||
|
||||
while True:
|
||||
now = datetime.now(UTC)
|
||||
# 顺手回收超时任务
|
||||
try:
|
||||
self._recover_timed_out_tasks(now)
|
||||
self.db.commit()
|
||||
except Exception as exc: # noqa: BLE001 - 回收失败不阻塞主流程
|
||||
logger.warning("wait_for_result 回收超时任务异常: %s", exc)
|
||||
self.db.rollback()
|
||||
|
||||
task = self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
if task is None:
|
||||
return None
|
||||
if task.status == "done":
|
||||
return task
|
||||
if task.status == "failed":
|
||||
return task
|
||||
# pending/processing 继续等
|
||||
if time.monotonic() >= deadline:
|
||||
logger.warning("GPU 任务 %s 等待超时(%ds),回退 MediaKit", task_id, timeout)
|
||||
return None
|
||||
time.sleep(interval)
|
||||
@@ -29,6 +29,7 @@ from app.services.mediakit_client import (
|
||||
MediaKitError,
|
||||
get_mediakit_client,
|
||||
)
|
||||
from app.tasks.lipsync_gpu import lipsync_gpu_process_async
|
||||
|
||||
# Celery 异步任务:TTS 合成 + MediaKit 提交(降级路径)
|
||||
from app.tasks.lipsync_tts import tts_synthesize_and_submit
|
||||
@@ -36,6 +37,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
from packages.application.cosyvoice_service import CosyVoiceError
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain.sentence_timings import (
|
||||
compute_sentence_timings,
|
||||
probe_audio_duration,
|
||||
@@ -63,6 +65,7 @@ class LipsyncService:
|
||||
self.client = client or get_mediakit_client()
|
||||
self._cosyvoice = cosyvoice_service
|
||||
self._voice_clone_repo = voice_clone_repo
|
||||
self.settings = get_api_settings()
|
||||
|
||||
def _get_cosyvoice(self):
|
||||
"""延迟获取 CosyVoiceService(与 tts 路由一致,含 OSS 预签名配置)."""
|
||||
@@ -215,7 +218,57 @@ class LipsyncService:
|
||||
if timings:
|
||||
job.sentence_timings = timings
|
||||
|
||||
# 4. 签名 URL 并提交 MediaKit
|
||||
# 4. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
||||
use_gpu = False
|
||||
if self.settings.use_gpu_lipsync:
|
||||
try:
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
|
||||
gpu_svc = GpuLipsyncService(self.db)
|
||||
if gpu_svc.has_available_worker():
|
||||
use_gpu = True
|
||||
logger.info("[lipsync] 检测到可用 GPU Worker,优先走 MuseTalk 本地推理: job_id=%s", job.id)
|
||||
else:
|
||||
logger.info("[lipsync] GPU 开关已开但无可用 Worker(心跳过期),回退 MediaKit: job_id=%s", job.id)
|
||||
except Exception as exc:
|
||||
logger.warning("[lipsync] GPU 服务初始化失败,回退 MediaKit: job_id=%s err=%s", job.id, exc)
|
||||
|
||||
if use_gpu:
|
||||
try:
|
||||
gpu_task = self._submit_to_gpu_create(job=job, gpu_svc=gpu_svc)
|
||||
if gpu_task is not None:
|
||||
# GPU 任务已创建,设为 processing 并异步等待结果
|
||||
job.mediakit_task_id = f"gpu:{gpu_task.id}"
|
||||
job.status = "processing"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
# 派发 Celery 异步任务处理 GPU 等待+结果回写
|
||||
try:
|
||||
lipsync_gpu_process_async.apply_async(args=(job.id, job.user_id, gpu_task.id))
|
||||
logger.info(
|
||||
"[lipsync] GPU 任务已异步派发: job_id=%s gpu_task=%s",
|
||||
job.id,
|
||||
gpu_task.id,
|
||||
)
|
||||
except Exception as celery_exc:
|
||||
logger.warning(
|
||||
"[lipsync] Celery 派发失败,降级同步等待: job_id=%s err=%s",
|
||||
job.id,
|
||||
celery_exc,
|
||||
)
|
||||
self._submit_to_gpu_wait(job=job, gpu_svc=gpu_svc, gpu_task=gpu_task)
|
||||
return
|
||||
# create 失败 → 回退 MediaKit
|
||||
logger.warning("[lipsync] GPU 任务创建失败,回退 MediaKit: job_id=%s", job.id)
|
||||
self.db.rollback()
|
||||
except Exception as exc:
|
||||
logger.exception("[lipsync] GPU 路径异常,回退 MediaKit: job_id=%s err=%s", job.id, exc)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 5. 签名 URL 并提交 MediaKit(兜底路径)
|
||||
video_url = self._sign_media_url(job.video_url)
|
||||
signed_audio_url = self._sign_media_url(job.audio_url)
|
||||
job.audio_url = signed_audio_url
|
||||
@@ -244,6 +297,120 @@ class LipsyncService:
|
||||
self.db.commit()
|
||||
raise
|
||||
|
||||
# ── GPU MuseTalk 路径 ────────────────────────────────────────────────
|
||||
|
||||
def _is_own_oss_url(self, url: str, storage) -> bool:
|
||||
"""判断 URL / 存储 key 是否属于自家 OSS。
|
||||
|
||||
- 裸存储 key(无 scheme):自家对象
|
||||
- host 与 storage.public_url host 一致:自家对象
|
||||
- 其余 http(s) 公网链接(如 dashscope-result 临时地址):外部对象
|
||||
"""
|
||||
if not url:
|
||||
return False
|
||||
parsed = urlparse(url)
|
||||
if not parsed.scheme:
|
||||
return True # 裸存储 key
|
||||
public_base = getattr(storage, "public_url", "")
|
||||
own_host = urlparse(public_base).netloc.lower() if public_base else ""
|
||||
return bool(own_host) and parsed.netloc.lower() == own_host
|
||||
|
||||
def _persist_external_audio_for_gpu(self, *, job, storage) -> Optional[str]:
|
||||
"""GPU 任务创建前,把外部域名的预合成 TTS 音频转存到自家 OSS。
|
||||
|
||||
Worker 部署在用户家庭网络,dashscope-result 等第三方临时 OSS 地址
|
||||
可能无法访问;转存后 gpu_svc 在 poll 时会签自家预签名 URL 给 Worker。
|
||||
已是自家 OSS 对象(含裸 key)直接返回 None(无需转存);
|
||||
转存失败返回 None,调用方回退使用原始 URL(最坏情况是 Worker 拉取失败,
|
||||
服务端重试耗尽后回退 MediaKit,不阻断业务)。
|
||||
"""
|
||||
if self._is_own_oss_url(job.audio_url, storage):
|
||||
return None
|
||||
try:
|
||||
audio_data = safe_download_bytes(
|
||||
job.audio_url,
|
||||
purpose="lipsync_gpu_tts_audio",
|
||||
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
|
||||
timeout=60.0,
|
||||
)
|
||||
storage_key = f"lipsync-tts/{job.user_id}/{job.id}.mp3"
|
||||
permanent_url = storage.upload_file(io.BytesIO(audio_data), storage_key, content_type="audio/mpeg")
|
||||
logger.info(
|
||||
"[lipsync] GPU 任务外部音频已转存自家 OSS: job_id=%s key=%s",
|
||||
job.id,
|
||||
storage_key,
|
||||
)
|
||||
return permanent_url
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[lipsync] GPU 任务外部音频转存 OSS 失败,回退原始 URL: job_id=%s err=%s",
|
||||
job.id,
|
||||
exc,
|
||||
)
|
||||
return None
|
||||
|
||||
def _submit_to_gpu_create(self, *, job, gpu_svc) -> Optional[object]:
|
||||
"""创建 GPU 任务并立即返回(异步模式)。
|
||||
|
||||
成功返回 gpu_task 对象;创建失败返回 None。
|
||||
不再同步等待结果,结果由 Celery 异步任务 lipsync_gpu_process_async 回写。
|
||||
"""
|
||||
storage = get_shared_storage_service()
|
||||
persisted_audio_url = self._persist_external_audio_for_gpu(job=job, storage=storage)
|
||||
audio_url_for_task = persisted_audio_url or job.audio_url
|
||||
gpu_task = gpu_svc.create_task(
|
||||
video_url=job.video_url,
|
||||
audio_url=audio_url_for_task,
|
||||
lipsync_job_id=job.id,
|
||||
user_id=job.user_id,
|
||||
project_id=job.project_id,
|
||||
)
|
||||
logger.info(
|
||||
"[lipsync] 已创建 GPU 任务(异步): job_id=%s gpu_task=%s",
|
||||
job.id,
|
||||
gpu_task.id,
|
||||
)
|
||||
return gpu_task
|
||||
|
||||
def _submit_to_gpu_wait(self, *, job, gpu_svc, gpu_task) -> None:
|
||||
"""同步等待 GPU 结果(Celery 派发失败时的降级路径)。"""
|
||||
final_task = gpu_svc.wait_for_result(gpu_task.id)
|
||||
if final_task is None:
|
||||
logger.warning("[lipsync] GPU 同步等待超时,回退 MediaKit: gpu_task=%s", gpu_task.id)
|
||||
return
|
||||
if final_task.status != "done":
|
||||
logger.warning(
|
||||
"[lipsync] GPU 同步等待失败: gpu_task=%s status=%s",
|
||||
gpu_task.id,
|
||||
final_task.status,
|
||||
)
|
||||
return
|
||||
try:
|
||||
storage = get_shared_storage_service()
|
||||
signed_result_url = storage.get_download_url(
|
||||
final_task.result_url, expires_seconds=MEDIAKIT_URL_TTL_SECONDS
|
||||
)
|
||||
if signed_result_url:
|
||||
final_task.result_url = signed_result_url
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[lipsync] GPU 结果签名失败: gpu_task=%s err=%s",
|
||||
gpu_task.id,
|
||||
exc,
|
||||
)
|
||||
job.mediakit_task_id = ""
|
||||
job.status = STATUS_COMPLETED
|
||||
job.output_video_url = final_task.result_url
|
||||
job.output_duration = final_task.result_duration or 0.0
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
logger.info(
|
||||
"[lipsync] GPU 同步等待完成: job_id=%s duration=%.2f",
|
||||
job.id,
|
||||
job.output_duration,
|
||||
)
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────────
|
||||
|
||||
def create_job(
|
||||
@@ -478,6 +645,29 @@ class LipsyncService:
|
||||
if job.status in (STATUS_COMPLETED, "failed"):
|
||||
return job
|
||||
|
||||
# GPU 异步路径:mediakit_task_id 以 "gpu:" 开头,由 Celery 任务异步更新
|
||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>30 分钟)则标失败
|
||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"):
|
||||
if job.status in ("processing", "gpu_processing"):
|
||||
_now = datetime.now(UTC)
|
||||
_upd = job.updated_at
|
||||
if _upd is not None and _upd.tzinfo is None:
|
||||
_upd = _upd.replace(tzinfo=UTC)
|
||||
stale_minutes = 30
|
||||
if _upd and (_now - _upd).total_seconds() > stale_minutes * 60:
|
||||
logger.warning(
|
||||
"GPU 异步任务超时(>%d 分钟),标记失败: job_id=%s",
|
||||
stale_minutes,
|
||||
job_id,
|
||||
)
|
||||
job.status = "failed"
|
||||
job.error_message = f"GPU 处理超时(>{stale_minutes} 分钟)"
|
||||
job.error_code = "GpuTimeout"
|
||||
job.completed_at = _now
|
||||
job.updated_at = _now
|
||||
self.db.commit()
|
||||
return job
|
||||
|
||||
# 未提交的任务不轮询
|
||||
if not job.mediakit_task_id:
|
||||
return job
|
||||
|
||||
@@ -0,0 +1,344 @@
|
||||
"""叙事剪辑前置服务 — #1970 PR3.
|
||||
|
||||
叙事模式(assembly_mode='narrative')在生成任务入队前同步完成:
|
||||
|
||||
1. 按 script_id 读取文案(归属校验);
|
||||
2. 按 tts_voice_source 解析音色(preset=CosyVoice 音色 id;clone=克隆档案 id,
|
||||
解析档案归属并取其 CosyVoice voice_id);
|
||||
3. 同步 TTS 合成(复用 tts_job 现有 workflow:提交即同步返回,未完成则轮询兜底),
|
||||
失败直接抛 NarrativeError(HTTP 层转 4xx,任务不入队);
|
||||
4. 把合成音频转存为配音库 audio asset(与 /tts/jobs/{id}/save-to-library 同一套
|
||||
存储路径与元信息约定),返回 asset_id —— 下游仍以 voice_library_id(实为
|
||||
audio asset id)消费,渲染链路零改动。
|
||||
|
||||
积分扣点与 /tts 合成端点保持一致(ai_voice 场景),失败退费。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import ScriptModel
|
||||
from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.use_cases import CreateTTSJobUseCase
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.shared.storage import SharedStorageService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_POINTS_SCENE = "ai_voice"
|
||||
_SYNTH_TIMEOUT = 180.0 # 叙事配音在 HTTP 请求内同步等待,长文案分段合成时留出余量
|
||||
_CONTENT_TYPE_MAP = {"mp3": "audio/mpeg", "wav": "audio/wav", "pcm": "audio/pcm", "opus": "audio/opus"}
|
||||
|
||||
|
||||
class NarrativeError(Exception):
|
||||
"""叙事模式前置处理失败(文案/音色/TTS/落库)。"""
|
||||
|
||||
def __init__(self, message: str, *, status_code: int = 400) -> None:
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class NarrativeContext:
|
||||
"""叙事模式前置处理结果。"""
|
||||
|
||||
script: ScriptModel
|
||||
voice_asset_id: str
|
||||
tts_job_id: str
|
||||
audio_duration: float
|
||||
|
||||
|
||||
def _find_or_create_voice_library(
|
||||
*,
|
||||
user_id: str,
|
||||
project_repository: Any,
|
||||
asset_library_repository: Any,
|
||||
) -> AssetLibrary:
|
||||
"""找到(或自动创建)用户 voice 素材库;与 tts.py 保存配音库逻辑一致。"""
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
if not projects:
|
||||
raise NarrativeError("没有可用的项目,无法保存叙事配音", status_code=400)
|
||||
|
||||
for project in projects:
|
||||
for lib in asset_library_repository.find_by_project(project.id):
|
||||
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
|
||||
if kind == AssetLibraryKind.VOICE.value:
|
||||
return lib
|
||||
|
||||
project = projects[0]
|
||||
library = AssetLibrary.create(project_id=project.id, name="配音素材库", kind=AssetLibraryKind.VOICE)
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
try:
|
||||
return asset_library_repository.create(library)
|
||||
except IntegrityError:
|
||||
session = getattr(asset_library_repository, "session", None)
|
||||
if session is not None:
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception: # noqa: BLE001 - 回滚失败不影响重查
|
||||
logger.warning("IntegrityError 后回滚 session 失败", exc_info=True)
|
||||
for lib in asset_library_repository.find_by_project(project.id):
|
||||
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
|
||||
if kind == AssetLibraryKind.VOICE.value:
|
||||
return lib
|
||||
raise NarrativeError("配音素材库创建失败,请重试", status_code=500) from None
|
||||
|
||||
|
||||
def _resolve_voice(
|
||||
*,
|
||||
user_id: str,
|
||||
tts_voice_id: str,
|
||||
tts_voice_source: str,
|
||||
voice_clone_repository: Any,
|
||||
) -> tuple[str, str]:
|
||||
"""解析音色 → (CosyVoice voice_id, voice_clone_profile_id)。"""
|
||||
if tts_voice_source == "clone":
|
||||
profile = voice_clone_repository.get(tts_voice_id)
|
||||
if profile is None:
|
||||
raise NarrativeError("克隆音色不存在", status_code=404)
|
||||
if profile.user_id != user_id:
|
||||
raise NarrativeError("无权使用该克隆音色", status_code=403)
|
||||
if not profile.voice_id:
|
||||
raise NarrativeError("音色克隆尚未完成,请稍后再试", status_code=400)
|
||||
return profile.voice_id, profile.id
|
||||
# preset:tts_voice_id 即 CosyVoice 音色 id;与 /tts 端点一致,
|
||||
# 若前端误传克隆档案 UUID,同样兼容解析。
|
||||
profile = voice_clone_repository.get(tts_voice_id)
|
||||
if profile is not None:
|
||||
if profile.user_id != user_id:
|
||||
raise NarrativeError("无权使用该音色", status_code=403)
|
||||
if not profile.voice_id:
|
||||
raise NarrativeError("音色克隆尚未完成,请稍后再试", status_code=400)
|
||||
return profile.voice_id, profile.id
|
||||
return tts_voice_id, ""
|
||||
|
||||
|
||||
def _save_tts_job_as_voice_asset(
|
||||
*,
|
||||
job: Any,
|
||||
user_id: str,
|
||||
name: str,
|
||||
project_repository: Any,
|
||||
asset_library_repository: Any,
|
||||
asset_repository: Any,
|
||||
storage_service: SharedStorageService,
|
||||
) -> Asset:
|
||||
"""把已完成 TTS job 的音频转存为配音库 audio asset(同 save-to-library 约定)。"""
|
||||
if not job.output_audio_url and not job.output_audio_key:
|
||||
raise NarrativeError("TTS 合成缺少输出音频", status_code=502)
|
||||
|
||||
library = _find_or_create_voice_library(
|
||||
user_id=user_id,
|
||||
project_repository=project_repository,
|
||||
asset_library_repository=asset_library_repository,
|
||||
)
|
||||
|
||||
audio_format = (job.format or "mp3").strip() or "mp3"
|
||||
content_type = _CONTENT_TYPE_MAP.get(audio_format, "audio/mpeg")
|
||||
storage_key = f"uploads/voice/tts/{job.id}.{audio_format}"
|
||||
|
||||
tmp_path: Path | None = None
|
||||
audio_duration: float | None = None
|
||||
file_size = 0
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp:
|
||||
tmp_path = Path(tmp.name)
|
||||
download_source = job.output_audio_key or job.output_audio_url
|
||||
downloaded = storage_service.download_asset(download_source, tmp_path)
|
||||
if not downloaded or not tmp_path.exists() or tmp_path.stat().st_size == 0:
|
||||
raise NarrativeError("叙事配音音频转存失败", status_code=502)
|
||||
file_size = tmp_path.stat().st_size
|
||||
storage_service.upload_file(tmp_path, storage_key, content_type=content_type)
|
||||
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"quiet",
|
||||
"-print_format",
|
||||
"json",
|
||||
"-show_format",
|
||||
str(tmp_path),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
if proc.returncode == 0:
|
||||
dur = float(json.loads(proc.stdout).get("format", {}).get("duration", 0))
|
||||
if dur > 0:
|
||||
audio_duration = dur
|
||||
except Exception: # noqa: BLE001 - ffprobe 仅用于时长兜底
|
||||
logger.warning("叙事配音 ffprobe 时长提取失败: job_id=%s", job.id, exc_info=True)
|
||||
except NarrativeError:
|
||||
raise
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("叙事配音转存失败: job_id=%s, error=%s", job.id, e, exc_info=True)
|
||||
raise NarrativeError("叙事配音音频转存失败", status_code=502) from e
|
||||
finally:
|
||||
if tmp_path and tmp_path.exists():
|
||||
try:
|
||||
tmp_path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
metadata_: dict[str, object] = {
|
||||
"source": "tts_job",
|
||||
"tts_job_id": job.id,
|
||||
"narrative": True,
|
||||
"format": job.format,
|
||||
"sample_rate": job.sample_rate,
|
||||
"voice_id": job.voice_id,
|
||||
"voice_name": job.voice_model or "",
|
||||
}
|
||||
if job.metadata:
|
||||
for key in ("speed", "language"):
|
||||
if key in job.metadata:
|
||||
metadata_[key] = job.metadata[key]
|
||||
|
||||
asset = Asset.create(
|
||||
project_id=library.project_id,
|
||||
library_id=library.id,
|
||||
name=name or f"叙事配音-{job.id[:8]}",
|
||||
storage_key=storage_key,
|
||||
mime_type=content_type,
|
||||
metadata=metadata_,
|
||||
file_size=file_size,
|
||||
duration=job.duration or audio_duration or None,
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.PENDING,
|
||||
uploaded_by_user_id=user_id,
|
||||
)
|
||||
try:
|
||||
return asset_repository.create(asset)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("叙事配音 asset 落库失败,清理 OSS: %s, error=%s", storage_key, e, exc_info=True)
|
||||
try:
|
||||
storage_service.delete_file(storage_key)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("清理孤儿 OSS 文件失败: %s", storage_key, exc_info=True)
|
||||
raise NarrativeError("叙事配音保存失败,请重试", status_code=502) from e
|
||||
|
||||
|
||||
def prepare_narrative_voice(
|
||||
*,
|
||||
db: Session,
|
||||
user_id: str,
|
||||
script_id: str,
|
||||
tts_voice_id: str,
|
||||
tts_voice_source: str,
|
||||
tts_repository: Any,
|
||||
cosyvoice_service: CosyVoiceService,
|
||||
voice_clone_repository: Any,
|
||||
asset_repository: Any,
|
||||
asset_library_repository: Any,
|
||||
project_repository: Any,
|
||||
storage_service: SharedStorageService,
|
||||
points_enabled: bool = False,
|
||||
is_member: bool = False,
|
||||
member_type: str | None = None,
|
||||
) -> NarrativeContext:
|
||||
"""叙事模式入队前同步合成配音并落为 audio asset。
|
||||
|
||||
Raises:
|
||||
NarrativeError: 文案缺失/归属不符、音色不可用、TTS 失败、转存失败。
|
||||
"""
|
||||
script = db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
|
||||
if script is None:
|
||||
raise NarrativeError("文案不存在或无权使用", status_code=404)
|
||||
content = (script.content or "").strip()
|
||||
if not content:
|
||||
raise NarrativeError("文案内容为空,无法合成配音", status_code=400)
|
||||
|
||||
actual_voice_id, clone_profile_id = _resolve_voice(
|
||||
user_id=user_id,
|
||||
tts_voice_id=tts_voice_id,
|
||||
tts_voice_source=tts_voice_source,
|
||||
voice_clone_repository=voice_clone_repository,
|
||||
)
|
||||
|
||||
# 积分扣点(与 /tts 合成端点同口径),失败时在合成失败分支退费
|
||||
points_svc = PointsService() if points_enabled else None
|
||||
points_deducted = 0
|
||||
if points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(content) / 240))
|
||||
points_deducted = calculate_points_cost(
|
||||
_POINTS_SCENE,
|
||||
is_member=is_member,
|
||||
duration_minutes=est_minutes,
|
||||
member_type=member_type,
|
||||
)
|
||||
deduct_res = points_svc.deduct_points(user_id, points_deducted, _POINTS_SCENE, db)
|
||||
if not deduct_res["success"]:
|
||||
raise NarrativeError(
|
||||
f"积分不足,需要 {points_deducted} 积分,当前余额 {deduct_res['balance']}",
|
||||
status_code=402,
|
||||
)
|
||||
|
||||
use_case = CreateTTSJobUseCase(tts_repository)
|
||||
job = use_case.execute(
|
||||
user_id=user_id,
|
||||
input_text=content,
|
||||
voice_id=actual_voice_id,
|
||||
voice_clone_profile_id=clone_profile_id,
|
||||
metadata={"speed": 1.0, "emotion": "", "language": "zh-CN", "narrative": True, "script_id": script_id},
|
||||
)
|
||||
|
||||
workflow = TTSWorkflowService(repository=tts_repository, cosyvoice_service=cosyvoice_service)
|
||||
try:
|
||||
job = workflow.start_synthesis(job.id)
|
||||
if not job.is_completed:
|
||||
job = workflow.poll_and_process_synthesis(job.id, timeout=_SYNTH_TIMEOUT)
|
||||
except Exception as e: # noqa: BLE001 - 同步合成异常统一转 NarrativeError
|
||||
logger.error("叙事配音 TTS 合成失败: job_id=%s, error=%s", job.id, e, exc_info=True)
|
||||
try:
|
||||
workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("标记叙事 TTS job 失败出错: job_id=%s", job.id, exc_info=True)
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 失败退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError(f"配音合成失败:{e}", status_code=502) from e
|
||||
|
||||
if not job.is_completed:
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 未完成退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError("配音合成未完成,请稍后重试", status_code=504)
|
||||
|
||||
asset = _save_tts_job_as_voice_asset(
|
||||
job=job,
|
||||
user_id=user_id,
|
||||
name=(script.title or "叙事配音")[:60],
|
||||
project_repository=project_repository,
|
||||
asset_library_repository=asset_library_repository,
|
||||
asset_repository=asset_repository,
|
||||
storage_service=storage_service,
|
||||
)
|
||||
|
||||
return NarrativeContext(
|
||||
script=script,
|
||||
voice_asset_id=asset.id,
|
||||
tts_job_id=job.id,
|
||||
audio_duration=float(job.duration or asset.duration or 0.0),
|
||||
)
|
||||
@@ -22,6 +22,11 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyEditPlanClipRepository,
|
||||
SQLAlchemyEditPlanRepository,
|
||||
)
|
||||
from packages.domain.atom_clip_resolver import load_atom_clips_for_assets
|
||||
from packages.domain.atom_clip_selector import (
|
||||
estimate_required_clip_count,
|
||||
select_atom_clips,
|
||||
)
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
from packages.domain.edit_plan import EditPlan
|
||||
from packages.domain.edit_plan_clip import EditPlanClip
|
||||
@@ -52,10 +57,12 @@ class PlanGeneratorService:
|
||||
基于模板 + 素材,自动生成 EditPlan 及 EditPlanClip 列表。
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session, asset_repo=None) -> None:
|
||||
def __init__(self, db: Session, asset_repo=None, atom_clip_repo=None) -> None:
|
||||
self._plan_repo = SQLAlchemyEditPlanRepository(db)
|
||||
self._clip_repo = SQLAlchemyEditPlanClipRepository(db)
|
||||
self._asset_repo = asset_repo
|
||||
# #1970 原子化切片:可选注入;未注入时走旧的整条素材选片路径(向后兼容)
|
||||
self._atom_clip_repo = atom_clip_repo
|
||||
|
||||
# ── 公开接口 ─────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -121,18 +128,34 @@ class PlanGeneratorService:
|
||||
|
||||
# 4. 按 editing_mode 分配素材
|
||||
if asset_ids:
|
||||
# 获取素材时长信息,用于随机起始时间
|
||||
asset_durations = None
|
||||
if self._asset_repo:
|
||||
asset_durations = self._fetch_asset_durations(asset_ids)
|
||||
self._distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
random_selection=random_preview,
|
||||
asset_durations=asset_durations,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
# #1970 原子化切片:素材 clip 从 atom_clips 表选取(未就绪自动内存兜底)。
|
||||
# 预览随机模式保持旧路径(整条素材 + 随机起点),与现有预览契约一致。
|
||||
atom_applied = False
|
||||
if not random_preview and self._atom_clip_repo is not None:
|
||||
try:
|
||||
atom_applied = self._distribute_atom_clips(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("原子片段选片失败,回退整条素材选片", exc_info=True)
|
||||
atom_applied = False
|
||||
|
||||
if not atom_applied:
|
||||
# 获取素材时长信息,用于随机起始时间
|
||||
asset_durations = None
|
||||
if self._asset_repo:
|
||||
asset_durations = self._fetch_asset_durations(asset_ids)
|
||||
self._distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
editing_mode,
|
||||
random_selection=random_preview,
|
||||
asset_durations=asset_durations,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
|
||||
# 5. 持久化所有 clips 并计算总时长
|
||||
created_clips: list[EditPlanClip] = []
|
||||
@@ -259,6 +282,103 @@ class PlanGeneratorService:
|
||||
external_used_segments=external_used_segments,
|
||||
)
|
||||
|
||||
def _distribute_atom_clips(
|
||||
self,
|
||||
clips: list[EditPlanClip],
|
||||
asset_ids: list[str],
|
||||
editing_mode: str,
|
||||
*,
|
||||
user_id: str = "",
|
||||
) -> bool:
|
||||
"""#1970 原子化切片选片(就地修改 clips,未持久化).
|
||||
|
||||
从 ``asset_atom_clips`` 表按原子片段选取;老素材/切片未就绪的素材
|
||||
内存兜底切片。同一原子片段在一次方案中只用一次;跨视频避让走
|
||||
edit_plan_clips.atom_clip_id 最近使用记录。
|
||||
|
||||
Returns:
|
||||
True 表示原子片段选片成功;False 表示无可用片段,调用方应回退
|
||||
到旧的整条素材 distribute_assets。
|
||||
"""
|
||||
# 1. 加载候选原子片段(DB + 兜底)
|
||||
clips_by_asset = load_atom_clips_for_assets(
|
||||
asset_ids,
|
||||
atom_clip_repo=self._atom_clip_repo,
|
||||
asset_repo=self._asset_repo,
|
||||
)
|
||||
if not clips_by_asset:
|
||||
return False
|
||||
|
||||
# 2. 最近使用片段(跨视频原子片段级避让)
|
||||
recently_used: set[str] = set()
|
||||
if user_id and hasattr(self._clip_repo, "list_recent_atom_clip_ids_by_user"):
|
||||
try:
|
||||
recently_used = set(self._clip_repo.list_recent_atom_clip_ids_by_user(user_id, limit=200))
|
||||
except Exception:
|
||||
logger.warning("跨视频原子片段避让查询失败", exc_info=True)
|
||||
|
||||
# 3. 片段需求估算:无配音时按 clips 数量;voice_over 的配音总时长存于
|
||||
# clip.config["voice_duration"],按 平均片段时长≈需要片段数 估算
|
||||
voice_total = 0.0
|
||||
for c in clips:
|
||||
cfg_vd = c.config.get("voice_duration") if c.config else None
|
||||
if cfg_vd:
|
||||
voice_total += float(cfg_vd)
|
||||
avg_clip_target = sum(float(c.duration or 0.0) for c in clips) / max(len(clips), 1)
|
||||
required_count = estimate_required_clip_count(
|
||||
voice_total or sum(float(c.duration or 0.0) for c in clips),
|
||||
avg_clip_target or 3.5,
|
||||
)
|
||||
required_count = max(required_count, len(clips))
|
||||
|
||||
rng = random.Random()
|
||||
|
||||
# 4. 正式生成:先按素材 smart_score 对素材池排序,再展开为片段池
|
||||
# (同素材的片段保持连续,高分素材的片段排在前面优先入选)
|
||||
if self._asset_repo:
|
||||
asset_order = self._sort_assets_by_smart_score(list(clips_by_asset.keys()))
|
||||
ordered: dict[str, list] = {}
|
||||
for aid in asset_order:
|
||||
if aid in clips_by_asset:
|
||||
ordered[aid] = clips_by_asset[aid]
|
||||
clips_by_asset = ordered
|
||||
|
||||
candidates: list = []
|
||||
for asset_clips in clips_by_asset.values():
|
||||
candidates.extend(asset_clips)
|
||||
|
||||
# 5. 逐虚拟片段选片:评分排序,同片段不重复使用
|
||||
used_atom_ids: set[str] = set()
|
||||
asset_usage: dict[str, int] = {}
|
||||
assigned = 0
|
||||
for clip in clips:
|
||||
# 对每个虚拟片段重新评分(usage_count 随选择动态变化)
|
||||
scored = select_atom_clips(
|
||||
candidates,
|
||||
target_duration=float(clip.duration or 0.0),
|
||||
used_atom_clip_ids=used_atom_ids,
|
||||
asset_usage_counts=asset_usage,
|
||||
recently_used_atom_ids=recently_used,
|
||||
required_count=required_count,
|
||||
limit=1,
|
||||
rng=rng,
|
||||
)
|
||||
if not scored:
|
||||
# 候选耗尽(同片段不可重复),交由调用方回退或留白
|
||||
continue
|
||||
picked = scored[0]
|
||||
clip.asset_id = picked.asset_id
|
||||
clip.atom_clip_id = picked.atom_clip_id
|
||||
clip.start_time = round(picked.start_time, 3)
|
||||
clip.duration = round(picked.duration, 3)
|
||||
used_atom_ids.add(picked.atom_clip_id)
|
||||
asset_usage[picked.asset_id] = asset_usage.get(picked.asset_id, 0) + 1
|
||||
assigned += 1
|
||||
|
||||
if assigned == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _fetch_asset_scene_points(self, asset_ids: list[str]) -> dict[str, list[float]]:
|
||||
"""从素材 metadata 读取场景切换点缓存(无缓存的素材不包含在结果中)。"""
|
||||
points_map: dict[str, list[float]] = {}
|
||||
|
||||
@@ -38,7 +38,12 @@ def transcribe_to_text(media_path: str | Path) -> str:
|
||||
ASRTranscriptionError: ASR 调用失败
|
||||
"""
|
||||
# 延迟导入,避免循环依赖和启动时副作用
|
||||
from apps.worker.services.asr_service_factory import get_asr_service
|
||||
try:
|
||||
from apps.worker.services.asr_service_factory import get_asr_service
|
||||
except ImportError as exc:
|
||||
# API 镜像未打包 worker 代码(本地 ASR 依赖 worker 的 asr_service_factory)
|
||||
logger.warning("本地 ASR 不可用(apps.worker 未安装): %s", exc)
|
||||
raise ASRNotConfiguredError("本地 ASR 服务不可用(worker 模块未安装)") from exc
|
||||
|
||||
asr = get_asr_service()
|
||||
if asr is None:
|
||||
|
||||
@@ -51,9 +51,6 @@ class ScriptService:
|
||||
content: str = "",
|
||||
segments: list | None = None,
|
||||
tags: list | None = None,
|
||||
title_text: str = "",
|
||||
title_category: str = "",
|
||||
title_config: dict | None = None,
|
||||
) -> ScriptModel:
|
||||
script = ScriptModel(
|
||||
id=str(uuid.uuid4()),
|
||||
@@ -62,9 +59,6 @@ class ScriptService:
|
||||
content=content,
|
||||
segments=segments if segments is not None else [],
|
||||
tags=tags if tags is not None else [],
|
||||
title_text=title_text or "",
|
||||
title_category=title_category or "",
|
||||
title_config=title_config if title_config is not None else {},
|
||||
)
|
||||
self.db.add(script)
|
||||
self.db.commit()
|
||||
@@ -89,9 +83,6 @@ class ScriptService:
|
||||
content: Optional[str] = None,
|
||||
segments: Optional[list] = None,
|
||||
tags: Optional[list] = None,
|
||||
title_text: Optional[str] = None,
|
||||
title_category: Optional[str] = None,
|
||||
title_config: Optional[dict] = None,
|
||||
) -> ScriptModel:
|
||||
script = self.get_script(script_id, user_id)
|
||||
if title is not None:
|
||||
@@ -102,27 +93,11 @@ class ScriptService:
|
||||
script.segments = segments
|
||||
if tags is not None:
|
||||
script.tags = tags
|
||||
if title_text is not None:
|
||||
script.title_text = title_text
|
||||
if title_category is not None:
|
||||
script.title_category = title_category
|
||||
if title_config is not None:
|
||||
script.title_config = title_config
|
||||
script.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
self.db.refresh(script)
|
||||
return script
|
||||
|
||||
# ── title config ─────────────────────────────────────────────────────
|
||||
|
||||
def get_title_config_for_script(self, script_id: str, user_id: str) -> dict:
|
||||
"""从 script 读取标题配置,返回可直接用于渲染的 title_config dict."""
|
||||
script = self.get_script(script_id, user_id)
|
||||
config = dict(script.title_config or {})
|
||||
if not config.get("text") and script.title_text:
|
||||
config["text"] = script.title_text
|
||||
return config
|
||||
|
||||
# ── delete ────────────────────────────────────────────────────────────
|
||||
|
||||
def delete_script(self, script_id: str, user_id: str) -> bool:
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
"""GPU MuseTalk 异步推理任务 — 将 GPU 推理等待从 HTTP 请求移至 Celery 后台执行.
|
||||
|
||||
优化目标:将 POST /lipsync/jobs 的 API 响应时间从 >200s 降到 <1s。
|
||||
任务流程:
|
||||
1. 加载 LipsyncJob,获取 gpu_task_id
|
||||
2. 调用 GpuLipsyncService.wait_for_result 轮询等待 GPU 完成
|
||||
3. 签名结果 URL(7 天),更新 job 为 completed
|
||||
4. 失败/超时时:尝试 MediaKit 兜底,若仍失败则标记 job 为 failed
|
||||
|
||||
使用 @shared_task 确保被 Worker 侧 celery_app 正确注册。
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from celery import shared_task
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 与 LipsyncService 保持一致
|
||||
_MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
|
||||
|
||||
|
||||
def _get_db_session() -> Session:
|
||||
"""获取 DB session(兼容 API 和 Worker 两种运行时)."""
|
||||
try:
|
||||
from worker_app.db import SessionLocal # type: ignore
|
||||
except ImportError:
|
||||
from app.db import SessionLocal # type: ignore
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
def _sign_media_url(url: str) -> str:
|
||||
"""对自家 OSS URL 签 7 天预签名。"""
|
||||
if not url:
|
||||
return url
|
||||
try:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
public_base = getattr(storage, "public_url", "")
|
||||
if not isinstance(public_base, str) or not public_base:
|
||||
return url
|
||||
own_host = urlparse(public_base).netloc.lower()
|
||||
host = urlparse(url).netloc.lower()
|
||||
if not own_host or host != own_host:
|
||||
return url
|
||||
return storage.get_download_url(url, expires_seconds=_MEDIAKIT_URL_TTL_SECONDS)
|
||||
except Exception:
|
||||
return url
|
||||
|
||||
|
||||
@shared_task(
|
||||
name="lipsync_gpu_process_async",
|
||||
bind=True,
|
||||
max_retries=0,
|
||||
acks_late=True,
|
||||
)
|
||||
def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str) -> None:
|
||||
"""异步处理 GPU MuseTalk 推理。
|
||||
|
||||
Args:
|
||||
job_id: LipsyncJob 的 ID
|
||||
user_id: 用户 ID
|
||||
gpu_task_id: GpuLipsyncTask 的 ID
|
||||
"""
|
||||
db: Session = _get_db_session()
|
||||
try:
|
||||
job = db.query(LipsyncJobModel).filter_by(id=job_id, user_id=user_id).first()
|
||||
if job is None:
|
||||
logger.error("[lipsync_gpu_async] job 不存在: job_id=%s", job_id)
|
||||
return
|
||||
|
||||
# 确保状态为 processing
|
||||
if job.status not in ("processing", "gpu_processing"):
|
||||
logger.warning(
|
||||
"[lipsync_gpu_async] job 状态异常,跳过: job_id=%s status=%s",
|
||||
job_id,
|
||||
job.status,
|
||||
)
|
||||
return
|
||||
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
|
||||
gpu_svc = GpuLipsyncService(db)
|
||||
final_task = gpu_svc.wait_for_result(gpu_task_id)
|
||||
|
||||
if final_task is None:
|
||||
logger.warning(
|
||||
"[lipsync_gpu_async] GPU 超时,回退 MediaKit: job_id=%s gpu_task=%s",
|
||||
job_id,
|
||||
gpu_task_id,
|
||||
)
|
||||
_fallback_to_mediakit(db, job)
|
||||
return
|
||||
|
||||
if final_task.status != "done":
|
||||
logger.warning(
|
||||
"[lipsync_gpu_async] GPU 失败,回退 MediaKit: job_id=%s gpu_task=%s status=%s",
|
||||
job_id,
|
||||
gpu_task_id,
|
||||
final_task.status,
|
||||
)
|
||||
_fallback_to_mediakit(db, job)
|
||||
return
|
||||
|
||||
# 签名结果 URL
|
||||
result_url = final_task.result_url or ""
|
||||
try:
|
||||
storage = get_shared_storage_service()
|
||||
signed = storage.get_download_url(result_url, expires_seconds=_MEDIAKIT_URL_TTL_SECONDS)
|
||||
if signed:
|
||||
result_url = signed
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[lipsync_gpu_async] 签名失败,用原 URL: job_id=%s err=%s",
|
||||
job_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
job.status = "completed"
|
||||
job.output_video_url = result_url
|
||||
job.output_duration = final_task.result_duration or 0.0
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"[lipsync_gpu_async] GPU 完成: job_id=%s duration=%.2f",
|
||||
job_id,
|
||||
job.output_duration,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("[lipsync_gpu_async] 异常: job_id=%s err=%s", job_id, exc)
|
||||
try:
|
||||
job = db.query(LipsyncJobModel).filter_by(id=job_id).first()
|
||||
if job:
|
||||
job.status = "failed"
|
||||
job.error_message = f"GPU 异步处理异常: {exc}"
|
||||
job.error_code = "GpuAsyncError"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _fallback_to_mediakit(db: Session, job: LipsyncJobModel) -> None:
|
||||
"""GPU 失败时回退到 MediaKit 云端渲染。"""
|
||||
try:
|
||||
from app.services.mediakit_client import MediaKitError, get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
video_url = _sign_media_url(job.video_url)
|
||||
audio_url = _sign_media_url(job.audio_url)
|
||||
|
||||
result = client.submit_lipsync(
|
||||
video_url=video_url,
|
||||
audio_url=audio_url,
|
||||
enable_video_loop=job.enable_video_loop,
|
||||
client_token=job.id,
|
||||
)
|
||||
job.mediakit_task_id = result["task_id"]
|
||||
job.status = "submitted"
|
||||
job.submitted_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"[lipsync_gpu_async] 已回退 MediaKit: job_id=%s task_id=%s",
|
||||
job.id,
|
||||
result["task_id"],
|
||||
)
|
||||
except MediaKitError as exc:
|
||||
job.status = "failed"
|
||||
job.error_message = str(exc)
|
||||
job.error_code = exc.code
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.error("[lipsync_gpu_async] MediaKit 也失败: job_id=%s err=%s", job.id, exc)
|
||||
except Exception as exc:
|
||||
job.status = "failed"
|
||||
job.error_message = f"GPU+MediaKit 均失败: {exc}"
|
||||
job.error_code = "FallbackFailed"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.error("[lipsync_gpu_async] 兜底异常: job_id=%s err=%s", job.id, exc)
|
||||
Executable
+117
@@ -0,0 +1,117 @@
|
||||
import { expect, test, type APIRequestContext, type Page } from "@playwright/test"
|
||||
|
||||
const PASSWORD = "SmokePass123!"
|
||||
const apiBase = process.env.E2E_API_BASE || "/api/v1"
|
||||
const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".length) : ""
|
||||
|
||||
async function routeBrowserApiToTestApi(page: Page) {
|
||||
if (!apiOrigin) return
|
||||
await page.route("**/api/v1/**", async (route) => {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
await route.fulfill({ response })
|
||||
})
|
||||
}
|
||||
|
||||
async function loginWithRetry(request: APIRequestContext, email: string, password: string) {
|
||||
for (let i = 0; i <= 2; i++) {
|
||||
const r = await request.post(`${apiBase}/auth/login`, { data: { email, password } })
|
||||
if (r.status() !== 429) {
|
||||
expect(r.ok(), `login: ${await r.text()}`).toBeTruthy()
|
||||
return (await r.json()).access_token as string
|
||||
}
|
||||
console.log(`[douyin] 429 retry ${i + 1}/2`)
|
||||
await new Promise((res) => setTimeout(res, 65000))
|
||||
}
|
||||
throw new Error("Login retries exhausted")
|
||||
}
|
||||
|
||||
/**
|
||||
* #1972 抖音文案提取冒烟
|
||||
*
|
||||
* 路径:文案库页面 → 点「🎬 从抖音提取」→ 粘贴分享文案 → 点「开始提取」
|
||||
* → mock /api/v1/scripts/extract-from-douyin 返回稳定文案 → 断言「新建文案」弹窗中预填了非空文案
|
||||
*/
|
||||
test.describe("Douyin Script Extraction (#1972)", () => {
|
||||
test("extract flow: open modal, paste link, text prefilled in create modal", async ({
|
||||
page,
|
||||
request,
|
||||
}) => {
|
||||
test.setTimeout(180_000)
|
||||
await page.setViewportSize({ width: 1440, height: 900 })
|
||||
|
||||
const suffix = Math.random().toString(36).slice(2, 8)
|
||||
const email = `e2e-douyin-${suffix}@example.com`
|
||||
await request.post(`${apiBase}/auth/register`, {
|
||||
data: { email, password: PASSWORD, username: `e2e_dy_${suffix}` },
|
||||
})
|
||||
const token = await loginWithRetry(request, email, PASSWORD)
|
||||
const authHeader = { Authorization: `Bearer ${token}` }
|
||||
|
||||
const proj = await request.post(`${apiBase}/projects`, {
|
||||
headers: authHeader,
|
||||
data: { name: `Smoke Douyin ${suffix}` },
|
||||
})
|
||||
const projectId = (await proj.json()).id ?? (await proj.json()).project_id
|
||||
await request.post(`${apiBase}/asset-libraries`, {
|
||||
headers: authHeader,
|
||||
data: { project_id: projectId, name: "Smoke", kind: "video" },
|
||||
})
|
||||
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
}, token)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
|
||||
// Mock 抖音提取接口返回稳定文案
|
||||
const extractedText = "大家好,今天给大家推荐一款超好用的产品,性价比非常高,快来看看吧!"
|
||||
await page.route("**/api/v1/scripts/extract-from-douyin", (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ text: extractedText, duration_seconds: 15 }),
|
||||
}),
|
||||
)
|
||||
// 文案列表空态
|
||||
await page.route(
|
||||
(url) => url.pathname.endsWith("/scripts") && !url.pathname.includes("extract-from-douyin"),
|
||||
(route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ items: [], total: 0, page: 1, page_size: 20 }),
|
||||
}),
|
||||
)
|
||||
|
||||
await page.goto("/app/scripts")
|
||||
// 文案库页面加载
|
||||
await expect(page.getByText(/文案库|文案/).first()).toBeVisible({ timeout: 30000 })
|
||||
|
||||
// 点「🎬 从抖音提取」按钮
|
||||
await page.getByRole("button", { name: /从抖音提取/ }).click()
|
||||
await expect(page.getByText("从抖音视频提取文案")).toBeVisible({ timeout: 5000 })
|
||||
|
||||
// 在 TextArea 粘贴"抖音分享文案"
|
||||
const textarea = page.locator(".ant-modal textarea").first()
|
||||
await expect(textarea).toBeVisible()
|
||||
await textarea.fill("8.88 复制打开抖音,看看【推荐视频】https://v.douyin.com/abcDEF/")
|
||||
|
||||
// 点「开始提取」
|
||||
await page.getByRole("button", { name: "开始提取" }).click()
|
||||
await expect(page.getByText(/提取中/)).toBeVisible({ timeout: 3000 })
|
||||
|
||||
// 等待抖音弹窗关闭,「新建文案」弹窗打开并预填提取文案
|
||||
await expect(page.getByText("从抖音视频提取文案")).not.toBeVisible({ timeout: 15000 })
|
||||
await expect(page.getByText("新建文案")).toBeVisible({ timeout: 5000 })
|
||||
const createTextarea = page.locator(".ant-modal textarea").first()
|
||||
await expect(createTextarea).toBeVisible()
|
||||
await expect(createTextarea).toHaveValue(new RegExp(extractedText.slice(0, 10)))
|
||||
console.log("[douyin] Extraction flow completed ✓, text length:", extractedText.length)
|
||||
})
|
||||
})
|
||||
@@ -1,4 +1,4 @@
|
||||
import { expect, test, type APIRequestContext } from "@playwright/test"
|
||||
import { expect, test, type APIRequestContext, type Page } from "@playwright/test"
|
||||
import * as fs from "node:fs"
|
||||
import * as path from "node:path"
|
||||
import { fileURLToPath } from "node:url"
|
||||
@@ -8,7 +8,8 @@ const PASSWORD = "SmokePass123!"
|
||||
const apiBase = process.env.E2E_API_BASE || "/api/v1"
|
||||
const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".length) : ""
|
||||
|
||||
const routeBrowserApiToTestApi = async (page: import("@playwright/test").Page) => {
|
||||
/** 将浏览器侧 /api/v1 请求路由到 Playwright request 源(支持跨域) */
|
||||
async function routeBrowserApiToTestApi(page: Page) {
|
||||
if (!apiOrigin) return
|
||||
await page.route("**/api/v1/**", async (route) => {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
@@ -24,309 +25,358 @@ async function loginWithRetry(
|
||||
email: string,
|
||||
password: string,
|
||||
maxRetries = 2,
|
||||
) {
|
||||
): Promise<string> {
|
||||
for (let i = 0; i <= maxRetries; i++) {
|
||||
const response = await request.post(`${apiBase}/auth/login`, {
|
||||
data: { email, password },
|
||||
})
|
||||
if (response.status() !== 429) return response
|
||||
console.log(`[login] 触发限流,等待 65s 后重试 (${i + 1}/${maxRetries})`)
|
||||
const resp = await request.post(`${apiBase}/auth/login`, { data: { email, password } })
|
||||
if (resp.status() !== 429) {
|
||||
expect(resp.ok(), `Login should succeed: ${await resp.text()}`).toBeTruthy()
|
||||
const data = await resp.json()
|
||||
return data.access_token
|
||||
}
|
||||
console.log(`[login] 429 rate limited, retry ${i + 1}/${maxRetries} after 65s`)
|
||||
await new Promise((r) => setTimeout(r, 65000))
|
||||
}
|
||||
return request.post(`${apiBase}/auth/login`, {
|
||||
data: { email, password },
|
||||
throw new Error("Login failed after retries")
|
||||
}
|
||||
|
||||
/**
|
||||
* 注册新用户 + 建项目/视频库/上传 sample.mp4,等素材 ready。返回 { token, projectId, libraryId, assetId }。
|
||||
*/
|
||||
async function setupFreshUser(
|
||||
request: APIRequestContext,
|
||||
label: string,
|
||||
): Promise<{ token: string; libraryId: string; assetId: string; suffix: string }> {
|
||||
const suffix = Math.random().toString(36).slice(2, 8)
|
||||
const email = `e2e-${label}-${suffix}@example.com`
|
||||
await request.post(`${apiBase}/auth/register`, {
|
||||
data: { email, password: PASSWORD, username: `e2e_${label}_${suffix}` },
|
||||
})
|
||||
const token = await loginWithRetry(request, email, PASSWORD)
|
||||
const auth = { Authorization: `Bearer ${token}` }
|
||||
|
||||
const proj = await request.post(`${apiBase}/projects`, {
|
||||
headers: auth,
|
||||
data: { name: `Smoke ${label} ${suffix}` },
|
||||
})
|
||||
expect(proj.ok(), `create project: ${await proj.text()}`).toBeTruthy()
|
||||
const projectId = (await proj.json()).id ?? (await proj.json()).project_id
|
||||
|
||||
const lib = await request.post(`${apiBase}/asset-libraries`, {
|
||||
headers: auth,
|
||||
data: { project_id: projectId, name: "Smoke", kind: "video" },
|
||||
})
|
||||
expect(lib.ok(), `create library: ${await lib.text()}`).toBeTruthy()
|
||||
const libraryId = (await lib.json()).id
|
||||
|
||||
const samplePath = path.join(__dirname, "fixtures", "sample.mp4")
|
||||
const sampleBuf = fs.readFileSync(samplePath)
|
||||
const up = await request.post(`${apiBase}/upload`, {
|
||||
headers: auth,
|
||||
multipart: {
|
||||
project_id: projectId,
|
||||
library_id: libraryId,
|
||||
file: {
|
||||
name: "sample.mp4",
|
||||
mimeType: "video/mp4",
|
||||
buffer: sampleBuf,
|
||||
},
|
||||
},
|
||||
})
|
||||
expect(up.ok(), `upload sample: ${await up.text()}`).toBeTruthy()
|
||||
const assetId = (await up.json()).asset_id
|
||||
await expect
|
||||
.poll(
|
||||
async () => {
|
||||
const r = await request.get(`${apiBase}/assets/${assetId}`, { headers: auth })
|
||||
return r.ok() ? (await r.json()).status : "pending"
|
||||
},
|
||||
{ timeout: 90_000, intervals: [3000, 3000, 5000] },
|
||||
)
|
||||
.toBe("ready")
|
||||
return { token, libraryId, assetId, suffix }
|
||||
}
|
||||
|
||||
type ProjectResponse = { id: string }
|
||||
type LibraryResponse = { id: string }
|
||||
type AssetListResponse = {
|
||||
items: Array<{
|
||||
id: string
|
||||
name: string
|
||||
status: string
|
||||
}>
|
||||
}
|
||||
|
||||
test.describe("Core generation flow", () => {
|
||||
test.describe.configure({ timeout: 600_000 })
|
||||
|
||||
test("walks through wizard with count modal and starts generation", async ({ page, request }) => {
|
||||
/**
|
||||
* #1970 智能剪辑核心冒烟(新 5 步向导)
|
||||
*
|
||||
* 新流程:选择模式 → 选择素材 → 选择标题 → 确认生成 → 选择封面
|
||||
*
|
||||
* 两条路径:
|
||||
* 1) 随机混剪(默认)→ Step1 下一步 → 配音选择弹窗 → Step2 选素材 → 数量弹窗
|
||||
* → Step3 标题 → Step4 确认生成 → 断言任务创建
|
||||
* 2) 叙事剪辑 → Step1 切模式 → 下一步 → 文案选择弹窗 → TTS 弹窗选音色(mock 合成)
|
||||
* → Step2 AI 提示卡可见 + 选素材 → 数量弹窗 → Step3 标题 → Step4 确认生成
|
||||
* → 断言任务创建
|
||||
*/
|
||||
test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
test("random mode: 5-step wizard creates generation task", async ({ page, request }) => {
|
||||
test.setTimeout(600_000)
|
||||
await page.setViewportSize({ width: 1440, height: 1000 })
|
||||
const { token, suffix } = await setupFreshUser(request, "random")
|
||||
const authHeader = { Authorization: `Bearer ${token}` }
|
||||
|
||||
await routeBrowserApiToTestApi(page)
|
||||
const suffix = Date.now().toString(36)
|
||||
const email = `e2e-gen-${suffix}@example.com`
|
||||
const username = `e2e_gen_${suffix}`
|
||||
const libraryName = `E2E Gen Lib ${suffix}`
|
||||
// 确保默认模板存在(智能剪辑页依赖模板)
|
||||
const tmpls = await request.get(`${apiBase}/templates`, { headers: authHeader })
|
||||
const tmplsJson = await tmpls.json()
|
||||
const templates = Array.isArray(tmplsJson)
|
||||
? tmplsJson
|
||||
: Array.isArray(tmplsJson.items)
|
||||
? tmplsJson.items
|
||||
: []
|
||||
expect(templates.length).toBeGreaterThan(0)
|
||||
|
||||
// Register
|
||||
const register = await request.post(`${apiBase}/auth/register`, {
|
||||
data: { email, username, password: PASSWORD, display_name: username },
|
||||
})
|
||||
expect(register.status()).toBe(201)
|
||||
const registerData = (await register.json()) as { user_id: string }
|
||||
|
||||
// Login
|
||||
const login = await loginWithRetry(request, email, PASSWORD)
|
||||
expect(login.status()).toBe(200)
|
||||
const loginData = (await login.json()) as { access_token: string }
|
||||
const headers = { Authorization: `Bearer ${loginData.access_token}` }
|
||||
|
||||
// Create project
|
||||
const project = await request.post(`${apiBase}/projects`, {
|
||||
headers,
|
||||
data: { name: `E2E Gen Proj ${suffix}` },
|
||||
})
|
||||
expect(project.status()).toBe(200)
|
||||
const projectData = (await project.json()) as ProjectResponse
|
||||
|
||||
// Create asset library
|
||||
const library = await request.post(`${apiBase}/asset-libraries`, {
|
||||
headers,
|
||||
data: { project_id: projectData.id, name: libraryName, kind: "video" },
|
||||
})
|
||||
expect(library.status()).toBe(200)
|
||||
const libraryData = (await library.json()) as LibraryResponse
|
||||
|
||||
// Upload source video
|
||||
const sourceFileName = "e2e-gen-source.mp4"
|
||||
const sampleVideoPath = path.join(__dirname, "fixtures", "sample.mp4")
|
||||
const sampleVideoBuffer = fs.readFileSync(sampleVideoPath)
|
||||
const upload = await request.post(`${apiBase}/upload`, {
|
||||
headers,
|
||||
multipart: {
|
||||
project_id: projectData.id,
|
||||
library_id: libraryData.id,
|
||||
file: {
|
||||
name: sourceFileName,
|
||||
mimeType: "video/mp4",
|
||||
buffer: sampleVideoBuffer,
|
||||
},
|
||||
},
|
||||
})
|
||||
expect(upload.status()).toBe(200)
|
||||
|
||||
// Wait for asset to be ready
|
||||
await expect
|
||||
.poll(
|
||||
async () => {
|
||||
const assets = await request.get(`${apiBase}/assets`, {
|
||||
headers,
|
||||
params: { library_id: libraryData.id },
|
||||
})
|
||||
if (!assets.ok()) return `http_${assets.status()}`
|
||||
const data = (await assets.json()) as AssetListResponse
|
||||
const asset = data.items.find((a) => a.name === sourceFileName)
|
||||
if (!asset) return "missing"
|
||||
return asset.status
|
||||
},
|
||||
{ timeout: 30_000, intervals: [1_000, 2_000, 3_000] },
|
||||
// 注入登录态 + 路由 API
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
.toBe("ready")
|
||||
}, token)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
|
||||
// #1926 P0 fix: POST /templates CRUD endpoint removed; GET /templates
|
||||
// now auto-creates a default template for new users. Use the first one.
|
||||
const templatesResp = await request.get(`${apiBase}/templates`, { headers })
|
||||
expect(templatesResp.status(), await templatesResp.text()).toBe(200)
|
||||
const templatesData = (await templatesResp.json()) as {
|
||||
items: Array<{ id: string }>
|
||||
}
|
||||
expect(Array.isArray(templatesData.items)).toBe(true)
|
||||
expect(templatesData.items.length).toBeGreaterThan(0)
|
||||
const templateId = templatesData.items[0].id
|
||||
expect(templateId).toBeTruthy()
|
||||
|
||||
// Set auth in localStorage
|
||||
await page.addInitScript(
|
||||
({ token, user }) => {
|
||||
localStorage.setItem("access_token", token)
|
||||
localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({
|
||||
state: { user, isAuthenticated: true },
|
||||
version: 0,
|
||||
// ── 提前 mock 配音列表(VoiceSelectModal 查询 /assets?kind=voice) ──
|
||||
await page.route(
|
||||
(url) => url.pathname.endsWith("/assets") && url.searchParams.get("kind") === "voice",
|
||||
(route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
items: [
|
||||
{
|
||||
id: `asset-voice-${suffix}`,
|
||||
name: "测试配音.mp3",
|
||||
file_url: "data:audio/mpeg;base64,",
|
||||
duration: 10,
|
||||
file_size: 1024,
|
||||
kind: "voice",
|
||||
status: "ready",
|
||||
},
|
||||
],
|
||||
total: 1,
|
||||
}),
|
||||
)
|
||||
},
|
||||
{
|
||||
token: loginData.access_token,
|
||||
user: {
|
||||
id: registerData.user_id,
|
||||
user_id: registerData.user_id,
|
||||
email,
|
||||
username,
|
||||
display_name: username,
|
||||
is_email_verified: true,
|
||||
email_verified: true,
|
||||
},
|
||||
},
|
||||
}),
|
||||
)
|
||||
|
||||
// Navigate to generate page
|
||||
await page.goto("/app/generate")
|
||||
await expect(page.getByRole("heading", { name: "智能剪辑" })).toBeVisible({
|
||||
timeout: 20_000,
|
||||
timeout: 30000,
|
||||
})
|
||||
|
||||
// 5步向导:素材→数量弹窗→配音→标题→确认生成→封面(#1911 删除选模板步骤,后端自动使用默认模板;
|
||||
// #1677 批量生成在选完素材后弹「要生成几个视频?」数量弹窗,默认1,回车确认)
|
||||
// Step 1: select material (card grid UI)
|
||||
await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible()
|
||||
const librarySelect = page.locator("select").first()
|
||||
await librarySelect.selectOption({ label: libraryName })
|
||||
// 新 UI: 素材以 9:16 竖屏卡片展示,点击卡片选中
|
||||
// 注意:卡片中心是播放按钮(stopPropagation 会阻止选中),所以点击左上角避开
|
||||
const materialCard = page.getByTestId("material-card").filter({ hasText: sourceFileName })
|
||||
await expect(materialCard).toBeVisible({ timeout: 10_000 })
|
||||
await materialCard.click({ position: { x: 15, y: 15 } })
|
||||
// 验证选中:卡片应出现勾选标记(用 testid 定位,避免 ✓ 字符文本匹配不稳定)
|
||||
await expect(materialCard.getByTestId("material-card-check")).toBeVisible({ timeout: 5_000 })
|
||||
await page.getByRole("button", { name: "下一步" }).click()
|
||||
// ── Step 1:默认随机混剪选中,点下一步 ──────────────────────────
|
||||
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
|
||||
await expect(page.getByText("随机混剪")).toBeVisible()
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// #1677 数量弹窗:默认值1,点击「生成 1 个视频」确认(新用户单视频冒烟路径)
|
||||
await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({
|
||||
timeout: 5_000,
|
||||
})
|
||||
// ── 配音选择弹窗:选第一个配音 → 确认 ─────────────────────────
|
||||
await expect(page.getByText("🎙️ 选择配音")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByText("测试配音.mp3").first().click()
|
||||
await page.getByRole("button", { name: "确认选择" }).click()
|
||||
await expect(page.getByText("🎙️ 选择配音")).not.toBeVisible()
|
||||
|
||||
// ── Step 2:选择素材 ──────────────────────────────────────────
|
||||
await expect(page.getByText("选择素材", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
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 2: voice(新注册用户无配音素材时展示空状态 h3「🎙️ 选择配音」,仍可点「下一步」跳过)
|
||||
await expect(page.getByRole("heading", { name: /选择配音/ })).toBeVisible({ timeout: 15000 })
|
||||
await page.getByRole("button", { name: "下一步" }).click()
|
||||
|
||||
// Step 3: title(新顺序:标题在预览之前)
|
||||
await expect(page.getByRole("heading", { name: /选择标题/ })).toBeVisible({ timeout: 15000 })
|
||||
// 等待组件完全渲染
|
||||
await page.waitForTimeout(2000)
|
||||
|
||||
// Antd AutoComplete 的 placeholder 渲染在 span 上,input 无 placeholder 属性
|
||||
// 使用 Antd AutoComplete 特有的 class 定位输入框
|
||||
const titleInput = page.locator(".ant-select-auto-complete input")
|
||||
// ── Step 3:填写标题 ──────────────────────────────────────────
|
||||
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
const titleInput = page.getByPlaceholder("输入或从标题库选择")
|
||||
await expect(titleInput).toBeVisible({ timeout: 5000 })
|
||||
await titleInput.fill(`测试随机剪辑 ${suffix}`)
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
const titleText = `E2E Test ${suffix}`
|
||||
await titleInput.fill(titleText)
|
||||
// ── Step 4:确认生成 ──────────────────────────────────────────
|
||||
await expect(page.getByText("📋 生成配置")).toBeVisible({ timeout: 10000 })
|
||||
await expect(page.getByText("随机混剪")).toBeVisible()
|
||||
const confirmBtn = page.getByRole("button", { name: /确认生成视频/ })
|
||||
await expect(confirmBtn).toBeEnabled({ timeout: 5000 })
|
||||
|
||||
// 步骤3(标题页)底部操作栏按钮是「下一步 →」,点击后进入步骤4
|
||||
// 步骤4底部才是「✨ 确认生成视频」按钮
|
||||
const nextBtn = page.locator(".xx-step-actions .xx-btn-primary").filter({ hasText: "下一步" })
|
||||
await expect(nextBtn).toBeVisible({ timeout: 15_000 })
|
||||
await nextBtn.click()
|
||||
const createTask = page.waitForResponse(
|
||||
(r) => r.url().includes("/generation/tasks") && r.request().method() === "POST",
|
||||
{ timeout: 30000 },
|
||||
)
|
||||
await confirmBtn.click()
|
||||
const taskResp = await createTask
|
||||
expect(taskResp.ok(), `Create task: ${await taskResp.text()}`).toBeTruthy()
|
||||
const taskId = (await taskResp.json()).id ?? (await taskResp.json()).task_id
|
||||
console.log("[random] Generation task created:", taskId)
|
||||
await expect(page.getByText(/正在生成|提交/)).toBeVisible({ timeout: 15000 })
|
||||
console.log("[random] Wizard flow completed ✓")
|
||||
})
|
||||
|
||||
// Step 4:「确认生成」页面——此处底部是「✨ 确认生成视频」按钮
|
||||
// 注意:Step4 主内容区是实时预览画布,没有 h3 「🎬 确认生成」标题,标题由顶部步骤条展示
|
||||
// 等待前端实时预览就绪:未就绪时右侧 FrontendPreviewPlayer 显示「准备预览素材...」占位,
|
||||
// 就绪(previewReady:素材已解析 + 模板已选中)后占位消失;否则按钮会被校验拦截弹 warning
|
||||
await page
|
||||
.getByText("准备预览素材")
|
||||
.waitFor({ state: "detached", timeout: 30_000 })
|
||||
.catch(() => {})
|
||||
test("narrative mode: select script + mock TTS, create generation task", async ({
|
||||
page,
|
||||
request,
|
||||
}) => {
|
||||
test.setTimeout(600_000)
|
||||
await page.setViewportSize({ width: 1440, height: 1000 })
|
||||
const { token, suffix } = await setupFreshUser(request, "narrative")
|
||||
|
||||
// 定位底部操作栏的「✨ 确认生成视频」按钮
|
||||
// 使用底部操作栏 xx-step-actions 作用域,避免命中其他 primary 按钮
|
||||
const confirmBtn = page
|
||||
.locator(".xx-step-actions .xx-btn-primary")
|
||||
.filter({ hasText: "确认生成" })
|
||||
await expect(confirmBtn).toBeVisible({ timeout: 30_000 })
|
||||
await expect(confirmBtn).toBeEnabled({ timeout: 30_000 })
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
}, token)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
|
||||
// Wait for generation API to be called — 先挂监听再点击,避免竞态
|
||||
const generatePromise = page.waitForResponse(
|
||||
(response) => {
|
||||
const url = response.url()
|
||||
const path = new URL(url).pathname
|
||||
return response.request().method() === "POST" && path.endsWith("/generation/tasks")
|
||||
},
|
||||
{ timeout: 60_000 },
|
||||
// ── Mock 文案列表、音色、TTS 合成(避免真实合成) ──────────────
|
||||
const mockScriptId = `script-mock-${suffix}`
|
||||
const mockVoiceId = `preset-voice-${suffix}`
|
||||
const mockJobId = `tts-job-${suffix}`
|
||||
|
||||
// 文案列表(ScriptSelectModal 查询 /scripts)
|
||||
await page.route("**/api/v1/scripts**", (route) => {
|
||||
const url = new URL(route.request().url())
|
||||
if (url.pathname.includes("/extract-from-douyin")) {
|
||||
route.continue()
|
||||
return
|
||||
}
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
items: [
|
||||
{
|
||||
id: mockScriptId,
|
||||
title: "测试带货文案",
|
||||
content: "这是一段测试用的带货文案内容,用于 E2E 冒烟测试。",
|
||||
tags: ["带货"],
|
||||
title_category: "daihuo",
|
||||
created_at: new Date().toISOString(),
|
||||
updated_at: new Date().toISOString(),
|
||||
},
|
||||
],
|
||||
total: 1,
|
||||
page: 1,
|
||||
page_size: 200,
|
||||
}),
|
||||
})
|
||||
})
|
||||
|
||||
// 预设音色(TtsVoiceModal 查询 GET /voices/presets)
|
||||
await page.route("**/api/v1/voices/presets**", (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
items: [
|
||||
{
|
||||
voice_id: mockVoiceId,
|
||||
name: "晓晓(女声)",
|
||||
description: "温柔女声",
|
||||
gender: "female",
|
||||
language: "zh-CN",
|
||||
preview_url: null,
|
||||
tags: ["温柔"],
|
||||
},
|
||||
],
|
||||
total: 1,
|
||||
}),
|
||||
}),
|
||||
)
|
||||
|
||||
await confirmBtn.click()
|
||||
// 克隆音色:空列表
|
||||
await page.route(
|
||||
(url) => url.pathname.endsWith("/voice-clones"),
|
||||
(route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ items: [] }),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify generation was triggered
|
||||
const genResp = await generatePromise
|
||||
if (!genResp.ok()) {
|
||||
const body = await genResp.text()
|
||||
console.error(
|
||||
`[E2E DEBUG] 触发生成接口失败: status=${genResp.status()} url=${genResp.url()} body=${body.slice(0, 500)}`,
|
||||
)
|
||||
}
|
||||
// Generate API may return 400 in test env if template has no ready segments
|
||||
// That is OK for a wizard flow smoke test
|
||||
if (genResp.ok()) {
|
||||
const genData = (await genResp.json()) as {
|
||||
items: Array<{ id: string; status: string }>
|
||||
total: number
|
||||
}
|
||||
expect(genData.items.length).toBeGreaterThan(0)
|
||||
expect(genData.items[0].id).toBeTruthy()
|
||||
// TTS 合成:直接返回 completed 任务
|
||||
await page.route("**/api/v1/tts/synthesize", (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ job_id: mockJobId, status: "queued" }),
|
||||
}),
|
||||
)
|
||||
await page.route(`**/api/v1/tts/jobs/${mockJobId}/status`, (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
job_id: mockJobId,
|
||||
status: "completed",
|
||||
progress: 100,
|
||||
audio_url: "data:audio/mpeg;base64,",
|
||||
duration: 5,
|
||||
}),
|
||||
}),
|
||||
)
|
||||
await page.route(`**/api/v1/tts/jobs/${mockJobId}/save-to-library`, (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ id: `tts-asset-${suffix}`, name: "AI合成配音" }),
|
||||
}),
|
||||
)
|
||||
|
||||
// 单视频(N=1):点击「确认生成视频」后跳步骤 5「确认生成」进度页,展示进度卡
|
||||
// 注意:进度页底部按钮变为 disabled 的「⏳ 视频渲染中…」
|
||||
await expect(page.getByText("视频渲染中")).toBeVisible({ timeout: 30_000 })
|
||||
|
||||
// 等待渲染完成:单视频成片播放器渲染(带「⬇️ 下载」按钮),最长等待 3 分钟
|
||||
// 注意:message.success「视频生成完成」toast 3秒后自动消失,不能作为稳定断言点
|
||||
await expect(page.getByRole("button", { name: "⬇️ 下载" })).toBeVisible({ timeout: 420_000 })
|
||||
|
||||
// #1954 修复:生成完成后步骤4底部应显示「下一步:选择封面」按钮
|
||||
// 等待底部主按钮从「⏳/确认生成」切换为「下一步:选择封面」
|
||||
const nextCoverBtn = page
|
||||
.locator(".xx-step-actions > .xx-btn-primary")
|
||||
.filter({ hasText: "选择封面" })
|
||||
await expect(nextCoverBtn).toBeVisible({ timeout: 15_000 })
|
||||
await nextCoverBtn.click()
|
||||
|
||||
// 断言进入步骤5封面页:主内容出现「选择封面」标题
|
||||
await expect(page.getByText("🖼️ 选择封面")).toBeVisible({ timeout: 10_000 })
|
||||
// 底部操作栏主按钮应消失(封面是最后一步,只剩「← 上一步」)
|
||||
await expect(page.locator(".xx-step-actions > .xx-btn-primary")).toHaveCount(0)
|
||||
} else {
|
||||
console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`)
|
||||
// 创建失败时停留在标题页并展示错误提示
|
||||
await page
|
||||
.getByText(/生成失败|重新生成/)
|
||||
.isVisible({ timeout: 15_000 })
|
||||
.catch(() => false)
|
||||
}
|
||||
|
||||
// Verify product library page loads (smoke: just verify page renders)
|
||||
await page.goto("/app/products")
|
||||
await expect(page).toHaveURL(/\/app\/products/)
|
||||
// Verify page container exists = page rendered correctly
|
||||
// (works in all states: loading/error/success - more reliable than checking search input)
|
||||
await expect(page.locator(".xx-products-page")).toBeVisible({
|
||||
timeout: 15_000,
|
||||
await page.goto("/app/generate")
|
||||
await expect(page.getByRole("heading", { name: "智能剪辑" })).toBeVisible({
|
||||
timeout: 30000,
|
||||
})
|
||||
|
||||
// 清理所有路由,避免页面关闭时飞地API请求导致测试报错
|
||||
await page.unrouteAll({ behavior: "ignoreErrors" })
|
||||
})
|
||||
// ── Step 1:切到叙事剪辑 → 下一步 ────────────────────────────
|
||||
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
|
||||
await page.getByText("叙事剪辑").click()
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
test("generation task API creates and lists tasks", async ({ request }) => {
|
||||
const suffix = Date.now().toString(36)
|
||||
const email = `e2e-gen-api-${suffix}@example.com`
|
||||
const username = `e2e_gen_api_${suffix}`
|
||||
// ── 文案选择弹窗:选第一条 → 确认 ─────────────────────────────
|
||||
await expect(page.getByText("📝 选择文案")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByText("测试带货文案").first().click()
|
||||
await page.getByRole("button", { name: "确认选择" }).click()
|
||||
await expect(page.getByText("📝 选择文案")).not.toBeVisible()
|
||||
|
||||
const register = await request.post(`${apiBase}/auth/register`, {
|
||||
data: { email, username, password: PASSWORD, display_name: username },
|
||||
})
|
||||
expect(register.status()).toBe(201)
|
||||
// ── TTS 音色弹窗:选系统音色 → 合成 ─────────────────────────
|
||||
await expect(page.getByText("🎙️ 合成配音")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByText("晓晓(女声)").first().click()
|
||||
await page.getByRole("button", { name: "🎧 合成配音" }).click()
|
||||
await expect(page.getByText("🎙️ 合成配音")).not.toBeVisible({ timeout: 30000 })
|
||||
|
||||
const login = await loginWithRetry(request, email, PASSWORD)
|
||||
expect(login.status()).toBe(200)
|
||||
const loginData = (await login.json()) as { access_token: string }
|
||||
const headers = { Authorization: `Bearer ${loginData.access_token}` }
|
||||
// ── Step 2:AI 匹配提示卡可见 + 选素材 ────────────────────────
|
||||
await expect(page.getByText("选择素材", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
await expect(page.getByText(/AI智能匹配/)).toBeVisible()
|
||||
await page.getByTestId("material-card").first().click()
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
const project = await request.post(`${apiBase}/projects`, {
|
||||
headers,
|
||||
data: { name: `E2E API Proj ${suffix}` },
|
||||
})
|
||||
expect(project.status()).toBe(200)
|
||||
// ── 数量弹窗 ─────────────────────────────────────────────────
|
||||
await expect(page.getByText("要生成几个视频?")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByRole("button", { name: "生成 1 个视频" }).click()
|
||||
|
||||
// List generation tasks via task center API
|
||||
const tasks = await request.get(`${apiBase}/tasks`, { headers })
|
||||
expect(tasks.status()).toBe(200)
|
||||
const tasksData = await tasks.json()
|
||||
expect(Array.isArray(tasksData.items)).toBe(true)
|
||||
// ── Step 3:填写标题(handleScriptModalConfirm 已预填 script.title,但我们再覆盖一次) ─
|
||||
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
const titleInput2 = page.getByPlaceholder("输入或从标题库选择")
|
||||
await expect(titleInput2).toBeVisible({ timeout: 5000 })
|
||||
await titleInput2.fill(`测试叙事剪辑 ${suffix}`)
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// ── Step 4:确认生成 ──────────────────────────────────────────
|
||||
await expect(page.getByText("📋 生成配置")).toBeVisible({ timeout: 10000 })
|
||||
await expect(page.getByText("叙事剪辑")).toBeVisible()
|
||||
const confirmBtn2 = page.getByRole("button", { name: /确认生成视频/ })
|
||||
await expect(confirmBtn2).toBeEnabled({ timeout: 5000 })
|
||||
|
||||
const createTask2 = page.waitForResponse(
|
||||
(r) => r.url().includes("/generation/tasks") && r.request().method() === "POST",
|
||||
{ timeout: 30000 },
|
||||
)
|
||||
await confirmBtn2.click()
|
||||
const taskResp2 = await createTask2
|
||||
expect(taskResp2.ok(), `Create task: ${await taskResp2.text()}`).toBeTruthy()
|
||||
console.log("[narrative] Generation task created:", (await taskResp2.json()).id)
|
||||
await expect(page.getByText(/正在生成|提交/)).toBeVisible({ timeout: 15000 })
|
||||
console.log("[narrative] Wizard flow completed ✓")
|
||||
})
|
||||
})
|
||||
|
||||
Executable
+105
@@ -0,0 +1,105 @@
|
||||
import { expect, test, type APIRequestContext, type Page } from "@playwright/test"
|
||||
|
||||
const PASSWORD = "SmokePass123!"
|
||||
const apiBase = process.env.E2E_API_BASE || "/api/v1"
|
||||
const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".length) : ""
|
||||
|
||||
async function routeBrowserApiToTestApi(page: Page) {
|
||||
if (!apiOrigin) return
|
||||
await page.route("**/api/v1/**", async (route) => {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
await route.fulfill({ response })
|
||||
})
|
||||
}
|
||||
|
||||
async function loginWithRetry(request: APIRequestContext, email: string, password: string) {
|
||||
for (let i = 0; i <= 2; i++) {
|
||||
const r = await request.post(`${apiBase}/auth/login`, { data: { email, password } })
|
||||
if (r.status() !== 429) {
|
||||
expect(r.ok(), `login: ${await r.text()}`).toBeTruthy()
|
||||
return (await r.json()).access_token as string
|
||||
}
|
||||
console.log(`[nav] 429 retry ${i + 1}/2`)
|
||||
await new Promise((res) => setTimeout(res, 65000))
|
||||
}
|
||||
throw new Error("Login retries exhausted")
|
||||
}
|
||||
|
||||
/**
|
||||
* 核心页面导航冒烟:侧边栏主要入口能访问、文案库/配音库页面能正常加载(不出白屏/无致命 js error)
|
||||
*/
|
||||
test.describe("Core Navigation", () => {
|
||||
let authToken: string
|
||||
|
||||
test.beforeAll(async ({ request }) => {
|
||||
const suffix = Math.random().toString(36).slice(2, 8)
|
||||
const email = `e2e-nav-${suffix}@example.com`
|
||||
await request.post(`${apiBase}/auth/register`, {
|
||||
data: { email, password: PASSWORD, username: `e2e_nav_${suffix}` },
|
||||
})
|
||||
authToken = await loginWithRetry(request, email, PASSWORD)
|
||||
const authHeader = { Authorization: `Bearer ${authToken}` }
|
||||
const proj = await request.post(`${apiBase}/projects`, {
|
||||
headers: authHeader,
|
||||
data: { name: `Smoke Nav ${suffix}` },
|
||||
})
|
||||
if (proj.ok()) {
|
||||
const projectId = (await proj.json()).id ?? (await proj.json()).project_id
|
||||
await request.post(`${apiBase}/asset-libraries`, {
|
||||
headers: authHeader,
|
||||
data: { project_id: projectId, name: "Nav Lib", kind: "video" },
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
test.beforeEach(async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 })
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
}, authToken)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
})
|
||||
|
||||
const navCases = [
|
||||
{ path: "/app/dashboard", marker: /概览|工作台|最近/i, name: "概览" },
|
||||
{ path: "/app/generate", marker: /智能剪辑|剪辑/, name: "智能剪辑" },
|
||||
{ path: "/app/assets", marker: /视频库|素材/, name: "视频库" },
|
||||
{ path: "/app/scripts", marker: /文案/, name: "文案库" },
|
||||
{ path: "/app/voices", marker: /配音|我的音色|配音库/, name: "配音库" },
|
||||
{ path: "/app/products", marker: /成品|作品/, name: "成品库" },
|
||||
{ path: "/app/history", marker: /历史|任务/, name: "任务历史" },
|
||||
{ path: "/app/tasks", marker: /任务中心|任务列表/, name: "任务中心" },
|
||||
{ path: "/app/points", marker: /积分|我的积分/, name: "积分中心" },
|
||||
]
|
||||
|
||||
for (const c of navCases) {
|
||||
test(`visit ${c.name} (${c.path}) loads without fatal pageerror`, async ({ page }) => {
|
||||
const errors: Error[] = []
|
||||
page.on("pageerror", (e) => errors.push(e))
|
||||
await page.goto(c.path)
|
||||
await expect(page.locator("body")).not.toBeEmpty({ timeout: 20000 })
|
||||
// 过滤掉常见第三方/非致命错误
|
||||
const fatal = errors.filter(
|
||||
(e) =>
|
||||
!/ResizeObserver|Loading chunk|network error|Failed to fetch|chunkLoadError/i.test(
|
||||
e.message,
|
||||
),
|
||||
)
|
||||
expect(fatal, `${c.name} pageerrors: ${fatal.map((e) => e.message).join("; ")}`).toHaveLength(
|
||||
0,
|
||||
)
|
||||
await expect(
|
||||
page.getByText(c.marker).first(),
|
||||
`${c.name} should show relevant text`,
|
||||
).toBeVisible({ timeout: 15000 })
|
||||
console.log(`[nav] ${c.name} loaded ✓`)
|
||||
})
|
||||
}
|
||||
})
|
||||
@@ -11,8 +11,12 @@ import type {
|
||||
ScriptCategory,
|
||||
} from "./types"
|
||||
|
||||
/** 是否启用 mock(后端合入后改为 false) */
|
||||
export const SCRIPTS_API_MOCK = true
|
||||
/**
|
||||
* 是否启用 mock。
|
||||
* #1894:文案库接口已上线,默认 false 走真实 API;
|
||||
* 通过 SCRIPTS_API_MOCK=true 环境变量可本地开启 mock 调试(行为同 POINTS_API_MOCK)。
|
||||
*/
|
||||
export const SCRIPTS_API_MOCK = (process.env.SCRIPTS_API_MOCK as string | undefined) === "true"
|
||||
|
||||
// ==================== Mock 数据 ====================
|
||||
|
||||
|
||||
@@ -71,6 +71,16 @@ export interface CreateGenerationTaskRequest {
|
||||
duration?: number
|
||||
/** 视频宽高比,如 "9:16" */
|
||||
video_ratio?: string
|
||||
/** #1970:剪辑模式 random/narrative */
|
||||
assembly_mode?: "random" | "narrative"
|
||||
/** #1970:叙事模式下的文案 ID */
|
||||
script_id?: string
|
||||
/** #1970:TTS 音色 ID */
|
||||
tts_voice_id?: string
|
||||
/** #1970:TTS 音色来源 preset/clone */
|
||||
tts_voice_source?: "preset" | "clone"
|
||||
/** #1970:智能降重开关(默认 true) */
|
||||
dedup_enabled?: boolean
|
||||
/** 标题烧录配置 */
|
||||
title_config?: {
|
||||
text?: string
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
/**
|
||||
* 标题相关 API — 目录化入口
|
||||
* 保持与原 titles.ts 相同导出,向后兼容
|
||||
*/
|
||||
|
||||
// 类型
|
||||
export type {
|
||||
TitleItem,
|
||||
BackendTitleResponse,
|
||||
BackendCreateTitleRequest,
|
||||
BackendUpdateTitleRequest,
|
||||
CreateTitleRequest,
|
||||
} from "./types"
|
||||
|
||||
// 工具函数
|
||||
export { toTitleItem } from "./utils"
|
||||
|
||||
// API 函数
|
||||
export { getTitles, createTitle, updateTitle, deleteTitle, batchImportTitles } from "./titles"
|
||||
@@ -1,65 +0,0 @@
|
||||
/**
|
||||
* 标题相关 API 函数
|
||||
* Phase 1 新增:全局标题库
|
||||
* 注意:后端 schema 使用 name + text 字段,前端 UI 用 content 展示
|
||||
*/
|
||||
import apiClient from "../client"
|
||||
import type {
|
||||
BackendCreateTitleRequest,
|
||||
BackendTitleResponse,
|
||||
BackendUpdateTitleRequest,
|
||||
CreateTitleRequest,
|
||||
TitleItem,
|
||||
} from "./types"
|
||||
import { toTitleItem } from "./utils"
|
||||
|
||||
/** 获取当前用户的所有标题 */
|
||||
export const getTitles = async (): Promise<TitleItem[]> => {
|
||||
const response = await apiClient.get<{ items: BackendTitleResponse[] } | BackendTitleResponse[]>(
|
||||
"/titles",
|
||||
)
|
||||
// 兼容两种后端返回格式:{ items: [...] } 或直接 [...]
|
||||
const items = Array.isArray(response.data) ? response.data : response.data.items || []
|
||||
return items.map(toTitleItem)
|
||||
}
|
||||
|
||||
/** 创建标题 */
|
||||
export const createTitle = async (data: CreateTitleRequest): Promise<TitleItem> => {
|
||||
// 后端要求 name(≤255)和 text(≤500),name 从 content 截取
|
||||
const payload: BackendCreateTitleRequest = {
|
||||
name: data.content.slice(0, 255),
|
||||
text: data.content.slice(0, 500),
|
||||
category: data.category || "default",
|
||||
}
|
||||
const response = await apiClient.post<BackendTitleResponse>("/titles", payload)
|
||||
return toTitleItem(response.data)
|
||||
}
|
||||
|
||||
/** 更新标题 */
|
||||
export const updateTitle = async (
|
||||
titleId: string,
|
||||
data: Partial<CreateTitleRequest>,
|
||||
): Promise<TitleItem> => {
|
||||
const payload: BackendUpdateTitleRequest = {}
|
||||
if (data.content !== undefined) {
|
||||
payload.name = data.content.slice(0, 255)
|
||||
payload.text = data.content.slice(0, 500)
|
||||
}
|
||||
if (data.category !== undefined) {
|
||||
payload.category = data.category
|
||||
}
|
||||
// 后端用 PUT,非 PATCH
|
||||
const response = await apiClient.put<BackendTitleResponse>(`/titles/${titleId}`, payload)
|
||||
return toTitleItem(response.data)
|
||||
}
|
||||
|
||||
/** 删除标题 */
|
||||
export const deleteTitle = async (titleId: string): Promise<void> => {
|
||||
await apiClient.delete(`/titles/${titleId}`)
|
||||
}
|
||||
|
||||
/** 批量导入标题 */
|
||||
export const batchImportTitles = async (titles: string[]): Promise<{ imported_count: number }> => {
|
||||
const response = await apiClient.post("/titles/batch-import", { titles })
|
||||
return response.data
|
||||
}
|
||||
@@ -1,54 +0,0 @@
|
||||
/**
|
||||
* 标题相关类型定义
|
||||
*/
|
||||
|
||||
/** 标题条目(前端展示用) */
|
||||
export interface TitleItem {
|
||||
id: string
|
||||
content: string
|
||||
category?: string
|
||||
source?: string
|
||||
word_count?: number
|
||||
is_favorite?: boolean
|
||||
created_at?: string
|
||||
updated_at?: string
|
||||
}
|
||||
|
||||
/** 后端标题响应格式 */
|
||||
export interface BackendTitleResponse {
|
||||
id: string
|
||||
user_id: string
|
||||
name: string
|
||||
text: string
|
||||
category: string
|
||||
description: string
|
||||
tags: string[]
|
||||
usage_count: number
|
||||
is_active: boolean
|
||||
created_at: string
|
||||
updated_at: string
|
||||
}
|
||||
|
||||
/** 后端创建标题请求格式 */
|
||||
export interface BackendCreateTitleRequest {
|
||||
name: string
|
||||
text: string
|
||||
category: string
|
||||
description?: string
|
||||
tags?: string[]
|
||||
}
|
||||
|
||||
/** 后端更新标题请求格式 */
|
||||
export interface BackendUpdateTitleRequest {
|
||||
name?: string
|
||||
text?: string
|
||||
category?: string
|
||||
description?: string
|
||||
tags?: string[]
|
||||
}
|
||||
|
||||
/** 创建标题请求(前端接口,保持向后兼容) */
|
||||
export interface CreateTitleRequest {
|
||||
content: string
|
||||
category?: string
|
||||
}
|
||||
@@ -1,14 +0,0 @@
|
||||
/**
|
||||
* 标题数据转换工具函数
|
||||
*/
|
||||
import type { BackendTitleResponse, TitleItem } from "./types"
|
||||
|
||||
/** 将后端响应映射为前端 TitleItem */
|
||||
export const toTitleItem = (item: BackendTitleResponse): TitleItem => ({
|
||||
id: item.id,
|
||||
content: item.text,
|
||||
category: item.category,
|
||||
word_count: item.text?.length || 0,
|
||||
created_at: item.created_at,
|
||||
updated_at: item.updated_at,
|
||||
})
|
||||
@@ -17,6 +17,7 @@ import {
|
||||
} from "@ant-design/icons"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import { usePointsStore } from "@/store/pointsStore"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
import "./PointsBadge.css"
|
||||
|
||||
const { Text, Paragraph } = Typography
|
||||
@@ -32,9 +33,13 @@ const PointsBadge: React.FC = () => {
|
||||
const { balance, membership, subscription, dailyUsage, init, loading } = usePointsStore()
|
||||
|
||||
useEffect(() => {
|
||||
if (!ENABLE_CREDIT_SYSTEM) return
|
||||
if (!balance) init()
|
||||
}, [balance, init])
|
||||
|
||||
// 功能开关:积分系统关闭时直接隐藏徽章
|
||||
if (!ENABLE_CREDIT_SYSTEM) return null
|
||||
|
||||
// 余额:优先用 membership.points_balance(冗余字段),降级 balance.balance
|
||||
const bal = membership?.points_balance ?? balance?.balance ?? 0
|
||||
const lowBalance = bal > 0 && bal < 10
|
||||
|
||||
@@ -15,6 +15,7 @@ import React, { useMemo } from "react"
|
||||
import { Tooltip } from "antd"
|
||||
import { WarningOutlined } from "@ant-design/icons"
|
||||
import { usePointsStore } from "@/store/pointsStore"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
import type { PointsSource } from "@/api/points/types"
|
||||
import "./PointsCost.css"
|
||||
|
||||
@@ -53,7 +54,7 @@ const PointsCost: React.FC<Props> = ({
|
||||
compact = false,
|
||||
showRechargeHint = true,
|
||||
className = "",
|
||||
}) => {
|
||||
}: Props) => {
|
||||
const { balance, dailyUsage, rules, membership } = usePointsStore()
|
||||
const qty = quantity ?? units ?? 1
|
||||
|
||||
@@ -118,6 +119,9 @@ const PointsCost: React.FC<Props> = ({
|
||||
}
|
||||
}, [rules, balance, dailyUsage, membership, scene, qty, durationMinutes])
|
||||
|
||||
// 积分系统关闭时不展示消耗提示(组件保留,hooks 必须在 return 前调用)
|
||||
if (!ENABLE_CREDIT_SYSTEM) return null
|
||||
|
||||
if (!rule || !balance) {
|
||||
return <span className={`xx-points-cost ${className}`} />
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ import { useLogout } from "@/hooks/useAuth"
|
||||
import type { MenuProps } from "antd"
|
||||
import { NAV_ITEMS } from "@/config/navigation"
|
||||
import PointsBadge from "@/components/common/PointsBadge"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
import { usePointsStore } from "@/store/pointsStore"
|
||||
import "./Header.css"
|
||||
|
||||
@@ -57,30 +58,36 @@ const Header: React.FC = () => {
|
||||
label: "订阅管理",
|
||||
onClick: () => navigate("/app/subscription"),
|
||||
},
|
||||
// v2: 我的积分入口
|
||||
{
|
||||
key: "points-center",
|
||||
icon: <ThunderboltOutlined />,
|
||||
label: (
|
||||
<Space>
|
||||
我的积分
|
||||
{balance && <span style={{ color: "#8b5cf6", fontWeight: 700 }}>{balance.balance}</span>}
|
||||
</Space>
|
||||
),
|
||||
onClick: () => navigate("/app/points"),
|
||||
},
|
||||
{
|
||||
key: "points-history",
|
||||
icon: <HistoryOutlined />,
|
||||
label: "积分明细",
|
||||
onClick: () => navigate("/app/points/transactions"),
|
||||
},
|
||||
{
|
||||
key: "recharge",
|
||||
icon: <WalletOutlined />,
|
||||
label: "充值积分",
|
||||
onClick: () => navigate("/app/points/recharge"),
|
||||
},
|
||||
// 积分系统开关关闭时隐藏积分相关菜单项(代码保留不删除)
|
||||
...(ENABLE_CREDIT_SYSTEM
|
||||
? [
|
||||
{
|
||||
key: "points-center",
|
||||
icon: <ThunderboltOutlined />,
|
||||
label: (
|
||||
<Space>
|
||||
我的积分
|
||||
{balance && (
|
||||
<span style={{ color: "#8b5cf6", fontWeight: 700 }}>{balance.balance}</span>
|
||||
)}
|
||||
</Space>
|
||||
),
|
||||
onClick: () => navigate("/app/points"),
|
||||
},
|
||||
{
|
||||
key: "points-history",
|
||||
icon: <HistoryOutlined />,
|
||||
label: "积分明细",
|
||||
onClick: () => navigate("/app/points/transactions"),
|
||||
},
|
||||
{
|
||||
key: "recharge",
|
||||
icon: <WalletOutlined />,
|
||||
label: "充值积分",
|
||||
onClick: () => navigate("/app/points/recharge"),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
{ type: "divider" },
|
||||
{
|
||||
key: "logout",
|
||||
@@ -130,7 +137,13 @@ const Header: React.FC = () => {
|
||||
|
||||
{/* v2: 升级会员入口(仅免费用户显示) */}
|
||||
{!isMember && (
|
||||
<Tooltip title="升级会员解锁无限混剪、批量导出,积分 8 折起">
|
||||
<Tooltip
|
||||
title={
|
||||
ENABLE_CREDIT_SYSTEM
|
||||
? "升级会员解锁无限混剪、批量导出,积分 8 折起"
|
||||
: "升级会员解锁无限混剪、批量导出"
|
||||
}
|
||||
>
|
||||
<Button
|
||||
type="primary"
|
||||
size="small"
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
/**
|
||||
* 功能开关配置
|
||||
* 集中管理前端特性的启用/隐藏,便于灰度与回滚。
|
||||
* 注意:仅控制 UI 展示与前端校验,后端扣减逻辑由后端对应开关控制。
|
||||
*/
|
||||
|
||||
/**
|
||||
* 积分系统 UI 开关(默认 false = 隐藏)
|
||||
* - false:隐藏所有积分相关入口/余额/消耗提示/不足弹窗/充值入口;会员标识保留;
|
||||
* 功能流程不做积分预校验,直接走生成。
|
||||
* - true:展示完整积分系统 UI。
|
||||
*/
|
||||
export const ENABLE_CREDIT_SYSTEM = false
|
||||
@@ -3,6 +3,7 @@
|
||||
* Header.tsx 和 Sidebar.tsx 共享此数据源,避免路由配置重复
|
||||
*/
|
||||
import React from "react"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "./features"
|
||||
import {
|
||||
DashboardOutlined,
|
||||
FileOutlined,
|
||||
@@ -105,12 +106,17 @@ export const NAV_ITEMS: NavItem[] = [
|
||||
path: "/app/subscription",
|
||||
icon: React.createElement(CrownOutlined),
|
||||
},
|
||||
{
|
||||
key: "points",
|
||||
label: "积分中心",
|
||||
path: "/app/points",
|
||||
icon: React.createElement(ThunderboltOutlined),
|
||||
},
|
||||
// 积分系统开关关闭时隐藏积分中心入口(代码保留不删除)
|
||||
...(ENABLE_CREDIT_SYSTEM
|
||||
? [
|
||||
{
|
||||
key: "points",
|
||||
label: "积分中心",
|
||||
path: "/app/points",
|
||||
icon: React.createElement(ThunderboltOutlined),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
]
|
||||
|
||||
/** 侧边栏导航分组(Sidebar 分组列表使用) */
|
||||
@@ -200,12 +206,17 @@ export const NAV_GROUPS: NavGroup[] = [
|
||||
path: "/app/subscription",
|
||||
icon: React.createElement(CrownOutlined),
|
||||
},
|
||||
{
|
||||
key: "points",
|
||||
label: "积分中心",
|
||||
path: "/app/points",
|
||||
icon: React.createElement(ThunderboltOutlined),
|
||||
},
|
||||
// 积分系统开关关闭时隐藏积分中心入口(代码保留不删除)
|
||||
...(ENABLE_CREDIT_SYSTEM
|
||||
? [
|
||||
{
|
||||
key: "points",
|
||||
label: "积分中心",
|
||||
path: "/app/points",
|
||||
icon: React.createElement(ThunderboltOutlined),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
@@ -619,6 +619,7 @@ const AiAvatarPage: React.FC = () => {
|
||||
scriptText={state.scriptText}
|
||||
onScriptTextChange={state.setScriptText}
|
||||
onOpenScriptModal={() => state.setShowScriptModal(true)}
|
||||
onScriptCreated={(s) => state.selectScript(s as import("./types").Script)}
|
||||
/>
|
||||
<div className="aa-step-btn-row">
|
||||
<button
|
||||
@@ -1146,8 +1147,9 @@ const ScriptSelectModalLazy: React.FC<{
|
||||
useEffect(() => {
|
||||
if (!open) return
|
||||
setLoading(true)
|
||||
getScripts()
|
||||
.then((items) => setScripts(Array.isArray(items) ? items : []))
|
||||
// #1894: getScripts 返回 { items, total } 分页结构,取 items 即可
|
||||
getScripts({ page_size: 200 })
|
||||
.then((res) => setScripts(Array.isArray(res) ? res : (res.items ?? [])))
|
||||
.catch(() => setScripts([]))
|
||||
.finally(() => setLoading(false))
|
||||
}, [open])
|
||||
|
||||
@@ -2,31 +2,15 @@
|
||||
* AI数字人 — API 封装(#1822 契约对齐)
|
||||
*/
|
||||
import apiClient from "@/api/client"
|
||||
import type { Script, LipsyncJob, RenderJob, BRollSegment, SentenceTiming } from "../types"
|
||||
// #1894: Script 类型统一从 @/api/scripts 取(ai-avatar 本地 Script 仅保留渲染/对口型等自有类型)
|
||||
import type { LipsyncJob, RenderJob, BRollSegment, SentenceTiming } from "../types"
|
||||
|
||||
/* ── 文案库 ── */
|
||||
export const getScripts = async (): Promise<Script[]> => {
|
||||
const response = await apiClient.get<{ items?: Script[] } | Script[]>("/scripts")
|
||||
// 后端列表返回 { items, total } 分页对象,做兼容解包 + 数组防御(#1809 白屏修复)
|
||||
const data = response.data as unknown
|
||||
if (Array.isArray(data)) return data
|
||||
const items = (data as { items?: Script[] })?.items
|
||||
return Array.isArray(items) ? items : []
|
||||
}
|
||||
|
||||
export const getScriptById = async (id: string): Promise<Script> => {
|
||||
const response = await apiClient.get<Script>(`/scripts/${id}`)
|
||||
return response.data
|
||||
}
|
||||
|
||||
export const createScript = async (data: { title: string; content: string }): Promise<Script> => {
|
||||
const response = await apiClient.post<Script>("/scripts", data)
|
||||
return response.data
|
||||
}
|
||||
|
||||
export const deleteScript = async (id: string): Promise<void> => {
|
||||
await apiClient.delete(`/scripts/${id}`)
|
||||
}
|
||||
/* ── 文案库 ──
|
||||
* #1894: 统一走 @/api/scripts 的 getScripts,不再各自封装;
|
||||
* 这样 mock 开关、分页/搜索参数、字段对齐都和文案库页面保持一致。
|
||||
*/
|
||||
// #1894: 统一复用文案库 API,不再在 ai-avatar 里重复实现
|
||||
export { getScripts, getScript as getScriptById, createScript, deleteScript } from "@/api/scripts"
|
||||
|
||||
/* ── 素材单查(拿到 file_url 作为对口型的 video_url) ── */
|
||||
export const getAssetById = async (id: string): Promise<{ file_url?: string; id: string }> => {
|
||||
@@ -60,7 +44,8 @@ export const createLipsyncJob = async (data: {
|
||||
enable_video_loop?: boolean
|
||||
project_id?: string
|
||||
}): Promise<LipsyncJob> => {
|
||||
const response = await apiClient.post<LipsyncJob>("/lipsync/jobs", data)
|
||||
// GPU 口型同步推理约 20s,留足余量到 120s 防止 10s 默认超时
|
||||
const response = await apiClient.post<LipsyncJob>("/lipsync/jobs", data, { timeout: 120_000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
/**
|
||||
* AI数字人 — 文案面板(步骤1用)
|
||||
* 文案库选择 / 手动输入 + 字数统计
|
||||
* #1894: 文案库选择走 @/api/scripts;手动输入支持一键「保存到文案库」
|
||||
*/
|
||||
import { useState } from "react"
|
||||
import { message } from "antd"
|
||||
import { createScript } from "../api/aiAvatar"
|
||||
|
||||
interface PanelScriptProps {
|
||||
scriptText: string
|
||||
onScriptTextChange: (text: string) => void
|
||||
onOpenScriptModal: () => void
|
||||
/** 手动保存到文案库后回调(把新脚本传入,父组件可更新 selectedScript) */
|
||||
onScriptCreated?: (script: { id: string; title: string; content: string }) => void
|
||||
}
|
||||
|
||||
type ScriptTab = "library" | "manual"
|
||||
@@ -16,8 +20,30 @@ export function PanelScript({
|
||||
scriptText,
|
||||
onScriptTextChange,
|
||||
onOpenScriptModal,
|
||||
onScriptCreated,
|
||||
}: PanelScriptProps) {
|
||||
const [scriptTab, setScriptTab] = useState<ScriptTab>("library")
|
||||
const [saving, setSaving] = useState(false)
|
||||
|
||||
const handleSaveToLibrary = async () => {
|
||||
const text = scriptText.trim()
|
||||
if (!text) {
|
||||
message.warning("请先输入文案内容")
|
||||
return
|
||||
}
|
||||
// 用正文前 20 字作为默认标题
|
||||
const autoTitle = text.slice(0, 20).replace(/\n+/g, " ").trim() || "手动输入文案"
|
||||
setSaving(true)
|
||||
try {
|
||||
const created = await createScript({ title: autoTitle, content: text, tags: [] })
|
||||
message.success({ content: "已保存到文案库", duration: 1 })
|
||||
onScriptCreated?.(created)
|
||||
} catch {
|
||||
message.error("保存到文案库失败,请稍后重试")
|
||||
} finally {
|
||||
setSaving(false)
|
||||
}
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="aa-script-lipsync">
|
||||
@@ -59,7 +85,20 @@ export function PanelScript({
|
||||
}
|
||||
onChange={(e) => onScriptTextChange(e.target.value)}
|
||||
/>
|
||||
<div className="aa-char-count">{scriptText.length} 字</div>
|
||||
<div style={{ display: "flex", justifyContent: "space-between", alignItems: "center" }}>
|
||||
<div className="aa-char-count">{scriptText.length} 字</div>
|
||||
{scriptTab === "manual" && scriptText.trim().length > 0 && (
|
||||
<button
|
||||
type="button"
|
||||
className="aa-btn aa-btn--text"
|
||||
disabled={saving}
|
||||
onClick={handleSaveToLibrary}
|
||||
style={{ fontSize: 12, padding: "2px 8px" }}
|
||||
>
|
||||
{saving ? "保存中..." : "💾 保存到文案库"}
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -15,7 +15,8 @@ import type { TitleOption } from "@/pages/generate/components/title/TitleLibrary
|
||||
import type { TitleSettings } from "@/pages/generate/types"
|
||||
import { POSITION_OPTIONS, FONT_OPTIONS, TITLE_PRESETS } from "@/pages/generate/constants"
|
||||
import type { AiAvatarTitleConfig } from "../types"
|
||||
import { getTitles } from "@/api/titles"
|
||||
// #1894: 标题数据源切换到文案库,取 script.title 作为候选
|
||||
import { getScripts } from "@/api/scripts"
|
||||
|
||||
const { TextArea } = Input
|
||||
|
||||
@@ -28,11 +29,23 @@ const PanelTitleConfig: React.FC<PanelTitleConfigProps> = ({ titleConfig, onUpda
|
||||
/** TitleStylePanel 内部高亮的预设 key(面板本地状态) */
|
||||
const [activePreset, setActivePreset] = useState<string | null>(null)
|
||||
|
||||
/** 标题库选项(复用智能剪辑的标题库) */
|
||||
/** 标题库选项(#1894:从文案库 scripts[].title 取候选) */
|
||||
const [titleOptions, setTitleOptions] = useState<TitleOption[]>([])
|
||||
useEffect(() => {
|
||||
getTitles()
|
||||
.then((items) => setTitleOptions(items.map((t) => ({ label: t.content, value: t.content }))))
|
||||
getScripts({ page_size: 200 })
|
||||
.then((res) => {
|
||||
const items = Array.isArray(res) ? res : (res.items ?? [])
|
||||
// 去重 + 过滤空标题
|
||||
const seen = new Set<string>()
|
||||
const opts: TitleOption[] = []
|
||||
for (const s of items) {
|
||||
const t = (s.title || "").trim()
|
||||
if (!t || seen.has(t)) continue
|
||||
seen.add(t)
|
||||
opts.push({ label: t, value: t })
|
||||
}
|
||||
setTitleOptions(opts)
|
||||
})
|
||||
.catch(() => setTitleOptions([]))
|
||||
}, [])
|
||||
|
||||
@@ -84,10 +97,12 @@ const PanelTitleConfig: React.FC<PanelTitleConfigProps> = ({ titleConfig, onUpda
|
||||
style={{ fontSize: 15 }}
|
||||
/>
|
||||
<div style={{ marginTop: 8, display: "flex", alignItems: "center", gap: 8 }}>
|
||||
<span style={{ fontSize: 12, color: "#8c8ca1", whiteSpace: "nowrap" }}>📚 标题库</span>
|
||||
<span style={{ fontSize: 12, color: "#8c8ca1", whiteSpace: "nowrap" }}>
|
||||
📚 文案库标题
|
||||
</span>
|
||||
<TitleLibraryAutoComplete
|
||||
key={titleConfig.title}
|
||||
placeholder="选择标题填入上方"
|
||||
placeholder="从文案库选择标题"
|
||||
value=""
|
||||
onChange={(val) => {
|
||||
if (val) onUpdate({ title: val })
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
/**
|
||||
* AI数字人 — 标题库选择弹窗
|
||||
* 复用智能剪辑的标题库 API,选择标题后填入输入框
|
||||
*/
|
||||
import React, { useEffect, useState } from "react"
|
||||
import { getTitles } from "@/api/titles"
|
||||
import type { TitleItem } from "@/api/titles/types"
|
||||
|
||||
interface TitleLibraryModalProps {
|
||||
open: boolean
|
||||
onClose: () => void
|
||||
onSelect: (title: string) => void
|
||||
}
|
||||
|
||||
const TitleLibraryModal: React.FC<TitleLibraryModalProps> = ({ open, onClose, onSelect }) => {
|
||||
const [titles, setTitles] = useState<TitleItem[]>([])
|
||||
const [loading, setLoading] = useState(false)
|
||||
const [search, setSearch] = useState("")
|
||||
|
||||
useEffect(() => {
|
||||
if (!open) return
|
||||
setLoading(true)
|
||||
getTitles()
|
||||
.then((items) => setTitles(items))
|
||||
.catch(() => setTitles([]))
|
||||
.finally(() => setLoading(false))
|
||||
}, [open])
|
||||
|
||||
const filtered = titles.filter(
|
||||
(t) => !search || t.content.toLowerCase().includes(search.toLowerCase()),
|
||||
)
|
||||
|
||||
if (!open) return null
|
||||
|
||||
return (
|
||||
<div className="aa-modal-overlay" onClick={onClose}>
|
||||
<div className="aa-modal" onClick={(e) => e.stopPropagation()} style={{ maxWidth: 600 }}>
|
||||
<div className="aa-modal__header">
|
||||
<span className="aa-modal__title">从标题库选择</span>
|
||||
<button className="aa-modal__close" onClick={onClose}></button>
|
||||
</div>
|
||||
<div className="aa-modal__body">
|
||||
<div style={{ marginBottom: 12 }}>
|
||||
<input
|
||||
className="aa-input"
|
||||
placeholder="搜索标题..."
|
||||
value={search}
|
||||
onChange={(e) => setSearch(e.target.value)}
|
||||
/>
|
||||
</div>
|
||||
{loading ? (
|
||||
<div style={{ textAlign: "center", padding: 40, color: "#8c8ca1" }}>加载中...</div>
|
||||
) : filtered.length === 0 ? (
|
||||
<div style={{ textAlign: "center", padding: 40, color: "#8c8ca1" }}>
|
||||
暂无标题,请先在标题库创建
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ maxHeight: 400, overflowY: "auto" }}>
|
||||
{filtered.map((t) => (
|
||||
<div
|
||||
key={t.id}
|
||||
style={{
|
||||
padding: "12px 16px",
|
||||
marginBottom: 8,
|
||||
background: "#f8f8fc",
|
||||
borderRadius: 8,
|
||||
cursor: "pointer",
|
||||
transition: "background 0.2s",
|
||||
}}
|
||||
onMouseEnter={(e) => (e.currentTarget.style.background = "#eef0ff")}
|
||||
onMouseLeave={(e) => (e.currentTarget.style.background = "#f8f8fc")}
|
||||
onClick={() => {
|
||||
onSelect(t.content)
|
||||
onClose()
|
||||
}}
|
||||
>
|
||||
<div style={{ fontSize: 14, color: "#1a1a2e", marginBottom: 4 }}>{t.content}</div>
|
||||
<div style={{ fontSize: 12, color: "#8c8ca1" }}>
|
||||
{t.word_count ?? t.content.length}字 ·{" "}
|
||||
{t.created_at ? new Date(t.created_at).toLocaleDateString() : ""}
|
||||
</div>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
<div className="aa-modal__footer">
|
||||
<button className="aa-btn" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default TitleLibraryModal
|
||||
@@ -56,15 +56,11 @@ export interface TtsPreviewResult {
|
||||
error: string | null
|
||||
}
|
||||
|
||||
/* ── 文案 ── */
|
||||
export interface Script {
|
||||
id: string
|
||||
title: string
|
||||
content: string
|
||||
char_count: number
|
||||
created_at: string
|
||||
updated_at?: string
|
||||
}
|
||||
/* ── 文案 ──
|
||||
* #1894: 直接复用文案库的 ScriptItem 类型,保证字段(title/content/tags/...)一致;
|
||||
* 个别 ai-avatar 专属属性如有需要再在此处扩展。
|
||||
*/
|
||||
export type Script = import("@/api/scripts").ScriptItem
|
||||
|
||||
/* ── 对口型任务 ── */
|
||||
export interface LipsyncJob {
|
||||
|
||||
@@ -11,6 +11,9 @@ import type { VoiceClone } from "@/api/voice-clone"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { useCloneProgress } from "@/hooks/useCloneProgress"
|
||||
import CloneModal from "@/components/voice/CloneModal"
|
||||
import VoiceSelectModal from "./components/VoiceSelectModal"
|
||||
import ScriptSelectModal from "./components/ScriptSelectModal"
|
||||
import TtsVoiceModal from "./components/TtsVoiceModal"
|
||||
import GenerateHeader from "./components/GenerateHeader"
|
||||
import FrontendPreviewPlayer from "./components/FrontendPreviewPlayer"
|
||||
import CanvasPreviewGrid from "./components/CanvasPreviewGrid"
|
||||
@@ -30,6 +33,7 @@ import { getAssetsByKind } from "@/api/assets"
|
||||
import { previewTts } from "@/api/tts"
|
||||
import { usePointsStore } from "@/store/pointsStore"
|
||||
import { hasEnoughPoints } from "./hooks/pointsCost"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
import "./generate.css"
|
||||
import "./generate-points.css"
|
||||
|
||||
@@ -62,11 +66,26 @@ const GeneratePage: React.FC = () => {
|
||||
selectedVoice,
|
||||
setSelectedVoice,
|
||||
voiceMode,
|
||||
setVoiceMode,
|
||||
selectedClonedVoice,
|
||||
setSelectedClonedVoice,
|
||||
editMode,
|
||||
setEditMode,
|
||||
selectedScript,
|
||||
setSelectedScript,
|
||||
ttsVoiceId,
|
||||
setTtsVoiceId,
|
||||
ttsVoiceSource,
|
||||
setTtsVoiceSource,
|
||||
ttsVoiceAssetId,
|
||||
setTtsVoiceAssetId,
|
||||
dedupEnabled,
|
||||
setDedupEnabled,
|
||||
|
||||
cloneModalOpen,
|
||||
setCloneModalOpen,
|
||||
videoRatio,
|
||||
setVideoRatio,
|
||||
duration,
|
||||
style,
|
||||
autoSubtitles,
|
||||
@@ -124,6 +143,11 @@ const GeneratePage: React.FC = () => {
|
||||
/* ── 数量选择弹窗 ── */
|
||||
const [countModalOpen, setCountModalOpen] = useState(false)
|
||||
|
||||
/* ── #1970 流程重构:分支弹窗 ── */
|
||||
const [voiceModalOpen, setVoiceModalOpen] = useState(false)
|
||||
const [scriptModalOpen, setScriptModalOpen] = useState(false)
|
||||
const [ttsModalOpen, setTtsModalOpen] = useState(false)
|
||||
|
||||
/* ── 标题样式回调 ── */
|
||||
const styleUpdaters = useTitleStyleUpdaters({
|
||||
titleSettings,
|
||||
@@ -301,6 +325,12 @@ const GeneratePage: React.FC = () => {
|
||||
selectedClonedVoice,
|
||||
coverSettings,
|
||||
videoRatio,
|
||||
editMode,
|
||||
selectedScript,
|
||||
ttsVoiceId,
|
||||
ttsVoiceSource,
|
||||
ttsVoiceAssetId,
|
||||
dedupEnabled,
|
||||
style,
|
||||
duration,
|
||||
autoSubtitles,
|
||||
@@ -340,7 +370,7 @@ const GeneratePage: React.FC = () => {
|
||||
return Array.from({ length: count }, (_, i) => list[i] ?? "")
|
||||
})
|
||||
setSelectedVariantIds(Array.from({ length: count }, (_, i) => i))
|
||||
setCurrentStep(2)
|
||||
setCurrentStep(3)
|
||||
},
|
||||
[
|
||||
setPreviewCount,
|
||||
@@ -354,21 +384,76 @@ const GeneratePage: React.FC = () => {
|
||||
],
|
||||
)
|
||||
|
||||
/* ── #1970:Step1 弹窗回调 ── */
|
||||
const handleVoiceModalConfirm = useCallback(
|
||||
(voiceAssetId: string) => {
|
||||
setSelectedVoice(voiceAssetId)
|
||||
setVoiceMode("custom")
|
||||
setVoiceModalOpen(false)
|
||||
setCurrentStep(2)
|
||||
},
|
||||
[setSelectedVoice, setVoiceMode, setCurrentStep],
|
||||
)
|
||||
|
||||
const handleScriptModalConfirm = useCallback(
|
||||
(script: import("@/api/scripts").ScriptItem) => {
|
||||
setSelectedScript(script)
|
||||
// 自动带入标题(若标题为空则预填)
|
||||
if (!titleSettings.title?.trim() && script.title) {
|
||||
setTitleSettings((prev) => ({ ...prev, title: script.title, aiAutoSelect: false }))
|
||||
}
|
||||
setScriptModalOpen(false)
|
||||
// 自动打开 TTS 弹窗
|
||||
setTtsModalOpen(true)
|
||||
},
|
||||
[setSelectedScript, setTitleSettings, titleSettings.title],
|
||||
)
|
||||
|
||||
const handleTtsSynthesized = useCallback(
|
||||
(payload: { voiceAssetId: string; ttsVoiceId: string; ttsVoiceSource: "preset" | "clone" }) => {
|
||||
setTtsVoiceId(payload.ttsVoiceId)
|
||||
setTtsVoiceSource(payload.ttsVoiceSource)
|
||||
setTtsVoiceAssetId(payload.voiceAssetId)
|
||||
if (payload.ttsVoiceSource === "clone") {
|
||||
setSelectedClonedVoice(payload.ttsVoiceId)
|
||||
setVoiceMode("clone")
|
||||
} else {
|
||||
setSelectedVoice(payload.ttsVoiceId)
|
||||
setVoiceMode("preset")
|
||||
}
|
||||
setTtsModalOpen(false)
|
||||
message.success("配音合成成功")
|
||||
setCurrentStep(2)
|
||||
},
|
||||
[
|
||||
setTtsVoiceId,
|
||||
setTtsVoiceSource,
|
||||
setTtsVoiceAssetId,
|
||||
setSelectedVoice,
|
||||
setSelectedClonedVoice,
|
||||
setVoiceMode,
|
||||
setCurrentStep,
|
||||
],
|
||||
)
|
||||
|
||||
/* ── 步骤3「确认生成视频」:校验通过 → 创建正式生成任务 → 跳步骤4看实时进展 ── */
|
||||
const handleConfirmGenerate = useCallback(async () => {
|
||||
// 积分预检查
|
||||
const units = isBatch ? Math.max(selectedVariantIds.length, 1) : 1
|
||||
const check = hasEnoughPoints(
|
||||
balance ?? null,
|
||||
units,
|
||||
dailyUsage ?? null,
|
||||
[],
|
||||
"free",
|
||||
rules?.free_user_multiplier ?? 1.15,
|
||||
)
|
||||
if (!check.sufficient) {
|
||||
message.error(check.reason ?? "积分不足,请充值")
|
||||
return
|
||||
// 积分预检查(积分系统关闭时跳过,直接走生成流程)
|
||||
let check: ReturnType<typeof hasEnoughPoints> = { sufficient: true, cost: 0 }
|
||||
if (ENABLE_CREDIT_SYSTEM) {
|
||||
const units = isBatch ? Math.max(selectedVariantIds.length, 1) : 1
|
||||
check = hasEnoughPoints(
|
||||
balance ?? null,
|
||||
units,
|
||||
dailyUsage ?? null,
|
||||
[],
|
||||
"free",
|
||||
rules?.free_user_multiplier ?? 1.15,
|
||||
)
|
||||
if (!check.sufficient) {
|
||||
message.error(check.reason ?? "积分不足,请充值")
|
||||
return
|
||||
}
|
||||
}
|
||||
if (isBatch) {
|
||||
if (selectedVariantIds.length === 0) {
|
||||
@@ -412,12 +497,20 @@ const GeneratePage: React.FC = () => {
|
||||
const { goNext, goPrev } = useStepNavigation({
|
||||
currentStep,
|
||||
setCurrentStep,
|
||||
editMode,
|
||||
materialMode,
|
||||
selectedMaterials,
|
||||
smartSelectedIds,
|
||||
titleSettings,
|
||||
generated,
|
||||
onOpenCountModal: () => setCountModalOpen(true),
|
||||
onOpenStep1Modal: () => {
|
||||
if (editMode === "random") {
|
||||
setVoiceModalOpen(true)
|
||||
} else {
|
||||
setScriptModalOpen(true)
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
/* ── 最终成片 ── */
|
||||
@@ -431,19 +524,18 @@ const GeneratePage: React.FC = () => {
|
||||
|
||||
/* ── 积分消耗估算(步骤3确认生成展示用) ── */
|
||||
const unitsForCost = isBatch ? Math.max(selectedVariantIds.length, 1) : 1
|
||||
const pointsEstimate = useMemo(
|
||||
() =>
|
||||
hasEnoughPoints(
|
||||
balance ?? null,
|
||||
unitsForCost,
|
||||
dailyUsage ?? null,
|
||||
[],
|
||||
"free",
|
||||
rules?.free_user_multiplier ?? 1.15,
|
||||
),
|
||||
[unitsForCost, balance, dailyUsage, rules],
|
||||
)
|
||||
const insufficientPoints = !pointsEstimate.sufficient
|
||||
const pointsEstimate = useMemo(() => {
|
||||
if (!ENABLE_CREDIT_SYSTEM) return { sufficient: true, cost: 0 }
|
||||
return hasEnoughPoints(
|
||||
balance ?? null,
|
||||
unitsForCost,
|
||||
dailyUsage ?? null,
|
||||
[],
|
||||
"free",
|
||||
rules?.free_user_multiplier ?? 1.15,
|
||||
)
|
||||
}, [unitsForCost, balance, dailyUsage, rules])
|
||||
const insufficientPoints = ENABLE_CREDIT_SYSTEM && !pointsEstimate.sufficient
|
||||
|
||||
/* ================================================================
|
||||
渲染
|
||||
@@ -462,7 +554,7 @@ const GeneratePage: React.FC = () => {
|
||||
{!isBatch ? (
|
||||
<FrontendPreviewPlayer
|
||||
assets={previewAssets}
|
||||
videoRatio={videoRatio}
|
||||
videoRatio={videoRatio as "9:16" | "16:9"}
|
||||
ready={previewAssets.length > 0}
|
||||
serverClips={serverClips}
|
||||
voiceAudioUrl={previewVoiceAudioUrl || undefined}
|
||||
@@ -497,7 +589,7 @@ const GeneratePage: React.FC = () => {
|
||||
<CanvasPreviewGrid
|
||||
count={previewCount}
|
||||
assets={previewAssets}
|
||||
videoRatio={videoRatio}
|
||||
videoRatio={videoRatio as "9:16" | "16:9"}
|
||||
titles={previewTitles}
|
||||
titleSettings={titleSettings}
|
||||
voiceAudioUrls={variantVoiceAudioUrls}
|
||||
@@ -545,6 +637,16 @@ const GeneratePage: React.FC = () => {
|
||||
coverSettings={coverSettings}
|
||||
onCoverSettingsChange={setCoverSettings}
|
||||
selectedVoice={selectedVoice}
|
||||
editMode={editMode}
|
||||
onEditModeChange={setEditMode}
|
||||
dedupEnabled={dedupEnabled}
|
||||
onDedupEnabledChange={setDedupEnabled}
|
||||
onPreviewCountChange={setPreviewCount}
|
||||
videoRatio={videoRatio as "9:16" | "16:9"}
|
||||
onVideoRatioChange={(r) => setVideoRatio(r)}
|
||||
selectedScript={selectedScript}
|
||||
ttsVoiceId={ttsVoiceId}
|
||||
ttsVoiceSource={ttsVoiceSource}
|
||||
onSelectedVoiceChange={setSelectedVoice}
|
||||
onServerClipsChange={setServerClips}
|
||||
generating={generating}
|
||||
@@ -671,6 +773,27 @@ const GeneratePage: React.FC = () => {
|
||||
onClose={() => setCloneModalOpen(false)}
|
||||
onSuccess={handleCloneSuccess}
|
||||
/>
|
||||
|
||||
{/* #1970 流程弹窗 */}
|
||||
<VoiceSelectModal
|
||||
open={voiceModalOpen}
|
||||
selectedVoice={selectedVoice}
|
||||
onCancel={() => setVoiceModalOpen(false)}
|
||||
onConfirm={handleVoiceModalConfirm}
|
||||
/>
|
||||
<ScriptSelectModal
|
||||
open={scriptModalOpen}
|
||||
selectedScriptId={selectedScript?.id ?? null}
|
||||
onCancel={() => setScriptModalOpen(false)}
|
||||
onConfirm={handleScriptModalConfirm}
|
||||
/>
|
||||
<TtsVoiceModal
|
||||
open={ttsModalOpen}
|
||||
scriptText={selectedScript?.content ?? ""}
|
||||
scriptTitle={selectedScript?.title ?? ""}
|
||||
onCancel={() => setTtsModalOpen(false)}
|
||||
onSynthesized={handleTtsSynthesized}
|
||||
/>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
/**
|
||||
* GeneratePage 步骤内容渲染(#1899 简化为 5 步,#1913 传递 selectedTemplate)
|
||||
* 步骤顺序:素材(1) → 配音(2) → 标题(3) → 确认生成(4) → 封面(5)
|
||||
* 步骤3预览(Canvas 网格)与步骤4进度(批量渲染网格)由 GeneratePage 直接渲染在左侧大区域。
|
||||
* GeneratePage 步骤内容渲染(#1970 流程重构)
|
||||
* 步骤顺序:选择模式(1) → 选择素材(2) → 选择标题(3) → 确认生成(4) → 选择封面(5)
|
||||
* 原步骤"选择配音"已从主流程移除,改为 Step1 下一步分支弹窗(VoiceSelectModal / ScriptSelectModal → TtsVoiceModal)。
|
||||
*/
|
||||
import React from "react"
|
||||
import type { EditPlanClip } from "@/api/template-editor"
|
||||
import type { CoverConfig } from "../types/cover"
|
||||
import type { TitleSettings } from "../types"
|
||||
import type { ScriptItem } from "@/api/scripts"
|
||||
import Step1EditMode from "./Step1EditMode"
|
||||
import type { EditMode } from "./Step1EditMode"
|
||||
import Step2MaterialSelect from "../components/Step2MaterialSelect"
|
||||
import Step3VoiceWithMode from "./Step3VoiceWithMode"
|
||||
import Step4TitleSettings from "../components/Step4TitleSettings"
|
||||
import Step6CoverSettings from "../components/Step6CoverSettings"
|
||||
import BatchGenerationGrid from "./BatchGenerationGrid"
|
||||
@@ -17,19 +19,28 @@ import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
export interface GenerateStepContentProps {
|
||||
currentStep: number
|
||||
/* 片段数量(#1899) */
|
||||
/* Step1:剪辑模式 + 生成设置 */
|
||||
editMode: EditMode
|
||||
onEditModeChange: (m: EditMode) => void
|
||||
dedupEnabled: boolean
|
||||
onDedupEnabledChange: (v: boolean) => void
|
||||
/* ── 片段数量(#1899) ── */
|
||||
clipCount: number
|
||||
onClipCountChange: (n: number) => void
|
||||
/* 素材 */
|
||||
/* ── 生成数量/比例(Step1 设置) ── */
|
||||
previewCount: number
|
||||
onPreviewCountChange: (n: number) => void
|
||||
videoRatio: "9:16" | "16:9"
|
||||
onVideoRatioChange: (r: "9:16" | "16:9") => void
|
||||
/* ── 素材 ── */
|
||||
materialMode: "manual" | "auto"
|
||||
onMaterialModeChange: (mode: "manual" | "auto") => void
|
||||
selectedMaterials: string[]
|
||||
onSelectedMaterialsChange: (ids: string[]) => void
|
||||
smartSelectedIds: string[]
|
||||
onSmartSelectedIdsChange: (ids: string[]) => void
|
||||
/* 当前选中的模板/草稿 ID;空串时由后端自动兜底(#1913) */
|
||||
selectedTemplate?: string
|
||||
/* 标题 */
|
||||
/* ── 标题 ── */
|
||||
titleSettings: TitleSettings
|
||||
onTitleSettingsChange: (settings: TitleSettings) => void
|
||||
onUpdatePosition: (position: string) => void
|
||||
@@ -42,14 +53,14 @@ export interface GenerateStepContentProps {
|
||||
onApplyPreset: (presetKey: string) => void
|
||||
activePreset: string | null
|
||||
titlePresets: { key: string; label: string; previewStyle: React.CSSProperties }[]
|
||||
/* 封面 */
|
||||
/* ── 封面 ── */
|
||||
coverSettings: CoverConfig
|
||||
onCoverSettingsChange: (settings: CoverConfig) => void
|
||||
/* 配音 */
|
||||
/* ── 配音 ── */
|
||||
selectedVoice: string
|
||||
onSelectedVoiceChange: (id: string) => void
|
||||
onServerClipsChange: (clips: EditPlanClip[]) => void
|
||||
/* 生成 */
|
||||
/* ── 生成 ── */
|
||||
generating: boolean
|
||||
generated: boolean
|
||||
generateError: string | null
|
||||
@@ -58,14 +69,10 @@ export interface GenerateStepContentProps {
|
||||
onRetry: () => void
|
||||
onRetryBatchTask: (taskId: string) => void
|
||||
onDismissError: () => void
|
||||
/** 批量:每个正式生成任务的独立状态(步骤4进度网格) */
|
||||
batchTasks: BatchTaskState[]
|
||||
/** BGM 开关 */
|
||||
bgm: boolean
|
||||
/** BGM 配置 */
|
||||
bgmConfig?: { enabled: boolean; music_id?: string }
|
||||
/* ── 批量生成(#1677)── */
|
||||
previewCount: number
|
||||
/* ── 批量生成 ── */
|
||||
previewTitles: string[]
|
||||
onPreviewTitlesChange: (titles: string[]) => void
|
||||
voiceModePerVideo: boolean
|
||||
@@ -74,15 +81,27 @@ export interface GenerateStepContentProps {
|
||||
onVoiceLibraryIdsChange: (ids: string[]) => void
|
||||
previewCovers: string[]
|
||||
onPreviewCoversChange: (urls: string[]) => void
|
||||
/** 批量模式勾选的变体索引 */
|
||||
selectedVariantIds?: number[]
|
||||
/* ── 摘要信息(#1970 Step4 展示用) ── */
|
||||
selectedScript: ScriptItem | null
|
||||
ttsVoiceId: string
|
||||
ttsVoiceSource: "preset" | "clone"
|
||||
}
|
||||
|
||||
export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) => {
|
||||
// Only destructure props actually referenced in JSX below
|
||||
const {
|
||||
currentStep,
|
||||
editMode,
|
||||
onEditModeChange,
|
||||
dedupEnabled,
|
||||
onDedupEnabledChange,
|
||||
clipCount,
|
||||
onClipCountChange,
|
||||
previewCount,
|
||||
onPreviewCountChange,
|
||||
videoRatio,
|
||||
onVideoRatioChange,
|
||||
materialMode,
|
||||
onMaterialModeChange,
|
||||
selectedMaterials,
|
||||
@@ -104,31 +123,24 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
titlePresets,
|
||||
coverSettings,
|
||||
onCoverSettingsChange,
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
onServerClipsChange,
|
||||
generating,
|
||||
generated,
|
||||
generateError,
|
||||
progress,
|
||||
onRetry,
|
||||
generatedVideos,
|
||||
batchTasks,
|
||||
onRetryBatchTask,
|
||||
previewCount,
|
||||
previewTitles,
|
||||
onPreviewTitlesChange,
|
||||
voiceModePerVideo,
|
||||
onVoiceModePerVideoChange,
|
||||
voiceLibraryIds,
|
||||
onVoiceLibraryIdsChange,
|
||||
previewCovers,
|
||||
onPreviewCoversChange,
|
||||
selectedVariantIds,
|
||||
selectedScript,
|
||||
ttsVoiceId,
|
||||
ttsVoiceSource,
|
||||
} = props
|
||||
|
||||
// #1913:包装 onServerClipsChange,适配 hook 的 (clips, templateId?) 签名
|
||||
// 如果 hook 传回了后端兜底创建的 templateId,同时通知外层更新 selectedTemplate
|
||||
const handleClipsChange = React.useCallback(
|
||||
(clips: EditPlanClip[], _templateId?: string) => {
|
||||
onServerClipsChange(clips)
|
||||
@@ -138,8 +150,22 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
|
||||
switch (currentStep) {
|
||||
case 1:
|
||||
return (
|
||||
<Step1EditMode
|
||||
editMode={editMode}
|
||||
onEditModeChange={onEditModeChange}
|
||||
previewCount={previewCount}
|
||||
onPreviewCountChange={onPreviewCountChange}
|
||||
videoRatio={videoRatio}
|
||||
onVideoRatioChange={onVideoRatioChange}
|
||||
dedupEnabled={dedupEnabled}
|
||||
onDedupEnabledChange={onDedupEnabledChange}
|
||||
/>
|
||||
)
|
||||
case 2:
|
||||
return (
|
||||
<Step2MaterialSelect
|
||||
editMode={editMode}
|
||||
materialMode={materialMode}
|
||||
onMaterialModeChange={onMaterialModeChange}
|
||||
selectedMaterials={selectedMaterials}
|
||||
@@ -152,18 +178,6 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
onServerClipsChange={handleClipsChange}
|
||||
/>
|
||||
)
|
||||
case 2:
|
||||
return (
|
||||
<Step3VoiceWithMode
|
||||
previewCount={previewCount}
|
||||
selectedVoice={selectedVoice}
|
||||
onSelectedVoiceChange={onSelectedVoiceChange}
|
||||
voiceModePerVideo={voiceModePerVideo}
|
||||
onVoiceModePerVideoChange={onVoiceModePerVideoChange}
|
||||
voiceLibraryIds={voiceLibraryIds}
|
||||
onVoiceLibraryIdsChange={onVoiceLibraryIdsChange}
|
||||
/>
|
||||
)
|
||||
case 3:
|
||||
return (
|
||||
<Step4TitleSettings
|
||||
@@ -185,20 +199,49 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
/>
|
||||
)
|
||||
case 4:
|
||||
/* 确认生成页:批量=逐任务进度网格;单视频=仅渲染进度/失败状态 */
|
||||
if (previewCount > 1) {
|
||||
return (
|
||||
<BatchGenerationGrid
|
||||
tasks={batchTasks}
|
||||
titles={previewTitles}
|
||||
onRetryTask={onRetryBatchTask}
|
||||
/>
|
||||
)
|
||||
}
|
||||
if (generated && !generating && !generateError) return null
|
||||
return (
|
||||
<div className="xx-form-section">
|
||||
{generating && (
|
||||
{/* 配置摘要(#1970) */}
|
||||
<div
|
||||
style={{
|
||||
padding: 14,
|
||||
background: "#f9fafb",
|
||||
borderRadius: 8,
|
||||
marginBottom: 16,
|
||||
fontSize: 13,
|
||||
lineHeight: 1.8,
|
||||
color: "#374151",
|
||||
}}
|
||||
>
|
||||
<div style={{ fontWeight: 600, fontSize: 14, marginBottom: 6, color: "#111" }}>
|
||||
📋 生成配置
|
||||
</div>
|
||||
<div>🎬 剪辑模式:{editMode === "random" ? "🎲 随机混剪" : "📖 叙事剪辑"}</div>
|
||||
{editMode === "random" ? (
|
||||
<div>🎙️ 配音来源:配音库音频</div>
|
||||
) : (
|
||||
<>
|
||||
<div>📝 文案:{selectedScript?.title ?? "未选择"}</div>
|
||||
<div>
|
||||
🎙️ 合成配音音色:
|
||||
{ttsVoiceId
|
||||
? `${ttsVoiceSource === "clone" ? "克隆音色" : "系统音色"}(${ttsVoiceId.slice(0, 8)}...)`
|
||||
: "未选择"}
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
<div>📱 视频比例:{videoRatio}</div>
|
||||
<div>🎯 智能降重:{dedupEnabled ? "已开启" : "已关闭"}</div>
|
||||
{previewCount > 1 && <div>📦 生成数量:{previewCount} 个</div>}
|
||||
</div>
|
||||
|
||||
{previewCount > 1 ? (
|
||||
<BatchGenerationGrid
|
||||
tasks={batchTasks}
|
||||
titles={previewTitles}
|
||||
onRetryTask={onRetryBatchTask}
|
||||
/>
|
||||
) : generating ? (
|
||||
<div className="xx-gen-progress-card">
|
||||
<div className="xx-gen-progress-header">
|
||||
<div className="xx-gen-progress-info">
|
||||
@@ -217,8 +260,7 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{generateError && !generating && (
|
||||
) : generateError ? (
|
||||
<div className="xx-gen-error-card">
|
||||
<div className="xx-gen-error-info">
|
||||
<div className="xx-gen-error-title">生成失败</div>
|
||||
@@ -228,7 +270,7 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
🔄 重试
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
) : null}
|
||||
</div>
|
||||
)
|
||||
case 5:
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
/**
|
||||
* 叙事剪辑 — 文案选择弹窗(#1970)
|
||||
* - 搜索框:防抖 300ms,命中文字黄色高亮
|
||||
* - 标签筛选行:全部/带货/工厂/测评/教程/口播/种草
|
||||
* - 数量统计 + 卡片列表(可滚动,max-height 420px)
|
||||
* - 调用 GET /api/v1/scripts?keyword=&tag=&page_size=200
|
||||
*/
|
||||
import React, { useState, useEffect, useMemo, useRef, useCallback } from "react"
|
||||
import { Modal, Input, Tag, Spin } from "antd"
|
||||
import { SearchOutlined, CheckCircleFilled } from "@ant-design/icons"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { getScripts } from "@/api/scripts"
|
||||
import type { ScriptItem } from "@/api/scripts"
|
||||
|
||||
interface ScriptSelectModalProps {
|
||||
open: boolean
|
||||
selectedScriptId: string | null
|
||||
onCancel: () => void
|
||||
onConfirm: (script: ScriptItem) => void
|
||||
}
|
||||
|
||||
const SCRIPT_TABS = [
|
||||
{ key: "all", label: "全部" },
|
||||
{ key: "带货", label: "带货" },
|
||||
{ key: "工厂", label: "工厂" },
|
||||
{ key: "测评", label: "测评" },
|
||||
{ key: "教程", label: "教程" },
|
||||
{ key: "口播", label: "口播" },
|
||||
{ key: "种草", label: "种草" },
|
||||
]
|
||||
|
||||
/** 在文本中用 <mark> 高亮关键词(黄色背景) */
|
||||
function highlight(text: string, keyword: string): React.ReactNode {
|
||||
if (!keyword) return text
|
||||
const idx = text.toLowerCase().indexOf(keyword.toLowerCase())
|
||||
if (idx < 0) return text
|
||||
return (
|
||||
<>
|
||||
{text.slice(0, idx)}
|
||||
<mark style={{ background: "#fef08a", color: "#713f12", padding: "0 2px", borderRadius: 2 }}>
|
||||
{text.slice(idx, idx + keyword.length)}
|
||||
</mark>
|
||||
{text.slice(idx + keyword.length)}
|
||||
</>
|
||||
)
|
||||
}
|
||||
|
||||
const ScriptSelectModal: React.FC<ScriptSelectModalProps> = ({
|
||||
open,
|
||||
selectedScriptId,
|
||||
onCancel,
|
||||
onConfirm,
|
||||
}) => {
|
||||
const [innerSelected, setInnerSelected] = useState<string | null>(selectedScriptId)
|
||||
const [activeTag, setActiveTag] = useState<string>("all")
|
||||
const [searchInput, setSearchInput] = useState("")
|
||||
const [debouncedKw, setDebouncedKw] = useState("")
|
||||
const debounceRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setInnerSelected(selectedScriptId)
|
||||
setActiveTag("all")
|
||||
setSearchInput("")
|
||||
setDebouncedKw("")
|
||||
}
|
||||
}, [open, selectedScriptId])
|
||||
|
||||
// 300ms 防抖
|
||||
useEffect(() => {
|
||||
if (debounceRef.current) clearTimeout(debounceRef.current)
|
||||
debounceRef.current = setTimeout(() => setDebouncedKw(searchInput.trim()), 300)
|
||||
return () => {
|
||||
if (debounceRef.current) clearTimeout(debounceRef.current)
|
||||
}
|
||||
}, [searchInput])
|
||||
|
||||
const { data, isLoading } = useQuery({
|
||||
queryKey: ["scripts", "select-modal", debouncedKw, activeTag],
|
||||
queryFn: () =>
|
||||
getScripts({
|
||||
page: 1,
|
||||
page_size: 200,
|
||||
keyword: debouncedKw || undefined,
|
||||
tag: activeTag === "all" ? undefined : activeTag,
|
||||
}),
|
||||
enabled: open,
|
||||
})
|
||||
|
||||
const scripts: ScriptItem[] = useMemo(() => data?.items ?? [], [data])
|
||||
const selected = useMemo(
|
||||
() => scripts.find((s) => s.id === innerSelected) ?? null,
|
||||
[scripts, innerSelected],
|
||||
)
|
||||
|
||||
const handleConfirm = useCallback(() => {
|
||||
if (selected) onConfirm(selected)
|
||||
}, [selected, onConfirm])
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="📝 选择文案"
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={handleConfirm}
|
||||
okText="确认选择"
|
||||
cancelText="取消"
|
||||
okButtonProps={{ disabled: !selected, style: { background: "#7c3aed" } }}
|
||||
width={680}
|
||||
destroyOnClose
|
||||
>
|
||||
{/* 搜索 */}
|
||||
<Input
|
||||
allowClear
|
||||
prefix={<SearchOutlined style={{ color: "#9ca3af" }} />}
|
||||
placeholder="搜索标题、内容或标签"
|
||||
value={searchInput}
|
||||
onChange={(e) => setSearchInput(e.target.value)}
|
||||
style={{ marginBottom: 12 }}
|
||||
/>
|
||||
|
||||
{/* 标签筛选 */}
|
||||
<div style={{ display: "flex", flexWrap: "wrap", gap: 8, marginBottom: 12 }}>
|
||||
{SCRIPT_TABS.map((t) => {
|
||||
const active = activeTag === t.key
|
||||
return (
|
||||
<Tag
|
||||
key={t.key}
|
||||
onClick={() => setActiveTag(t.key)}
|
||||
style={{
|
||||
cursor: "pointer",
|
||||
padding: "4px 14px",
|
||||
borderRadius: 16,
|
||||
border: active ? "1px solid #7c3aed" : "1px solid #e5e7eb",
|
||||
background: active ? "#ede9fe" : "#fff",
|
||||
color: active ? "#7c3aed" : "#4b5563",
|
||||
margin: 0,
|
||||
fontSize: 13,
|
||||
}}
|
||||
>
|
||||
{t.label}
|
||||
</Tag>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
|
||||
{/* 数量统计 */}
|
||||
<div style={{ fontSize: 12, color: "#6b7280", marginBottom: 8 }}>
|
||||
共 {data?.total ?? scripts.length} 条文案
|
||||
</div>
|
||||
|
||||
{/* 卡片列表 */}
|
||||
<div style={{ maxHeight: 420, overflowY: "auto", paddingRight: 4 }}>
|
||||
{isLoading ? (
|
||||
<div style={{ textAlign: "center", padding: "40px 0" }}>
|
||||
<Spin />
|
||||
</div>
|
||||
) : scripts.length === 0 ? (
|
||||
<div style={{ textAlign: "center", padding: "40px 0", color: "#9ca3af" }}>
|
||||
暂无匹配文案
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ display: "flex", flexDirection: "column", gap: 10 }}>
|
||||
{scripts.map((s) => {
|
||||
const isSel = innerSelected === s.id
|
||||
const preview = (s.content || "").replace(/\s+/g, " ").slice(0, 80)
|
||||
return (
|
||||
<div
|
||||
key={s.id}
|
||||
onClick={() => setInnerSelected(s.id)}
|
||||
style={{
|
||||
padding: 14,
|
||||
borderRadius: 8,
|
||||
border: isSel ? "2px solid #7c3aed" : "1px solid #e5e7eb",
|
||||
background: isSel ? "#faf5ff" : "#fff",
|
||||
cursor: "pointer",
|
||||
transition: "all 0.2s",
|
||||
position: "relative",
|
||||
}}
|
||||
>
|
||||
{isSel && (
|
||||
<CheckCircleFilled
|
||||
style={{
|
||||
position: "absolute",
|
||||
top: 12,
|
||||
right: 12,
|
||||
color: "#7c3aed",
|
||||
fontSize: 18,
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
<div
|
||||
style={{
|
||||
fontSize: 14,
|
||||
fontWeight: 600,
|
||||
color: isSel ? "#6d28d9" : "#111",
|
||||
marginBottom: 4,
|
||||
paddingRight: 24,
|
||||
}}
|
||||
>
|
||||
{highlight(s.title || "未命名", debouncedKw)}
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 12,
|
||||
color: "#6b7280",
|
||||
lineHeight: 1.6,
|
||||
marginBottom: 8,
|
||||
}}
|
||||
>
|
||||
{highlight(preview + ((s.content || "").length > 80 ? "..." : ""), debouncedKw)}
|
||||
</div>
|
||||
{s.tags && s.tags.length > 0 && (
|
||||
<div style={{ display: "flex", gap: 4, flexWrap: "wrap" }}>
|
||||
{s.tags.slice(0, 5).map((tg) => (
|
||||
<Tag
|
||||
key={tg}
|
||||
style={{
|
||||
margin: 0,
|
||||
fontSize: 11,
|
||||
padding: "1px 8px",
|
||||
borderRadius: 10,
|
||||
background: "#f3f4f6",
|
||||
border: "none",
|
||||
color: "#6b7280",
|
||||
}}
|
||||
>
|
||||
{tg}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default ScriptSelectModal
|
||||
@@ -0,0 +1,264 @@
|
||||
/**
|
||||
* Step 1 选择剪辑模式 + 生成设置(#1970 新流程第一步)
|
||||
* - 剪辑模式:🎲随机混剪 / 📖叙事剪辑,二选一,选中紫底紫框
|
||||
* - 生成设置:生成数量(-/+ 1-10 默认1)、视频比例(9:16/16:9 默认9:16)、智能降重开关(默认开)
|
||||
*/
|
||||
import React from "react"
|
||||
import { MinusOutlined, PlusOutlined } from "@ant-design/icons"
|
||||
|
||||
export type EditMode = "random" | "narrative"
|
||||
|
||||
interface Step1EditModeProps {
|
||||
editMode: EditMode
|
||||
onEditModeChange: (mode: EditMode) => void
|
||||
/** 生成数量(1-10,默认1) */
|
||||
previewCount: number
|
||||
onPreviewCountChange: (n: number) => void
|
||||
/** 视频比例 */
|
||||
videoRatio: "9:16" | "16:9"
|
||||
onVideoRatioChange: (ratio: "9:16" | "16:9") => void
|
||||
/** 智能降重开关(默认 true) */
|
||||
dedupEnabled: boolean
|
||||
onDedupEnabledChange: (v: boolean) => void
|
||||
}
|
||||
|
||||
const PURPLE = "#7c3aed"
|
||||
const PURPLE_BG = "linear-gradient(135deg, #ede9fe, #ddd6fe)"
|
||||
const PURPLE_BORDER = "2px solid #7c3aed"
|
||||
|
||||
const MODE_CARDS: Array<{
|
||||
key: EditMode
|
||||
emoji: string
|
||||
title: string
|
||||
desc: string
|
||||
features: string[]
|
||||
}> = [
|
||||
{
|
||||
key: "random",
|
||||
emoji: "🎲",
|
||||
title: "随机混剪",
|
||||
desc: "根据配音时长随机抽取素材片段,灵活组合",
|
||||
features: ["随机抽帧组合", "每次画面不同", "适合批量生成"],
|
||||
},
|
||||
{
|
||||
key: "narrative",
|
||||
emoji: "📖",
|
||||
title: "叙事剪辑",
|
||||
desc: "按文案内容匹配相关画面,有逻辑组织镜头",
|
||||
features: ["画面匹配文案", "叙事感更强", "需要素材标签"],
|
||||
},
|
||||
]
|
||||
|
||||
const Step1EditMode: React.FC<Step1EditModeProps> = ({
|
||||
editMode,
|
||||
onEditModeChange,
|
||||
previewCount,
|
||||
onPreviewCountChange,
|
||||
videoRatio,
|
||||
onVideoRatioChange,
|
||||
dedupEnabled,
|
||||
onDedupEnabledChange,
|
||||
}) => {
|
||||
return (
|
||||
<div className="xx-form-section">
|
||||
<h3>🎬 选择剪辑模式</h3>
|
||||
<p style={{ color: "#666", fontSize: 14, marginBottom: 16 }}>
|
||||
选择适合您的剪辑方式,后续流程会根据模式自动调整
|
||||
</p>
|
||||
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(auto-fit, minmax(240px, 1fr))",
|
||||
gap: 16,
|
||||
marginBottom: 24,
|
||||
}}
|
||||
>
|
||||
{MODE_CARDS.map((card) => {
|
||||
const selected = editMode === card.key
|
||||
return (
|
||||
<div
|
||||
key={card.key}
|
||||
onClick={() => onEditModeChange(card.key)}
|
||||
style={{
|
||||
padding: 20,
|
||||
borderRadius: 12,
|
||||
border: selected ? PURPLE_BORDER : "1px solid #e5e7eb",
|
||||
background: selected ? PURPLE_BG : "#fff",
|
||||
cursor: "pointer",
|
||||
transition: "all 0.2s",
|
||||
}}
|
||||
>
|
||||
<div style={{ fontSize: 36, marginBottom: 8 }}>{card.emoji}</div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 18,
|
||||
fontWeight: 600,
|
||||
color: selected ? PURPLE : "#111",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
{card.title}
|
||||
</div>
|
||||
<div style={{ fontSize: 13, color: "#666", marginBottom: 12 }}>{card.desc}</div>
|
||||
<div style={{ display: "flex", flexDirection: "column", gap: 4 }}>
|
||||
{card.features.map((f) => (
|
||||
<div key={f} style={{ fontSize: 12, color: selected ? "#6d28d9" : "#6b7280" }}>
|
||||
✅ {f}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
|
||||
<h3 style={{ marginTop: 8 }}>⚙️ 生成设置</h3>
|
||||
|
||||
<div className="xx-form-field" style={{ marginTop: 12 }}>
|
||||
<label>生成数量</label>
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 12 }}>
|
||||
<div
|
||||
style={{
|
||||
display: "inline-flex",
|
||||
alignItems: "center",
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 8,
|
||||
overflow: "hidden",
|
||||
background: "#fff",
|
||||
}}
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => onPreviewCountChange(Math.max(1, previewCount - 1))}
|
||||
disabled={previewCount <= 1}
|
||||
style={{
|
||||
width: 36,
|
||||
height: 36,
|
||||
border: "none",
|
||||
background: "transparent",
|
||||
cursor: previewCount <= 1 ? "not-allowed" : "pointer",
|
||||
color: previewCount <= 1 ? "#d1d5db" : "#374151",
|
||||
fontSize: 16,
|
||||
}}
|
||||
>
|
||||
<MinusOutlined />
|
||||
</button>
|
||||
<span
|
||||
style={{
|
||||
minWidth: 40,
|
||||
textAlign: "center",
|
||||
fontSize: 16,
|
||||
fontWeight: 600,
|
||||
color: "#111",
|
||||
}}
|
||||
>
|
||||
{previewCount}
|
||||
</span>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => onPreviewCountChange(Math.min(10, previewCount + 1))}
|
||||
disabled={previewCount >= 10}
|
||||
style={{
|
||||
width: 36,
|
||||
height: 36,
|
||||
border: "none",
|
||||
background: "transparent",
|
||||
cursor: previewCount >= 10 ? "not-allowed" : "pointer",
|
||||
color: previewCount >= 10 ? "#d1d5db" : "#374151",
|
||||
fontSize: 16,
|
||||
}}
|
||||
>
|
||||
<PlusOutlined />
|
||||
</button>
|
||||
</div>
|
||||
<span style={{ fontSize: 12, color: "#6b7280" }}>最多一次生成 10 个</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="xx-form-field" style={{ marginTop: 16 }}>
|
||||
<label>视频比例</label>
|
||||
<div style={{ display: "flex", gap: 12, marginTop: 4 }}>
|
||||
{[
|
||||
{ key: "9:16" as const, emoji: "📱", label: "竖屏 9:16" },
|
||||
{ key: "16:9" as const, emoji: "🖥️", label: "横屏 16:9" },
|
||||
].map((opt) => {
|
||||
const selected = videoRatio === opt.key
|
||||
return (
|
||||
<button
|
||||
key={opt.key}
|
||||
type="button"
|
||||
onClick={() => onVideoRatioChange(opt.key)}
|
||||
style={{
|
||||
padding: "10px 20px",
|
||||
borderRadius: 8,
|
||||
border: selected ? PURPLE_BORDER : "1px solid #e5e7eb",
|
||||
background: selected ? PURPLE_BG : "#fff",
|
||||
color: selected ? PURPLE : "#374151",
|
||||
cursor: "pointer",
|
||||
fontSize: 14,
|
||||
fontWeight: selected ? 600 : 400,
|
||||
transition: "all 0.2s",
|
||||
}}
|
||||
>
|
||||
{opt.emoji} {opt.label}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
className="xx-form-field"
|
||||
style={{
|
||||
marginTop: 16,
|
||||
padding: "12px 16px",
|
||||
background: "#f9fafb",
|
||||
borderRadius: 8,
|
||||
}}
|
||||
>
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 8 }}>
|
||||
<span style={{ fontSize: 14, fontWeight: 500, color: "#111" }}>
|
||||
🎯 智能降重 {dedupEnabled ? "已开启" : "已关闭"}
|
||||
</span>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => onDedupEnabledChange(!dedupEnabled)}
|
||||
style={{
|
||||
width: 44,
|
||||
height: 24,
|
||||
borderRadius: 12,
|
||||
border: "none",
|
||||
background: dedupEnabled ? PURPLE : "#d1d5db",
|
||||
position: "relative",
|
||||
cursor: "pointer",
|
||||
transition: "background 0.2s",
|
||||
padding: 0,
|
||||
flexShrink: 0,
|
||||
}}
|
||||
aria-label="toggle dedup"
|
||||
>
|
||||
<span
|
||||
style={{
|
||||
position: "absolute",
|
||||
top: 2,
|
||||
left: dedupEnabled ? 22 : 2,
|
||||
width: 20,
|
||||
height: 20,
|
||||
borderRadius: "50%",
|
||||
background: "#fff",
|
||||
transition: "left 0.2s",
|
||||
boxShadow: "0 1px 3px rgba(0,0,0,0.2)",
|
||||
}}
|
||||
/>
|
||||
</button>
|
||||
</div>
|
||||
<div style={{ fontSize: 12, color: "#6b7280", marginTop: 4 }}>
|
||||
自动对画面做微调,避免查重不过
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default Step1EditMode
|
||||
@@ -11,6 +11,8 @@ import SmartMatchInput from "./material/SmartMatchInput"
|
||||
import SmartMatchResults from "./material/SmartMatchResults"
|
||||
|
||||
interface Step2MaterialSelectProps {
|
||||
/** 剪辑模式:random 随机混剪 / narrative 叙事剪辑(#1970) */
|
||||
editMode?: "random" | "narrative"
|
||||
materialMode: "manual" | "auto"
|
||||
onMaterialModeChange: (mode: "manual" | "auto") => void
|
||||
selectedMaterials: string[]
|
||||
@@ -43,6 +45,27 @@ const Step2MaterialSelect: React.FC<Step2MaterialSelectProps> = (props) => {
|
||||
<div className="xx-form-section">
|
||||
<h3>📦 选择素材</h3>
|
||||
|
||||
{/* 叙事剪辑:AI 智能匹配提示卡(#1970) */}
|
||||
{props.editMode === "narrative" && (
|
||||
<div
|
||||
style={{
|
||||
marginTop: 12,
|
||||
padding: "12px 16px",
|
||||
background: "linear-gradient(135deg,#ede9fe,#f5f3ff)",
|
||||
border: "1px solid #c4b5fd",
|
||||
borderRadius: 8,
|
||||
fontSize: 13,
|
||||
color: "#5b21b6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 8,
|
||||
}}
|
||||
>
|
||||
<span style={{ fontSize: 18 }}>🤖</span>
|
||||
<span>AI智能匹配:系统将根据您的文案内容,从素材库自动匹配合适的视频片段</span>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 片段数量(#1899) */}
|
||||
<div className="xx-form-field" style={{ marginTop: 12 }}>
|
||||
<label>片段数量</label>
|
||||
|
||||
@@ -0,0 +1,491 @@
|
||||
/**
|
||||
* 叙事剪辑 — TTS 音色选择 + 合成配音弹窗(#1970)
|
||||
* - Tabs:✨系统音色 / 🎙️我的克隆音色
|
||||
* - 2列音色卡片(头像emoji+名称+描述+标签+▶试听+选中✓)
|
||||
* - 底部:取消 / 🎧 合成配音(主按钮,必须选音色才能点)
|
||||
* - 合成中:紫色 spinner + "正在合成配音..." + "请稍候,通常需要10-30秒"
|
||||
* - 合成成功:保存到配音库并回调(voiceAssetId + ttsVoiceId + ttsVoiceSource)
|
||||
*
|
||||
* 复用现有 /api/tts 的 synthesizeSpeech + 轮询 getTTSJobStatus 逻辑;
|
||||
* 不直接复用 TtsModal(它是页面配音弹窗,含文本输入/语速/情感等字段,叙事模式文本来自文案)。
|
||||
*/
|
||||
import React, { useState, useEffect, useMemo, useRef, useCallback } from "react"
|
||||
import { Modal, Tabs, Spin, message } from "antd"
|
||||
import { CheckCircleFilled, SoundOutlined } from "@ant-design/icons"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { fetchPresetVoices } from "@/api/voices"
|
||||
import { getVoiceClones } from "@/api/voice-clone"
|
||||
import { synthesizeSpeech, getTTSJobStatus, saveTtsToLibrary } from "@/api/tts"
|
||||
import type { PresetVoiceItem } from "@/api/voices"
|
||||
import type { VoiceClone } from "@/api/voice-clone"
|
||||
import { VOICE_GENDER_ICON } from "../constants"
|
||||
|
||||
interface TtsVoiceModalProps {
|
||||
open: boolean
|
||||
/** 需要合成的文本(来自选中的文案 content) */
|
||||
scriptText: string
|
||||
scriptTitle: string
|
||||
onCancel: () => void
|
||||
/** 合成成功回调:asset_id 为保存到配音库后的素材ID */
|
||||
onSynthesized: (payload: {
|
||||
voiceAssetId: string
|
||||
ttsVoiceId: string
|
||||
ttsVoiceSource: "preset" | "clone"
|
||||
}) => void
|
||||
}
|
||||
|
||||
type TtsSynthStatus = "idle" | "synthesizing" | "saving" | "done" | "error"
|
||||
|
||||
const TtsVoiceModal: React.FC<TtsVoiceModalProps> = ({
|
||||
open,
|
||||
scriptText,
|
||||
scriptTitle,
|
||||
onCancel,
|
||||
onSynthesized,
|
||||
}) => {
|
||||
const [activeTab, setActiveTab] = useState<"preset" | "clone">("preset")
|
||||
const [selectedVoiceId, setSelectedVoiceId] = useState<string>("")
|
||||
const [status, setStatus] = useState<TtsSynthStatus>("idle")
|
||||
const [error, setError] = useState<string | null>(null)
|
||||
const [previewingId, setPreviewingId] = useState<string | null>(null)
|
||||
const audioRef = useRef<HTMLAudioElement | null>(null)
|
||||
const timerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
|
||||
/* 系统音色 */
|
||||
const { data: presetData } = useQuery({
|
||||
queryKey: ["preset-voices", "modal"],
|
||||
queryFn: fetchPresetVoices,
|
||||
enabled: open,
|
||||
})
|
||||
const presetVoices: PresetVoiceItem[] = useMemo(() => presetData?.items ?? [], [presetData])
|
||||
|
||||
/* 克隆音色(仅 ready 状态可用) */
|
||||
const { data: cloneListRaw = [] } = useQuery({
|
||||
queryKey: ["voice-clones", "ready"],
|
||||
queryFn: () => getVoiceClones({ status: "ready" }),
|
||||
enabled: open,
|
||||
})
|
||||
const cloneVoices: VoiceClone[] = useMemo(
|
||||
() => cloneListRaw.filter((v: VoiceClone) => v.status === "ready"),
|
||||
[cloneListRaw],
|
||||
)
|
||||
|
||||
/* 打开时重置状态 */
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setSelectedVoiceId("")
|
||||
setStatus("idle")
|
||||
setError(null)
|
||||
setActiveTab("preset")
|
||||
} else {
|
||||
if (timerRef.current) {
|
||||
clearInterval(timerRef.current)
|
||||
timerRef.current = null
|
||||
}
|
||||
if (audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
audioRef.current = null
|
||||
}
|
||||
setPreviewingId(null)
|
||||
}
|
||||
return () => {
|
||||
if (timerRef.current) clearInterval(timerRef.current)
|
||||
}
|
||||
}, [open])
|
||||
|
||||
const handlePreview = useCallback(
|
||||
(voiceId: string, previewUrl: string | null | undefined) => {
|
||||
if (!previewUrl) {
|
||||
message.info("该音色暂无试听音频")
|
||||
return
|
||||
}
|
||||
if (previewingId === voiceId && audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
setPreviewingId(null)
|
||||
return
|
||||
}
|
||||
if (audioRef.current) audioRef.current.pause()
|
||||
const a = new Audio(previewUrl)
|
||||
audioRef.current = a
|
||||
setPreviewingId(voiceId)
|
||||
a.onended = () => {
|
||||
setPreviewingId(null)
|
||||
audioRef.current = null
|
||||
}
|
||||
a.play().catch(() => {
|
||||
setPreviewingId(null)
|
||||
audioRef.current = null
|
||||
})
|
||||
},
|
||||
[previewingId],
|
||||
)
|
||||
|
||||
const textToSynth = useMemo(() => {
|
||||
// 文案内容取首段(过长会被 TTS 截断,保持和用户感知一致)
|
||||
const t = (scriptText || "").trim()
|
||||
return t.length > 500 ? t.slice(0, 500) : t
|
||||
}, [scriptText])
|
||||
|
||||
const handleSynthesize = useCallback(async () => {
|
||||
if (!selectedVoiceId) {
|
||||
message.warning("请先选择一个音色")
|
||||
return
|
||||
}
|
||||
if (!textToSynth) {
|
||||
message.warning("文案内容为空,无法合成")
|
||||
return
|
||||
}
|
||||
setStatus("synthesizing")
|
||||
setError(null)
|
||||
try {
|
||||
const isClone = activeTab === "clone"
|
||||
const payload: Record<string, unknown> = {
|
||||
text: textToSynth,
|
||||
speed: 1.0,
|
||||
language: "zh-CN",
|
||||
}
|
||||
if (isClone) {
|
||||
payload.voice_clone_profile_id = selectedVoiceId
|
||||
} else {
|
||||
payload.voice_id = selectedVoiceId
|
||||
}
|
||||
const resp = await synthesizeSpeech(
|
||||
payload as unknown as Parameters<typeof synthesizeSpeech>[0],
|
||||
)
|
||||
const jobId = resp.job_id
|
||||
|
||||
await new Promise<void>((resolve, reject) => {
|
||||
timerRef.current = setInterval(async () => {
|
||||
try {
|
||||
const job = await getTTSJobStatus(jobId)
|
||||
if (job.status === "completed") {
|
||||
if (timerRef.current) clearInterval(timerRef.current)
|
||||
timerRef.current = null
|
||||
resolve()
|
||||
} else if (job.status === "failed") {
|
||||
if (timerRef.current) clearInterval(timerRef.current)
|
||||
timerRef.current = null
|
||||
reject(new Error(job.error_message || "合成失败"))
|
||||
}
|
||||
} catch (e) {
|
||||
if (timerRef.current) clearInterval(timerRef.current)
|
||||
timerRef.current = null
|
||||
reject(e)
|
||||
}
|
||||
}, 2000)
|
||||
})
|
||||
|
||||
// 保存到配音库
|
||||
setStatus("saving")
|
||||
await saveTtsToLibrary(jobId, { name: scriptTitle?.slice(0, 30) || "AI合成配音" })
|
||||
setStatus("done")
|
||||
|
||||
// 合成成功后回调;voiceAssetId 由后端在保存时产出,这里用 ttsVoiceId 占位,
|
||||
// 父流程会在下一次 asset 列表刷新后重新选取;前端直接以 ttsVoiceId 为 key 传给后端
|
||||
// (叙事模式后端通过 script_id + tts_voice_id 自行再合成,不依赖 asset_id)。
|
||||
onSynthesized({
|
||||
voiceAssetId: jobId,
|
||||
ttsVoiceId: selectedVoiceId,
|
||||
ttsVoiceSource: isClone ? "clone" : "preset",
|
||||
})
|
||||
} catch (err: unknown) {
|
||||
setStatus("error")
|
||||
const msg = err instanceof Error ? err.message : "合成失败,请稍后重试"
|
||||
setError(msg)
|
||||
}
|
||||
}, [selectedVoiceId, textToSynth, activeTab, scriptTitle, onSynthesized])
|
||||
|
||||
const renderVoiceCard = (v: {
|
||||
id: string
|
||||
name: string
|
||||
description?: string
|
||||
gender?: string
|
||||
tags?: string[]
|
||||
preview_url?: string | null
|
||||
}) => {
|
||||
const isSel = selectedVoiceId === v.id
|
||||
const isPlaying = previewingId === v.id
|
||||
const emoji = v.gender ? (VOICE_GENDER_ICON[v.gender] ?? "🎤") : "🎤"
|
||||
return (
|
||||
<div
|
||||
key={v.id}
|
||||
onClick={() => setSelectedVoiceId(v.id)}
|
||||
style={{
|
||||
padding: 12,
|
||||
borderRadius: 8,
|
||||
border: isSel ? "2px solid #7c3aed" : "1px solid #e5e7eb",
|
||||
background: isSel ? "#faf5ff" : "#fff",
|
||||
cursor: "pointer",
|
||||
transition: "all 0.2s",
|
||||
position: "relative",
|
||||
}}
|
||||
>
|
||||
{isSel && (
|
||||
<CheckCircleFilled
|
||||
style={{
|
||||
position: "absolute",
|
||||
top: 10,
|
||||
right: 10,
|
||||
color: "#7c3aed",
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 10, marginBottom: 8 }}>
|
||||
<div
|
||||
style={{
|
||||
width: 36,
|
||||
height: 36,
|
||||
borderRadius: "50%",
|
||||
background: isSel ? "linear-gradient(135deg,#7c3aed,#a78bfa)" : "#f3f4f6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
fontSize: 18,
|
||||
}}
|
||||
>
|
||||
{emoji}
|
||||
</div>
|
||||
<div style={{ flex: 1, minWidth: 0 }}>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 14,
|
||||
fontWeight: 600,
|
||||
color: isSel ? "#6d28d9" : "#111",
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
}}
|
||||
>
|
||||
{v.name}
|
||||
</div>
|
||||
{v.description && (
|
||||
<div
|
||||
style={{
|
||||
fontSize: 11,
|
||||
color: "#6b7280",
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
}}
|
||||
>
|
||||
{v.description}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{v.preview_url && (
|
||||
<button
|
||||
type="button"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
handlePreview(v.id, v.preview_url)
|
||||
}}
|
||||
style={{
|
||||
width: 28,
|
||||
height: 28,
|
||||
borderRadius: "50%",
|
||||
border: "none",
|
||||
background: isPlaying ? "#ef4444" : "#7c3aed",
|
||||
color: "#fff",
|
||||
cursor: "pointer",
|
||||
fontSize: 11,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
}}
|
||||
>
|
||||
<SoundOutlined />
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
{v.tags && v.tags.length > 0 && (
|
||||
<div style={{ display: "flex", gap: 4, flexWrap: "wrap" }}>
|
||||
{v.tags.slice(0, 3).map((tg) => (
|
||||
<span
|
||||
key={tg}
|
||||
style={{
|
||||
fontSize: 10,
|
||||
padding: "1px 6px",
|
||||
borderRadius: 8,
|
||||
background: "#f3f4f6",
|
||||
color: "#6b7280",
|
||||
}}
|
||||
>
|
||||
{tg}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
/* 合成中 loading 覆盖层 */
|
||||
const renderSynthOverlay = () => {
|
||||
if (status !== "synthesizing" && status !== "saving") return null
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
position: "absolute",
|
||||
inset: 0,
|
||||
background: "rgba(255,255,255,0.92)",
|
||||
zIndex: 10,
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
gap: 12,
|
||||
borderRadius: 8,
|
||||
}}
|
||||
>
|
||||
<Spin size="large" style={{ color: "#7c3aed" }} />
|
||||
<div style={{ fontSize: 16, fontWeight: 600, color: "#6d28d9" }}>
|
||||
{status === "synthesizing" ? "正在合成配音..." : "正在保存到配音库..."}
|
||||
</div>
|
||||
<div style={{ fontSize: 12, color: "#6b7280" }}>请稍候,通常需要 10-30 秒</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="🎙️ 合成配音"
|
||||
open={open}
|
||||
onCancel={status === "synthesizing" || status === "saving" ? undefined : onCancel}
|
||||
cancelText="取消"
|
||||
okText="🎧 合成配音"
|
||||
okButtonProps={{
|
||||
disabled: !selectedVoiceId || status === "synthesizing" || status === "saving",
|
||||
style: { background: "#7c3aed" },
|
||||
}}
|
||||
onOk={handleSynthesize}
|
||||
width={680}
|
||||
destroyOnClose
|
||||
confirmLoading={status === "synthesizing" || status === "saving"}
|
||||
>
|
||||
<div style={{ position: "relative" }}>
|
||||
{error && (
|
||||
<div
|
||||
style={{
|
||||
padding: "10px 12px",
|
||||
background: "#fef2f2",
|
||||
border: "1px solid #fecaca",
|
||||
color: "#b91c1c",
|
||||
borderRadius: 6,
|
||||
fontSize: 13,
|
||||
marginBottom: 12,
|
||||
}}
|
||||
>
|
||||
{error}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div
|
||||
style={{
|
||||
fontSize: 12,
|
||||
color: "#6b7280",
|
||||
marginBottom: 12,
|
||||
padding: "8px 12px",
|
||||
background: "#f9fafb",
|
||||
borderRadius: 6,
|
||||
}}
|
||||
>
|
||||
将根据文案《{scriptTitle?.slice(0, 30) || "所选文案"}》合成配音,文本长度:
|
||||
{textToSynth.length} 字
|
||||
</div>
|
||||
|
||||
<Tabs
|
||||
activeKey={activeTab}
|
||||
onChange={(k) => {
|
||||
setActiveTab(k as "preset" | "clone")
|
||||
setSelectedVoiceId("")
|
||||
}}
|
||||
items={[
|
||||
{
|
||||
key: "preset",
|
||||
label: "✨ 系统音色",
|
||||
children: (
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "1fr 1fr",
|
||||
gap: 10,
|
||||
maxHeight: 420,
|
||||
overflowY: "auto",
|
||||
paddingRight: 4,
|
||||
}}
|
||||
>
|
||||
{presetVoices.length === 0 ? (
|
||||
<div
|
||||
style={{
|
||||
gridColumn: "1/-1",
|
||||
textAlign: "center",
|
||||
padding: 30,
|
||||
color: "#9ca3af",
|
||||
}}
|
||||
>
|
||||
正在加载系统音色...
|
||||
</div>
|
||||
) : (
|
||||
presetVoices.map((v) =>
|
||||
renderVoiceCard({
|
||||
id: v.voice_id,
|
||||
name: v.name,
|
||||
description: v.description,
|
||||
gender: v.gender,
|
||||
tags: v.tags,
|
||||
preview_url: v.preview_url,
|
||||
}),
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "clone",
|
||||
label: "🎙️ 我的克隆音色",
|
||||
children: (
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "1fr 1fr",
|
||||
gap: 10,
|
||||
maxHeight: 420,
|
||||
overflowY: "auto",
|
||||
paddingRight: 4,
|
||||
}}
|
||||
>
|
||||
{cloneVoices.length === 0 ? (
|
||||
<div
|
||||
style={{
|
||||
gridColumn: "1/-1",
|
||||
textAlign: "center",
|
||||
padding: 30,
|
||||
color: "#9ca3af",
|
||||
}}
|
||||
>
|
||||
暂无就绪的克隆音色,请先在配音库完成音色克隆
|
||||
</div>
|
||||
) : (
|
||||
cloneVoices.map((v) =>
|
||||
renderVoiceCard({
|
||||
id: v.id,
|
||||
name: v.name,
|
||||
description: v.description,
|
||||
gender: "neutral",
|
||||
tags: ["克隆"],
|
||||
preview_url: v.sample_url || null,
|
||||
}),
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
),
|
||||
},
|
||||
]}
|
||||
/>
|
||||
{renderSynthOverlay()}
|
||||
</div>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default TtsVoiceModal
|
||||
@@ -0,0 +1,241 @@
|
||||
/**
|
||||
* 随机混剪 — 配音选择弹窗(#1970)
|
||||
* 内容复用 Step5VoiceSelect 的配音库音频卡片(图标+文件名+时长/大小+▶试听),
|
||||
* 无 TTS / 克隆音色入口;确认后进入 Step2。
|
||||
*/
|
||||
import React from "react"
|
||||
import { Modal } from "antd"
|
||||
import { AudioOutlined } from "@ant-design/icons"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { useState, useRef, useCallback } from "react"
|
||||
import { getAssetsByKind } from "@/api/assets"
|
||||
import type { AssetItem } from "@/api/assets"
|
||||
|
||||
interface VoiceSelectModalProps {
|
||||
open: boolean
|
||||
selectedVoice: string
|
||||
onCancel: () => void
|
||||
onConfirm: (voiceAssetId: string) => void
|
||||
}
|
||||
|
||||
const getDuration = (item: AssetItem): number =>
|
||||
item.duration ?? (item.metadata?.duration as number) ?? 0
|
||||
const getFileSize = (item: AssetItem): number =>
|
||||
item.file_size ?? (item.metadata?.file_size as number) ?? 0
|
||||
const isAiVoice = (item: AssetItem): boolean => {
|
||||
const d = getDuration(item)
|
||||
const s = getFileSize(item)
|
||||
return (!d || d <= 0) && (!s || s <= 0)
|
||||
}
|
||||
const fmtDur = (s?: number): string => {
|
||||
if (!s || s <= 0) return "时长未知"
|
||||
return `${s.toFixed(1)}秒`
|
||||
}
|
||||
const fmtSize = (b?: number): string => {
|
||||
if (!b || b <= 0) return "未知"
|
||||
if (b < 1024) return `${b} B`
|
||||
if (b < 1024 * 1024) return `${(b / 1024).toFixed(1)} KB`
|
||||
if (b < 1024 * 1024 * 1024) return `${(b / (1024 * 1024)).toFixed(1)} MB`
|
||||
return `${(b / (1024 * 1024 * 1024)).toFixed(1)} GB`
|
||||
}
|
||||
|
||||
const VoiceSelectModal: React.FC<VoiceSelectModalProps> = ({
|
||||
open,
|
||||
selectedVoice,
|
||||
onCancel,
|
||||
onConfirm,
|
||||
}) => {
|
||||
const navigate = useNavigate()
|
||||
const [innerSelected, setInnerSelected] = React.useState(selectedVoice)
|
||||
const [playingId, setPlayingId] = useState<string | null>(null)
|
||||
const audioRef = useRef<HTMLAudioElement | null>(null)
|
||||
|
||||
React.useEffect(() => {
|
||||
if (open) setInnerSelected(selectedVoice)
|
||||
}, [open, selectedVoice])
|
||||
|
||||
const { data: materials = [], isLoading } = useQuery({
|
||||
queryKey: ["assets", "voice", "modal"],
|
||||
queryFn: () => getAssetsByKind("voice", { limit: 50 }),
|
||||
enabled: open,
|
||||
})
|
||||
|
||||
const togglePlay = useCallback(
|
||||
(item: AssetItem) => {
|
||||
if (playingId === item.id && audioRef.current) {
|
||||
audioRef.current.pause()
|
||||
setPlayingId(null)
|
||||
return
|
||||
}
|
||||
if (audioRef.current) audioRef.current.pause()
|
||||
if (!item.file_url) return
|
||||
const audio = new Audio(item.file_url)
|
||||
audioRef.current = audio
|
||||
setPlayingId(item.id)
|
||||
audio.onended = () => {
|
||||
setPlayingId(null)
|
||||
audioRef.current = null
|
||||
}
|
||||
audio.play().catch(() => {
|
||||
setPlayingId(null)
|
||||
audioRef.current = null
|
||||
})
|
||||
},
|
||||
[playingId],
|
||||
)
|
||||
|
||||
const handleGoUpload = () => navigate("/app/voices?tab=material&upload=1")
|
||||
|
||||
const handleConfirm = () => {
|
||||
if (!innerSelected) return
|
||||
onConfirm(innerSelected)
|
||||
}
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="🎙️ 选择配音"
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={handleConfirm}
|
||||
okText="确认选择"
|
||||
cancelText="取消"
|
||||
okButtonProps={{ disabled: !innerSelected, style: { background: "#7c3aed" } }}
|
||||
width={720}
|
||||
destroyOnClose
|
||||
>
|
||||
<p style={{ color: "#666", fontSize: 13, marginBottom: 12 }}>
|
||||
从配音库中选择已上传的音频素材,点击 ▶ 可试听
|
||||
</p>
|
||||
{isLoading ? (
|
||||
<div style={{ textAlign: "center", padding: "40px 0", color: "#999" }}>加载中...</div>
|
||||
) : materials.length === 0 ? (
|
||||
<div style={{ textAlign: "center", padding: "40px 0", color: "#999" }}>
|
||||
<AudioOutlined style={{ fontSize: 48, color: "#d9d9d9", marginBottom: 12 }} />
|
||||
<p style={{ marginBottom: 12 }}>暂无配音素材</p>
|
||||
<button
|
||||
type="button"
|
||||
onClick={handleGoUpload}
|
||||
style={{
|
||||
padding: "8px 20px",
|
||||
background: "#7c3aed",
|
||||
color: "#fff",
|
||||
border: "none",
|
||||
borderRadius: 6,
|
||||
cursor: "pointer",
|
||||
}}
|
||||
>
|
||||
去配音库上传
|
||||
</button>
|
||||
</div>
|
||||
) : (
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(auto-fill, minmax(200px, 1fr))",
|
||||
gap: 12,
|
||||
maxHeight: 460,
|
||||
overflowY: "auto",
|
||||
paddingRight: 4,
|
||||
}}
|
||||
>
|
||||
{materials.map((item) => {
|
||||
const isSel = innerSelected === item.id
|
||||
const isPlaying = playingId === item.id
|
||||
return (
|
||||
<div
|
||||
key={item.id}
|
||||
onClick={() => setInnerSelected(item.id)}
|
||||
style={{
|
||||
padding: 14,
|
||||
borderRadius: 8,
|
||||
border: isSel ? "2px solid #7c3aed" : "1px solid #e8e8e8",
|
||||
background: isSel ? "#ede9fe" : "#fff",
|
||||
cursor: "pointer",
|
||||
transition: "all 0.2s",
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
width: 36,
|
||||
height: 36,
|
||||
borderRadius: 8,
|
||||
background: isSel
|
||||
? "linear-gradient(135deg,#7c3aed,#a78bfa)"
|
||||
: "linear-gradient(135deg,#f0f0f0,#e8e8e8)",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
}}
|
||||
>
|
||||
<AudioOutlined style={{ color: isSel ? "#fff" : "#666" }} />
|
||||
</div>
|
||||
{item.file_url && (
|
||||
<button
|
||||
type="button"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
togglePlay(item)
|
||||
}}
|
||||
style={{
|
||||
width: 30,
|
||||
height: 30,
|
||||
borderRadius: "50%",
|
||||
border: "none",
|
||||
background: isPlaying ? "#ef4444" : "#7c3aed",
|
||||
color: "#fff",
|
||||
cursor: "pointer",
|
||||
fontSize: 12,
|
||||
}}
|
||||
>
|
||||
▶
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
fontWeight: 500,
|
||||
marginTop: 8,
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
color: isSel ? "#6d28d9" : "#333",
|
||||
}}
|
||||
title={item.name}
|
||||
>
|
||||
{item.name}
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "space-between",
|
||||
fontSize: 11,
|
||||
color: "#999",
|
||||
marginTop: 4,
|
||||
}}
|
||||
>
|
||||
{isAiVoice(item) ? (
|
||||
<span style={{ color: "#7c3aed", fontWeight: 500 }}>AI 音色</span>
|
||||
) : (
|
||||
<span>{fmtDur(getDuration(item))}</span>
|
||||
)}
|
||||
<span>{isAiVoice(item) ? "按文本合成" : fmtSize(getFileSize(item))}</span>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default VoiceSelectModal
|
||||
@@ -27,10 +27,10 @@ export const VOICE_GENDER_ICON: Record<string, string> = {
|
||||
neutral: "✨",
|
||||
}
|
||||
|
||||
/* ── 步骤定义(5步,#1899 简化:删除选模板步骤) ── */
|
||||
/* ── 步骤定义(5步,#1970 流程重构:选择模式 → 素材 → 标题 → 确认 → 封面) ── */
|
||||
export const STEPS = [
|
||||
{ key: 1, label: "选择素材" },
|
||||
{ key: 2, label: "选择配音" },
|
||||
{ key: 1, label: "选择模式" },
|
||||
{ key: 2, label: "选择素材" },
|
||||
{ key: 3, label: "选择标题" },
|
||||
{ key: 4, label: "确认生成" },
|
||||
{ key: 5, label: "选择封面" },
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
import type { UseGenerateVideoProps } from "./types"
|
||||
|
||||
/**
|
||||
* 生成前置校验
|
||||
* 生成前置校验(#1970 适配新流程)
|
||||
* - 随机混剪:需选配音(selectedVoice,配音库音频)
|
||||
* - 叙事剪辑:需选文案 + TTS 音色
|
||||
* 返回错误信息,通过则返回 null
|
||||
*/
|
||||
export const validateGenerateInputs = (props: UseGenerateVideoProps): string | null => {
|
||||
@@ -12,19 +14,30 @@ export const validateGenerateInputs = (props: UseGenerateVideoProps): string | n
|
||||
smartSelectedIds,
|
||||
voiceMode,
|
||||
selectedClonedVoice,
|
||||
editMode = "random",
|
||||
selectedScript,
|
||||
ttsVoiceId,
|
||||
selectedVoice,
|
||||
} = props
|
||||
|
||||
// AI 自动选择模式下,标题可以为空(后端会自行生成)
|
||||
if (!titleSettings.aiAutoSelect && !titleSettings.title?.trim()) {
|
||||
return "请先选择或输入标题"
|
||||
}
|
||||
// 无论手动还是自动模式,都必须有素材
|
||||
const materialIds = materialMode === "auto" ? smartSelectedIds || [] : selectedMaterials || []
|
||||
if (materialIds.length === 0) {
|
||||
return materialMode === "auto" ? "AI 未匹配到素材,请手动选择素材后重试" : "请至少选择一个素材"
|
||||
}
|
||||
if (voiceMode === "clone" && !selectedClonedVoice) {
|
||||
return "请先选择一个克隆音色"
|
||||
if (editMode === "narrative") {
|
||||
if (!selectedScript?.id) return "请先选择文案"
|
||||
if (!ttsVoiceId) return "请先合成配音"
|
||||
} else {
|
||||
// 随机混剪:配音库音频
|
||||
if (!selectedVoice && voiceMode !== "clone") {
|
||||
return "请先选择配音"
|
||||
}
|
||||
if (voiceMode === "clone" && !selectedClonedVoice) {
|
||||
return "请先选择一个克隆音色"
|
||||
}
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
@@ -13,7 +13,19 @@ export interface UseGenerateVideoProps {
|
||||
selectedVoice: string
|
||||
selectedClonedVoice: string
|
||||
coverSettings: CoverConfig
|
||||
videoRatio: string
|
||||
videoRatio: "9:16" | "16:9" | string
|
||||
/** #1970 剪辑模式 */
|
||||
editMode?: "random" | "narrative"
|
||||
/** 叙事模式下选中的文案 */
|
||||
selectedScript?: { id: string; title?: string; content?: string } | null
|
||||
/** TTS 音色 ID(叙事模式) */
|
||||
ttsVoiceId?: string
|
||||
/** TTS 音色来源 */
|
||||
ttsVoiceSource?: "preset" | "clone"
|
||||
/** 合成后保存到配音库的 asset id / job id(叙事模式) */
|
||||
ttsVoiceAssetId?: string
|
||||
/** 智能降重开关(默认 true) */
|
||||
dedupEnabled?: boolean
|
||||
style: string
|
||||
duration: number
|
||||
autoSubtitles: boolean
|
||||
|
||||
@@ -12,6 +12,7 @@ import { getEditingTemplates } from "@/api/editing-planner"
|
||||
import type { EditPlanClip } from "@/api/template-editor"
|
||||
import type { CoverConfig } from "../../types/cover"
|
||||
import type { PresetVoiceItem } from "@/api/voices"
|
||||
import type { ScriptItem } from "@/api/scripts"
|
||||
import { DEFAULT_COVER_SETTINGS, DEFAULT_CLIP_COUNT } from "../../constants"
|
||||
import type { TitleSettings } from "../../types"
|
||||
import { usePlanConfigLoader } from "./usePlanConfigLoader"
|
||||
@@ -82,8 +83,28 @@ export interface GenerateFormState {
|
||||
cloneModalOpen: boolean
|
||||
setCloneModalOpen: (open: boolean) => void
|
||||
|
||||
/* ── 剪辑模式(#1970 流程重构)── */
|
||||
editMode: "random" | "narrative"
|
||||
setEditMode: (mode: "random" | "narrative") => void
|
||||
/** 叙事模式下选中的文案 */
|
||||
selectedScript: ScriptItem | null
|
||||
setSelectedScript: (s: ScriptItem | null) => void
|
||||
/** TTS 音色 ID */
|
||||
ttsVoiceId: string
|
||||
setTtsVoiceId: (id: string) => void
|
||||
/** TTS 音色来源:preset 系统 / clone 克隆 */
|
||||
ttsVoiceSource: "preset" | "clone"
|
||||
setTtsVoiceSource: (src: "preset" | "clone") => void
|
||||
/** 合成后配音库 asset id(叙事模式保存到库后获得;随机模式 = selectedVoice) */
|
||||
ttsVoiceAssetId: string
|
||||
setTtsVoiceAssetId: (id: string) => void
|
||||
/** 智能降重开关(默认 true) */
|
||||
dedupEnabled: boolean
|
||||
setDedupEnabled: (v: boolean) => void
|
||||
|
||||
/* 高级设置 */
|
||||
videoRatio: string
|
||||
videoRatio: "9:16" | "16:9" | string
|
||||
setVideoRatio: (r: "9:16" | "16:9") => void
|
||||
duration: number
|
||||
style: string
|
||||
autoSubtitles: boolean
|
||||
@@ -201,13 +222,21 @@ export const useGenerateFormState = (): GenerateFormState => {
|
||||
/* ── 克隆声音弹窗 ── */
|
||||
const [cloneModalOpen, setCloneModalOpen] = useState(false)
|
||||
|
||||
/* ── 高级设置(隐藏但保留) ── */
|
||||
const [videoRatio] = useState("9:16")
|
||||
/* ── 高级设置 ── */
|
||||
const [videoRatio, setVideoRatio] = useState<"9:16" | "16:9">("9:16")
|
||||
const [duration] = useState(30)
|
||||
const [style] = useState("business")
|
||||
const [autoSubtitles] = useState(true)
|
||||
const [bgm] = useState(true)
|
||||
|
||||
/* ── 剪辑模式状态(#1970) ── */
|
||||
const [editMode, setEditMode] = useState<"random" | "narrative">("random")
|
||||
const [selectedScript, setSelectedScript] = useState<ScriptItem | null>(null)
|
||||
const [ttsVoiceId, setTtsVoiceId] = useState<string>("")
|
||||
const [ttsVoiceSource, setTtsVoiceSource] = useState<"preset" | "clone">("preset")
|
||||
const [ttsVoiceAssetId, setTtsVoiceAssetId] = useState<string>("")
|
||||
const [dedupEnabled, setDedupEnabled] = useState<boolean>(true)
|
||||
|
||||
/* ── 预览任务 ID ── */
|
||||
const previewStorageKey = editPlanId
|
||||
? `preview_task_id_${editPlanId}`
|
||||
@@ -274,9 +303,22 @@ export const useGenerateFormState = (): GenerateFormState => {
|
||||
selectedClonedVoice,
|
||||
setSelectedClonedVoice,
|
||||
presetVoices,
|
||||
editMode,
|
||||
setEditMode,
|
||||
selectedScript,
|
||||
setSelectedScript,
|
||||
ttsVoiceId,
|
||||
setTtsVoiceId,
|
||||
ttsVoiceSource,
|
||||
setTtsVoiceSource,
|
||||
ttsVoiceAssetId,
|
||||
setTtsVoiceAssetId,
|
||||
dedupEnabled,
|
||||
setDedupEnabled,
|
||||
cloneModalOpen,
|
||||
setCloneModalOpen,
|
||||
videoRatio,
|
||||
setVideoRatio,
|
||||
duration,
|
||||
style,
|
||||
autoSubtitles,
|
||||
|
||||
@@ -123,6 +123,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { width: outputWidth, height: outputHeight } = calculateResolution(
|
||||
props.videoRatio || "9:16",
|
||||
)
|
||||
const editMode = props.editMode ?? "random"
|
||||
const dedupEnabled = props.dedupEnabled !== false
|
||||
|
||||
const assetIds =
|
||||
props.materialMode === "auto" ? props.smartSelectedIds : props.selectedMaterials
|
||||
@@ -151,10 +153,13 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
|
||||
const coverUrl = props.coverSettings?.thumbnail_url || props.coverSettings?.upload_url || ""
|
||||
|
||||
// #1970:叙事模式下 ttsVoiceId 作为配音 id;随机模式用 selectedVoice
|
||||
const voiceLibraryId =
|
||||
props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
editMode === "narrative"
|
||||
? props.ttsVoiceId || ""
|
||||
: props.voiceMode === "clone"
|
||||
? props.selectedClonedVoice || props.selectedVoice || ""
|
||||
: props.selectedVoice || ""
|
||||
|
||||
/* ── 批量变体数组(长度1=共用,长度=count=独立,空=回退单值) ── */
|
||||
const indexes =
|
||||
@@ -197,6 +202,15 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
custom_title: props.titleSettings?.title || "",
|
||||
duration: props.duration || undefined,
|
||||
video_ratio: props.videoRatio,
|
||||
assembly_mode: editMode,
|
||||
...(editMode === "narrative" && props.selectedScript?.id
|
||||
? {
|
||||
script_id: props.selectedScript.id,
|
||||
tts_voice_id: props.ttsVoiceId || undefined,
|
||||
tts_voice_source: props.ttsVoiceSource || undefined,
|
||||
}
|
||||
: {}),
|
||||
dedup_enabled: dedupEnabled,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(props.selectedVoice && !voiceLibraryId ? { voice_ids: [props.selectedVoice] } : {}),
|
||||
bgm_config: {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { useEffect, useRef } from "react"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { getTitles } from "@/api/titles"
|
||||
// #1894: 标题候选从文案库 scripts[].title 获取,不再调用废弃的 /api/titles
|
||||
import { getScripts } from "@/api/scripts"
|
||||
import type { TitleSettings } from "../../types"
|
||||
import { useAiTitleGenerator } from "./useAiTitleGenerator"
|
||||
import { useTitleStyleUpdaters } from "./useTitleStyleUpdaters"
|
||||
@@ -22,10 +23,14 @@ export function useStep4Title({
|
||||
onTitleSettingsChange,
|
||||
selectedTemplate,
|
||||
}: UseStep4TitleProps) {
|
||||
// 标题库数据
|
||||
// 标题候选(#1894:统一从文案库取 scripts[].title,去重)
|
||||
const { data: userTitles = [] } = useQuery({
|
||||
queryKey: ["titles"],
|
||||
queryFn: () => getTitles(),
|
||||
queryKey: ["scripts", "titles-source"],
|
||||
queryFn: async () => {
|
||||
const res = await getScripts({ page_size: 200 })
|
||||
const items = Array.isArray(res) ? res : (res.items ?? [])
|
||||
return items.map((s) => ({ content: (s.title || "").trim() })).filter((s) => !!s.content)
|
||||
},
|
||||
staleTime: 30_000,
|
||||
})
|
||||
|
||||
|
||||
@@ -1,17 +1,22 @@
|
||||
/**
|
||||
* GeneratePage 步骤导航(#1899 简化为 5 步,单视频与批量一致)
|
||||
* 步骤:素材(1) → 配音(2) → 标题(3) → 确认生成(4) → 封面(5)
|
||||
* GeneratePage 步骤导航(#1970 流程重构)
|
||||
* 步骤:选择模式(1) → 选择素材(2) → 选择标题(3) → 确认生成(4) → 选择封面(5)
|
||||
*
|
||||
* - 步骤3底部按钮是「确认生成视频」(由 GenerateStepActions 调 onConfirmGenerate),
|
||||
* 创建成功后跳转步骤4;本 hook 的 goNext 只负责 1→2→3 和 4→5 的「下一步」。
|
||||
* - 步骤4(确认生成进度页):渲染全部完成(generated)后「下一步」解锁进封面。
|
||||
* - 步骤1(选择模式):下一步分支由外层弹窗处理(VoiceSelectModal / ScriptSelectModal),
|
||||
* 本 hook 的 goNext 仅在未选模式时拦截;外层 Modal onConfirm 里主动 setCurrentStep(2)。
|
||||
* - 步骤2(选择素材):弹数量选择弹窗(PreviewCountModal),确认后跳步骤3。
|
||||
* - 步骤3 底部按钮是「确认生成视频」(由 GenerateStepActions 调 onConfirmGenerate),
|
||||
* 创建成功后跳步骤4;本 hook 的 goNext 只负责 2→3 和 4→5 的「下一步」。
|
||||
* - 步骤4(确认生成进度页):全部渲染完成后「下一步」解锁进封面。
|
||||
*/
|
||||
import { message } from "antd"
|
||||
import type { TitleSettings } from "../types"
|
||||
import type { EditMode } from "../components/Step1EditMode"
|
||||
|
||||
export interface UseStepNavigationOptions {
|
||||
currentStep: number
|
||||
setCurrentStep: (step: number | ((prev: number) => number)) => void
|
||||
editMode: EditMode
|
||||
materialMode: "manual" | "auto"
|
||||
selectedMaterials: string[]
|
||||
smartSelectedIds: string[]
|
||||
@@ -20,6 +25,8 @@ export interface UseStepNavigationOptions {
|
||||
generated: boolean
|
||||
/** 点素材下一步时弹出数量选择弹窗 */
|
||||
onOpenCountModal: () => void
|
||||
/** 步骤1下一步:根据 editMode 打开对应弹窗(随机→配音 / 叙事→文案) */
|
||||
onOpenStep1Modal: () => void
|
||||
}
|
||||
|
||||
export interface UseStepNavigationReturn {
|
||||
@@ -36,22 +43,29 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
|
||||
smartSelectedIds,
|
||||
generated,
|
||||
onOpenCountModal,
|
||||
onOpenStep1Modal,
|
||||
} = options
|
||||
|
||||
const goNext = () => {
|
||||
if (currentStep === 1) {
|
||||
// 选完素材弹数量选择弹窗
|
||||
// 步骤1:先校验素材/配音等由弹窗负责,goNext 只负责触发弹窗
|
||||
onOpenStep1Modal()
|
||||
return
|
||||
}
|
||||
if (currentStep === 2) {
|
||||
// 素材校验
|
||||
if (materialMode === "manual" && selectedMaterials.length === 0) {
|
||||
message.warning("请至少选择一个素材")
|
||||
return
|
||||
}
|
||||
if (materialMode === "auto" && smartSelectedIds.length === 0) {
|
||||
message.warning("请先进行智能匹配并选择素材")
|
||||
return
|
||||
}
|
||||
// 弹数量选择弹窗
|
||||
onOpenCountModal()
|
||||
return
|
||||
}
|
||||
if (currentStep === 1 && materialMode === "manual" && selectedMaterials.length === 0) {
|
||||
message.warning("请至少选择一个素材")
|
||||
return
|
||||
}
|
||||
if (currentStep === 1 && materialMode === "auto" && smartSelectedIds.length === 0) {
|
||||
message.warning("请先进行智能匹配并选择素材")
|
||||
return
|
||||
}
|
||||
// 步骤4(确认生成):全部渲染完成后才能下一步进封面
|
||||
if (currentStep === 4) {
|
||||
if (!generated) {
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
* 操作:编辑 / 删除 / 复制 / 使用(跳创作页预填)
|
||||
* - 新建/编辑弹窗:标题(原"名称")、正文(含 AI 改写)、分类、标签
|
||||
* - #1893/#1894 AI 能力:
|
||||
* - 顶部「🎬 从抖音提取」按钮 → 输入抖音链接 → ASR 提取文案 → 自动填充到新建弹窗
|
||||
* - 顶部「🎬 从抖音提取」按钮 → 粘贴分享文案/链接(后端自动提取URL) → ASR 提取文案 → 自动填充到新建弹窗
|
||||
* - 正文下方「✨ AI 改写」按钮 → 点击直接执行(美化 loading spinner + "正在改写..."),
|
||||
* 成功自动替换正文并 toast「改写成功」1s 自动关闭;失败 toast 错误
|
||||
* - 标题旁「✨ AI 生成标题」按钮 → 候选列表一键填入
|
||||
@@ -272,20 +272,16 @@ const ScriptLibrary: React.FC = () => {
|
||||
setDouyinModalOpen(true)
|
||||
}
|
||||
|
||||
/** 执行抖音提取,成功后打开新建弹窗并预填 content */
|
||||
/** #1894:执行抖音提取。前端不再做 URL 前缀校验,直接把用户粘贴的原文(含分享文案+链接)交给后端 _extract_url_from_text 自动提取。后端 400 错误(未找到链接/非抖音域名等)直接透传给用户。 */
|
||||
const handleDouyinExtract = async () => {
|
||||
const url = douyinUrl.trim()
|
||||
if (!url) {
|
||||
message.warning("请粘贴抖音视频链接")
|
||||
return
|
||||
}
|
||||
if (!/^https?:\/\//i.test(url)) {
|
||||
message.warning("请输入以 http(s):// 开头的完整链接")
|
||||
const raw = douyinUrl.trim()
|
||||
if (!raw) {
|
||||
message.warning("请粘贴抖音视频链接或分享文案")
|
||||
return
|
||||
}
|
||||
setDouyinLoading(true)
|
||||
try {
|
||||
const res = await extractScriptFromDouyin({ url })
|
||||
const res = await extractScriptFromDouyin({ url: raw })
|
||||
message.success(`提取成功${res.duration_seconds ? `(时长 ${res.duration_seconds}s)` : ""}`)
|
||||
setDouyinModalOpen(false)
|
||||
setDouyinUrl("")
|
||||
@@ -300,6 +296,7 @@ const ScriptLibrary: React.FC = () => {
|
||||
})
|
||||
setModalOpen(true)
|
||||
} catch (err) {
|
||||
// 后端 400(未找到有效链接/仅支持抖音域名等)直接透传错误信息
|
||||
message.error(extractErrMsg(err, "抖音文案提取失败"))
|
||||
} finally {
|
||||
setDouyinLoading(false)
|
||||
@@ -637,11 +634,12 @@ const ScriptLibrary: React.FC = () => {
|
||||
destroyOnClose
|
||||
>
|
||||
<Paragraph type="secondary" style={{ marginBottom: 12, fontSize: 13 }}>
|
||||
粘贴抖音分享链接(支持 v.douyin.com 短链和 www.douyin.com/video/ 长链), AI
|
||||
将自动下载音频并识别文案。首次识别可能需要 5-15 秒。
|
||||
粘贴抖音分享文案或链接即可,系统会自动从文本中识别链接(支持 v.douyin.com 短链、
|
||||
www.douyin.com/video/ 长链,以及 App「复制链接」带的分享文案)。AI
|
||||
将自动下载音频并识别文案, 首次识别可能需要 5-15 秒。
|
||||
</Paragraph>
|
||||
<Input.TextArea
|
||||
placeholder="例如:https://v.douyin.com/xxxxx/ 或 https://www.douyin.com/video/xxxxx"
|
||||
placeholder="直接粘贴 App「复制链接」的全部内容即可,例如:8.88 复制打开抖音... https://v.douyin.com/xxxxx/"
|
||||
value={douyinUrl}
|
||||
onChange={(e) => setDouyinUrl(e.target.value)}
|
||||
rows={2}
|
||||
@@ -651,7 +649,7 @@ const ScriptLibrary: React.FC = () => {
|
||||
{douyinLoading && (
|
||||
<div className="xx-ai-loading-hint">
|
||||
<Spin size="small" style={{ marginRight: 8 }} />
|
||||
正在下载视频并识别文案,可能需要数秒,请稍候…
|
||||
提取中,正在下载视频并识别文案,可能需要数秒,请稍候…
|
||||
</div>
|
||||
)}
|
||||
</Modal>
|
||||
|
||||
@@ -39,6 +39,7 @@ import { getDiscountPriceCents } from "@/api/points/types"
|
||||
import type { SubscriptionPlan } from "@/api/subscription/types"
|
||||
import { PLAN_LABEL, BILLING_CYCLE_LABEL } from "@/api/subscription/types"
|
||||
import "./Plans.css"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
|
||||
const { Title, Text, Paragraph } = Typography
|
||||
|
||||
@@ -249,17 +250,23 @@ const Plans: React.FC = () => {
|
||||
return (
|
||||
<div className="xx-plans-page">
|
||||
<PageHead
|
||||
title="会员与积分"
|
||||
description="开通会员解锁全部功能,按需充值积分灵活使用 AI 能力"
|
||||
title={ENABLE_CREDIT_SYSTEM ? "会员与积分" : "会员订阅"}
|
||||
description={
|
||||
ENABLE_CREDIT_SYSTEM
|
||||
? "开通会员解锁全部功能,按需充值积分灵活使用 AI 能力"
|
||||
: "开通会员解锁全部功能"
|
||||
}
|
||||
actions={
|
||||
<Space>
|
||||
<Button
|
||||
icon={<ThunderboltOutlined />}
|
||||
onClick={() => navigate("/app/points/transactions")}
|
||||
>
|
||||
积分明细
|
||||
</Button>
|
||||
</Space>
|
||||
ENABLE_CREDIT_SYSTEM ? (
|
||||
<Space>
|
||||
<Button
|
||||
icon={<ThunderboltOutlined />}
|
||||
onClick={() => navigate("/app/points/transactions")}
|
||||
>
|
||||
积分明细
|
||||
</Button>
|
||||
</Space>
|
||||
) : null
|
||||
}
|
||||
/>
|
||||
|
||||
@@ -296,13 +303,15 @@ const Plans: React.FC = () => {
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
<div>
|
||||
<Text type="secondary">可用积分</Text>
|
||||
<div className="xx-current-balance">
|
||||
<ThunderboltOutlined style={{ color: "#8b5cf6" }} />
|
||||
<span className="xx-current-balance-val">{bal}</span>
|
||||
{ENABLE_CREDIT_SYSTEM && (
|
||||
<div>
|
||||
<Text type="secondary">可用积分</Text>
|
||||
<div className="xx-current-balance">
|
||||
<ThunderboltOutlined style={{ color: "#8b5cf6" }} />
|
||||
<span className="xx-current-balance-val">{bal}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{!isMember && freeLimit > 0 && (
|
||||
<div>
|
||||
<Text type="secondary">今日免费混剪</Text>
|
||||
@@ -319,18 +328,20 @@ const Plans: React.FC = () => {
|
||||
)}
|
||||
</Space>
|
||||
</Col>
|
||||
<Col>
|
||||
<Button
|
||||
type="primary"
|
||||
icon={<ThunderboltOutlined />}
|
||||
onClick={() => {
|
||||
const el = document.getElementById("points-packages")
|
||||
el?.scrollIntoView({ behavior: "smooth" })
|
||||
}}
|
||||
>
|
||||
充值积分
|
||||
</Button>
|
||||
</Col>
|
||||
{ENABLE_CREDIT_SYSTEM && (
|
||||
<Col>
|
||||
<Button
|
||||
type="primary"
|
||||
icon={<ThunderboltOutlined />}
|
||||
onClick={() => {
|
||||
const el = document.getElementById("points-packages")
|
||||
el?.scrollIntoView({ behavior: "smooth" })
|
||||
}}
|
||||
>
|
||||
充值积分
|
||||
</Button>
|
||||
</Col>
|
||||
)}
|
||||
</Row>
|
||||
</Card>
|
||||
|
||||
@@ -461,69 +472,71 @@ const Plans: React.FC = () => {
|
||||
</Col>
|
||||
</Row>
|
||||
|
||||
{/* 积分充值 */}
|
||||
<div id="points-packages">
|
||||
<Title level={4} style={{ marginTop: 40 }}>
|
||||
<ThunderboltOutlined style={{ color: "#8b5cf6", marginRight: 8 }} />
|
||||
积分充值
|
||||
<Tooltip title="积分永久有效,可用于所有 AI 功能;付费会员享折扣">
|
||||
<Text type="secondary" style={{ fontSize: 13, marginLeft: 8, fontWeight: "normal" }}>
|
||||
(永久有效)
|
||||
</Text>
|
||||
</Tooltip>
|
||||
</Title>
|
||||
{/* 积分充值(积分系统关闭时隐藏,代码保留不删除) */}
|
||||
{ENABLE_CREDIT_SYSTEM && (
|
||||
<div id="points-packages">
|
||||
<Title level={4} style={{ marginTop: 40 }}>
|
||||
<ThunderboltOutlined style={{ color: "#8b5cf6", marginRight: 8 }} />
|
||||
积分充值
|
||||
<Tooltip title="积分永久有效,可用于所有 AI 功能;付费会员享折扣">
|
||||
<Text type="secondary" style={{ fontSize: 13, marginLeft: 8, fontWeight: "normal" }}>
|
||||
(永久有效)
|
||||
</Text>
|
||||
</Tooltip>
|
||||
</Title>
|
||||
|
||||
<Row gutter={[16, 16]}>
|
||||
{packages.map((pkg) => {
|
||||
const priceCents = getDiscountPriceCents(pkg, userDiscount)
|
||||
const originalCents = pkg.price_cents
|
||||
const discount =
|
||||
priceCents < originalCents ? Math.round((1 - priceCents / originalCents) * 100) : 0
|
||||
const unit = priceCents / 100 / pkg.points
|
||||
const isHot = pkg.unit_price < 0.1
|
||||
return (
|
||||
<Col xs={24} sm={8} key={pkg.code}>
|
||||
<Card
|
||||
className={`xx-pkg-card ${discount > 0 ? "has-discount" : ""} ${isHot ? "recommended" : ""}`}
|
||||
hoverable
|
||||
>
|
||||
{isHot && <div className="xx-pkg-badge">热门</div>}
|
||||
{discount > 0 && (
|
||||
<Tag color="gold" className="xx-pkg-discount">
|
||||
{Math.round((priceCents / originalCents) * 10) / 1}折
|
||||
</Tag>
|
||||
)}
|
||||
<div className="xx-pkg-name">{pkg.name}</div>
|
||||
<div className="xx-pkg-points">
|
||||
<ThunderboltOutlined /> {pkg.points.toLocaleString()} 积分
|
||||
</div>
|
||||
<div className="xx-pkg-price">
|
||||
<span className="currency">¥</span>
|
||||
<span className="amount">
|
||||
{(priceCents / 100)
|
||||
.toFixed(priceCents % 100 === 0 ? 0 : 1)
|
||||
.replace(/\.0$/, "")}
|
||||
</span>
|
||||
{discount > 0 && (
|
||||
<span className="xx-pkg-origin">¥{(originalCents / 100).toFixed(0)}</span>
|
||||
)}
|
||||
</div>
|
||||
<div className="xx-pkg-unit">≈¥{unit.toFixed(3)}/积分</div>
|
||||
<Button
|
||||
block
|
||||
type={isHot ? "primary" : "default"}
|
||||
loading={buying === pkg.code}
|
||||
onClick={() => handleBuyPoints(pkg)}
|
||||
style={{ marginTop: 12 }}
|
||||
<Row gutter={[16, 16]}>
|
||||
{packages.map((pkg) => {
|
||||
const priceCents = getDiscountPriceCents(pkg, userDiscount)
|
||||
const originalCents = pkg.price_cents
|
||||
const discount =
|
||||
priceCents < originalCents ? Math.round((1 - priceCents / originalCents) * 100) : 0
|
||||
const unit = priceCents / 100 / pkg.points
|
||||
const isHot = pkg.unit_price < 0.1
|
||||
return (
|
||||
<Col xs={24} sm={8} key={pkg.code}>
|
||||
<Card
|
||||
className={`xx-pkg-card ${discount > 0 ? "has-discount" : ""} ${isHot ? "recommended" : ""}`}
|
||||
hoverable
|
||||
>
|
||||
立即购买
|
||||
</Button>
|
||||
</Card>
|
||||
</Col>
|
||||
)
|
||||
})}
|
||||
</Row>
|
||||
</div>
|
||||
{isHot && <div className="xx-pkg-badge">热门</div>}
|
||||
{discount > 0 && (
|
||||
<Tag color="gold" className="xx-pkg-discount">
|
||||
{Math.round((priceCents / originalCents) * 10) / 1}折
|
||||
</Tag>
|
||||
)}
|
||||
<div className="xx-pkg-name">{pkg.name}</div>
|
||||
<div className="xx-pkg-points">
|
||||
<ThunderboltOutlined /> {pkg.points.toLocaleString()} 积分
|
||||
</div>
|
||||
<div className="xx-pkg-price">
|
||||
<span className="currency">¥</span>
|
||||
<span className="amount">
|
||||
{(priceCents / 100)
|
||||
.toFixed(priceCents % 100 === 0 ? 0 : 1)
|
||||
.replace(/\.0$/, "")}
|
||||
</span>
|
||||
{discount > 0 && (
|
||||
<span className="xx-pkg-origin">¥{(originalCents / 100).toFixed(0)}</span>
|
||||
)}
|
||||
</div>
|
||||
<div className="xx-pkg-unit">≈¥{unit.toFixed(3)}/积分</div>
|
||||
<Button
|
||||
block
|
||||
type={isHot ? "primary" : "default"}
|
||||
loading={buying === pkg.code}
|
||||
onClick={() => handleBuyPoints(pkg)}
|
||||
style={{ marginTop: 12 }}
|
||||
>
|
||||
立即购买
|
||||
</Button>
|
||||
</Card>
|
||||
</Col>
|
||||
)
|
||||
})}
|
||||
</Row>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
* - subscription: GET /subscription/current(plan_id + billing_cycle)
|
||||
*/
|
||||
import { create } from "zustand"
|
||||
import { ENABLE_CREDIT_SYSTEM } from "@/config/features"
|
||||
import { getPointsBalance, getPointsRules, getDailyUsage, getMembership } from "@/api/points"
|
||||
import { getCurrentSubscription } from "@/api/subscription"
|
||||
import type {
|
||||
@@ -49,14 +50,29 @@ export const usePointsStore = create<PointsState>((set, get) => ({
|
||||
|
||||
init: async () => {
|
||||
// 已加载过不重复拉取
|
||||
if (get().balance && get().rules && get().subscription) return
|
||||
// 积分系统关闭时:只要 subscription/membership 已有值就跳过;开启时需 balance+rules+subscription 齐了才跳过
|
||||
if (ENABLE_CREDIT_SYSTEM) {
|
||||
if (get().balance && get().rules && get().subscription) return
|
||||
} else {
|
||||
if (get().subscription && get().membership) return
|
||||
}
|
||||
set({ loading: true, error: null })
|
||||
try {
|
||||
// 积分系统关闭时不拉取余额/规则/每日额度,但仍拉会员/订阅用于 VIP 标识展示
|
||||
const balancePromise = ENABLE_CREDIT_SYSTEM
|
||||
? getPointsBalance().catch(() => null)
|
||||
: Promise.resolve(null)
|
||||
const rulesPromise = ENABLE_CREDIT_SYSTEM
|
||||
? getPointsRules().catch(() => null)
|
||||
: Promise.resolve(null)
|
||||
const dailyUsagePromise = ENABLE_CREDIT_SYSTEM
|
||||
? getDailyUsage().catch(() => null)
|
||||
: Promise.resolve(null)
|
||||
const [balance, rules, subscription, dailyUsage, membership] = await Promise.all([
|
||||
getPointsBalance().catch(() => null),
|
||||
getPointsRules().catch(() => null),
|
||||
balancePromise,
|
||||
rulesPromise,
|
||||
getCurrentSubscription().catch(() => null),
|
||||
getDailyUsage().catch(() => null),
|
||||
dailyUsagePromise,
|
||||
getMembership().catch(() => null),
|
||||
])
|
||||
set({
|
||||
|
||||
@@ -1,112 +0,0 @@
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest"
|
||||
import { getTitles, createTitle, updateTitle, deleteTitle, batchImportTitles } from "@/api/titles"
|
||||
|
||||
const mockGet = vi.fn()
|
||||
const mockPost = vi.fn()
|
||||
const mockPut = vi.fn()
|
||||
const mockDelete = vi.fn()
|
||||
const mockPatch = vi.fn()
|
||||
|
||||
vi.mock("@/api/client", () => ({
|
||||
default: {
|
||||
get: (...args: unknown[]) => mockGet(...args),
|
||||
post: (...args: unknown[]) => mockPost(...args),
|
||||
put: (...args: unknown[]) => mockPut(...args),
|
||||
delete: (...args: unknown[]) => mockDelete(...args),
|
||||
patch: (...args: unknown[]) => mockPatch(...args),
|
||||
},
|
||||
}))
|
||||
|
||||
vi.mock("antd", () => ({ message: { error: vi.fn(), success: vi.fn() } }))
|
||||
vi.mock("@/store/authStore", () => ({ useAuthStore: { getState: vi.fn(() => ({})) } }))
|
||||
|
||||
describe("titles API", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
mockGet.mockResolvedValue({ data: { success: true, items: [] } })
|
||||
mockPost.mockResolvedValue({ data: { success: true, items: [] } })
|
||||
mockPut.mockResolvedValue({ data: { success: true, items: [] } })
|
||||
mockDelete.mockResolvedValue({ data: { success: true, items: [] } })
|
||||
mockPatch.mockResolvedValue({ data: { success: true, items: [] } })
|
||||
})
|
||||
|
||||
describe("getTitles", () => {
|
||||
it("should resolve successfully", async () => {
|
||||
await expect(getTitles()).resolves.not.toThrow()
|
||||
})
|
||||
|
||||
it("should reject on API error", async () => {
|
||||
mockGet.mockRejectedValue(new Error("Network error"))
|
||||
mockPost.mockRejectedValue(new Error("Network error"))
|
||||
mockPut.mockRejectedValue(new Error("Network error"))
|
||||
mockDelete.mockRejectedValue(new Error("Network error"))
|
||||
mockPatch.mockRejectedValue(new Error("Network error"))
|
||||
|
||||
await expect(getTitles()).rejects.toThrow()
|
||||
})
|
||||
})
|
||||
|
||||
describe("createTitle", () => {
|
||||
it("should resolve successfully", async () => {
|
||||
await expect(createTitle({ title: "测试标题", content: "测试内容" })).resolves.not.toThrow()
|
||||
})
|
||||
|
||||
it("should reject on API error", async () => {
|
||||
mockGet.mockRejectedValue(new Error("Network error"))
|
||||
mockPost.mockRejectedValue(new Error("Network error"))
|
||||
mockPut.mockRejectedValue(new Error("Network error"))
|
||||
mockDelete.mockRejectedValue(new Error("Network error"))
|
||||
mockPatch.mockRejectedValue(new Error("Network error"))
|
||||
|
||||
await expect(createTitle({ name: "test-item" })).rejects.toThrow()
|
||||
})
|
||||
})
|
||||
|
||||
describe("updateTitle", () => {
|
||||
it("should resolve successfully", async () => {
|
||||
await expect(updateTitle("test-titleId", { title: "新标题" })).resolves.not.toThrow()
|
||||
})
|
||||
|
||||
it("should reject on API error", async () => {
|
||||
mockGet.mockRejectedValue(new Error("Network error"))
|
||||
mockPost.mockRejectedValue(new Error("Network error"))
|
||||
mockPut.mockRejectedValue(new Error("Network error"))
|
||||
mockDelete.mockRejectedValue(new Error("Network error"))
|
||||
mockPatch.mockRejectedValue(new Error("Network error"))
|
||||
|
||||
await expect(updateTitle("test-titleId")).rejects.toThrow()
|
||||
})
|
||||
})
|
||||
|
||||
describe("deleteTitle", () => {
|
||||
it("should resolve successfully", async () => {
|
||||
await expect(deleteTitle("test-titleId")).resolves.not.toThrow()
|
||||
})
|
||||
|
||||
it("should reject on API error", async () => {
|
||||
mockGet.mockRejectedValue(new Error("Network error"))
|
||||
mockPost.mockRejectedValue(new Error("Network error"))
|
||||
mockPut.mockRejectedValue(new Error("Network error"))
|
||||
mockDelete.mockRejectedValue(new Error("Network error"))
|
||||
mockPatch.mockRejectedValue(new Error("Network error"))
|
||||
|
||||
await expect(deleteTitle("test-titleId")).rejects.toThrow()
|
||||
})
|
||||
})
|
||||
|
||||
describe("batchImportTitles", () => {
|
||||
it("should resolve successfully", async () => {
|
||||
await expect(batchImportTitles("test-titles")).resolves.not.toThrow()
|
||||
})
|
||||
|
||||
it("should reject on API error", async () => {
|
||||
mockGet.mockRejectedValue(new Error("Network error"))
|
||||
mockPost.mockRejectedValue(new Error("Network error"))
|
||||
mockPut.mockRejectedValue(new Error("Network error"))
|
||||
mockDelete.mockRejectedValue(new Error("Network error"))
|
||||
mockPatch.mockRejectedValue(new Error("Network error"))
|
||||
|
||||
await expect(batchImportTitles("test-titles")).rejects.toThrow()
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -39,7 +39,6 @@ describe("navigation config", () => {
|
||||
expect(keys).toContain("dashboard")
|
||||
expect(keys).toContain("assets")
|
||||
expect(keys).toContain("voices")
|
||||
expect(keys).toContain("titles")
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
@@ -215,8 +215,14 @@ vi.mock("@/api/editing-planner", () => ({
|
||||
MODE_LABELS: { pip: "画中画" },
|
||||
}))
|
||||
|
||||
vi.mock("@/api/titles", () => ({
|
||||
getTitles: vi.fn().mockResolvedValue({ items: [] }),
|
||||
// #1894: 标题数据源已切到 @/api/scripts,mock scripts 返回空数组作为默认
|
||||
vi.mock("@/api/scripts", () => ({
|
||||
getScripts: vi.fn().mockResolvedValue({ items: [], total: 0, page: 1, page_size: 20 }),
|
||||
aiRewriteScript: vi.fn(),
|
||||
aiGenerateTitles: vi.fn(),
|
||||
SCRIPTS_API_MOCK: false,
|
||||
SCRIPT_CATEGORY_LABEL: {},
|
||||
REWRITE_STYLE_OPTIONS: [],
|
||||
}))
|
||||
|
||||
vi.mock("@/api/template-editor", () => ({
|
||||
|
||||
@@ -32,6 +32,7 @@ vi.mock("antd", () => ({
|
||||
|
||||
vi.mock("@/api/subscription", () => ({
|
||||
getCurrentSubscription: vi.fn().mockResolvedValue({ plan: "free", status: "active" }),
|
||||
getSubscriptionPlans: vi.fn().mockResolvedValue({ items: [{ plan_id: "free", name: "Free" }] }),
|
||||
changePlan: vi.fn().mockResolvedValue({ success: true }),
|
||||
toggleAutoRenew: vi.fn().mockResolvedValue({ success: true }),
|
||||
cancelSubscription: vi.fn().mockResolvedValue({ success: true }),
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
"""智能降重微变换纯逻辑模块 — #1970 PR2.
|
||||
|
||||
所有函数均为纯函数:不调用 FFmpeg、不读写文件,只负责按可复现种子
|
||||
生成每个片段 / 整片的微变换参数与 filter_complex 片段。
|
||||
|
||||
6 个维度:
|
||||
1. hflip 水平翻转(每片段 50%,有字幕/文字的片段不翻转)
|
||||
2. 播放速度 0.97~1.03x(视频 setpts + 音频 atempo)
|
||||
3. 亮度 ±2%(eq=brightness)
|
||||
4. 对比度 ±2%(eq=contrast)
|
||||
5. 饱和度 ±2%(eq=saturation)
|
||||
6. BGM 起始偏移 2~8 秒(音频 atrim 起点)
|
||||
|
||||
随机种子 = hash(task_id + video_index) % 10000,保证同一任务同一视频
|
||||
可复现;dedup_enabled=False 时不生成本模块任何输出。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
# ── 常量(与需求文档 §2 对齐)──────────────────────────────────────────────────
|
||||
|
||||
SPEED_MIN = 0.97
|
||||
SPEED_MAX = 1.03
|
||||
COLOR_DELTA = 0.02
|
||||
HFLIP_PROBABILITY = 0.5
|
||||
BGM_OFFSET_MIN = 2.0
|
||||
BGM_OFFSET_MAX = 8.0
|
||||
SEED_MODULO = 10000
|
||||
|
||||
|
||||
def make_video_seed(task_id: str, video_index: int) -> int:
|
||||
"""生成视频级可复现种子:hash(task_id+video_index) % 10000。
|
||||
|
||||
用 sha256 而非内置 hash():内置 hash 对字符串带进程级随机盐(PYTHONHASHSEED),
|
||||
跨进程不可复现。结果映射到 0~9999。
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
raw = f"{task_id or ''}:{int(video_index)}"
|
||||
digest = hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
return int(digest[:8], 16) % SEED_MODULO
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ClipMicroTransform:
|
||||
"""单个片段的微变换参数。"""
|
||||
|
||||
clip_index: int
|
||||
hflip: bool = False
|
||||
speed: float = 1.0
|
||||
brightness: float = 0.0
|
||||
contrast: float = 1.0
|
||||
saturation: float = 1.0
|
||||
has_text: bool = False
|
||||
|
||||
def video_filter_suffix(self) -> str:
|
||||
"""返回追加在片段视频处理链上的 filter 后缀(无末尾标签)。
|
||||
|
||||
顺序:trim/setpts(已有)→ 调速 setpts → hflip → eq → format。
|
||||
调速的 setpts 必须位于 trim 之后;hflip/eq 在缩放之后即可,
|
||||
concat_engine 按「调速 → hflip → eq」顺序拼接到 scale/fps 之前的
|
||||
trim 之后、scale 之后均可,这里只产出独立步骤、由引擎决定插入点。
|
||||
"""
|
||||
parts: list[str] = []
|
||||
# 速度:setpts=PTS/speed(speed>1 时画面加速,时间戳变小)
|
||||
if abs(self.speed - 1.0) > 1e-4:
|
||||
parts.append(f"setpts=PTS/{self.speed:.5f}")
|
||||
# 水平翻转:有文字/字幕片段不翻转
|
||||
if self.hflip and not self.has_text:
|
||||
parts.append("hflip")
|
||||
# 色彩微调:brightness 取值 -1~1(±0.02),contrast/saturation 围绕 1.0
|
||||
if abs(self.brightness) > 1e-4 or abs(self.contrast - 1.0) > 1e-4 or abs(self.saturation - 1.0) > 1e-4:
|
||||
parts.append(
|
||||
f"eq=brightness={self.brightness:+.4f}:"
|
||||
f"contrast={self.contrast:.4f}:saturation={self.saturation:.4f}"
|
||||
)
|
||||
return ",".join(parts)
|
||||
|
||||
def audio_filter_suffix(self) -> str:
|
||||
"""返回片段音频链上的调速 filter(atempo),无调速时返回空串。"""
|
||||
if abs(self.speed - 1.0) <= 1e-4:
|
||||
return ""
|
||||
return f"atempo={self.speed:.5f}"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VideoMicroTransformPlan:
|
||||
"""一个成片视频的全部微变换参数。"""
|
||||
|
||||
task_id: str
|
||||
video_index: int
|
||||
seed: int
|
||||
clips: list[ClipMicroTransform] = field(default_factory=list)
|
||||
bgm_start_offset: float = 0.0
|
||||
|
||||
def clip(self, index: int) -> ClipMicroTransform | None:
|
||||
for c in self.clips:
|
||||
if c.clip_index == index:
|
||||
return c
|
||||
return None
|
||||
|
||||
|
||||
def _draw_speed(rng: random.Random) -> float:
|
||||
return round(rng.uniform(SPEED_MIN, SPEED_MAX), 5)
|
||||
|
||||
|
||||
def _draw_signed_delta(rng: random.Random) -> float:
|
||||
return round(rng.uniform(-COLOR_DELTA, COLOR_DELTA), 4)
|
||||
|
||||
|
||||
def build_micro_transform_plan(
|
||||
task_id: str,
|
||||
video_index: int,
|
||||
clip_count: int,
|
||||
*,
|
||||
clip_has_text: list[bool] | None = None,
|
||||
enable_bgm_offset: bool = True,
|
||||
) -> VideoMicroTransformPlan:
|
||||
"""按可复现种子生成整片的微变换计划。
|
||||
|
||||
Args:
|
||||
task_id: 生成任务 ID(种子输入)
|
||||
video_index: 视频在批次中的序号(0 起)
|
||||
clip_count: 片段数量
|
||||
clip_has_text: 每个片段是否有字幕/文字轨道(True 的片段不翻转);
|
||||
None 时按 P1 约定视为无可靠文字检测——保守起见 hflip 一律关闭
|
||||
enable_bgm_offset: 是否生成 BGM 起始偏移(无 BGM 时调用方可忽略该值)
|
||||
|
||||
Returns:
|
||||
VideoMicroTransformPlan
|
||||
"""
|
||||
seed = make_video_seed(task_id, video_index)
|
||||
rng = random.Random(seed)
|
||||
|
||||
# P1 字幕检测约定:无法判断片段是否有文字时,一律不翻转(宁可少一个维度也不误翻字幕)
|
||||
safe_has_text = clip_has_text if clip_has_text is not None else [True] * max(clip_count, 0)
|
||||
|
||||
clips: list[ClipMicroTransform] = []
|
||||
for i in range(max(clip_count, 0)):
|
||||
has_text = bool(safe_has_text[i]) if i < len(safe_has_text) else True
|
||||
do_hflip = (not has_text) and rng.random() < HFLIP_PROBABILITY
|
||||
clips.append(
|
||||
ClipMicroTransform(
|
||||
clip_index=i,
|
||||
hflip=do_hflip,
|
||||
speed=_draw_speed(rng),
|
||||
brightness=_draw_signed_delta(rng),
|
||||
contrast=round(1.0 + _draw_signed_delta(rng), 4),
|
||||
saturation=round(1.0 + _draw_signed_delta(rng), 4),
|
||||
has_text=has_text,
|
||||
)
|
||||
)
|
||||
|
||||
bgm_offset = rng.uniform(BGM_OFFSET_MIN, BGM_OFFSET_MAX) if enable_bgm_offset else 0.0
|
||||
return VideoMicroTransformPlan(
|
||||
task_id=task_id,
|
||||
video_index=video_index,
|
||||
seed=seed,
|
||||
clips=clips,
|
||||
bgm_start_offset=round(bgm_offset, 3),
|
||||
)
|
||||
|
||||
|
||||
def build_bgm_offset_trim(start_offset: float, bgm_duration: float) -> str:
|
||||
"""生成 BGM 起始偏移的 atrim 片段。
|
||||
|
||||
偏移超出 BGM 长度时回退为 0(从头播放),避免空输入。
|
||||
返回的字符串形如 "atrim=start=3.200,",可拼到 BGM filter chain 最前面;
|
||||
无需偏移时返回空串。
|
||||
"""
|
||||
if start_offset <= 0 or bgm_duration <= 0 or start_offset >= bgm_duration - 0.5:
|
||||
return ""
|
||||
return f"atrim=start={start_offset:.3f},"
|
||||
@@ -493,6 +493,41 @@ class RenderAdapter:
|
||||
logger.warning("ASR 服务初始化失败,自动字幕将不可用: %s", e)
|
||||
return None
|
||||
|
||||
def _resolve_clip_has_text(self, clips: list[Any]) -> list[bool] | None:
|
||||
"""#1970:按源视频片段顺序解析 atom_clip.ai_tags.has_text。
|
||||
|
||||
顺序与 UnifiedRenderService 的「非 audio 源片段」口径一致。
|
||||
仅当 atom_clip 存在 ai_tags 字典且 has_text 显式为 False 时标记为
|
||||
无文字(允许 hflip);atom_clip_id 缺失、ai_tags 未生成、has_text 为
|
||||
true/null/非布尔值时一律按有文字处理(保守不翻转)。
|
||||
查询失败时返回 None,渲染层回退到全保守路径。
|
||||
"""
|
||||
video_clips = [c for c in clips if getattr(c, "clip_type", "main") != "audio"]
|
||||
atom_ids: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for c in video_clips:
|
||||
atom_id = getattr(c, "atom_clip_id", "") or ""
|
||||
if atom_id and atom_id not in seen:
|
||||
seen.add(atom_id)
|
||||
atom_ids.append(atom_id)
|
||||
if not atom_ids:
|
||||
return None
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
|
||||
atom_clips = SQLAlchemyAssetAtomClipRepository(self._db).find_by_ids(atom_ids)
|
||||
except Exception as exc:
|
||||
logger.warning("[render-adapter] atom_clip ai_tags 查询失败,hflip 全量保守处理: %s", exc)
|
||||
return None
|
||||
has_text_map: dict[str, bool] = {}
|
||||
for ac in atom_clips:
|
||||
ai_tags = getattr(ac, "ai_tags", None)
|
||||
no_text = isinstance(ai_tags, dict) and ai_tags.get("has_text") is False
|
||||
has_text_map[ac.id] = not no_text
|
||||
return [has_text_map.get((getattr(c, "atom_clip_id", "") or ""), True) for c in video_clips]
|
||||
|
||||
def _do_render(
|
||||
self,
|
||||
plan: Any,
|
||||
@@ -542,6 +577,7 @@ class RenderAdapter:
|
||||
)
|
||||
|
||||
# 4. 执行统一渲染
|
||||
clip_has_text = self._resolve_clip_has_text(clips)
|
||||
render_svc = UnifiedRenderService(
|
||||
plan=plan,
|
||||
clips=clips,
|
||||
@@ -552,6 +588,7 @@ class RenderAdapter:
|
||||
bgm_path=bgm_path,
|
||||
asr_service=asr_service,
|
||||
voiceover_audio_path=voiceover_audio_path,
|
||||
clip_has_text=clip_has_text,
|
||||
)
|
||||
result = render_svc.render()
|
||||
|
||||
|
||||
@@ -98,6 +98,7 @@ def mix_audio(
|
||||
bgm_path: str | None = None,
|
||||
bgm_config: dict | None = None,
|
||||
audio_tracks_config: dict | None = None,
|
||||
bgm_start_offset: float = 0.0,
|
||||
) -> Path | None:
|
||||
"""音频后处理混音.
|
||||
|
||||
@@ -157,7 +158,10 @@ def mix_audio(
|
||||
if bgm_path and bgm_config and isinstance(bgm_config, dict) and bgm_config.get("enabled", False):
|
||||
from video_processing.bgm_mixer import BGMConfig, build_bgm_only
|
||||
|
||||
bgm_cfg = BGMConfig.from_config_dict(bgm_path, bgm_config)
|
||||
_bgm_cfg_dict = dict(bgm_config or {})
|
||||
if bgm_start_offset and not _bgm_cfg_dict.get("audio_offset"):
|
||||
_bgm_cfg_dict["audio_offset"] = round(float(bgm_start_offset), 3)
|
||||
bgm_cfg = BGMConfig.from_config_dict(bgm_path, _bgm_cfg_dict)
|
||||
try:
|
||||
return build_bgm_only(ctx, bgm_cfg, video_duration)
|
||||
except Exception:
|
||||
@@ -187,7 +191,10 @@ def mix_audio(
|
||||
if bgm_path and bgm_config and isinstance(bgm_config, dict) and bgm_config.get("enabled", False):
|
||||
from video_processing.bgm_mixer import BGMConfig, mix_bgm_with_main
|
||||
|
||||
bgm_cfg = BGMConfig.from_config_dict(bgm_path, bgm_config)
|
||||
_bgm_cfg_dict = dict(bgm_config or {})
|
||||
if bgm_start_offset and not _bgm_cfg_dict.get("audio_offset"):
|
||||
_bgm_cfg_dict["audio_offset"] = round(float(bgm_start_offset), 3)
|
||||
bgm_cfg = BGMConfig.from_config_dict(bgm_path, _bgm_cfg_dict)
|
||||
|
||||
try:
|
||||
# 这里 main_audio 就是 output_path,先有主音频再混 BGM
|
||||
|
||||
@@ -155,6 +155,7 @@ class UnifiedRenderService:
|
||||
asr_service: Any = None, # ASRService 实例,用于自动生成字幕
|
||||
bgm_path: str | None = None, # BGM 本地文件路径
|
||||
voiceover_audio_path: str | None = None, # 配音素材库音频本地路径
|
||||
clip_has_text: list[bool] | None = None, # 源视频片段是否有文字(来自 atom_clip.ai_tags.has_text)
|
||||
):
|
||||
self.plan = plan
|
||||
self.clips = clips
|
||||
@@ -167,10 +168,98 @@ class UnifiedRenderService:
|
||||
self.asr_service = asr_service
|
||||
self.bgm_path = bgm_path
|
||||
self.voiceover_audio_path = voiceover_audio_path
|
||||
# #1970:片段级文字检测(顺序与非 audio 的源视频片段一致);None 表示无可靠检测,保守不翻转
|
||||
self._clip_has_text = clip_has_text
|
||||
self._transition_engine = TransitionEngine(default_duration=transition_duration)
|
||||
self._speed_engine = SpeedEngine()
|
||||
self._asr_timeline_cache: Any = None # ASR 字幕结果缓存,避免重复调用
|
||||
self._asr_timeline_cached = False
|
||||
# #1970 PR2:片段级微变换计划缓存(懒构建,dedup_enabled=False 时为 None)
|
||||
self._micro_plan_cache: Any = None
|
||||
self._micro_plan_loaded = False
|
||||
|
||||
# ── #1970 PR2 智能降重:片段级微变换 ───────────────────────────────────
|
||||
def _dedup_enabled(self) -> bool:
|
||||
"""读取 plan.config.dedup_enabled,缺省视为 True(向后兼容)。"""
|
||||
cfg = self.plan.config or {}
|
||||
return bool(cfg.get("dedup_enabled", True))
|
||||
|
||||
def _get_micro_transform_plan(self, clip_count: int) -> Any:
|
||||
"""按 task_id+视频序号构建可复现的片段级微变换计划。
|
||||
|
||||
种子 hash(generation_task_id + video_index)%10000,同一任务重渲结果一致。
|
||||
dedup_enabled=False 时返回 None,调用方不注入任何微变换。
|
||||
hflip 放开(#1970):clip_has_text 来自 atom_clip.ai_tags.has_text,
|
||||
仅 AI 明确判定无文字的片段可参与 50% 翻转;未打标签 / has_text 为
|
||||
true/null 或缺位时一律视为有文字,保持保守不翻转。
|
||||
"""
|
||||
if self._micro_plan_loaded:
|
||||
return self._micro_plan_cache
|
||||
self._micro_plan_loaded = True
|
||||
if not self._dedup_enabled() or clip_count <= 0:
|
||||
self._micro_plan_cache = None
|
||||
return None
|
||||
try:
|
||||
from video_processing.micro_transform_pure import build_micro_transform_plan
|
||||
|
||||
cfg = self.plan.config or {}
|
||||
task_id = str(cfg.get("generation_task_id", "") or "")
|
||||
video_index = int(cfg.get("video_index", 0) or 0)
|
||||
# self._clip_has_text 顺序与非 audio 源片段一致;
|
||||
# None(未提供检测,如内存直渲/旧任务)→ 纯函数层按全有文字保守处理;
|
||||
# 列表短于片段数时缺位片段同样按有文字处理
|
||||
self._micro_plan_cache = build_micro_transform_plan(
|
||||
task_id,
|
||||
video_index,
|
||||
clip_count,
|
||||
clip_has_text=self._clip_has_text,
|
||||
enable_bgm_offset=bool(cfg.get("bgm")),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[unified-render] 微变换计划构建失败,本次不注入: %s", e)
|
||||
self._micro_plan_cache = None
|
||||
return self._micro_plan_cache
|
||||
|
||||
@staticmethod
|
||||
def _apply_micro_transform_video(filters: list[str], mt: Any) -> None:
|
||||
"""把片段视频微变换就地追加到 filter 链(post-scale 阶段调用)。
|
||||
|
||||
顺序:hflip 在 pre-scale 阶段由 _apply_micro_hflip 处理,这里只加
|
||||
eq 亮度/对比度/饱和度。速度 setpts 与既有 clip speed 相乘(见调用点),
|
||||
避免出现两条 setpts 互相覆盖。
|
||||
"""
|
||||
if mt is None:
|
||||
return
|
||||
if abs(mt.brightness) > 1e-4 or abs(mt.contrast - 1.0) > 1e-4 or abs(mt.saturation - 1.0) > 1e-4:
|
||||
filters.append(
|
||||
f"eq=brightness={mt.brightness:+.4f}:" f"contrast={mt.contrast:.4f}:saturation={mt.saturation:.4f}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _apply_micro_hflip(filters: list[str], mt: Any) -> None:
|
||||
"""片段级水平翻转(pre-scale 阶段)。P1 有文字/无法判定时 mt.hflip=False。"""
|
||||
if mt is not None and mt.hflip and not mt.has_text:
|
||||
filters.append("hflip")
|
||||
|
||||
@staticmethod
|
||||
def _micro_speed_factor(mt: Any) -> float:
|
||||
"""片段微变换速度因子(0.97~1.03),无计划返回 1.0。"""
|
||||
if mt is None:
|
||||
return 1.0
|
||||
return float(getattr(mt, "speed", 1.0) or 1.0)
|
||||
|
||||
def _get_micro_bgm_offset(self) -> float:
|
||||
"""#1970 PR2:读取本视频 BGM 起始偏移(秒),无 BGM/禁用时为 0。"""
|
||||
if not self.plan.config:
|
||||
return 0.0
|
||||
try:
|
||||
count = len([c for c in (self.plan.clips or []) if getattr(c, "clip_type", "main") != "audio"])
|
||||
plan = self._get_micro_transform_plan(count)
|
||||
if plan:
|
||||
return round(float(plan.bgm_start_offset or 0.0), 3)
|
||||
except Exception:
|
||||
logger.debug("微变换 BGM 偏移读取失败,按 0 处理: plan_id=%s", getattr(self.plan, "id", "?"))
|
||||
return 0.0
|
||||
|
||||
def render(self) -> RenderResult:
|
||||
"""执行渲染,返回 RenderResult.
|
||||
@@ -316,6 +405,9 @@ class UnifiedRenderService:
|
||||
ctx = RenderContext(work_dir=self.work_dir, plan_id=self.plan.id)
|
||||
from video_processing.bgm_mixer import BGMConfig, mix_bgm_with_main
|
||||
|
||||
_bgm_off = self._get_micro_bgm_offset()
|
||||
if _bgm_off and not (bgm_config or {}).get("audio_offset"):
|
||||
bgm_config = {**bgm_config, "audio_offset": _bgm_off}
|
||||
bgm_cfg = BGMConfig.from_config_dict(self.bgm_path, bgm_config)
|
||||
# 从直通输出中提取音频
|
||||
main_audio_path = self.work_dir / f"pass_through_audio_{self.plan.id}.aac"
|
||||
@@ -365,6 +457,7 @@ class UnifiedRenderService:
|
||||
bgm_path=self.bgm_path,
|
||||
bgm_config=bgm_config,
|
||||
audio_tracks_config=audio_tracks_config,
|
||||
bgm_start_offset=self._get_micro_bgm_offset(),
|
||||
)
|
||||
t_audio_end = time.time()
|
||||
audio_mix_ms = int((t_audio_end - t_audio_start) * 1000)
|
||||
@@ -1112,6 +1205,28 @@ class UnifiedRenderService:
|
||||
if ass_path is not None:
|
||||
return False, "有字幕叠加"
|
||||
|
||||
# #1970 PR2:片段级微变换(变速/hflip/亮度/对比度/饱和度)需要重编码
|
||||
try:
|
||||
_video_sources = [c for c in (self.clips or []) if getattr(c, "clip_type", "main") != "audio"]
|
||||
_ordinal = -1
|
||||
for _i, _c in enumerate(_video_sources):
|
||||
if getattr(_c, "id", None) == getattr(clip, "clip_id", None):
|
||||
_ordinal = _i
|
||||
break
|
||||
_mt_plan = self._get_micro_transform_plan(len(_video_sources))
|
||||
if _mt_plan and 0 <= _ordinal < len(_mt_plan.clips):
|
||||
_mt = _mt_plan.clips[_ordinal]
|
||||
if (
|
||||
abs(UnifiedRenderService._micro_speed_factor(_mt) - 1.0) >= 1e-6
|
||||
or (_mt.hflip and not _mt.has_text)
|
||||
or abs(_mt.brightness) > 1e-4
|
||||
or abs(_mt.contrast - 1.0) > 1e-4
|
||||
or abs(_mt.saturation - 1.0) > 1e-4
|
||||
):
|
||||
return False, "启用了片段级微变换"
|
||||
except Exception:
|
||||
logger.debug("stream copy 微变换门控检查异常,按可 copy 处理", exc_info=True)
|
||||
|
||||
# 有调速 → 需要重编码 → 不能 copy
|
||||
speed = UnifiedRenderService._clip_speed(clip)
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
@@ -1318,11 +1433,16 @@ class UnifiedRenderService:
|
||||
|
||||
# 视觉扰动(plan 级别,直通模式同样适用)
|
||||
vp = self._get_visual_perturbation()
|
||||
# #1970 PR2:单片段直通;计划按源视频片段数构建,序号取 config._micro_index
|
||||
_src_video_count = len([c for c in (self.clips or []) if getattr(c, "clip_type", "main") != "audio"])
|
||||
mt_plan = self._get_micro_transform_plan(max(1, _src_video_count))
|
||||
_mi = int(clip.config.get("_micro_index", 0)) if isinstance(clip.config, dict) else 0
|
||||
mt = mt_plan.clips[_mi] if mt_plan and 0 <= _mi < len(mt_plan.clips) else None
|
||||
|
||||
# 调速 — 与 filter_complex 路径一致(叠加视觉扰动 speed_factor)
|
||||
# 调速 — 与 filter_complex 路径一致(叠加视觉扰动 speed_factor 与 #1970 微变换速度)
|
||||
speed = UnifiedRenderService._clip_speed(clip)
|
||||
vp_speed = vp.get("speed_factor", 1.0) if vp else 1.0
|
||||
effective_speed = speed * vp_speed
|
||||
effective_speed = speed * vp_speed # 微变换速度已烘焙进 playback_speed
|
||||
if abs(effective_speed - 1.0) >= 1e-6:
|
||||
filters.append(f"setpts=PTS/{effective_speed:.4f}")
|
||||
|
||||
@@ -1336,6 +1456,8 @@ class UnifiedRenderService:
|
||||
# 视觉扰动:hflip(在 scale 之前)
|
||||
if vp:
|
||||
self._apply_visual_perturbation_pre_scale(filters, vp)
|
||||
# #1970 PR2:片段级 hflip(P1 保守:有文字/无法判定时不翻转)
|
||||
UnifiedRenderService._apply_micro_hflip(filters, mt)
|
||||
|
||||
# scale + pad(等比缩放+留黑边)
|
||||
if role in ("overlay", "corner_voice"):
|
||||
@@ -1354,6 +1476,8 @@ class UnifiedRenderService:
|
||||
# 视觉扰动:zoom + brightness(在 scale+pad 之后、调色之前)
|
||||
if vp:
|
||||
self._apply_visual_perturbation_post_scale(filters, vp)
|
||||
# #1970 PR2:片段级亮度/对比度/饱和度微调
|
||||
UnifiedRenderService._apply_micro_transform_video(filters, mt)
|
||||
|
||||
# 调色滤镜
|
||||
color_grade = ColorGradeConfig.from_dict(clip.config.get("color_grade"))
|
||||
@@ -1450,7 +1574,8 @@ class UnifiedRenderService:
|
||||
# 音频调速(在降噪之后、音量之前,与 render_audio.py concat 路径保持一致)
|
||||
# SpeedEngine.build_audio_filter 内部已实现多级 atempo 串联,
|
||||
# 自动处理超出 [0.5, 2.0] 范围的速度(如 0.25x → atempo=0.5,atempo=0.5)。
|
||||
speed = UnifiedRenderService._clip_speed(clip)
|
||||
# #1970 PR2:叠加片段微变换速度因子,保持音画同步。
|
||||
speed = UnifiedRenderService._clip_speed(clip) # 微变换速度已烘焙进 playback_speed
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
try:
|
||||
from video_processing.speed_engine import SpeedConfig, SpeedEngine
|
||||
@@ -1522,11 +1647,20 @@ class UnifiedRenderService:
|
||||
支持多段裁剪:一个 clip 配置了 trim_segments 时会展开为多个 ResolvedClip。
|
||||
"""
|
||||
resolved: list[ResolvedClip] = []
|
||||
# #1970 PR2:预建片段级微变换计划,按源视频片段序号取速度因子,
|
||||
# 烘焙进 playback_speed,保证视频 setpts 与音频 atempo 一致。
|
||||
video_source_clips = [c for c in self.clips if getattr(c, "clip_type", "main") != "audio"]
|
||||
mt_plan = self._get_micro_transform_plan(len(video_source_clips))
|
||||
_video_ordinal = {id(c): i for i, c in enumerate(video_source_clips)}
|
||||
|
||||
for clip in self.clips:
|
||||
asset_id = clip.asset_id
|
||||
if not asset_id:
|
||||
logger.warning("片段无素材: clip_id=%s", clip.id)
|
||||
continue
|
||||
_mt_idx = _video_ordinal.get(id(clip), -1)
|
||||
_mt = mt_plan.clips[_mt_idx] if mt_plan and 0 <= _mt_idx < len(mt_plan.clips) else None
|
||||
_micro_speed = UnifiedRenderService._micro_speed_factor(_mt)
|
||||
|
||||
local_path = self.asset_path_map.get(asset_id)
|
||||
if local_path is None or not local_path.exists():
|
||||
@@ -1555,7 +1689,7 @@ class UnifiedRenderService:
|
||||
seg_duration = seg.trim.duration
|
||||
|
||||
# 多段裁剪:如果段的时长超过素材实际时长,减速补偿
|
||||
seg_speed = configured_speed
|
||||
seg_speed = configured_speed * _micro_speed
|
||||
if actual_duration > 0 and seg_duration > actual_duration + 0.05:
|
||||
seg_speed = max(0.25, round(configured_speed * actual_duration / seg_duration, 4))
|
||||
logger.info(
|
||||
@@ -1578,7 +1712,7 @@ class UnifiedRenderService:
|
||||
transition_effect=clip.transition_effect or "cut",
|
||||
transition_duration=getattr(clip, "transition_duration", 0.0) or 0.0,
|
||||
playback_speed=seg_speed,
|
||||
config={**clip_config, "_segment_id": seg.segment_id},
|
||||
config={**clip_config, "_segment_id": seg.segment_id, "_micro_index": _mt_idx},
|
||||
actual_duration=actual_duration,
|
||||
trim_config=seg.trim,
|
||||
)
|
||||
@@ -1633,12 +1767,13 @@ class UnifiedRenderService:
|
||||
avail_in_asset,
|
||||
freeze_seconds,
|
||||
)
|
||||
final_speed = configured_speed
|
||||
final_speed = configured_speed * _micro_speed
|
||||
|
||||
# freeze 标记写入 config,供视频 tpad / 音频 apad 读取
|
||||
resolved_config = dict(clip_config)
|
||||
if freeze_seconds > 0:
|
||||
resolved_config["_freeze_seconds"] = freeze_seconds
|
||||
resolved_config["_micro_index"] = _mt_idx
|
||||
|
||||
rc = ResolvedClip(
|
||||
clip_id=clip.id,
|
||||
@@ -1754,9 +1889,15 @@ class UnifiedRenderService:
|
||||
preprocessed_labels: list[str] = []
|
||||
# 视觉扰动(plan 级别,所有 clip 共享同一套扰动参数)
|
||||
vp = self._get_visual_perturbation()
|
||||
# #1970 PR2:片段级微变换(每片段独立参数,dedup_enabled=False 时为 None)
|
||||
# 计划按源视频片段数构建,trim 多段展开时各段通过 config._micro_index 找参数
|
||||
_src_video_count = len([c for c in (self.clips or []) if getattr(c, "clip_type", "main") != "audio"])
|
||||
mt_plan = self._get_micro_transform_plan(_src_video_count)
|
||||
for i, clip in enumerate(all_clips):
|
||||
label = f"v{i}"
|
||||
role = _resolve_layer_role(clip.clip_type, clip.config)
|
||||
_mi = int(clip.config.get("_micro_index", i)) if isinstance(clip.config, dict) else i
|
||||
mt = mt_plan.clips[_mi] if mt_plan and 0 <= _mi < len(mt_plan.clips) else None
|
||||
|
||||
filters: list[str] = []
|
||||
|
||||
@@ -1774,10 +1915,10 @@ class UnifiedRenderService:
|
||||
filters.append(f"trim=duration={trim_dur:.3f}")
|
||||
filters.append("setpts=PTS-STARTPTS")
|
||||
|
||||
# 调速 — 基于 setpts 改变播放速度(叠加视觉扰动 speed_factor)
|
||||
# 调速 — 基于 setpts 改变播放速度(叠加视觉扰动 speed_factor 与 #1970 微变换速度)
|
||||
speed = UnifiedRenderService._clip_speed(clip)
|
||||
vp_speed = vp.get("speed_factor", 1.0) if vp else 1.0
|
||||
effective_speed = speed * vp_speed
|
||||
effective_speed = speed * vp_speed # 微变换速度已烘焙进 playback_speed
|
||||
if abs(effective_speed - 1.0) >= 1e-6:
|
||||
filters.append(f"setpts=PTS/{effective_speed:.4f}")
|
||||
|
||||
@@ -1791,6 +1932,8 @@ class UnifiedRenderService:
|
||||
# 视觉扰动:hflip(在 scale 之前,翻转原始画面)
|
||||
if vp:
|
||||
self._apply_visual_perturbation_pre_scale(filters, vp)
|
||||
# #1970 PR2:片段级 hflip(P1 保守:有文字/无法判定时不翻转)
|
||||
UnifiedRenderService._apply_micro_hflip(filters, mt)
|
||||
|
||||
# scale
|
||||
if role in ("overlay", "corner_voice"):
|
||||
@@ -1809,6 +1952,8 @@ class UnifiedRenderService:
|
||||
# 视觉扰动:zoom + brightness(在 scale+pad 之后、调色之前)
|
||||
if vp:
|
||||
self._apply_visual_perturbation_post_scale(filters, vp)
|
||||
# #1970 PR2:片段级亮度/对比度/饱和度微调
|
||||
UnifiedRenderService._apply_micro_transform_video(filters, mt)
|
||||
|
||||
# 调色滤镜(每个 clip 独立的 color grade 配置)
|
||||
color_grade = ColorGradeConfig.from_dict(clip.config.get("color_grade"))
|
||||
|
||||
@@ -27,6 +27,11 @@ celery_app.conf.broker_transport_options = {"visibility_timeout": 4 * 60 * 60}
|
||||
celery_app.conf.imports = (
|
||||
"worker_app.tasks.health",
|
||||
"worker_app.tasks.ingest",
|
||||
"worker_app.tasks.atom_clips",
|
||||
# #1970 片段级 AI 标签:必须显式 import 注册,否则 worker 报
|
||||
# "Received unregistered task of type 'worker.tag_atom_clip'"
|
||||
"worker_app.tasks.atom_clip_tagging",
|
||||
"worker_app.tasks.backfill_atom_clip_tags",
|
||||
"worker_app.tasks.classification",
|
||||
"worker_app.tasks.generation",
|
||||
"worker_app.tasks.voice_extraction",
|
||||
|
||||
@@ -53,12 +53,25 @@ def __getattr__(name: str):
|
||||
from .batch_thumbnail import batch_generate_thumbnails
|
||||
|
||||
return batch_generate_thumbnails
|
||||
elif name == "generate_atom_clips":
|
||||
from .atom_clips import generate_atom_clips
|
||||
|
||||
return generate_atom_clips
|
||||
elif name == "tag_atom_clip_task":
|
||||
from .atom_clip_tagging import tag_atom_clip_task
|
||||
|
||||
return tag_atom_clip_task
|
||||
elif name == "backfill_atom_clip_tags":
|
||||
from .backfill_atom_clip_tags import backfill_atom_clip_tags
|
||||
|
||||
return backfill_atom_clip_tags
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"batch_generate_thumbnails",
|
||||
"classify_asset",
|
||||
"generate_atom_clips",
|
||||
"generate_video",
|
||||
"healthcheck",
|
||||
"ingest_asset",
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
"""片段级 AI 标签 Celery 任务 — #1970 智能剪辑流程重构 P2.
|
||||
|
||||
为单个 atom_clip 调用视觉 AI 生成结构化标签,并更新到 ai_tags 字段。
|
||||
失败不阻断流程(降级为仅继承素材标签)。
|
||||
|
||||
任务名:worker.tag_atom_clip
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.atom_clip_tagger import tag_atom_clip
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.tag_atom_clip", bind=True, max_retries=2, default_retry_delay=10)
|
||||
def tag_atom_clip_task(self, atom_clip_id: str, force: bool = False) -> dict:
|
||||
"""为单个原子片段生成 AI 标签.
|
||||
|
||||
Args:
|
||||
atom_clip_id: 原子片段 ID。
|
||||
force: True 时允许覆盖只有 inherited_tags 的降级记录
|
||||
(视觉 API 曾失败写入的占位标签,#1970)。
|
||||
已有完整标签(含 has_text)始终跳过,保证幂等。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:status / clip_id / ai_tags(部分字段)。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
|
||||
clip = atom_repo.find_by_id(atom_clip_id)
|
||||
if clip is None:
|
||||
return {"status": "skipped", "reason": "clip not found", "clip_id": atom_clip_id}
|
||||
|
||||
# 已有完整标签则跳过(幂等);force 仅放行缺失 has_text 的降级记录
|
||||
if clip.ai_tags is not None:
|
||||
has_real_tags = isinstance(clip.ai_tags, dict) and "has_text" in clip.ai_tags
|
||||
if has_real_tags or not force:
|
||||
return {"status": "skipped", "reason": "already tagged", "clip_id": atom_clip_id}
|
||||
|
||||
# 获取素材信息
|
||||
asset = asset_repo.find_by_id(clip.asset_id)
|
||||
if asset is None:
|
||||
return {"status": "skipped", "reason": "asset not found", "clip_id": atom_clip_id}
|
||||
|
||||
# 获取视频可访问 URL
|
||||
storage = get_shared_storage_service()
|
||||
video_url = storage.get_download_url(asset.storage_key, expires_seconds=3600)
|
||||
|
||||
# 初始化客户端
|
||||
doubao_client = get_doubao_client()
|
||||
mediakit_client = get_mediakit_client()
|
||||
|
||||
# 调用 tagger
|
||||
ai_tags = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url=video_url,
|
||||
doubao_client=doubao_client,
|
||||
mediakit_client=mediakit_client,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
# 更新数据库
|
||||
atom_repo.update_ai_tags(atom_clip_id, ai_tags)
|
||||
|
||||
logger.info(
|
||||
"[atom_clip_tagging] clip_id=%s ai_tags=%s",
|
||||
atom_clip_id,
|
||||
{k: v for k, v in ai_tags.items() if k != "inherited_tags"},
|
||||
)
|
||||
return {
|
||||
"status": "completed",
|
||||
"clip_id": atom_clip_id,
|
||||
"has_ai_tags": any(v for k, v in ai_tags.items() if k != "inherited_tags" and v),
|
||||
}
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception("[atom_clip_tagging] clip_id=%s 失败: %s", atom_clip_id, exc)
|
||||
# 可重试异常
|
||||
if self.request.retries < self.max_retries:
|
||||
raise self.retry(exc=exc) from None
|
||||
return {"status": "failed", "clip_id": atom_clip_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -0,0 +1,109 @@
|
||||
"""素材原子切片 Celery 任务 — #1970 智能剪辑流程重构 P1.
|
||||
|
||||
素材入库预处理完成(ingest 置 READY)后异步触发:
|
||||
根据素材时长和已缓存的 scdet 切换点计算原子片段并落库。
|
||||
失败不阻断素材入库主流程(atom_clips 未就绪时选片有内存兜底)。
|
||||
|
||||
P2 增强:切片完成后自动链式触发 AI 标签任务(每个 clip 一个 tag_atom_clip 任务)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.atom_clip_service import compute_atom_clips
|
||||
from packages.domain.plan_generator_utils import extract_scene_points_from_metadata
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.generate_atom_clips")
|
||||
def generate_atom_clips(asset_id: str) -> dict:
|
||||
"""为单条视频素材生成原子片段。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:status / asset_id / clips_count。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
|
||||
asset = asset_repo.find_by_id(asset_id)
|
||||
if asset is None:
|
||||
return {"status": "skipped", "reason": "asset not found", "asset_id": asset_id}
|
||||
|
||||
# 仅视频素材切片
|
||||
if asset.mime_type and not asset.mime_type.startswith("video/"):
|
||||
return {"status": "skipped", "reason": "not a video", "asset_id": asset_id}
|
||||
if not asset.duration or asset.duration <= 0:
|
||||
return {"status": "skipped", "reason": "invalid duration", "asset_id": asset_id}
|
||||
|
||||
# 已生成过则幂等跳过(重新切片需先显式删除)
|
||||
existing = atom_repo.count_by_asset(asset_id)
|
||||
if existing > 0:
|
||||
return {
|
||||
"status": "skipped",
|
||||
"reason": "already generated",
|
||||
"asset_id": asset_id,
|
||||
"clips_count": existing,
|
||||
}
|
||||
|
||||
scene_points = extract_scene_points_from_metadata(asset.metadata)
|
||||
# P1 阶段继承素材的标签 ID;片段级语义标签是 P2 功能
|
||||
tags = list(getattr(asset, "tag_ids", []) or [])
|
||||
|
||||
clips = compute_atom_clips(
|
||||
asset_id=asset_id,
|
||||
duration=float(asset.duration),
|
||||
scene_change_points=scene_points,
|
||||
tags=tags,
|
||||
)
|
||||
if not clips:
|
||||
return {"status": "skipped", "reason": "no clips computed", "asset_id": asset_id}
|
||||
|
||||
atom_repo.batch_create(clips)
|
||||
logger.info(
|
||||
"[atom_clips] asset_id=%s 生成 %d 个原子片段",
|
||||
asset_id,
|
||||
len(clips),
|
||||
)
|
||||
|
||||
# P2 增强:链式触发 AI 标签任务(每个 clip 一个异步任务)
|
||||
_dispatch_tagging_tasks(clips)
|
||||
|
||||
return {"status": "completed", "asset_id": asset_id, "clips_count": len(clips)}
|
||||
except Exception as exc: # noqa: BLE001 - 后台任务兜底,失败不阻断主流程
|
||||
db.rollback()
|
||||
logger.exception("[atom_clips] asset_id=%s 生成失败: %s", asset_id, exc)
|
||||
return {"status": "failed", "asset_id": asset_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _dispatch_tagging_tasks(clips: list) -> None:
|
||||
"""为每个新建片段发送 AI 标签异步任务.
|
||||
|
||||
失败不阻断(标签任务是锦上添花,不影响核心流程)。
|
||||
"""
|
||||
try:
|
||||
for clip in clips:
|
||||
celery_app.send_task(
|
||||
"worker.tag_atom_clip",
|
||||
args=[clip.id],
|
||||
)
|
||||
logger.info(
|
||||
"[atom_clips] 已发送 %d 个 AI 标签任务",
|
||||
len(clips),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[atom_clips] 发送 AI 标签任务失败(不影响切片结果): %s",
|
||||
e,
|
||||
)
|
||||
@@ -0,0 +1,106 @@
|
||||
"""批量回填 AI 标签 Celery 任务 — #1970 智能剪辑流程重构 P2.
|
||||
|
||||
查找所有 ai_tags IS NULL 的 atom_clips,分批触发 tag_atom_clip 任务。
|
||||
可通过 API 路由触发(管理员权限)。
|
||||
|
||||
任务名:worker.backfill_atom_clip_tags
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
# 默认批量参数
|
||||
DEFAULT_BATCH_SIZE = 10
|
||||
DEFAULT_BATCH_INTERVAL = 5 # 秒
|
||||
|
||||
|
||||
@celery_app.task(name="worker.backfill_atom_clip_tags")
|
||||
def backfill_atom_clip_tags(
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
batch_interval: int = DEFAULT_BATCH_INTERVAL,
|
||||
max_clips: int = 0,
|
||||
force: bool = False,
|
||||
) -> dict:
|
||||
"""批量回填未打标的 atom_clips.
|
||||
|
||||
Args:
|
||||
batch_size: 每批处理数量,默认 10。
|
||||
batch_interval: 每批间隔秒数,默认 5。
|
||||
max_clips: 最大处理总数,0 表示不限。
|
||||
force: True 时连同只有 inherited_tags 的降级记录一起强制重打
|
||||
(视觉 API 曾失败、DOUBAO_VISION_MODEL 修复后重跑用,#1970)。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:total_submitted / batches。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
total_submitted = 0
|
||||
batches = 0
|
||||
|
||||
while True:
|
||||
# 查找未打标的片段
|
||||
remaining = max_clips - total_submitted if max_clips > 0 else batch_size
|
||||
fetch_limit = min(batch_size, remaining) if max_clips > 0 else batch_size
|
||||
|
||||
untagged = atom_repo.find_untagged(limit=fetch_limit, include_downgraded=force)
|
||||
if not untagged:
|
||||
break
|
||||
|
||||
# 逐个发送 tag 任务
|
||||
for clip in untagged:
|
||||
try:
|
||||
celery_app.send_task(
|
||||
"worker.tag_atom_clip",
|
||||
args=[clip.id],
|
||||
kwargs={"force": force},
|
||||
)
|
||||
total_submitted += 1
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[backfill] 提交任务失败 clip_id=%s: %s",
|
||||
clip.id,
|
||||
e,
|
||||
)
|
||||
|
||||
batches += 1
|
||||
logger.info(
|
||||
"[backfill] 第 %d 批完成,已提交 %d 个任务",
|
||||
batches,
|
||||
total_submitted,
|
||||
)
|
||||
|
||||
# 检查是否达到上限
|
||||
if max_clips > 0 and total_submitted >= max_clips:
|
||||
break
|
||||
|
||||
# 批间间隔
|
||||
time.sleep(batch_interval)
|
||||
|
||||
logger.info(
|
||||
"[backfill] 回填完成: total_submitted=%d batches=%d",
|
||||
total_submitted,
|
||||
batches,
|
||||
)
|
||||
return {
|
||||
"status": "completed",
|
||||
"total_submitted": total_submitted,
|
||||
"batches": batches,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.exception("[backfill] 回填失败: %s", exc)
|
||||
return {"status": "failed", "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -890,25 +890,51 @@ def generate_video(self, task_id: str) -> dict:
|
||||
_flush_logs(task_id, gen_task)
|
||||
_update_task_progress(task_id, 80, "渲染完成")
|
||||
|
||||
# ── 3.5 随机边缘裁剪降重(#1664) ──────────────────────────
|
||||
from video_processing.ffmpeg_utils import random_edge_crop
|
||||
|
||||
# ── 3.5 随机边缘裁剪降重(#1664;#1970 dedup_enabled=False 时跳过) ──
|
||||
_dedup_enabled = True
|
||||
try:
|
||||
cropped_path = random_edge_crop(output_path)
|
||||
if cropped_path != output_path:
|
||||
output_path = cropped_path
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("边缘裁剪", "已应用随机 2-5% 边缘裁剪降重")
|
||||
_flush_logs(task_id, gen_task)
|
||||
logger.info("[task_id=%s] 随机边缘裁剪完成: %s", task_id, output_path)
|
||||
except Exception as crop_err:
|
||||
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
|
||||
|
||||
with SessionLocal() as _dedup_db:
|
||||
_plan_row = (
|
||||
_dedup_db.query(EditPlanModel.config)
|
||||
.filter(EditPlanModel.id == current_plan_id)
|
||||
.first()
|
||||
)
|
||||
if _plan_row is not None:
|
||||
_cfg = _plan_row[0] if isinstance(_plan_row[0], dict) else {}
|
||||
_dedup_enabled = bool(_cfg.get("dedup_enabled", True))
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"[task_id=%s] 随机边缘裁剪失败,使用原始视频继续: %s",
|
||||
"[task_id=%s] 读取 plan dedup_enabled 失败,按开启处理",
|
||||
task_id,
|
||||
crop_err,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
if not _dedup_enabled:
|
||||
logger.info("[task_id=%s] dedup_enabled=False,跳过边缘裁剪与微变换", task_id)
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("降重", "已关闭边缘裁剪与微变换(确定性渲染)")
|
||||
_flush_logs(task_id, gen_task)
|
||||
else:
|
||||
from video_processing.ffmpeg_utils import random_edge_crop
|
||||
|
||||
try:
|
||||
cropped_path = random_edge_crop(output_path)
|
||||
if cropped_path != output_path:
|
||||
output_path = cropped_path
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("边缘裁剪", "已应用随机 2-5% 边缘裁剪降重")
|
||||
_flush_logs(task_id, gen_task)
|
||||
logger.info("[task_id=%s] 随机边缘裁剪完成: %s", task_id, output_path)
|
||||
except Exception as crop_err:
|
||||
logger.warning(
|
||||
"[task_id=%s] 随机边缘裁剪失败,使用原始视频继续: %s",
|
||||
task_id,
|
||||
crop_err,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
# ── 4. 上传 OSS(不落库) ───────────────────────────────
|
||||
_update_task_progress(task_id, 85, "开始上传")
|
||||
file_url, _storage_key = _upload_rendered_video(
|
||||
|
||||
@@ -808,6 +808,21 @@ def ingest_asset(job_id: str) -> dict:
|
||||
|
||||
db.commit()
|
||||
|
||||
# ── #1970 素材原子切片:视频 READY 后异步触发,失败不阻断入库 ──
|
||||
# atom_clips 未就绪时选片逻辑有内存兜底(compute_fallback_clips)。
|
||||
try:
|
||||
if media_type == "video" and float(asset.duration or 0) > 0:
|
||||
celery_app.send_task(
|
||||
"worker.generate_atom_clips",
|
||||
args=[asset.id],
|
||||
)
|
||||
except Exception as atom_err: # noqa: BLE001
|
||||
logger.warning(
|
||||
"触发原子切片任务失败(不影响入库): asset_id=%s err=%s",
|
||||
asset.id,
|
||||
atom_err,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"job_id": job.id,
|
||||
|
||||
@@ -234,9 +234,32 @@ DOUBAO_TIMEOUT=60
|
||||
# 最大重试次数
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
|
||||
# 视觉模型 Endpoint ID(支持图片/视频理解的模型)
|
||||
DOUBAO_VISION_MODEL=${DOUBAO_VISION_MODEL}
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
|
||||
WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
|
||||
WECHAT_OPEN_REDIRECT_URI=https://saas.xiaoxiajianji.com/auth/wechat/callback
|
||||
|
||||
# 抖音 cookies 文件路径(yt-dlp 已废弃,保留兼容)
|
||||
DOUYIN_COOKIES_FILE=/app/configs/douyin_cookies.txt
|
||||
DOUYIN_DEBUG_ERRORS=false
|
||||
|
||||
|
||||
# ==================== 抖音视频解析(三层兜底)====================
|
||||
# P0: App Feed API(免费,零 Key)— 内置,无需配置
|
||||
# P1: TikHub API(付费,https://tikhub.io)
|
||||
TIKHUB_API_KEY=${TIKHUB_API_KEY}
|
||||
# P2: apizero.cn(国内付费,https://apizero.cn)
|
||||
APIZERO_API_KEY=${APIZERO_API_KEY}
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询) ====================
|
||||
GPU_WORKER_TOKEN=${GPU_WORKER_TOKEN}
|
||||
GPU_TASK_TIMEOUT_SECONDS=900
|
||||
USE_GPU_LIPSYNC=false
|
||||
GPU_LIPSYNC_POLL_INTERVAL=5
|
||||
GPU_LIPSYNC_WAIT_TIMEOUT=1200
|
||||
GPU_WORKER_STALE_SECONDS=300
|
||||
|
||||
@@ -251,9 +251,32 @@ DOUBAO_TIMEOUT=60
|
||||
# 最大重试次数
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
|
||||
# 视觉模型 Endpoint ID(支持图片/视频理解的模型)
|
||||
DOUBAO_VISION_MODEL=${DOUBAO_VISION_MODEL}
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
|
||||
WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
|
||||
WECHAT_OPEN_REDIRECT_URI=https://staging.xiaoxiajianji.com/auth/wechat/callback
|
||||
|
||||
# 抖音 cookies 文件路径(yt-dlp 已废弃,保留兼容)
|
||||
DOUYIN_COOKIES_FILE=/app/configs/douyin_cookies.txt
|
||||
DOUYIN_DEBUG_ERRORS=false
|
||||
|
||||
|
||||
# ==================== 抖音视频解析(三层兜底)====================
|
||||
# P0: App Feed API(免费,零 Key)— 内置,无需配置
|
||||
# P1: TikHub API(付费,https://tikhub.io)
|
||||
TIKHUB_API_KEY=${TIKHUB_API_KEY}
|
||||
# P2: apizero.cn(国内付费,https://apizero.cn)
|
||||
APIZERO_API_KEY=${APIZERO_API_KEY}
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询) ====================
|
||||
GPU_WORKER_TOKEN=${GPU_WORKER_TOKEN}
|
||||
GPU_TASK_TIMEOUT_SECONDS=900
|
||||
USE_GPU_LIPSYNC=true
|
||||
GPU_LIPSYNC_POLL_INTERVAL=5
|
||||
GPU_LIPSYNC_WAIT_TIMEOUT=1200
|
||||
GPU_WORKER_STALE_SECONDS=300
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# Netscape HTTP Cookie File
|
||||
# 抖音 cookies 占位。CI 部署时会通过 scp 上传真实 cookies。
|
||||
# 若本文件被使用说明 CI 上传失败,请检查 deploy-staging job。
|
||||
@@ -0,0 +1,30 @@
|
||||
# ============================================================
|
||||
# MuseTalk GPU Worker 环境变量
|
||||
# 部署到 RTX2060 电脑后,复制为 .env 并修改值
|
||||
# ============================================================
|
||||
|
||||
# SaaS API 基础 URL(staging / production)
|
||||
API_BASE_URL=https://staging-api.xiaoxiajianji.com
|
||||
# API_BASE_URL=https://api.xiaoxiajianji.com # 生产
|
||||
|
||||
# 长期 API Token,必须与服务端 GPU_WORKER_TOKEN 一致(找后端拿)
|
||||
GPU_WORKER_TOKEN=replace-with-real-token
|
||||
|
||||
# 本机 Worker 唯一 ID(默认自动生成 hostname+MAC 后4位,可手动指定)
|
||||
# WORKER_ID=rtx2060-0193
|
||||
|
||||
# 本地 MuseTalk 地址(默认 http://127.0.0.1:7861)
|
||||
MUSE_TALK_URL=http://127.0.0.1:7861
|
||||
|
||||
# 轮询/心跳/超时(秒)
|
||||
POLL_INTERVAL=5
|
||||
HEARTBEAT_INTERVAL=15
|
||||
# 下载/推理/上传 HTTP 超时,需与服务端 GPU_TASK_TIMEOUT_SECONDS 对齐(默认 900)
|
||||
REQUEST_TIMEOUT=900
|
||||
|
||||
# 单个任务本地最大重试次数(仅网络/MuseTalk 瞬时错误才重试,默认 1)
|
||||
TASK_MAX_RETRY=1
|
||||
# 推理期间任务心跳间隔(秒,独立线程,无需改动)
|
||||
TASK_HEARTBEAT_INTERVAL=30
|
||||
# 输入视频最短时长(秒),小于则直接上报失败,不调用 MuseTalk
|
||||
MIN_VIDEO_DURATION_SECONDS=3
|
||||
@@ -0,0 +1,353 @@
|
||||
# MuseTalk GPU Worker 部署指南
|
||||
|
||||
本目录包含两个组件:
|
||||
|
||||
1. **gpu_worker.py**:反向轮询客户端,部署在 RTX2060 本地,轮询 SaaS API 拉取口型任务,调用本地 MuseTalk 服务推理,上传结果回 SaaS。
|
||||
2. **musetalk_server.py**:MuseTalk Flask HTTP 服务端,接收 gpu_worker.py 的推理请求,调用 MuseTalk 模型生成口型同步视频。
|
||||
|
||||
---
|
||||
|
||||
## 一、环境准备
|
||||
|
||||
### 1.1 硬件要求
|
||||
|
||||
- GPU: NVIDIA RTX 2060 或更高(显存 ≥ 6GB)
|
||||
- CUDA: 11.8+
|
||||
- Python: 3.10+
|
||||
- ffmpeg: 需安装并加入 PATH
|
||||
|
||||
### 1.2 安装依赖
|
||||
|
||||
```bash
|
||||
cd deploy/gpu_worker
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 二、MuseTalk 服务端部署(musetalk_server.py)
|
||||
|
||||
### 2.1 配置环境变量
|
||||
|
||||
复制 `.env.example` 为 `.env`,修改配置:
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
vim .env
|
||||
```
|
||||
|
||||
关键配置:
|
||||
|
||||
| 变量 | 说明 | 默认值 |
|
||||
|------|------|--------|
|
||||
| `MUSE_PORT` | 监听端口 | `7861` |
|
||||
| `MUSE_INFERENCE_TIMEOUT` | 推理超时秒数 | `600` |
|
||||
| `MUSE_VIDEO_MAX_MB` | 视频上传大小限制 MB | `100` |
|
||||
| `MUSE_AUDIO_MAX_MB` | 音频上传大小限制 MB | `20` |
|
||||
| `MUSE_DEFAULT_FPS` | 视频 fps 兜底值 | `25.0` |
|
||||
| `MUSE_TEMP_DIR` | 临时文件目录 | `/tmp/musetalk_$$` |
|
||||
| `MUSE_VIDEO_ENCODER` | 兜底循环视频时的编码器:`auto`(优先 h264_nvenc,失败回退 libx264)/`h264_nvenc`/`libx264` | `auto` |
|
||||
|
||||
### 2.2 更新部署(v2 性能修复,必做)
|
||||
|
||||
> ⚠️ 2026-09-20 v2 架构:修复 16 倍性能回归。旧版在推理前 loop 视频导致 MuseTalk 处理帧数翻倍、RTX2060 推理 >200s、nginx 504。**必须重新拉取并重启**:
|
||||
|
||||
```bash
|
||||
# 在 RTX2060 上备份旧文件并拉取新版本
|
||||
cp ~/projects/MuseTalk/musetalk_server.py ~/projects/MuseTalk/musetalk_server.py.bak
|
||||
wget -O ~/projects/MuseTalk/musetalk_server.py \
|
||||
"https://git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/raw/branch/develop/deploy/gpu_worker/musetalk_server.py"
|
||||
|
||||
# 重启服务
|
||||
sudo systemctl restart musetalk-server
|
||||
sudo systemctl status musetalk-server
|
||||
curl http://127.0.0.1:7861/health
|
||||
```
|
||||
|
||||
v2 架构核心变化:
|
||||
|
||||
- **MuseTalk 直传全量音频**:不再在推理前用 ffmpeg 循环视频。MuseTalk 原生支持长音频输入,内部自动循环视频帧。推理时间不变(~14s/5s 视频)
|
||||
- **ffmpeg 只做快速封装**:`-c:v copy -c:a aac -shortest`,秒级完成,不重编码
|
||||
- **循环仅兜底**:仅当 MuseTalk 输出画面短于音频时(极端情况),才 `-stream_loop` + NVENC 兜底
|
||||
- **删除 `MUSE_ENABLE_VIDEO_LOOP`**:不再需要此开关,MuseTalk 原生处理
|
||||
|
||||
### 2.3 启动服务
|
||||
|
||||
```bash
|
||||
# 前台运行(调试用)
|
||||
python musetalk_server.py
|
||||
|
||||
# 后台运行(生产用 systemd)
|
||||
sudo systemctl start musetalk-server
|
||||
sudo systemctl enable musetalk-server
|
||||
```
|
||||
|
||||
### 2.4 验证健康检查
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:7861/health
|
||||
```
|
||||
|
||||
应返回:
|
||||
|
||||
```json
|
||||
{
|
||||
"status": "healthy",
|
||||
"gpu": {
|
||||
"gpu_name": "NVIDIA GeForce RTX 2060",
|
||||
"memory_total_mb": 6144,
|
||||
"memory_used_mb": 1024,
|
||||
"memory_free_mb": 5120
|
||||
},
|
||||
"current_task": {
|
||||
"task_id": null,
|
||||
"running": false,
|
||||
"elapsed_seconds": 0.0
|
||||
},
|
||||
"timestamp": 1700000000.0
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 三、GPU Worker 客户端部署(gpu_worker.py)
|
||||
|
||||
### 3.1 配置环境变量
|
||||
|
||||
复制 `.env.example` 为 `.env`,修改配置:
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
vim .env
|
||||
```
|
||||
|
||||
关键配置:
|
||||
|
||||
| 变量 | 说明 | 默认值 |
|
||||
|------|------|--------|
|
||||
| `API_BASE_URL` | SaaS API 基础 URL | `https://staging-api.xiaoxiajianji.com` |
|
||||
| `GPU_WORKER_TOKEN` | 长期 API Token(与服务端一致) | - |
|
||||
| `MUSE_TALK_URL` | 本地 MuseTalk 服务地址 | `http://127.0.0.1:7861` |
|
||||
| `POLL_INTERVAL` | 轮询间隔秒 | `5` |
|
||||
| `HEARTBEAT_INTERVAL` | 空闲心跳间隔秒 | `15` |
|
||||
| `REQUEST_TIMEOUT` | HTTP 请求超时秒 | `900` |
|
||||
| `TASK_MAX_RETRY` | 本地最大重试次数 | `1` |
|
||||
| `TASK_HEARTBEAT_INTERVAL` | 推理期间任务心跳间隔秒 | `30` |
|
||||
| `MIN_VIDEO_DURATION_SECONDS` | 最短输入视频时长秒 | `3` |
|
||||
|
||||
### 3.2 启动 Worker
|
||||
|
||||
```bash
|
||||
# 前台运行(调试用)
|
||||
python gpu_worker.py
|
||||
|
||||
# 后台运行(生产用 systemd)
|
||||
sudo systemctl start xiaoxia-gpu-worker
|
||||
sudo systemctl enable xiaoxia-gpu-worker
|
||||
```
|
||||
|
||||
### 3.3 验证启动日志
|
||||
|
||||
应看到:
|
||||
|
||||
```
|
||||
============================================================
|
||||
MuseTalk GPU Worker 启动
|
||||
worker_id = rtx2060-xxxx
|
||||
api_base = https://staging-api.xiaoxiajianji.com
|
||||
muse_talk = http://127.0.0.1:7861
|
||||
poll = 5.0s / heartbeat = 15.0s
|
||||
============================================================
|
||||
MuseTalk 健康检查通过: {...}
|
||||
注册/心跳成功
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 四、常见问题排查
|
||||
|
||||
| 现象 | 可能原因 / 排查 |
|
||||
|---|---|
|
||||
| 日志 401 `Invalid GPU worker token` | `.env` 的 `GPU_WORKER_TOKEN` 与服务端不一致 |
|
||||
| 日志 `MuseTalk 健康检查未通过` | 本地 MuseTalk 没启动,或端口不是 7861;`curl http://127.0.0.1:7861/health` 验证 |
|
||||
| 任务长时间不被拉取 | Worker 和服务端连不上;检查 API_BASE_URL 是否可达、Token 是否正确 |
|
||||
| 推理后上传 OSS 失败 | 本地出口网络被防火墙拦截 OSS 域名(oss-cn-hangzhou.aliyuncs.com) |
|
||||
| 服务端看到任务回退到 pending 重试 | 任务心跳真正超时(默认 900s):Worker 进程崩溃/断网,或推理彻底卡死;正常长推理期间心跳线程每 30s 续期,不会回退 |
|
||||
| 日志 `MuseTalk 推理超时或连接失败` | 视频太长或显存不足;可临时调大 REQUEST_TIMEOUT(服务端 GPU_TASK_TIMEOUT_SECONDS 需同步调大),或限制输入视频时长 |
|
||||
| 日志 `视频过短(x.xxs < 3s)` | 输入视频不足 3s,MuseTalk 对短视频会 division by zero,已在本地直接上报失败;可用 MIN_VIDEO_DURATION_SECONDS 调整阈值 |
|
||||
| MuseTalk 服务端 503 `GPU 正在处理其他任务` | 并发请求被锁拒绝,等当前推理完成即可 |
|
||||
| MuseTalk 服务端 504 `推理超时` | 推理超过 MUSE_INFERENCE_TIMEOUT,客户端会调 /cancel 终止服务端任务 |
|
||||
|
||||
---
|
||||
|
||||
## 五、安全注意事项
|
||||
|
||||
- `.env` 包含长期 Token,文件权限设为 600(`chmod 600 .env`)
|
||||
- Token 泄露要立即在服务端更换 `GPU_WORKER_TOKEN` 并重启 Worker
|
||||
- Worker 只需要出站访问 SaaS API 和 OSS,不需要开放任何入站端口
|
||||
- MuseTalk 服务端只监听本地 127.0.0.1(或 0.0.0.0 但通过防火墙限制),不暴露到公网
|
||||
- 临时文件自动清理(推理完成/失败后),无需手动维护
|
||||
|
||||
---
|
||||
|
||||
## 六、工程改进记录(musetalk_server.py)
|
||||
|
||||
相比原 `worker.py`,修复了以下 8 个 bug:
|
||||
|
||||
1. **Flask 单线程阻塞**:`app.run(threaded=True)`,推理时 `/health` 仍可响应
|
||||
2. **fps=0 除零崩溃**:`_get_video_fps()` 兜底 `MUSE_DEFAULT_FPS`
|
||||
3. **ffmpeg 不检查返回码**:`subprocess.run(check=True)` + 超时检查,失败立即报错
|
||||
4. **无并发锁**:`threading.Lock` 控制并发,第二请求立即 503
|
||||
5. **无推理超时**:线程 join timeout,超时返回 504 并调 `/cancel`
|
||||
6. **结果文件不清理**:推理完成/失败后自动删除临时目录
|
||||
7. **无人脸检测兜底**:MuseTalk 推理内部处理(TODO: 可在 `_run_inference` 前置检查)
|
||||
8. **上传无大小限制**:`_check_file_size()` 校验,超限返回 413
|
||||
|
||||
新增:
|
||||
- `/cancel` 端点:终止当前推理任务,清理临时文件
|
||||
- `/health` 端点:返回 GPU 显存信息和当前任务状态
|
||||
|
||||
2026-09-20 追加修复(音轨正确性,上线阻断级):
|
||||
|
||||
9. **音轨未替换(严重)**:旧最终封装让 ffmpeg 默认选流,结果保留了源视频自带音轨(与画面相关系数 0.9998,与 TTS 无关)。改为 `_mux_video_with_audio()` 统一封装,强制 `-map 0:v:0 -map 1:a:0`,画面取 MuseTalk 无声产物、音轨只取驱动音频
|
||||
10. **音视频时长不对齐**:TTS 长于原视频时 `-shortest` 会截短语音。改为探测双方时长,音频更长时 `-stream_loop -1` 循环画面 + `h264_nvenc` 硬件重编码(`MUSE_VIDEO_ENCODER=auto`,失败回退 libx264)+ `-t <音频时长>`;不循环时 `-c:v copy` 秒封装
|
||||
- 开关 `MUSE_ENABLE_VIDEO_LOOP=0` 可关闭循环;请求也支持 form 参数 `enable_video_loop` 单任务覆盖
|
||||
|
||||
2026-09-20 v2 架构重构(性能回归修复,上线阻断级):
|
||||
|
||||
11. **16 倍性能回归**:#9/#10 的实现虽然音轨正确,但在某些集成场景下(推理前 loop 视频再喂 MuseTalk)导致推理帧数 ×2.2 + 叠加 ffmpeg 软编码预处理,5s 视频 +11s 音频推理 >200s,nginx 60s 超时 504
|
||||
- **正确架构**:MuseTalk 原生支持长音频输入,内部自动循环视频帧。把【原视频】+【全量音频】直传 MuseTalk,输出时长=音频时长
|
||||
- **ffmpeg 后置快速封装**:`-c:v copy -c:a aac -shortest` 秒级完成,不重编码
|
||||
- **循环仅兜底**:仅当 MuseTalk 输出画面短于音频时(极端情况),才 `-stream_loop` + NVENC 兜底补齐
|
||||
- **业务侧异步化**:POST /lipsync/jobs 创建 GPU 任务后立即返回 `job.status="processing"`,Celery 异步等待结果回写。前端 GET /jobs/{id} 轮询。避免同步阻塞 HTTP 请求 >200s
|
||||
- **删除 `MUSE_ENABLE_VIDEO_LOOP`**:不再需要此开关
|
||||
|
||||
---
|
||||
|
||||
## 七、自动部署
|
||||
|
||||
从 2026-09-20 起,GPU 节点配置文件和脚本全部入库到 `deploy/gpu_worker/`,支持一键初始化新节点 + develop 分支 push 后 30 秒内自动拉取更新。
|
||||
|
||||
### 7.1 服务架构
|
||||
|
||||
每个 GPU 渲染节点运行三个 systemd 单元:
|
||||
|
||||
| 单元 | 类型 | 作用 |
|
||||
|---|---|---|
|
||||
| `musetalk-worker.service` | simple(常驻) | MuseTalk Flask 推理 API(监听 127.0.0.1:7861) |
|
||||
| `xiaoxia-gpu-worker.service` | simple(常驻) | 反向轮询 SaaS API 拉口型任务的 Worker 客户端 |
|
||||
| `gpu-poll.timer` + `gpu-poll.service` | timer(每 30s 触发 oneshot) | 轮询 Gitea `deploy/gpu_worker/` 最新 commit,有变更自动执行 update 脚本 |
|
||||
|
||||
脚本目录(节点本地):
|
||||
|
||||
| 路径 | 来源 | 作用 |
|
||||
|---|---|---|
|
||||
| `~/projects/update-gpu-worker.sh` | `scripts/update-gpu-worker.sh` | 备份 → 拉代码 → 重启两个服务 → 健康检查 → 失败回滚 |
|
||||
| `~/projects/gpu-webhook/poll_and_update.sh` | `scripts/poll_and_update.sh` | 轮询 Gitea API 比对 SHA,有新 commit 时触发 update |
|
||||
|
||||
### 7.2 新节点部署步骤
|
||||
|
||||
**前置准备**(手动,首次部署必做):
|
||||
|
||||
1. 安装 NVIDIA 驱动 + CUDA 11.8+,`nvidia-smi` 能看到 GPU
|
||||
2. 克隆 MuseTalk 代码到 `~/projects/MuseTalk/`,下载模型权重到 `~/projects/MuseTalk/models/musetalk/`(权重约几 GB,不适合自动下载)
|
||||
3. 创建 Python 虚拟环境 `~/projects/MuseTalk/venv/` 并安装 MuseTalk 依赖(PyTorch CUDA 版等)
|
||||
4. 创建 Worker 虚拟环境 `/opt/xiaoxia-gpu-worker/venv/` 并 `pip install -r requirements.txt`
|
||||
5. 准备 `.env` 文件(Worker 端):`/opt/xiaoxia-gpu-worker/.env`,填好 `API_BASE_URL`、`GPU_WORKER_TOKEN`、`MUSE_TALK_URL` 等(参考 `.env.example`)
|
||||
|
||||
> ⚠️ 模型权重和 Python 虚拟环境(含 CUDA 版 PyTorch)体积大、安装慢,首次部署必须手动准备;后续脚本只更新 `.py` 文件和配置,不碰权重和 venv。
|
||||
|
||||
**一键初始化**:
|
||||
|
||||
```bash
|
||||
# 从仓库拉取 setup 脚本并执行(在全新 GPU 机器上以 ying 用户执行)
|
||||
wget -q -O /tmp/setup-gpu-node.sh \
|
||||
"https://git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/raw/branch/develop/deploy/gpu_worker/scripts/setup-gpu-node.sh"
|
||||
bash /tmp/setup-gpu-node.sh
|
||||
```
|
||||
|
||||
脚本自动完成:
|
||||
|
||||
1. apt 安装系统依赖(python3、ffmpeg、wget、curl、git)
|
||||
2. 创建必要目录(`~/projects/MuseTalk`、`~/projects/gpu-webhook`、`/opt/xiaoxia-gpu-worker`)
|
||||
3. 从仓库拉取三个 systemd 单元文件 + update/poll 脚本到本地
|
||||
4. 安装 systemd 服务到 `/etc/systemd/system/`
|
||||
5. 配置 sudo 免密(仅允许 `ying` 用户免密 restart 两个服务、status、journalctl、cp、chmod、tee)
|
||||
6. 首次执行 update 脚本拉取最新 `musetalk_server.py` 和 `gpu_worker.py`
|
||||
7. `systemctl daemon-reload` + enable + start 三个单元
|
||||
|
||||
**初始化后检查**:
|
||||
|
||||
```bash
|
||||
sudo systemctl status musetalk-worker # 应 active (running)
|
||||
sudo systemctl status xiaoxia-gpu-worker # 应 active (running)
|
||||
sudo systemctl status gpu-poll.timer # 应 active (waiting)
|
||||
curl http://127.0.0.1:7861/health # 应返回 healthy + GPU 显存信息
|
||||
```
|
||||
|
||||
### 7.3 自动更新机制
|
||||
|
||||
push 到 `develop` 分支且修改了 `deploy/gpu_worker/` 下任何文件后:
|
||||
|
||||
1. `gpu-poll.timer` 每 30 秒触发 `gpu-poll.service`
|
||||
2. `poll_and_update.sh` 调用 Gitea API 取 `deploy/gpu_worker/` 路径最新 commit SHA
|
||||
3. 与本地 `~/projects/gpu-webhook/.last_commit` 比对,无变更直接退出
|
||||
4. 有变更:写入新 SHA → 执行 `update-gpu-worker.sh`
|
||||
5. `update-gpu-worker.sh` 执行流程:
|
||||
- 备份当前 `musetalk_server.py` / `gpu_worker.py`(带时间戳后缀)
|
||||
- wget 拉取最新 `musetalk_server.py`、`gpu_worker.py`
|
||||
- 比对 `requirements.txt`,有变化则 pip install
|
||||
- `sudo systemctl restart musetalk-worker`,等 5 秒
|
||||
- `sudo systemctl restart xiaoxia-gpu-worker`,等 8 秒
|
||||
- `curl http://127.0.0.1:7861/health` 健康检查
|
||||
- 健康 → 写日志退出 0
|
||||
- 不健康 → 回滚到最新备份 → 重启 → 退出 1(日志记录 rolled back)
|
||||
|
||||
端到端延迟:从 push 到节点拉到新代码并重启,约 30~60 秒。
|
||||
|
||||
### 7.4 手动更新命令
|
||||
|
||||
```bash
|
||||
# 立即手动触发一次更新(不依赖 timer)
|
||||
bash ~/projects/update-gpu-worker.sh
|
||||
|
||||
# 查看更新日志
|
||||
tail -f /tmp/gpu-worker-update.log
|
||||
|
||||
# 查看轮询日志
|
||||
tail -f /tmp/gpu-poll.log
|
||||
|
||||
# 查看服务运行日志
|
||||
journalctl -u musetalk-worker -f # MuseTalk 推理服务日志
|
||||
journalctl -u xiaoxia-gpu-worker -f # GPU Worker 客户端日志
|
||||
journalctl -u gpu-poll.service -f # 轮询/更新触发日志
|
||||
```
|
||||
|
||||
### 7.5 仓库文件清单(自动部署相关)
|
||||
|
||||
```
|
||||
deploy/gpu_worker/
|
||||
├── musetalk-worker.service # MuseTalk 推理 API 的 systemd 服务
|
||||
├── gpu-poll.service # 自动更新轮询 oneshot service
|
||||
├── gpu-poll.timer # 每 30 秒触发轮询的 timer
|
||||
├── xiaoxia-gpu-worker.service # GPU Worker 客户端 systemd 服务(已有)
|
||||
├── gpu_worker.py # GPU Worker 客户端脚本(已有,自动更新)
|
||||
├── musetalk_server.py # MuseTalk Flask 服务端(已有,自动更新)
|
||||
├── requirements.txt # Worker Python 依赖(已有)
|
||||
├── .env.example # Worker 环境变量模板(已有)
|
||||
├── README.md # 本文档
|
||||
└── scripts/
|
||||
├── update-gpu-worker.sh # 更新脚本:备份→拉取→重启→健康检查→回滚
|
||||
├── poll_and_update.sh # 轮询脚本:SHA 比对→触发更新
|
||||
└── setup-gpu-node.sh # 新节点一键初始化脚本
|
||||
```
|
||||
|
||||
### 7.6 注意事项
|
||||
|
||||
- **首次部署必须手动准备**:MuseTalk 代码仓库、模型权重(`models/musetalk/`,几 GB)、MuseTalk 的 Python 虚拟环境(`venv/`,含 CUDA 版 PyTorch)。这些体积大、安装耗时长,不在自动更新范围内。
|
||||
- **脚本路径写死**:当前脚本路径固定为 `/home/ying/projects/` 和 `/opt/xiaoxia-gpu-worker/`,用户名固定 `ying`。后续如有多节点/多用户需求再做参数化。
|
||||
- **sudo 免密范围最小化**:setup 脚本写入 `/etc/sudoers.d/ying-gpu-update`,仅放行 restart/status 两个 GPU 相关服务、daemon-reload、journalctl、cp、chmod、tee,不开放全量 root。
|
||||
- **回滚只回滚 .py 文件**:健康检查失败只回滚 `musetalk_server.py` 和 `gpu_worker.py`,不回滚 pip 依赖(requirements.txt 变化概率低,且 pip 操作本身可能失败)。如需完全回滚,手动 `pip install -r requirements.txt` 指定旧版本。
|
||||
- **poll 脚本容错**:Gitea API 请求失败直接跳过,不触发更新,不会因为网络抖动误重启服务。
|
||||
@@ -0,0 +1,9 @@
|
||||
[Unit]
|
||||
Description=GPU Worker Auto-Update Poller
|
||||
|
||||
[Service]
|
||||
Type=oneshot
|
||||
User=ying
|
||||
ExecStart=/bin/bash /home/ying/projects/gpu-webhook/poll_and_update.sh
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
@@ -0,0 +1,10 @@
|
||||
[Unit]
|
||||
Description=Poll Gitea for GPU worker updates every 30 seconds
|
||||
|
||||
[Timer]
|
||||
OnBootSec=30
|
||||
OnUnitActiveSec=30
|
||||
AccuracySec=5
|
||||
|
||||
[Install]
|
||||
WantedBy=timers.target
|
||||
@@ -0,0 +1,485 @@
|
||||
"""MuseTalk GPU Worker — 反向轮询模式.
|
||||
|
||||
部署在有 RTX2060 的本地电脑上(192.168.0.193),
|
||||
主动轮询 SaaS API 拉取口型任务、调用本地 MuseTalk 推理、上传结果回 SaaS。
|
||||
|
||||
环境变量:
|
||||
API_BASE_URL SaaS API 基础 URL(不含 /api/v1),如 https://staging-api.xiaoxiajianji.com
|
||||
GPU_WORKER_TOKEN 长期 API Token(服务端 GPU_WORKER_TOKEN 需一致)
|
||||
WORKER_ID 本机唯一 ID(默认 hostname+网卡MAC 后4位)
|
||||
MUSE_TALK_URL 本地 MuseTalk 地址,默认 http://127.0.0.1:7861
|
||||
POLL_INTERVAL 轮询间隔秒,默认 5
|
||||
HEARTBEAT_INTERVAL 空闲心跳间隔秒,默认 15
|
||||
REQUEST_TIMEOUT HTTP 请求超时秒(下载/推理/上传统一使用),默认 900
|
||||
需与服务端 GPU_TASK_TIMEOUT_SECONDS(默认 900)对齐
|
||||
TASK_MAX_RETRY 单任务本地最大重试次数(仅对瞬时错误重试),默认 1
|
||||
TASK_HEARTBEAT_INTERVAL 推理期间任务心跳间隔秒,默认 30
|
||||
MIN_VIDEO_DURATION_SECONDS 最短输入视频时长秒,小于则直接上报失败,默认 3
|
||||
|
||||
用法:
|
||||
python gpu_worker.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger("musetalk-worker")
|
||||
|
||||
# ── 配置 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _env(name: str, default: str = "") -> str:
|
||||
v = os.environ.get(name, default)
|
||||
return v.strip() if isinstance(v, str) else default
|
||||
|
||||
|
||||
class Config:
|
||||
api_base_url: str = _env("API_BASE_URL", "https://staging-api.xiaoxiajianji.com").rstrip("/")
|
||||
gpu_worker_token: str = _env("GPU_WORKER_TOKEN")
|
||||
muse_talk_url: str = _env("MUSE_TALK_URL", "http://127.0.0.1:7861").rstrip("/")
|
||||
poll_interval: float = float(_env("POLL_INTERVAL", "5"))
|
||||
heartbeat_interval: float = float(_env("HEARTBEAT_INTERVAL", "15"))
|
||||
# #1970:RTX2060 6G 处理 720p 长视频可能 >5min;与服务端
|
||||
# GPU_TASK_TIMEOUT_SECONDS 默认值对齐为 900,避免推理被本地/服务端先掐断。
|
||||
request_timeout: float = float(_env("REQUEST_TIMEOUT", "900"))
|
||||
# 本地只在网络/MuseTalk 瞬时错误时重试 1 次;服务端 MAX_ATTEMPTS=3
|
||||
# 负责跨 worker/真正超时后的重派发,总尝试次数不再相乘放大。
|
||||
task_max_retry: int = int(_env("TASK_MAX_RETRY", "1"))
|
||||
# 推理期间任务心跳间隔(独立线程 POST /gpu/register 带 task_id)
|
||||
task_heartbeat_interval: float = float(_env("TASK_HEARTBEAT_INTERVAL", "30"))
|
||||
# 输入视频最短时长(秒):过短(如 1s)MuseTalk 会 division by zero,
|
||||
# 本地前置拦截,直接上报 failed,不浪费 GPU 时间
|
||||
min_video_duration_seconds: float = float(_env("MIN_VIDEO_DURATION_SECONDS", "3"))
|
||||
worker_id: str = _env("WORKER_ID", "")
|
||||
|
||||
@classmethod
|
||||
def derived_worker_id(cls) -> str:
|
||||
if cls.worker_id:
|
||||
return cls.worker_id
|
||||
# hostname + MAC 后4位 → 稳定唯一 ID
|
||||
try:
|
||||
mac = uuid.getnode()
|
||||
mac_suffix = f"{mac:012x}"[-4:]
|
||||
except Exception:
|
||||
mac_suffix = "0000"
|
||||
host = platform.node() or socket.gethostname() or "rtx2060"
|
||||
return f"{host}-{mac_suffix}"
|
||||
|
||||
|
||||
# ── 辅助 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _api_headers() -> dict[str, str]:
|
||||
token = Config.gpu_worker_token
|
||||
if not token:
|
||||
logger.warning("GPU_WORKER_TOKEN 未配置,开发模式下会被服务端拒绝(生产环境必须配置)")
|
||||
return {"Authorization": f"Bearer {token}"} if token else {}
|
||||
|
||||
|
||||
def _check_musetalk_health() -> tuple[bool, dict]:
|
||||
"""检查本地 MuseTalk 健康状态,返回 (ok, info)."""
|
||||
try:
|
||||
r = requests.get(f"{Config.muse_talk_url}/health", timeout=5)
|
||||
if r.status_code == 200:
|
||||
try:
|
||||
return True, r.json()
|
||||
except Exception:
|
||||
return True, {}
|
||||
return False, {"status_code": r.status_code, "body": r.text[:200]}
|
||||
except Exception as exc:
|
||||
return False, {"error": str(exc)}
|
||||
|
||||
|
||||
def _register(task_id: Optional[str] = None) -> bool:
|
||||
"""向服务端注册 / 心跳,附带 GPU 信息。
|
||||
|
||||
推理期间的心跳线程传 task_id:服务端会同步刷新该 processing 任务的
|
||||
last_heartbeat_at,防止长推理被误判超时回收。
|
||||
"""
|
||||
ok, info = _check_musetalk_health()
|
||||
free_vram = int(info.get("free_vram_mb", 0) or 0) if isinstance(info, dict) else 0
|
||||
gpu_name = info.get("gpu_name", "") if isinstance(info, dict) else ""
|
||||
if not gpu_name:
|
||||
# 尝试在 Windows 上读 nvidia-smi
|
||||
gpu_name = _probe_gpu_name()
|
||||
payload = {
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"hostname": platform.node(),
|
||||
"gpu_name": gpu_name,
|
||||
"free_vram_mb": free_vram,
|
||||
"capabilities": "musetalk",
|
||||
}
|
||||
if task_id:
|
||||
payload["task_id"] = task_id
|
||||
try:
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/register",
|
||||
json=payload,
|
||||
headers=_api_headers(),
|
||||
timeout=15,
|
||||
)
|
||||
if r.status_code == 200:
|
||||
return True
|
||||
logger.error("注册/心跳失败: HTTP %d body=%s", r.status_code, r.text[:300])
|
||||
return False
|
||||
except Exception as exc:
|
||||
logger.error("注册/心跳异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
def _probe_gpu_name() -> str:
|
||||
"""尽力探测 GPU 型号(不强制依赖 pynvml)."""
|
||||
try:
|
||||
import subprocess
|
||||
|
||||
out = subprocess.check_output(
|
||||
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=5,
|
||||
)
|
||||
return out.decode("utf-8", errors="ignore").strip().splitlines()[0].strip()
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
|
||||
def _poll_task() -> Optional[dict]:
|
||||
"""轮询拉取一条待处理任务;无任务返回 None."""
|
||||
try:
|
||||
r = requests.get(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/poll",
|
||||
params={"worker_id": Config.derived_worker_id()},
|
||||
headers=_api_headers(),
|
||||
timeout=30,
|
||||
)
|
||||
if r.status_code == 204:
|
||||
return None
|
||||
if r.status_code == 200:
|
||||
data = r.json()
|
||||
return data.get("task")
|
||||
logger.error("poll 返回 %d: %s", r.status_code, r.text[:300])
|
||||
return None
|
||||
except Exception as exc:
|
||||
logger.error("poll 异常: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def _download(url: str, path: Path) -> bool:
|
||||
"""下载文件到本地,支持预签名 URL."""
|
||||
try:
|
||||
with requests.get(url, stream=True, timeout=Config.request_timeout) as r:
|
||||
if r.status_code >= 400:
|
||||
logger.error("下载失败 HTTP %d: %s", r.status_code, url[:120])
|
||||
return False
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "wb") as f:
|
||||
for chunk in r.iter_content(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
return path.stat().st_size > 0
|
||||
except Exception as exc:
|
||||
logger.error("下载异常 %s: %s", url[:120], exc)
|
||||
return False
|
||||
|
||||
|
||||
def _call_musetalk(video_path: Path, audio_path: Path, out_path: Path) -> tuple[bool, float, str, bool]:
|
||||
"""调用本地 MuseTalk /inference.
|
||||
|
||||
返回 (success, duration_seconds, error_msg, retryable)。
|
||||
duration 用 ffprobe 读结果视频,失败填 0。
|
||||
retryable 仅对瞬时错误(连接失败/超时/5xx)为 True;HTTP 4xx、结果过小
|
||||
等确定性失败不重试,直接上报服务端(服务端 MAX_ATTEMPTS 再决定是否重派发)。
|
||||
"""
|
||||
try:
|
||||
with open(video_path, "rb") as vf, open(audio_path, "rb") as af:
|
||||
files = {
|
||||
"video": (video_path.name, vf, "video/mp4"),
|
||||
"audio": (audio_path.name, af, "application/octet-stream"),
|
||||
}
|
||||
r = requests.post(
|
||||
f"{Config.muse_talk_url}/inference",
|
||||
files=files,
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
retryable = r.status_code >= 500
|
||||
return False, 0.0, f"MuseTalk HTTP {r.status_code}: {r.text[:500]}", retryable
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
out_path.write_bytes(r.content)
|
||||
if out_path.stat().st_size < 1024:
|
||||
# 确定性失败(推理产物异常),本地重试大概率还是坏的,不重试
|
||||
return False, 0.0, f"MuseTalk 返回结果过小 ({out_path.stat().st_size} bytes)", False
|
||||
duration = _probe_duration(out_path)
|
||||
return True, duration, "", False
|
||||
except (requests.exceptions.Timeout, requests.exceptions.ConnectionError):
|
||||
# 瞬时网络/超时错误,允许本地重试 1 次;同时调 /cancel 让服务端终止僵尸推理
|
||||
_cancel_musetalk()
|
||||
return False, 0.0, f"MuseTalk 推理超时或连接失败(>{Config.request_timeout}s)", True
|
||||
except Exception as exc:
|
||||
return False, 0.0, f"MuseTalk 调用异常: {exc}", False
|
||||
|
||||
|
||||
def _cancel_musetalk() -> None:
|
||||
"""调 MuseTalk /cancel 端点终止服务端僵尸推理进程,避免超时后任务还在跑占显存."""
|
||||
try:
|
||||
r = requests.post(f"{Config.muse_talk_url}/cancel", timeout=10)
|
||||
if r.status_code == 200:
|
||||
logger.info("已调 MuseTalk /cancel,服务端终止推理")
|
||||
else:
|
||||
logger.warning("MuseTalk /cancel 返回 %d: %s", r.status_code, r.text[:200])
|
||||
except Exception as exc:
|
||||
# /cancel 失败不应影响主流程上报
|
||||
logger.warning("调 MuseTalk /cancel 异常(忽略): %s", exc)
|
||||
|
||||
|
||||
def _probe_duration(path: Path) -> float:
|
||||
"""用 ffprobe 读视频时长(若系统装了 ffmpeg);否则返回 0."""
|
||||
try:
|
||||
import subprocess
|
||||
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
)
|
||||
return float(out.decode().strip() or 0)
|
||||
except Exception:
|
||||
return 0.0
|
||||
|
||||
|
||||
def _upload_result(upload_url: str, file_path: Path) -> bool:
|
||||
"""PUT 上传结果视频到预签名 URL."""
|
||||
try:
|
||||
with open(file_path, "rb") as f:
|
||||
r = requests.put(
|
||||
upload_url,
|
||||
data=f,
|
||||
headers={"Content-Type": "video/mp4"},
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code >= 400:
|
||||
logger.error("上传结果失败 HTTP %d: %s", r.status_code, r.text[:500])
|
||||
return False
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.error("上传结果异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
def _report_result(task_id: str, success: bool, duration: float = 0.0, error_msg: str = "") -> bool:
|
||||
"""通知服务端结果。失败时也尝试上报错误(不含视频文件)."""
|
||||
try:
|
||||
data = {
|
||||
"task_id": task_id,
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"success": "true" if success else "false",
|
||||
"duration_seconds": str(duration),
|
||||
"error_msg": error_msg,
|
||||
}
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/result",
|
||||
data=data,
|
||||
headers=_api_headers(),
|
||||
timeout=30,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
logger.error("上报结果失败 HTTP %d: %s", r.status_code, r.text[:300])
|
||||
return False
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.error("上报结果异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
class TaskHeartbeat(threading.Thread):
|
||||
"""推理期间的任务心跳线程。
|
||||
|
||||
主循环的空闲心跳在 ``_handle_task`` 同步阻塞(下载/推理/上传最长 900s)
|
||||
期间无法发送,服务端会因任务 last_heartbeat_at 停滞而误判超时回退 pending。
|
||||
本线程每 task_heartbeat_interval 秒(默认 30s)POST /gpu/register 并
|
||||
携带当前 task_id,让服务端持续续期任务心跳;任务处理结束 stop()。
|
||||
"""
|
||||
|
||||
def __init__(self, task_id: str, interval: float):
|
||||
super().__init__(daemon=True, name=f"hb-{task_id[:8]}")
|
||||
self.task_id = task_id
|
||||
self.interval = max(5.0, interval)
|
||||
self._stop_event = threading.Event()
|
||||
|
||||
def run(self) -> None:
|
||||
# 先立即发一次,再按间隔循环(首次心跳失败不影响主流程)
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
if _register(self.task_id):
|
||||
logger.debug("任务 %s 心跳已发送", self.task_id)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("任务 %s 心跳异常(忽略): %s", self.task_id, exc)
|
||||
self._stop_event.wait(self.interval)
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop_event.set()
|
||||
|
||||
|
||||
def _handle_task(task: dict) -> None:
|
||||
"""处理一条任务(整个串行流程:下载→时长校验→推理→上传→上报)。"""
|
||||
task_id = task["task_id"]
|
||||
logger.info("开始处理任务 %s", task_id)
|
||||
# 领取任务后立即启动任务级心跳线程,覆盖下载/推理/上报全过程
|
||||
hb = TaskHeartbeat(task_id, Config.task_heartbeat_interval)
|
||||
hb.start()
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="musetalk_") as tmpdir:
|
||||
tmp = Path(tmpdir)
|
||||
video_path = tmp / "input.mp4"
|
||||
audio_path = tmp / "input_audio.bin"
|
||||
out_path = tmp / "output.mp4"
|
||||
|
||||
# 1. 下载
|
||||
if not _download(task["video_url"], video_path):
|
||||
_report_result(task_id, False, 0.0, "下载人物视频失败")
|
||||
return
|
||||
if not _download(task["audio_url"], audio_path):
|
||||
_report_result(task_id, False, 0.0, "下载驱动音频失败")
|
||||
return
|
||||
|
||||
# 2. 输入时长前置校验:短视频 MuseTalk 会 division by zero,
|
||||
# 直接上报 failed,不浪费 GPU 时间。ffprobe 不可用/读失败(0.0)
|
||||
# 时不拦截,交给 MuseTalk 处理,避免误杀。
|
||||
video_duration = _probe_duration(video_path)
|
||||
if video_duration and video_duration < Config.min_video_duration_seconds:
|
||||
msg = (
|
||||
f"视频过短({video_duration:.2f}s < {Config.min_video_duration_seconds:.0f}s),"
|
||||
"MuseTalk 无法处理"
|
||||
)
|
||||
logger.error("任务 %s %s", task_id, msg)
|
||||
_report_result(task_id, False, 0.0, msg)
|
||||
return
|
||||
|
||||
# 3. 推理(本地仅对瞬时错误重试)
|
||||
success = False
|
||||
duration = 0.0
|
||||
err = ""
|
||||
retryable = False
|
||||
for attempt in range(Config.task_max_retry + 1):
|
||||
if attempt > 0:
|
||||
logger.info("任务 %s 第 %d 次重试(瞬时错误)...", task_id, attempt + 1)
|
||||
time.sleep(2)
|
||||
success, duration, err, retryable = _call_musetalk(video_path, audio_path, out_path)
|
||||
if success or not retryable:
|
||||
break
|
||||
if not success:
|
||||
logger.error("任务 %s 推理失败: %s", task_id, err)
|
||||
_report_result(task_id, False, 0.0, err)
|
||||
return
|
||||
|
||||
# 4. 上报结果(multipart 同时上传文件 → API 代为 PUT 到 OSS,逻辑最稳)
|
||||
_report_success_with_file(task_id, duration, out_path)
|
||||
finally:
|
||||
hb.stop()
|
||||
|
||||
|
||||
def _report_success_with_file(task_id: str, duration: float, file_path: Path) -> None:
|
||||
"""上报成功并 multipart 附带结果视频."""
|
||||
try:
|
||||
data = {
|
||||
"task_id": task_id,
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"success": "true",
|
||||
"duration_seconds": str(duration),
|
||||
"error_msg": "",
|
||||
}
|
||||
with open(file_path, "rb") as f:
|
||||
files = {"result": (f"{task_id}.mp4", f, "video/mp4")}
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/result",
|
||||
data=data,
|
||||
files=files,
|
||||
headers=_api_headers(),
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
logger.error("上报成功结果失败 HTTP %d: %s", r.status_code, r.text[:300])
|
||||
return
|
||||
logger.info("任务 %s 完成,duration=%.1fs", task_id, duration)
|
||||
except Exception as exc:
|
||||
logger.error("上报成功结果异常: %s", exc)
|
||||
|
||||
|
||||
# ── 主循环 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def main() -> int:
|
||||
logger.info("=" * 60)
|
||||
logger.info("MuseTalk GPU Worker 启动")
|
||||
logger.info(" worker_id = %s", Config.derived_worker_id())
|
||||
logger.info(" api_base = %s", Config.api_base_url)
|
||||
logger.info(" muse_talk = %s", Config.muse_talk_url)
|
||||
logger.info(" poll = %.1fs / heartbeat = %.1fs", Config.poll_interval, Config.heartbeat_interval)
|
||||
logger.info("=" * 60)
|
||||
|
||||
if not Config.gpu_worker_token:
|
||||
logger.warning("GPU_WORKER_TOKEN 未配置(开发模式),生产环境必须设置")
|
||||
|
||||
# 先检查一次 MuseTalk
|
||||
ok, info = _check_musetalk_health()
|
||||
if ok:
|
||||
logger.info("MuseTalk 健康检查通过: %s", info)
|
||||
else:
|
||||
logger.warning("MuseTalk 健康检查未通过: %s(继续运行,等待服务可用)", info)
|
||||
|
||||
# 启动时立即注册
|
||||
_register()
|
||||
last_heartbeat = time.time()
|
||||
|
||||
while True:
|
||||
try:
|
||||
# 心跳
|
||||
now = time.time()
|
||||
if now - last_heartbeat >= Config.heartbeat_interval:
|
||||
if _register():
|
||||
last_heartbeat = now
|
||||
|
||||
# 轮询任务
|
||||
task = _poll_task()
|
||||
if task is not None:
|
||||
_handle_task(task)
|
||||
# 处理完立即再 poll(不 sleep),尽可能拉满 GPU
|
||||
continue
|
||||
|
||||
time.sleep(Config.poll_interval)
|
||||
except KeyboardInterrupt:
|
||||
logger.info("收到中断信号,退出")
|
||||
return 0
|
||||
except Exception as exc:
|
||||
logger.exception("主循环异常: %s", exc)
|
||||
time.sleep(Config.poll_interval)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,19 @@
|
||||
[Unit]
|
||||
Description=MuseTalk Inference API Server
|
||||
After=network.target nvidia-persistenced.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=ying
|
||||
WorkingDirectory=/home/ying/projects/MuseTalk
|
||||
Environment=PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128
|
||||
Environment=PATH=/home/ying/projects/MuseTalk/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin
|
||||
ExecStart=/home/ying/projects/MuseTalk/venv/bin/python /home/ying/projects/MuseTalk/musetalk_server.py
|
||||
Restart=always
|
||||
RestartSec=10
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
SyslogIdentifier=musetalk-server
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -0,0 +1,647 @@
|
||||
"""MuseTalk Flask HTTP 服务 — 反向轮询架构的服务端部分.
|
||||
|
||||
部署在 RTX2060 本地,接收 gpu_worker.py 的推理请求,调用 MuseTalk 生成口型同步视频。
|
||||
本文件修复了原 worker.py 的 8 个工程 bug,并新增 /cancel 端点。
|
||||
|
||||
#1978 性能修复(v2 架构):
|
||||
MuseTalk 原生支持长音频输入(内部循环视频帧),不需要我们先 loop 视频。
|
||||
正确流程:原视频 + 全量音频 → MuseTalk 推理 → 输出时长=音频时长的无声画面
|
||||
→ ffmpeg 快速 -c:v copy 替换音轨。推理时间不变(~14s),后处理几秒。
|
||||
禁止在推理前用 ffmpeg 循环视频(会导致 MuseTalk 处理 2x+ 帧数,慢 16 倍)。
|
||||
|
||||
环境变量:
|
||||
MUSE_PORT 监听端口,默认 7861
|
||||
MUSE_MAX_CONCURRENT 最大并发推理数,默认 1(GPU 一次只能处理一个)
|
||||
MUSE_INFERENCE_TIMEOUT 推理超时秒数,默认 600
|
||||
MUSE_VIDEO_MAX_MB 视频上传大小限制 MB,默认 100
|
||||
MUSE_AUDIO_MAX_MB 音频上传大小限制 MB,默认 20
|
||||
MUSE_DEFAULT_FPS 视频 fps 兜底值,默认 25.0
|
||||
MUSE_TEMP_DIR 临时文件目录,默认 /tmp/musetalk_$$
|
||||
MUSE_VIDEO_ENCODER 循环视频时的编码器(仅兜底):auto(默认)/h264_nvenc/libx264
|
||||
|
||||
接口:
|
||||
GET /health 健康检查 + GPU 显存信息
|
||||
POST /inference 推理请求(multipart: video + audio)
|
||||
POST /cancel 终止当前推理任务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import atexit
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from flask import Flask, jsonify, request, send_file
|
||||
|
||||
# ── 日志 ──────────────────────────────────────────────────────────────
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger("musetalk-server")
|
||||
|
||||
# ── 配置 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _env(name: str, default: str = "") -> str:
|
||||
v = os.environ.get(name, default)
|
||||
return v.strip() if isinstance(v, str) else default
|
||||
|
||||
|
||||
class Config:
|
||||
port: int = int(_env("MUSE_PORT", "7861"))
|
||||
max_concurrent: int = int(_env("MUSE_MAX_CONCURRENT", "1"))
|
||||
inference_timeout: float = float(_env("MUSE_INFERENCE_TIMEOUT", "600"))
|
||||
video_max_mb: int = int(_env("MUSE_VIDEO_MAX_MB", "100"))
|
||||
audio_max_mb: int = int(_env("MUSE_AUDIO_MAX_MB", "20"))
|
||||
default_fps: float = float(_env("MUSE_DEFAULT_FPS", "25.0"))
|
||||
temp_dir: str = _env("MUSE_TEMP_DIR", f"/tmp/musetalk_{os.getpid()}")
|
||||
# 循环视频时的编码器(仅当 MuseTalk 输出画面短于音频时的兜底)
|
||||
video_encoder: str = _env("MUSE_VIDEO_ENCODER", "auto") or "auto"
|
||||
# 判定音视频时长差异的容差(秒)
|
||||
duration_epsilon: float = 0.25
|
||||
|
||||
|
||||
# ── 全局状态 ──────────────────────────────────────────────────────────
|
||||
inference_lock = threading.Lock()
|
||||
current_task: dict = {"task_id": None, "process": None, "start_time": 0.0}
|
||||
shutdown_event = threading.Event()
|
||||
|
||||
# ── Flask App ─────────────────────────────────────────────────────────
|
||||
app = Flask(__name__)
|
||||
|
||||
|
||||
def _cleanup_temp_dir():
|
||||
"""退出时清理临时目录."""
|
||||
if os.path.exists(Config.temp_dir):
|
||||
try:
|
||||
shutil.rmtree(Config.temp_dir)
|
||||
logger.info("已清理临时目录: %s", Config.temp_dir)
|
||||
except Exception as exc:
|
||||
logger.warning("清理临时目录失败: %s", exc)
|
||||
|
||||
|
||||
atexit.register(_cleanup_temp_dir)
|
||||
|
||||
|
||||
def _signal_handler(signum, frame):
|
||||
"""优雅退出."""
|
||||
logger.info("收到信号 %s,准备退出...", signum)
|
||||
shutdown_event.set()
|
||||
if current_task["process"]:
|
||||
logger.info("终止正在进行的推理进程...")
|
||||
try:
|
||||
current_task["process"].terminate()
|
||||
current_task["process"].wait(timeout=5)
|
||||
except Exception:
|
||||
pass
|
||||
_cleanup_temp_dir()
|
||||
exit(0)
|
||||
|
||||
|
||||
signal.signal(signal.SIGTERM, _signal_handler)
|
||||
signal.signal(signal.SIGINT, _signal_handler)
|
||||
|
||||
|
||||
# ── 工具函数 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_gpu_info() -> dict:
|
||||
"""获取 GPU 显存信息(通过 nvidia-smi)."""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"nvidia-smi",
|
||||
"--query-gpu=name,memory.total,memory.used,memory.free",
|
||||
"--format=csv,noheader,nounits",
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=5,
|
||||
)
|
||||
parts = out.decode().strip().split(",")
|
||||
if len(parts) >= 4:
|
||||
return {
|
||||
"gpu_name": parts[0].strip(),
|
||||
"memory_total_mb": int(parts[1].strip()),
|
||||
"memory_used_mb": int(parts[2].strip()),
|
||||
"memory_free_mb": int(parts[3].strip()),
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.warning("nvidia-smi 失败: %s", exc)
|
||||
return {"gpu_name": "unknown", "memory_total_mb": 0, "memory_used_mb": 0, "memory_free_mb": 0}
|
||||
|
||||
|
||||
def _get_video_fps(video_path: Path) -> float:
|
||||
"""用 ffprobe 读视频帧率,失败或为 0 时返回 default_fps."""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=r_frame_rate",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(video_path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
)
|
||||
fps_str = out.decode().strip()
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
fps = float(num) / float(den) if float(den) != 0 else 0.0
|
||||
else:
|
||||
fps = float(fps_str) if fps_str else 0.0
|
||||
return fps if fps > 0 else Config.default_fps
|
||||
except Exception as exc:
|
||||
logger.warning("ffprobe 读 fps 失败: %s,使用默认 %.1f", exc, Config.default_fps)
|
||||
return Config.default_fps
|
||||
|
||||
|
||||
def _get_media_duration(path: Path) -> float:
|
||||
"""用 ffprobe 读媒体时长(秒),失败返回 0.0."""
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
)
|
||||
duration = float(out.decode().strip())
|
||||
return duration if duration > 0 else 0.0
|
||||
except Exception as exc:
|
||||
logger.warning("ffprobe 读时长失败 %s: %s", path, exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _pick_video_encoder() -> str:
|
||||
"""选择视频编码器:配置指定则用指定值;auto 时探测 NVENC 是否可用,不可用回退 libx264."""
|
||||
configured = Config.video_encoder.strip()
|
||||
if configured in ("h264_nvenc", "libx264"):
|
||||
return configured
|
||||
# auto:探测本机 ffmpeg 是否编译了 h264_nvenc
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["ffmpeg", "-hide_banner", "-encoders"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
check=False,
|
||||
)
|
||||
if b"h264_nvenc" in result.stdout:
|
||||
return "h264_nvenc"
|
||||
except Exception as exc:
|
||||
logger.warning("探测 ffmpeg 编码器失败,回退 libx264: %s", exc)
|
||||
return "libx264"
|
||||
|
||||
|
||||
def _mux_video_with_audio(
|
||||
video_path: Path,
|
||||
audio_path: Path,
|
||||
output_path: Path,
|
||||
timeout: float = 300,
|
||||
) -> None:
|
||||
"""把无声画面视频与驱动音频封装为最终结果.
|
||||
|
||||
#1978 v2 架构:MuseTalk 已处理全量音频,输出视频时长=音频时长。
|
||||
此处仅做快速封装:-map 0:v:0 -map 1:a:0 强制取画面+驱动音频,
|
||||
-c:v copy 无损秒级封装(不重编码),-shortest 以较短流为准。
|
||||
|
||||
仅当 MuseTalk 输出画面短于音频时(极端兜底),才启用 -stream_loop + NVENC
|
||||
循环视频到音频长度。正常情况下走 copy 快速路径。
|
||||
"""
|
||||
video_duration = _get_media_duration(video_path)
|
||||
audio_duration = _get_media_duration(audio_path)
|
||||
|
||||
# 判断是否需要兜底循环(正常情况下 MuseTalk 输出已 >= 音频时长)
|
||||
need_loop_fallback = bool(
|
||||
audio_duration > 0 and video_duration > 0 and video_duration < audio_duration - Config.duration_epsilon
|
||||
)
|
||||
|
||||
if need_loop_fallback:
|
||||
# 兜底:MuseTalk 输出画面不足,循环补齐
|
||||
encoder = _pick_video_encoder()
|
||||
preset = "p4" if encoder == "h264_nvenc" else "veryfast"
|
||||
logger.warning(
|
||||
"MuseTalk 输出(%.2fs)短于音频(%.2fs),兜底循环视频以 %s 重编码",
|
||||
video_duration,
|
||||
audio_duration,
|
||||
encoder,
|
||||
)
|
||||
|
||||
def build_cmd(enc: str, pre: str) -> list:
|
||||
return [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-i",
|
||||
str(audio_path),
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-c:v",
|
||||
enc,
|
||||
"-preset",
|
||||
pre,
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
"-t",
|
||||
f"{audio_duration:.3f}",
|
||||
str(output_path),
|
||||
]
|
||||
|
||||
try:
|
||||
_run_ffmpeg(build_cmd(encoder, preset), timeout=timeout)
|
||||
except RuntimeError:
|
||||
if encoder == "h264_nvenc":
|
||||
logger.warning("h264_nvenc 兜底失败,回退 libx264 重试")
|
||||
_run_ffmpeg(build_cmd("libx264", "veryfast"), timeout=timeout)
|
||||
else:
|
||||
raise
|
||||
else:
|
||||
# 正常快速路径:-c:v copy 无损封装,仅替换音轨为驱动音频
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-i",
|
||||
str(audio_path),
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-c:v",
|
||||
"copy",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
"-shortest",
|
||||
str(output_path),
|
||||
]
|
||||
_run_ffmpeg(cmd, timeout=timeout)
|
||||
|
||||
|
||||
def _check_file_size(file, max_mb: int, label: str) -> Optional[str]:
|
||||
"""检查文件大小,超限返回错误信息,否则返回 None."""
|
||||
file.seek(0, 2)
|
||||
size = file.tell()
|
||||
file.seek(0)
|
||||
max_bytes = max_mb * 1024 * 1024
|
||||
if size > max_bytes:
|
||||
return f"{label} 文件大小 {size / (1024*1024):.1f}MB 超过限制 {max_mb}MB"
|
||||
if size == 0:
|
||||
return f"{label} 文件为空"
|
||||
return None
|
||||
|
||||
|
||||
def _run_ffmpeg(cmd: list, timeout: float = 120) -> subprocess.CompletedProcess:
|
||||
"""运行 ffmpeg 命令,检查返回码和超时."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
timeout=timeout,
|
||||
check=True,
|
||||
)
|
||||
return result
|
||||
except subprocess.CalledProcessError as exc:
|
||||
stderr = exc.stderr.decode(errors="ignore") if exc.stderr else ""
|
||||
raise RuntimeError(f"ffmpeg 失败 (code={exc.returncode}): {stderr[:500]}") from exc
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
raise RuntimeError(f"ffmpeg 超时(>{timeout}s)") from exc
|
||||
|
||||
|
||||
def _run_inference(
|
||||
video_path: Path,
|
||||
audio_path: Path,
|
||||
output_path: Path,
|
||||
) -> None:
|
||||
"""执行 MuseTalk 推理(v2 架构:全量音频直传,不在推理前 loop 视频).
|
||||
|
||||
#1978 性能修复核心:
|
||||
MuseTalk 原生支持长音频输入,内部会自动循环视频帧。
|
||||
我们只需把【原视频】和【全量音频】传给 MuseTalk,
|
||||
输出视频时长 = 音频时长(MuseTalk 自行处理帧循环)。
|
||||
禁止在推理前用 ffmpeg 循环视频(会导致慢 16 倍)。
|
||||
|
||||
实际部署时替换为 MuseTalk 真实推理逻辑。
|
||||
此处为示例实现:提取帧 → 模拟 MuseTalk 产出音频时长的无声画面 → 快速封装。
|
||||
"""
|
||||
fps = _get_video_fps(video_path)
|
||||
audio_duration = _get_media_duration(audio_path)
|
||||
video_duration = _get_media_duration(video_path)
|
||||
logger.info(
|
||||
"推理开始: video=%.2fs, audio=%.2fs, fps=%.2f",
|
||||
video_duration,
|
||||
audio_duration,
|
||||
fps,
|
||||
)
|
||||
|
||||
frames_dir = video_path.parent / "frames"
|
||||
frames_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1. 从原视频提取帧(仅原视频长度,不循环)
|
||||
_run_ffmpeg(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-r",
|
||||
str(fps),
|
||||
str(frames_dir / "frame_%05d.png"),
|
||||
],
|
||||
timeout=120,
|
||||
)
|
||||
|
||||
frame_files = sorted(frames_dir.glob("*.png"))
|
||||
if not frame_files:
|
||||
raise RuntimeError("未从视频中提取到帧")
|
||||
|
||||
# 2. 模拟 MuseTalk 推理:输入原视频帧 + 全量音频,输出音频时长的无声画面。
|
||||
# TODO: 替换为 MuseTalk 真实推理逻辑。
|
||||
# MuseTalk 真实调用示例(伪代码):
|
||||
# from musetalk import MuseTalkModel
|
||||
# model = MuseTalkModel(...)
|
||||
# silent_video = model.infer(video_path=video_path, audio_path=audio_path)
|
||||
# # MuseTalk 内部会循环视频帧匹配音频长度,输出时长=音频时长
|
||||
logger.warning("使用示例推理逻辑,未实际调用 MuseTalk 模型")
|
||||
|
||||
# 示例:生成音频时长的无声画面(循环原视频帧到音频长度)
|
||||
# 真实部署时 silent_video_path 应替换为 MuseTalk 输出的无声视频路径
|
||||
silent_video_path = video_path.parent / "visual_silent.mp4"
|
||||
|
||||
if audio_duration > video_duration + Config.duration_epsilon:
|
||||
# 音频更长:循环视频帧到音频长度(仅用于示例,真实 MuseTalk 内部处理)
|
||||
encoder = _pick_video_encoder()
|
||||
preset = "p4" if encoder == "h264_nvenc" else "veryfast"
|
||||
logger.info(
|
||||
"示例:循环视频帧到音频长度 %.2fs(真实 MuseTalk 内部处理,无需此步骤)",
|
||||
audio_duration,
|
||||
)
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-an",
|
||||
"-c:v",
|
||||
encoder,
|
||||
"-preset",
|
||||
preset,
|
||||
"-t",
|
||||
f"{audio_duration:.3f}",
|
||||
str(silent_video_path),
|
||||
]
|
||||
try:
|
||||
_run_ffmpeg(cmd, timeout=300)
|
||||
except RuntimeError:
|
||||
if encoder == "h264_nvenc":
|
||||
cmd[cmd.index(encoder)] = "libx264"
|
||||
cmd[cmd.index(preset) + 1] = "veryfast"
|
||||
_run_ffmpeg(cmd, timeout=300)
|
||||
else:
|
||||
raise
|
||||
else:
|
||||
# 音频不长:直接生成无声视频(原视频长度)
|
||||
_run_ffmpeg(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-an",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"veryfast",
|
||||
str(silent_video_path),
|
||||
],
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
# 3. 快速封装:-map 取推理画面 + 驱动音频,-c:v copy 无损秒级封装
|
||||
# MuseTalk 输出已匹配音频长度,此处无需循环,仅替换音轨
|
||||
_mux_video_with_audio(silent_video_path, audio_path, output_path)
|
||||
|
||||
if not output_path.exists() or output_path.stat().st_size < 1024:
|
||||
raise RuntimeError("推理产物不存在或过小")
|
||||
|
||||
logger.info(
|
||||
"推理完成: output=%.2fs (audio=%.2fs)",
|
||||
_get_media_duration(output_path),
|
||||
audio_duration,
|
||||
)
|
||||
|
||||
|
||||
# ── 路由 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@app.route("/health", methods=["GET"])
|
||||
def health():
|
||||
"""健康检查 + GPU 显存信息."""
|
||||
gpu_info = _get_gpu_info()
|
||||
task_info = {
|
||||
"task_id": current_task["task_id"],
|
||||
"running": current_task["process"] is not None,
|
||||
"elapsed_seconds": time.time() - current_task["start_time"] if current_task["start_time"] else 0.0,
|
||||
}
|
||||
return jsonify(
|
||||
{
|
||||
"status": "healthy",
|
||||
"gpu": gpu_info,
|
||||
"current_task": task_info,
|
||||
"timestamp": time.time(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@app.route("/inference", methods=["POST"])
|
||||
def inference():
|
||||
"""推理请求:multipart form 包含 video 和 audio 文件.
|
||||
|
||||
#1978 v2:MuseTalk 直接处理全量音频,输出时长=音频时长,无需预处理循环。
|
||||
"""
|
||||
# 并发控制:检查锁
|
||||
if not inference_lock.acquire(blocking=False):
|
||||
return jsonify({"error": "GPU 正在处理其他任务,请稍后重试", "status": "busy"}), 503
|
||||
|
||||
task_id = None
|
||||
video_path = None
|
||||
audio_path = None
|
||||
output_path = None
|
||||
|
||||
try:
|
||||
# 解析参数
|
||||
if "video" not in request.files or "audio" not in request.files:
|
||||
return jsonify({"error": "缺少 video 或 audio 文件"}), 400
|
||||
|
||||
video_file = request.files["video"]
|
||||
audio_file = request.files["audio"]
|
||||
task_id = request.form.get("task_id", f"task_{int(time.time())}")
|
||||
|
||||
# 文件大小检查
|
||||
err = _check_file_size(video_file, Config.video_max_mb, "视频")
|
||||
if err:
|
||||
return jsonify({"error": err}), 413
|
||||
err = _check_file_size(audio_file, Config.audio_max_mb, "音频")
|
||||
if err:
|
||||
return jsonify({"error": err}), 413
|
||||
|
||||
# 保存到临时目录
|
||||
task_dir = Path(Config.temp_dir) / task_id
|
||||
task_dir.mkdir(parents=True, exist_ok=True)
|
||||
video_path = task_dir / "input.mp4"
|
||||
audio_path = task_dir / "input_audio.wav"
|
||||
output_path = task_dir / "output.mp4"
|
||||
|
||||
video_file.save(str(video_path))
|
||||
audio_file.save(str(audio_path))
|
||||
|
||||
logger.info("开始推理 task_id=%s, video=%s, audio=%s", task_id, video_path.name, audio_path.name)
|
||||
|
||||
# 更新当前任务信息
|
||||
current_task["task_id"] = task_id
|
||||
current_task["start_time"] = time.time()
|
||||
current_task["process"] = "inference_thread" # 标记为运行中
|
||||
|
||||
# 在线程中运行推理(支持超时)
|
||||
result_container = {"error": None}
|
||||
|
||||
def inference_thread():
|
||||
try:
|
||||
_run_inference(video_path, audio_path, output_path)
|
||||
except Exception as exc:
|
||||
result_container["error"] = str(exc)
|
||||
|
||||
thread = threading.Thread(target=inference_thread)
|
||||
thread.start()
|
||||
thread.join(timeout=Config.inference_timeout)
|
||||
|
||||
if thread.is_alive():
|
||||
# 超时,终止
|
||||
logger.error("推理超时 (>%ds),终止任务 %s", Config.inference_timeout, task_id)
|
||||
return jsonify({"error": f"推理超时(>{Config.inference_timeout}s)", "task_id": task_id}), 504
|
||||
|
||||
if result_container["error"]:
|
||||
logger.error("推理失败 task_id=%s: %s", task_id, result_container["error"])
|
||||
return jsonify({"error": result_container["error"], "task_id": task_id}), 500
|
||||
|
||||
# 返回结果文件
|
||||
logger.info("推理完成 task_id=%s, output=%s", task_id, output_path)
|
||||
return send_file(str(output_path), mimetype="video/mp4", as_attachment=True, download_name=f"{task_id}.mp4")
|
||||
|
||||
except Exception as exc:
|
||||
logger.exception("推理异常: %s", exc)
|
||||
return jsonify({"error": str(exc)}), 500
|
||||
|
||||
finally:
|
||||
# 释放锁,清理当前任务信息
|
||||
inference_lock.release()
|
||||
current_task["task_id"] = None
|
||||
current_task["process"] = None
|
||||
current_task["start_time"] = 0.0
|
||||
|
||||
# 清理临时文件
|
||||
if video_path and video_path.parent.exists():
|
||||
try:
|
||||
shutil.rmtree(video_path.parent)
|
||||
logger.info("已清理临时目录: %s", video_path.parent)
|
||||
except Exception as exc:
|
||||
logger.warning("清理临时目录失败: %s", exc)
|
||||
|
||||
|
||||
@app.route("/cancel", methods=["POST"])
|
||||
def cancel():
|
||||
"""终止当前正在进行的推理任务."""
|
||||
if current_task["task_id"] is None:
|
||||
return jsonify({"message": "当前无正在运行的任务"})
|
||||
|
||||
task_id = current_task["task_id"]
|
||||
logger.info("收到取消请求,终止任务 %s", task_id)
|
||||
|
||||
# 终止推理进程(如果是 subprocess)
|
||||
if current_task["process"] and current_task["process"] != "inference_thread":
|
||||
try:
|
||||
current_task["process"].terminate()
|
||||
current_task["process"].wait(timeout=5)
|
||||
logger.info("已终止推理进程")
|
||||
except Exception as exc:
|
||||
logger.warning("终止进程失败: %s", exc)
|
||||
|
||||
# 清理临时文件
|
||||
task_dir = Path(Config.temp_dir) / task_id
|
||||
if task_dir.exists():
|
||||
try:
|
||||
shutil.rmtree(task_dir)
|
||||
logger.info("已清理临时目录: %s", task_dir)
|
||||
except Exception as exc:
|
||||
logger.warning("清理临时目录失败: %s", exc)
|
||||
|
||||
# 重置当前任务
|
||||
current_task["task_id"] = None
|
||||
current_task["process"] = None
|
||||
current_task["start_time"] = 0.0
|
||||
|
||||
return jsonify({"message": f"已取消任务 {task_id}"})
|
||||
|
||||
|
||||
# ── 主入口 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def main():
|
||||
"""启动 Flask 服务."""
|
||||
# 创建临时目录
|
||||
Path(Config.temp_dir).mkdir(parents=True, exist_ok=True)
|
||||
logger.info("临时目录: %s", Config.temp_dir)
|
||||
|
||||
gpu_info = _get_gpu_info()
|
||||
logger.info(
|
||||
"GPU: %s (显存 %dMB / %dMB)",
|
||||
gpu_info["gpu_name"],
|
||||
gpu_info["memory_used_mb"],
|
||||
gpu_info["memory_total_mb"],
|
||||
)
|
||||
logger.info(
|
||||
"启动 MuseTalk Server: port=%d, timeout=%.0fs, max_concurrent=%d",
|
||||
Config.port,
|
||||
Config.inference_timeout,
|
||||
Config.max_concurrent,
|
||||
)
|
||||
|
||||
app.run(host="0.0.0.0", port=Config.port, threaded=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1 @@
|
||||
requests>=2.31.0
|
||||
@@ -0,0 +1,47 @@
|
||||
#!/bin/bash
|
||||
|
||||
REPO_API="https://git.xiaoxiajianji.com/api/v1/repos/xiaoxia/xiaoxia-saas/commits?sha=develop&path=deploy/gpu_worker&limit=1"
|
||||
STATE_FILE="/home/ying/projects/gpu-webhook/.last_commit"
|
||||
UPDATE_SCRIPT="/home/ying/projects/update-gpu-worker.sh"
|
||||
LOG_FILE="/tmp/gpu-poll.log"
|
||||
|
||||
log() {
|
||||
echo "[$(date +"%Y-%m-%d %H:%M:%S")] $*" >> "$LOG_FILE"
|
||||
}
|
||||
|
||||
LATEST_SHA=$(curl -sk --max-time 10 "$REPO_API" | python3 -c "
|
||||
import sys, json
|
||||
try:
|
||||
data = json.load(sys.stdin)
|
||||
if isinstance(data, list) and len(data) > 0:
|
||||
print(data[0].get('sha', ''))
|
||||
else:
|
||||
print('')
|
||||
except:
|
||||
print('')
|
||||
" 2>/dev/null)
|
||||
|
||||
if [ -z "$LATEST_SHA" ]; then
|
||||
log "get latest commit failed, skip"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
LAST_SHA=""
|
||||
if [ -f "$STATE_FILE" ]; then
|
||||
LAST_SHA=$(cat "$STATE_FILE")
|
||||
fi
|
||||
|
||||
if [ "$LATEST_SHA" = "$LAST_SHA" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
if [ -z "$LAST_SHA" ]; then
|
||||
echo "$LATEST_SHA" > "$STATE_FILE"
|
||||
log "first run, recording SHA: $LATEST_SHA"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
log "new commit detected: $LAST_SHA -> $LATEST_SHA, triggering update"
|
||||
echo "$LATEST_SHA" > "$STATE_FILE"
|
||||
bash "$UPDATE_SCRIPT" >> "$LOG_FILE" 2>&1
|
||||
log "update completed"
|
||||
@@ -0,0 +1,62 @@
|
||||
#!/bin/bash
|
||||
# GPU节点一键初始化脚本 - 在全新GPU机器上执行
|
||||
|
||||
set -e
|
||||
|
||||
echo "=== 1. 安装系统依赖 ==="
|
||||
sudo apt-get update -qq
|
||||
sudo apt-get install -y -qq python3 python3-pip python3-venv ffmpeg wget curl git
|
||||
|
||||
echo "=== 2. 创建目录 ==="
|
||||
mkdir -p ~/projects/MuseTalk ~/projects/gpu-webhook /opt/xiaoxia-gpu-worker
|
||||
|
||||
echo "=== 3. 安装nvidia-container-toolkit(如需要Docker)==="
|
||||
# 可选,当前不使用Docker,跳过
|
||||
# distribution=$(. /etc/os-release;echo $ID$VERSION_ID)
|
||||
# curl -s -L https://nvidia.github.io/nvidia-docker/gpgkey | sudo apt-key add -
|
||||
# curl -s -L https://nvidia.github.io/nvidia-docker/$distribution/nvidia-docker.list | sudo tee /etc/apt/sources.list.d/nvidia-docker.list
|
||||
# sudo apt-get update && sudo apt-get install -y nvidia-container-toolkit
|
||||
# sudo nvidia-ctk runtime configure --runtime=docker
|
||||
# sudo systemctl restart docker
|
||||
|
||||
echo "=== 4. 拉取服务配置和脚本 ==="
|
||||
REPO_URL="https://git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/raw/branch/develop/deploy/gpu_worker"
|
||||
wget -q -O /tmp/musetalk-worker.service "$REPO_URL/musetalk-worker.service"
|
||||
wget -q -O /tmp/gpu-poll.service "$REPO_URL/gpu-poll.service"
|
||||
wget -q -O /tmp/gpu-poll.timer "$REPO_URL/gpu-poll.timer"
|
||||
wget -q -O ~/projects/update-gpu-worker.sh "$REPO_URL/scripts/update-gpu-worker.sh"
|
||||
wget -q -O ~/projects/gpu-webhook/poll_and_update.sh "$REPO_URL/scripts/poll_and_update.sh"
|
||||
chmod +x ~/projects/update-gpu-worker.sh ~/projects/gpu-webhook/poll_and_update.sh
|
||||
|
||||
echo "=== 5. 安装systemd服务 ==="
|
||||
sudo cp /tmp/musetalk-worker.service /etc/systemd/system/
|
||||
sudo cp /tmp/gpu-poll.service /etc/systemd/system/
|
||||
sudo cp /tmp/gpu-poll.timer /etc/systemd/system/
|
||||
|
||||
echo "=== 6. 配置sudo免密 ==="
|
||||
sudo bash -c 'cat > /etc/sudoers.d/ying-gpu-update << EOF
|
||||
ying ALL=(ALL) NOPASSWD: /bin/systemctl restart musetalk-worker
|
||||
ying ALL=(ALL) NOPASSWD: /bin/systemctl restart xiaoxia-gpu-worker
|
||||
ying ALL=(ALL) NOPASSWD: /bin/systemctl status musetalk-worker
|
||||
ying ALL=(ALL) NOPASSWD: /bin/systemctl status xiaoxia-gpu-worker
|
||||
ying ALL=(ALL) NOPASSWD: /bin/systemctl daemon-reload
|
||||
ying ALL=(ALL) NOPASSWD: /usr/bin/journalctl
|
||||
ying ALL=(ALL) NOPASSWD: /bin/cp
|
||||
ying ALL=(ALL) NOPASSWD: /bin/chmod
|
||||
ying ALL=(ALL) NOPASSWD: /usr/bin/tee
|
||||
EOF'
|
||||
sudo chmod 440 /etc/sudoers.d/ying-gpu-update
|
||||
|
||||
echo "=== 7. 首次拉取代码并启动服务 ==="
|
||||
bash ~/projects/update-gpu-worker.sh
|
||||
sudo systemctl daemon-reload
|
||||
sudo systemctl enable musetalk-worker xiaoxia-gpu-worker gpu-poll.timer
|
||||
sudo systemctl start musetalk-worker xiaoxia-gpu-worker gpu-poll.timer
|
||||
|
||||
echo "=== 完成! ==="
|
||||
echo "检查服务状态:"
|
||||
echo " sudo systemctl status musetalk-worker"
|
||||
echo " sudo systemctl status xiaoxia-gpu-worker"
|
||||
echo " sudo systemctl status gpu-poll.timer"
|
||||
echo "健康检查:curl http://127.0.0.1:7861/health"
|
||||
echo "更新日志:tail -f /tmp/gpu-worker-update.log"
|
||||
@@ -0,0 +1,62 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
REPO_URL="https://git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/raw/branch/develop/deploy/gpu_worker"
|
||||
MUSE_DIR="/home/ying/projects/MuseTalk"
|
||||
WORKER_DIR="/opt/xiaoxia-gpu-worker"
|
||||
LOG_FILE="/tmp/gpu-worker-update.log"
|
||||
|
||||
log() {
|
||||
local NOW
|
||||
NOW=$(date +"%Y-%m-%d %H:%M:%S")
|
||||
echo "[$NOW] $*" | tee -a "$LOG_FILE"
|
||||
}
|
||||
|
||||
log "========== start update =========="
|
||||
|
||||
BAK_SUFFIX=$(date +"%Y%m%d%H%M%S")
|
||||
cp "$MUSE_DIR/musetalk_server.py" "$MUSE_DIR/musetalk_server.py.bak.$BAK_SUFFIX"
|
||||
cp "$WORKER_DIR/gpu_worker.py" "$WORKER_DIR/gpu_worker.py.bak.$BAK_SUFFIX"
|
||||
log "backup done ($BAK_SUFFIX)"
|
||||
|
||||
wget -q -O "$MUSE_DIR/musetalk_server.py" "$REPO_URL/musetalk_server.py"
|
||||
log "musetalk_server.py updated"
|
||||
|
||||
wget -q -O "$WORKER_DIR/gpu_worker.py" "$REPO_URL/gpu_worker.py"
|
||||
log "gpu_worker.py updated"
|
||||
|
||||
wget -q -O /tmp/gpu-requirements.txt "$REPO_URL/requirements.txt"
|
||||
if [ -f "$WORKER_DIR/requirements.txt" ] && ! diff -q "$WORKER_DIR/requirements.txt" /tmp/gpu-requirements.txt > /dev/null 2>&1; then
|
||||
log "requirements changed, updating..."
|
||||
cp /tmp/gpu-requirements.txt "$WORKER_DIR/requirements.txt"
|
||||
"$WORKER_DIR/venv/bin/pip" install -r "$WORKER_DIR/requirements.txt" -q
|
||||
log "pip install done"
|
||||
else
|
||||
log "requirements no change, skip pip"
|
||||
fi
|
||||
|
||||
sudo systemctl restart musetalk-worker
|
||||
log "musetalk restarted"
|
||||
sleep 5
|
||||
|
||||
sudo systemctl restart xiaoxia-gpu-worker
|
||||
log "gpu-worker restarted"
|
||||
sleep 8
|
||||
|
||||
HEALTH=$(curl -s http://127.0.0.1:7861/health 2>/dev/null)
|
||||
if echo "$HEALTH" | grep -q "healthy\|ok"; then
|
||||
log "health check OK"
|
||||
log "========== update done =========="
|
||||
exit 0
|
||||
else
|
||||
log "health check FAILED, rolling back..."
|
||||
LATEST_MUSE_BAK=$(ls -t "$MUSE_DIR/musetalk_server.py.bak."* 2>/dev/null | head -1)
|
||||
LATEST_WORKER_BAK=$(ls -t "$WORKER_DIR/gpu_worker.py.bak."* 2>/dev/null | head -1)
|
||||
[ -n "$LATEST_MUSE_BAK" ] && cp "$LATEST_MUSE_BAK" "$MUSE_DIR/musetalk_server.py"
|
||||
[ -n "$LATEST_WORKER_BAK" ] && cp "$LATEST_WORKER_BAK" "$WORKER_DIR/gpu_worker.py"
|
||||
sudo systemctl restart musetalk-worker
|
||||
sleep 5
|
||||
sudo systemctl restart xiaoxia-gpu-worker
|
||||
log "rolled back"
|
||||
exit 1
|
||||
fi
|
||||
@@ -0,0 +1,21 @@
|
||||
[Unit]
|
||||
Description=MuseTalk GPU Worker (xiaoxia-saas 反向轮询)
|
||||
After=network.target musetalk.service
|
||||
# 本地 MuseTalk 服务启动后再启动本 Worker;若 MuseTalk 没有 systemd 服务则删除 musetalk.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=%i
|
||||
WorkingDirectory=/opt/xiaoxia-gpu-worker
|
||||
# 读取环境变量(API 地址、Token、轮询间隔等)
|
||||
EnvironmentFile=/opt/xiaoxia-gpu-worker/.env
|
||||
ExecStart=/opt/xiaoxia-gpu-worker/venv/bin/python /opt/xiaoxia-gpu-worker/gpu_worker.py
|
||||
Restart=always
|
||||
RestartSec=10
|
||||
# 日志走 journal,用 journalctl -u xiaoxia-gpu-worker -f 查看
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
SyslogIdentifier=xiaoxia-gpu-worker
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -19,6 +19,16 @@ COPY alembic/ ./alembic/
|
||||
COPY scripts/ ./scripts/
|
||||
COPY packages/ ./packages/
|
||||
COPY apps/api/ ./apps/api/
|
||||
# 抖音 cookies 文件:镜像内 baked-in 兜底 + host 挂载可覆盖
|
||||
# - /app/configs/douyin_cookies_default.txt: 镜像构建时 COPY 的兜底 cookies(始终有效)
|
||||
# - /app/configs/douyin_cookies.txt: host volume 挂载点(部署脚本 scp 覆盖,过期需更新)
|
||||
RUN mkdir -p /app/configs
|
||||
COPY deploy/configs/douyin_cookies.txt /app/configs/douyin_cookies_default.txt
|
||||
# 初始 COPY 一份到挂载点,host 挂载为空文件时 Python 代码会自动 fallback 到 default
|
||||
COPY deploy/configs/douyin_cookies.txt /app/configs/douyin_cookies.txt
|
||||
|
||||
# 强制升级 yt-dlp 到最新(抖音反爬经常变更,旧版 cookies 支持失效;#1968/#1963)
|
||||
RUN pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com --upgrade "yt-dlp>=2026.8.19"
|
||||
|
||||
# 设置环境变量
|
||||
ENV PATH="/opt/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
|
||||
|
||||
@@ -67,9 +67,10 @@ services:
|
||||
ports:
|
||||
- "127.0.0.1:${API_PORT:-8000}:8000"
|
||||
|
||||
# 共享生成文件目录
|
||||
# 共享生成文件目录 + 抖音 cookies 等运行时配置
|
||||
volumes:
|
||||
- generated-files:/app/generated
|
||||
- ../../deploy/configs:/app/configs:ro
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user