Compare commits

..

1 Commits

Author SHA1 Message Date
xiaoxia 99c5777514 debug: AI数字人克隆音色试听加 Network/Console 调试日志
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 43s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m44s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m48s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 1m50s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m0s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m0s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 2m21s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 4m42s
CI/CD Pipeline / CI Gate (pull_request) Successful in 4s
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 3m27s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 25s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 53s
AI Code Review / AI Code Review (pull_request) Successful in 6m30s
CI/CD Pipeline / Canary Release to Production (pull_request) Failing after 170h44m19s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 170h44m23s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 170h44m23s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 170h44m20s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 170h48m57s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 170h48m57s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 170h49m2s
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Failing after 170h49m5s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Failing after 170h49m6s
CI/CD Pipeline / PR Build API Image (pull_request) Failing after 170h49m6s
CI/CD Pipeline / Integration Tests (pull_request) Failing after 170h49m7s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 170h49m7s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 170h49m7s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 170h49m7s
CI/CD Pipeline / Check push changed paths (pull_request) Failing after 170h49m9s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 171h19m8s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 171h23m43s
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Failing after 171h23m51s
CI/CD Pipeline / PR Build Worker Image (pull_request) Failing after 171h23m52s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 171h23m53s
- 请求前打印 voice_id/voice_name/voice_clone_profile_id
- 响应后打印 audio_url 前80字符 + duration
- catch 打印 HTTP status + response data + message
- 帮助定位 /tts/preview 400/422/500/502 具体原因
2026-09-09 12:34:40 +08:00
783 changed files with 33577 additions and 39487 deletions
+1 -1
View File
@@ -1,2 +1,2 @@
CI trigger file - safe to delete
retrigger at 2026-09-15 20:31:24 UTC
updated!
-34
View File
@@ -196,37 +196,3 @@ DOUBAO_MODEL=doubao-seed-1-6-250615
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
DOUBAO_TIMEOUT=30
DOUBAO_MAX_RETRIES=2
# ==================== 积分/会员系统 (#1895) ====================
# 积分扣点总开关:默认 false(对现有用户零影响)。
# P2 阶段各业务路由逐个接入 @points_gate 时,用
# `if settings.points_enabled: ...`
# 包裹扣点逻辑;所有路由接入完成并验证通过后再在 staging/prod 打开。
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 长期鉴权 TokenWorker 端 .env 的 GPU_WORKER_TOKEN 必须与此一致
# 留空时 development 环境允许匿名访问(仅本地调试),staging/production 必须配置
GPU_WORKER_TOKEN=
# 单任务超时(秒),processing 超过此时长无任务心跳才回退 pending 或标记 failed
# #1970RTX2060 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
File diff suppressed because it is too large Load Diff
-60
View File
@@ -1,60 +0,0 @@
name: "Debug: Web container v2 (mount conflict)"
on:
push:
branches: [debug/web-crash-v2]
workflow_dispatch:
jobs:
web-diag:
runs-on: runtime-builder
timeout-minutes: 10
steps:
- name: Setup SSH and diagnose
shell: bash
env:
STAGING_SSH_KEY: ${{ secrets.PREVIEW_SSH_KEY }}
run: |
set -x
which ssh || (apt-get update -qq && apt-get install -y -qq openssh-client)
mkdir -p ~/.ssh && chmod 700 ~/.ssh
printf "%s" "$STAGING_SSH_KEY" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
H=47.98.113.167; P=22222
ssh-keyscan -p $P -H $H >> ~/.ssh/known_hosts 2>/dev/null
ssh -p $P -i ~/.ssh/id_rsa -o StrictHostKeyChecking=no root@$H 'bash -s' <<'REMOTE'
set -x
echo "=== Current staging containers ==="
docker ps -a --filter name=xiaoxia-*-staging --format "table {{.Names}}\t{{.Status}}\t{{.Image}}"
echo ""
echo "=== Web container logs (current/current-rolledback) ==="
docker logs xiaoxia-web-staging 2>&1 | tail -40
echo ""
echo "=== Web inspect: env & mounts ==="
docker inspect xiaoxia-web-staging --format 'Entrypoint: {{.Config.Entrypoint}} Cmd: {{.Config.Cmd}}'
docker inspect xiaoxia-web-staging --format '{{range .Config.Env}}{{.}}{{"\n"}}{{end}}' | grep -E "APP_ENV|VERSION"
echo "Mounts:"
docker inspect xiaoxia-web-staging --format '{{range .Mounts}}{{.Type}} {{.Source}} -> {{.Destination}} (rw={{.RW}}){{"\n"}}{{end}}'
echo ""
echo "=== Reproduce: rm on read-only bind mount ==="
docker run --rm --name nginx-ro-test \
-v /var/lib/xiaoxia-saas-staging/nginx-staging.conf:/etc/nginx/conf.d/default.conf:ro \
git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/xiaoxia-saas-web:387514c \
sh -c '
set -x
echo "Before:"
ls -la /etc/nginx/conf.d/
echo "Try rm (as entrypoint does):"
rm -f /etc/nginx/conf.d/default.conf
echo "rm exitcode=$?"
echo "After rm:"
ls -la /etc/nginx/conf.d/
echo "Test ln:"
ln -s /etc/nginx/nginx-staging.conf /etc/nginx/conf.d/default.conf
echo "ln exitcode=$?"
ls -la /etc/nginx/conf.d/
echo "nginx -t:"
nginx -t 2>&1
' 2>&1
echo ""
echo "=== Also test with NEW fixed image (9c0d4b1 if present) ==="
docker images | grep xiaoxia-saas-web | head -5
REMOTE
-2
View File
@@ -494,5 +494,3 @@
- [Fixed] Bug 修复
- [Security] 安全相关更新
- [Performance] 性能优化
---
- 2026-09-16: fix extract-from-douyin 异常路径全部返回业务码(消除500) #1963
-1
View File
@@ -1 +0,0 @@
retrigger3
-1
View File
@@ -263,4 +263,3 @@ pytest --cov=packages --cov-report=html
---
**License**: MIT
<!-- CI trigger: 1788229339 -->
@@ -1,45 +0,0 @@
"""lipsync_jobs 增加 TTS 直生字段(voice_id/script_text/speed/emotion
Revision ID: 073_add_lipsync_tts_fields
Revises: 072_add_ai_avatar_render
Create Date: 2026-09-09
"""
import sqlalchemy as sa
from alembic import op
revision = "073_add_lipsync_tts_fields"
down_revision = "072_add_ai_avatar_render"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 对口型支持「传音色 + 文案直接生成」:后端内部先 TTS 合成音频再提交对口型
op.add_column(
"lipsync_jobs",
sa.Column("voice_id", sa.String(200), nullable=False, server_default=""),
)
op.add_column(
"lipsync_jobs",
sa.Column("script_text", sa.Text(), nullable=False, server_default=""),
)
op.add_column(
"lipsync_jobs",
sa.Column("speed", sa.Float(), nullable=False, server_default=sa.text("1.0")),
)
op.add_column(
"lipsync_jobs",
sa.Column("emotion", sa.String(20), nullable=False, server_default=""),
)
# audio_url 改为可空:直生模式下音频由后端 TTS 合成后回填
op.alter_column("lipsync_jobs", "audio_url", existing_type=sa.Text(), nullable=True)
def downgrade() -> None:
op.alter_column("lipsync_jobs", "audio_url", existing_type=sa.Text(), nullable=False)
op.drop_column("lipsync_jobs", "emotion")
op.drop_column("lipsync_jobs", "speed")
op.drop_column("lipsync_jobs", "script_text")
op.drop_column("lipsync_jobs", "voice_id")
@@ -1,36 +0,0 @@
"""ai_avatar_render_jobs.script_id 放宽为可空串(手动文案直生场景不关联文案库)
Revision ID: 074_render_script_id_optional
Revises: 073_add_lipsync_tts_fields
Create Date: 2026-09-09
"""
import sqlalchemy as sa
from alembic import op
revision = "074_render_script_id_optional"
down_revision = "073_add_lipsync_tts_fields"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 列保持 NOT NULL(空串占位),仅应用层允许不传;这里显式补 server_default 防止历史约束歧义
with op.batch_alter_table("ai_avatar_render_jobs") as batch:
batch.alter_column(
"script_id",
existing_type=sa.String(length=36),
nullable=False,
server_default="",
)
def downgrade() -> None:
with op.batch_alter_table("ai_avatar_render_jobs") as batch:
batch.alter_column(
"script_id",
existing_type=sa.String(length=36),
nullable=False,
server_default=None,
)
@@ -1,27 +0,0 @@
"""add sentence_timings to lipsync_jobs
Revision ID: 075_add_sentence_timings
Revises: 074_ai_avatar_render_script_id_optional
Create Date: 2026-09-12
"""
import sqlalchemy as sa
from alembic import op
revision = "075_add_sentence_timings"
down_revision = "074_render_script_id_optional"
branch_labels = None
depends_on = None
def upgrade() -> None:
with op.batch_alter_table("lipsync_jobs") as batch:
batch.add_column(
sa.Column("sentence_timings", sa.JSON(), nullable=True),
)
def downgrade() -> None:
with op.batch_alter_table("lipsync_jobs") as batch:
batch.drop_column("sentence_timings")
-133
View File
@@ -1,133 +0,0 @@
"""add membership & points system
Revision ID: 076_membership_points
Revises: 075_add_sentence_timings
Create Date: 2026-09-15
"""
import sqlalchemy as sa
from sqlalchemy import text
from alembic import op
revision = "076_membership_points"
down_revision = "075_add_sentence_timings"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 1. users 表新增字段
with op.batch_alter_table("users") as batch:
batch.add_column(
sa.Column("is_member", sa.Boolean(), nullable=False, server_default=sa.text("false")),
)
batch.add_column(
sa.Column("member_type", sa.String(20), nullable=True),
)
batch.add_column(
sa.Column("member_expires_at", sa.DateTime(), nullable=True),
)
batch.add_column(
sa.Column("points_balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
)
# 2. points_accounts 积分账户表
op.create_table(
"points_accounts",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, unique=True, index=True),
sa.Column("balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_earned", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_spent", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 3. points_transactions 积分流水表
op.create_table(
"points_transactions",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("account_id", sa.String(36), nullable=False, index=True),
sa.Column("type", sa.String(20), nullable=False, index=True),
sa.Column("source", sa.String(50), nullable=False, index=True),
sa.Column("amount", sa.Integer(), nullable=False),
sa.Column("balance_after", sa.Integer(), nullable=False),
sa.Column("description", sa.String(255), nullable=False, server_default=""),
sa.Column("ref_id", sa.String(100), nullable=False, server_default=""),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 4. points_orders 积分/会员订单表
op.create_table(
"points_orders",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("order_type", sa.String(20), nullable=False),
sa.Column("product_code", sa.String(50), nullable=False),
sa.Column("amount_cents", sa.Integer(), nullable=False),
sa.Column("original_amount_cents", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("discount", sa.Float(), nullable=False, server_default=sa.text("1.0")),
sa.Column("points_amount", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("payment_method", sa.String(50), nullable=True),
sa.Column("payment_id", sa.String(100), nullable=True),
sa.Column("paid_at", sa.DateTime(), nullable=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 5. daily_usage_records 每日使用记录表
op.create_table(
"daily_usage_records",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("usage_date", sa.DateTime(), nullable=False),
sa.Column("usage_type", sa.String(50), nullable=False, server_default="free_clip"),
sa.Column("count", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.UniqueConstraint(
"user_id",
"usage_date",
"usage_type",
name="uq_daily_usage_user_date_type",
),
)
def downgrade() -> None:
op.drop_table("daily_usage_records")
op.drop_table("points_orders")
op.drop_table("points_transactions")
op.drop_table("points_accounts")
with op.batch_alter_table("users") as batch:
batch.drop_column("points_balance")
batch.drop_column("member_expires_at")
batch.drop_column("member_type")
batch.drop_column("is_member")
-67
View File
@@ -1,67 +0,0 @@
"""#1894: merge title_libraries into scripts — add title_text/title_category/title_config
Revision ID: 077_merge_title_libs
Revises: 076_membership_points
Create Date: 2026-09-15
"""
import sqlalchemy as sa
from alembic import context, op
revision = "077_merge_title_libs"
down_revision = "076_membership_points"
branch_labels = None
depends_on = None
def upgrade() -> 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="{}"),
)
if context.get_context().dialect.name == "postgresql":
conn = op.get_bind()
result = conn.execute(sa.text("SELECT to_regclass('public.title_libraries')"))
if result.scalar() is not None:
conn.execute(sa.text("""
INSERT INTO scripts
(id, user_id, title, content, segments, tags,
title_text, title_category, title_config,
created_at, updated_at)
SELECT
gen_random_uuid()::TEXT,
tl.user_id,
COALESCE(tl.name, '迁移标题'),
COALESCE(tl.text, ''),
'[]'::JSONB,
COALESCE(tl.tags, '[]'::JSONB),
COALESCE(tl.text, ''),
COALESCE(tl.category, ''),
COALESCE(tl."metadata", '{}'::JSONB),
tl.created_at,
tl.updated_at
FROM title_libraries tl
WHERE tl.is_active = true
AND NOT EXISTS (
SELECT 1 FROM scripts s
WHERE s.user_id = tl.user_id
AND s.title_text = COALESCE(tl.text, '')
AND s.title_category = COALESCE(tl.category, '')
AND s.created_at = tl.created_at
)
"""))
def downgrade() -> None:
with op.batch_alter_table("scripts") as batch:
batch.drop_column("title_config")
batch.drop_column("title_category")
batch.drop_column("title_text")
@@ -1,33 +0,0 @@
"""#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="{}"))
-58
View File
@@ -1,58 +0,0 @@
"""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")
@@ -1,37 +0,0 @@
"""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")
@@ -1,58 +0,0 @@
"""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")
-26
View File
@@ -1,26 +0,0 @@
"""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")
View File
-28
View File
@@ -6,7 +6,6 @@ from app.api.routes.assets import router as assets_router
from app.api.routes.auth import router as auth_router
from app.api.routes.chunked_upload import router as chunked_upload_router
from app.api.routes.classification_jobs import router as classification_jobs_router
from app.api.routes.clips_standalone import router as clips_standalone_router
from app.api.routes.cover_templates import router as cover_templates_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.feature_flags import router as feature_flags_router
@@ -14,15 +13,12 @@ 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
from app.api.routes.lipsync import router as lipsync_router
from app.api.routes.points import points_router, usage_router
from app.api.routes.projects import router as projects_router
from app.api.routes.scripts import router as scripts_router
from app.api.routes.scripts_ai import router as scripts_ai_router
from app.api.routes.share import router as share_router
from app.api.routes.subscription import router as subscription_router
from app.api.routes.tags import router as tags_router
@@ -160,10 +156,6 @@ api_router.include_router(
prefix="/templates",
tags=["Template"],
)
api_router.include_router(
clips_standalone_router,
tags=["Clips"],
)
api_router.include_router(
templates_editor_router,
prefix="/templates/{template_id}/editor",
@@ -192,28 +184,8 @@ api_router.include_router(
prefix="/scripts",
tags=["ScriptLibrary"],
)
api_router.include_router(
scripts_ai_router,
prefix="/scripts",
tags=["ScriptLibrary AI"],
)
api_router.include_router(
ai_avatar_render_router,
prefix="/ai-avatar/render",
tags=["AI Avatar Render"],
)
api_router.include_router(
points_router,
prefix="/points",
tags=["Points"],
)
api_router.include_router(
usage_router,
prefix="/usage",
tags=["Usage"],
)
api_router.include_router(
gpu_lipsync_router,
prefix="/gpu",
tags=["GPU Worker"],
)
@@ -1,91 +0,0 @@
"""默认模板兜底共享逻辑(P0 #1922).
提供 get_or_create_default_template_id(db, user_id) 共享函数,
供 templates.py 列表查询、clips_standalone.py 独立端点、dependencies.py
resolve_draft_plan_id 三处复用,避免三处各写一套兜底逻辑产生分叉。
根因:PR#1918 清理模板管理 API 时误删了 GET /templates 自动创建默认模板
兜底,前端 PR#1913 去掉空 tid 拦截后首次进入生成页拼出
/templates//editor/clips/from-assets(双斜杠)→ FastAPI 404,阻断新用户首次
生成。
"""
from __future__ import annotations
import logging
from typing import Optional
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
def get_or_create_default_template_id(db: Session, user_id: str) -> Optional[str]:
"""获取或自动创建默认配音模板的 id。
判定逻辑(不做异常降级,只有确实创建失败时才回滚重查):
1. 查用户名下 is_active=True 且有 TemplateClipConfig 的模板 → 返回其 id
2. 无则调用 CreateTemplateUseCase 创建一条默认 voice_over 模板;
3. 创建异常时 rollback 再重查一次(防并发唯一键冲突),重查仍无返回 None。
"""
from packages.adapters.sqlalchemy_impl.models import (
TemplateClipConfigModel,
TemplateModel,
)
from packages.adapters.sqlalchemy_impl.template_repository import (
SQLAlchemyTemplateRepository,
)
from packages.application.template.commands import (
CreateTemplateCommand,
SegmentCommand,
)
from packages.application.template.use_cases import CreateTemplateUseCase
existing = (
db.query(TemplateModel)
.filter(TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True))
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
try:
repo = SQLAlchemyTemplateRepository(db)
cmd = CreateTemplateCommand(
user_id=user_id,
name="默认配音模板",
mode="voice_over",
category="default",
tags=[],
title_config={},
subtitle_config={},
bgm_config={},
estimated_duration=0.0,
segments=[SegmentCommand(segment_order=0, duration_min=1.0, duration_max=30.0)],
)
tpl = CreateTemplateUseCase(repo).execute(cmd)
db.commit()
logger.info("auto-created default voice_over template: id=%s user=%s", tpl.id, user_id)
return tpl.id
except Exception:
db.rollback()
# 重查:可能并发请求已建好
existing = (
db.query(TemplateModel)
.filter(TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True))
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
logger.exception("failed to auto-create default template user=%s", user_id)
return None
+4 -19
View File
@@ -1,6 +1,6 @@
"""路由层共享辅助函数 — 消除跨文件重复定义。"""
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Any
from fastapi import HTTPException, status
@@ -25,27 +25,12 @@ 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"
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
return getattr(user, "subscription_plan", "free") or "free"
def require_project_and_library(
@@ -153,4 +138,4 @@ def format_utc_datetime(dt: datetime | None) -> str:
return dt
if dt.tzinfo is None:
return dt.isoformat() + "Z"
return dt.astimezone(UTC).isoformat().replace("+00:00", "Z")
return dt.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
+6 -6
View File
@@ -5,7 +5,7 @@
from __future__ import annotations
from typing import Literal
from typing import List, Literal
from app.services.ai_service import TITLE_STYLES, generate_smart_titles, semantic_match_assets
from fastapi import APIRouter
@@ -31,7 +31,7 @@ class GenerateTitlesRequest(BaseModel):
class GenerateTitlesResponse(BaseModel):
"""智能标题生成响应."""
titles: list[str] = Field(..., description="生成的标题列表")
titles: List[str] = Field(..., description="生成的标题列表")
style: str = Field(..., description="实际使用的风格")
source: str = Field(..., description="来源:doubao 或 fallback")
description: str = Field(..., description="原始描述")
@@ -53,7 +53,7 @@ class AssetMatchItem(BaseModel):
id: str = Field(..., description="素材ID")
name: str = Field(default="", description="素材名称")
tags: list[str] = Field(default_factory=list, description="标签列表")
tags: List[str] = Field(default_factory=list, description="标签列表")
description: str = Field(default="", description="素材描述")
@@ -61,7 +61,7 @@ class SemanticMatchRequest(BaseModel):
"""语义匹配请求."""
description: str = Field(..., min_length=1, max_length=500, description="目标视频内容描述")
assets: list[AssetMatchItem] = Field(..., min_length=1, max_length=100, description="待匹配素材列表")
assets: List[AssetMatchItem] = Field(..., min_length=1, max_length=100, description="待匹配素材列表")
top_k: int = Field(default=0, ge=0, le=100, description="返回前K个,0返回全部")
@@ -75,7 +75,7 @@ class SemanticMatchResultItem(AssetMatchItem):
class SemanticMatchResponse(BaseModel):
"""语义匹配响应."""
matches: list[SemanticMatchResultItem] = Field(..., description="按匹配度降序排列的素材列表")
matches: List[SemanticMatchResultItem] = Field(..., description="按匹配度降序排列的素材列表")
source: str = Field(..., description="来源:doubao / fallback")
description: str = Field(..., description="原始描述")
total: int = Field(..., description="输入素材总数")
@@ -99,7 +99,7 @@ def generate_titles(request: GenerateTitlesRequest):
return GenerateTitlesResponse(**result)
@router.get("/titles/styles", response_model=list[TitleStyleInfo])
@router.get("/titles/styles", response_model=List[TitleStyleInfo])
def list_title_styles():
"""获取支持的标题风格列表."""
return [
+11 -164
View File
@@ -11,17 +11,13 @@
from __future__ import annotations
import logging
from datetime import UTC, datetime
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.ai_avatar_render import (
AiAvatarRenderJobResponse,
CreateAiAvatarRenderRequest,
FinalizeRenderResponse,
SmartCoverResponse,
)
from app.services.ai_avatar_cover_service import generate_smart_cover
from app.services.ai_avatar_render_service import (
AiAvatarRenderError,
AiAvatarRenderService,
@@ -29,8 +25,6 @@ from app.services.ai_avatar_render_service import (
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.middleware.points_gate import points_gate
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -44,12 +38,10 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
@points_gate("ai_digital_human", per_unit=15)
def create_render_job(
body: CreateAiAvatarRenderRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
db: Session = Depends(get_db_session),
):
"""提交 AI 数字人渲染任务.
@@ -57,7 +49,7 @@ def create_render_job(
"""
try:
job = svc.create_render_job(
user_id=current_user.user.id,
user_id=current_user.id,
lipsync_job_id=body.lipsync_job_id,
script_id=body.script_id,
b_roll_segments=[s.model_dump() for s in body.b_roll_segments],
@@ -82,16 +74,10 @@ def create_render_job(
from app.tasks.ai_avatar_render import execute_ai_avatar_render
execute_ai_avatar_render.delay(job.id)
except Exception as exc:
logger.exception("Celery 任务投递失败(创建): job_id=%s err=%s", job.id, exc)
job.status = "failed"
job.error_message = f"任务提交失败:{exc}"
job.updated_at = datetime.now(UTC)
svc.db.commit()
svc.db.refresh(job)
return AiAvatarRenderJobResponse.model_validate(job)
except Exception:
logger.warning("Celery 任务提交失败,渲染任务已创建但未触发执行: %s", job.id)
return AiAvatarRenderJobResponse.model_validate(job)
return job
# ── GET /jobs — 任务列表 ─────────────────────────────────────────────────
@@ -108,7 +94,7 @@ def list_render_jobs(
):
"""获取 AI 数字人渲染任务列表."""
items, total = svc.list_render_jobs(
user_id=current_user.user.id,
user_id=current_user.id,
project_id=project_id,
status=status,
offset=offset,
@@ -132,7 +118,7 @@ def get_render_job(
svc: AiAvatarRenderService = Depends(_get_service),
):
"""获取渲染任务详情."""
job = svc.get_render_job(job_id, current_user.user.id)
job = svc.get_render_job(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
return job
@@ -148,7 +134,7 @@ def cancel_render_job(
svc: AiAvatarRenderService = Depends(_get_service),
):
"""取消渲染任务(仅 pending 状态可取消)."""
job = svc.cancel_render_job(job_id, current_user.user.id)
job = svc.cancel_render_job(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "cancelled":
@@ -169,7 +155,7 @@ def retry_render_job(
svc: AiAvatarRenderService = Depends(_get_service),
):
"""重试失败的渲染任务."""
job = svc.retry_render_job(job_id, current_user.user.id)
job = svc.retry_render_job(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "pending":
@@ -183,146 +169,7 @@ def retry_render_job(
from app.tasks.ai_avatar_render import execute_ai_avatar_render
execute_ai_avatar_render.delay(job.id)
except Exception as exc:
logger.exception("Celery 任务投递失败重试: job_id=%s err=%s", job.id, exc)
job.status = "failed"
job.error_message = f"任务提交失败:{exc}"
job.updated_at = datetime.now(UTC)
svc.db.commit()
svc.db.refresh(job)
return AiAvatarRenderJobResponse.model_validate(job)
except Exception:
logger.warning("Celery 任务提交失败重试任务已重置但未触发执行: %s", job.id)
return AiAvatarRenderJobResponse.model_validate(job)
# ── POST /{job_id}/smart-cover — 从最终成片智能抽封面(步骤②)────────
@router.post("/{job_id}/smart-cover", response_model=SmartCoverResponse)
def generate_render_smart_cover(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""从最终渲染成片智能抽帧生成封面(MediaKit 抽帧 + 评分选最佳帧 + 转存 OSS).
- 必须等渲染任务 completed 后才可调用(否则返回 400)
- 生成成功后自动更新 render_job 的 cover_config 与 output_cover_url
"""
from app.services.ai_avatar_render_service import AiAvatarRenderService
svc = AiAvatarRenderService(db)
job = svc.get_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "completed":
raise HTTPException(status_code=400, detail="请先完成视频生成")
video_url = (job.output_video_url or "").strip()
if not video_url:
raise HTTPException(status_code=400, detail="渲染成片视频 URL 为空")
try:
# 从最终成片抽帧,帧本身已含标题/B-roll,直接转存 OSS
cover_url = generate_smart_cover(video_url, job_id=job_id, max_frames=5)
except Exception as exc:
logger.error(
"渲染成片智能封面生成异常: user=%s render_id=%s video_url=%s err=%s",
current_user.user.id,
job_id,
video_url[:80],
exc,
exc_info=True,
)
cover_url = ""
if not cover_url:
return SmartCoverResponse(
cover_url="",
status="fallback_failed",
message="智能抽帧失败(MediaKit 不可用或抽帧异常),请稍后重试",
)
# 更新 render_job 的封面字段(异步写入 DB;失败不影响返回)
try:
job.cover_config = {
**(job.cover_config if isinstance(job.cover_config, dict) else {}),
"mode": "auto_frame",
"url": cover_url,
}
job.output_cover_url = cover_url
job.updated_at = datetime.now(UTC)
db.commit()
except Exception as exc:
logger.warning("更新 render_job 封面字段失败(不影响返回): job_id=%s err=%s", job_id, exc)
logger.info(
"渲染成片智能封面生成成功: user=%s render_id=%s cover_url=%s",
current_user.user.id,
job_id,
cover_url[:120],
)
return SmartCoverResponse(cover_url=cover_url, status="completed")
# ── POST /{job_id}/finalize — 封面选定后正式入库成片库 ────────────────────
@router.post("/{job_id}/finalize", response_model=FinalizeRenderResponse)
def finalize_render_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""用户完成封面选择后,将视频正式保存到成片库.
- 必须等渲染任务 completed 后才可调用
- 如果已通过 smart-cover/custom-cover 设置了封面,会自动带上
- 返回成片库视频ID
- 幂等:已 finalize 的任务重复调用会返回 existing 记录
"""
from app.services.ai_avatar_render_service import AiAvatarRenderError, AiAvatarRenderService
svc = AiAvatarRenderService(db)
job = svc.get_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "completed":
raise HTTPException(status_code=400, detail="请先完成视频生成")
# 幂等检查(通过 generation_task_id=job_id 识别,finalize_job 内部也做了一次,这里提前返回简化)
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
existing = (
db.query(GeneratedVideoModel)
.filter(
GeneratedVideoModel.user_id == current_user.user.id,
GeneratedVideoModel.generation_task_id == job_id,
)
.first()
)
if existing is not None:
return FinalizeRenderResponse(
video_id=existing.id,
cover_url=existing.thumbnail_url or "",
status="already_finalized",
)
try:
video = svc.finalize_job(job_id, current_user.user.id)
return FinalizeRenderResponse(
video_id=video.id,
cover_url=video.thumbnail_url or job.output_cover_url or "",
status="success",
)
except AiAvatarRenderError as exc:
status_map = {
"RenderJobNotFound": 404,
"RenderNotCompleted": 400,
"OutputVideoMissing": 400,
}
raise HTTPException(
status_code=status_map.get(exc.code, 400),
detail=str(exc),
) from exc
except Exception as exc:
logger.error("渲染任务finalize失败: job_id=%s err=%s", job_id, exc, exc_info=True)
raise HTTPException(status_code=500, detail=f"保存到成片库失败: {str(exc)}") from exc
return job
+2 -2
View File
@@ -1,5 +1,5 @@
import logging
from typing import Any, Optional
from typing import Any, List, Optional
from app.api.routes._helpers import check_project_access, format_utc_datetime
from app.auth import AuthenticatedUser, get_current_user
@@ -390,7 +390,7 @@ def update_asset_review_status(
return _to_asset_response(updated)
@router.post("/batch", response_model=list[AssetResponse])
@router.post("/batch", response_model=List[AssetResponse])
def batch_get_assets(
request: BatchGetRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
+1 -18
View File
@@ -13,7 +13,7 @@ from typing import Optional
import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_db_session, get_user_repository
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, EmailStr, field_validator
@@ -126,7 +126,6 @@ async def register(
request: RegisterRequest,
user_repository: UserRepository = Depends(get_user_repository),
email_service=Depends(get_auth_email_service),
db=Depends(get_db_session),
) -> RegisterResponse:
use_case = RegisterUserUseCase(
user_repository=user_repository,
@@ -144,22 +143,6 @@ async def register(
if error or response is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=_translate_auth_error(error))
# 新用户注册赠送 50 积分(失败不影响注册)
if settings.points_enabled:
try:
from packages.domain.points_service import PointsService
_svc = PointsService()
_svc.add_points(
user_id=response.user_id,
amount=50,
source="task_reward",
db=db,
description="新用户注册赠送",
)
except Exception as _bonus_err:
import logging
logging.getLogger(__name__).warning("注册送积分失败: user_id=%s err=%s", response.user_id, _bonus_err)
return RegisterResponse(
user_id=response.user_id,
email=response.email,
+8 -8
View File
@@ -8,7 +8,7 @@ import json
import logging
import shutil
import tempfile
from datetime import UTC, datetime, timedelta
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any
from uuid import uuid4
@@ -156,7 +156,7 @@ def _cleanup_expired_uploads() -> int:
if not CHUNK_STORAGE_ROOT.exists():
return 0
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
cleaned = 0
for meta_file in CHUNK_STORAGE_ROOT.glob("*.meta.json"):
@@ -166,7 +166,7 @@ def _cleanup_expired_uploads() -> int:
expires_at = datetime.fromisoformat(meta["expires_at"])
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=UTC)
expires_at = expires_at.replace(tzinfo=timezone.utc)
# Only cleanup uploads that are not actively being uploaded
if expires_at < now and meta.get("status") != "uploading":
@@ -177,8 +177,8 @@ def _cleanup_expired_uploads() -> int:
meta_file.unlink()
cleaned += 1
logger.info(f"Cleaned up expired upload: {upload_id}")
except Exception:
logger.exception("Failed to cleanup upload metadata: %s", meta_file)
except Exception as e:
logger.warning(f"Failed to cleanup upload metadata {meta_file}: {e}")
return cleaned
@@ -226,7 +226,7 @@ async def init_chunked_upload(
# Generate upload ID
upload_id = uuid4().hex
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
expires_at = now + timedelta(hours=CHUNK_EXPIRY_HOURS)
# Create chunk directory
@@ -421,9 +421,9 @@ async def upload_chunk(
# Check expiry
expires_at = datetime.fromisoformat(meta["expires_at"])
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=UTC)
expires_at = expires_at.replace(tzinfo=timezone.utc)
if expires_at < datetime.now(UTC):
if expires_at < datetime.now(timezone.utc):
raise HTTPException(status_code=status.HTTP_410_GONE, detail="Upload has expired")
# Validate chunk index
@@ -1,90 +0,0 @@
"""独立的从素材创建片段端点(不依赖 template_id 路径参数).
POST /api/v1/clips/from-assets
- 与 /api/v1/templates/{template_id}/editor/clips/from-assets 功能一致
- 区别:template_id 从 body 传入(可选),为空时后端自动创建/查找默认模板
- 解决前端首次加载时 templateId 为空导致双斜杠 404 的问题(P0 #1922
- 内部复用 resolve_draft_plan_id 和 create_clips_from_assets_editor 的核心逻辑
"""
from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_asset_repository, get_db_session
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from ._default_template import get_or_create_default_template_id
from .templates_editor.clips import create_clips_from_assets_editor
from .templates_editor.dependencies import resolve_draft_plan_id
from .templates_editor.schemas import ClipsFromAssetsRequest, ClipsFromAssetsResponse
logger = logging.getLogger(__name__)
router = APIRouter(tags=["Clips"])
class StandaloneClipsRequest(ClipsFromAssetsRequest):
"""扩展请求:template_id 可选(不传则后端自动兜底默认模板)。"""
template_id: str | None = None
def _get_editor_services_direct(db: Session) -> tuple[EditTemplateService, EditPlanService]:
"""直接构造服务实例(非 Depends 版本,供独立端点内部调用)。"""
return EditTemplateService(db), EditPlanService(db)
@router.post("/clips/from-assets", response_model=ClipsFromAssetsResponse)
def create_clips_from_assets(
body: StandaloneClipsRequest,
background_tasks: BackgroundTasks,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
asset_repo: SQLAlchemyAssetRepository = Depends(get_asset_repository),
) -> ClipsFromAssetsResponse:
"""从素材批量创建片段(template_id 可选,为空自动兜底)。"""
user_id = str(current_user.user.id)
services = _get_editor_services_direct(db)
# 1. 解析/兜底 template_id,拿到 plan_id
template_id = (body.template_id or "").strip()
if not template_id:
template_id = get_or_create_default_template_id(db, user_id)
if not template_id:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="无法自动创建默认模板,请刷新页面重试",
)
plan_id = resolve_draft_plan_id(
template_id=template_id,
services=services,
current_user=current_user,
db=db,
auto_create_default=False, # 上面已兜底过
)
# 2. 构造标准化请求(去除独立端扩展字段),复用原端点核心逻辑
core_body = ClipsFromAssetsRequest(
asset_ids=body.asset_ids,
clip_type=body.clip_type,
clip_count=body.clip_count,
required_clips_count=body.required_clips_count,
)
# 3. 直接调用原端点函数(此时所有 Depends 依赖已手动传入)
return create_clips_from_assets_editor(
template_id=template_id,
body=core_body,
background_tasks=background_tasks,
plan_id=plan_id,
services=services,
asset_repo=asset_repo,
db=db,
current_user=current_user,
)
+2 -9
View File
@@ -11,7 +11,7 @@ from __future__ import annotations
import ipaddress
import logging
import re
from typing import Any, Optional
from typing import Any, List, Optional
from urllib.parse import urlparse
from app.auth import AuthenticatedUser, get_current_user
@@ -27,7 +27,6 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
)
from packages.application import ListGeneratedVideosByTaskUseCase
from packages.domain.config_schemas import normalize_plan_config
from packages.middleware.points_gate import points_gate
from packages.shared.storage import get_shared_storage_service
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
@@ -42,7 +41,7 @@ router = APIRouter(tags=["Generation"])
class GenerateCoverRequest(BaseModel):
"""AI 封面生成请求体"""
asset_ids: list[str] = Field(default_factory=list, description="素材 ID 列表(确定视频来源)")
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表(确定视频来源)")
cover_type: str = Field(
default="ai_frame",
description="封面类型: ai_frame / manual / upload / ai_regenerate",
@@ -332,7 +331,6 @@ def _is_trusted_media_url(url: str) -> bool:
@router.post("/generate-cover", response_model=GenerateCoverResponse)
@points_gate("ai_cover")
def generate_cover(
body: GenerateCoverRequest,
template_id: str = Query(..., description="模板 ID"),
@@ -768,12 +766,7 @@ def generate_cover(
storage_svc = get_shared_storage_service()
mk_client = get_mediakit_client()
# 从 plan.config 读取完整标题样式,E2 从源素材抽帧时叠加(源素材本身无标题)
# #1901 统一读 "title",兼容老数据 "title_config"
_e2_title_cfg = (plan.config or {}).get("title", {}) or {}
if not isinstance(_e2_title_cfg, dict) or not (_e2_title_cfg.get("text") or "").strip():
_alt = (plan.config or {}).get("title_config", {}) or {}
if isinstance(_alt, dict):
_e2_title_cfg = _alt
if not isinstance(_e2_title_cfg, dict):
_e2_title_cfg = {}
_e2_title_text = (_e2_title_cfg.get("text", "") or "").strip() if _e2_title_cfg.get("enabled", True) else ""
@@ -43,7 +43,6 @@ from packages.application import (
GetGenerationTaskUseCase,
ListGeneratedVideosByTaskUseCase,
)
from packages.middleware.points_gate import points_gate
logger = logging.getLogger(__name__)
@@ -272,7 +271,6 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
@points_gate("ai_video", quantity_field="preview_count")
def create_preview_generation_task(
request: CreatePreviewGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
+114 -196
View File
@@ -16,12 +16,10 @@ 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,
@@ -44,7 +42,6 @@ from packages.application import (
ListGeneratedVideosByTaskUseCase,
)
from packages.domain.smart_match import smart_select_assets
from packages.middleware.points_gate import points_gate
logger = logging.getLogger(__name__)
@@ -61,10 +58,27 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
def _query_voice_durations(db: Session, voice_ids: list[str]) -> list[float]:
"""[已下沉] 路由层兼容别名 → app.services.generation_common.query_voice_durations。"""
from app.services.generation_common import query_voice_durations
"""批量查询配音素材时长(秒),#1749 配音时长分配用。
return query_voice_durations(db, voice_ids)
逐项 try/float 硬化:MagicMock/异常/缺失 → 0.0(无配音不分配,不阻断)。
"""
ids = [v for v in dict.fromkeys(voice_ids or []) if v]
if not ids:
return []
try:
from packages.adapters.sqlalchemy_impl.models import AssetModel
rows = db.query(AssetModel.id, AssetModel.duration).filter(AssetModel.id.in_(ids)).all()
dur_map: dict[str, float] = {}
for row in rows:
try:
dur_map[row[0]] = float(row[1] or 0.0)
except (TypeError, ValueError):
dur_map[row[0]] = 0.0
return [dur_map.get(v, 0.0) for v in ids]
except Exception:
logger.warning("[生成任务] 配音时长查询失败(按无配音处理,不阻断)", exc_info=True)
return [0.0 for _ in ids]
def _to_generation_task_response(task) -> GenerationTaskResponse:
@@ -134,8 +148,6 @@ 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。
@@ -145,8 +157,6 @@ def _select_assets_from_library(
count: 选取数量,0 表示全部(仅 smart 模式有效)
rng: 可选随机源(smart 模式排序噪声用),生产环境不传则内部随机;
测试可注入固定种子或零噪声随机源获得确定性结果。
script_tags: #1970 叙事模式文案标签;非空时标签命中素材优先,不足再用其余素材兜底。
tag_names_by_id: asset_id → 素材标签名列表(素材只存 tag_ids 时由调用方查名称注入)。
Returns:
选中的素材 ID 列表
@@ -156,20 +166,6 @@ 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%)
@@ -182,78 +178,67 @@ def _select_assets_from_library(
return [a.id for a in ready_video_assets]
# #1970 PR3video_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
"""任务入队成功后,回写 EditPlan.configgeneration_task_id + title_config。
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,
)
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
失败只记日志,不影响任务创建。
"""
if not plan_id:
return
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first()
if plan_model is None:
logger.warning("[生成任务] 回写plan.config失败: plan不存在 plan_id=%s", plan_id)
return
current_config = plan_model.config if isinstance(plan_model.config, dict) else {}
merged = dict(current_config)
merged["generation_task_id"] = task_id
# 检查标题是否发生变化,如果变化则清除 cover 字段强制重新生成封面
if title_config:
old_title_config = merged.get("title_config", {}) or {}
old_title_text = (old_title_config.get("text") or "").strip()
new_title_text = (title_config.get("text") or "").strip()
if old_title_text != new_title_text:
# 标题变化,清除旧封面
if "cover" in merged:
del merged["cover"]
logger.info(
"[生成任务] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s",
plan_id,
old_title_text,
new_title_text,
)
merged["title_config"] = title_config
plan_model.config = merged
db.commit()
logger.info(
"[生成任务] 回写plan.config成功: plan_id=%s task_id=%s keys=%s",
plan_id,
task_id,
list(merged.keys()),
)
except Exception as e:
logger.warning(
"[生成任务] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s",
plan_id,
e,
exc_info=True,
)
try:
db.rollback()
except Exception:
pass
def _resolve_project_and_library(
@@ -294,7 +279,6 @@ def _resolve_project_and_library(
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
@points_gate("ai_video", quantity_field="count")
def create_generation_task(
request: CreateGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
@@ -303,63 +287,16 @@ 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, assembly=%s, count=%d",
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%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
@@ -385,29 +322,19 @@ 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",) or narrative_script_tags):
# 项目级模式:未指定 asset_ids 且选择了 smart 模式(或叙事模式按标签匹配)时自动选取
elif project_id and not resolved_asset_ids and request.asset_select_mode in ("smart",):
# 项目级模式:未指定 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(
@@ -471,10 +398,6 @@ 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(
@@ -566,14 +489,25 @@ def create_generation_task(
# 各变体配音时长(查询硬化:异常 → 0.0 不阻断)
voice_durations = _query_voice_durations(db, variant_voices)
# 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(公共函数
from app.services.generation_common import resolve_latest_plan_by_template
# 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(与单任务兜底同源
batch_source_plan_id = request.source_edit_plan_id
if not batch_source_plan_id and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
batch_source_plan_id = (
request.source_edit_plan_id
or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id)
or ""
)
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id,
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
batch_source_plan_id = _latest.id
except Exception:
logger.warning("[生成任务] 批量源 plan 解析失败", exc_info=True)
if not batch_source_plan_id and not request.variant_plan_ids:
# 无任何可用源 plan:批量变体无从选片,明确报错,严禁静默共用/同源
@@ -618,15 +552,7 @@ def create_generation_task(
) from clone_err
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做素材区间避让)
# 变体 1..N-1 独立选片
for task_index in range(1, count):
variant = None
last_err: Exception | None = None
@@ -638,8 +564,6 @@ def create_generation_task(
created_by_user_id=user_id,
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:
@@ -671,18 +595,7 @@ def create_generation_task(
) from last_err
variant_plan_ids.append(variant.id)
# #1855 P0:把新变体的clips区间追加到batch_segments,供下一变体避让
try:
_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)
# ③ 配音时长分配(回传 plan / clone 变体0 均需幂等分配;reselect 已在选片时分配,
# #1855apply_voice_duration_to_plan 已内置幂等判断,重复调用安全)
# ③ 配音时长分配(回传 plan / clone 变体0 均需幂等分配;reselect 已在选片时分配)
for _vi, _pid in enumerate(variant_plan_ids):
_vd = voice_durations[_vi] if _vi < len(voice_durations) else 0.0
if _vd > 0:
@@ -703,13 +616,24 @@ def create_generation_task(
)
_single_vd: list[float] = _query_voice_durations(db, _voices)
_single_dur = _single_vd[0] if _single_vd else 0.0
from app.services.generation_common import resolve_latest_plan_by_template
_single_plan = request.source_edit_plan_id
if not _single_plan and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_single_plan = (
request.source_edit_plan_id
or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id)
or ""
)
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id,
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
_single_plan = _latest.id
except Exception:
logger.warning("[生成任务] 单任务源 plan 解析失败", exc_info=True)
if _single_dur > 0 and _single_plan:
from app.services.edit_plan_service import EditPlanService
@@ -821,11 +745,6 @@ 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(
@@ -916,7 +835,6 @@ 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,
@@ -30,74 +30,6 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _get_or_create_default_template_id(db: Session, user_id: str) -> str | None:
"""为用户查找一个有效模板;若不存在则自动创建默认配音模板。
前端 #1911 删除了模板选择 UI,当调用方未传 template_id/source_edit_plan_id
时(如剪辑页首次进入直接选片),后端兜底查找/创建默认模板,避免 400。
Returns:
template_id(字符串);失败时返回 None。
"""
from packages.adapters.sqlalchemy_impl.models import TemplateClipConfigModel, TemplateModel
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import CreateTemplateCommand, SegmentCommand
from packages.application.template.use_cases import CreateTemplateUseCase
# 1. 先查已有有效模板(is_active=True 且存在片段配置)
existing = (
db.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
# 验证该模板是否有片段配置;若没有继续尝试创建默认
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
# 2. 无有效模板 → 自动创建默认配音模板
try:
repo = SQLAlchemyTemplateRepository(db)
cmd = CreateTemplateCommand(
user_id=user_id,
name="默认配音模板",
mode="voice_over",
category="default",
tags=[],
title_config={},
subtitle_config={},
bgm_config={},
estimated_duration=0.0,
segments=[
SegmentCommand(
segment_order=0,
duration_min=1.0,
duration_max=30.0,
material_type=None,
),
],
)
use_case = CreateTemplateUseCase(repo)
tpl = use_case.execute(cmd)
logger.info(
"[variant-plans] 自动创建默认模板: user=%s tpl=%s",
user_id,
tpl.id,
)
return tpl.id
except Exception:
logger.exception("[variant-plans] 自动创建默认模板失败: user=%s", user_id)
return None
class VariantPlanRequest(BaseModel):
"""轻量选片请求体(与前端 variantPlans.ts 契约一致)。"""
@@ -111,8 +43,8 @@ class VariantPlanRequest(BaseModel):
@model_validator(mode="after")
def _validate(self) -> "VariantPlanRequest":
# 不再强制要求 template_id / source_edit_plan_id
# 后端在路由内会自动查找/创建默认模板兜底(#1911 后前端不再显式选模板)。
if not self.template_id.strip() and not self.source_edit_plan_id.strip():
raise ValueError("template_id 与 source_edit_plan_id 至少需要提供一个")
try:
resolve_variant_voice_ids(
count=self.count,
@@ -158,19 +90,25 @@ def create_variant_plans(
except VariantVoiceError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
# 解析源 plan:显式传入优先;否则按 template_id + user 查最新(公共函数)
from app.services.generation_common import resolve_latest_plan_by_template
# 解析源 plan:显式传入优先;否则按 template_id + user 查最新
source_plan_id = request.source_edit_plan_id.strip()
template_id = request.template_id.strip()
if not source_plan_id and request.template_id.strip():
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
# P0 兜底:前端 #1911 已删除模板选择 UI,调用方可能不传 template_id
# 此时自动为该用户查找/创建默认模板。
if not source_plan_id and not template_id:
template_id = _get_or_create_default_template_id(db, user_id) or ""
if not source_plan_id and template_id:
source_plan_id = resolve_latest_plan_by_template(db, template_id=template_id, user_id=user_id) or ""
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id.strip(),
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
source_plan_id = _latest.id
except Exception:
logger.warning("[variant-plans] 源 plan 解析失败", exc_info=True)
if not source_plan_id:
raise HTTPException(
@@ -184,7 +122,7 @@ def create_variant_plans(
voice_durations = _query_voice_durations(db, voices)
except Exception:
logger.exception("[variant-plans] 配音时长查询失败(按占位段长选片)")
logger.warning("[variant-plans] 配音时长查询失败(按占位段长选片)", exc_info=True)
voice_durations = [0.0] * request.count
from app.services.edit_plan_service import EditPlanService
@@ -205,7 +143,7 @@ def create_variant_plans(
except HTTPException:
raise
except Exception as e:
logger.exception("[variant-plans] 选片异常")
logger.error("[variant-plans] 选片异常: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="选片失败,请稍后重试") from e
# 组装 clips 响应
-231
View File
@@ -1,231 +0,0 @@
"""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,
)
+3 -3
View File
@@ -1,4 +1,4 @@
from datetime import UTC, datetime
from datetime import datetime, timezone
import psycopg
import redis
@@ -13,7 +13,7 @@ router = APIRouter(tags=["Health"])
async def health_check():
return {
"status": "healthy",
"timestamp": datetime.now(UTC).isoformat(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"version": settings.APP_VERSION,
}
@@ -33,7 +33,7 @@ async def startup_check():
all_ready = all(check["status"] == "healthy" for check in checks.values())
response = {
"status": "started" if all_ready else "starting",
"timestamp": datetime.now(UTC).isoformat(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"checks": checks,
}
if not all_ready:
+59 -216
View File
@@ -1,39 +1,26 @@
"""对口型 API 路由 — #1796 MediaKit 对口型, #1809 参数调整, #1845 配音前置.
"""对口型 API 路由 — #1796 MediaKit 对口型, #1809 参数调整.
接口:
POST /api/v1/lipsync/jobs 提交对口型任务(支持 TTS/直传/预合成 三种模式)
POST /api/v1/lipsync/jobs 提交对口型任务
GET /api/v1/lipsync/jobs 任务列表
GET /api/v1/lipsync/jobs/{id} 任务详情
POST /api/v1/lipsync/jobs/{id}/refresh 刷新任务状态
POST /api/v1/lipsync/jobs/{id}/cancel 取消任务
POST /api/v1/lipsync/tts-preview #1845 步骤1 TTS 预合成(同步 HTTP~2-3s
"""
from __future__ import annotations
import logging
import math
from datetime import UTC
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.dependencies import (
get_db_session,
get_voice_clone_profile_repository,
)
from app.schemas.lipsync import (
AiAvatarTtsPreviewRequest,
AiAvatarTtsPreviewResponse,
CreateLipsyncJobRequest,
LipsyncJobResponse,
)
from app.dependencies import get_cosyvoice_service, get_db_session, get_voice_clone_profile_repository
from app.schemas.lipsync import CreateLipsyncJobRequest, LipsyncJobResponse
from app.services.lipsync_service import LipsyncService
from app.services.mediakit_client import MediaKitError
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
logger = logging.getLogger(__name__)
@@ -42,13 +29,35 @@ router = APIRouter()
def _get_service(
db: Session = Depends(get_db_session),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
) -> LipsyncService:
# voice_clone_repo 用于克隆音色 profile 解析
return LipsyncService(
db,
voice_clone_repo=voice_clone_repo,
)
return LipsyncService(db, cosyvoice_service=cosyvoice_service)
def _resolve_voice_id(
raw_voice_id: str,
user_id: str,
voice_clone_repo,
) -> str:
"""解析 voice_id:支持预设音色 ID 或克隆音色 profile UUID.
与 TTS 路由保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id。
"""
try:
profile = voice_clone_repo.get(raw_voice_id)
except Exception as exc:
logger.error("查询克隆音色失败: voice_id=%s, error=%s", raw_voice_id, exc)
raise HTTPException(
status_code=400,
detail=f"voice_id 无效: {raw_voice_id}",
) from exc
if profile is not None:
if profile.user_id != user_id:
raise HTTPException(status_code=403, detail="无权访问该音色")
if not profile.voice_id:
raise HTTPException(status_code=400, detail="音色克隆尚未完成,请稍后再试")
return profile.voice_id
return raw_voice_id
# ── POST /jobs — 提交对口型任务 ───────────────────────────────────────────
@@ -58,194 +67,55 @@ def _get_service(
def create_lipsync_job(
body: CreateLipsyncJobRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
svc: LipsyncService = Depends(_get_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
):
user_id = current_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_digital_human"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
# 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟)
if body.audio_url and body.audio_duration and body.audio_duration > 0:
est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0))
elif body.script_text:
est_minutes = max(1.0, math.ceil(len(body.script_text) / 240))
else:
est_minutes = 1.0
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(current_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(current_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
"""提交对口型任务.
三种模式:
- TTS 直生(旧版/降级):传 {video_url, voice_id, script_text, speed?, emotion?}
后端 dispatch Celery 异步任务。
- 直接音频:传 {video_url, audio_url},后端同步下载+算timings+提交MediaKit。
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings}
后端同步ffprobe+写入timings+直接提交MediaKit~2-3s)。
#1809: 前端传 {voice_id, script_text, video_url}
后端内部调 TTS 合成音频,再提交 MediaKit。
"""
# 解析 voice_id(支持克隆音色 profile UUID
actual_voice_id = _resolve_voice_id(body.voice_id, current_user.id, voice_clone_repo)
try:
job = svc.create_job(
user_id=user_id,
user_id=current_user.id,
video_url=body.video_url,
audio_url=body.audio_url,
audio_duration=body.audio_duration,
sentence_timings=body.sentence_timings,
voice_id=body.voice_id,
voice_id=actual_voice_id,
script_text=body.script_text,
speed=body.speed,
emotion=body.emotion,
enable_video_loop=body.enable_video_loop,
project_id=body.project_id,
)
except ValueError as exc:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}")
# 参数无效(如 voice_id 格式不对、文本过长等)
raise HTTPException(status_code=400, detail=str(exc)) from exc
except MediaKitError as exc:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}")
status_code = 502
if exc.code in ("VoiceForbidden",):
status_code = 403
elif exc.code in ("InvalidInput", "TTSInvalidParam", "VoiceNotReady"):
status_code = 400
except CosyVoiceError as exc:
# TTS 合成基础设施失败(API/网络/认证)
raise HTTPException(
status_code=status_code,
status_code=502,
detail={"code": "TTSSynthesisFailed", "message": str(exc)},
) from exc
except MediaKitError as exc:
raise HTTPException(
status_code=502,
detail={
"code": exc.code,
"message": str(exc),
"request_id": getattr(exc, "request_id", ""),
"request_id": exc.request_id,
},
) from exc
except Exception as exc:
# 兜底:任何未预期的错误返回 400 而非 500
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型异常退积分异常: err={refund_err}")
raise HTTPException(
status_code=400,
detail=f"创建对口型任务失败: {exc}",
) from exc
# 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等)
# 若任务已创建且状态为 failed,退费
if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed":
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}")
return job
# ── POST /tts-preview — #1845 步骤1 TTS 预合成 ──────────────────────────
@router.post("/tts-preview", response_model=AiAvatarTtsPreviewResponse)
def preview_tts(
body: AiAvatarTtsPreviewRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
svc: LipsyncService = Depends(_get_service),
):
user_id = current_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_digital_human"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(current_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(current_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
"""步骤1「生成配音」同步 TTS 预合成.
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL~24h 有效)。
耗时约 2-3 秒。
"""
try:
result = svc.preview_tts(
user_id=user_id,
voice_id=body.voice_id,
script_text=body.script_text,
speed=body.speed,
emotion=body.emotion,
)
except MediaKitError as exc:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
status_code = 400
if exc.code in ("VoiceForbidden",):
status_code = 403
elif exc.code in ("TTSNoAudio",):
status_code = 502
raise HTTPException(
status_code=status_code,
detail={
"code": exc.code,
"message": str(exc),
},
) from exc
except Exception as exc:
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}")
raise HTTPException(
status_code=400,
detail=f"TTS 合成失败: {exc}",
) from exc
return result
# ── GET /jobs — 任务列表 ─────────────────────────────────────────────────
@@ -260,7 +130,7 @@ def list_lipsync_jobs(
):
"""获取对口型任务列表."""
items, total = svc.list_jobs(
user_id=current_user.user.id,
user_id=current_user.id,
project_id=project_id,
status=status,
offset=offset,
@@ -280,40 +150,13 @@ def list_lipsync_jobs(
@router.get("/jobs/{job_id}", response_model=LipsyncJobResponse)
def get_lipsync_job(
job_id: str,
background: BackgroundTasks,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""获取对口型任务详情."""
job = svc.get_job(job_id, current_user.user.id)
job = svc.get_job(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.status not in ("completed", "failed"):
# 三层防御 ①:如果距上次更新超过 30 秒,同步刷新一次(避免 background task
# 静默失败导致前端永远看到 running);否则挂后台异步刷新(避免阻塞轮询)。
from datetime import datetime as _dt
_now = _dt.now(UTC)
_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)
if refreshed is not None:
job = refreshed
except Exception as exc: # noqa: BLE001
logger.error("同步刷新对口型状态失败 job_id=%s err=%s", job_id, exc, exc_info=True)
background.add_task(svc.refresh_job_status, job_id, current_user.user.id)
else:
background.add_task(svc.refresh_job_status, job_id, current_user.user.id)
return job
@@ -327,7 +170,7 @@ def refresh_lipsync_job(
svc: LipsyncService = Depends(_get_service),
):
"""从 MediaKit 拉取最新状态并更新."""
job = svc.refresh_job_status(job_id, current_user.user.id)
job = svc.refresh_job_status(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
return job
@@ -342,13 +185,13 @@ def cancel_lipsync_job(
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""取消对口型任务(仅 pending/tts_processing/submitted 状态可取消)."""
job = svc.cancel_job(job_id, current_user.user.id)
"""取消对口型任务(仅 pending/submitted 状态可取消)."""
job = svc.cancel_job(job_id, current_user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.status != "cancelled":
raise HTTPException(
status_code=400,
detail=f"任务状态 {job.status} 不可取消,仅 pending/tts_processing/submitted 可取消",
detail=f"任务状态 {job.status} 不可取消,仅 pending/submitted 可取消",
)
return job
-339
View File
@@ -1,339 +0,0 @@
"""积分 & 会员 API 路由 (#1895)
导出两个 router
- points_router: 积分相关路由,前缀 /points
- usage_router: 每日额度路由,前缀 /usage
"""
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.points import (
DailyUsageResponse,
MembershipStatusResponse,
PointRuleItem,
PointsBalanceResponse,
PointsCheckRequest,
PointsCheckResponse,
PointsDeductRequest,
PointsOrderResponse,
PointsPackageItem,
PointsPackagesResponse,
PointsRechargeRequest,
PointsRefundRequest,
PointsRulesResponse,
PointsTransactionsResponse,
SimpleMessageResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.domain.points_rules import (
FREE_USER_MULTIPLIER,
MEMBER_DISCOUNT,
POINTS_PACKAGES,
POINTS_SCENES,
calculate_points_cost,
)
from packages.domain.points_service import PointsService
logger = logging.getLogger(__name__)
# ── 两个 router ──
points_router = APIRouter()
usage_router = APIRouter()
def _get_service() -> PointsService:
return PointsService()
def _is_member(user: AuthenticatedUser) -> bool:
"""判断用户是否为付费会员。"""
return getattr(user.user, "is_member", False)
def _member_type(user: AuthenticatedUser) -> str | None:
return getattr(user.user, "member_type", None)
# ════════════════════════════════════════════════════════════════
# 积分相关路由 (prefix=/points)
# ════════════════════════════════════════════════════════════════
@points_router.get("/balance", response_model=PointsBalanceResponse)
def get_balance(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询当前用户积分余额 + 会员状态。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
return PointsBalanceResponse(
balance=account["balance"],
total_earned=account["total_earned"],
total_spent=account["total_spent"],
is_member=_is_member(current_user),
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
)
@points_router.get("/transactions", response_model=PointsTransactionsResponse)
def get_transactions(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
type: Optional[str] = Query(None, description="筛选类型: add/deduct"),
source: Optional[str] = Query(None, description="筛选来源场景"),
start_date: Optional[datetime] = Query(None),
end_date: Optional[datetime] = Query(None),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询积分流水(分页+筛选)。"""
svc = _get_service()
result = svc.get_transactions(
user_id=current_user.user.id,
db=db,
page=page,
page_size=page_size,
type_filter=type,
source_filter=source,
start_date=start_date,
end_date=end_date,
)
return PointsTransactionsResponse(**result)
@points_router.get("/rules", response_model=PointsRulesResponse)
def get_rules(
_current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询所有积分消耗规则。"""
rules = []
for scene_key, scene_data in POINTS_SCENES.items():
rules.append(
PointRuleItem(
scene_key=scene_key,
name=scene_data["name"],
base_points=scene_data["base_points"],
unit=scene_data["unit"],
extra_per_30s=scene_data.get("extra_per_30s"),
description=scene_data.get("description", ""),
)
)
return PointsRulesResponse(
rules=rules,
free_user_multiplier=FREE_USER_MULTIPLIER,
)
@points_router.get("/packages", response_model=PointsPackagesResponse)
def get_packages(
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询可购买的积分包列表。"""
packages = []
for code, pkg in POINTS_PACKAGES.items():
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
packages.append(
PointsPackageItem(
code=code,
name=pkg["name"],
points=pkg["points"],
price_cents=pkg["price_cents"],
unit_price=unit_price,
)
)
mt = _member_type(current_user)
discount = MEMBER_DISCOUNT.get(mt) if mt else None
return PointsPackagesResponse(packages=packages, user_discount=discount)
@points_router.post("/check", response_model=PointsCheckResponse)
def check_points(
body: PointsCheckRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""消费前检查余额是否足够。未知 scene_key 返回 400(而非 500)。"""
if body.scene_key not in POINTS_SCENES:
raise HTTPException(
status_code=400,
detail={
"code": "UNKNOWN_SCENE",
"message": f"未知场景: {body.scene_key}",
"valid_scenes": sorted(POINTS_SCENES.keys()),
},
)
is_mem = _is_member(current_user)
mt = _member_type(current_user)
# 混剪场景先检查免费额度
is_free_quota = False
if body.scene_key == "ai_video" and not is_mem:
svc = _get_service()
if svc.check_daily_free_clip(current_user.user.id, db):
is_free_quota = True
required = calculate_points_cost(
body.scene_key,
is_mem,
quantity=body.quantity or 1,
duration_minutes=body.duration_minutes or 0,
member_type=mt,
)
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
balance = account["balance"]
return PointsCheckResponse(
allowed=is_free_quota or balance >= required,
required_points=required,
current_balance=balance,
remaining_after=balance - required,
is_free_quota=is_free_quota,
)
@points_router.post("/deduct", response_model=SimpleMessageResponse)
def deduct_points(
body: PointsDeductRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分扣减(内部服务调用)。"""
svc = _get_service()
result = svc.deduct_points(
user_id=current_user.user.id,
amount=body.amount,
source=body.scene_key,
db=db,
description=body.description or "",
ref_id=body.ref_id or "",
)
if not result["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {body.amount},余额 {result['balance']}",
},
)
return SimpleMessageResponse(
success=True,
message=f"扣减 {body.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/refund", response_model=SimpleMessageResponse)
def refund_points(
body: PointsRefundRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分退还(内部服务调用)。"""
from packages.adapters.sqlalchemy_impl.models import PointsTransactionModel
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:
raise HTTPException(status_code=403, detail="无权退还他人积分")
svc = _get_service()
result = svc.refund_points(
user_id=current_user.user.id,
amount=txn.amount,
source=txn.source,
db=db,
ref_id=body.transaction_id,
description=body.reason or f"退还: {txn.description}",
)
if not result["success"]:
raise HTTPException(status_code=500, detail="退还失败")
return SimpleMessageResponse(
success=True,
message=f"退还 {txn.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/recharge", response_model=PointsOrderResponse)
def create_recharge_order(
body: PointsRechargeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""创建积分充值订单。pay_params 在支付通道接入后填入 prepay_id/payment_url;当前为空 dict。"""
svc = _get_service()
try:
order = svc.create_order(
user_id=current_user.user.id,
order_type="points",
product_code=body.package_id,
db=db,
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from None
package = POINTS_PACKAGES.get(body.package_id, {})
now = datetime.now(timezone.utc)
expire_at = now + timedelta(hours=48)
# TODO: 接入微信/支付宝后填充真实 prepay_id / payment_url
order["points_amount"] = package.get("points", 0)
order["pay_params"] = {}
order["expire_at"] = expire_at.isoformat()
return PointsOrderResponse(**order)
@points_router.get("/subscription/membership", response_model=MembershipStatusResponse)
def get_membership_status(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""获取当前用户会员状态(聚合信息)。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
is_mem = _is_member(current_user)
max_resolution = "1080p" if is_mem else "720p"
return MembershipStatusResponse(
is_member=is_mem,
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
points_balance=account["balance"],
max_resolution=max_resolution,
)
# ════════════════════════════════════════════════════════════════
# 每日额度路由 (prefix=/usage)
# ════════════════════════════════════════════════════════════════
@usage_router.get("/daily", response_model=DailyUsageResponse)
def get_daily_usage(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询今日免费混剪额度使用情况。"""
svc = _get_service()
result = svc.get_daily_usage(current_user.user.id, db)
return DailyUsageResponse(**result)
# 为了向后兼容,也导出一个不带后缀的 router(方便旧引用)
router = points_router
-554
View File
@@ -1,554 +0,0 @@
"""Scripts AI 能力路由 — Issue #1893/#1963.
三个 AI 工具接口(均挂载在 /api/v1/scripts 前缀下):
- 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
"""
from __future__ import annotations
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
from app.schemas.scripts_ai import (
AiGenerateTitlesRequest,
AiGenerateTitlesResponse,
AiRewriteRequest,
AiRewriteResponse,
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,
transcribe_to_text,
)
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.orm import Session
from packages.middleware.points_gate import points_gate
from packages.shared.ai_client import get_doubao_client
logger = logging.getLogger(__name__)
router = APIRouter()
_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 _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="无效的抖音链接,仅支持 http(s) 协议",
)
if not _DOUYIN_HOST_RE.search(host):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="无效的抖音链接,仅支持 douyin.com 域名(v.douyin.com 短链或 www.douyin.com 长链)",
)
return url
# ── MediaKitClient ASR 扩展(monkey patch) ────────────────────────────
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),
):
page_url = _extract_and_validate_douyin_url(request.url)
_dbg("page_url", page_url)
# ── 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())
direct_url = result.video_url if result else None
feed_desc = (result.desc or "").strip() if result else ""
# 图文视频(无 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,
)
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)
# ── Phase BASR 转文字 ──
mk_client = get_mediakit_client()
text = ""
duration = 0.0
# B1MediaKit 云端 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,
)
else:
logger.info("抖音 MediaKit ASR 返回空文本(无旁白/BGM视频)")
except MediaKitError as exc:
logger.warning("MediaKit ASR 失败,回退本地 ASR: %s", exc)
text = ""
# 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 = ""
# ── Phase C:结果判定 & 兜底 ──
# ASR 空结果(无旁白视频)→ 使用解析源 desc 兜底
if not text and feed_desc:
text = feed_desc
logger.info("抖音 ASR 空结果,使用解析源 desc 兜底: desc_len=%d", len(text))
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=page_url,
)
# ── 2. AI 文案改写 ────────────────────────────────────────────────────
@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),
):
content = (request.content or "").strip()
if not content:
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 = system_prompt + "\n风格要求:" + style
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": "请改写以下文案:\n\n" + content},
]
try:
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="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)
# ── 3. AI 标题生成 ────────────────────────────────────────────────────
@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),
):
content = (request.content or "").strip()
if not content:
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)
titles = result.get("titles", [])[:count]
return AiGenerateTitlesResponse(titles=titles)
+73 -90
View File
@@ -4,17 +4,15 @@ from __future__ import annotations
import logging
from dataclasses import replace
from datetime import UTC, datetime
from typing import Any
from datetime import datetime, timezone
from typing import List
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,
@@ -28,23 +26,48 @@ logger = logging.getLogger(__name__)
router = APIRouter()
# ============ 会员展示名称(与 packages.domain.points_rules.MEMBERSHIP_PRICES 对应)============
# ============ 配额定义(硬编码,后续可迁移到配置中心) ============
_PLAN_NAMES: dict[str, str] = {
MembershipType.FREE: "免费用户",
MembershipType.MONTHLY: "月卡会员",
MembershipType.QUARTERLY: "季卡会员",
MembershipType.YEARLY: "年卡会员",
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},
}
# ============ Helper Functions ============
def _get_plan_name(plan_id: str) -> str:
return _PLAN_NAMES.get(plan_id, "免费用户")
"""获取套餐显示名称"""
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)
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
"""构建订阅信息响应"""
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
if user.user.subscription_expires_at:
period_end = user.user.subscription_expires_at.isoformat()
period_start = now.isoformat()
@@ -52,20 +75,15 @@ 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=plan_id,
plan_name=_get_plan_name(plan_id),
plan_id=user.user.subscription_plan or "free",
plan_name=_get_plan_name(user.user.subscription_plan or "free"),
status=user.user.subscription_status or "active",
billing_cycle=plan_id if plan_id != MembershipType.FREE else BillingCycle.MONTHLY,
billing_cycle="monthly",
current_period_start=period_start,
current_period_end=period_end,
amount=0 if plan_id == MembershipType.FREE else 0, # 金额由前端 /plans 接口展示
amount=_get_plan_price(user.user.subscription_plan or "free", "monthly"),
auto_renew=True,
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
)
@@ -82,43 +100,10 @@ async def get_current_subscription(
return _build_subscription_info(current_user)
@router.get("/plans")
def list_membership_plans(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> dict[str, list[dict[str, Any]]]:
"""查询所有会员档位(供前端会员购买页展示)。
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
"""
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
plans: list[dict[str, Any]] = []
for plan_id, info in MEMBERSHIP_PRICES.items():
days = info["duration_days"]
monthly_cents = round(info["price_cents"] * 30 / days)
features: dict[str, Any] = {"max_resolution": "1080p"}
if plan_id == MembershipType.MONTHLY:
features.update({"free_clips_daily": 2})
elif plan_id == MembershipType.QUARTERLY:
features.update({"free_clips_daily": 5})
elif plan_id == MembershipType.YEARLY:
features.update({"free_clips_daily": "unlimited"})
plans.append({
"plan_id": plan_id,
"name": info["name"],
"price_cents": info["price_cents"],
"monthly_price_cents": monthly_cents,
"duration_days": days,
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
"features": features,
})
return {"plans": plans}
@router.get("/billing-records", response_model=list[BillingRecord])
@router.get("/billing-records", response_model=List[BillingRecord])
async def get_billing_records(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> list[BillingRecord]:
) -> List[BillingRecord]:
"""获取账单记录列表"""
from packages.adapters.sqlalchemy_impl.billing_repository import SQLAlchemyBillingRepository
from packages.adapters.sqlalchemy_impl.session import SessionLocal
@@ -133,7 +118,7 @@ async def get_billing_records(
return [
BillingRecord(
id=r.id,
plan_name=_get_plan_name(r.plan_name),
plan_name=r.plan_name,
amount=r.amount,
billing_cycle=r.billing_cycle,
status=r.status,
@@ -147,10 +132,6 @@ 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,
@@ -159,45 +140,47 @@ async def change_plan(
) -> ChangePlanResponse:
"""变更订阅套餐(升级/降级)"""
# TODO: 接入支付验证(支付宝/微信支付)
target_plan = request.target_plan_id
if target_plan not in _VALID_PLANS:
valid_plans = {"free", "standard", "pro", "enterprise"}
if request.target_plan_id not in valid_plans:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的会员类型。支持: {', '.join(sorted(_VALID_PLANS))}",
detail=f"无效的套餐ID。支持的套餐: {', '.join(valid_plans)}",
)
if request.billing_cycle not in _VALID_CYCLES:
valid_cycles = {"monthly", "yearly"}
if request.billing_cycle not in valid_cycles:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的计费周期。支持: {', '.join(sorted(_VALID_CYCLES))}",
detail="无效的计费周期。支持: monthly, yearly",
)
user = current_user.user
current_plan = user.subscription_plan or MembershipType.FREE
# 旧档位归一化,避免永远显示"您已经是xxx"
if current_plan in {"standard", "pro", "enterprise"}:
current_plan = MembershipType.MONTHLY
current_plan = user.subscription_plan or "free"
target_plan = request.target_plan_id
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=-1, # 付费会员不限项目数
max_storage_gb=100,
max_projects=quotas["max_projects"],
max_storage_gb=quotas["max_storage_gb"],
)
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),
)
@@ -209,11 +192,10 @@ async def cancel_subscription(
) -> SimpleResponse:
"""取消订阅"""
user = current_user.user
plan_id = user.subscription_plan or MembershipType.FREE
if plan_id == MembershipType.FREE:
if user.subscription_plan == "free":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="免费用户无需取消订阅",
detail="体验版无需取消",
)
updated_user = replace(user, subscription_status="cancelled")
@@ -221,7 +203,7 @@ async def cancel_subscription(
return SimpleResponse(
success=True,
message="订阅已取消,当前周期结束后将降级为免费用户",
message="订阅已取消,当前周期结束后停止服务",
)
@@ -247,14 +229,11 @@ 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(
{
@@ -267,20 +246,19 @@ async def payment_callback(
}
)
# 在事务中标记支付成功并更新订阅
repo.mark_paid(record_id, payment_method, payment_id)
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)
# 计算到期时间
days = 365 if billing_cycle == "yearly" else 30
expires_at = datetime.now(timezone.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("支付回调处理失败: user_id=%s, plan=%s, error=%s", user_id, plan, e)
logger.error(f"支付回调处理失败: user_id={user_id}, plan={plan}, error={e}")
# 不返回原始异常信息,避免泄漏内部实现细节
raise HTTPException(status_code=500, detail="支付处理失败,请稍后重试") from e
finally:
session.close()
@@ -292,5 +270,10 @@ 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,
)
+386 -39
View File
@@ -1,13 +1,4 @@
"""Template 列表路由(供生成页自动选模板).
保留:
- GET /templates:列表查询(生成页使用)
- 默认模板自动创建兜底逻辑(复用 _default_template.get_or_create_default_template_id
其他模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等 HTTP 端点
已在 PR#1918 中删除(前端 PR#1911 已删除 my-templates / editing-planner /
templates 管理页面)。
"""
"""Template CRUD + generate + category routes."""
from __future__ import annotations
@@ -16,19 +7,53 @@ import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.template import (
CategoryResponse,
CopyTemplateRequest,
CreateCategoryRequest,
CreateTemplateRequest,
GenerateWarningResponse,
ListCategoriesResponse,
ListTagsResponse,
ListTemplatesResponse,
SegmentResponse,
TemplateResponse,
TemplateUsageResponse,
ToggleFavoriteResponse,
UpdateTemplateRequest,
ValidateTemplateRequest,
ValidateTemplateResponse,
)
from fastapi import APIRouter, Depends, Query
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import ListTemplatesFilter
from packages.application.template.use_cases import CountTemplatesUseCase, ListTemplatesUseCase
from ._default_template import get_or_create_default_template_id
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
SegmentCommand,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
from packages.application.template.use_cases import (
CopyTemplateUseCase,
CountTemplatesUseCase,
CreateCategoryUseCase,
CreateTemplateUseCase,
DeleteCategoryUseCase,
DeleteTemplateUseCase,
GetTemplateUseCase,
ListCategoriesUseCase,
ListTagsUseCase,
ListTemplatesUseCase,
NotFoundError,
UpdateTemplateUseCase,
ValidateTemplateUseCase,
ValidationError,
)
router = APIRouter()
@@ -37,32 +62,354 @@ def _get_template_repository(session: Session = Depends(get_db_session)) -> SQLA
return SQLAlchemyTemplateRepository(session)
@router.get("", response_model=ListTemplatesResponse, summary="获取模板列表")
def _segment_to_response(seg) -> SegmentResponse:
return SegmentResponse(
id=seg.id,
template_id=seg.template_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
created_at=seg.created_at,
updated_at=seg.updated_at,
)
def _to_response(template, usage_count: int = 0) -> TemplateResponse:
return TemplateResponse(
id=template.id,
user_id=template.user_id,
name=template.name,
mode=template.mode,
category=template.category,
tags=template.tags,
title_config=template.title_config,
subtitle_config=template.subtitle_config,
bgm_config=template.bgm_config,
estimated_duration=template.estimated_duration,
segments=[_segment_to_response(s) for s in getattr(template, "segments", [])],
is_active=template.is_active,
usage_count=usage_count,
created_at=template.created_at,
updated_at=template.updated_at,
)
# ── Template CRUD ──
@router.get("", response_model=ListTemplatesResponse)
def list_templates(
mode: str | None = Query(None, description="编辑模式:generic/vlog/storyboard,不传返回全部"),
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
category: str | None = Query(None, description="按分类筛选"),
tag: str | None = Query(None, description="按标签筛选"),
page: int = Query(1, ge=1, description="页码,从 1 开始"),
page_size: int = Query(20, ge=1, le=100, description="每页条数,默认 20"),
current_user: AuthenticatedUser = Depends(get_current_user),
repo: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
db: Session = Depends(get_db_session),
):
"""获取用户可用的模板列表(仅返回 active 状态)。"""
user_id = str(current_user.user.id)
# P0 兜底:无有效模板时自动创建默认配音模板(解决新用户首次进入生成页 404)
get_or_create_default_template_id(db, user_id)
keyword: str | None = Query(None, description="按名称关键词搜索"),
mode: str | None = Query(None, description="按剪辑模式筛选"),
valid_only: bool = Query(
False,
description="仅返回已配置片段的模板(剪辑页传 true;模板编辑器不传,可查看全部模板含草稿)",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTemplatesResponse:
user_id = authenticated_user.user.id
try:
tpl_filter = ListTemplatesFilter(
category=category,
tag=tag,
keyword=keyword,
mode=mode,
valid_only=valid_only,
)
use_case = ListTemplatesUseCase(template_repository)
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
count_use_case = CountTemplatesUseCase(template_repository)
total = count_use_case.execute(user_id, filter=tpl_filter)
list_uc = ListTemplatesUseCase(repo)
count_uc = CountTemplatesUseCase(repo)
filters = ListTemplatesFilter(
category=category,
tag=tag,
mode=mode,
valid_only=True, # 仅返回 active + 有片段配置
# 批量查询使用次数
items = []
for t in templates:
usage = template_repository.get_usage_count(t.id)
items.append(_to_response(t, usage_count=usage))
except Exception:
logger.exception("list_templates 查询失败: user_id=%s", user_id)
return ListTemplatesResponse(items=[], total=0)
return ListTemplatesResponse(
items=items,
total=total,
)
skip = (page - 1) * page_size
templates = list_uc.execute(user_id, skip=skip, limit=page_size, filter=filters)
total = count_uc.execute(user_id, filter=filters)
items = [TemplateResponse.model_validate(tpl, from_attributes=True) for tpl in templates]
return ListTemplatesResponse(items=items, total=total)
@router.get("/{template_id}", response_model=TemplateResponse)
def get_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
try:
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
usage = template_repository.get_usage_count(template_id)
except Exception as _e:
logger.exception("get_template 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="模板查询失败") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return _to_response(template, usage_count=usage)
@router.post("", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def create_template(
request: CreateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = CreateTemplateCommand(
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
],
)
use_case = CreateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.patch("/{template_id}", response_model=TemplateResponse)
def update_template(
template_id: str,
request: UpdateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = UpdateTemplateCommand(
template_id=template_id,
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=(
[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
]
if request.segments is not None
else None
),
)
use_case = UpdateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteTemplateUseCase(template_repository)
deleted = use_case.execute(template_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return
@router.post("/{template_id}/copy", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def copy_template(
template_id: str,
request: CopyTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
"""复制模板(含所有片段配置)"""
user_id = authenticated_user.user.id
command = CopyTemplateCommand(
template_id=template_id,
user_id=user_id,
new_name=request.new_name,
)
use_case = CopyTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.get("/{template_id}/usage", response_model=TemplateUsageResponse)
def get_template_usage(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateUsageResponse:
"""获取模板使用次数(关联的剪辑计划数量)"""
user_id = authenticated_user.user.id
# 鉴权:确保模板存在且属于当前用户
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
usage = template_repository.get_usage_count(template_id)
return TemplateUsageResponse(template_id=template_id, usage_count=usage)
@router.post("/{template_id}/toggle-favorite", response_model=ToggleFavoriteResponse)
def toggle_favorite(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ToggleFavoriteResponse:
"""切换模板收藏状态(当前为兼容端点,始终返回 false)"""
user_id = authenticated_user.user.id
use_case = GetTemplateUseCase(template_repository)
try:
template = use_case.execute(template_id, user_id)
except Exception as _e:
logger.exception("toggle_favorite 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return ToggleFavoriteResponse(id=template_id, is_favorite=False)
# ── Validate template ──
@router.post("/{template_id}/validate", response_model=ValidateTemplateResponse)
def validate_template(
template_id: str,
request: ValidateTemplateRequest = ValidateTemplateRequest(),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ValidateTemplateResponse:
user_id = authenticated_user.user.id
command = ValidateTemplateCommand(
template_id=template_id,
user_id=user_id,
voiceover_duration=request.voiceover_duration,
)
use_case = ValidateTemplateUseCase(template_repository)
try:
result = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return ValidateTemplateResponse(
template=_to_response(result.template),
warnings=[GenerateWarningResponse(code=w.code, message=w.message, details=w.details) for w in result.warnings],
)
# ── Category CRUD ──
@router.get("/categories/list", response_model=ListCategoriesResponse)
def list_categories(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListCategoriesResponse:
user_id = authenticated_user.user.id
try:
use_case = ListCategoriesUseCase(template_repository)
categories = use_case.execute(user_id)
except Exception:
logger.exception("list_categories 查询失败: user_id=%s", user_id)
return ListCategoriesResponse(items=[])
return ListCategoriesResponse(
items=[CategoryResponse(id=c.id, user_id=c.user_id, name=c.name, created_at=c.created_at) for c in categories],
)
@router.post("/categories", response_model=CategoryResponse, status_code=status.HTTP_201_CREATED)
def create_category(
request: CreateCategoryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> CategoryResponse:
user_id = authenticated_user.user.id
command = CreateCategoryCommand(user_id=user_id, name=request.name)
use_case = CreateCategoryUseCase(template_repository)
category = use_case.execute(command)
return CategoryResponse(
id=category.id,
user_id=category.user_id,
name=category.name,
created_at=category.created_at,
)
@router.delete(
"/categories/{category_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response
)
def delete_category(
category_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteCategoryUseCase(template_repository)
deleted = use_case.execute(category_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Category not found")
return Response(status_code=204)
# ── Tags ──
@router.get("/tags/list", response_model=ListTagsResponse)
def list_tags(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTagsResponse:
"""获取用户所有模板标签(去重排序)"""
user_id = authenticated_user.user.id
try:
use_case = ListTagsUseCase(template_repository)
tags = use_case.execute(user_id)
except Exception:
logger.exception("list_tags 查询失败: user_id=%s", user_id)
return ListTagsResponse(items=[])
return ListTagsResponse(items=tags)
@@ -28,7 +28,7 @@ from .adjustments import router as adjustments_router
from .ai_features import router as ai_features_router
from .bgm import router as bgm_router
from .clips import router as clips_router
from .dependencies import get_draft_plan_id, get_editor_services, resolve_draft_plan_id # noqa: F401
from .dependencies import get_draft_plan_id, get_editor_services # noqa: F401
from .draft import router as draft_router
from .effects import router as effects_router
from .export import router as export_router
@@ -65,8 +65,8 @@ def _build_asset_analyses(
if url:
video_urls.append(url)
valid_asset_ids.append(aid)
except Exception:
logger.exception("获取素材URL失败: asset_id=%s", aid)
except Exception as e:
logger.warning("获取素材URL失败: asset_id=%s error=%s", aid, str(e))
if not video_urls:
logger.info("无可用视频素材,跳过视频理解分析")
@@ -108,7 +108,7 @@ def _build_asset_analyses(
return analyses
except Exception as e:
logger.exception("MediaKit 视频理解异常,将降级到无分析模式: %s", e)
logger.warning("MediaKit 视频理解异常,将降级到无分析模式: %s", str(e))
return {}
@@ -177,7 +177,7 @@ def editor_ai_recommend(
try:
db.rollback()
except Exception:
logger.exception("db rollback failed in ai_recommend")
pass
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="AI推荐结果保存失败,请稍后重试",
@@ -132,8 +132,8 @@ def _build_asset_url_map(
result: dict[str, str | None] = {}
try:
storage = get_storage_service()
except Exception as e:
logger.exception("获取存储服务失败,跳过asset_url生成: %s", e)
except Exception:
logger.warning("获取存储服务失败,跳过asset_url生成")
return {aid: None for aid in asset_ids}
# 批量查询所有 Asset(单次 SQL IN 查询,避免 N+1)
@@ -141,7 +141,7 @@ def _build_asset_url_map(
assets = asset_repo.find_by_ids(unique_ids)
asset_map = {a.id: a for a in assets}
except Exception:
logger.exception("批量查询素材失败: asset_ids=%s", asset_ids)
logger.warning("批量查询素材失败: asset_ids=%s", asset_ids, exc_info=True)
return {aid: None for aid in asset_ids if aid}
for aid in unique_ids:
@@ -156,7 +156,7 @@ def _build_asset_url_map(
continue
result[aid] = storage.get_download_url(storage_key, expires_seconds=3600)
except Exception:
logger.exception("生成素材签名URL失败: asset_id=%s", aid)
logger.warning("生成素材签名URL失败: asset_id=%s", aid, exc_info=True)
result[aid] = None
return result
@@ -486,8 +486,8 @@ def _get_mediakit_recommendations(
if url:
video_urls.append(url)
valid_asset_ids.append(asset_id)
except Exception:
logger.exception("获取素材URL失败: asset_id=%s", asset_id)
except Exception as e:
logger.warning("获取素材URL失败: asset_id=%s error=%s", asset_id, e)
if not video_urls:
return {}
@@ -563,7 +563,7 @@ def _get_mediakit_recommendations(
return recommendations
except Exception as e:
logger.exception("MediaKit 智能选片异常,降级为随机选择: %s", e)
logger.warning("MediaKit 智能选片异常,降级为随机选择: %s", e)
return {}
@@ -622,10 +622,8 @@ def create_clips_from_assets_editor(
"""从素材批量创建片段(按模板segment配置创建,MediaKit异步更新).
逻辑:
1. 从模板读取 segments,片段数量优先级:显式 clip_count1-10)→ 旧字段
required_clips_count(兼容,超10截断)→ 默认 3(产品默认 3 段)。
片段数大于模板 segment 数时按顺序循环复用 segment 配置。
2. 每个片段时长在对应 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
1. 从模板读取 segments,片段数量 = segment 数量(忽略前端传的 required_clips_count
2. 每个片段时长在 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
3. 素材按片段顺序轮询分配,素材不够时同一素材切多个片段
4. 使用 replace_all_clips_transactional 原子性地清空旧片段并创建新的(随机起始时间)
5. 立即返回响应(目标 <1秒)
@@ -650,22 +648,6 @@ def create_clips_from_assets_editor(
detail="模板未配置片段",
)
# 1.5 归一化片段数量:
# 优先级:显式 clip_count → 旧字段 required_clips_count(由 schema 归一化到 clip_count
# → 默认 3(产品默认 3 段)。按 N 循环复用 segment 配置;N <= len(segments) 时截取前 N 个
# (保持向后兼容:原模板有 N 个 segment、前端不传 clip_count 且 N<=10 时按模板段数创建;
# 默认模板仅有 1 个通用 segment 时按 clip_count=3 循环生成 3 段)。
requested_clip_count = getattr(body, "clip_count", None)
if requested_clip_count is None:
# schema 未显式传 clip_count 且无 legacy:使用模板 segments 数量,若超出 10 则截断
requested_clip_count = len(segments) if 1 <= len(segments) <= 10 else 3
requested_clip_count = max(1, min(int(requested_clip_count), 10))
effective_segments: list[tuple[int, float, float]] = []
for i in range(requested_clip_count):
src = segments[i % len(segments)]
effective_segments.append((i, float(src[1]), float(src[2])))
segments = effective_segments
# 防御:schema validator 已过滤 null/空串,这里再归一化一次,
# 避免异常入参(undefined → null)导致后续 /assets/{id} 404 / 422
asset_ids = [str(aid).strip() for aid in (body.asset_ids or []) if isinstance(aid, str) and aid.strip()]
@@ -878,7 +860,7 @@ def create_clips_from_assets_editor(
duplicate_warning = None
if dup_rate > 50:
duplicate_warning = f"查重率 {dup_rate:.1f}% 超过50%,建议更换素材或模板"
logger.exception(
logger.warning(
"from-assets 成片查重率超标: plan_id=%s dup_rate=%.1f%%",
plan_id,
dup_rate,
@@ -978,8 +960,8 @@ def _update_mediakit_recommendations_async( # pragma: no cover
# 尝试获取存储服务(用于生成视频 URL)
try:
storage = get_storage_service()
except Exception as e:
logger.exception("后台任务: 获取存储服务失败,跳过 SceneChange 更新: %s", e)
except Exception:
logger.warning("后台任务: 获取存储服务失败,跳过 SceneChange 更新")
return
# 获取 MediaKit 客户端
@@ -1005,8 +987,8 @@ def _update_mediakit_recommendations_async( # pragma: no cover
if storage_key and mime.startswith("video/"):
try:
video_url = storage.get_download_url(storage_key)
except Exception:
logger.exception("后台任务: 获取素材URL失败: asset_id=%s", asset_id)
except Exception as e:
logger.warning("后台任务: 获取素材URL失败: asset_id=%s error=%s", asset_id, e)
# 构建该素材的占用区间列表(排除已更新片段)
def _get_other_segments(asset_id_inner, clip_id_inner):
@@ -1057,11 +1039,12 @@ def _update_mediakit_recommendations_async( # pragma: no cover
asset_id,
len(scene_changes),
)
except Exception:
except Exception as cache_err:
# 缓存写入失败不影响本次片段更新
logger.exception(
"后台任务: 场景点缓存写入失败: asset_id=%s",
logger.warning(
"后台任务: 场景点缓存写入失败: asset_id=%s error=%s",
asset_id,
cache_err,
)
# SceneChange 未获得有效结果 → 尝试 analyze_videos 作为 fallback
@@ -1141,10 +1124,11 @@ def _update_mediakit_recommendations_async( # pragma: no cover
recommended_start + clip_duration,
plan_id,
)
except Exception:
logger.exception(
"后台任务: 同步素材区间记录失败,回滚本次片段更新: clip_id=%s",
except Exception as me:
logger.warning(
"后台任务: 同步素材区间记录失败,回滚本次片段更新: clip_id=%s error=%s",
clip.id,
me,
)
db.rollback()
continue
@@ -1160,8 +1144,8 @@ def _update_mediakit_recommendations_async( # pragma: no cover
asset_id,
recommended_start,
)
except Exception:
logger.exception("后台任务: 单个片段更新失败: clip_id=%s", clip.id)
except Exception as ue:
logger.warning("后台任务: 单个片段更新失败: clip_id=%s error=%s", clip.id, ue)
try:
db.rollback()
except Exception:
@@ -1170,9 +1154,9 @@ def _update_mediakit_recommendations_async( # pragma: no cover
logger.info("后台任务完成: plan_id=%s 成功更新 %d 个片段", plan_id, updated_count)
except Exception:
except Exception as e:
# 后台任务失败不影响已创建的片段,静默处理
logger.exception("后台任务异常: plan_id=%s", plan_id)
logger.warning("后台任务异常: plan_id=%s error=%s", plan_id, e, exc_info=True)
if db:
try:
db.rollback()
@@ -2,16 +2,13 @@
核心依赖:
- get_editor_services: 获取模板+计划服务
- get_draft_plan_id: Depends 形式的路径依赖(template_id 路径参数必填)
- resolve_draft_plan_id: 纯函数版本,供 clips_standalone 等非路径参数场景复用
(支持空 tid 时自动兜底创建默认模板)
- get_draft_plan_id: 根据 template_id 获取或创建草稿,返回 plan_id
"""
from __future__ import annotations
import logging
from app.api.routes._default_template import get_or_create_default_template_id
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.services.edit_plan_service import EditPlanService
@@ -33,61 +30,50 @@ def get_editor_services(
return EditTemplateService(db), EditPlanService(db)
def resolve_draft_plan_id(
def get_draft_plan_id(
template_id: str,
services: tuple[EditTemplateService, EditPlanService],
current_user: AuthenticatedUser,
db: Session,
auto_create_default: bool = True,
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> str:
"""根据 template_id 获取或创建草稿,返回 plan_id(纯函数,不带 Depends)。
"""路径依赖:根据 template_id 获取或创建草稿,返回 plan_id.
当 auto_create_default=True 且 template_id 为空时,自动调用
get_or_create_default_template_id 创建默认模板(用于 clips_standalone
等非路径参数场景)。
这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里,
确保 template_id → plan_id 的映射始终存在。
模板读取遵循单一数据源、显式判定(不使用异常降级):
- 用户自建模板在旧表 ``templates``(归属 user_idis_active=True);
- 全局模板在新表 ``edit_templates``(无 user_id,全局可读)。
模板不存在、已删除或不归属于当前用户时,一律返回 404。
"""
tpl_svc, plan_svc = services
user_id = str(current_user.user.id)
# 0. 空 tid 兜底
if not template_id:
if auto_create_default:
tid = get_or_create_default_template_id(db, user_id)
if not tid:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="无法自动创建默认模板,请刷新页面重试",
)
template_id = tid
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="template_id 不能为空",
)
# 1. 门禁:校验模板存在且可访问
# 0. 门禁:校验模板存在且可访问(即使草稿已缓存命中也要校验,
# 避免模板被删除/无权访问后仍可通过既有草稿 plan 继续操作)。
old_repo = SQLAlchemyTemplateRepository(db)
old_template = old_repo.get_active(template_id, user_id)
is_global_template = tpl_svc.get_template(template_id) is not None
if old_template is None and not is_global_template:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
# 2. 草稿已存在 → 直接返回
# 1. 草稿已存在 → 直接返回
draft = tpl_svc.get_template_draft(template_id)
if draft is not None:
return draft.id
# 3. 全局模板(新系统)→ 用新服务创建草稿
# 2. 全局模板(新系统)→ 用新服务创建草稿
if is_global_template:
draft = tpl_svc.create_template_draft(template_id, user_id=user_id)
return draft.id
# 4. 旧模板(templates 表)→ 基于旧模板创建草稿计划
# 3. 旧模板(templates 表)→ 基于旧模板创建草稿计划
from app.services.plan_generator_service import PlanGeneratorService
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
# 构造伪 EditTemplate 对象(只填 generate_from_template 需要的字段)
pseudo_template = EditTemplate(
id=old_template.id,
name=old_template.name,
@@ -95,6 +81,7 @@ def resolve_draft_plan_id(
status=EditTemplateStatus.ACTIVE,
)
# 将旧模板 segments 转换为 clip_configs
clip_configs: list[TemplateClipConfig] = []
for seg in old_template.segments or []:
clip_configs.append(
@@ -118,6 +105,7 @@ def resolve_draft_plan_id(
)
plan = result["plan"]
# 标记为模板草稿(后续可复用 tpl_svc.get_template_draft 的查找逻辑)
plan_svc.update_plan_config(plan.id, {"is_template_draft": True})
logger.info(
@@ -127,23 +115,3 @@ def resolve_draft_plan_id(
user_id,
)
return plan.id
def get_draft_plan_id(
template_id: str,
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> str:
"""路径依赖:根据 template_id 获取或创建草稿,返回 plan_id.
Depends 版本:路径参数 template_id 由 FastAPI 保证非空,不自动兜底。
兜底逻辑走 resolve_draft_plan_id(auto_create_default=False)。
"""
return resolve_draft_plan_id(
template_id=template_id,
services=services,
current_user=current_user,
db=db,
auto_create_default=False,
)
@@ -6,9 +6,9 @@
from __future__ import annotations
import re as _re
from typing import Any, Optional
from typing import Any, List, Optional
from pydantic import BaseModel, Field, model_validator, validator
from pydantic import BaseModel, Field, validator
_EXPORT_RESOLUTION_PATTERN = _re.compile(r"^\d+x\d+$")
_EXPORT_VALID_QUALITY_PRESETS = {"ultra_fast", "fast", "balanced", "high", "best"}
@@ -21,7 +21,7 @@ _EXPORT_VALID_FORMATS = {"mp4", "mov"}
class AIRecommendRequest(BaseModel):
"""AI 推荐片段方案请求体"""
asset_ids: list[str] = Field(default_factory=list, description="素材 ID 列表")
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表")
editing_mode: str = Field(default="one_take", description="剪辑模式: one_take / pip / voice_over / voice_pip")
target_duration: float = Field(default=30.0, ge=1.0, le=600.0, description="目标时长(秒)")
@@ -44,7 +44,7 @@ class AIRecommendResponse(BaseModel):
"""AI 推荐片段方案响应体"""
plan_id: str = Field(..., description="剪辑计划 ID")
clips: list[AIRecommendClipItem] = Field(..., description="推荐的片段列表")
clips: List[AIRecommendClipItem] = Field(..., description="推荐的片段列表")
config: dict[str, Any] = Field(..., description="推荐的 plan configcover/title/subtitle/bgm")
total_duration: float = Field(..., ge=0.0, description="推荐方案总时长(秒)")
confidence: float = Field(..., ge=0.0, le=1.0, description="AI 推荐置信度 (0~1)")
@@ -137,7 +137,7 @@ class ClipReorderItem(BaseModel):
class ClipReorderRequest(BaseModel):
"""片段重排序请求"""
items: list[ClipReorderItem] = Field(..., min_length=1, max_length=500, description="重排序条目列表")
items: List[ClipReorderItem] = Field(..., min_length=1, max_length=500, description="重排序条目列表")
class ClipReorderResponse(BaseModel):
@@ -151,7 +151,7 @@ class ClipReorderResponse(BaseModel):
class ClipBatchDeleteRequest(BaseModel):
"""批量删除片段请求"""
clip_ids: list[str] = Field(..., min_length=1, max_length=500, description="要删除的片段ID列表")
clip_ids: List[str] = Field(..., min_length=1, max_length=500, description="要删除的片段ID列表")
class ClipBatchDeleteResponse(BaseModel):
@@ -162,26 +162,13 @@ class ClipBatchDeleteResponse(BaseModel):
message: str = ""
# sentinel:区分「前端未传 clip_count」和「显式传 0/None」
_UNSET = object()
class ClipsFromAssetsRequest(BaseModel):
"""从素材批量创建片段请求"""
asset_ids: list[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
asset_ids: List[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
clip_type: str = Field(default="main", description="片段类型,默认 main")
clip_count: Optional[int] = Field(
default=None,
ge=1,
le=10,
description="片段数量(1-10);不传时使用旧字段 required_clips_count;两者都不传时回退为模板 segments 数量(默认 3 段)。",
)
required_clips_count: Optional[int] = Field(
default=None,
ge=1,
le=200,
description="[已废弃] 旧字段,请使用 clip_count;仅作向后兼容——clip_count 未显式传入时才回退本字段(超10截断到10)。",
default=None, ge=1, le=200, description="要求创建的片段数量;不传则等于素材数量"
)
@validator("asset_ids", pre=True)
@@ -193,23 +180,6 @@ class ClipsFromAssetsRequest(BaseModel):
return v
return [x for x in v if isinstance(x, str) and x.strip()]
@model_validator(mode="before")
@classmethod
def _backfill_clip_count(cls, data: Any) -> Any:
"""兼容旧字段 required_clips_count:仅当新字段 clip_count 未显式传入时才回退旧字段;
两者都没传时保持 clip_count=None,路由层按模板 segments 数量兜底。旧字段超 10 截断到 10。"""
if not isinstance(data, dict):
return data
has_new = "clip_count" in data and data["clip_count"] is not None
if not has_new:
legacy = data.get("required_clips_count")
if legacy is not None:
try:
data["clip_count"] = max(1, min(int(legacy), 10))
except (TypeError, ValueError):
pass
return data
class ClipsFromAssetsResponse(BaseModel):
"""从素材批量创建片段响应"""
@@ -218,7 +188,7 @@ class ClipsFromAssetsResponse(BaseModel):
created_count: int
plan_id: str = ""
message: str = ""
clip_ids: list[str] = Field(default_factory=list, description="创建的片段ID列表")
clip_ids: List[str] = Field(default_factory=list, description="创建的片段ID列表")
duplicate_warning: Optional[str] = Field(default=None, description="查重率超标警告")
exhaustion_warning: Optional[str] = Field(default=None, description="素材耗尽警告")
@@ -302,7 +272,7 @@ class ExportPresetItem(BaseModel):
class ExportPresetListResponse(BaseModel):
"""导出预设列表响应"""
items: list[ExportPresetItem]
items: List[ExportPresetItem]
total: int
@@ -316,7 +286,7 @@ class FilterPresetResponse(BaseModel):
name: str
category: str
description: str
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class FilterConfigResponse(BaseModel):
@@ -346,7 +316,7 @@ class FilterUpdateRequest(BaseModel):
class FilterPresetListResponse(BaseModel):
"""滤镜预设列表响应"""
items: list[FilterPresetResponse]
items: List[FilterPresetResponse]
total: int
@@ -360,7 +330,7 @@ class TransitionPresetResponse(BaseModel):
name: str
category: str
description: str
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
default_duration: float
min_duration: float
max_duration: float
@@ -402,7 +372,7 @@ class BatchTransitionResponse(BaseModel):
class TransitionPresetListResponse(BaseModel):
"""转场预设列表响应"""
items: list[TransitionPresetResponse]
items: List[TransitionPresetResponse]
total: int
@@ -458,7 +428,7 @@ class EditorClipResponse(BaseModel):
class EditorClipListResponse(BaseModel):
"""片段列表响应"""
items: list[EditorClipResponse]
items: List[EditorClipResponse]
total: int
@@ -496,7 +466,7 @@ class EditorClipBatchItem(BaseModel):
class EditorClipBatchUpdateRequest(BaseModel):
"""批量替换clips请求(全量覆盖)"""
clips: list[EditorClipBatchItem] = Field(default_factory=list)
clips: List[EditorClipBatchItem] = Field(default_factory=list)
class EditorClipBatchUpdateResponse(BaseModel):
@@ -584,4 +554,4 @@ class EditorTimelineResponse(BaseModel):
plan_id: str
total_duration: float
scenes: list[EditorTimelineSceneResponse]
scenes: List[EditorTimelineSceneResponse]
+179 -23
View File
@@ -1,35 +1,191 @@
"""Title library routes — DEPRECATED (#1894).
独立标题库已废弃。前端应直接调用 GET /api/v1/scripts 获取文案列表,
取每条文案的 `title` 字段作为标题候选。
所有 /api/v1/titles 端点统一返回 HTTP 410 Gone。
"""
"""Title library CRUD routes."""
from __future__ import annotations
from fastapi import APIRouter, Response, status
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
router = APIRouter()
_GONE_MESSAGE = (
"标题库 API 已废弃(#1894):独立标题库已合并进文案库,"
"请使用 GET /api/v1/scripts 获取文案列表并取 title 字段作为标题。"
)
def _get_title_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTitleLibraryRepository:
return SQLAlchemyTitleLibraryRepository(session)
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 _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.api_route("", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
def titles_root_gone(response: Response) -> dict:
return _gone(response)
@router.get("", response_model=ListTitleLibraryResponse)
def list_titles(
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:
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.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
def titles_subpath_gone(response: Response, path: str) -> dict:
return _gone(response)
@router.post("/pick", response_model=TitleLibraryItemResponse)
def pick_title(
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:
"""智能选择一个标题。
策略:优先使用次数少的,从最少的前5个中随机选一个,兼顾公平和多样性。
"""
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,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
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(
request: CreateTitleLibraryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
user_repository: UserRepository = Depends(get_user_repository),
) -> TitleLibraryItemResponse:
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,
request: UpdateTitleLibraryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
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,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> Response:
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
Regular → Executable
+14 -103
View File
@@ -4,14 +4,12 @@ from __future__ import annotations
import json
import logging
import math
import subprocess
import tempfile
from pathlib import Path
from typing import Any, Optional
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.celery_app import celery_app
from app.core.storage import get_storage_service
from app.dependencies import (
@@ -53,8 +51,6 @@ from packages.application.tts_job.use_cases import (
)
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.domain.voice_presets import list_voices
from packages.ports.asset_library_repository import AssetLibraryRepository
from packages.ports.asset_repository import AssetRepository
@@ -132,7 +128,6 @@ def _to_response(job, sign_url=None) -> TTSJobResponse:
def synthesize(
request: TTSSynthesizeRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
@@ -144,31 +139,6 @@ def synthesize(
"""
user_id = authenticated_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_voice"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
# 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
actual_voice_id = request.voice_id
@@ -203,15 +173,6 @@ def synthesize(
# job.voice_id 统一存解析后的 CosyVoice voice_id
actual_voice_id = resolved_profile.voice_id
# 语速/情绪等合成参数随 metadata 落库,workflow 提交 CosyVoice 时读取透传
synthesis_meta = {
"speed": request.speed,
"emotion": request.emotion or "",
"language": request.language or "zh-CN",
}
if request.metadata_:
synthesis_meta.update(request.metadata_)
use_case = CreateTTSJobUseCase(repository)
job = use_case.execute(
user_id=user_id,
@@ -219,7 +180,7 @@ def synthesize(
voice_id=actual_voice_id,
voice_model=request.voice_model,
voice_clone_profile_id=voice_clone_profile_id,
metadata=synthesis_meta,
metadata=request.metadata_,
)
# 提交 CosyVoice 合成任务
@@ -228,7 +189,6 @@ def synthesize(
cosyvoice_service=cosyvoice_service,
)
synthesis_error: Exception | None = None
try:
job = workflow.start_synthesis(job.id)
except Exception as e:
@@ -236,17 +196,10 @@ def synthesize(
# 但 DB 异常、网络异常等意外错误可能逃逸。
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
synthesis_error = e
try:
job = workflow.process_synthesis_failure(job.id, str(e))
except Exception as inner_e:
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}")
# 合成失败且已扣积分 → 退费
if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}")
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
if job.status.value == "processing":
@@ -261,17 +214,10 @@ def synthesize(
celery_app.send_task("worker.process_tts_synthesis", args=[job.id])
except Exception as e:
# Celery 调度失败,标记 job 为 failed
# e used below for refund context
try:
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
except Exception as inner_e:
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
# 调度失败退费
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}")
return TTSSynthesizeResponse(
job_id=job.id,
@@ -508,17 +454,10 @@ def save_tts_job_to_library(
try:
proc = subprocess.run(
[
"ffprobe",
"-v",
"quiet",
"-print_format",
"json",
"-show_format",
str(tmp_path),
"ffprobe", "-v", "quiet", "-print_format", "json",
"-show_format", str(tmp_path),
],
capture_output=True,
text=True,
timeout=10,
capture_output=True, text=True, timeout=10,
)
if proc.returncode == 0:
fmt = json.loads(proc.stdout).get("format", {})
@@ -598,7 +537,6 @@ def save_tts_job_to_library(
def preview_tts(
request: TTSPreviewRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
) -> TTSPreviewResponse:
@@ -607,31 +545,6 @@ def preview_tts(
用于前端预览配音效果,限制文本长度 200 字以内。
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
"""
user_id = authenticated_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_voice"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
actual_voice_id = request.voice_id
profile = voice_clone_repo.get(request.voice_id)
@@ -654,19 +567,17 @@ def preview_tts(
text=request.text,
voice_id=actual_voice_id,
speed=request.speed,
emotion=request.emotion,
language=getattr(request, "language", "zh-CN"),
)
except (CosyVoiceError, ValueError) as e:
# 合成失败退费
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
if isinstance(e, CosyVoiceError):
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
except CosyVoiceError as e:
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail=f"TTS 合成失败: {e}",
) from e
except ValueError as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=str(e),
) from e
return TTSPreviewResponse(
audio_url=result.audio_url,
+14 -110
View File
@@ -3,17 +3,14 @@
from __future__ import annotations
import logging
import math
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.celery_app import celery_app
from app.core.storage import get_storage_service
from app.dependencies import (
get_asset_repository,
get_cosyvoice_service,
get_db_session,
get_project_repository,
get_voice_clone_profile_repository,
)
@@ -25,7 +22,6 @@ from app.schemas.voice_clone import (
VoiceCloneStatusResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
@@ -42,11 +38,6 @@ from packages.application.voice_clone.use_cases import (
from packages.application.voice_clone.workflow import (
VoiceCloneWorkflowService,
)
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
# remove duplicate
_DUMMY_DELETED = ()
from packages.ports.asset_repository import AssetRepository
from packages.ports.project_repository import ProjectRepository
from packages.shared.storage import SharedStorageService
@@ -172,12 +163,12 @@ def create_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone {profile.id}")
except Exception as e:
logger.exception("Failed to dispatch Celery task")
logger.error(f"Failed to dispatch Celery task: {e}")
# P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing
try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception:
logger.exception("Failed to mark profile as failed after dispatch error")
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile)
@@ -286,111 +277,32 @@ def retry_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone retry {profile.id}")
except Exception as e:
logger.exception("Failed to dispatch Celery task")
logger.error(f"Failed to dispatch Celery task: {e}")
# P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing
try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception:
logger.exception("Failed to mark profile as failed after dispatch error")
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile)
_ALLOWED_PREVIEW_EMOTIONS = {
"",
# 7 种标准英文枚举(CosyVoice v3 官方值)
"neutral",
"happy",
"sad",
"angry",
"surprised",
"fearful",
"disgusted",
# 前端中文 7 标签
"中立",
"开心",
"难过",
"生气",
"惊讶",
"恐惧",
"厌恶",
# 旧英文 4 枚举 + 常见中文别名兼容
"natural",
"excited",
"calm",
"friendly",
"自然",
"愉快",
"高兴",
"快乐",
"兴奋",
"悲伤",
"愤怒",
"惊奇",
"吃惊",
"害怕",
"讨厌",
# 灵应 P1 指定别名
"中性",
"伤心",
"沉稳",
"亲切",
}
@router.get("/{clone_id}/preview", response_model=VoiceClonePreviewResponse)
def get_voice_clone_preview(
clone_id: str,
text: str = Query("", description="自定义试听文本,为空则使用默认示例"),
speed: float = Query(1.0, ge=0.5, le=2.0, description="语速,0.5-2.0,默认 1.0"),
emotion: str = Query(
"",
description="情绪:neutral/happy/sad/angry/surprised/fearful/disgusted,兼容旧值 natural/excited/calm/friendly,空为默认自然",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
cosyvoice: CosyVoiceService = Depends(get_cosyvoice_service),
) -> VoiceClonePreviewResponse:
"""获取克隆音色试听音频(实时 TTS 合成)。
- 克隆音色必须处于 ready 状态
- 使用默认试听文本时,结果缓存 7 天(仅默认 text+speed=1.0+emotion=空 组合缓存)
- 可传入自定义 text/speed/emotion 试听不同效果
- 使用默认试听文本时,结果缓存 7 天
- 可传入自定义 text 参数试听不同文本
"""
import time
user_id = authenticated_user.user.id
_points_deducted = 0
_points_scene = "voice_clone_synth"
_points_svc = PointsService() if settings.points_enabled else None
_preview_text_for_points = text.strip() or CLONE_PREVIEW_TEMPLATE
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(_preview_text_for_points) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
if emotion not in _ALLOWED_PREVIEW_EMOTIONS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"不支持的 emotion 值: {emotion},可选: neutral/happy/sad/angry/surprised/fearful/disgusted 或中文 中立/中性/开心/难过/伤心/生气/愤怒/惊讶/吃惊/恐惧/害怕/厌恶/讨厌 或留空",
)
use_case = GetVoiceCloneUseCase(repository)
try:
profile = use_case.execute(clone_id, authenticated_user.user.id)
@@ -403,8 +315,8 @@ def get_voice_clone_preview(
detail=f"Voice clone is not ready (current status: {profile.status})",
)
# 仅默认试听文本 + 默认 speed + 默认 emotion 时使用缓存
use_cache = (not text.strip()) and abs(speed - 1.0) < 1e-6 and (not emotion)
# 有自定义文本时不缓存
use_cache = not text.strip()
if use_cache and clone_id in _clone_preview_cache:
audio_url, duration, file_size, cached_text, cached_at = _clone_preview_cache[clone_id]
@@ -425,20 +337,12 @@ def get_voice_clone_preview(
text=preview_text,
voice_id=profile.voice_id,
format="mp3",
speed=speed,
emotion=emotion,
speed=1.0,
)
except (CosyVoiceError, ValueError) as e:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"克隆音色试听失败退积分异常: clone_id={clone_id}, err={refund_err}")
if isinstance(e, CosyVoiceError):
raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
except CosyVoiceError as e:
raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e
# 缓存(仅默认参数组合
# 缓存(仅默认试听文本
if use_cache:
_clone_preview_cache[clone_id] = (
result.audio_url,
+3 -4
View File
@@ -105,8 +105,8 @@ def _resolve_preset_preview_url(
_preset_preview_cache[voice_id] = (audio_url, time.time())
logger.info("Preset voice preview generated: %s", voice_id)
return audio_url
except Exception:
logger.exception("Failed to generate preset voice preview: voice_id=%s", voice_id)
except Exception as e:
logger.warning("Failed to generate preview for %s, using fallback: %s", voice_id, e)
return fallback_url
@@ -127,7 +127,6 @@ def _resolve_all_preset_preview_urls(
try:
result_map[p.voice_id] = _resolve_preset_preview_url(p.voice_id, p.preview_url, cosyvoice)
except Exception:
logger.exception("Failed to resolve preset preview URL: voice_id=%s", p.voice_id)
result_map[p.voice_id] = p.preview_url
return result_map
@@ -733,7 +732,7 @@ def _find_or_create_voice_library_for_extract(*, user_id, project_repository, as
try:
session.rollback()
except Exception:
logger.exception("session rollback failed in _find_or_create_voice_library")
pass
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:
View File
+1 -1
View File
@@ -6,7 +6,7 @@ ensuring proper lifecycle management and testability.
from __future__ import annotations
from collections.abc import Generator
from typing import Generator
import redis
from app.config import settings
+1 -1
View File
@@ -4,7 +4,7 @@
import logging
import time
from collections.abc import Callable
from typing import Callable
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
@@ -10,7 +10,7 @@ Exposes:
import re
import time
from collections.abc import Callable
from typing import Callable
from fastapi import Request, Response
from prometheus_client import (
+7 -22
View File
@@ -50,11 +50,9 @@ class CreateAiAvatarRenderRequest(BaseModel):
"""创建渲染任务请求."""
lipsync_job_id: str = Field(..., description="对口型任务 ID")
script_id: str = Field("", description="文案 ID(选自文案库时传;手动输入文案直生场景可留空)")
script_id: str = Field(..., description="文案 ID")
b_roll_segments: list[BRollSegment] = Field(default_factory=list, description="B-roll 片段列表")
title_config: dict[str, Any] = Field(
default_factory=dict, description="标题配置(可含 title_image_dataurl:前端 Canvas 渲染的标题 PNG dataURL"
)
title_config: dict[str, Any] = Field(default_factory=dict, description="标题配置")
cover_config: dict[str, Any] = Field(default_factory=dict, description="封面配置")
project_id: str = Field("", description="项目 ID")
@@ -69,7 +67,10 @@ class CreateAiAvatarRenderRequest(BaseModel):
@field_validator("script_id")
@classmethod
def validate_script_id(cls, v: str) -> str:
return (v or "").strip()
v = v.strip()
if not v:
raise ValueError("script_id 不能为空")
return v
class AiAvatarRenderJobResponse(BaseModel):
@@ -79,7 +80,7 @@ class AiAvatarRenderJobResponse(BaseModel):
user_id: str
project_id: str
lipsync_job_id: str
script_id: str = ""
script_id: str
b_roll_segments: list[dict[str, Any]]
title_config: dict[str, Any]
cover_config: dict[str, Any]
@@ -108,19 +109,3 @@ class AiAvatarRenderProgressResponse(BaseModel):
output_cover_url: str
output_duration: float
error_message: str
class SmartCoverResponse(BaseModel):
"""智能封面响应(封面从最终成片抽帧,不再叠加标题)."""
cover_url: str = Field("", description="封面图公网 URL(OSS,非临时);失败为空")
status: str = Field("completed", description="completed / fallback_failed")
message: str = Field("", description="失败原因(如有)")
class FinalizeRenderResponse(BaseModel):
"""封面选好后点「完成」,正式入库成片库的响应."""
video_id: str = Field(..., description="成片库视频ID")
cover_url: str = Field("", description="封面URL")
status: str = Field("success", description="success/already_finalized")
-38
View File
@@ -98,24 +98,6 @@ 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 音色 IDpreset 为 CosyVoice 音色 idclone 为克隆档案 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 配音严格守卫。
@@ -145,26 +127,6 @@ 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_idTTS 音色 ID")
return self
@model_validator(mode="after")
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
has_project = bool(self.project_id.strip())
-111
View File
@@ -1,111 +0,0 @@
"""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="人物视频预签名下载 URLGET")
audio_url: str = Field(..., description="驱动音频预签名下载 URLGET")
lipsync_job_id: str = ""
user_id: str = ""
project_id: str = ""
created_at: datetime
upload_url: str = Field(..., description="结果视频预签名上传 URLPUT, 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 = ""
+33 -94
View File
@@ -1,20 +1,11 @@
"""对口型 API Schema 定义 — #1796 / #1809 / #1822 / #1845(配音前置).
支持三种输入模式:
1. TTS 直生模式(兼容旧版前端):传 voice_id + script_text+ speed/emotion),
后端 Celery 异步做 TTS 合成 + MediaKit 提交。
2. 直接音频模式:传 video_url + audio_url(音频已由调用方准备好)。
3. 预合成音频模式(#1845 配音前置新主路径):前端先调 POST /lipsync/tts-preview
拿到 audio_url + sentence_timings,再在 create_job 时传 audio_url + audio_duration
+ sentence_timings,后端跳过 TTS 和时间戳计算,直接 ffprobe 校验后提交 MediaKit。
"""
"""对口型 API Schema 定义 — #1796, #1809 参数调整."""
from __future__ import annotations
from datetime import datetime
from typing import Optional
from pydantic import BaseModel, Field, model_validator
from pydantic import BaseModel, Field, field_validator
class LipsyncJobResponse(BaseModel):
@@ -26,17 +17,12 @@ class LipsyncJobResponse(BaseModel):
video_url: str
audio_url: str
enable_video_loop: bool
voice_id: str = ""
script_text: str = ""
speed: float = 1.0
emotion: str = ""
mediakit_task_id: str
status: str
output_video_url: str
output_duration: float
error_message: str
error_code: str
sentence_timings: Optional[list] = None
submitted_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
@@ -47,92 +33,45 @@ class LipsyncJobResponse(BaseModel):
class CreateLipsyncJobRequest(BaseModel):
"""创建对口型任务请求.
"""创建对口型任务请求 — #1809.
三种模式(三选一):
- TTS 直生(旧版/降级):voice_id + script_text 必填;audio_url 留空
- 直接音频:video_url + audio_url 必填。
- 预合成音频(#1845 新主路径):audio_url 必填 + 可选 audio_duration/sentence_timings
后端同步 ffprobe 校验时长、写入 timings,直接提交 MediaKit。
前端传 {voice_id, script_text, video_url}
后端内部调 TTS 生成 audio_url 再提交 MediaKit
"""
video_url: str = Field(..., description="人物视频 URL(MP4,≤30min,单人真人)")
# 模式 2/3:直接/预合成音频
audio_url: str = Field("", description="驱动音频 URLmp3/aac/wav/m4a/flac);直生模式留空")
audio_duration: Optional[float] = Field(None, ge=0, description="预合成音频时长(秒),可选;后端会 ffprobe 校验")
sentence_timings: Optional[list] = Field(None, description="预合成接口返回的句子时间戳,可选;若传入则直接写入 job")
# 模式 1TTS 直生
voice_id: str = Field("", description="音色 ID(预置音色或克隆音色 profile UUID")
script_text: str = Field("", description="要合成的文案(直生模式必填,最长 5000 字符)")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速(0.5-2.0),默认 1.0")
emotion: str = Field(
"",
description="情绪(英文枚举 neutral/happy/sad/angry/surprised/fearful/disgusted,或中文 中立/开心/难过/生气/惊讶/恐惧/厌恶;空为默认自然)",
)
enable_video_loop: bool = Field(
True, description="音频长于视频时是否循环画面(AI数字人默认开启,防止音频长于视频被截断)"
)
voice_id: str = Field(..., description="音色 ID(预设音色或克隆音色 profile ID)")
script_text: str = Field(..., description="要合成的脚本文本")
enable_video_loop: bool = Field(False, description="音频长于视频时是否循环画面")
project_id: str = Field("", description="项目 ID(可选)")
@model_validator(mode="after")
def _validate_input_mode(self) -> "CreateLipsyncJobRequest":
video = (self.video_url or "").strip()
if not video:
@field_validator("video_url")
@classmethod
def validate_video_url(cls, v: str) -> str:
v = v.strip()
if not v:
raise ValueError("video_url 不能为空")
if not video.startswith(("http://", "https://")):
if not v.startswith(("http://", "https://")):
raise ValueError("video_url 必须是 HTTP/HTTPS URL")
lower = video.lower().split("?")[0]
allowed_video_exts = (".mp4", ".mov", ".m4v", ".webm", ".avi", ".mkv", ".3gp")
if not any(lower.endswith(ext) for ext in allowed_video_exts):
raise ValueError("video_url 格式不支持,仅支持: " + ", ".join(allowed_video_exts))
lower = v.lower().split("?")[0]
if not lower.endswith(".mp4"):
raise ValueError("video_url 仅支持 MP4 格式")
return v
has_audio = bool((self.audio_url or "").strip())
has_tts = bool((self.voice_id or "").strip()) and bool((self.script_text or "").strip())
@field_validator("voice_id")
@classmethod
def validate_voice_id(cls, v: str) -> str:
v = v.strip()
if not v:
raise ValueError("voice_id 不能为空")
return v
if not has_audio and not has_tts:
raise ValueError(
"必须提供驱动音频:要么传 audio_url(直接/预合成音频模式),"
"要么同时传 voice_id + script_textTTS 直生模式)"
)
if has_tts and len(self.script_text) > 5000:
@field_validator("script_text")
@classmethod
def validate_script_text(cls, v: str) -> str:
v = v.strip()
if not v:
raise ValueError("script_text 不能为空")
if len(v) > 5000:
raise ValueError("script_text 最长 5000 字符")
if has_audio:
au = self.audio_url.strip()
if not au.startswith(("http://", "https://")):
raise ValueError("audio_url 必须是 HTTP/HTTPS URL")
au_lower = au.lower().split("?")[0]
allowed = (".mp3", ".aac", ".wav", ".m4a", ".flac")
if not any(au_lower.endswith(ext) for ext in allowed):
raise ValueError(f"audio_url 格式不支持,仅支持: {', '.join(allowed)}")
self.audio_url = au
return self
# ── #1845 TTS 预合成接口 ────────────────────────────────────────────────
class AiAvatarTtsPreviewRequest(BaseModel):
"""步骤1「生成配音」预合成请求(同步 HTTP,~2-3s)."""
voice_id: str = Field(..., min_length=1, max_length=128, description="音色 ID")
script_text: str = Field(..., min_length=1, max_length=5000, description="要合成的文案")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速(0.5-2.0),默认 1.0")
emotion: str = Field(
"neutral",
max_length=32,
description="情绪(英文枚举 neutral/happy/sad/angry/surprised/fearful/disgusted,或中文 中立/开心/难过/生气/惊讶/恐惧/厌恶;默认 neutral)",
)
class AiAvatarTtsPreviewResponse(BaseModel):
"""TTS 预合成响应(临时 URL,24h 内有效,足够当前会话使用)."""
audio_url: str = Field(..., description="CosyVoice 临时音频 URL")
duration: float = Field(..., ge=0, description="音频总时长(秒),ffprobe 测得")
sentence_timings: list[dict] = Field(..., description="句子级精确时间戳")
return v
-209
View File
@@ -1,209 +0,0 @@
"""积分 & 会员相关 Pydantic Schema (#1895)"""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from pydantic import BaseModel, Field
# ============ 余额 & 账户 ============
class PointsBalanceResponse(BaseModel):
"""积分余额 + 会员状态"""
balance: int = Field(..., description="当前积分余额")
total_earned: int = Field(..., description="累计获得积分")
total_spent: int = Field(..., description="累计消耗积分")
is_member: bool = Field(default=False, description="是否付费会员")
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
# ============ 流水 ============
class PointsTransactionItem(BaseModel):
"""单条积分流水"""
id: str
type: str = Field(..., description="类型: add/deduct")
source: str = Field(..., description="来源场景")
amount: int
balance_after: int
description: str = ""
ref_id: str = ""
created_at: Optional[str] = None
class PointsTransactionsResponse(BaseModel):
"""积分流水分页响应"""
items: list[PointsTransactionItem]
total: int
page: int
page_size: int
# ============ 规则 & 积分包 ============
class PointRuleItem(BaseModel):
"""单条积分规则"""
scene_key: str
name: str
base_points: int
unit: str
extra_per_30s: Optional[int] = None
description: str = Field(default="", description="规则中文说明,例如 AI 配音每分钟消耗 X 积分")
class PointsRulesResponse(BaseModel):
"""所有积分消耗规则"""
rules: list[PointRuleItem]
free_user_multiplier: float = Field(..., description="免费用户积分上浮系数")
class PointsPackageItem(BaseModel):
"""积分包信息"""
code: str
name: str
points: int
price_cents: int
unit_price: str = Field("", description="单价描述,如 ¥0.099/积分")
class PointsPackagesResponse(BaseModel):
"""可购买的积分包列表"""
packages: list[PointsPackageItem]
user_discount: Optional[float] = Field(None, description="当前用户折扣(会员)")
# ============ 消费前检查 ============
class PointsCheckRequest(BaseModel):
"""消费前余额检查请求"""
scene_key: str
duration_minutes: Optional[float] = None
quantity: Optional[int] = 1
class PointsCheckResponse(BaseModel):
"""消费前余额检查响应"""
allowed: bool
required_points: int
current_balance: int
remaining_after: int
is_free_quota: bool = False
# ============ 手动扣减 / 退还(内部接口) ============
class PointsDeductRequest(BaseModel):
"""积分扣减请求"""
scene_key: str
amount: int
description: Optional[str] = ""
ref_id: Optional[str] = ""
class PointsRefundRequest(BaseModel):
"""积分退还请求"""
transaction_id: str
reason: Optional[str] = ""
class PointsRechargeRequest(BaseModel):
"""积分充值请求"""
package_id: str = Field(..., description="积分包 code,如 starter_pack")
# ============ 订单 ============
class PointsOrderResponse(BaseModel):
"""订单信息"""
id: str
order_type: str
product_code: str
amount_cents: int
points_amount: int = Field(0, description="本次充值/购买可获得的积分(仅 points 类型订单有意义)")
status: str
pay_params: dict[str, Any] = Field(
default_factory=dict, description="拉起支付所需参数(payment_url/prepay_id 等),支付通道接入后填充"
)
expire_at: Optional[str] = Field(None, description="订单过期时间(ISO 8601),默认创建后 48 小时")
created_at: Optional[str] = None
# ============ 每日额度 ============
class DailyUsageResponse(BaseModel):
"""今日免费额度使用情况"""
free_clips_used: int
free_clips_limit: int
free_clips_remaining: int
reset_at: str
# ============ 会员状态(聚合) ============
class MembershipStatusResponse(BaseModel):
"""当前用户会员状态(聚合信息)"""
is_member: bool
member_type: Optional[str] = None
member_expires_at: Optional[datetime] = None
points_balance: int
max_resolution: str = Field(
default="1080p",
description="可用最高分辨率: 720p(free) / 1080p(paid)",
)
# ============ 订阅档位 ============
class MembershipPlanItem(BaseModel):
"""单个会员档位"""
plan_id: str = Field(..., description="档位标识: monthly/quarterly/yearly")
name: str = Field(..., description="档位名称,例如 月卡")
monthly_price_cents: int = Field(..., description="折算月价(分)")
price_cents: int = Field(..., description="该档位总价(分)")
duration_days: int = Field(..., description="时长(天)")
points_discount: float = Field(..., description="该档位积分折扣,如 0.9 表示 9 折")
features: dict[str, Any] = Field(default_factory=dict, description="档位权益(max_resolution 等)")
class MembershipPlansResponse(BaseModel):
"""所有会员档位列表"""
plans: list[MembershipPlanItem]
# ============ 通用响应 ============
class SimpleMessageResponse(BaseModel):
"""简单消息响应"""
success: bool
message: str
data: Optional[dict[str, Any]] = None
+7 -7
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -20,8 +20,8 @@ class ScriptResponse(BaseModel):
user_id: str
title: str
content: str
segments: list[ScriptSegment] = Field(default_factory=list)
tags: list[str] = Field(default_factory=list)
segments: List[ScriptSegment] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
created_at: datetime
updated_at: datetime
@@ -34,12 +34,12 @@ class ScriptListResponse(BaseModel):
class CreateScriptRequest(BaseModel):
title: str = Field(..., min_length=1, max_length=255)
content: str = ""
segments: list[ScriptSegment] = Field(default_factory=list)
tags: list[str] = Field(default_factory=list)
segments: List[ScriptSegment] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateScriptRequest(BaseModel):
title: Optional[str] = Field(None, min_length=1, max_length=255)
content: Optional[str] = None
segments: Optional[list[ScriptSegment]] = None
tags: Optional[list[str]] = None
segments: Optional[List[ScriptSegment]] = None
tags: Optional[List[str]] = None
-60
View File
@@ -1,60 +0,0 @@
"""Scripts AI 能力 Pydantic schemas — Issue #1893.
抖音文案提取、AI 改写、AI 标题生成的请求/响应模型。
"""
from __future__ import annotations
from typing import List, Optional
from pydantic import BaseModel, Field
# ── 抖音文案提取 ─────────────────────────────────────────────────────────────
class ExtractFromDouyinRequest(BaseModel):
"""从抖音视频提取文案请求."""
url: str = Field(..., description="抖音视频链接(短链或长链)")
class ExtractFromDouyinResponse(BaseModel):
"""从抖音视频提取文案响应."""
text: str = Field(..., description="ASR 识别出的文案文本")
duration_seconds: float = Field(..., description="视频时长(秒)")
source_url: str = Field(..., description="原始视频链接")
# ── AI 改写 ─────────────────────────────────────────────────────────────────
class AiRewriteRequest(BaseModel):
"""AI 文案改写请求."""
content: str = Field(..., description="原文内容")
style: Optional[str] = Field("口语化", description="改写风格,如 口语化/正式/活泼")
class AiRewriteResponse(BaseModel):
"""AI 文案改写响应."""
original: str = Field(..., description="原文")
rewritten: str = Field(..., description="改写后的文案")
style: str = Field(..., description="使用的改写风格")
# ── AI 标题生成 ──────────────────────────────────────────────────────────────
class AiGenerateTitlesRequest(BaseModel):
"""AI 标题生成请求."""
content: str = Field(..., description="文案内容")
count: int = Field(3, ge=1, le=5, description="生成标题数量(1-5,默认3")
class AiGenerateTitlesResponse(BaseModel):
"""AI 标题生成响应."""
titles: List[str] = Field(..., description="生成的标题列表")
+7 -14
View File
@@ -7,21 +7,15 @@ from typing import Optional
from pydantic import BaseModel, Field
# ============ Enums / Types ============
# 会员体系(#1951/#1955 实装):
# free — 免费用户
# monthly — 月卡
# quarterly — 季卡
# yearly — 年卡
# 已废弃档位:standard / pro / enterprise(保留常量名便于识别旧字段,但不在 API 中暴露)
class MembershipType(str):
"""会员类型(与 packages.domain.points_rules.MEMBERSHIP_PRICES 一致)"""
class PlanType(str):
"""套餐类型"""
FREE = "free"
MONTHLY = "monthly"
QUARTERLY = "quarterly"
YEARLY = "yearly"
STANDARD = "standard"
PRO = "pro"
ENTERPRISE = "enterprise"
class SubscriptionStatus(str):
@@ -46,7 +40,6 @@ class BillingCycle(str):
"""计费周期"""
MONTHLY = "monthly"
QUARTERLY = "quarterly"
YEARLY = "yearly"
@@ -102,8 +95,8 @@ class SimpleResponse(BaseModel):
class ChangePlanRequest(BaseModel):
"""升级/降级请求"""
target_plan_id: str = Field(..., description="目标会员类型: monthly/quarterly/yearly")
billing_cycle: str = Field(..., description="计费周期: monthly/quarterly/yearly")
target_plan_id: str = Field(..., description="目标套餐ID")
billing_cycle: str = Field(..., description="计费周期: monthly/yearly")
class ToggleAutoRenewRequest(BaseModel):
+84 -22
View File
@@ -1,14 +1,9 @@
"""Template API schemas(精简版:仅保留列表接口 + 默认模板自动兜底所需字段).
前端 PR#1911 删除 my-templates / editing-planner / templates 管理页后,
模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等端点全部下线,
对应 Request/Response 模型也一并清理。
"""
"""Template API schemas."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -42,12 +37,12 @@ class TemplateResponse(BaseModel):
name: str
mode: str
category: str = ""
tags: list[str] = Field(default_factory=list)
title_config: dict[str, Any] = Field(default_factory=dict)
subtitle_config: dict[str, Any] = Field(default_factory=dict)
bgm_config: dict[str, Any] = Field(default_factory=dict)
tags: List[str] = Field(default_factory=list)
title_config: Dict[str, Any] = Field(default_factory=dict)
subtitle_config: Dict[str, Any] = Field(default_factory=dict)
bgm_config: Dict[str, Any] = Field(default_factory=dict)
estimated_duration: float = 0.0
segments: list[SegmentResponse] = Field(default_factory=list)
segments: List[SegmentResponse] = Field(default_factory=list)
is_active: bool = True
is_favorite: bool = False
usage_count: int = 0
@@ -55,29 +50,96 @@ class TemplateResponse(BaseModel):
updated_at: datetime
class ToggleFavoriteResponse(BaseModel):
id: str
is_favorite: bool
class ListTemplatesResponse(BaseModel):
items: list[TemplateResponse]
items: List[TemplateResponse]
total: int = 0
# ── Template Request(保留给内部 _get_or_create_default_template_id 兜底创建默认模板使用)──
# ── Template Request ──
class CreateTemplateRequest(BaseModel):
name: str
mode: str
category: str = ""
tags: list[str] = Field(default_factory=list)
title_config: dict[str, Any] = Field(default_factory=dict)
subtitle_config: dict[str, Any] = Field(default_factory=dict)
bgm_config: dict[str, Any] = Field(default_factory=dict)
tags: List[str] = Field(default_factory=list)
title_config: Dict[str, Any] = Field(default_factory=dict)
subtitle_config: Dict[str, Any] = Field(default_factory=dict)
bgm_config: Dict[str, Any] = Field(default_factory=dict)
estimated_duration: float = 0.0
segments: list[SegmentRequest] = Field(default_factory=list)
segments: List[SegmentRequest] = Field(default_factory=list)
class UpdateTemplateRequest(BaseModel):
name: Optional[str] = None
mode: Optional[str] = None
category: Optional[str] = None
tags: Optional[List[str]] = None
title_config: Optional[Dict[str, Any]] = None
subtitle_config: Optional[Dict[str, Any]] = None
bgm_config: Optional[Dict[str, Any]] = None
estimated_duration: Optional[float] = None
segments: Optional[List[SegmentRequest]] = None
# ── Validate ──
class ValidateTemplateRequest(BaseModel):
voiceover_duration: Optional[float] = None # 配音实际时长(秒)
class GenerateWarningResponse(BaseModel):
"""兼容老 import(如校验逻辑内部复用);模板管理页已下线,可按需进一步清理。"""
code: str
message: str
details: dict[str, Any] = Field(default_factory=dict)
details: Dict[str, Any] = Field(default_factory=dict)
class ValidateTemplateResponse(BaseModel):
template: TemplateResponse
warnings: List[GenerateWarningResponse] = Field(default_factory=list)
# ── Category ──
class CategoryResponse(BaseModel):
id: str
user_id: str
name: str
created_at: datetime
class CreateCategoryRequest(BaseModel):
name: str
class ListCategoriesResponse(BaseModel):
items: List[CategoryResponse]
# ── Copy Template ──
class CopyTemplateRequest(BaseModel):
new_name: str
# ── Tags ──
class ListTagsResponse(BaseModel):
items: List[str]
# ── Usage Stats ──
class TemplateUsageResponse(BaseModel):
template_id: str
usage_count: int
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -15,7 +15,7 @@ class TitleLibraryItemResponse(BaseModel):
text: str
category: str = "default"
description: str = ""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
usage_count: int = 0
is_active: bool = True
created_at: datetime
@@ -32,7 +32,7 @@ class CreateTitleLibraryRequest(BaseModel):
text: str = Field(..., min_length=1, max_length=500)
category: str = "default"
description: str = ""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateTitleLibraryRequest(BaseModel):
@@ -40,4 +40,4 @@ class UpdateTitleLibraryRequest(BaseModel):
text: Optional[str] = Field(None, min_length=1, max_length=500)
category: Optional[str] = None
description: Optional[str] = None
tags: Optional[list[str]] = None
tags: Optional[List[str]] = None
+4 -10
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -16,14 +16,10 @@ class TTSSynthesizeRequest(BaseModel):
output_name: str = Field("", description="输出文件名")
language: str = Field("zh-CN", description="语言")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field(
"",
description="情绪(中文/英文:自然/兴奋/沉稳/亲切/开心/悲伤/愤怒/惊讶/恐惧/厌恶 等;通过 instruction 自然语言指令控制)",
)
voice_model: str = Field("", description="语音模型名称")
voice_clone_profile_id: str = Field("", description="关联的音色克隆档案 ID")
format: str = Field("mp3", description="输出格式(mp3/wav/pcm")
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
class Config:
populate_by_name = True
@@ -49,7 +45,7 @@ class TTSJobResponse(BaseModel):
error_message: str = ""
retry_count: int = 0
max_retries: int = 3
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
started_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
@@ -83,7 +79,7 @@ class TTSSynthesizeResponse(BaseModel):
class ListTTSJobResponse(BaseModel):
"""TTS 任务列表响应。"""
items: list[TTSJobResponse]
items: List[TTSJobResponse]
total: int
page: int
page_size: int
@@ -113,8 +109,6 @@ class TTSPreviewRequest(BaseModel):
text: str = Field(..., min_length=1, max_length=200, description="合成文本,限制 200 字")
voice_id: str = Field(..., min_length=1, description="音色 ID")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field("", description="情绪(中文/英文:自然/兴奋/沉稳/亲切/开心/悲伤/愤怒/惊讶/恐惧/厌恶 等)")
language: str = Field("zh-CN", description="语言(zh-CN/en-US 等)")
pitch: float = Field(1.0, ge=0.5, le=2.0, description="音调(预留,当前未使用)")
+2 -2
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -61,7 +61,7 @@ class ShareResponse(BaseModel):
class ShareListResponse(BaseModel):
"""分享列表响应."""
items: list[ShareResponse]
items: List[ShareResponse]
total: int = 0
skip: int = 0
limit: int = 20
+3 -3
View File
@@ -6,7 +6,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Literal, Optional
from typing import List, Literal, Optional
from pydantic import BaseModel, Field
@@ -56,7 +56,7 @@ class UnifiedVoiceItemResponse(BaseModel):
status: str = "completed"
"""状态"""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
"""标签列表"""
# 克隆音色特有字段
@@ -113,7 +113,7 @@ class PresetVoiceItemResponse(BaseModel):
preview_url: str = ""
"""预览音频 URL"""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
"""标签列表"""
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -19,7 +19,7 @@ class CreateVoiceCloneRequest(BaseModel):
language: str = Field("zh-CN", description="语言")
gender: str = Field("unknown", description="性别")
max_retries: int = Field(3, ge=1, le=10, description="最大重试次数")
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
class Config:
populate_by_name = True
@@ -41,7 +41,7 @@ class VoiceCloneProfileResponse(BaseModel):
error_message: str = ""
retry_count: int = 0
max_retries: int = 3
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
created_at: datetime
updated_at: datetime
@@ -62,7 +62,7 @@ class VoiceCloneStatusResponse(BaseModel):
class ListVoiceCloneResponse(BaseModel):
"""音色克隆列表响应。"""
items: list[VoiceCloneProfileResponse]
items: List[VoiceCloneProfileResponse]
total: int
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -21,7 +21,7 @@ class VoiceLibraryItemResponse(BaseModel):
file_size: int = 0
status: str = "completed"
project_id: Optional[str] = None
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
created_at: datetime
updated_at: datetime
@@ -42,7 +42,7 @@ class CreateVoiceLibraryRequest(BaseModel):
file_size: int = 0
status: str = "completed"
project_id: Optional[str] = None
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateVoiceLibraryRequest(BaseModel):
@@ -55,4 +55,4 @@ class UpdateVoiceLibraryRequest(BaseModel):
duration: Optional[float] = None
file_size: Optional[int] = None
status: Optional[str] = None
tags: Optional[list[str]] = None
tags: Optional[List[str]] = None
@@ -1,215 +0,0 @@
"""AI 数字人封面服务 — MediaKit 抽帧 + 质量评分选最佳帧 + 转存 OSS.
与 generation_cover.py 的智能选帧能力对齐(不再用 FFmpeg 简单截帧):
1. MediaKit extract_frames 抽取多帧(默认 5 帧,SpecifiedFrames 策略)
2. cover_frame_scorer.score_frames 按清晰度/亮度/色彩评分选最佳
3. 下载最佳帧并转存 OSS,返回公网封面 URL
设计原则:封面一律从最终成片(已叠加标题/B-roll)抽帧,帧本身已含标题,
本服务**不再叠加标题**。对口型阶段的裸视频封面入口已删除(废弃)。
降级:MediaKit 不可用或抽帧失败时返回空字符串,由调用方决定回退策略。
"""
from __future__ import annotations
import logging
import tempfile
import uuid
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
logger = logging.getLogger(__name__)
# MediaKit 抽帧轮询参数:poll_interval=2s × max_poll=30 → 最长 60s(与 mediakit_client 默认值/lipsync 轮询保持一致,防止合成视频下载+抽帧超时)
COVER_POLL_INTERVAL = 2.0
COVER_MAX_POLL_ATTEMPTS = 30
# 帧图片下载超时(秒)
FRAME_DOWNLOAD_TIMEOUT = 20
# 最佳帧下载超时(用于 persist)
BEST_FRAME_DOWNLOAD_TIMEOUT = 30
# 自家 OSS 私有桶 URL 重签有效期(供 MediaKit GPU worker 拉取)
MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
def _sign_video_url_for_mediakit(video_url: str) -> str:
"""如果 video_url 是自家 OSS 私有桶 URL,重新签名为长有效期预签名 URL。
MediaKit GPU worker 需要能公网访问 video_url,裸 public_url 在私有桶下会 403。
"""
if not video_url:
return video_url
try:
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
public_base = getattr(storage, "public_url", "")
if not isinstance(public_base, str) or not public_base:
return video_url
own_host = urlparse(public_base).netloc.lower()
url_host = urlparse(video_url).netloc.lower()
if own_host and url_host == own_host:
signed = storage.get_download_url(video_url, expires_seconds=MEDIAKIT_URL_TTL_SECONDS)
if signed:
logger.info("[数字人封面] video_url 已重签(自家 OSS 私有桶)")
return signed
except Exception:
logger.warning("[数字人封面] video_url 重签失败,使用原始 URL", exc_info=True)
return video_url
def select_best_cover_frame(video_url: str, *, max_frames: int = 5) -> str:
"""从视频抽取多帧并评分选最佳帧,返回最佳帧的临时 URL."""
if not video_url:
return ""
video_url = _sign_video_url_for_mediakit(video_url)
try:
from packages.shared.cover_frame_scorer import score_frames
from packages.shared.mediakit_client import get_mediakit_client
mk = get_mediakit_client()
if not mk.is_available:
logger.warning("[数字人封面] MediaKit 未配置,无法智能抽帧")
return ""
logger.info(
"[数字人封面] 开始抽帧: video_url=%s max_frames=%d",
video_url[:80],
max_frames,
)
snapshots = mk.extract_frames(
video_url=video_url,
strategy="SpecifiedFrames",
max_frames=max_frames,
poll_interval=COVER_POLL_INTERVAL,
max_poll_attempts=COVER_MAX_POLL_ATTEMPTS,
max_retries=1,
)
if not snapshots:
logger.warning("[数字人封面] MediaKit 未返回帧: %s", video_url[:80])
return ""
if len(snapshots) == 1:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
import httpx
candidates = []
with httpx.Client(timeout=FRAME_DOWNLOAD_TIMEOUT, follow_redirects=True) as client:
for snap in snapshots:
url = snap.get("image_url") or snap.get("url") or ""
if not url:
continue
tmp_path: Optional[str] = None
try:
resp = client.get(url)
resp.raise_for_status()
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
candidates.append({"image_path": tmp_path, "url": url})
except Exception as e:
logger.warning("[数字人封面] 帧下载失败,跳过: url=%s err=%s", url[:80], e)
candidates.append({"image_path": None, "url": url, "score": 0.0})
if not candidates:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
scored = score_frames(candidates)
best = scored[0] if scored else None
best_url = best.get("url", "") if best else ""
for c in candidates:
p = c.get("image_path")
if p:
try:
Path(p).unlink(missing_ok=True)
except Exception:
pass
logger.info(
"[数字人封面] 智能选帧完成: candidates=%d best_score=%s",
len(candidates),
best.get("score") if best else "n/a",
)
return best_url
except Exception:
logger.warning("[数字人封面] 智能选帧失败", exc_info=True)
return ""
def persist_cover_to_oss(
frame_url: str,
*,
job_id: str = "",
prefix: str = "ai-avatar/covers",
) -> str:
"""下载最佳帧图并转存到 OSS,返回公网封面 URL(预签名).
封面来自最终成片抽帧,帧本身已含标题,本函数不再做任何文字/图片叠加。
"""
if not frame_url:
return ""
tmp_path: Optional[str] = None
try:
import httpx
with httpx.Client(timeout=BEST_FRAME_DOWNLOAD_TIMEOUT, follow_redirects=True) as client:
resp = client.get(frame_url)
resp.raise_for_status()
if not resp.content:
logger.warning("[数字人封面] 帧图内容为空: %s", frame_url[:80])
return frame_url
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
token = job_id or uuid.uuid4().hex[:12]
cover_key = f"{prefix}/{token}/cover_{uuid.uuid4().hex[:8]}.jpg"
public_url = storage.upload_file(
file_or_path=tmp_path,
storage_key=cover_key,
content_type="image/jpeg",
)
logger.info("[数字人封面] 封面已转存 OSS: key=%s", cover_key)
if public_url:
signed = storage.get_download_url(cover_key, expires_seconds=86400)
return signed
return frame_url
except Exception:
logger.warning("[数字人封面] 封面转存 OSS 失败,返回原始 URL", exc_info=True)
return frame_url
finally:
if tmp_path:
try:
Path(tmp_path).unlink(missing_ok=True)
except Exception:
pass
def generate_smart_cover(
video_url: str,
*,
job_id: str = "",
max_frames: int = 5,
) -> str:
"""一站式:MediaKit 智能抽帧选最佳 → 转存 OSS。失败返回空字符串。
封面从最终成片抽帧,不再叠加任何标题(帧本身已含)。
"""
best_frame = select_best_cover_frame(video_url, max_frames=max_frames)
if not best_frame:
return ""
return persist_cover_to_oss(best_frame, job_id=job_id)
+76 -357
View File
@@ -9,14 +9,11 @@
from __future__ import annotations
import base64
import binascii
import logging
import os
import subprocess
import tempfile
import uuid
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Any, Optional
from sqlalchemy.orm import Session
@@ -27,9 +24,8 @@ from packages.adapters.sqlalchemy_impl.models import (
ScriptModel,
)
from packages.domain.video_filter_builder import (
build_broll_overlay_filter,
build_cover_extract_command,
build_title_drawtext_filter,
build_title_overlay_filter,
)
from packages.shared.storage import get_shared_storage_service
@@ -57,8 +53,8 @@ class AiAvatarRenderService:
*,
user_id: str,
lipsync_job_id: str,
script_id: str = "",
b_roll_segments: list[dict[str, Any]] | None = None,
script_id: str,
b_roll_segments: list[dict[str, Any]],
title_config: dict[str, Any],
cover_config: dict[str, Any],
project_id: str = "",
@@ -87,19 +83,17 @@ class AiAvatarRenderService:
if not lipsync_job.output_video_url:
raise AiAvatarRenderError("对口型任务输出视频 URL 为空", code="LipsyncJobNoOutput")
# 2. 验证文案归属(仅当选了文案库条目时;手动输入文案直生场景 script_id 可空)
script_id = (script_id or "").strip()
if script_id:
script = (
self.db.query(ScriptModel)
.filter(
ScriptModel.id == script_id,
ScriptModel.user_id == user_id,
)
.first()
# 2. 验证文案归属
script = (
self.db.query(ScriptModel)
.filter(
ScriptModel.id == script_id,
ScriptModel.user_id == user_id,
)
if script is None:
raise AiAvatarRenderError("文案不存在或无权访问", code="ScriptNotFound")
.first()
)
if script is None:
raise AiAvatarRenderError("文案不存在或无权访问", code="ScriptNotFound")
# 3. 创建渲染任务
job_id = str(uuid.uuid4())
@@ -109,7 +103,7 @@ class AiAvatarRenderService:
project_id=project_id,
lipsync_job_id=lipsync_job_id,
script_id=script_id,
b_roll_segments=[s if isinstance(s, dict) else s.model_dump() for s in (b_roll_segments or [])],
b_roll_segments=[s if isinstance(s, dict) else s.model_dump() for s in b_roll_segments],
title_config=title_config,
cover_config=cover_config,
status="pending",
@@ -117,7 +111,7 @@ class AiAvatarRenderService:
self.db.add(job)
self.db.flush()
job.submitted_at = datetime.now(UTC)
job.submitted_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
@@ -164,7 +158,7 @@ class AiAvatarRenderService:
return None
if job.status in ("pending", "submitted"):
job.status = "cancelled"
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
@@ -186,7 +180,7 @@ class AiAvatarRenderService:
job.output_duration = 0.0
job.started_at = None
job.completed_at = None
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
@@ -200,8 +194,9 @@ class AiAvatarRenderService:
1. 下载对口型输出视频 (20%)
2. 构建 FFmpeg 滤镜链 (40%)
3. 执行 FFmpeg 渲染 (80%)
4. 上传到 OSS (95%) — 封面不再自动生成,改由前端主动抽帧
5. 更新任务状态 (100%)
4. 提取封面 (90%)
5. 上传到 OSS (95%)
6. 更新任务状态 (100%)
"""
job = self.db.query(AiAvatarRenderJob).filter(AiAvatarRenderJob.id == job_id).first()
if job is None:
@@ -215,9 +210,9 @@ class AiAvatarRenderService:
try:
# 更新状态为 processing
job.status = "processing"
job.started_at = datetime.now(UTC)
job.started_at = datetime.now(timezone.utc)
job.progress = 5
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
# 获取对口型任务信息
@@ -231,32 +226,27 @@ class AiAvatarRenderService:
self.db.commit()
# 2. 构建 FFmpeg 滤镜链 (40%)
# 用 ffprobe 探测输入视频分辨率,确保 B-roll 缩放与标题位置与实际输出一致。
# AI 数字人对口型输出为 9:16 竖屏,默认兜底 720x1280;探测失败时使用默认值不阻断渲染。
output_width, output_height = self._probe_video_resolution(input_video_path)
if output_width <= 0 or output_height <= 0:
output_width, output_height = 720, 1280
logger.info(
"[数字人渲染] ffprobe 探测分辨率失败或无效,使用默认竖屏尺寸 %sx%s",
output_width,
output_height,
)
else:
logger.info("[数字人渲染] 探测输入视频分辨率: %sx%s", output_width, output_height)
from packages.domain.video_filter_builder import build_broll_overlay_filter
broll_filter, broll_label = build_broll_overlay_filter(
filter_complex = build_broll_overlay_filter(
b_roll_segments=job.b_roll_segments,
video_duration=lipsync_job.output_duration,
output_width=output_width,
output_height=output_height,
)
# 标题叠加路径:优先前端 Canvas 渲染的 PNG 图层(所见即所得),
# 无 title_image_dataurl 时降级到 drawtext 重画文字。
title_cfg = job.title_config if isinstance(job.title_config, dict) else {}
title_dataurl = (title_cfg or {}).get("title_image_dataurl") if title_cfg else None
use_title_png = isinstance(title_dataurl, str) and title_dataurl.startswith("data:image/")
title_input_index = 1 + len(job.b_roll_segments or []) if use_title_png else None
# 标题叠加
title_filter = build_title_drawtext_filter(job.title_config)
if title_filter:
if filter_complex:
filter_complex += f"[vout]{title_filter}[vout_titled];"
else:
filter_complex = f"[0:v]{title_filter}[vout_titled];"
# 清理末尾分号
if filter_complex.endswith(";"):
filter_complex = filter_complex[:-1]
# 最终输出标签
final_label = "vout_titled" if title_filter else ("vout" if filter_complex else None)
job.progress = 40
self.db.commit()
@@ -265,126 +255,42 @@ class AiAvatarRenderService:
with tempfile.TemporaryDirectory() as tmpdir:
output_video_path = os.path.join(tmpdir, "output.mp4")
# 在临时目录里解码保存标题 PNG(with 退出自动清理)
title_png_path: Optional[str] = None
extra_inputs: list[str] = []
title_filter = None
if use_title_png:
try:
title_png_path = os.path.join(tmpdir, f"title_{job.id}.png")
self._save_title_dataurl_to_file(title_dataurl, dst_path=title_png_path)
extra_inputs.append(title_png_path)
logger.info(
"[数字人渲染] 标题 PNG 已保存: %s (input index %d)", title_png_path, title_input_index
)
except Exception as exc:
logger.warning("[数字人渲染] 标题 PNG 解码/保存失败,降级 drawtext: %s", exc)
title_png_path = None
extra_inputs = []
# 构建标题滤镜
final_label = None
if title_png_path and title_input_index is not None:
title_input_label = f"[{title_input_index}:v]"
base_label = f"[{broll_label}]" if broll_label else "[0:v]"
title_filter = build_title_overlay_filter(
title_cfg,
output_width=output_width,
output_height=output_height,
title_png_path=title_png_path,
title_input_label=title_input_label,
base_label=base_label,
output_label="vout_titled",
)
if not title_filter:
# build 返回 None → 文件不存在(极端并发情况),降级 drawtext
title_png_path = None
extra_inputs = []
if title_png_path:
# overlay 路径
if broll_filter and title_filter:
filter_complex = broll_filter + f";{title_filter}"
elif broll_filter:
filter_complex = broll_filter
final_label = broll_label
elif title_filter:
filter_complex = title_filter
else:
filter_complex = ""
if title_filter:
final_label = "vout_titled"
elif not final_label:
final_label = None
else:
# 降级:drawtext 重画文字
title_filter = build_title_drawtext_filter(
title_cfg,
output_width=output_width,
output_height=output_height,
)
if broll_filter and title_filter:
filter_complex = broll_filter + f";[{broll_label}]{title_filter}[vout_titled]"
final_label = "vout_titled"
elif broll_filter:
filter_complex = broll_filter
final_label = broll_label
elif title_filter:
filter_complex = f"[0:v]{title_filter}[vout_titled]"
final_label = "vout_titled"
else:
filter_complex = ""
final_label = None
cmd_list = self._build_ffmpeg_command(
cmd = self._build_ffmpeg_command(
input_video=input_video_path,
b_roll_segments=job.b_roll_segments,
extra_inputs=extra_inputs,
filter_complex=filter_complex,
final_label=final_label,
output_path=output_video_path,
)
try:
render_result = subprocess.run(
cmd_list,
capture_output=True,
text=True,
timeout=600,
)
except subprocess.TimeoutExpired as exc:
raise AiAvatarRenderError(
"FFmpeg 渲染超时(600s",
code="FFmpegTimeout",
) from exc
if render_result.returncode != 0:
stderr_tail = (render_result.stderr or "").strip()[-800:]
raise AiAvatarRenderError(
f"FFmpeg 渲染失败,退出码: {render_result.returncode}, stderr: {stderr_tail}",
code="FFmpegFailed",
)
exit_code = os.system(cmd)
if exit_code != 0:
raise AiAvatarRenderError(f"FFmpeg 渲染失败,退出码: {exit_code}", code="FFmpegFailed")
job.progress = 80
self.db.commit()
# 4/5. 上传成片到 OSS (95%) —— 已砍掉自动抽封面逻辑(步骤⑤);
# 封面由前端在渲染完成后通过 /smart-cover 接口主动从成片抽帧,不阻塞渲染链路。
# 4. 提取封面 (90%)
cover_path = ""
if job.cover_config:
cover_path = os.path.join(tmpdir, "cover.jpg")
cover_cmd = build_cover_extract_command(job.cover_config, cover_path)
cover_cmd = cover_cmd.replace("INPUT_VIDEO", output_video_path)
cover_exit = os.system(cover_cmd)
if cover_exit != 0:
logger.warning("封面提取失败,跳过: %s", cover_cmd)
cover_path = ""
job.progress = 90
self.db.commit()
# 5. 上传到 OSS (95%)
output_video_url = self._upload_to_oss(output_video_path, f"ai-avatar/{job_id}/output.mp4")
job.output_video_url = output_video_url
# 封面透传:如果用户已在 cover_config 中选定封面 URLmode=upload 的自定义上传 或
# mode=auto_frame 已有的智能封面结果),直接透传到 output_cover_url,不再重新截帧。
if isinstance(job.cover_config, dict):
_pre_cover_url = (
job.cover_config.get("url")
or job.cover_config.get("imageUrl")
or job.cover_config.get("cover_url")
or ""
)
if _pre_cover_url:
job.output_cover_url = _pre_cover_url
logger.info("[数字人渲染] 使用用户已选定封面 URL: job_id=%s", job_id)
if cover_path:
output_cover_url = self._upload_to_oss(cover_path, f"ai-avatar/{job_id}/cover.jpg")
job.output_cover_url = output_cover_url
# 获取输出视频时长
job.output_duration = lipsync_job.output_duration
@@ -394,110 +300,23 @@ class AiAvatarRenderService:
# 6. 完成
job.status = "completed"
job.progress = 100
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
job.completed_at = datetime.now(timezone.utc)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.info("渲染任务完成: %s", job_id)
# 7. 渲染完成,停留在「待选封面」状态:不自动入库。
# 用户在前端选好封面、点「完成」后,由 /{job_id}/finalize 接口显式入库。
logger.info("渲染任务完成,等待用户选择封面后入库: job_id=%s", job_id)
except AiAvatarRenderError as exc:
job.status = "failed"
job.error_message = str(exc)
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.error("渲染任务失败 [%s]: %s", job_id, exc)
raise
except Exception as exc:
job.status = "failed"
job.error_message = f"渲染异常: {str(exc)}"
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.exception("渲染任务异常 [%s]", job_id)
raise
def _persist_to_library(self, job: AiAvatarRenderJob, cover_url: Optional[str] = None):
"""将渲染结果写入成片库,返回 GeneratedVideo 领域对象.
Args:
job: 渲染任务(必须 status=completed 且 output_video_url 非空)
cover_url: 可选的封面 URL 覆盖(finalize 时传入即优先使用,否则取 job.output_cover_url
"""
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.domain.generated_video import GeneratedVideo
clip_name = f"AI数字人_{job.id[:8]}"
# AI数字人入口是独立页面,前端可能不传 project_id(无项目概念),
# 兜底为 "ai_avatar" 避免 DB 非空约束/查询问题;generation_task_id 用 render_job_id 便于反查。
clip_project_id = (job.project_id or "").strip() or "ai_avatar"
clip_generation_task_id = job.id
effective_cover = (cover_url or "").strip() if cover_url else (job.output_cover_url or "").strip()
clip = GeneratedVideo.create(
project_id=clip_project_id,
generation_task_id=clip_generation_task_id,
name=clip_name,
file_url=job.output_video_url,
user_id=job.user_id,
duration=job.output_duration or 0.0,
thumbnail_url=effective_cover or None,
generation_params={
"source": "ai_avatar_render",
"render_job_id": job.id,
},
)
video_repo = SQLAlchemyGeneratedVideoRepository(self.db)
video_repo.create(clip)
logger.info("[数字人渲染] 成片已入库: clip_id=%s render_job=%s", clip.id, job.id)
return clip
def finalize_job(self, job_id: str, user_id: str, cover_url: Optional[str] = None):
"""用户在前端点「完成」后调用:将已 completed 的渲染任务正式入库到成片库.
- 必须 status=completed 才可调用
- cover_url 若传入则优先使用并回写 job.output_cover_url;否则使用 job.output_cover_urlsmart-cover/custom-cover 已写入)
- 幂等:已入库则返回已存在的 GeneratedVideo
"""
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
job = self.get_render_job(job_id, user_id)
if job is None:
raise AiAvatarRenderError("渲染任务不存在", code="RenderJobNotFound")
if job.status != "completed":
raise AiAvatarRenderError(f"渲染任务未完成(当前状态: {job.status}),无法入库", code="RenderNotCompleted")
if not (job.output_video_url or "").strip():
raise AiAvatarRenderError("渲染成片视频 URL 为空,无法入库", code="OutputVideoMissing")
# 幂等检查:已入库直接返回现有记录(通过 generation_task_id=job_id 识别,
# 因为入库时 generation_task_id 被设置为 render_job_id 自身)
existing = (
self.db.query(GeneratedVideoModel)
.filter(
GeneratedVideoModel.user_id == user_id,
GeneratedVideoModel.generation_task_id == job_id,
)
.first()
)
if existing is not None:
logger.info("[数字人渲染] finalize 幂等命中,返回已存在记录: clip_id=%s job_id=%s", existing.id, job_id)
return SQLAlchemyGeneratedVideoRepository(self.db).get(existing.id)
# 传入 cover_url 时回写到 job
if cover_url and cover_url.strip():
job.output_cover_url = cover_url.strip()
# 同步更新 cover_config,保持 smart-cover 路径一致
if isinstance(job.cover_config, dict):
job.cover_config = {**job.cover_config, "mode": "auto_frame", "url": cover_url.strip()}
job.updated_at = datetime.now(UTC)
self.db.commit()
return self._persist_to_library(job, cover_url=cover_url)
def _download_video(self, url: str) -> str:
"""下载视频到临时文件."""
@@ -515,132 +334,32 @@ class AiAvatarRenderService:
os.unlink(tmp.name)
raise
@staticmethod
def _save_title_dataurl_to_file(dataurl: str, *, dst_path: str | None = None, job_id: str = "") -> str:
"""解码前端传来的 data:image/png;base64,... 并保存为本地 PNG 文件。
Args:
dataurl: 完整 dataURL 字符串
dst_path: 指定输出路径;为 None 时创建临时文件并返回路径
job_id: 仅在 dst_path 为空时用于临时文件命名
Returns:
保存后的本地文件路径
"""
if not isinstance(dataurl, str) or not dataurl.startswith("data:image/"):
raise ValueError("title_image_dataurl 不是合法的 data:image URL")
# 拆分 data:image/png;base64,<payload>
try:
header, b64 = dataurl.split(",", 1)
except ValueError as exc:
raise ValueError("title_image_dataurl 缺少 base64 payload") from exc
if "base64" not in header:
raise ValueError("title_image_dataurl 不是 base64 编码")
try:
png_bytes = base64.b64decode(b64, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"title_image_dataurl base64 解码失败: {exc}") from exc
if not png_bytes:
raise ValueError("title_image_dataurl 解码后为空")
if dst_path:
out_path = dst_path
with open(out_path, "wb") as f:
f.write(png_bytes)
return out_path
suffix = f"_title_{job_id}.png" if job_id else "_title.png"
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
tmp.write(png_bytes)
return tmp.name
@staticmethod
def _probe_video_resolution(video_path: str) -> tuple[int, int]:
"""用 ffprobe 探测视频分辨率,返回 (width, height);失败返回 (0, 0)。"""
try:
result = subprocess.run(
[
"ffprobe",
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height",
"-of",
"csv=p=0:s=x",
video_path,
],
capture_output=True,
text=True,
timeout=15,
)
if result.returncode == 0 and result.stdout.strip():
parts = result.stdout.strip().split("x")
if len(parts) == 2:
w, h = int(parts[0]), int(parts[1])
if w > 0 and h > 0:
return w, h
except Exception as exc:
logger.warning("[数字人渲染] ffprobe 探测分辨率失败: %s", exc)
return 0, 0
def _build_ffmpeg_command(
self,
*,
input_video: str,
b_roll_segments: list[dict[str, Any]],
extra_inputs: list[str] | None = None,
filter_complex: str,
final_label: Optional[str],
output_path: str,
) -> list[str]:
"""构建 FFmpeg 命令list 形式,shell=False.
根因修复 #1798 P0OSS 预签名 URL 含 `&Expires=...&Signature=...` 特殊字符,
os.system(shell=True) 会把 `&` 解释为后台命令分隔符,导致 -filter_complex 被
当成独立命令报 sh: -filter_complex: not foundexit 127 → Python 32512)。
list + shell=False 彻底规避 shell 转义问题。
"""
cmd: list[str] = ["ffmpeg", "-i", input_video]
) -> str:
"""构建 FFmpeg 命令."""
# 输入文件
inputs = f"-i {input_video}"
for seg in b_roll_segments:
asset_url = seg.get("asset_url", "")
if asset_url:
cmd.extend(["-i", asset_url])
# 额外输入(例如前端 Canvas 渲染的标题 PNG)
for extra in extra_inputs or []:
cmd.extend(["-i", extra])
inputs += f" -i {asset_url}"
# 滤镜
if filter_complex and final_label:
cmd.extend(
[
"-filter_complex",
filter_complex,
"-map",
f"[{final_label}]",
"-map",
"0:a?",
]
)
filter_arg = f'-filter_complex "{filter_complex}" -map "[{final_label}]"'
elif filter_complex:
cmd.extend(["-filter_complex", filter_complex])
filter_arg = f'-filter_complex "{filter_complex}"'
else:
filter_arg = ""
cmd.extend(
[
"-c:v",
"libx264",
"-preset",
"veryfast",
"-crf",
"23",
"-c:a",
"aac",
"-b:a",
"128k",
"-y",
output_path,
]
)
return cmd
return f"ffmpeg {inputs} {filter_arg} -c:v libx264 -preset fast -crf 23 -y {output_path}"
def _upload_to_oss(self, local_path: str, oss_key: str) -> str:
"""上传文件到 OSS,返回 URL.
+12 -12
View File
@@ -13,7 +13,7 @@
from __future__ import annotations
import logging
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from packages.domain.ai_parsing import generate_titles_fallback as _generate_titles_fallback_base
from packages.domain.ai_parsing import keyword_match_fallback as _semantic_match_fallback_base
@@ -64,7 +64,7 @@ def _generate_titles_fallback(
description: str,
style: str = "viral",
count: int = 5,
) -> list[str]:
) -> List[str]:
"""本地降级:基于模板规则生成标题(薄包装,转发到 ai_parsing 模块)."""
style_info = TITLE_STYLES.get(style, TITLE_STYLES["viral"])
return _generate_titles_fallback_base(description, style_info, count)
@@ -74,7 +74,7 @@ def generate_smart_titles(
description: str,
style: str = "viral",
count: int = 5,
) -> dict[str, Any]:
) -> Dict[str, Any]:
"""生成智能标题.
Args:
@@ -164,16 +164,16 @@ def generate_smart_titles(
def _semantic_match_fallback(
description: str,
assets: list[dict[str, Any]],
) -> list[dict[str, Any]]:
assets: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""本地降级:基于关键词的简单匹配(薄包装,转发到 ai_parsing 模块)."""
return _semantic_match_fallback_base(description, assets)
def _parse_semantic_match_response(
content: str,
asset_ids: list[str],
) -> Optional[dict[str, float]]:
asset_ids: List[str],
) -> Optional[Dict[str, float]]:
"""从模型返回中解析素材匹配度(薄包装,转发到 ai_parsing 模块)."""
result = _parse_semantic_match_base(content, asset_ids)
if result is None:
@@ -183,9 +183,9 @@ def _parse_semantic_match_response(
def semantic_match_assets(
description: str,
assets: list[dict[str, Any]],
assets: List[Dict[str, Any]],
top_k: int = 0,
) -> dict[str, Any]:
) -> Dict[str, Any]:
"""智能素材语义匹配.
根据用户描述,评估每个素材的语义匹配度并排序。
@@ -336,13 +336,13 @@ class AIService:
description: str,
style: str = "viral",
count: int = 5,
) -> dict[str, Any]:
) -> Dict[str, Any]:
return generate_smart_titles(description, style, count)
def semantic_match(
self,
description: str,
assets: list[dict[str, Any]],
assets: List[Dict[str, Any]],
top_k: int = 0,
) -> dict[str, Any]:
) -> Dict[str, Any]:
return semantic_match_assets(description, assets, top_k)
@@ -28,8 +28,8 @@ from __future__ import annotations
import json
import logging
from collections.abc import Callable
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Callable
from sqlalchemy.orm import Session
@@ -57,7 +57,7 @@ _REUSE_OVERLAP_RATIO = 0.6
def _now_iso() -> str:
return datetime.now(UTC).isoformat()
return datetime.now(timezone.utc).isoformat()
def _read_meta(model) -> dict:
@@ -156,7 +156,7 @@ def record_used_segments(
r["plan_id"] = plan_id
meta[USED_RANGES_KEY] = ranges
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
model.updated_at = datetime.now(timezone.utc)
return
ranges.append(
@@ -171,7 +171,7 @@ def record_used_segments(
)
meta[USED_RANGES_KEY] = ranges
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
model.updated_at = datetime.now(timezone.utc)
def remove_used_segment(
@@ -215,7 +215,7 @@ def remove_used_segment(
if removed:
meta[USED_RANGES_KEY] = remaining
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
model.updated_at = datetime.now(timezone.utc)
return removed
@@ -231,7 +231,7 @@ def reset_used_segments(db: Session, asset_id: str) -> None:
if meta.get(USED_RANGES_KEY):
meta[USED_RANGES_KEY] = []
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
model.updated_at = datetime.now(timezone.utc)
logger.info("[片段追踪] 素材区间记录手动清空: asset_id=%s", asset_id)
-243
View File
@@ -1,243 +0,0 @@
"""抖音视频解析多源轮询服务。
优先级(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
+89 -214
View File
@@ -7,7 +7,7 @@
from __future__ import annotations
import logging
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from app.services.asset_segment_tracker import (
REUSE_RATIO_LIMIT,
@@ -206,7 +206,7 @@ class EditPlanService:
status: Optional[EditPlanClipStatus] = None,
skip: int = 0,
limit: int = 100,
) -> list[EditPlanClip]:
) -> List[EditPlanClip]:
"""列出计划的片段"""
# 确保计划存在
self.get_plan_or_raise(plan_id)
@@ -423,7 +423,6 @@ 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,8 +473,6 @@ class EditPlanService:
name_suffix: str = "变体",
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)。
@@ -492,8 +489,6 @@ class EditPlanService:
created_by_user_id: 新 plan 归属用户。
name_suffix: plan 名后缀。
rng: 可选随机数(测试注入种子)。
batch_segments: 可选,外部传入的批次内已使用素材区间(前序变体避让用)。
传入时作为初始避让对象;未传则保持原逻辑从源 plan clips 自建(向后兼容)。
Raises:
ValueError: 源 plan 不存在/无片段、素材池为空或时长全未知。
@@ -505,7 +500,7 @@ class EditPlanService:
source = self.get_plan_or_raise(source_plan_id)
# 分页读取源 plan 全部片段
clips: list[EditPlanClip] = []
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(source_plan_id, skip=skip, limit=page)
@@ -542,26 +537,16 @@ class EditPlanService:
voice = float(voice_duration or 0.0)
except (TypeError, ValueError):
voice = 0.0
rhythm_template_for_reselect = None
if source.config:
rhythm_template_for_reselect = source.config.get("rhythm_template")
if voice > 0 and source_clips_data:
from packages.domain.voice_duration_planner import plan_clip_durations
_effects: list[str | None] = [c.get("transition_effect") for c in source_clips_data]
_tdurs: list[float] = [float(c.get("transition_duration") or 0.0) for c in source_clips_data]
# #1855 P0:先占位durations为空dict,真正查durations在后面pool_ids确定后执行;
# plan_clip_durations 的 asset_durations 参数在该函数中仅作最大段长钳制,
# 这里先不依赖它(durations 还没查),传 None 让planner用默认策略;
# 真正的asset_durations会在后面 clips_data 生成时传入 reselect_clips_for_variant
target_durations = plan_clip_durations(
len(source_clips_data),
voice,
transition_effects=_effects,
transition_durations=_tdurs,
rhythm_template=rhythm_template_for_reselect,
asset_durations=None,
)
if target_durations:
for _c, _d in zip(source_clips_data, target_durations, strict=False):
@@ -597,82 +582,25 @@ class EditPlanService:
created_by_user_id=created_by_user_id or (source.created_by_user_id or ""),
)
# 批次内区间:外部传入时使用外部传入(含前序变体已用区间);
# 否则保持原逻辑从源 plan clips 自建(向后兼容)
if batch_segments is not None:
batch_segments_resolved: dict[str, list[tuple[float, float]]] = {
k: list(v) for k, v in batch_segments.items()
}
else:
batch_segments_resolved = {}
for c in clips:
if c.asset_id and float(c.duration or 0) > 0:
st = float(c.start_time or 0.0)
batch_segments_resolved.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
# 批次内区间:以源 plan(变体 0)片段为初始避让对象
batch_segments: dict[str, list[tuple[float, float]]] = {}
for c in clips:
if c.asset_id and float(c.duration or 0) > 0:
st = float(c.start_time or 0.0)
batch_segments.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
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
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,
target_durations=target_durations,
rng=rng,
)
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
# 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit
for item in clips_data:
aid = item.get("asset_id", "")
if aid:
@@ -718,7 +646,7 @@ class EditPlanService:
source = self.get_plan_or_raise(source_plan_id)
# 分页读取源 plan 全部片段
clips: list[EditPlanClip] = []
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(source_plan_id, skip=skip, limit=page)
@@ -839,18 +767,7 @@ class EditPlanService:
if plan is None:
return None
# #1855 P0:幂等判断——如果已成功分配过且当前 total_duration 已接近 voice_duration,直接返回
try:
existing_mark = None
if plan.config:
existing_mark = plan.config.get("voice_duration_applied")
cur_total = float(plan.total_duration or 0.0)
if existing_mark is not None and abs(existing_mark - voice) < 1e-6 and abs(cur_total - voice) < 0.5:
return plan
except Exception:
pass
clips: list[EditPlanClip] = []
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(plan_id, skip=skip, limit=page)
@@ -921,10 +838,6 @@ class EditPlanService:
)
try:
plan.total_duration = net
# #1855 P0:写入幂等标记,避免二次调用时只重分配 duration 不重算 start_time
new_cfg = dict(plan.config or {})
new_cfg["voice_duration_applied"] = voice
plan.config = new_cfg
db = self._clip_repo.session
db.commit()
except Exception:
@@ -965,69 +878,6 @@ class EditPlanService:
rng = rng or _random.Random()
plan_ids: list[str] = []
# #1855 P0:先确定片段数 clip_count(用于节奏模板生成长度匹配)
from packages.domain.bgm_pool import allocate_bgm_pool_for_variants
from packages.domain.variant_plan_selector import (
generate_pixel_perturbation,
generate_visual_perturbation,
)
from packages.domain.voice_duration_planner import RHYTHM_TEMPLATES, adapt_template_length
clip_count = 0
# 从源 plan 获取片段数(分页读,避免关系加载问题)
_sclips: list = []
_sk, _pg = 0, 500
while True:
_b = self._clip_repo.list_by_plan(source_plan_id, skip=_sk, limit=_pg)
if not _b:
break
_sclips.extend(_b)
if len(_b) < _pg:
break
_sk += _pg
clip_count = len(_sclips)
# 预先生成所有 N 个变体的节奏模板/BGM/扰动参数(时机提前到选片前写入config)
rhythm_templates_for_variants: list = []
for _idx in range(count):
if clip_count > 0:
variant_seed = rng.randint(0, 999999)
_tpl = adapt_template_length(RHYTHM_TEMPLATES[variant_seed % len(RHYTHM_TEMPLATES)], clip_count)
rhythm_templates_for_variants.append(_tpl)
else:
rhythm_templates_for_variants.append(None)
source_bgm_config: dict = {}
source_plan = self.get_plan(source_plan_id)
if source_plan and source_plan.config:
source_bgm_config = source_plan.config.get("bgm", {}) or {}
variant_seeds_for_bgm = [rng.randint(0, 999999) for _ in range(count)]
bgm_pool_assignments = allocate_bgm_pool_for_variants(source_bgm_config, variant_seeds_for_bgm)
def _build_variant_config_update(idx: int) -> dict:
"""构建单个变体的 config 更新(节奏模板/BGM/视觉/像素扰动)。"""
upd: dict = {}
try:
perturbation = generate_visual_perturbation(rng)
if idx == 0:
perturbation["hflip"] = False
upd["visual_perturbation"] = perturbation
except Exception:
logger.exception("变体 %d 视觉扰动生成失败(不阻断)", idx)
try:
pixel_pert = generate_pixel_perturbation(rng)
upd["pixel_perturbation"] = pixel_pert
except Exception:
logger.exception("变体 %d 像素扰动生成失败(不阻断)", idx)
rt = rhythm_templates_for_variants[idx] if idx < len(rhythm_templates_for_variants) else None
if rt is not None:
upd["rhythm_template"] = rt
if idx < len(bgm_pool_assignments):
existing_bgm = dict((source_plan.config or {}).get("bgm", {}) or {})
existing_bgm.update(bgm_pool_assignments[idx])
upd["bgm"] = existing_bgm
return upd
# 变体 0:clone(片段结构同源 plan,起点重算),不污染源 plan
plan0 = self.clone_plan_for_variant(
source_plan_id,
@@ -1040,15 +890,6 @@ class EditPlanService:
v0_voice = float(voice_durations[0] or 0.0)
except (TypeError, ValueError):
v0_voice = 0.0
# #1855 P0:在配音分配前先写入变体0的节奏模板/扰动/BGM,确保 apply_voice_duration_to_plan 能读到 rhythm_template
try:
_cfg0 = _build_variant_config_update(0)
if _cfg0:
self.update_plan_config(plan0.id, _cfg0)
except Exception:
logger.exception("变体0 配置写入失败(不阻断): plan=%s", plan0.id)
if v0_voice > 0:
try:
self.apply_voice_duration_to_plan(plan0.id, v0_voice)
@@ -1056,12 +897,7 @@ class EditPlanService:
logger.exception("变体0 配音分配失败(不阻断): plan=%s", plan0.id)
plan_ids.append(plan0.id)
# #1855 P0:批次内素材区间避让表——从变体0实际落库的clips构建初始值(公共函数)
from app.services.generation_common import collect_plan_segments as _collect_plan_segments
batch_segments_acc: dict[str, list[tuple[float, float]]] = _collect_plan_segments(plan0.id, self._clip_repo)
# 变体 1..N-1:独立选片(传入累积的 batch_segments 做区间避让)
# 变体 1..N-1:独立选片
for i in range(1, count):
voice = 0.0
if voice_durations and i < len(voice_durations):
@@ -1069,14 +905,6 @@ class EditPlanService:
voice = float(voice_durations[i] or 0.0)
except (TypeError, ValueError):
voice = 0.0
# #1855 P0:在reselect前先为"变体i"准备配置更新——但reselect内部复制的是source.config
# 所以每个变体独立的节奏模板需要在reselect后单独写入config
# 但 plan_clip_durations 用的是 source.config.rhythm_template(即源plan的节奏模板),
# 为了让每个变体在选片阶段就使用自己的节奏模板分配段长,这里采用:
# - reselect 仍使用源 plan 的 rhythm_template(保持片段骨架一致)
# - 选片完成后立即写入该变体自己的 rhythm_template/扰动/BGM 到config
# 后续不再二次 apply_voice_duration_to_plan(由幂等标记跳过)
variant = self.reselect_plan_for_variant(
source_plan_id,
candidate_asset_ids,
@@ -1084,26 +912,73 @@ class EditPlanService:
name_suffix=f"变体{i + 1}",
voice_duration=voice,
rng=rng,
batch_segments=batch_segments_acc,
)
# 选片完成后写入该变体的独立配置(节奏模板/扰动/BGM)
try:
_cfgi = _build_variant_config_update(i)
if _cfgi:
self.update_plan_config(variant.id, _cfgi)
except Exception:
logger.exception("变体 %d 配置写入失败(不阻断): plan=%s", i, variant.id)
plan_ids.append(variant.id)
# #1855 P0:把当前新变体的 clips 区间追加到 batch_segments,供下一变体避让
# #1764:为每个变体生成独立节奏模板(让批量视频片段时长分布不同)
from packages.domain.voice_duration_planner import RHYTHM_TEMPLATES, adapt_template_length
clip_count = 0
if voice_durations and len(voice_durations) > 0:
# 从源 plan 获取片段数
source_plan = self.get_plan(source_plan_id)
if source_plan and hasattr(source_plan, "clips"):
clip_count = len(list(source_plan.clips)) if source_plan.clips else 0
rhythm_templates_for_variants = []
if clip_count > 0:
for idx in range(len(plan_ids)):
# 每个变体用不同的 seed 选择节奏模板
variant_seed = rng.randint(0, 999999)
template = adapt_template_length(RHYTHM_TEMPLATES[variant_seed % len(RHYTHM_TEMPLATES)], clip_count)
rhythm_templates_for_variants.append(template)
logger.info("变体 %d 节奏模板: plan=%s template=%s", idx, plan_ids[idx], template)
# #1767:BGM 池差异化分配(让批量变体使用不同 BGM / 段落 / 音量)
from packages.domain.bgm_pool import allocate_bgm_pool_for_variants
source_bgm_config = {}
source_plan = self.get_plan(source_plan_id)
if source_plan and source_plan.config:
source_bgm_config = source_plan.config.get("bgm", {}) or {}
variant_seeds_for_bgm = [rng.randint(0, 999999) for _ in plan_ids]
bgm_pool_assignments = allocate_bgm_pool_for_variants(source_bgm_config, variant_seeds_for_bgm)
# 为每个变体生成独立视觉扰动参数(让批量视频画面本身更不同)
from packages.domain.variant_plan_selector import generate_visual_perturbation
for idx, pid in enumerate(plan_ids):
try:
_new_segs = _collect_plan_segments(variant.id, self._clip_repo)
for _aid, _ivs in _new_segs.items():
batch_segments_acc.setdefault(_aid, []).extend(_ivs)
perturbation = generate_visual_perturbation(rng)
# 变体 0 不做 hflip(保持预览 plan 原始画面方向)
if idx == 0:
perturbation["hflip"] = False
config_update = {"visual_perturbation": perturbation}
# #1764:写入节奏模板
if idx < len(rhythm_templates_for_variants):
config_update["rhythm_template"] = rhythm_templates_for_variants[idx]
# #1765:写入像素级扰动滤镜
from packages.domain.variant_plan_selector import generate_pixel_perturbation
pixel_pert = generate_pixel_perturbation(rng)
config_update["pixel_perturbation"] = pixel_pert
# #1767:写入 BGM 池分配(覆盖 bgm 配置中的 preset_id / audio_offset / volume_adjust_db
if idx < len(bgm_pool_assignments):
existing_bgm = dict((source_plan.config or {}).get("bgm", {}) or {})
existing_bgm.update(bgm_pool_assignments[idx])
config_update["bgm"] = existing_bgm
self.update_plan_config(pid, config_update)
logger.info(
"变体 %d 视觉扰动+像素扰动+BGM池: plan=%s vis=%s pix=%s bgm=%s",
idx,
pid,
perturbation,
pixel_pert,
bgm_pool_assignments[idx] if idx < len(bgm_pool_assignments) else None,
)
except Exception:
logger.exception("变体 %d 区间收集失败(不阻断): plan=%s", i, variant.id)
logger.exception("变体 %d 视觉扰动生成失败(不阻断): plan=%s", idx, pid)
# 标记所有变体 plan 的 clips 为 ready(已分配素材+起点,语义上就是 ready)
for pid in plan_ids:
@@ -1116,7 +991,7 @@ class EditPlanService:
# ── 片段分割与合并 ──────────────────────────────────────────────────────
def split_clip(self, clip_id: str, split_time: float) -> dict[str, Any]:
def split_clip(self, clip_id: str, split_time: float) -> Dict[str, Any]:
"""将一个片段从指定位置分割为两个片段
Args:
@@ -1204,7 +1079,7 @@ class EditPlanService:
"right_clip": created_right,
}
def merge_clips(self, clip_ids: list[str]) -> EditPlanClip:
def merge_clips(self, clip_ids: List[str]) -> EditPlanClip:
"""合并多个连续片段为一个片段
Args:
@@ -1270,7 +1145,7 @@ class EditPlanService:
# ── 渲染生成流程 ────────────────────────────────────────────────────────
def get_generation_status(self, plan_id: str) -> dict[str, Any]:
def get_generation_status(self, plan_id: str) -> Dict[str, Any]:
"""获取渲染进度状态
Returns:
@@ -1417,7 +1292,7 @@ class EditPlanService:
)
return count
def update_plan_config(self, plan_id: str, config_updates: dict[str, Any]) -> EditPlan:
def update_plan_config(self, plan_id: str, config_updates: Dict[str, Any]) -> EditPlan:
"""更新计划配置(合并更新)
Args:
@@ -7,7 +7,7 @@
from __future__ import annotations
import logging
from typing import Any, Optional
from typing import Any, List, Optional
from sqlalchemy.orm import Session
@@ -76,7 +76,7 @@ class EditTemplateService:
active_only: bool = False,
skip: int = 0,
limit: int = 50,
) -> list[EditTemplate]:
) -> List[EditTemplate]:
"""列出模板
Args:
@@ -227,7 +227,7 @@ class EditTemplateService:
clip_type: Optional[ClipType] = None,
skip: int = 0,
limit: int = 100,
) -> list[TemplateClipConfig]:
) -> List[TemplateClipConfig]:
"""列出模板的片段配置
注意:本方法要求模板存在于新表 ``edit_templates``(全局模板库),
@@ -253,7 +253,7 @@ class EditTemplateService:
clip_type: Optional[ClipType] = None,
skip: int = 0,
limit: int = 100,
) -> list[TemplateClipConfig]:
) -> List[TemplateClipConfig]:
"""编辑器读取模板片段配置的单一数据源入口.
片段配置主表是 ``template_clip_configs``(直接读取,不抛异常、不降级)。
@@ -404,8 +404,8 @@ class EditTemplateService:
def reorder_clip_configs(
self,
template_id: str,
config_ids: list[str],
) -> list[TemplateClipConfig]:
config_ids: List[str],
) -> List[TemplateClipConfig]:
"""重新排序片段配置
Args:
@@ -560,7 +560,7 @@ class EditTemplateService:
)
# 5. 转换每个片段为模板片段配置
created_configs: list[TemplateClipConfig] = []
created_configs: List[TemplateClipConfig] = []
for clip_config_obj in clips_to_template_clip_configs(created_template.id, clips):
created = self._clip_config_repo.create(clip_config_obj)
created_configs.append(created)
-233
View File
@@ -1,233 +0,0 @@
"""智能剪辑公共服务辅助函数(从 route 层下沉)。
集中管理:
- query_voice_durations:批量查询配音素材时长
- writeback_edit_plan_config:任务入队后回写 EditPlan.config
- collect_plan_segments:分页读取 plan clips 构建素材区间表(变体避让用)
- resolve_latest_plan_by_template:按 template_id + user_id 查最新 EditPlan
设计原则:
- 无副作用的纯查询 / 幂等写回;失败一律不阻断主流程(记日志 + 返回安全默认值)
- 不依赖 FastAPI / HTTPException,便于 service 层和 worker 复用
"""
from __future__ import annotations
import logging
from typing import Any, Optional
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
def query_voice_durations(db: Session, voice_ids: list[str]) -> list[float]:
"""批量查询配音素材时长(秒),#1749 配音时长分配用。
逐项 try/float 硬化:MagicMock/异常/缺失 → 0.0(无配音不分配,不阻断)。
#1855 P0修复:不再对 voice_ids 去重,保持与调用方传入顺序/长度一致,
允许同配音id多次出现时返回相同时长(支持"同配音N变体"的时长对齐)。
"""
raw_ids = list(voice_ids or [])
if not raw_ids:
return []
unique_ids: list[str] = []
_seen: set[str] = set()
for v in raw_ids:
if v and v not in _seen:
_seen.add(v)
unique_ids.append(v)
if not unique_ids:
return [0.0 for _ in raw_ids]
try:
from packages.adapters.sqlalchemy_impl.models import AssetModel
rows = db.query(AssetModel.id, AssetModel.duration).filter(AssetModel.id.in_(unique_ids)).all()
dur_map: dict[str, float] = {}
for row in rows:
try:
dur_map[row[0]] = float(row[1] or 0.0)
except (TypeError, ValueError):
dur_map[row[0]] = 0.0
return [dur_map.get(v, 0.0) if v else 0.0 for v in raw_ids]
except Exception:
logger.warning("[generation_common] 配音时长查询失败(按无配音处理,不阻断)", exc_info=True)
return [0.0 for _ in raw_ids]
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:
"""任务入队成功后,回写 EditPlan.configgeneration_task_id + title_config。
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
#1970dedup_enabled 非 None 时一并写入,worker 据此决定 edge_crop/微变换;
PR3 叙事模式再写 assembly_mode/script_id/video_ratio(可追溯,不影响渲染)。
失败只记日志,不影响任务创建。
"""
if not plan_id:
return
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first()
if plan_model is None:
logger.warning("[generation_common] 回写plan.config失败: plan不存在 plan_id=%s", plan_id)
return
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"
# 先读取新旧两个 key,判断标题文字是否变化
old_title_cfg = merged.get("title", {}) or {}
if not isinstance(old_title_cfg, dict) or not (old_title_cfg.get("text") or "").strip():
old_title_cfg = merged.get("title_config", {}) or {}
old_title_text = (old_title_cfg.get("text") or "").strip() if isinstance(old_title_cfg, dict) else ""
new_title_text = (title_config.get("text") or "").strip()
if old_title_text != new_title_text:
if "cover" in merged:
del merged["cover"]
logger.info(
"[generation_common] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s",
plan_id,
old_title_text,
new_title_text,
)
# 字段名归一化(font_size→size, font_preset→font, font_color→color),与 worker sync_configs_to_plan 保持一致
normalized = dict(title_config)
if "font_size" in normalized and "size" not in normalized:
normalized["size"] = normalized["font_size"]
if "font_preset" in normalized and "font" not in normalized:
normalized["font"] = normalized["font_preset"]
if "font_color" in normalized and "color" not in normalized:
normalized["color"] = normalized["font_color"]
merged["title"] = normalized
# 清掉旧 key,避免双字段并存
merged.pop("title_config", None)
plan_model.config = merged
db.commit()
logger.info(
"[generation_common] 回写plan.config成功: plan_id=%s task_id=%s keys=%s",
plan_id,
task_id,
list(merged.keys()),
)
except Exception as e:
logger.warning(
"[generation_common] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s",
plan_id,
e,
exc_info=True,
)
try:
db.rollback()
except Exception:
pass
def collect_plan_segments(
plan_id: str,
clip_repo: Any,
*,
page_size: int = 500,
) -> dict[str, list[tuple[float, float]]]:
"""分页读取 plan 所有 clips,构建 {asset_id: [(start, end), ...]} 素材区间表。
用于 #1855 P0 批次内素材区间避让(变体间素材片段重叠控制)。
"""
segs: dict[str, list[tuple[float, float]]] = {}
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:
if c.asset_id and float(c.duration or 0) > 0:
st = float(c.start_time or 0.0)
segs.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
if len(batch) < pg:
break
sk += pg
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,
*,
template_id: str,
user_id: str,
) -> Optional[str]:
"""按 template_id + user_id 查找最新的 EditPlan.id(模板兜底用)。找不到返回 None。"""
if not (template_id or "").strip():
return None
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == template_id.strip(),
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
return latest.id if latest else None
except Exception:
logger.warning(
"[generation_common] 按template查找最新plan失败: template=%s user=%s",
template_id,
user_id,
exc_info=True,
)
return None
@@ -1,382 +0,0 @@
"""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_secondsWorker 推理期会通过 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=%sstatus=%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_atclaim 时写入,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 心跳超时,回退 pendingattempt=%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:
终态 taskstatus=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)
+89 -549
View File
@@ -1,25 +1,19 @@
"""对口型 Service — #1796 MediaKit 对口型业务逻辑, #1809 参数调整, #1845 配音前置.
"""对口型 Service — #1796 MediaKit 对口型业务逻辑, #1809 参数调整.
职责:
- 创建/查询对口型任务
- 三输入模式:
1. TTS 直生(voice_id + script_text)→ 走 Celery 异步(降级路径)
2. 直接音频(audio_url,前端未传 timings)→ 同步下载 + 算 timings + 提交 MediaKit
3. 预合成音频(audio_url + sentence_timings#1845 新主路径)→ 同步 ffprobe 校验时长 +
写入前端传来的 timings → 直接提交 MediaKit~2-3s
- 创建/查询/取消对口型任务
- 调用 TTS 合成音频(#1809:前端不再传 audio_url
- 调用 MediaKit 客户端提交异步任务
- 轮询更新任务状态(中间状态同步 DB,成片转存自家 OSS)
- 轮询更新任务状态
- 用户隔离(每个用户只能操作自己的任务)
"""
from __future__ import annotations
import io
import logging
import uuid
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Optional
from urllib.parse import urlparse
from app.services.mediakit_client import (
STATUS_COMPLETED,
@@ -29,26 +23,13 @@ from app.services.mediakit_client import (
MediaKitError,
get_mediakit_client,
)
# Celery 异步任务:TTS 合成 + MediaKit 提交(降级路径)
from app.tasks.lipsync_tts import tts_synthesize_and_submit
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,
)
from packages.shared.storage import get_shared_storage_service
from packages.shared.url_security import ALLOWED_AUDIO_MIME_TYPES, safe_download_bytes
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
logger = logging.getLogger(__name__)
# 传给 MediaKit GPU worker / 回给前端播放的 OSS 预签名有效期:7 天。
MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
class LipsyncService:
"""对口型任务 Service."""
@@ -57,284 +38,19 @@ class LipsyncService:
self,
db: Session,
client: Optional[MediaKitClient] = None,
cosyvoice_service=None,
voice_clone_repo=None,
cosyvoice_service: Optional[CosyVoiceService] = None,
):
self.db = db
self.client = client or get_mediakit_client()
self._cosyvoice = cosyvoice_service
self._voice_clone_repo = voice_clone_repo
self.settings = get_api_settings()
self._cosyvoice_service = cosyvoice_service
def _get_cosyvoice(self):
"""延迟获取 CosyVoiceService(与 tts 路由一致,含 OSS 预签名配置)."""
if self._cosyvoice is None:
@property
def cosyvoice_service(self) -> CosyVoiceService:
if self._cosyvoice_service is None:
from app.dependencies import get_cosyvoice_service
self._cosyvoice = get_cosyvoice_service()
return self._cosyvoice
def _resolve_voice_id(self, voice_id: str, user_id: str) -> str:
"""将克隆音色 profile UUID 解析为 CosyVoice voice_id。
与 /tts/synthesize 保持一致:命中 profile → 校验归属 → 返回其 voice_id;
未命中(预置音色 ID 或克隆 CosyVoice voice_id)原样返回。
"""
if not voice_id:
return ""
if self._voice_clone_repo is None:
try:
from app.dependencies import get_voice_clone_profile_repository
self._voice_clone_repo = get_voice_clone_profile_repository(self.db)
except Exception:
return voice_id
try:
profile = self._voice_clone_repo.get(voice_id)
except Exception:
return voice_id
if profile is None:
return voice_id
if getattr(profile, "user_id", "") != user_id:
raise MediaKitError("无权访问该音色", code="VoiceForbidden")
if not getattr(profile, "voice_id", ""):
raise MediaKitError("音色克隆尚未完成,请稍后再试", code="VoiceNotReady")
return profile.voice_id
def _synthesize_and_persist_audio(
self,
*,
user_id: str,
job_id: str,
voice_id: str,
script_text: str,
speed: float,
emotion: str,
) -> str:
"""TTS 直生:调 CosyVoice 合成音频并转存 OSS,返回可公网访问的音频 URL.
Raises:
MediaKitError: 合成失败
"""
actual_voice_id = self._resolve_voice_id(voice_id, user_id)
cosyvoice = self._get_cosyvoice()
try:
result = cosyvoice.submit_synthesize_task(
text=script_text,
voice_id=actual_voice_id,
speed=speed,
emotion=emotion, # normalize 在 CosyVoiceService 内部完成
language="zh",
)
except CosyVoiceError as exc:
raise MediaKitError(f"TTS 合成失败: {exc}", code="TTSSynthesisFailed") from exc
except ValueError as exc:
raise MediaKitError(f"TTS 参数错误: {exc}", code="TTSInvalidParam") from exc
temp_url = result.get("audio_url", "")
if not temp_url:
raise MediaKitError("TTS 未返回音频 URL", code="TTSNoAudio")
# 转存到自家 OSS,避免临时 URL 过期导致 MediaKit 拉取失败
try:
audio_data = safe_download_bytes(
temp_url,
purpose="lipsync_tts_audio",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
storage = get_shared_storage_service()
storage_key = f"lipsync-tts/{user_id}/{job_id}.mp3"
permanent_url = storage.upload_file(io.BytesIO(audio_data), storage_key, content_type="audio/mpeg")
logger.info("对口型 TTS 音频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
return permanent_url
except Exception as exc:
logger.warning("TTS 音频转存 OSS 失败,回退临时 URL: job_id=%s err=%s", job_id, exc)
return temp_url
def _submit_audio_direct(
self,
*,
job: LipsyncJobModel,
supplied_timings: Optional[list] = None,
supplied_duration: Optional[float] = None,
) -> None:
"""音频直传模式(包含 #1845 预合成路径):同步下载 → ffprobe → timings → 提交 MediaKit.
直接在 HTTP 请求内完成,不走 Celery。job.status 成功后置为 submitted。
失败时把 job 标成 failed 并 commit,然后抛 MediaKitError。
Args:
job: 已 commit 的 LipsyncJobModelaudio_url / video_url 已写入)
supplied_timings: 前端传来的预合成 timings(可选,可信时直接用)
supplied_duration: 前端传来的预合成时长(可选,用于优先避免重复探测)
"""
# 1. 下载音频
audio_data: bytes | None = None
try:
audio_data = safe_download_bytes(
job.audio_url,
purpose="lipsync_direct_audio",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
logger.info(
"[lipsync] 直传音频下载完成: job_id=%s size=%d",
job.id,
len(audio_data) if audio_data else 0,
)
except Exception as exc:
logger.warning("[lipsync] 直传音频下载失败,跳过 timings 计算: job_id=%s err=%s", job.id, exc)
# 2. ffprobe 探测时长(优先用前端传入的预合成时长,但以 ffprobe 为准做兜底校验)
audio_duration = 0.0
if audio_data:
audio_duration = probe_audio_duration(audio_data)
if audio_duration <= 0 and supplied_duration and supplied_duration > 0:
audio_duration = supplied_duration
logger.info(
"[lipsync] ffprobe 失败,使用前端传入的预合成时长: job_id=%s duration=%.2f", job.id, audio_duration
)
# 3. 句子时间戳:优先用前端预合成传入的 timings(后端预合成接口已经算过,可信);
# 否则若音频下载成功则重算;否则不设置(不阻塞主流程)
timings: Optional[list] = None
if supplied_timings:
timings = supplied_timings
logger.info("[lipsync] 使用前端预合成句子时间戳: job_id=%s sentences=%d", job.id, len(timings))
elif audio_data and audio_duration > 0 and job.script_text:
try:
timings = compute_sentence_timings(audio_data, job.script_text, audio_duration)
logger.info(
"[lipsync] 后端重算句子时间戳: job_id=%s sentences=%d duration=%.2f",
job.id,
len(timings) if timings else 0,
audio_duration,
)
except Exception as exc:
logger.warning("[lipsync] 句子时间戳计算失败(不阻塞): job_id=%s err=%s", job.id, exc)
if timings:
job.sentence_timings = timings
# 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(job=job, gpu_svc=gpu_svc)
if gpu_task is not None:
# GPU 任务完成:直接把结果写入 job,标为 completed
job.mediakit_task_id = "" # GPU 路径不走 MediaKit
job.status = STATUS_COMPLETED
job.output_video_url = gpu_task.result_url
job.output_duration = gpu_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 MuseTalk 推理完成: job_id=%s gpu_task=%s duration=%.2f",
job.id,
gpu_task.id,
job.output_duration,
)
# 转存到持久 OSS 路径(GPU 结果已在 gpu-lipsync/results/ 下,直接签短链)
return
# wait_for_result 返回 None 表示超时/最终失败 → 继续走 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
try:
result = self.client.submit_lipsync(
video_url=video_url,
audio_url=signed_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)
self.db.commit()
logger.info(
"[lipsync] 直传音频已提交 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
logger.error("[lipsync] 直传音频提交 MediaKit 失败: job_id=%s err=%s", job.id, exc)
self.db.commit()
raise
# ── GPU MuseTalk 路径 ────────────────────────────────────────────────
def _submit_to_gpu(self, *, job, gpu_svc) -> Optional[object]:
"""创建 GPU 任务并同步等待结果。
成功返回终态 task 对象(status=done);超时或 GPU 最终失败返回 None,
调用方回退 MediaKit。
注意:job.video_url / job.audio_url 可能是:
- 自家 OSS 存储 keystorage.is_own_url 判断,gpu_svc.create_task 内部
get_download_url 会自动签预签名 URL 给 Worker)
- 外部公网 URLCosyVoice 临时链接等):poll 返回时原样透传给 Worker
Worker 可直接 GET 下载。
"""
# 创建 GPU 任务
gpu_task = gpu_svc.create_task(
video_url=job.video_url,
audio_url=job.audio_url,
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,
)
# 同步等待 Worker 处理完成(轮询 DB)
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 None
if final_task.status != "done":
logger.warning(
"[lipsync] GPU 任务失败: gpu_task=%s status=%s err=%s",
gpu_task.id,
final_task.status,
final_task.error_msg,
)
return None
# result_url 是 OSS 存储 key;签一个长有效期 URL 写回 job.output_video_url
result_signed = self._sign_media_url(final_task.result_url)
final_task.result_url = result_signed or final_task.result_url
return final_task
self._cosyvoice_service = get_cosyvoice_service()
return self._cosyvoice_service
# ── 创建任务 ──────────────────────────────────────────────────────────
@@ -343,49 +59,47 @@ class LipsyncService:
*,
user_id: str,
video_url: str,
audio_url: str = "",
audio_duration: Optional[float] = None,
sentence_timings: Optional[list] = None,
voice_id: str = "",
script_text: str = "",
speed: float = 1.0,
emotion: str = "",
enable_video_loop: bool = True,
voice_id: str,
script_text: str,
enable_video_loop: bool = False,
project_id: str = "",
) -> LipsyncJobModel:
"""创建对口型任务.
"""创建对口型任务并提交到 MediaKit.
三种输入模式:
- TTS 直生:voice_id + script_textaudio_url 留空)
→ 创建 DB 记录(状态 tts_processing),dispatch Celery 异步任务(降级路径)。
API 响应 <1s。
- 直接音频:audio_url 非空 + 无 sentence_timings
→ 同步下载音频 + 重算 timings + 提交 MediaKit(几秒完成)。
- 预合成音频(#1845 新主路径):audio_url 非空 + 传 sentence_timings
→ 同步 ffprobe 校验时长 + 写入 timings + 提交 MediaKit~2-3s)。
#1809: 内部调 TTS 合成音频,不再由前端传 audio_url。
Raises:
MediaKitError: 参数校验失败或 MediaKit 提交失败
CosyVoiceError: TTS 合成失败
MediaKitError: API 调用失败
"""
# 0. 输入校验
is_pre_synth = bool(audio_url) and bool(sentence_timings)
bool(audio_url) and not is_pre_synth
is_tts_mode = not bool(audio_url)
# 1. 调 TTS 合成音频
try:
tts_result = self.cosyvoice_service.synthesize_speech(
text=script_text,
voice_id=voice_id,
)
audio_url = tts_result.audio_url
except CosyVoiceError as exc:
logger.error("TTS 合成失败: voice_id=%s, error=%s", voice_id, exc)
# 创建失败记录
job_id = str(uuid.uuid4())
job = LipsyncJobModel(
id=job_id,
user_id=user_id,
project_id=project_id,
video_url=video_url,
audio_url="",
enable_video_loop=enable_video_loop,
status="failed",
error_message=f"TTS 合成失败: {exc}",
error_code="TTSSynthesisFailed",
)
self.db.add(job)
self.db.commit()
self.db.refresh(job)
raise
if is_tts_mode:
if not (voice_id and script_text):
raise MediaKitError(
"必须提供 audio_url 或 voice_id+script_text",
code="InvalidInput",
)
# TTS 模式:在 HTTP 请求中同步校验音色归属,快速失败
self._resolve_voice_id(voice_id, user_id)
elif is_pre_synth:
# 预合成模式:script_text 可空(因为 timings 已自带句子文本),但仍建议传
if not isinstance(sentence_timings, list) or len(sentence_timings) == 0:
raise MediaKitError("预合成模式 sentence_timings 不能为空", code="InvalidInput")
# 1. 创建数据库记录
# 2. 创建数据库记录
job_id = str(uuid.uuid4())
job = LipsyncJobModel(
id=job_id,
@@ -394,140 +108,33 @@ class LipsyncService:
video_url=video_url,
audio_url=audio_url,
enable_video_loop=enable_video_loop,
voice_id=voice_id or "",
script_text=script_text or "",
speed=speed,
emotion=emotion or "",
# 音频直传(含预合成)直接进入 pending(后续同步改为 submitted);TTS 模式进入 tts_processing
status="tts_processing" if is_tts_mode else "pending",
status="pending",
)
self.db.add(job)
self.db.flush()
# ⚠️ 必须先 commit 再发 Celery 任务 / 后续同步操作,避免事务竞态
# 3. 提交到 MediaKit
try:
result = self.client.submit_lipsync(
video_url=video_url,
audio_url=audio_url,
enable_video_loop=enable_video_loop,
client_token=job_id, # 幂等控制
)
job.mediakit_task_id = result["task_id"]
job.status = "submitted"
job.submitted_at = datetime.now(timezone.utc)
except MediaKitError as exc:
job.status = "failed"
job.error_message = str(exc)
job.error_code = exc.code
logger.error("提交对口型任务失败: %s", exc)
raise
self.db.commit()
self.db.refresh(job)
if is_tts_mode:
# 2a. TTS 模式:dispatch Celery 异步任务处理 TTS 合成 + MediaKit 提交(降级路径)
try:
tts_synthesize_and_submit.apply_async(
args=(
job_id,
user_id,
voice_id,
script_text,
speed,
emotion or "",
)
)
except Exception as exc:
logger.exception(
"Celery 任务提交失败,TTS 任务已创建但未触发执行: job_id=%s err=%s",
job_id,
exc,
)
job.status = "failed"
job.error_message = f"Celery 任务投递失败: {exc}"
job.error_code = "AsyncDispatchFailed"
job.updated_at = datetime.now(UTC)
self.db.commit()
else:
# 2b/2c. 直接音频 / 预合成音频:同步路径
self._submit_audio_direct(
job=job,
supplied_timings=sentence_timings,
supplied_duration=audio_duration,
)
self.db.refresh(job)
return job
# ── TTS 预合成(#1845 步骤1「生成配音」同步接口使用) ──────────────────
def preview_tts(
self,
*,
user_id: str,
voice_id: str,
script_text: str,
speed: float = 1.0,
emotion: str = "neutral",
) -> dict:
"""同步做 TTS 合成 + 下载 + ffprobe + 句子时间戳计算.
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL~24h 有效期)。
耗时约 2-3 秒,由前端在步骤1点「生成配音」时同步等待。
Returns:
{"audio_url": str, "duration": float, "sentence_timings": list[dict]}
Raises:
MediaKitError: TTS 合成失败 / 下载失败 / ffprobe 失败
"""
# 1. 音色解析(校验克隆音色归属)
actual_voice_id = self._resolve_voice_id(voice_id, user_id)
cosyvoice = self._get_cosyvoice()
# 2. TTS 合成(同步,~2-3s
try:
result = cosyvoice.submit_synthesize_task(
text=script_text,
voice_id=actual_voice_id,
speed=speed,
emotion=emotion, # normalize 在 CosyVoiceService 内部完成
language="zh",
)
except CosyVoiceError as exc:
raise MediaKitError(f"TTS 合成失败: {exc}", code="TTSSynthesisFailed") from exc
except ValueError as exc:
raise MediaKitError(f"TTS 参数错误: {exc}", code="TTSInvalidParam") from exc
temp_url = result.get("audio_url", "")
if not temp_url:
raise MediaKitError("TTS 未返回音频 URL", code="TTSNoAudio")
# 3. 下载音频到内存(用于 ffprobe + 静音检测)
try:
audio_data = safe_download_bytes(
temp_url,
purpose="tts_preview_audio",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
except Exception as exc:
logger.warning("[tts-preview] TTS 音频下载失败,仍返回 audio_url: user_id=%s err=%s", user_id, exc)
return {
"audio_url": temp_url,
"duration": 0.0,
"sentence_timings": [],
}
# 4. ffprobe 时长
duration = probe_audio_duration(audio_data)
if duration <= 0:
logger.warning("[tts-preview] ffprobe 未返回有效时长,timings 留空: user_id=%s", user_id)
return {
"audio_url": temp_url,
"duration": 0.0,
"sentence_timings": [],
}
# 5. 句子时间戳
timings = compute_sentence_timings(audio_data, script_text, duration)
logger.info(
"[tts-preview] TTS 预合成完成: user_id=%s duration=%.2f sentences=%d",
user_id,
duration,
len(timings),
)
return {
"audio_url": temp_url,
"duration": round(duration, 2),
"sentence_timings": timings,
}
# ── 查询任务 ──────────────────────────────────────────────────────────
def get_job(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
@@ -561,7 +168,11 @@ class LipsyncService:
# ── 更新任务状态(轮询) ──────────────────────────────────────────────
def refresh_job_status(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
"""从 MediaKit 拉取最新状态并更新本地记录."""
"""从 MediaKit 拉取最新状态并更新本地记录.
Returns:
更新后的 Job,或 None(任务不存在/不属于该用户)
"""
job = self.get_job(job_id, user_id)
if job is None:
return None
@@ -581,107 +192,36 @@ class LipsyncService:
return job
mk_status = status_data.get("status", STATUS_RUNNING)
logger.info("MediaKit 对口型状态 [%s]: %s", job_id, mk_status)
try:
if mk_status == STATUS_COMPLETED:
result = status_data.get("result", {})
job.status = STATUS_COMPLETED
temp_url = result.get("video_url", "")
job.output_video_url = temp_url
job.output_duration = result.get("duration", 0.0)
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
self.db.commit()
# 异步转存自家 OSS
try:
from app.tasks.lipsync_tts import persist_output_video_task
persist_output_video_task.apply_async(args=(job_id, user_id, temp_url))
except Exception as exc:
logger.warning(
"提交输出视频异步转存任务失败,保留临时 URL: job_id=%s err=%s",
job_id,
exc,
)
elif mk_status == STATUS_FAILED:
error = status_data.get("error", {})
job.status = "failed"
job.error_message = error.get("message", "任务执行失败")
job.error_code = error.get("code", "TaskFailed")
job.completed_at = datetime.now(UTC)
else:
# 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted
if isinstance(mk_status, str) and mk_status:
job.status = mk_status
job.updated_at = datetime.now(UTC)
self.db.commit()
except Exception as exc: # noqa: BLE001 - DB 提交失败必须记录日志并重试,否则后台任务静默失败
logger.error(
"refresh_job_status 提交 DB 失败 job_id=%s mk_status=%s err=%s",
job_id,
mk_status,
exc,
exc_info=True,
)
try:
self.db.rollback()
except Exception:
pass
# DB commit 失败不 raise,返回当前 job 对象让下次轮询再试
if mk_status == STATUS_COMPLETED:
result = status_data.get("result", {})
job.status = STATUS_COMPLETED
job.output_video_url = result.get("video_url", "")
job.output_duration = result.get("duration", 0.0)
job.completed_at = datetime.now(timezone.utc)
elif mk_status == STATUS_FAILED:
error = status_data.get("error", {})
job.status = "failed"
job.error_message = error.get("message", "任务执行失败")
job.error_code = error.get("code", "TaskFailed")
job.completed_at = datetime.now(timezone.utc)
# running 状态只更新时间戳
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
def _persist_output_video(self, temp_url: str, job_id: str, user_id: str) -> str:
"""将 MediaKit 输出的临时视频 URL 转存到自家 OSS. 失败时回退返回原始临时 URL."""
if not temp_url:
return ""
try:
import httpx
with httpx.Client(timeout=180.0, follow_redirects=True) as client:
resp = client.get(temp_url)
resp.raise_for_status()
data = resp.content
storage = get_shared_storage_service()
storage_key = f"lipsync-outputs/{user_id}/{job_id}.mp4"
permanent_url = storage.upload_file(io.BytesIO(data), storage_key, content_type="video/mp4")
logger.info("对口型输出视频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
return self._sign_media_url(permanent_url) or temp_url
except Exception as exc:
logger.warning("对口型输出视频转存 OSS 失败,回退临时 URL: job_id=%s err=%s", job_id, exc)
return temp_url
def _sign_media_url(self, url: str) -> str:
"""对自家 OSS 私有桶 URL 重签长有效期预签名."""
if not url:
return url
try:
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 # 外部临时链接原样透传
signed = storage.get_download_url(url, expires_seconds=MEDIAKIT_URL_TTL_SECONDS)
return signed or url
except Exception as exc:
logger.warning("对口型 URL 重签失败,原样返回: url_prefix=%s err=%s", url[:80], exc)
return url
# ── 取消任务 ──────────────────────────────────────────────────────────
def cancel_job(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
"""取消任务(仅 pending/tts_processing/submitted 状态可取消)."""
"""取消任务(仅 pending/submitted 状态可取消)."""
job = self.get_job(job_id, user_id)
if job is None:
return None
if job.status in ("pending", "tts_processing", "submitted"):
if job.status in ("pending", "submitted"):
job.status = "cancelled"
job.updated_at = datetime.now(UTC)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
+3 -2
View File
@@ -75,7 +75,7 @@ class MediaKitClient:
*,
video_url: str,
audio_url: str,
enable_video_loop: bool = True,
enable_video_loop: bool = False,
callback_url: Optional[str] = None,
callback_args: Optional[str] = None,
client_token: Optional[str] = None,
@@ -103,7 +103,8 @@ class MediaKitClient:
"video_url": video_url,
"audio_url": audio_url,
}
payload["enable_video_loop"] = bool(enable_video_loop)
if enable_video_loop:
payload["enable_video_loop"] = True
if callback_url:
payload["callback_url"] = callback_url
if callback_args:
-344
View File
@@ -1,344 +0,0 @@
"""叙事剪辑前置服务 — #1970 PR3.
叙事模式(assembly_mode='narrative')在生成任务入队前同步完成:
1. 按 script_id 读取文案(归属校验);
2. 按 tts_voice_source 解析音色(preset=CosyVoice 音色 idclone=克隆档案 id
解析档案归属并取其 CosyVoice voice_id);
3. 同步 TTS 合成(复用 tts_job 现有 workflow:提交即同步返回,未完成则轮询兜底),
失败直接抛 NarrativeErrorHTTP 层转 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
# presettts_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),
)
+26 -146
View File
@@ -14,7 +14,7 @@ from __future__ import annotations
import logging
import random
from typing import Any
from typing import Any, List
from sqlalchemy.orm import Session
@@ -22,11 +22,6 @@ 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
@@ -57,20 +52,18 @@ class PlanGeneratorService:
基于模板 + 素材,自动生成 EditPlan 及 EditPlanClip 列表。
"""
def __init__(self, db: Session, asset_repo=None, atom_clip_repo=None) -> None:
def __init__(self, db: Session, asset_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
# ── 公开接口 ─────────────────────────────────────────────────────────────
def generate_from_template(
self,
template: EditTemplate,
clip_configs: list[TemplateClipConfig],
asset_ids: list[str],
clip_configs: List[TemplateClipConfig],
asset_ids: List[str],
*,
project_id: str = "",
created_by_user_id: str = "",
@@ -128,37 +121,21 @@ class PlanGeneratorService:
# 4. 按 editing_mode 分配素材
if asset_ids:
# #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,
)
# 获取素材时长信息,用于随机起始时间
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] = []
created_clips: List[EditPlanClip] = []
total_duration = 0.0
for clip in clips:
saved = self._clip_repo.create(clip)
@@ -207,15 +184,15 @@ class PlanGeneratorService:
def _create_clips_from_configs(
self,
plan_id: str,
clip_configs: list[TemplateClipConfig],
) -> list[EditPlanClip]:
clip_configs: List[TemplateClipConfig],
) -> List[EditPlanClip]:
"""从 TemplateClipConfig 列表创建 EditPlanClip 列表(未持久化).
委托给 plan_generator_utils.create_clips_from_configs 纯函数。
"""
return create_clips_from_configs(plan_id, clip_configs)
def _map_clip_types_for_mode(self, clips: list[EditPlanClip], editing_mode: str) -> None:
def _map_clip_types_for_mode(self, clips: List[EditPlanClip], editing_mode: str) -> None:
"""将 MAIN 类型片段按 editing_mode 映射为对应角色类型.
委托给 plan_generator_utils.map_clip_types_for_mode 纯函数。
@@ -227,7 +204,7 @@ class PlanGeneratorService:
plan_id: str,
editing_mode: str,
asset_count: int,
) -> list[EditPlanClip]:
) -> List[EditPlanClip]:
"""无 clip_configs 时,根据 editing_mode 生成默认 clip 结构.
委托给 plan_generator_utils.generate_default_clips 纯函数。
@@ -236,8 +213,8 @@ class PlanGeneratorService:
def _distribute_assets(
self,
clips: list[EditPlanClip],
asset_ids: list[str],
clips: List[EditPlanClip],
asset_ids: List[str],
editing_mode: str,
*,
random_selection: bool = False,
@@ -282,104 +259,7 @@ 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]]:
def _fetch_asset_scene_points(self, asset_ids: List[str]) -> dict[str, list[float]]:
"""从素材 metadata 读取场景切换点缓存(无缓存的素材不包含在结果中)。"""
points_map: dict[str, list[float]] = {}
if not self._asset_repo:
@@ -392,7 +272,7 @@ class PlanGeneratorService:
points_map[asset_id] = points
return points_map
def _sort_assets_by_smart_score(self, asset_ids: list[str]) -> list[str]:
def _sort_assets_by_smart_score(self, asset_ids: List[str]) -> List[str]:
"""按 smart_match 综合评分降序排列素材 ID(注入随机噪声)。
评分高的素材(质量好、时长合适、新鲜、使用次数少)倾向排在前面;
@@ -415,7 +295,7 @@ class PlanGeneratorService:
)
return [aid for aid, _ in scored]
def _fetch_asset_durations(self, asset_ids: list[str]) -> dict[str, float]:
def _fetch_asset_durations(self, asset_ids: List[str]) -> dict[str, float]:
"""从数据库获取素材时长信息.
Args:
@@ -1,64 +0,0 @@
"""文案提取 ASR 服务封装 — Issue #1893.
将已有的 ASR 服务工厂封装为面向文案提取场景的简单接口:
- transcribe_to_text(video_path) -> str:将视频/音频转写为纯文本
- 未配置 ASR 时抛 ASRNotConfiguredError(路由层映射为 503
- ASR 调用失败时抛 ASRTranscriptionError(路由层映射为 502
"""
from __future__ import annotations
import logging
from pathlib import Path
from packages.ports.asr_service import ASRServiceError
logger = logging.getLogger(__name__)
class ASRNotConfiguredError(Exception):
"""ASR 服务未配置."""
class ASRTranscriptionError(Exception):
"""ASR 转写失败."""
def transcribe_to_text(media_path: str | Path) -> str:
"""将视频/音频文件转写为纯文本.
Args:
media_path: 媒体文件路径
Returns:
转写出的文本
Raises:
ASRNotConfiguredError: ASR 服务未配置
ASRTranscriptionError: ASR 调用失败
"""
# 延迟导入,避免循环依赖和启动时副作用
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:
raise ASRNotConfiguredError("ASR 服务未配置,请联系管理员配置火山 MediaKit 或阿里云 ASR 密钥")
try:
timeline = asr.transcribe(Path(media_path))
# 拼接所有分段的文本
text = "".join(seg.text for seg in timeline.segments)
return text.strip()
except ASRNotConfiguredError:
raise
except ASRServiceError as exc:
logger.error("ASR 转写失败: %s", exc)
raise ASRTranscriptionError(f"语音识别失败: {exc}") from exc
except Exception as exc:
logger.error("ASR 转写异常: %s", exc)
raise ASRTranscriptionError(f"语音识别失败: {exc}") from exc
+2 -2
View File
@@ -6,7 +6,7 @@
from __future__ import annotations
import uuid
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Optional
from sqlalchemy.orm import Session
@@ -93,7 +93,7 @@ class ScriptService:
script.segments = segments
if tags is not None:
script.tags = tags
script.updated_at = datetime.now(UTC)
script.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(script)
return script
-461
View File
@@ -1,461 +0,0 @@
"""AI 数字人对口型 TTS 异步任务 — 将 TTS 合成从 HTTP 请求移至 Celery 后台执行.
优化目标:将 create_job 的 API 响应时间从 6~35s 降到 <1s。
任务流程:
1. 创建新 DB session,加载 job 记录
2. 调用 CosyVoice 合成音频
3. 下载音频并转存到自家 OSS
4. 更新 job 的 audio_url
5. 签名 URL 并提交到 MediaKit
6. 更新 job 状态为 submitted
7. 异常时标记 job 为 failed
注意:使用 @shared_task 而非绑定到某个 celery_app 实例,
确保任务能被 Worker 侧 celery_app 正确注册,同时 API 侧 send_task/apply_async 仍可正常调用。
#1845:句子时间戳计算已提取至 packages/domain/sentence_timings.py,本模块保留
_ 开头别名兼容历史导入,但 _compute_sentence_timings/_split_script_into_sentences/
_estimate_sentence_timings_by_chars 等内部函数已复用共享实现,避免重复代码。
"""
import io
import logging
from datetime import UTC, datetime
from urllib.parse import urlparse
from celery import shared_task
# 复用共享的句子时间戳工具(#1845 配音前置)
from packages.domain.sentence_timings import compute_sentence_timings as _compute_sentence_timings
from packages.domain.sentence_timings import (
probe_audio_duration,
)
logger = logging.getLogger(__name__)
# MediaKit 预签名 URL 有效期(7天,秒),与 LipsyncService._sign_media_url 保持一致
_MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
def _sign_media_url(url: str) -> str:
"""对自家 OSS 私有桶 URL 重签长有效期预签名.
- 自家 OSS URL → 重签 7 天有效期
- 外部临时 URL → 原样透传
- 任何异常降级原样返回,不阻断主流程
"""
if not url:
return url
try:
from packages.shared.storage import get_shared_storage_service
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
signed = storage.get_download_url(url, expires_seconds=_MEDIAKIT_URL_TTL_SECONDS)
return signed or url
except Exception as exc: # noqa: BLE001
logger.warning("[lipsync_tts] URL 重签失败,原样返回: url_prefix=%s err=%s", url[:80], exc)
return url
@shared_task(
bind=True,
name="lipsync_tts.synthesize_and_submit",
max_retries=5, # 事务竞态重试3次(job not found+ TTS偶发错误2次
default_retry_delay=30,
autoretry_for=(OSError, ConnectionError), # 网络/连接错误自动重试
retry_backoff=True,
retry_backoff_max=30,
soft_time_limit=180,
time_limit=200,
)
def tts_synthesize_and_submit(
self,
job_id: str,
user_id: str,
voice_id: str,
script_text: str,
speed: float,
emotion: str,
):
"""异步执行 TTS 合成 + OSS 转存 + MediaKit 提交.
在 Celery worker 中运行,不阻塞 HTTP 请求。保留作为降级路径
(预合成失败 / 旧版前端未传 audio_url 时走此路径)。
"""
from app.services.mediakit_client import MediaKitError, get_mediakit_client
from sqlalchemy.orm import Session as DBSession
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.shared.url_security import safe_download_bytes
# SessionLocal 获取:
# - API 容器:app.db.SessionLocal(环境变量完整,导入即建引擎)
# - Worker 容器:worker_app.db.SessionLocalWorker 自己的 settings 初始化引擎)
# API 侧没有 worker_app 模块 → ImportError 直接回退;
# Worker 侧 app.db 会因缺少 API 专有环境变量抛 pydantic ValidationError
# 此时也要回退到 worker_app.db。
try:
from worker_app.db import SessionLocal # type: ignore
except Exception: # noqa: BLE001
from app.db import SessionLocal # type: ignore
db: DBSession = SessionLocal()
try:
job = (
db.query(LipsyncJobModel)
.filter(
LipsyncJobModel.id == job_id,
LipsyncJobModel.user_id == user_id,
)
.first()
)
if job is None:
# 事务竞态防御:API 在 commit 前投递了任务,worker 消费时事务尚未提交。
retries = getattr(self.request, "retries", 0)
max_retries = 3
if retries < max_retries:
backoff = (2**retries) + (retries * 1) # 1s, 3s, 7s
logger.warning(
"[lipsync_tts] Job not found yet (retry %d/%d, backoff %ds): job_id=%s",
retries + 1,
max_retries,
backoff,
job_id,
)
self.db.close()
raise self.retry(countdown=backoff, max_retries=max_retries)
logger.error(
"[lipsync_tts] Job not found after %d retries, giving up: job_id=%s",
max_retries,
job_id,
)
return
# 已取消的任务不再处理
if job.status == "cancelled":
logger.info("[lipsync_tts] Job already cancelled, skipping: job_id=%s", job_id)
return
# 1. TTS 合成
logger.info(
"[lipsync_tts] 开始 TTS 合成: job_id=%s voice_id=%s text_len=%d speed=%.2f",
job_id,
voice_id,
len(script_text),
speed,
)
try:
cosyvoice = CosyVoiceService()
result = cosyvoice.submit_synthesize_task(
text=script_text,
voice_id=voice_id,
speed=speed,
emotion=emotion,
language="zh",
)
except CosyVoiceError as exc:
logger.error("[lipsync_tts] TTS 合成失败: job_id=%s err=%s", job_id, exc)
job.status = "failed"
job.error_message = f"TTS 合成失败: {exc}"
job.error_code = "TTSSynthesisFailed"
job.updated_at = datetime.now(UTC)
db.commit()
return
except ValueError as exc:
logger.error("[lipsync_tts] TTS 参数错误: job_id=%s err=%s", job_id, exc)
job.status = "failed"
job.error_message = f"TTS 参数错误: {exc}"
job.error_code = "TTSInvalidParam"
job.updated_at = datetime.now(UTC)
db.commit()
return
temp_url = result.get("audio_url", "")
if not temp_url:
logger.error("[lipsync_tts] TTS 未返回音频 URL: job_id=%s", job_id)
job.status = "failed"
job.error_message = "TTS 未返回音频 URL"
job.error_code = "TTSNoAudio"
job.updated_at = datetime.now(UTC)
db.commit()
return
# 2. 下载 TTS 音频到内存(用于 2.5 静音检测;不转存自家 OSS,直接使用 CosyVoice 临时 URL
audio_data: bytes | None = None
try:
audio_data = safe_download_bytes(
temp_url,
purpose="lipsync_tts_audio",
allowed_mime_types={
"audio/mpeg",
"audio/mp3",
"audio/wav",
"audio/x-wav", # CosyVoice 部分接口返回 audio/x-wav
"audio/mp4",
"audio/x-m4a",
},
timeout=60.0,
)
logger.info(
"[lipsync_tts] TTS 音频已下载到内存: job_id=%s size=%d",
job_id,
len(audio_data) if audio_data else 0,
)
except Exception as exc:
logger.warning(
"[lipsync_tts] TTS 音频下载失败,跳过静音检测,直接使用临时 URL 提交: job_id=%s err=%s",
job_id,
exc,
)
# TTS 音频使用 CosyVoice 临时 URL,跳过自家 OSS 转存(加速,步骤⑥)
job.audio_url = temp_url
logger.info("[lipsync_tts] TTS 音频使用 CosyVoice 临时 URL(跳过 OSS 转存): job_id=%s", job_id)
db.commit()
# 2.5 计算精确句子时间戳(基于 TTS 音频静音检测)—— 复用共享工具
try:
if not audio_data:
logger.warning("[lipsync_tts] 无音频数据,跳过句子时间戳计算: job_id=%s", job_id)
else:
_audio_duration = probe_audio_duration(audio_data)
logger.info(
"[lipsync_tts] 音频时长探测: job_id=%s duration=%.2f",
job_id,
_audio_duration,
)
if _audio_duration > 0:
_timings = _compute_sentence_timings(audio_data, script_text, _audio_duration)
if _timings:
job.sentence_timings = _timings
logger.info(
"[lipsync_tts] 句子时间戳已计算: job_id=%s sentences=%d duration=%.1f",
job_id,
len(_timings),
_audio_duration,
)
else:
logger.warning("[lipsync_tts] 句子时间戳计算返回空结果: job_id=%s", job_id)
else:
logger.warning(
"[lipsync_tts] ffprobe 未获取到有效时长,跳过句子时间戳: job_id=%s",
job_id,
)
db.commit()
except Exception as _st_err:
logger.warning(
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
)
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
audio_url = _sign_media_url(job.audio_url)
video_url = _sign_media_url(job.video_url)
client = get_mediakit_client()
try:
mk_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 = mk_result["task_id"]
job.status = "submitted"
job.submitted_at = datetime.now(UTC)
logger.info(
"[lipsync_tts] 已提交 MediaKit: job_id=%s task_id=%s",
job_id,
mk_result["task_id"],
)
except MediaKitError as exc:
job.status = "failed"
job.error_message = str(exc)
job.error_code = exc.code
logger.error("[lipsync_tts] 提交 MediaKit 失败: job_id=%s err=%s", job_id, exc)
# 三层防御 ③:链式触发 Celery 兜底轮询——MediaKit 提交成功后由 worker
# 主动拉取状态到终态,不依赖前端轮询触发的 FastAPI background task
# background task 可能静默失败导致永久卡 running)。
if job.status == "submitted" and job.mediakit_task_id:
try:
poll_mediakit_status.apply_async(
kwargs={"job_id": job_id, "user_id": user_id},
countdown=10, # 10 秒后开始轮询,给 MediaKit 一点处理时间
)
except Exception as exc: # noqa: BLE001
logger.warning("[lipsync_tts] 提交兜底轮询任务失败(不影响主流程): job_id=%s err=%s", job_id, exc)
db.commit()
except Exception:
logger.exception("[lipsync_tts] 未预期的异常: job_id=%s", job_id)
try:
job = db.query(LipsyncJobModel).filter(LipsyncJobModel.id == job_id).first()
if job and job.status not in ("cancelled", "failed", "completed"):
job.status = "failed"
job.error_message = "TTS 异步任务执行异常"
job.error_code = "AsyncTaskError"
job.updated_at = datetime.now(UTC)
db.commit()
except Exception:
logger.exception("[lipsync_tts] 回写失败状态时异常: job_id=%s", job_id)
finally:
db.close()
@shared_task(
bind=True,
name="lipsync_tts.poll_mediakit_status",
max_retries=60, # 最多轮询 60 次
default_retry_delay=10, # 每次间隔 10 秒(总兜底时长 10 分钟)
)
def poll_mediakit_status(self, job_id: str, user_id: str):
"""Celery 兜底轮询:TTS 提交 MediaKit 后,由 worker 主动拉取状态直到终态。
不依赖前端轮询,避免 background task 静默失败导致任务永久卡 running/submitted。
"""
from sqlalchemy.orm import Session as DBSession
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
try:
from worker_app.db import SessionLocal # type: ignore
except Exception: # noqa: BLE001
from app.db import SessionLocal # type: ignore
db: DBSession = SessionLocal()
try:
job = db.query(LipsyncJobModel).filter(LipsyncJobModel.id == job_id, LipsyncJobModel.user_id == user_id).first()
if job is None:
logger.warning("[lipsync_poll] Job not found: job_id=%s", job_id)
return
# 已终态,不需要再轮询
if job.status in ("completed", "failed", "cancelled"):
return
if not job.mediakit_task_id:
logger.warning("[lipsync_poll] Job has no mediakit_task_id: job_id=%s status=%s", job_id, job.status)
return
from app.services.lipsync_service import STATUS_COMPLETED as _SC
from app.services.lipsync_service import STATUS_FAILED as _SF
from app.services.lipsync_service import LipsyncService
from app.services.mediakit_client import MediaKitError, get_mediakit_client
client = get_mediakit_client()
try:
status_data = client.get_task_status(job.mediakit_task_id)
except MediaKitError as exc:
logger.warning("[lipsync_poll] 拉取 MediaKit 状态失败,将重试: job_id=%s err=%s", job_id, exc)
raise self.retry(exc=exc) from exc
mk_status = status_data.get("status", "running")
if mk_status in ("succeeded", _SC):
svc = LipsyncService(db)
result = status_data.get("result", {})
job.status = "completed"
output_url = result.get("video_url", "")
try:
job.output_video_url = svc._persist_output_video(output_url, job_id, user_id)
except Exception as exc: # noqa: BLE001
logger.warning("[lipsync_poll] 转存 OSS 失败,保留临时 URL: job_id=%s err=%s", job_id, exc)
job.output_video_url = output_url
job.output_duration = result.get("duration", 0.0)
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
db.commit()
logger.info("[lipsync_poll] 任务完成: job_id=%s", job_id)
elif mk_status in ("failed", "error", _SF):
error = status_data.get("error", {})
job.status = "failed"
job.error_message = error.get("message", "任务执行失败")
job.error_code = error.get("code", "TaskFailed")
job.completed_at = datetime.now(UTC)
job.updated_at = datetime.now(UTC)
db.commit()
logger.info("[lipsync_poll] 任务失败: job_id=%s err=%s", job_id, job.error_message)
else:
# 中间状态,更新时间戳,继续重试
job.updated_at = datetime.now(UTC)
if isinstance(mk_status, str) and mk_status:
job.status = mk_status
db.commit()
logger.debug("[lipsync_poll] 任务仍在 %s,继续轮询: job_id=%s", mk_status, job_id)
raise self.retry()
except Exception as exc:
logger.exception("[lipsync_poll] 未预期异常: job_id=%s", job_id)
try:
db.rollback()
except Exception:
pass
raise self.retry(exc=exc) from exc
finally:
db.close()
@shared_task(
name="lipsync_tts.persist_output_video",
max_retries=2,
default_retry_delay=30,
)
def persist_output_video_task(job_id: str, user_id: str, temp_url: str):
"""异步转存对口型输出视频到自家 OSS(步骤⑦ — 将同步阻塞挪到后台,加速前端响应)."""
try:
from worker_app.db import SessionLocal # type: ignore
except Exception: # noqa: BLE001
from app.db import SessionLocal # type: ignore
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
from packages.shared.storage import get_shared_storage_service
db = SessionLocal()
try:
job = db.query(LipsyncJobModel).filter(LipsyncJobModel.id == job_id, LipsyncJobModel.user_id == user_id).first()
if job is None:
logger.error("[lipsync_tts.persist] Job not found: job_id=%s", job_id)
return
if not temp_url:
logger.warning("[lipsync_tts.persist] temp_url 为空,跳过转存: job_id=%s", job_id)
return
try:
import httpx
with httpx.Client(timeout=180.0, follow_redirects=True) as client:
resp = client.get(temp_url)
resp.raise_for_status()
data = resp.content
storage = get_shared_storage_service()
storage_key = f"lipsync-outputs/{user_id}/{job_id}.mp4"
permanent_url = storage.upload_file(io.BytesIO(data), storage_key, content_type="video/mp4")
final_url = _sign_media_url(permanent_url) if permanent_url else temp_url
job.output_video_url = final_url
job.updated_at = datetime.now(UTC)
db.commit()
logger.info("[lipsync_tts.persist] 输出视频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
except Exception as exc:
logger.warning(
"[lipsync_tts.persist] 输出视频转存失败,保留临时 URL: job_id=%s err=%s",
job_id,
exc,
)
except Exception:
logger.exception("[lipsync_tts.persist] 未预期异常: job_id=%s", job_id)
finally:
db.close()
-117
View File
@@ -1,117 +0,0 @@
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)
})
})
+268 -323
View File
@@ -1,4 +1,4 @@
import { expect, test, type APIRequestContext, type Page } from "@playwright/test"
import { expect, test, type APIRequestContext } from "@playwright/test"
import * as fs from "node:fs"
import * as path from "node:path"
import { fileURLToPath } from "node:url"
@@ -8,8 +8,7 @@ const PASSWORD = "SmokePass123!"
const apiBase = process.env.E2E_API_BASE || "/api/v1"
const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".length) : ""
/** 将浏览器侧 /api/v1 请求路由到 Playwright request 源(支持跨域) */
async function routeBrowserApiToTestApi(page: Page) {
const routeBrowserApiToTestApi = async (page: import("@playwright/test").Page) => {
if (!apiOrigin) return
await page.route("**/api/v1/**", async (route) => {
const sourceUrl = new URL(route.request().url())
@@ -25,358 +24,304 @@ async function loginWithRetry(
email: string,
password: string,
maxRetries = 2,
): Promise<string> {
) {
for (let i = 0; i <= maxRetries; i++) {
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`)
const response = await request.post(`${apiBase}/auth/login`, {
data: { email, password },
})
if (response.status() !== 429) return response
console.log(`[login] 触发限流,等待 65s 后重试 (${i + 1}/${maxRetries})`)
await new Promise((r) => setTimeout(r, 65000))
}
throw new Error("Login failed after retries")
return request.post(`${apiBase}/auth/login`, {
data: { email, password },
})
}
/**
* 注册新用户 + 建项目/视频库/上传 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 TemplateResponse = { id: string }
type AssetListResponse = {
items: Array<{
id: string
name: string
status: string
}>
}
/**
* #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}` }
test.describe("Core generation flow", () => {
test.describe.configure({ timeout: 360_000 })
// 确保默认模板存在(智能剪辑页依赖模板)
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)
test("walks through 6-step wizard and starts generation", async ({ page, request }) => {
test.setTimeout(360_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 } }),
)
}, 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}`
// ── 提前 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,
// 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] },
)
.toBe("ready")
// Create an editing template so the generate page has at least one template
// (templates are now loaded from API; new users have none by default)
const template = await request.post(`${apiBase}/templates`, {
headers,
data: {
name: `E2E 测试模板 ${suffix}`,
mode: "pip",
estimated_duration: 30,
segments: [
{
segment_order: 1,
duration_min: 5,
duration_max: 30,
material_type: "video",
},
],
tags: ["e2e"],
},
})
expect(template.status(), await template.text()).toBe(201)
const templateData = (await template.json()) as TemplateResponse
expect(templateData.id).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,
}),
}),
)
},
{
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: 30000,
timeout: 20_000,
})
// ── Step 1:默认随机混剪选中,点下一步 ──────────────────────────
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
await expect(page.getByText("随机混剪")).toBeVisible()
await page.getByRole("button", { name: /下一步/ }).click()
// Step 1: template - default selected, click next
await expect(page.locator(".xx-choice-item.selected")).toBeVisible()
await page.getByRole("button", { name: "下一步" }).click()
// ── 配音选择弹窗:选第一个配音 → 确认 ─────────────────────────
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 })
// Step1 下一步弹出数量选择弹窗(Issue #1677 固定6步:模板→素材→配音→标题→确认生成→封面)
// 单视频流程:默认 1 个,点击「生成 1 个视频」进入步骤2
await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({
timeout: 10_000,
})
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// ── Step 3:填写标题 ──────────────────────────────────────────
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
const titleInput = page.getByPlaceholder("输入或从标题库选择")
// Step 2: 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 3: voice (可选步骤,新注册用户无配音素材,直接跳过)
await expect(page.getByRole("heading", { name: /选择配音/ })).toBeVisible({ timeout: 15000 })
await page.getByRole("button", { name: "下一步" }).click()
// Step 4: 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")
await expect(titleInput).toBeVisible({ timeout: 5000 })
await titleInput.fill(`测试随机剪辑 ${suffix}`)
await page.getByRole("button", { name: /下一步/ }).click()
// ── 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 })
const titleText = `E2E Test ${suffix}`
await titleInput.fill(titleText)
const createTask = page.waitForResponse(
(r) => r.url().includes("/generation/tasks") && r.request().method() === "POST",
{ timeout: 30000 },
// Step 4(标题+实时预览):确认生成按钮已移到标题页,点击直接创建最终渲染任务
// 等待前端实时预览就绪:未就绪时右侧 FrontendPreviewPlayer 显示「准备预览素材...」占位,
// 就绪(previewReady:素材已解析 + 模板已选中)后占位消失;否则按钮会被校验拦截弹 warning
await page
.getByText("准备预览素材")
.waitFor({ state: "detached", timeout: 30_000 })
.catch(() => {})
// Wait for generation API to be called
// 前端直接创建生成任务:POST /generation/tasks
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: 30_000 },
)
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 ✓")
// 点击「确认生成视频」
await page.locator(".xx-btn-primary").filter({ hasText: "确认生成视频" }).first().click()
// 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()
// 单视频(N=1):点击「确认生成视频」后跳 Step 5「确认生成」,展示实时渲染进度
await expect(page.getByRole("heading", { name: "🎬 确认生成" })).toBeVisible({
timeout: 30_000,
})
// 等待渲染完成:进度卡变为「视频生成完成」(最长等待 3 分钟)
await expect(page.getByText("视频生成完成")).toBeVisible({ timeout: 180_000 })
// 全部完成后「下一步:选择封面」解锁,点击进入 Step 6
await page.getByRole("button", { name: /下一步:选择封面/ }).click()
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({
timeout: 30_000,
})
} 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,
})
// 清理所有路由,避免页面关闭时飞地API请求导致测试报错
await page.unrouteAll({ behavior: "ignoreErrors" })
})
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")
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 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 文案列表、音色、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,
}),
})
const register = await request.post(`${apiBase}/auth/register`, {
data: { email, username, password: PASSWORD, display_name: username },
})
expect(register.status()).toBe(201)
// 预设音色(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,
}),
}),
)
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}` }
// 克隆音色:空列表
await page.route(
(url) => url.pathname.endsWith("/voice-clones"),
(route) =>
route.fulfill({
status: 200,
contentType: "application/json",
body: JSON.stringify({ items: [] }),
}),
)
// 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合成配音" }),
}),
)
await page.goto("/app/generate")
await expect(page.getByRole("heading", { name: "智能剪辑" })).toBeVisible({
timeout: 30000,
const project = await request.post(`${apiBase}/projects`, {
headers,
data: { name: `E2E API Proj ${suffix}` },
})
expect(project.status()).toBe(200)
// ── Step 1:切到叙事剪辑 → 下一步 ────────────────────────────
await expect(page.getByText("选择模式", { exact: true })).toBeVisible()
await page.getByText("叙事剪辑").click()
await page.getByRole("button", { name: /下一步/ }).click()
// ── 文案选择弹窗:选第一条 → 确认 ─────────────────────────────
await expect(page.getByText("📝 选择文案")).toBeVisible({ timeout: 5000 })
await page.getByText("测试带货文案").first().click()
await page.getByRole("button", { name: "确认选择" }).click()
await expect(page.getByText("📝 选择文案")).not.toBeVisible()
// ── 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 })
// ── 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()
// ── 数量弹窗 ─────────────────────────────────────────────────
await expect(page.getByText("要生成几个视频?")).toBeVisible({ timeout: 5000 })
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// ── 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 ✓")
// 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)
})
})
-105
View File
@@ -1,105 +0,0 @@
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 ✓`)
})
}
})
@@ -26,10 +26,6 @@ export interface BatchVariantPlansRequest {
count: number
/** 源剪辑计划 ID:优先取预览/草稿关联的 plan;不传由后端按 template_id+user 兜底最新 plan */
source_edit_plan_id?: string
/** 统一配音 ID(共用配音模式);独立配音模式不传,改传 voice_library_ids */
voice_library_id?: string
/** 独立配音 ID 列表(长度=count,按变体序号一一对应);共用配音模式不传 */
voice_library_ids?: string[]
}
/** 单个变体的计划片段 */
@@ -40,8 +36,6 @@ export interface VariantPlan {
plan_id: string
/** 该变体的真实片段(顺序/素材/起点与正式成片一致) */
clips: EditPlanClip[]
/** 该变体实际配音时长(秒),用于前端预览按配音时长对齐音画;后端暂未返回时缺省 */
voice_duration?: number
}
/** 批量变体计划响应 */
-345
View File
@@ -1,345 +0,0 @@
/**
* 积分系统 API 封装
* 对齐后端 staging 实测最终契约(2026-09-16
*
* 当前 POINTS_API_MOCK=true:使用 MOCK_* 常量 + setTimeout 模拟延迟,
* 等后端 P0(支付通道接入、change-plan 校验)稳定后切 false 联调。
*
* 会员/订阅 API 在 @/api/subscription 中定义,避免重复封装。
*/
import apiClient from "../client"
import type {
PointsBalance,
PointsRulesResponse,
PointsPackagesResponse,
PointsTransaction,
PointsTransactionsResponse,
PointsCheckRequest,
PointsCheckResponse,
CreateRechargeOrderRequest,
CreateRechargeOrderResponse,
DailyUsage,
MembershipResponse,
} from "./types"
/** 模拟网络延迟(ms */
const MOCK_DELAY = 500
/* ================================================================
* Mock 数据
* ================================================================ */
/** mock 余额(无 free_clips_* 字段,已拆分到 dailyUsage */
const MOCK_BALANCE: PointsBalance = {
balance: 258,
total_earned: 500,
total_spent: 242,
is_member: false,
member_type: null,
member_expires_at: null,
}
const MOCK_RULES: PointsRulesResponse = {
rules: [
{
scene_key: "ai_voice",
name: "AI 配音",
base_points: 2,
unit: "次",
description: "单次配音消耗 2 积分,超 30 秒每 30 秒 +1 积分",
extra_per_30s: 1,
},
{
scene_key: "ai_video",
name: "AI 视频生成",
base_points: 8,
unit: "条",
description: "单条视频 8 积分起,按视频时长加收",
extra_per_30s: 3,
},
{
scene_key: "ai_digital_human",
name: "AI 数字人",
base_points: 15,
unit: "次",
description: "数字人生成 15 积分起",
extra_per_30s: 5,
},
{
scene_key: "voice_clone_train",
name: "声音克隆训练",
base_points: 20,
unit: "次",
description: "声音模型训练一次性消耗 20 积分",
},
{
scene_key: "voice_clone_synth",
name: "声音克隆合成",
base_points: 3,
unit: "次",
description: "使用克隆声音合成音频每次 3 积分",
},
{
scene_key: "douyin_extract",
name: "抖音文案提取",
base_points: 1,
unit: "次",
description: "提取抖音视频文案每次 1 积分",
},
{
scene_key: "ai_rewrite",
name: "AI 文案改写",
base_points: 2,
unit: "次",
description: "AI 改写文案每次 2 积分",
},
{
scene_key: "ai_title",
name: "AI 标题生成",
base_points: 1,
unit: "次",
description: "AI 生成标题每次 1 积分,一次生成多条",
},
{
scene_key: "ai_cover",
name: "AI 封面生成",
base_points: 3,
unit: "次",
description: "AI 生成封面每次 3 积分",
},
],
free_user_multiplier: 1.15,
}
const MOCK_PACKAGES: PointsPackagesResponse = {
packages: [
{ code: "points_100", name: "100 积分", points: 100, price_cents: 990, unit_price: 0.099 },
{ code: "points_500", name: "500 积分", points: 500, price_cents: 4490, unit_price: 0.0898 },
{ code: "points_1000", name: "1000 积分", points: 1000, price_cents: 7990, unit_price: 0.0799 },
{
code: "points_3000",
name: "3000 积分",
points: 3000,
price_cents: 19900,
unit_price: 0.0663,
},
],
user_discount: null,
}
const MOCK_TRANSACTIONS: PointsTransaction[] = [
{
id: 1,
type: "deduct",
source: "ai_video",
amount: 10,
balance_after: 248,
description: "AI 视频生成 ×1(非会员倍率)",
ref_id: "task_abc123",
created_at: "2026-09-16T08:30:00Z",
},
{
id: 2,
type: "add",
source: "recharge",
amount: 100,
balance_after: 258,
description: "充值 100 积分",
ref_id: "order_xyz789",
created_at: "2026-09-15T14:20:00Z",
},
{
id: 3,
type: "deduct",
source: "ai_voice",
amount: 3,
balance_after: 158,
description: "AI 配音 ×145s 加收)",
ref_id: "",
created_at: "2026-09-15T10:15:00Z",
},
{
id: 4,
type: "add",
source: "sign_up",
amount: 60,
balance_after: 161,
description: "新用户注册赠送",
ref_id: "",
created_at: "2026-09-10T09:00:00Z",
},
{
id: 5,
type: "deduct",
source: "ai_title",
amount: 1,
balance_after: 101,
description: "AI 标题生成 ×1",
ref_id: "",
created_at: "2026-09-14T16:45:00Z",
},
]
const MOCK_DAILY_USAGE: DailyUsage = {
free_clips_used: 1,
free_clips_limit: 3,
free_clips_remaining: 2,
reset_at: new Date(Date.now() + 8 * 3600_000).toISOString(),
}
const MOCK_MEMBERSHIP: MembershipResponse = {
is_member: false,
member_type: null,
member_expires_at: null,
points_balance: 258,
max_resolution: "720p",
}
/* ================================================================
* 积分 API
* ================================================================ */
/** 获取积分余额 */
export async function getPointsBalance(): Promise<PointsBalance> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
return { ...MOCK_BALANCE }
}
const { data } = await apiClient.get(`/points/balance`)
return data
}
/** 获取积分消耗规则 */
export async function getPointsRules(): Promise<PointsRulesResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
return { rules: [...MOCK_RULES.rules], free_user_multiplier: MOCK_RULES.free_user_multiplier }
}
const { data } = await apiClient.get(`/points/rules`)
return data
}
/** 获取充值包列表 */
export async function getPointsPackages(): Promise<PointsPackagesResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
return { packages: MOCK_PACKAGES.packages.map((p) => ({ ...p })), user_discount: null }
}
const { data } = await apiClient.get(`/points/packages`)
return data
}
/**
* 获取积分流水(分页)
*/
export async function getPointsTransactions(
page = 1,
pageSize = 20,
): Promise<PointsTransactionsResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
const start = (page - 1) * pageSize
const items = MOCK_TRANSACTIONS.slice(start, start + pageSize)
return {
items: items.map((t) => ({ ...t })),
total: MOCK_TRANSACTIONS.length,
page,
page_size: pageSize,
}
}
const { data } = await apiClient.get(`/points/transactions`, {
params: { page, page_size: pageSize },
})
return data
}
/**
* 创建充值订单
* 注意:当前 pay_params 返回空对象 {}(支付通道未接入),
* 前端可以完成订单创建 UI,但无法发起真实支付,待后续支付通道接入后联调。
*/
export async function createPointsOrder(
data: CreateRechargeOrderRequest,
): Promise<CreateRechargeOrderResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY * 2))
const pkg = MOCK_PACKAGES.packages.find((p) => p.code === data.package_id)
if (!pkg) throw new Error("充值包不存在")
return {
id: `mock_order_${Date.now()}`,
order_type: "points_recharge",
product_code: pkg.code,
amount_cents: pkg.price_cents,
points_amount: pkg.points,
status: "pending",
pay_params: {},
expire_at: new Date(Date.now() + 30 * 60_000).toISOString(),
created_at: new Date().toISOString(),
}
}
const { data: d } = await apiClient.post(`/points/recharge`, data)
return d
}
/**
* 积分预检查(消耗前调用)
*/
export async function checkPoints(data: PointsCheckRequest): Promise<PointsCheckResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
const rule = MOCK_RULES.rules.find((r) => r.scene_key === data.scene_key)
if (!rule) {
throw {
error: {
code: 400,
message: `未知场景:${data.scene_key}`,
valid_scenes: MOCK_RULES.rules.map((r) => r.scene_key),
},
}
}
const durationExtra =
data.duration_minutes && data.duration_minutes > 0.5 && rule.extra_per_30s
? Math.ceil((data.duration_minutes * 60 - 30) / 30) * rule.extra_per_30s
: 0
const base = (rule.base_points + durationExtra) * data.quantity
const balance = MOCK_BALANCE.balance
const multiplier = MOCK_BALANCE.is_member ? 1 : MOCK_RULES.free_user_multiplier
const required = Math.ceil(base * multiplier)
// 免费额度抵扣
const isFreeQuota = !MOCK_BALANCE.is_member && MOCK_DAILY_USAGE.free_clips_remaining > 0
const finalRequired = isFreeQuota ? 0 : required
return {
allowed: balance >= finalRequired,
required_points: finalRequired,
current_balance: balance,
remaining_after: balance - finalRequired,
is_free_quota: isFreeQuota,
}
}
const { data: d2 } = await apiClient.post(`/points/check`, data)
return d2
}
/* ================================================================
* 每日免费额度 + 会员聚合信息(新接口)
* ================================================================ */
/** 获取每日免费额度使用情况 */
export async function getDailyUsage(): Promise<DailyUsage> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
return { ...MOCK_DAILY_USAGE }
}
const { data } = await apiClient.get(`/usage/daily`)
return data
}
/** 获取会员聚合信息(创作页可用来判断 max_resolution */
export async function getMembership(): Promise<MembershipResponse> {
if (process.env.POINTS_API_MOCK === "true") {
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
return { ...MOCK_MEMBERSHIP }
}
const { data } = await apiClient.get(`/points/subscription/membership`)
return data
}
-227
View File
@@ -1,227 +0,0 @@
/**
* 积分系统类型定义
* 对齐后端 staging 实测最终契约(2026-09-16
*
* Base path: /api/v1/
* 会员/订阅相关类型请从 @/api/subscription/types 引入,本文件仅保留积分核心类型。
*/
/* ================================================================
* 场景键
* ================================================================ */
/**
* 积分消耗场景键(9 个)
* - ai_script 已拆分为 douyin_extract / ai_rewrite / ai_title,前端禁止再传 ai_script
*/
export type PointsSource =
| "ai_voice" // AI 配音
| "ai_video" // AI 视频生成
| "ai_digital_human" // AI 数字人
| "voice_clone_train" // 声音克隆训练
| "voice_clone_synth" // 声音克隆合成
| "douyin_extract" // 抖音提取文案
| "ai_rewrite" // AI 文案改写
| "ai_title" // AI 标题生成
| "ai_cover" // AI 封面生成
/** 非消耗场景 source 前缀(用于流水 source 字段) */
export type PointsSourceExtra =
PointsSource | `refund:${string}` | "recharge" | "sign_up" | "bind_phone" | "gift" | "admin"
/* ================================================================
* 通用
* ================================================================ */
/** ISO 8601 时间字符串 */
export type ISODate = string
/* ================================================================
* 积分余额(GET /points/balance
* ================================================================ */
export interface PointsBalance {
/** 当前可用积分 */
balance: number
/** 累计获得积分 */
total_earned: number
/** 累计消耗积分 */
total_spent: number
/** 是否为付费会员 */
is_member: boolean
/** 会员类型(monthly/quarterly/yearly,非会员 null)。推荐使用 /subscription/current 的 plan_id+billing_cycle 做判断 */
member_type: "monthly" | "quarterly" | "yearly" | null
/** 会员到期时间 */
member_expires_at: ISODate | null
}
/* ================================================================
* 积分规则(GET /points/rules
* ================================================================ */
export interface PointsRule {
scene_key: PointsSource
/** 场景中文名 */
name: string
/** 基准消耗积分(points_per_use 改名) */
base_points: number
/** 单位描述,如「次」「分钟」「个」 */
unit: string
/** 超过30秒后每30秒额外积分(视频/语音类) */
extra_per_30s?: number
/** 场景说明(后端已补回) */
description?: string
}
export interface PointsRulesResponse {
rules: PointsRule[]
/** 非会员消耗倍率(如 1.15) */
free_user_multiplier: number
}
/* ================================================================
* 充值包(GET /points/packages
* ================================================================ */
export interface PointsPackage {
/** 包编码(id 改名) */
code: string
name: string
points: number
/** 原价,单位分 */
price_cents: number
/** 每积分单价(元),展示用 */
unit_price: number
}
export interface PointsPackagesResponse {
packages: PointsPackage[]
/** 当前用户折扣(会员折扣或活动折扣),null 表示无折扣 */
user_discount: number | null
}
/**
* 充值包前端展示辅助:折后价(分)
* 后端废弃 4 档 discounted_price_for_*,前端按 price_cents * (user_discount ?? 1) 计算。
*/
export function getDiscountPriceCents(pkg: PointsPackage, userDiscount: number | null): number {
return Math.round(pkg.price_cents * (userDiscount ?? 1))
}
/* ================================================================
* 积分流水(GET /points/transactions
* ================================================================ */
export type PointsTxType = "add" | "deduct"
export interface PointsTransaction {
id: number
/** 流水类型:add=获得/退款,deduct=消耗 */
type: PointsTxType
/**
* 消耗/获得来源:
* - 消耗场景直接用 PointsSource 值
* - 充值/退款/赠送使用 recharge / refund:<source> / sign_up / bind_phone / gift / admin
*/
source: string
/** 变动数量(绝对值,正负由 type 决定) */
amount: number
/** 变动后余额 */
balance_after: number
/** 中文描述 */
description: string
/** 关联订单/任务 ID,空字符串 "" 表示无关联(不是 null */
ref_id: string
created_at: ISODate
}
export interface PointsTransactionsResponse {
items: PointsTransaction[]
total: number
page: number
page_size: number
}
/* ================================================================
* 创建充值订单(POST /points/recharge
* ================================================================ */
export interface CreateRechargeOrderRequest {
/** 充值包 code(字段名保留 package_id 与后端一致) */
package_id: string
}
export interface CreateRechargeOrderResponse {
id: string
order_type: string
product_code: string
/** 订单金额(分) */
amount_cents: number
/** 充值积分数量 */
points_amount: number
status: string
/**
* 支付参数(支付通道未接入时返回空对象 {},前端可透传)
*/
pay_params: Record<string, unknown>
/** 订单过期时间 */
expire_at: ISODate
created_at: ISODate
}
/* ================================================================
* 积分预检查(POST /points/check
* ================================================================ */
export interface PointsCheckRequest {
scene_key: PointsSource
/** 数量(units 改名) */
quantity: number
/** 预计时长(分钟),可选 */
duration_minutes?: number
}
export interface PointsCheckResponse {
/** 是否可以执行 */
allowed: boolean
/** 需要消耗积分 */
required_points: number
/** 当前余额 */
current_balance: number
/** 扣除后剩余 */
remaining_after: number
/** 是否走免费额度 */
is_free_quota: boolean
}
/* ================================================================
* 每日使用情况(GET /usage/daily,新接口)
* ================================================================ */
export interface DailyUsage {
/** 今日已用免费次数 */
free_clips_used: number
/** 每日免费次数上限 */
free_clips_limit: number
/** 今日剩余免费次数 */
free_clips_remaining: number
/** 额度重置时间 */
reset_at: ISODate
}
/* ================================================================
* 会员聚合信息(GET /points/subscription/membership,新接口)
* ================================================================ */
export interface MembershipResponse {
is_member: boolean
/** 会员类型(monthly/quarterly/yearly,非会员 null */
member_type: "monthly" | "quarterly" | "yearly" | null
member_expires_at: ISODate | null
/** 当前积分余额(冗余,可与 balance 互校) */
points_balance: number
/** 最大分辨率,如 "720p" / "1080p" / "4k" */
max_resolution: string
}
/* ================================================================
* 错误响应(统一格式 {error:{code,message}}
* ================================================================ */
export interface ApiError {
error: {
code: number
message: string
/** 部分场景会返回,如 unknown scene_key */
valid_scenes?: PointsSource[]
}
}
+11 -34
View File
@@ -1,6 +1,6 @@
/**
* 成品 / 视频相关 API 函数
* 后端实际接口:/videos(分页:page/page_size,返回 {items, total, page, page_size}
* 后端实际接口:/videos
*/
import apiClient from "../client"
import type {
@@ -12,39 +12,16 @@ import type {
} from "./types"
import { mapVideoToProductItem } from "./utils"
/** 分页列表响应(前端消费用 */
export interface ProductListResult {
items: ProductItem[]
total: number
page: number
page_size: number
}
/**
* 获取成品列表(分页)
* @param params 分页与筛选参数:page 默认 1page_size 默认 20
*/
export const getProducts = async (params?: ProductListParams): Promise<ProductListResult> => {
const response = await apiClient.get("/videos", {
params: {
page: 1,
page_size: 20,
...params,
},
})
const data = response.data as {
items?: VideoItem[]
total?: number
page?: number
page_size?: number
}
const items: VideoItem[] = Array.isArray(data?.items) ? data.items : []
return {
items: items.map(mapVideoToProductItem),
total: data.total ?? items.length,
page: data.page ?? params?.page ?? 1,
page_size: data.page_size ?? params?.page_size ?? 20,
}
/** 获取成品列表(支持分页和筛选 */
export const getProducts = async (params?: ProductListParams): Promise<ProductItem[]> => {
const response = await apiClient.get("/videos", { params })
const data = response.data
const videos: VideoItem[] = Array.isArray(data?.items)
? data.items
: Array.isArray(data)
? data
: []
return videos.map(mapVideoToProductItem)
}
/** 获取单个成品详情 */
-1
View File
@@ -1,3 +1,2 @@
export * from "./scripts"
export * from "./types"
export * from "./scripts-ai"
-87
View File
@@ -1,87 +0,0 @@
/**
* 文案库 AI 能力 API#1893
* 三个端点均走真实后端,不参与 SCRIPTS_API_MOCK 开关。
*/
import apiClient from "../client"
/** ── 1. 从抖音视频提取文案(下载 + ASR) */
export interface ExtractFromDouyinRequest {
url: string
}
export interface ExtractFromDouyinResponse {
text: string
duration_seconds?: number
source_url?: string
}
export async function extractScriptFromDouyin(
body: ExtractFromDouyinRequest,
opts?: { signal?: AbortSignal },
): Promise<ExtractFromDouyinResponse> {
const res = await apiClient.post<ExtractFromDouyinResponse>(
"/scripts/extract-from-douyin",
body,
{
// ASR 可能较慢,给足超时
timeout: 60_000,
signal: opts?.signal,
},
)
return res.data
}
/** ── 2. AI 改写文案 */
export type RewriteStyle = "口语化" | "正式" | "活泼" | "治愈" | "励志"
export const REWRITE_STYLE_OPTIONS: { value: RewriteStyle; label: string }[] = [
{ value: "口语化", label: "口语化" },
{ value: "正式", label: "正式" },
{ value: "活泼", label: "活泼" },
{ value: "治愈", label: "治愈" },
{ value: "励志", label: "励志" },
]
export interface AiRewriteRequest {
content: string
style?: RewriteStyle
}
export interface AiRewriteResponse {
original: string
rewritten: string
style: RewriteStyle
}
export async function aiRewriteScript(
body: AiRewriteRequest,
opts?: { signal?: AbortSignal },
): Promise<AiRewriteResponse> {
const res = await apiClient.post<AiRewriteResponse>("/scripts/ai-rewrite", body, {
timeout: 60_000,
signal: opts?.signal,
})
return res.data
}
/** ── 3. AI 生成标题 */
export interface AiGenerateTitlesRequest {
content: string
count?: number
}
export interface AiGenerateTitlesResponse {
titles: string[]
}
export async function aiGenerateTitles(
body: AiGenerateTitlesRequest,
opts?: { signal?: AbortSignal },
): Promise<AiGenerateTitlesResponse> {
const res = await apiClient.post<AiGenerateTitlesResponse>(
"/scripts/ai-generate-titles",
{ content: body.content, count: body.count ?? 3 },
{
timeout: 30_000,
signal: opts?.signal,
},
)
return res.data
}
+18 -180
View File
@@ -1,199 +1,37 @@
/**
* 文案库 API#1811 v2
* CRUD + 搜索/分类/分页;后端未就绪时使用 mock 数据(SCRIPTS_API_MOCK=true
* 文案库 API
* 对接后端 /api/v1/scriptsCRUD + 列表解包
*/
import apiClient from "../client"
import type {
ScriptItem,
ScriptListParams,
ScriptListResponse,
ScriptUpsertRequest,
ScriptCategory,
CreateScriptRequest,
UpdateScriptRequest,
} from "./types"
/**
* 是否启用 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 数据 ====================
const MOCK_CATEGORIES: ScriptCategory[] = [
"promo",
"vlog",
"knowledge",
"story",
"emotion",
"other",
]
const SAMPLE_TITLES: Record<ScriptCategory, string[]> = {
promo: ["新品上市限时特惠", "618大促开场", "品牌故事宣传片"],
vlog: ["周末citywalk记录", "打工人的一天", "探店vlog"],
knowledge: ["3分钟学会XX", "冷知识科普", "行业深度解读"],
story: ["励志小故事", "情感故事一则", "反转剧情"],
emotion: ["深夜emo时刻", "治愈系文案", "朋友圈金句"],
other: ["通用开场白", "节日祝福", "万能结尾"],
}
const SAMPLE_TAGS = ["热门", "新品", "节日", "情感", "干货", "搞笑", "治愈", "励志"]
function genMockScripts(): ScriptItem[] {
const items: ScriptItem[] = []
const now = Date.now()
let idx = 0
for (const cat of MOCK_CATEGORIES) {
const titles = SAMPLE_TITLES[cat]
for (let i = 0; i < titles.length; i++) {
idx++
const title = titles[i]
const content = `这是一条【${cat}】分类下的示例文案,标题为「${title}」。\n\n正文可以包含多段落,每段对应一个片段(segments)。\n\n此为 mock 数据,后端接口就绪后会自动切换为真实数据。`
const segments = content.split(/\n\n+/).filter(Boolean)
const tagCount = 1 + (idx % 3)
const tags: string[] = []
for (let t = 0; t < tagCount; t++) {
tags.push(SAMPLE_TAGS[(idx + t) % SAMPLE_TAGS.length])
}
items.push({
id: `mock_${idx}`,
title,
content,
segments,
tags,
title_text: title,
title_category: cat,
title_config: {
font: "default",
color: "#ffffff",
stroke: "#000000",
position: (["top", "center", "bottom"] as const)[idx % 3],
size: 48,
bold: idx % 2 === 0,
italic: false,
},
char_count: content.length,
use_count: Math.floor(Math.random() * 50),
created_at: new Date(now - idx * 86400_000 * 2).toISOString(),
updated_at: new Date(now - idx * 86400_000).toISOString(),
})
}
}
return items
}
const MOCK_SCRIPTS = genMockScripts()
// ==================== 真实 API ====================
/** 获取文案列表(支持分页/搜索/分类) */
export async function getScripts(params: ScriptListParams = {}): Promise<ScriptListResponse> {
if (SCRIPTS_API_MOCK) {
const page = params.page ?? 1
const pageSize = params.page_size ?? 20
let items = [...MOCK_SCRIPTS]
if (params.keyword) {
const kw = params.keyword.toLowerCase()
items = items.filter(
(s) => s.title.toLowerCase().includes(kw) || s.content.toLowerCase().includes(kw),
)
}
if (params.category && params.category !== "all") {
items = items.filter((s) => s.title_category === params.category)
}
if (params.tag) {
items = items.filter((s) => s.tags?.includes(params.tag as string))
}
const total = items.length
const start = (page - 1) * pageSize
const pageItems = items.slice(start, start + pageSize)
return new Promise((r) =>
setTimeout(() => r({ items: pageItems, total, page, page_size: pageSize }), 200),
)
}
const res = await apiClient.get<ScriptListResponse>("/scripts", { params })
return res.data
}
/** 获取单条文案详情 */
export async function getScript(id: string): Promise<ScriptItem> {
if (SCRIPTS_API_MOCK) {
const item = MOCK_SCRIPTS.find((s) => s.id === id)
return new Promise((r) => setTimeout(() => r(item ?? MOCK_SCRIPTS[0]), 120))
}
const res = await apiClient.get<ScriptItem>(`/scripts/${id}`)
return res.data
/** 获取文案列表 — 必须解包 items(后端返回 {items,total}*/
export const getScripts = async (): Promise<ScriptItem[]> => {
const response = await apiClient.get<ScriptListResponse | ScriptItem[]>("/scripts")
const data = response.data as unknown
if (Array.isArray(data)) return data
const items = (data as { items?: ScriptItem[] })?.items
return Array.isArray(items) ? items : []
}
/** 新建文案 */
export async function createScript(data: ScriptUpsertRequest): Promise<ScriptItem> {
if (SCRIPTS_API_MOCK) {
const segments =
data.segments && data.segments.length > 0
? data.segments
: data.content.split(/\n\n+/).filter(Boolean)
const item: ScriptItem = {
id: `mock_${Date.now()}`,
...data,
segments,
char_count: data.content.length,
use_count: 0,
tags: data.tags ?? [],
created_at: new Date().toISOString(),
updated_at: new Date().toISOString(),
}
MOCK_SCRIPTS.unshift(item)
return new Promise((r) => setTimeout(() => r(item), 200))
}
const res = await apiClient.post<ScriptItem>("/scripts", data)
return res.data
export const createScript = async (data: CreateScriptRequest): Promise<ScriptItem> => {
const response = await apiClient.post<ScriptItem>("/scripts", data)
return response.data
}
/** 更新文案 */
export async function updateScript(id: string, data: ScriptUpsertRequest): Promise<ScriptItem> {
if (SCRIPTS_API_MOCK) {
const idx = MOCK_SCRIPTS.findIndex((s) => s.id === id)
const segments =
data.segments && data.segments.length > 0
? data.segments
: data.content.split(/\n\n+/).filter(Boolean)
const updated: ScriptItem = {
...MOCK_SCRIPTS[idx],
...data,
segments,
char_count: data.content.length,
tags: data.tags ?? MOCK_SCRIPTS[idx]?.tags ?? [],
updated_at: new Date().toISOString(),
}
if (idx >= 0) MOCK_SCRIPTS[idx] = updated
return new Promise((r) => setTimeout(() => r(updated), 200))
}
const res = await apiClient.put<ScriptItem>(`/scripts/${id}`, data)
return res.data
export const updateScript = async (id: string, data: UpdateScriptRequest): Promise<ScriptItem> => {
const response = await apiClient.put<ScriptItem>(`/scripts/${id}`, data)
return response.data
}
/** 删除文案 */
export async function deleteScript(id: string): Promise<void> {
if (SCRIPTS_API_MOCK) {
const idx = MOCK_SCRIPTS.findIndex((s) => s.id === id)
if (idx >= 0) MOCK_SCRIPTS.splice(idx, 1)
return new Promise((r) => setTimeout(r, 150))
}
export const deleteScript = async (id: string): Promise<void> => {
await apiClient.delete(`/scripts/${id}`)
}
/** 复制文案(返回新副本) */
export async function duplicateScript(id: string): Promise<ScriptItem> {
const orig = await getScript(id)
const copy = await createScript({
title: `${orig.title}(副本)`,
content: orig.content,
segments: orig.segments,
tags: orig.tags,
title_text: orig.title_text,
title_category: orig.title_category,
title_config: orig.title_config,
})
return copy
}
+6 -79
View File
@@ -1,97 +1,24 @@
/**
* 文案库 API — 类型定义#1811 v2 完整字段版)
* 字段对齐后端契约:title / content / segments / tags / title_text / title_category / title_config
* 同时保留 char_count / use_count / timestamps 等展示字段
* 文案库 API — 类型定义
* 对接后端 /api/v1/scripts
*/
/** 标题配置(字体、颜色、位置、字号) */
export interface ScriptTitleConfig {
/** 字体预设 key,如 "default" / "bold" / "handwritten" */
font?: string
/** 文字颜色(CSS color */
color?: string
/** 描边色 */
stroke?: string
/** 位置:top / center / bottom */
position?: "top" | "center" | "bottom"
/** 字号(px */
size?: number
/** 是否加粗 */
bold?: boolean
/** 是否斜体 */
italic?: boolean
}
/** 文案分类(可枚举,也支持自定义) */
export type ScriptCategory =
| "promo" // 营销推广
| "vlog" // Vlog/日常
| "knowledge" // 知识科普
| "story" // 故事剧情
| "emotion" // 情感语录
| "other" // 其他
export const SCRIPT_CATEGORY_LABEL: Record<ScriptCategory, string> = {
promo: "营销推广",
vlog: "Vlog 日常",
knowledge: "知识科普",
story: "故事剧情",
emotion: "情感语录",
other: "其他",
}
/** 文案条目 */
export interface ScriptItem {
id: string
/** 名称(标题) */
title: string
/** 正文 */
content: string
/** 分段(按段落切分,供后端/生成步骤逐段使用) */
segments?: string[]
/** 标签(逗号分隔或数组,列表展示用 Tag) */
tags?: string[]
/** 配套标题文本(选填,"使用"跳创作页时会预填到标题) */
title_text?: string
/** 分类 */
title_category?: ScriptCategory
/** 标题样式配置(字体/颜色/位置/字号) */
title_config?: ScriptTitleConfig
/** 正文字符数(后端返回,前端用于展示) */
char_count?: number
/** 使用次数(后端返回) */
use_count?: number
char_count: number
created_at: string
updated_at?: string
}
/** 列表查询参数(支持搜索/分类/分页) */
export interface ScriptListParams {
page?: number
page_size?: number
/** 标题/正文模糊搜索 */
keyword?: string
/** 分类筛选 */
category?: ScriptCategory | "all"
/** 标签筛选 */
tag?: string
}
/** 列表响应 */
export interface ScriptListResponse {
items: ScriptItem[]
total: number
page: number
page_size: number
}
/** 创建/编辑请求 */
export interface ScriptUpsertRequest {
export interface CreateScriptRequest {
title: string
content: string
segments?: string[]
tags?: string[]
title_text?: string
title_category?: ScriptCategory
title_config?: ScriptTitleConfig
}
export type UpdateScriptRequest = Partial<CreateScriptRequest>
+2 -8
View File
@@ -1,30 +1,24 @@
/**
* 订阅 API — 目录化入口
* 对齐后端 staging 最终契约(2026-09-16
* 保持与原 subscription.ts 相同导出,向后兼容
*/
// 类型
export type {
PlanId,
PlanType,
SubscriptionStatus,
BillingStatus,
BillingCycle,
Plan,
SubscriptionInfo,
SubscriptionPlan,
SubscriptionPlansResponse,
BillingRecord,
ChangePlanRequest,
ChangePlanResponse,
ToggleAutoRenewRequest,
} from "./types"
export { PLAN_LABEL, BILLING_CYCLE_LABEL } from "./types"
// API 函数
export {
getCurrentSubscription,
getSubscriptionPlans,
getBillingRecords,
changePlan,
cancelSubscription,
+22 -129
View File
@@ -1,154 +1,47 @@
/**
* 订阅/会员 API 封装
* 对齐后端 staging 实测最终契约(2026-09-16
*
* Base path: /api/v1/
* 所有请求走 apiClient(已配置 baseURL=/api/v1 和 token 拦截器)。
* 订阅相关 API 函数
*/
import apiClient from "../client"
import type {
SubscriptionInfo,
SubscriptionPlan,
SubscriptionPlansResponse,
BillingRecord,
ChangePlanRequest,
ChangePlanResponse,
ToggleAutoRenewRequest,
SubscriptionInfo,
} from "./types"
const MOCK_DELAY = 500
const MOCK_SUBSCRIPTION: SubscriptionInfo = {
id: "sub_mock_001",
plan_id: "free",
plan_name: "免费版",
status: "active",
billing_cycle: "monthly",
current_period_start: new Date(Date.now() - 30 * 86400_000).toISOString(),
current_period_end: new Date(Date.now() + 30 * 86400_000).toISOString(),
amount: 0,
auto_renew: false,
created_at: new Date(Date.now() - 30 * 86400_000).toISOString(),
}
const MOCK_PLANS: SubscriptionPlan[] = [
{
plan_id: "free",
name: "免费版",
price_cents: 0,
monthly_price_cents: 0,
duration_days: 0,
points_discount: 1,
features: { max_resolution: "720p", free_clips_daily: 3 },
},
{
plan_id: "monthly",
name: "月度会员",
price_cents: 1990,
monthly_price_cents: 1990,
duration_days: 30,
points_discount: 0.9,
features: { max_resolution: "1080p", free_clips_daily: 10 },
},
{
plan_id: "quarterly",
name: "季度会员",
price_cents: 3990,
monthly_price_cents: 1330,
duration_days: 90,
points_discount: 0.85,
features: { max_resolution: "1080p", free_clips_daily: 15 },
},
{
plan_id: "yearly",
name: "年度会员",
price_cents: 15900,
monthly_price_cents: 1325,
duration_days: 365,
points_discount: 0.8,
features: { max_resolution: "4k", free_clips_daily: 30 },
},
]
const MOCK_BILLING: BillingRecord[] = []
const isMock = () => (process.env.POINTS_API_MOCK as string | undefined) === "true"
/** 获取当前订阅 */
/** 获取当前订阅信息 */
export const getCurrentSubscription = async (): Promise<SubscriptionInfo> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY))
return { ...MOCK_SUBSCRIPTION }
}
const { data } = await apiClient.get("/subscription/current")
return data
const response = await apiClient.get("/subscription/current")
return response.data
}
/** 获取所有订阅档位 */
export const getSubscriptionPlans = async (): Promise<SubscriptionPlansResponse> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY))
return { plans: MOCK_PLANS.map((p) => ({ ...p, features: { ...p.features } })) }
}
const { data } = await apiClient.get("/subscription/plans")
return data
}
/** 获取账单记录 */
/** 获取账单记录列表 */
export const getBillingRecords = async (): Promise<BillingRecord[]> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY))
return MOCK_BILLING.map((r) => ({ ...r }))
}
const { data } = await apiClient.get("/subscription/billing-records")
return data
const response = await apiClient.get("/subscription/billing-records")
return response.data
}
/** 升级/降级套餐 */
export const changePlan = async (request: ChangePlanRequest): Promise<ChangePlanResponse> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY * 2))
const plan = MOCK_PLANS.find((p) => p.plan_id === request.target_plan_id)
if (!plan) return { success: false, message: "套餐不存在" }
const newSub: SubscriptionInfo = {
...MOCK_SUBSCRIPTION,
plan_id: plan.plan_id,
plan_name: plan.name,
billing_cycle: request.billing_cycle,
amount: plan.price_cents,
status: "pending",
current_period_start: new Date().toISOString(),
current_period_end: new Date(Date.now() + plan.duration_days * 86400_000).toISOString(),
auto_renew: true,
}
return {
success: true,
message: "订阅变更成功(mock,支付通道待接入)",
new_subscription: newSub,
}
}
const { data } = await apiClient.post("/subscription/change-plan", request)
return data
const response = await apiClient.post("/subscription/change-plan", request)
return response.data
}
/** 取消订阅(到期后失效) */
export const cancelSubscription = async (): Promise<{ success: boolean; message: string }> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY))
return { success: true, message: "已取消订阅,到期后将不再续费" }
}
const { data } = await apiClient.post("/subscription/cancel")
return data
/** 取消订阅 */
export const cancelSubscription = async (): Promise<{
success: boolean
message: string
}> => {
const response = await apiClient.post("/subscription/cancel")
return response.data
}
/** 切换自动续费 */
export const toggleAutoRenew = async (
req: ToggleAutoRenewRequest,
enabled: boolean,
): Promise<{ success: boolean; message: string }> => {
if (isMock()) {
await new Promise((r) => setTimeout(r, MOCK_DELAY))
return { success: true, message: req.enabled ? "已开启自动续费" : "已关闭自动续费" }
}
const { data } = await apiClient.post("/subscription/toggle-auto-renew", req)
return data
const response = await apiClient.post("/subscription/toggle-auto-renew", {
enabled,
})
return response.data
}
+29 -79
View File
@@ -1,115 +1,65 @@
/**
* 订阅/会员类型定义
* 对齐后端 staging 实测最终契约(2026-09-16
*
* Base path: /api/v1/
* 订阅相关类型定义
*/
/** 订阅计划 ID */
export type PlanId = "free" | "monthly" | "quarterly" | "yearly"
/** 计费周期 */
export type BillingCycle = "monthly" | "yearly"
/** 套餐类型 */
export type PlanType = "free" | "standard" | "pro" | "enterprise"
/** 订阅状态 */
export type SubscriptionStatus = "active" | "expired" | "cancelled" | "pending"
export type SubscriptionStatus = "active" | "expired" | "cancelled" | "trial"
/** 账单状态 */
export type BillingStatus = "paid" | "pending" | "failed" | "refunded"
/* ================================================================
* 当前订阅(GET /subscription/current
* ================================================================ */
/** 计费周期 */
export type BillingCycle = "monthly" | "yearly"
/** 套餐信息 */
export interface Plan {
id: PlanType
name: string
price: number | null
yearly_price?: number | null
description: string
recommended: boolean
features: string[]
}
/** 当前订阅信息 */
export interface SubscriptionInfo {
id: string
plan_id: PlanId
plan_id: PlanType
plan_name: string
status: SubscriptionStatus
/** 当前计费周期:monthly 对月卡/季卡按自然月续费;yearly 对年卡 */
billing_cycle: BillingCycle
current_period_start: string
current_period_end: string
/** 本期金额(分) */
amount: number
auto_renew: boolean
created_at: string
}
/* ================================================================
* 订阅计划(GET /subscription/plans
* ================================================================ */
export interface SubscriptionPlan {
plan_id: PlanId
/** 中文名 */
name: string
/** 价格(分),年卡/季卡为总价 */
price_cents: number
/** 折算月价(分),对比用 */
monthly_price_cents: number
/** 时长(天) */
duration_days: number
/** 积分折扣(0.9 = 9折,1 = 无折扣) */
points_discount: number
features: {
max_resolution: string
free_clips_daily: number
[key: string]: unknown
}
}
export interface SubscriptionPlansResponse {
plans: SubscriptionPlan[]
}
/* ================================================================
* 账单(GET /subscription/billing-records
* ================================================================ */
/** 账单记录 */
export interface BillingRecord {
id: string
/** 订单类型:subscribe/renew/upgrade/refund */
order_type: string
plan_id: PlanId
/** 金额(分) */
amount_cents: number
plan_name: string
amount: number
billing_cycle: BillingCycle
status: BillingStatus
payment_method: string
created_at: string
paid_at?: string
invoice_url?: string
}
/* ================================================================
* 变更/取消/开关自动续费
* ================================================================ */
/** 升级/降级请求 */
export interface ChangePlanRequest {
target_plan_id: PlanId
target_plan_id: PlanType
billing_cycle: BillingCycle
}
/** 升级/降级响应 */
export interface ChangePlanResponse {
success: boolean
message: string
new_subscription?: SubscriptionInfo
}
export interface ToggleAutoRenewRequest {
enabled: boolean
}
/* ================================================================
* 中文标签映射
* ================================================================ */
export const PLAN_LABEL: Record<PlanId, string> = {
free: "免费版",
monthly: "月度会员",
quarterly: "季度会员",
yearly: "年度会员",
}
export const BILLING_CYCLE_LABEL: Record<BillingCycle, string> = {
monthly: "月付",
yearly: "年付",
}
/**
* @deprecated 旧命名保留别名,新代码请直接用 PlanId
*/
export type PlanType = PlanId
-10
View File
@@ -71,16 +71,6 @@ export interface CreateGenerationTaskRequest {
duration?: number
/** 视频宽高比,如 "9:16" */
video_ratio?: string
/** #1970:剪辑模式 random/narrative */
assembly_mode?: "random" | "narrative"
/** #1970:叙事模式下的文案 ID */
script_id?: string
/** #1970TTS 音色 ID */
tts_voice_id?: string
/** #1970TTS 音色来源 preset/clone */
tts_voice_source?: "preset" | "clone"
/** #1970:智能降重开关(默认 true) */
dedup_enabled?: boolean
/** 标题烧录配置 */
title_config?: {
text?: string
+12 -15
View File
@@ -85,13 +85,9 @@ export async function batchDeleteEditPlanClips(
return response.data
}
/**
* 从素材批量创建片段(追加到时间线末尾)。
* #1921 修复:templateId 为空时调用新端点 POST /clips/from-assets,避免拼出双斜杠
* `/templates//editor/clips/from-assets` 导致 404;有 templateId 时保持原路径向后兼容。
*/
/** 从素材批量创建片段(追加到时间线末尾) */
export async function createClipsFromAssets(
templateId: string | undefined | null,
templateId: string,
assetIds: string[],
clipType = "main",
requiredClipsCount?: number,
@@ -104,15 +100,16 @@ export async function createClipsFromAssets(
if (requiredClipsCount !== undefined) {
body.required_clips_count = requiredClipsCount
}
// 新端点(#1921):templateId 为空时,body 不传 template_id,由后端兜底创建默认模板
const hasTid = !!templateId
const url = hasTid ? `/templates/${templateId}/editor/clips/from-assets` : "/clips/from-assets"
// from-assets 后端会调用 MediaKit 智能选片(最长 60s),单独延长超时
const response = await apiClient.post<ClipsFromAssetsResponse>(url, body, {
timeout: 60000,
signal: opts?.signal,
// _silentErrorToast 由 api/client.ts 响应拦截器读取(抑制全局错误 toast,#1777
...(opts?.silentErrorToast ? ({ _silentErrorToast: true } as Record<string, unknown>) : {}),
})
const response = await apiClient.post<ClipsFromAssetsResponse>(
`/templates/${templateId}/editor/clips/from-assets`,
body,
{
timeout: 60000,
signal: opts?.signal,
// _silentErrorToast 由 api/client.ts 响应拦截器读取(抑制全局错误 toast,#1777
...(opts?.silentErrorToast ? ({ _silentErrorToast: true } as Record<string, unknown>) : {}),
},
)
return response.data
}
+19
View File
@@ -0,0 +1,19 @@
/**
* 标题相关 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"
+65
View File
@@ -0,0 +1,65 @@
/**
* 标题相关 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
}

Some files were not shown because too many files have changed in this diff Show More