Compare commits
45 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e11e4f0e99 | |||
| 966da04c9c | |||
| fdeb792bab | |||
| ff1d878c62 | |||
| eeb8a05b69 | |||
| efb7fa5729 | |||
| 5e1520230f | |||
| 96bcec5fdd | |||
| dfc5e5a5b6 | |||
| 3be06c5763 | |||
| 02199d80ee | |||
| ca6803e1a5 | |||
| d449496f90 | |||
| a04e363d1b | |||
| da59a6c9a6 | |||
| 11c554e43a | |||
| f9f3e6bfb9 | |||
| 3904a8f3b5 | |||
| ba3e97c986 | |||
| 3106496c12 | |||
| f79f75b863 | |||
| b7a439d319 | |||
| 4d7c80ae07 | |||
| 26c0140d79 | |||
| c507f76c14 | |||
| 8a5cfe831e | |||
| d7d3f3184b | |||
| ec9240b52a | |||
| d0e5ef1753 | |||
| 6e8199581d | |||
| 1e23a3f094 | |||
| b04a803655 | |||
| e496f127a3 | |||
| 423be1446f | |||
| 6d9d2e8179 | |||
| b233529eee | |||
| 6de6971e7b | |||
| 548be6aced | |||
| 588a4b7320 | |||
| 1a93a9c00e | |||
| 01991f14d7 | |||
| 8abdeb9551 | |||
| 97ad0ae2e5 | |||
| 59c05148ab | |||
| a00031e100 |
+21
-2
@@ -79,14 +79,33 @@ CELERY_BROKER_URL=redis://localhost:6379/0
|
||||
CELERY_RESULT_BACKEND=redis://localhost:6379/1
|
||||
|
||||
|
||||
# ==================== Worker 配置 ====================
|
||||
# ==================== Worker 配置(#2073 队列分流) ====================
|
||||
#
|
||||
# 容器内跑三个独立进程:beat(只发定时任务)+ generation worker(实时高优)
|
||||
# + transcode worker(后台批量/清理)。三个进程的并发与开关独立配置。
|
||||
|
||||
# Worker 进程名称
|
||||
WORKER_NAME=xiaoxia-saas-worker
|
||||
|
||||
# Worker 并发数(同时执行的任务数)
|
||||
# 总并发参考(兼容旧变量):
|
||||
# - 若 GENERATION_CONCURRENCY 与 TRANSCODE_CONCURRENCY 都未显式设置,
|
||||
# entrypoint 会按此总数对半分配(gen=ceil(total/2), trans=剩余,各至少 1);
|
||||
# - 任一个 *_CONCURRENCY 显式设置后,按显式值生效,忽略此变量对应部分。
|
||||
WORKER_CONCURRENCY=4
|
||||
|
||||
# Generation worker 并发数(用户实时任务:视频生成/TTS/音色克隆/lipsync/数字人)
|
||||
# 实时链路对延迟敏感,建议 2C 以上机器设为 2;高负载场景可加到 4。
|
||||
GENERATION_CONCURRENCY=2
|
||||
|
||||
# Transcode worker 并发数(后台批量:素材入库转码/AI 分类打标/质量评分/查重/批量下载)
|
||||
# 后台任务可排队,独立伸缩;素材入库量大时可加到 4。
|
||||
TRANSCODE_CONCURRENCY=2
|
||||
|
||||
# 是否在本容器启动 celery beat 进程(默认 1)。
|
||||
# 默认 beat 与 worker 同容器部署;若要独立 beat 容器部署,worker 容器设为 0、
|
||||
# beat 容器单独跑 `celery -A worker_app.celery_app beat` 并设 BEAT_ENABLED=1。
|
||||
BEAT_ENABLED=1
|
||||
|
||||
# 每个子进程最多处理多少任务后重启(防止内存泄漏)
|
||||
WORKER_MAX_TASKS_PER_CHILD=1000
|
||||
|
||||
|
||||
@@ -1302,8 +1302,16 @@ jobs:
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/configs/douyin_cookies.txt"
|
||||
echo "✅ Douyin cookies uploaded"
|
||||
|
||||
# 上传 infra/docker 配置到服务器(compose 单一事实来源)
|
||||
echo "Uploading infra/docker configs to staging server..."
|
||||
ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" \
|
||||
"mkdir -p /var/lib/xiaoxia-saas-staging/infra/docker"
|
||||
scp -P "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no infra/docker/compose.yml \
|
||||
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/infra/docker/compose.yml"
|
||||
echo "✅ infra/docker/compose.yml uploaded"
|
||||
|
||||
# 通过环境变量传递凭证,避免命令行引号转义问题
|
||||
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh"
|
||||
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} COMPOSE_SYNC=0 sh"
|
||||
|
||||
# 清理 CI runner 上的渲染文件
|
||||
rm -f .env.rendered
|
||||
|
||||
+100
@@ -0,0 +1,100 @@
|
||||
"""add viral video tables
|
||||
|
||||
Revision ID: 086_add_viral_video_tables
|
||||
Revises: 085_atom_clip_caption_embedding
|
||||
Create Date: 2026-09-28
|
||||
|
||||
新增爆款视频相关表:
|
||||
- viral_video_jobs: 爆款视频任务
|
||||
- viral_video_style_templates: 风格模板配置
|
||||
- viral_video_prompt_templates: Prompt 模板(由 #2040 seed)
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "086_add_viral_video_tables"
|
||||
down_revision = "085_atom_clip_caption_embedding"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# viral_video_jobs
|
||||
op.create_table(
|
||||
"viral_video_jobs",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, index=True),
|
||||
sa.Column("images", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("industry", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("target_customer", sa.String(500), nullable=False, server_default=""),
|
||||
sa.Column("persona_id", sa.String(36), nullable=False, server_default=""),
|
||||
sa.Column("viral_structure", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("marketing_purpose", sa.String(100), nullable=False, server_default=""),
|
||||
sa.Column("bgm_preference", sa.String(50), nullable=False, server_default=""),
|
||||
sa.Column("duration", sa.Integer(), nullable=False, server_default="30"),
|
||||
sa.Column("user_copy_text", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("fusion_level", sa.String(20), nullable=False, server_default="ai_polish"),
|
||||
sa.Column("reference_audio_path", sa.String(1000), nullable=False, server_default=""),
|
||||
# v1.3 新增
|
||||
sa.Column("reference_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_strength", sa.String(20), nullable=False, server_default="medium"),
|
||||
sa.Column("style_guide", sa.JSON(), nullable=True),
|
||||
sa.Column("style_template_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
# 状态与结果
|
||||
sa.Column("status", sa.String(30), nullable=False, server_default="pending", index=True),
|
||||
sa.Column("intent_result", sa.JSON(), nullable=True),
|
||||
sa.Column("result_video_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("credits_cost", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error_msg", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("retry_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_style_templates
|
||||
op.create_table(
|
||||
"viral_video_style_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("thumbnail_url", sa.String(1000), nullable=False, server_default=""),
|
||||
sa.Column("style_config", sa.JSON(), nullable=False, server_default="{}"),
|
||||
sa.Column("is_system", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("sort_order", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# viral_video_prompt_templates
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("prompt_type", sa.String(50), nullable=False, index=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("content", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("variables", sa.JSON(), nullable=False, server_default="[]"),
|
||||
sa.Column("version", sa.Integer(), nullable=False, server_default="1"),
|
||||
sa.Column("is_active", sa.Boolean(), nullable=False, server_default=sa.text("true"), index=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# Seed 默认风格模板
|
||||
op.execute("""
|
||||
INSERT INTO viral_video_style_templates (id, name, description, style_config, is_system, sort_order)
|
||||
VALUES
|
||||
('style-tpl-001', '快节奏冲击', '高频切镜+动感BGM,适合食品饮料等快消品', '{"cut_speed": "fast", "transition": "jump_cut", "energy": "high"}', true, 1),
|
||||
('style-tpl-002', '质感慢镜', '慢节奏+电影感调色,适合美妆护肤珠宝', '{"cut_speed": "slow", "transition": "dissolve", "energy": "low", "color_grade": "cinematic"}', true, 2),
|
||||
('style-tpl-003', '口播种草', '数字人口播+产品特写穿插', '{"cut_speed": "medium", "transition": "cross_dissolve", "has_talking_head": true}', true, 3),
|
||||
('style-tpl-004', '场景叙事', '多场景切换+故事线叙述', '{"cut_speed": "medium", "transition": "wipe", "narrative": true}', true, 4)
|
||||
""")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
op.drop_table("viral_video_style_templates")
|
||||
op.drop_table("viral_video_jobs")
|
||||
@@ -11,9 +11,7 @@ from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.core.task_enqueue import (
|
||||
GLOBAL_PENDING_LIMIT,
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
@@ -302,26 +300,17 @@ def create_preview_generation_task(
|
||||
count,
|
||||
)
|
||||
|
||||
# 预检查队列限流(按变体总数计)
|
||||
try:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending + count > USER_PENDING_LIMIT:
|
||||
raise UserPendingLimitExceeded(
|
||||
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
|
||||
)
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
# 预检查队列限流(按变体总数计)——仅保留全局硬上限,用户上限改为软 warning 在 safe_enqueue 内处理(#2098)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
|
||||
) from e
|
||||
detail=build_rate_limit_detail(
|
||||
GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT),
|
||||
generation_task_repository,
|
||||
scope="global",
|
||||
),
|
||||
)
|
||||
|
||||
# 确定视频比例:优先前端传入,否则从模板 mode 推断
|
||||
video_ratio = request.video_ratio or ""
|
||||
@@ -589,9 +578,6 @@ def create_preview_generation_task(
|
||||
if not enqueued:
|
||||
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
|
||||
_mark_task_failed(generation_task_repository, task, "任务入队失败")
|
||||
except UserPendingLimitExceeded as e:
|
||||
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
|
||||
rate_limit_exc = rate_limit_exc or e
|
||||
except GlobalQueueFull as e:
|
||||
_mark_task_failed(generation_task_repository, task, "系统队列已满")
|
||||
rate_limit_exc = rate_limit_exc or e
|
||||
@@ -603,11 +589,6 @@ def create_preview_generation_task(
|
||||
|
||||
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败")
|
||||
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
|
||||
if isinstance(rate_limit_exc, UserPendingLimitExceeded):
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"),
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"),
|
||||
|
||||
@@ -7,7 +7,6 @@ from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.core.task_enqueue import (
|
||||
GLOBAL_PENDING_LIMIT,
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
@@ -50,12 +49,98 @@ from packages.domain.smart_match import smart_select_assets
|
||||
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
|
||||
# AssetClassification 枚举: scenic / product / person / animal / food / tech / sport / music / other
|
||||
_CATEGORY_KEYWORDS: dict[str, set[str]] = {
|
||||
"scenic": {"风景", "自然", "山水", "大海", "天空", "日落", "日出", "森林", "城市", "建筑", "夜景", "街道", "公园", "景区", "旅行", "旅游", "户外"},
|
||||
"product": {"产品", "商品", "展示", "演示", "开箱", "评测", "好物", "推荐", "种草", "购物", "电商", "带货", "品牌", "广告", "包装"},
|
||||
"person": {"人物", "人物采访", "对话", "说话", "讲解", "演讲", "采访", "聊天", "开会", "工作", "办公室", "团队", "员工", "老板", "女性", "男性", "美女", "帅哥"},
|
||||
"scenic": {
|
||||
"风景",
|
||||
"自然",
|
||||
"山水",
|
||||
"大海",
|
||||
"天空",
|
||||
"日落",
|
||||
"日出",
|
||||
"森林",
|
||||
"城市",
|
||||
"建筑",
|
||||
"夜景",
|
||||
"街道",
|
||||
"公园",
|
||||
"景区",
|
||||
"旅行",
|
||||
"旅游",
|
||||
"户外",
|
||||
},
|
||||
"product": {
|
||||
"产品",
|
||||
"商品",
|
||||
"展示",
|
||||
"演示",
|
||||
"开箱",
|
||||
"评测",
|
||||
"好物",
|
||||
"推荐",
|
||||
"种草",
|
||||
"购物",
|
||||
"电商",
|
||||
"带货",
|
||||
"品牌",
|
||||
"广告",
|
||||
"包装",
|
||||
},
|
||||
"person": {
|
||||
"人物",
|
||||
"人物采访",
|
||||
"对话",
|
||||
"说话",
|
||||
"讲解",
|
||||
"演讲",
|
||||
"采访",
|
||||
"聊天",
|
||||
"开会",
|
||||
"工作",
|
||||
"办公室",
|
||||
"团队",
|
||||
"员工",
|
||||
"老板",
|
||||
"女性",
|
||||
"男性",
|
||||
"美女",
|
||||
"帅哥",
|
||||
},
|
||||
"animal": {"动物", "宠物", "狗", "猫", "鸟", "鱼", "马", "牛", "羊", "野生动物", "动物园"},
|
||||
"food": {"美食", "食物", "餐饮", "餐厅", "做饭", "烹饪", "厨房", "菜品", "饮料", "水果", "甜点", "蛋糕", "咖啡", "茶", "零食", "吃"},
|
||||
"tech": {"科技", "数码", "电脑", "手机", "屏幕", "软件", "APP", "互联网", "AI", "人工智能", "机器人", "办公", "程序员", "代码", "屏幕录制"},
|
||||
"food": {
|
||||
"美食",
|
||||
"食物",
|
||||
"餐饮",
|
||||
"餐厅",
|
||||
"做饭",
|
||||
"烹饪",
|
||||
"厨房",
|
||||
"菜品",
|
||||
"饮料",
|
||||
"水果",
|
||||
"甜点",
|
||||
"蛋糕",
|
||||
"咖啡",
|
||||
"茶",
|
||||
"零食",
|
||||
"吃",
|
||||
},
|
||||
"tech": {
|
||||
"科技",
|
||||
"数码",
|
||||
"电脑",
|
||||
"手机",
|
||||
"屏幕",
|
||||
"软件",
|
||||
"APP",
|
||||
"互联网",
|
||||
"AI",
|
||||
"人工智能",
|
||||
"机器人",
|
||||
"办公",
|
||||
"程序员",
|
||||
"代码",
|
||||
"屏幕录制",
|
||||
},
|
||||
"sport": {"运动", "健身", "跑步", "篮球", "足球", "游泳", "瑜伽", "户外", "锻炼", "体育", "比赛", "球场"},
|
||||
"music": {"音乐", "歌曲", "演唱会", "乐器", "唱歌", "跳舞", "舞蹈", "MV", "演出", "乐队", "钢琴", "吉他", "节奏"},
|
||||
}
|
||||
@@ -77,6 +162,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
break
|
||||
return matched or None
|
||||
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -193,10 +279,13 @@ def _select_assets_from_library(
|
||||
# #2035:加载片段级 AI 标签,供叙事模式 AI 加权和 smart 模式语义匹配使用。
|
||||
# 失败降级为空(不影响选片主流程)。
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] = {}
|
||||
ai_tags_by_asset: dict[str, dict] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重)
|
||||
ai_tags_by_asset: dict[
|
||||
str, dict
|
||||
] = {} # asset_id → 聚合后的 ai_tags dict(取首个有 has_text 的片段;合并 scene/objects/action 去重)
|
||||
try:
|
||||
if db is not None:
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel
|
||||
|
||||
ready_ids = [a.id for a in ready_video_assets]
|
||||
clip_rows = (
|
||||
db.query(AssetAtomClipModel.asset_id, AssetAtomClipModel.ai_tags)
|
||||
@@ -594,21 +683,13 @@ def create_generation_task(
|
||||
# 同批次任务共享 batch_id,用于视频查重时批次内比对
|
||||
batch_id = uuid.uuid4().hex if count > 1 else ""
|
||||
|
||||
# 预检查:批量提交前先看会不会超限,避免建一半才拒
|
||||
# 预检查(Bug B #2098):只保留全局 503 保护,用户级不再硬拒 429;
|
||||
# 超额任务直接入队等待 worker 自然消费,前端展示排队位置而非阻止提交。
|
||||
# USER_PENDING_LIMIT 作为软上限(safe_enqueue 兜底),提高到 20 支持批量提交。
|
||||
try:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending + count > USER_PENDING_LIMIT:
|
||||
raise UserPendingLimitExceeded(
|
||||
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
|
||||
)
|
||||
if global_pending + count > GLOBAL_PENDING_LIMIT:
|
||||
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -923,14 +1004,10 @@ def create_generation_task(
|
||||
else:
|
||||
failed_tasks.append(task)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
# 兜底:如果预检查后又并发提交了,在这里也拦住
|
||||
failed_tasks.append(task)
|
||||
if not created_tasks:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from _e
|
||||
break
|
||||
# Bug B #2098: 用户级限流已改为软限制,此分支理论上不再触发;
|
||||
# 极端并发兜底仍入队(safe_enqueue 内部会打 warning 日志),不 429 拒绝
|
||||
logger.warning("[生成任务] 用户 pending 超软限制,仍允许入队: task_id=%s", task.id)
|
||||
created_tasks.append(task)
|
||||
except GlobalQueueFull as _e:
|
||||
failed_tasks.append(task)
|
||||
if not created_tasks:
|
||||
@@ -1081,10 +1158,8 @@ def confirm_generation(
|
||||
):
|
||||
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
# Bug B #2098: 用户级限流已软处理,理论上不再触发;作为防御仍放行
|
||||
logger.warning("[任务] 用户 pending 超软限制,任务已入队")
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -1263,22 +1338,8 @@ def retry_generation_task(
|
||||
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
|
||||
|
||||
user_id = authenticated_user.user.id
|
||||
# 预检查:创建前判断,>= 上限就拒绝
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
# 预检查(Bug B #2098):只保留全局 503,用户级不再硬拒
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
if user_pending >= USER_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(
|
||||
UserPendingLimitExceeded(
|
||||
user_id=user_id,
|
||||
pending_count=user_pending,
|
||||
limit=USER_PENDING_LIMIT,
|
||||
),
|
||||
generation_task_repository,
|
||||
scope="user",
|
||||
),
|
||||
)
|
||||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
@@ -1321,10 +1382,8 @@ def retry_generation_task(
|
||||
):
|
||||
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
# Bug B #2098: 用户级限流已软处理,理论上不再触发;作为防御仍放行
|
||||
logger.warning("[任务] 用户 pending 超软限制,任务已入队")
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
|
||||
@@ -290,6 +290,43 @@ def _submit_ingest_job(
|
||||
return job
|
||||
|
||||
|
||||
def _find_active_ingest_job(ingest_job_repository: Any, asset_id: str) -> Any | None:
|
||||
"""查询 asset 上是否存在"仍在跑或已成功"的 ingest job(FAILED 视为不存在,需重提)。"""
|
||||
if not asset_id:
|
||||
return None
|
||||
find = getattr(ingest_job_repository, "find_by_asset_id", None)
|
||||
if not callable(find):
|
||||
# 旧仓储未实现 find_by_asset_id,无法判断 → 保守返回 None(走正常流程,
|
||||
# _submit_ingest_job 自身有数据库唯一约束/幂等兜底,不会重复建 job)
|
||||
return None
|
||||
try:
|
||||
return find(asset_id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[upload] find_by_asset_id 查询失败,按无 job 处理: asset=%s", asset_id, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def _is_true_duplicate(existing_asset: Asset, ingest_job_repository: Any) -> tuple[bool, Any | None]:
|
||||
"""判断 `existing_asset` 是真重复(应短路返 duplicated)还是占位(应补提 ingest)。
|
||||
|
||||
返回 (is_duplicate, existing_job):
|
||||
- READY 素材:真重复,job 可能为 None(已就绪不需要 job_id)
|
||||
- PROCESSING/UPLOADING 且已有在跑/已完成 ingest job:幂等重试,真重复,job 返回给前端轮询
|
||||
- PROCESSING/UPLOADING 且无 job:prepare 建的占位 / 之前 ingest 创建失败 → 非重复,需补提 ingest
|
||||
- ERROR/DELETED:非重复(允许重新上传覆盖)
|
||||
"""
|
||||
status = getattr(existing_asset, "status", None)
|
||||
if status == AssetStatus.READY:
|
||||
return True, None
|
||||
if status in (AssetStatus.PROCESSING, AssetStatus.UPLOADING):
|
||||
job = _find_active_ingest_job(ingest_job_repository, existing_asset.id)
|
||||
if job is not None:
|
||||
return True, job
|
||||
return False, None
|
||||
# ERROR / DELETED / 其它:走正常流程重新 ingest
|
||||
return False, None
|
||||
|
||||
|
||||
@router.post("/direct/prepare", response_model=DirectUploadPrepareResponse)
|
||||
async def prepare_direct_upload(
|
||||
request: DirectUploadPrepareRequest,
|
||||
@@ -444,12 +481,25 @@ async def complete_direct_upload(
|
||||
file_size=request.file_size,
|
||||
)
|
||||
if existing is not None:
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id="",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
url=storage_service.get_url(existing.storage_key),
|
||||
is_dup, existing_job = _is_true_duplicate(existing, ingest_job_repository)
|
||||
if is_dup:
|
||||
logger.info(
|
||||
"[upload] complete 幂等命中真重复: asset=%s status=%s job=%s",
|
||||
existing.id,
|
||||
getattr(existing, "status", None),
|
||||
getattr(existing_job, "id", None),
|
||||
)
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id=getattr(existing_job, "id", "") or "",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
url=storage_service.get_url(existing.storage_key),
|
||||
)
|
||||
logger.info(
|
||||
"[upload] complete 命中占位 asset(status=%s 无 ingest job),继续补提 ingest: asset=%s",
|
||||
getattr(existing, "status", None),
|
||||
existing.id,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -479,14 +529,24 @@ async def complete_direct_upload(
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
storage_key=normalized_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=request.file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
# 幂等保护:补提占位场景下可能已有 job(极端竞态),先查一次
|
||||
existing_job = _find_active_ingest_job(ingest_job_repository, pending_asset.id)
|
||||
if existing_job is not None:
|
||||
logger.info(
|
||||
"[upload] complete 补提时发现 job 已存在(竞态/并发重试),复用: asset=%s job=%s",
|
||||
pending_asset.id,
|
||||
existing_job.id,
|
||||
)
|
||||
job = existing_job
|
||||
else:
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
storage_key=normalized_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=request.file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=normalized_key,
|
||||
ingest_job_id=job.id,
|
||||
@@ -533,12 +593,27 @@ async def upload_asset(
|
||||
file_size=0,
|
||||
)
|
||||
if existing is not None:
|
||||
return UploadAssetResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id="",
|
||||
url="",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
is_dup, existing_job = _is_true_duplicate(existing, ingest_job_repository)
|
||||
if is_dup:
|
||||
logger.info(
|
||||
"[upload] multipart 幂等命中真重复: asset=%s status=%s job=%s",
|
||||
existing.id,
|
||||
getattr(existing, "status", None),
|
||||
getattr(existing_job, "id", None),
|
||||
)
|
||||
return UploadAssetResponse(
|
||||
storage_key=existing.storage_key,
|
||||
ingest_job_id=getattr(existing_job, "id", "") or "",
|
||||
url=storage_service.get_url(existing.storage_key)
|
||||
if getattr(existing, "status", None) == AssetStatus.READY
|
||||
else "",
|
||||
duplicated=True,
|
||||
asset_id=existing.id,
|
||||
)
|
||||
logger.info(
|
||||
"[upload] multipart 命中占位 asset(status=%s 无 ingest job),继续补提 ingest: asset=%s",
|
||||
getattr(existing, "status", None),
|
||||
existing.id,
|
||||
)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
@@ -574,14 +649,23 @@ async def upload_asset(
|
||||
)
|
||||
# Issue #1776: 计数由 asset_repository.create() 自动维护
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
existing_job = _find_active_ingest_job(ingest_job_repository, pending_asset.id)
|
||||
if existing_job is not None:
|
||||
logger.info(
|
||||
"[upload] multipart 补提时发现 job 已存在(竞态/并发重试),复用: asset=%s job=%s",
|
||||
pending_asset.id,
|
||||
existing_job.id,
|
||||
)
|
||||
job = existing_job
|
||||
else:
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=file_hash,
|
||||
asset_id=pending_asset.id,
|
||||
)
|
||||
|
||||
return UploadAssetResponse(
|
||||
storage_key=storage_key,
|
||||
|
||||
@@ -6,7 +6,7 @@ from app.core.celery_app import celery_app
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
|
||||
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
|
||||
USER_PENDING_LIMIT = 20 # 单用户 pending 上限(#2098: 从 3 提到 20,支持批量任务自动排队)
|
||||
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
|
||||
WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值)
|
||||
|
||||
@@ -154,19 +154,18 @@ def check_queue_limits(
|
||||
user_pending_limit: int = USER_PENDING_LIMIT,
|
||||
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
|
||||
) -> None:
|
||||
"""检查队列限流(预检查用,任务创建前调用),超限抛对应异常。
|
||||
"""检查队列限流(预检查用,任务创建前调用)。
|
||||
|
||||
边界语义:>= 上限即拒绝(达到上限就不能再加新任务)。
|
||||
#2098 语义变更:用户级限流改为软提示,不再抛异常拒绝;仅全局硬上限抛 GlobalQueueFull。
|
||||
|
||||
Args:
|
||||
user_id: 用户 ID
|
||||
user_id: 用户 ID(保留参数,当前不做用户级硬拒)
|
||||
generation_task_repository: 任务仓储
|
||||
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
|
||||
user_pending_limit: 单用户 pending 上限(保留,当前未硬拒)
|
||||
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
|
||||
|
||||
Raises:
|
||||
GlobalQueueFull: 全局超限时抛出(优先级更高,先查全局)
|
||||
UserPendingLimitExceeded: 用户超限时抛出
|
||||
GlobalQueueFull: 全局超限时抛出
|
||||
"""
|
||||
# 先查全局(系统级保护优先级更高)
|
||||
global_pending = generation_task_repository.count_pending_total()
|
||||
@@ -179,17 +178,9 @@ def check_queue_limits(
|
||||
)
|
||||
raise GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
|
||||
|
||||
# 再查用户级
|
||||
if user_id:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
if user_pending >= user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
|
||||
user_id,
|
||||
user_pending,
|
||||
user_pending_limit,
|
||||
)
|
||||
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
|
||||
# #2098: 用户级限流改为软提示,不在预检查阶段拒绝(超额任务仍入队排队)。
|
||||
# 真正的系统保护由全局 GLOBAL_PENDING_LIMIT 硬上限承担。
|
||||
# UserPendingLimitExceeded 保留以兼容历史 import/except,但预检查与 safe_enqueue 均不再 raise。
|
||||
|
||||
|
||||
def _mark_task_failed_safely(
|
||||
@@ -246,7 +237,6 @@ def safe_enqueue_generation_task(
|
||||
|
||||
Raises:
|
||||
GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed
|
||||
UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed
|
||||
"""
|
||||
# ── 入队前检查:任务已是 pending,用 > 判断(包含当前任务) ──
|
||||
|
||||
@@ -263,19 +253,18 @@ def safe_enqueue_generation_task(
|
||||
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
||||
raise exc
|
||||
|
||||
# 用户级限流检查(传了 user_id 才做)
|
||||
# Bug B #2098: 用户级限流改为软提示,不再硬拒;所有任务都入队等待 worker 自然消费。
|
||||
# user_pending_limit 作为兜底阈值保留(默认 20),达到时打 warning 日志但仍入队,
|
||||
# 避免极端情况下恶意用户无限堆积任务。真正的系统保护由全局 GLOBAL_PENDING_LIMIT 承担。
|
||||
if user_id:
|
||||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||||
if user_pending > user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 任务数超限(入队前): user_id=%s, count=%d/%d",
|
||||
"[队列限流] 用户 pending 任务数超过软上限(入队): user_id=%s, count=%d/%d, 仍允许入队排队",
|
||||
user_id,
|
||||
user_pending,
|
||||
user_pending_limit,
|
||||
)
|
||||
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
|
||||
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
||||
raise exc
|
||||
|
||||
# ── 发送 Celery 任务 ──
|
||||
try:
|
||||
@@ -317,16 +306,18 @@ def safe_enqueue_generation_task(
|
||||
user_after = generation_task_repository.count_pending_by_user(user_id) if user_id else 0
|
||||
|
||||
global_over = global_after > global_pending_limit
|
||||
user_over = bool(user_id and user_after > user_pending_limit)
|
||||
|
||||
if global_over or user_over:
|
||||
if global_over:
|
||||
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
|
||||
exc = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
|
||||
else:
|
||||
reason = f"用户 pending 超限(入队后): {user_after}/{user_pending_limit}"
|
||||
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_after, limit=user_pending_limit)
|
||||
# Bug B #2098: 用户超限仅日志警告,不回滚任务
|
||||
if user_id and user_after > user_pending_limit:
|
||||
logger.warning(
|
||||
"[队列限流] 用户 pending 超软上限(入队后): user_id=%s, count=%d/%d",
|
||||
user_id,
|
||||
user_after,
|
||||
user_pending_limit,
|
||||
)
|
||||
|
||||
if global_over:
|
||||
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
|
||||
exc = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
|
||||
logger.warning(
|
||||
"[队列限流] %s, task_id=%s, user_id=%s — 回滚状态为 failed",
|
||||
reason,
|
||||
|
||||
@@ -3,7 +3,8 @@
|
||||
*/
|
||||
|
||||
/** 任务状态 */
|
||||
export type TaskStatus = "pending" | "waiting" | "running" | "completed" | "failed" | "cancelled"
|
||||
export type TaskStatus =
|
||||
"pending" | "waiting" | "running" | "awaiting_cover" | "completed" | "failed" | "cancelled"
|
||||
|
||||
/** 任务类型 */
|
||||
export type TaskType = "ingest" | "generation" | string
|
||||
|
||||
@@ -0,0 +1,271 @@
|
||||
/**
|
||||
* 标题迷你 Canvas 预览(#2001)
|
||||
*
|
||||
* 渲染一张指定宽度的小 Canvas 预览标题效果,用于:
|
||||
* - 预设卡片缩略图
|
||||
* - 样式面板顶部的实时预览
|
||||
*
|
||||
* 与 titleCanvas.ts 渲染逻辑保持一致,但:
|
||||
* - 固定分辨率(width × 宽高比约 2:1)
|
||||
* - 不调用 ffmpeg,只做视觉预览
|
||||
* - 支持背景色块、描边宽度/颜色、阴影参数化、行距、自动换行
|
||||
*/
|
||||
import React, { useEffect, useRef } from "react"
|
||||
import type { TitleStyleSettings } from "@/components/title/settings"
|
||||
import { getFontFamily } from "@/components/title/constants"
|
||||
|
||||
interface Props {
|
||||
settings: TitleStyleSettings
|
||||
width?: number
|
||||
sampleText?: string
|
||||
/** 背景(预览用,默认深色渐变模拟视频底),transparent=true 时忽略 */
|
||||
background?: string
|
||||
/** 高度(可选,默认按 portrait 选比例) */
|
||||
height?: number
|
||||
/** 透明背景(卡片/编辑器预览叠加在图片上时使用) */
|
||||
transparent?: boolean
|
||||
/** 纵向竖屏预览(9:16),true 时 aspect=16/9 适配手机视频比例 */
|
||||
portrait?: boolean
|
||||
}
|
||||
|
||||
/** 按 maxCharsPerLine 自动换行 */
|
||||
function wrapLines(text: string, maxChars: number): string[] {
|
||||
const manual = text
|
||||
.split(/[//\n]/)
|
||||
.map((l) => l.trim())
|
||||
.filter(Boolean)
|
||||
if (!maxChars || maxChars <= 0) return manual
|
||||
const out: string[] = []
|
||||
for (const line of manual) {
|
||||
if (line.length <= maxChars) {
|
||||
out.push(line)
|
||||
continue
|
||||
}
|
||||
let cur = ""
|
||||
for (const ch of line) {
|
||||
cur += ch
|
||||
if (cur.length >= maxChars) {
|
||||
out.push(cur)
|
||||
cur = ""
|
||||
}
|
||||
}
|
||||
if (cur) out.push(cur)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
const TitleMiniPreview: React.FC<Props> = ({
|
||||
settings,
|
||||
width = 200,
|
||||
sampleText,
|
||||
background = "linear-gradient(135deg,#1f2937,#111827)",
|
||||
height,
|
||||
transparent = false,
|
||||
portrait = false,
|
||||
}) => {
|
||||
const canvasRef = useRef<HTMLCanvasElement>(null)
|
||||
const h = height ?? Math.round(width * (portrait ? 16 / 9 : 1 / 1.8))
|
||||
const text = (sampleText || "预览标题").trim() || "预览标题"
|
||||
|
||||
useEffect(() => {
|
||||
let cancelled = false
|
||||
const draw = () => {
|
||||
if (cancelled) return
|
||||
const cvs = canvasRef.current
|
||||
if (!cvs) return
|
||||
const dpr = window.devicePixelRatio || 1
|
||||
cvs.width = width * dpr
|
||||
cvs.height = h * dpr
|
||||
cvs.style.width = `${width}px`
|
||||
cvs.style.height = `${h}px`
|
||||
const ctx = cvs.getContext("2d")
|
||||
if (!ctx) return
|
||||
ctx.scale(dpr, dpr)
|
||||
ctx.clearRect(0, 0, width, h)
|
||||
|
||||
// 背景(transparent 时跳过,用于叠加在图片上)
|
||||
if (!transparent) {
|
||||
ctx.fillStyle = "#111827"
|
||||
ctx.fillRect(0, 0, width, h)
|
||||
}
|
||||
|
||||
// 分辨率缩放:以 360 宽为基准(对应 720p 的一半),与外层 previewScale/previewR 保持一致
|
||||
const r = previewR
|
||||
|
||||
// 字体
|
||||
const size = r(settings.size)
|
||||
const ff = getFontFamily(settings.font)
|
||||
const parts: string[] = []
|
||||
if (settings.italic) parts.push("italic")
|
||||
if (settings.bold) parts.push("bold")
|
||||
parts.push(`${size}px`, ff)
|
||||
ctx.font = parts.join(" ")
|
||||
ctx.textAlign = "center"
|
||||
ctx.textBaseline = "middle"
|
||||
ctx.fillStyle = settings.color
|
||||
ctx.lineJoin = "round"
|
||||
|
||||
// 阴影
|
||||
const shadowEnabled = !!settings.shadow
|
||||
const prevShadow = {
|
||||
c: ctx.shadowColor,
|
||||
b: ctx.shadowBlur,
|
||||
ox: ctx.shadowOffsetX,
|
||||
oy: ctx.shadowOffsetY,
|
||||
}
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
|
||||
// 换行
|
||||
const lines = wrapLines(text, settings.maxCharsPerLine ?? 0)
|
||||
const lineH = size * (settings.lineHeight ?? 1.2)
|
||||
const totalH = lines.length * lineH
|
||||
let startY: number
|
||||
if (settings.position === "top") {
|
||||
startY = size / 2 + r(settings.marginTop ?? 24)
|
||||
} else if (settings.position === "center") {
|
||||
startY = h / 2 - totalH / 2 + size / 2
|
||||
} else {
|
||||
// bottom
|
||||
const botMargin = portrait ? r(24) : r(16)
|
||||
startY = h - totalH - botMargin + size / 2
|
||||
}
|
||||
let centerX = width / 2
|
||||
if (settings.position === "custom" && settings.posX != null) {
|
||||
centerX = (settings.posX / 100) * width
|
||||
}
|
||||
|
||||
// 背景块
|
||||
if (settings.bgEnabled) {
|
||||
const pad = r(settings.bgPadding ?? 12)
|
||||
const rad = r(settings.bgRadius ?? 8)
|
||||
let maxLineW = 0
|
||||
for (const l of lines) {
|
||||
const m = ctx.measureText(l)
|
||||
if (m.width > maxLineW) maxLineW = m.width
|
||||
}
|
||||
const bw = maxLineW + pad * 2
|
||||
const bh = totalH + pad * 2
|
||||
const bx = centerX - bw / 2
|
||||
const by = startY - size / 2 - pad + (size - lineH) / 2
|
||||
ctx.shadowColor = "rgba(0,0,0,0)"
|
||||
ctx.shadowBlur = 0
|
||||
ctx.fillStyle = settings.bgColor ?? "rgba(0,0,0,0.5)"
|
||||
roundRect(ctx, bx, by, bw, bh, rad)
|
||||
ctx.fill()
|
||||
// 关键修复:画完背景块后必须把 fillStyle 重置为文字颜色,
|
||||
// 否则后续 fillText 会用 bgColor 填充文字,导致「文字看不见只剩色块」
|
||||
ctx.fillStyle = settings.color
|
||||
// 恢复阴影
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
}
|
||||
|
||||
// 描边(先画,再画填充)
|
||||
const strokeEnabled = !!settings.stroke && (settings.strokeWidth ?? 0) > 0
|
||||
lines.forEach((line, i) => {
|
||||
const y = startY + i * lineH
|
||||
if (strokeEnabled) {
|
||||
ctx.shadowColor = "rgba(0,0,0,0)"
|
||||
ctx.shadowBlur = 0
|
||||
ctx.lineWidth = r(settings.strokeWidth ?? 4)
|
||||
ctx.strokeStyle = settings.strokeColor ?? "#000000"
|
||||
ctx.strokeText(line, centerX, y)
|
||||
// 恢复阴影
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
}
|
||||
ctx.fillText(line, centerX, y)
|
||||
})
|
||||
|
||||
// 恢复
|
||||
ctx.shadowColor = prevShadow.c
|
||||
ctx.shadowBlur = prevShadow.b
|
||||
ctx.shadowOffsetX = prevShadow.ox
|
||||
ctx.shadowOffsetY = prevShadow.oy
|
||||
}
|
||||
// 计算当前字号(draw() 内部同样逻辑,抽出来供 fontString 复用)
|
||||
const previewScale = width / 360
|
||||
const previewR = (v: number) => Math.round(v * previewScale)
|
||||
const buildFontString = () => {
|
||||
const size = previewR(settings.size)
|
||||
const ff = getFontFamily(settings.font)
|
||||
const parts: string[] = []
|
||||
if (settings.italic) parts.push("italic")
|
||||
if (settings.bold) parts.push("bold")
|
||||
parts.push(`${size}px`, ff)
|
||||
return parts.join(" ")
|
||||
}
|
||||
|
||||
// Web Font 加载保障:
|
||||
// 1) 等 document.fonts.ready(CSS @font-face 首次可用)
|
||||
// 2) 显式 FontFaceSet.load(fontString, text) 触发浏览器真正下载并加载
|
||||
// 当前字体到 Canvas 可用,避免首次绘制用 fallback 字体画出错字/色块
|
||||
const doDrawWhenReady = async () => {
|
||||
try {
|
||||
if (typeof document !== "undefined" && document.fonts) {
|
||||
await document.fonts.ready
|
||||
try {
|
||||
await document.fonts.load(buildFontString(), text)
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
if (!cancelled) draw()
|
||||
}
|
||||
}
|
||||
doDrawWhenReady()
|
||||
return () => {
|
||||
cancelled = true
|
||||
}
|
||||
}, [settings, width, h, text, transparent, portrait, background])
|
||||
|
||||
return (
|
||||
<canvas
|
||||
ref={canvasRef}
|
||||
style={{
|
||||
borderRadius: 6,
|
||||
display: "block",
|
||||
maxWidth: "100%",
|
||||
background: transparent ? "transparent" : background,
|
||||
}}
|
||||
/>
|
||||
)
|
||||
}
|
||||
|
||||
function roundRect(
|
||||
ctx: CanvasRenderingContext2D,
|
||||
x: number,
|
||||
y: number,
|
||||
w: number,
|
||||
h: number,
|
||||
r: number,
|
||||
) {
|
||||
const rr = Math.min(r, w / 2, h / 2)
|
||||
ctx.beginPath()
|
||||
ctx.moveTo(x + rr, y)
|
||||
ctx.lineTo(x + w - rr, y)
|
||||
ctx.quadraticCurveTo(x + w, y, x + w, y + rr)
|
||||
ctx.lineTo(x + w, y + h - rr)
|
||||
ctx.quadraticCurveTo(x + w, y + h, x + w - rr, y + h)
|
||||
ctx.lineTo(x + rr, y + h)
|
||||
ctx.quadraticCurveTo(x, y + h, x, y + h - rr)
|
||||
ctx.lineTo(x, y + rr)
|
||||
ctx.quadraticCurveTo(x, y, x + rr, y)
|
||||
ctx.closePath()
|
||||
}
|
||||
|
||||
export default TitleMiniPreview
|
||||
@@ -0,0 +1,458 @@
|
||||
/* ============================================================
|
||||
TitleStylePanel 标题样式面板 — 独立共用样式(#1809 ⑦)
|
||||
|
||||
从 generate.css 抽取的标题样式区块,供「智能剪辑」与「AI数字人」
|
||||
两个页面共用。AI数字人页面不引入 generate.css,直接由
|
||||
TitleStylePanel.tsx import 本文件,保证 24 个 T 预设格子的网格布局、
|
||||
配色描边、选中态与智能剪辑页面完全一致。
|
||||
|
||||
注意:本文件规则与 generate.css 中同名规则一一对应、取值相同;
|
||||
智能剪辑页面两处同时存在时同优先级同值,不改变其原有呈现。
|
||||
============================================================ */
|
||||
|
||||
/* ── 区块容器 ── */
|
||||
.xx-title-style-section {
|
||||
margin-top: 22px;
|
||||
padding-top: 20px;
|
||||
border-top: 1px solid var(--border-light);
|
||||
}
|
||||
|
||||
.xx-section-subtitle {
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: var(--text-primary);
|
||||
margin: 0 0 16px;
|
||||
}
|
||||
|
||||
.xx-title-style-row {
|
||||
display: grid;
|
||||
grid-template-columns: 1fr 1fr;
|
||||
gap: 14px;
|
||||
margin-bottom: 14px;
|
||||
}
|
||||
|
||||
.xx-half-field {
|
||||
margin-bottom: 0;
|
||||
}
|
||||
|
||||
.xx-field-label-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
|
||||
.xx-field-label-row label {
|
||||
margin-bottom: 0;
|
||||
}
|
||||
|
||||
.xx-field-value {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: var(--primary-color);
|
||||
}
|
||||
|
||||
/* ── 共用表单字段(位置/字体下拉) ── */
|
||||
.xx-title-style-section .xx-form-field {
|
||||
margin-bottom: 14px;
|
||||
}
|
||||
|
||||
.xx-title-style-section .xx-form-field:last-child {
|
||||
margin-bottom: 0;
|
||||
}
|
||||
|
||||
.xx-title-style-section .xx-form-field label {
|
||||
display: block;
|
||||
font-weight: 600;
|
||||
margin-bottom: 8px;
|
||||
font-size: 13px;
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.xx-title-style-section .xx-form-field select,
|
||||
.xx-title-style-section .xx-form-field input {
|
||||
width: 100%;
|
||||
height: 44px;
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: var(--radius-sm);
|
||||
background: var(--bg-primary);
|
||||
padding: 0 14px;
|
||||
font-size: 14px;
|
||||
outline: 0;
|
||||
transition: 0.15s ease;
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.xx-title-style-section .xx-form-field select:focus,
|
||||
.xx-title-style-section .xx-form-field input:focus {
|
||||
border-color: var(--primary-color);
|
||||
box-shadow: 0 0 0 3px rgba(79, 70, 229, 0.1);
|
||||
}
|
||||
|
||||
/* ── 字号滑块 ── */
|
||||
.xx-slider {
|
||||
width: 100%;
|
||||
height: 6px;
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
background: var(--border-color);
|
||||
border-radius: 3px;
|
||||
outline: none;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.xx-slider::-webkit-slider-thumb {
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 18px;
|
||||
height: 18px;
|
||||
background: var(--primary-color);
|
||||
border-radius: 50%;
|
||||
cursor: pointer;
|
||||
box-shadow: 0 2px 6px rgba(79, 70, 229, 0.3);
|
||||
}
|
||||
|
||||
.xx-slider::-moz-range-thumb {
|
||||
width: 18px;
|
||||
height: 18px;
|
||||
background: var(--primary-color);
|
||||
border-radius: 50%;
|
||||
cursor: pointer;
|
||||
border: none;
|
||||
box-shadow: 0 2px 6px rgba(79, 70, 229, 0.3);
|
||||
}
|
||||
|
||||
/* ── 标题预设卡片网格(24 个 T 格子) ── */
|
||||
.xx-title-presets-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(6, 52px);
|
||||
gap: 1px;
|
||||
}
|
||||
|
||||
.xx-title-preset-card {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
width: 52px;
|
||||
height: 52px;
|
||||
padding: 0;
|
||||
background: #404040;
|
||||
border: 2px solid transparent;
|
||||
border-radius: 8px;
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.xx-title-preset-card:hover {
|
||||
border-color: #666;
|
||||
background: #4d4d4d;
|
||||
}
|
||||
|
||||
.xx-title-preset-card.active {
|
||||
border-color: #409eff;
|
||||
background: #4d4d4d;
|
||||
}
|
||||
|
||||
.xx-title-preset-preview-text {
|
||||
font-size: 32px;
|
||||
line-height: 1;
|
||||
user-select: none;
|
||||
}
|
||||
|
||||
/* ── 样式按钮组(加粗/斜体/描边/阴影) ── */
|
||||
.xx-style-btns {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.xx-style-btn {
|
||||
width: 40px;
|
||||
height: 40px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: var(--radius-sm);
|
||||
background: var(--bg-primary);
|
||||
cursor: pointer;
|
||||
font-size: 15px;
|
||||
color: var(--text-secondary);
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.xx-style-btn:hover {
|
||||
border-color: var(--primary-300);
|
||||
color: var(--primary-color);
|
||||
}
|
||||
|
||||
.xx-style-btn.active {
|
||||
background: var(--primary-color);
|
||||
border-color: var(--primary-color);
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
#2001 爆款标题样式面板升级 — 新增样式(ts- 前缀)
|
||||
============================================================ */
|
||||
|
||||
.ts-panel {
|
||||
position: relative;
|
||||
}
|
||||
|
||||
/* 预览 */
|
||||
.ts-preview-wrap {
|
||||
margin-bottom: 14px;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
padding: 10px;
|
||||
background: #0f172a;
|
||||
border-radius: 8px;
|
||||
}
|
||||
|
||||
/* 表单字段 */
|
||||
.ts-form-field {
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
.ts-form-field label {
|
||||
display: block;
|
||||
font-weight: 600;
|
||||
margin-bottom: 6px;
|
||||
font-size: 12px;
|
||||
color: var(--text-primary, #1f2937);
|
||||
}
|
||||
.ts-field-label-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.ts-field-value {
|
||||
font-size: 12px;
|
||||
font-weight: 600;
|
||||
color: var(--primary-color, #7c3aed);
|
||||
}
|
||||
.ts-row-2 {
|
||||
display: grid;
|
||||
grid-template-columns: 1fr 1fr;
|
||||
gap: 10px;
|
||||
}
|
||||
.ts-half {
|
||||
margin-bottom: 0;
|
||||
}
|
||||
|
||||
.ts-select {
|
||||
width: 100%;
|
||||
height: 34px;
|
||||
border: 1px solid var(--border-color, #e5e7eb);
|
||||
border-radius: 6px;
|
||||
background: var(--bg-primary, #fff);
|
||||
padding: 0 10px;
|
||||
font-size: 13px;
|
||||
outline: 0;
|
||||
color: var(--text-primary, #1f2937);
|
||||
}
|
||||
.ts-select:focus {
|
||||
border-color: var(--primary-color, #7c3aed);
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.1);
|
||||
}
|
||||
.ts-input {
|
||||
width: 100%;
|
||||
height: 34px;
|
||||
border: 1px solid var(--border-color, #e5e7eb);
|
||||
border-radius: 6px;
|
||||
padding: 0 10px;
|
||||
font-size: 13px;
|
||||
outline: 0;
|
||||
}
|
||||
|
||||
.ts-slider {
|
||||
width: 100%;
|
||||
height: 4px;
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
background: #e5e7eb;
|
||||
border-radius: 2px;
|
||||
outline: none;
|
||||
}
|
||||
.ts-slider::-webkit-slider-thumb {
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
cursor: pointer;
|
||||
border: 2px solid #fff;
|
||||
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.2);
|
||||
}
|
||||
.ts-slider::-moz-range-thumb {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
background: #7c3aed;
|
||||
cursor: pointer;
|
||||
border: 2px solid #fff;
|
||||
}
|
||||
|
||||
/* 样式按钮 B/I/S/☁ */
|
||||
.ts-style-btns {
|
||||
display: flex;
|
||||
gap: 6px;
|
||||
}
|
||||
.ts-style-btn {
|
||||
width: 34px;
|
||||
height: 34px;
|
||||
border-radius: 6px;
|
||||
border: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
cursor: pointer;
|
||||
font-size: 14px;
|
||||
transition: 0.15s;
|
||||
color: #374151;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
.ts-style-btn:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.ts-style-btn.active {
|
||||
background: #faf5ff;
|
||||
color: #6d28d9;
|
||||
border-color: #7c3aed;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
/* 色板 */
|
||||
.ts-color-row {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 6px;
|
||||
align-items: center;
|
||||
}
|
||||
.ts-color-swatch {
|
||||
width: 24px;
|
||||
height: 24px;
|
||||
border-radius: 4px;
|
||||
border: 2px solid #fff;
|
||||
box-shadow: 0 0 0 1px #e5e7eb;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
transition: 0.15s;
|
||||
}
|
||||
.ts-color-swatch:hover {
|
||||
transform: scale(1.1);
|
||||
}
|
||||
.ts-color-swatch.active {
|
||||
box-shadow: 0 0 0 2px #7c3aed;
|
||||
transform: scale(1.1);
|
||||
}
|
||||
.ts-color-custom {
|
||||
background: repeating-conic-gradient(#ccc 0% 25%, #fff 0% 50%) 50%/8px 8px;
|
||||
color: #666;
|
||||
font-size: 14px;
|
||||
line-height: 20px;
|
||||
}
|
||||
.ts-color-native {
|
||||
width: 0;
|
||||
height: 0;
|
||||
border: 0;
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
/* 预设网格 10个 - 5列 */
|
||||
.ts-presets-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(5, 1fr);
|
||||
gap: 6px;
|
||||
}
|
||||
.ts-preset-card {
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 6px;
|
||||
background: #fff;
|
||||
padding: 4px;
|
||||
cursor: pointer;
|
||||
transition: 0.15s;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 4px;
|
||||
}
|
||||
.ts-preset-card:hover {
|
||||
border-color: #7c3aed;
|
||||
}
|
||||
.ts-preset-card.active {
|
||||
border-color: #7c3aed;
|
||||
background: #faf5ff;
|
||||
box-shadow: 0 0 0 1px #7c3aed;
|
||||
}
|
||||
.ts-preset-preview {
|
||||
height: 34px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
overflow: hidden;
|
||||
border-radius: 4px;
|
||||
background: #0f172a;
|
||||
}
|
||||
.ts-preset-preview canvas {
|
||||
max-width: 100%;
|
||||
max-height: 100%;
|
||||
}
|
||||
.ts-preset-meta {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 2px;
|
||||
font-size: 10px;
|
||||
color: #4b5563;
|
||||
justify-content: center;
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
padding: 0 2px 2px;
|
||||
}
|
||||
.ts-preset-emoji {
|
||||
font-size: 11px;
|
||||
}
|
||||
.ts-preset-label {
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
}
|
||||
|
||||
.ts-toggle-row label {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
cursor: pointer;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.ts-toggle-row input[type="checkbox"] {
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
accent-color: #7c3aed;
|
||||
}
|
||||
|
||||
/* Tabs 紧凑样式 */
|
||||
.xx-title-style-section .ant-tabs-nav {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.xx-title-style-section .ant-tabs-tab {
|
||||
font-size: 12px !important;
|
||||
padding: 6px 8px !important;
|
||||
}
|
||||
|
||||
/* 标题模板入口按钮(#2003) */
|
||||
.ts-template-btn {
|
||||
border: none;
|
||||
background: transparent;
|
||||
color: var(--primary-color, #7c3aed);
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
padding: 2px 0;
|
||||
font-weight: 500;
|
||||
}
|
||||
.ts-template-btn:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
@@ -0,0 +1,445 @@
|
||||
/**
|
||||
* 标题样式参数 Tab 面板(共享组件)
|
||||
*
|
||||
* 包含:基础/描边/阴影/背景/排版/封面 共 6 个 Tab
|
||||
* 仅负责 UI 渲染和参数 patch 回调,不维护 state、不调 API
|
||||
*/
|
||||
import React, { useState } from "react"
|
||||
import { Tabs } from "antd"
|
||||
import type { TitleStyleSettings } from "./settings"
|
||||
import {
|
||||
FONT_OPTIONS,
|
||||
TITLE_COLOR_PALETTE,
|
||||
STROKE_COLOR_PALETTE,
|
||||
BG_COLOR_PALETTE,
|
||||
} from "./constants"
|
||||
|
||||
export interface PositionOption {
|
||||
value: string
|
||||
label: string
|
||||
}
|
||||
|
||||
export interface FontOption {
|
||||
value: string
|
||||
label: string
|
||||
family: string
|
||||
tag?: "hot" | "new"
|
||||
}
|
||||
|
||||
export interface TitleStyleParamsTabProps {
|
||||
settings: TitleStyleSettings
|
||||
onUpdatePosition: (p: string) => void
|
||||
onUpdateFont: (f: string) => void
|
||||
onUpdateSize: (v: number) => void
|
||||
onToggleBold: () => void
|
||||
onToggleItalic: () => void
|
||||
onToggleStroke: () => void
|
||||
onToggleShadow: () => void
|
||||
onUpdatePatch: (patch: Partial<TitleStyleSettings>) => void
|
||||
positionOptions: PositionOption[]
|
||||
fontOptions?: FontOption[]
|
||||
/** 是否显示「封面」Tab(独立封面标题开关) */
|
||||
showCoverToggle?: boolean
|
||||
/** 封面独立标题开关状态 */
|
||||
coverEnabled?: boolean
|
||||
/** 封面开关变化 */
|
||||
onToggleCover?: (enabled: boolean) => void
|
||||
}
|
||||
|
||||
/* ── Slider 行 ── */
|
||||
const SliderRow: React.FC<{
|
||||
label: string
|
||||
value: number
|
||||
min: number
|
||||
max: number
|
||||
step?: number
|
||||
unit?: string
|
||||
onChange: (v: number) => void
|
||||
}> = ({ label, value, min, max, step = 1, unit = "px", onChange }) => (
|
||||
<div className="ts-form-field">
|
||||
<div className="ts-field-label-row">
|
||||
<label>{label}</label>
|
||||
<span className="ts-field-value">
|
||||
{value}
|
||||
{unit}
|
||||
</span>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
className="ts-slider"
|
||||
min={min}
|
||||
max={max}
|
||||
step={step}
|
||||
value={value}
|
||||
onChange={(e) => onChange(Number(e.target.value))}
|
||||
/>
|
||||
</div>
|
||||
)
|
||||
|
||||
/* ── 色板 ── */
|
||||
const ColorPicker: React.FC<{
|
||||
label?: string
|
||||
value: string
|
||||
palette: string[]
|
||||
onChange: (c: string) => void
|
||||
}> = ({ label, value, palette, onChange }) => {
|
||||
const [customOpen, setCustomOpen] = useState(false)
|
||||
return (
|
||||
<div className="ts-form-field">
|
||||
{label && <label>{label}</label>}
|
||||
<div className="ts-color-row">
|
||||
{palette.map((c) => (
|
||||
<button
|
||||
key={c}
|
||||
type="button"
|
||||
className={`ts-color-swatch${value.toLowerCase() === c.toLowerCase() ? " active" : ""}`}
|
||||
style={{ background: c }}
|
||||
onClick={() => onChange(c)}
|
||||
title={c}
|
||||
/>
|
||||
))}
|
||||
<button
|
||||
type="button"
|
||||
className="ts-color-swatch ts-color-custom"
|
||||
onClick={() => setCustomOpen((v) => !v)}
|
||||
title="自定义颜色"
|
||||
>
|
||||
+
|
||||
</button>
|
||||
<input
|
||||
type="color"
|
||||
className="ts-color-native"
|
||||
value={value.startsWith("rgba") ? "#000000" : value}
|
||||
onChange={(e) => {
|
||||
onChange(e.target.value)
|
||||
setCustomOpen(false)
|
||||
}}
|
||||
style={{
|
||||
opacity: customOpen ? 1 : 0,
|
||||
position: customOpen ? "static" : "absolute",
|
||||
pointerEvents: customOpen ? "auto" : "none",
|
||||
width: customOpen ? 28 : 0,
|
||||
height: customOpen ? 28 : 0,
|
||||
border: "none",
|
||||
padding: 0,
|
||||
cursor: "pointer",
|
||||
background: "transparent",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
<div style={{ fontSize: 11, color: "#9ca3af", marginTop: 2 }}>
|
||||
当前:<code style={{ fontSize: 11 }}>{value}</code>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
const TitleStyleParamsTab: React.FC<TitleStyleParamsTabProps> = ({
|
||||
settings,
|
||||
onUpdatePosition,
|
||||
onUpdateFont,
|
||||
onUpdateSize,
|
||||
onToggleBold,
|
||||
onToggleItalic,
|
||||
onToggleStroke,
|
||||
onToggleShadow,
|
||||
onUpdatePatch,
|
||||
positionOptions,
|
||||
fontOptions = FONT_OPTIONS,
|
||||
showCoverToggle = false,
|
||||
coverEnabled = false,
|
||||
onToggleCover,
|
||||
}) => {
|
||||
const upd = onUpdatePatch
|
||||
return (
|
||||
<Tabs
|
||||
size="small"
|
||||
defaultActiveKey="basic"
|
||||
items={[
|
||||
{
|
||||
key: "basic",
|
||||
label: "基础",
|
||||
children: (
|
||||
<>
|
||||
<div className="ts-row-2">
|
||||
<div className="ts-form-field ts-half">
|
||||
<label>位置</label>
|
||||
<select
|
||||
className="ts-select"
|
||||
value={settings.position}
|
||||
onChange={(e) => onUpdatePosition(e.target.value)}
|
||||
>
|
||||
{positionOptions.map((o) => (
|
||||
<option key={o.value} value={o.value}>
|
||||
{o.label}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</div>
|
||||
<div className="ts-form-field ts-half">
|
||||
<label>字体</label>
|
||||
<select
|
||||
className="ts-select"
|
||||
value={settings.font}
|
||||
onChange={(e) => onUpdateFont(e.target.value)}
|
||||
>
|
||||
{fontOptions.map((f) => (
|
||||
<option key={f.value} value={f.value}>
|
||||
{f.tag === "hot" ? "🔥 " : f.tag === "new" ? "🆕 " : ""}
|
||||
{f.label}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
<SliderRow
|
||||
label="字号"
|
||||
value={settings.size}
|
||||
min={16}
|
||||
max={120}
|
||||
onChange={onUpdateSize}
|
||||
/>
|
||||
<div className="ts-form-field">
|
||||
<label>样式</label>
|
||||
<div className="ts-style-btns">
|
||||
<button
|
||||
type="button"
|
||||
className={`ts-style-btn${settings.bold ? " active" : ""}`}
|
||||
onClick={onToggleBold}
|
||||
>
|
||||
<b>B</b>
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className={`ts-style-btn${settings.italic ? " active" : ""}`}
|
||||
onClick={onToggleItalic}
|
||||
>
|
||||
<i>I</i>
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className={`ts-style-btn${settings.stroke ? " active" : ""}`}
|
||||
onClick={() => {
|
||||
onToggleStroke()
|
||||
if (!settings.stroke && (settings.strokeWidth ?? 0) < 2)
|
||||
upd({ strokeWidth: 4 })
|
||||
}}
|
||||
title="描边"
|
||||
>
|
||||
S
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className={`ts-style-btn${settings.shadow ? " active" : ""}`}
|
||||
onClick={() => {
|
||||
onToggleShadow()
|
||||
if (!settings.shadow) {
|
||||
upd({
|
||||
shadowOffsetX: 2,
|
||||
shadowOffsetY: 2,
|
||||
shadowBlur: 4,
|
||||
shadowColor: "rgba(0,0,0,0.8)",
|
||||
})
|
||||
}
|
||||
}}
|
||||
title="阴影"
|
||||
>
|
||||
☁
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
<ColorPicker
|
||||
label="字色"
|
||||
value={settings.color}
|
||||
palette={TITLE_COLOR_PALETTE}
|
||||
onChange={(c) => upd({ color: c })}
|
||||
/>
|
||||
</>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "stroke",
|
||||
label: "描边",
|
||||
children: (
|
||||
<>
|
||||
<div className="ts-toggle-row">
|
||||
<label>
|
||||
<input type="checkbox" checked={settings.stroke} onChange={onToggleStroke} />
|
||||
启用描边
|
||||
</label>
|
||||
</div>
|
||||
{settings.stroke && (
|
||||
<>
|
||||
<SliderRow
|
||||
label="描边宽度"
|
||||
value={settings.strokeWidth ?? 4}
|
||||
min={0}
|
||||
max={20}
|
||||
onChange={(v) => upd({ strokeWidth: v })}
|
||||
/>
|
||||
<ColorPicker
|
||||
label="描边颜色"
|
||||
value={settings.strokeColor ?? "#000000"}
|
||||
palette={STROKE_COLOR_PALETTE}
|
||||
onChange={(c) => upd({ strokeColor: c })}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
</>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "shadow",
|
||||
label: "阴影",
|
||||
children: (
|
||||
<>
|
||||
<div className="ts-toggle-row">
|
||||
<label>
|
||||
<input type="checkbox" checked={settings.shadow} onChange={onToggleShadow} />
|
||||
启用阴影
|
||||
</label>
|
||||
</div>
|
||||
{settings.shadow && (
|
||||
<>
|
||||
<SliderRow
|
||||
label="X偏移"
|
||||
value={settings.shadowOffsetX ?? 2}
|
||||
min={-20}
|
||||
max={20}
|
||||
onChange={(v) => upd({ shadowOffsetX: v })}
|
||||
/>
|
||||
<SliderRow
|
||||
label="Y偏移"
|
||||
value={settings.shadowOffsetY ?? 2}
|
||||
min={-20}
|
||||
max={20}
|
||||
onChange={(v) => upd({ shadowOffsetY: v })}
|
||||
/>
|
||||
<SliderRow
|
||||
label="模糊半径"
|
||||
value={settings.shadowBlur ?? 4}
|
||||
min={0}
|
||||
max={30}
|
||||
onChange={(v) => upd({ shadowBlur: v })}
|
||||
/>
|
||||
<div className="ts-form-field">
|
||||
<label>阴影颜色</label>
|
||||
<input
|
||||
type="text"
|
||||
className="ts-input"
|
||||
value={settings.shadowColor ?? "rgba(0,0,0,0.8)"}
|
||||
onChange={(e) => upd({ shadowColor: e.target.value })}
|
||||
placeholder="rgba(0,0,0,0.8)"
|
||||
/>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "bg",
|
||||
label: "背景",
|
||||
children: (
|
||||
<>
|
||||
<div className="ts-toggle-row">
|
||||
<label>
|
||||
<input
|
||||
type="checkbox"
|
||||
checked={settings.bgEnabled}
|
||||
onChange={() => upd({ bgEnabled: !settings.bgEnabled })}
|
||||
/>
|
||||
启用背景色块
|
||||
</label>
|
||||
</div>
|
||||
{settings.bgEnabled && (
|
||||
<>
|
||||
<ColorPicker
|
||||
label="背景颜色(含透明度)"
|
||||
value={settings.bgColor}
|
||||
palette={BG_COLOR_PALETTE}
|
||||
onChange={(c) => upd({ bgColor: c })}
|
||||
/>
|
||||
<SliderRow
|
||||
label="内边距"
|
||||
value={settings.bgPadding}
|
||||
min={0}
|
||||
max={40}
|
||||
onChange={(v) => upd({ bgPadding: v })}
|
||||
/>
|
||||
<SliderRow
|
||||
label="圆角"
|
||||
value={settings.bgRadius}
|
||||
min={0}
|
||||
max={30}
|
||||
onChange={(v) => upd({ bgRadius: v })}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
</>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "layout",
|
||||
label: "排版",
|
||||
children: (
|
||||
<>
|
||||
<SliderRow
|
||||
label="每行最大字符数"
|
||||
value={settings.maxCharsPerLine ?? 0}
|
||||
min={0}
|
||||
max={20}
|
||||
unit=""
|
||||
onChange={(v) => upd({ maxCharsPerLine: v })}
|
||||
/>
|
||||
<div
|
||||
className="ts-form-field"
|
||||
style={{ fontSize: 11, color: "#9ca3af", marginTop: -4 }}
|
||||
>
|
||||
0 = 不自动换行(按 / 手动分行)
|
||||
</div>
|
||||
<SliderRow
|
||||
label="行距倍数"
|
||||
value={Math.round((settings.lineHeight ?? 1.2) * 100) / 100}
|
||||
min={1}
|
||||
max={2}
|
||||
step={0.05}
|
||||
unit=""
|
||||
onChange={(v) => upd({ lineHeight: Number(v.toFixed(2)) })}
|
||||
/>
|
||||
<SliderRow
|
||||
label="顶部边距"
|
||||
value={settings.marginTop ?? 24}
|
||||
min={0}
|
||||
max={200}
|
||||
onChange={(v) => upd({ marginTop: v })}
|
||||
/>
|
||||
</>
|
||||
),
|
||||
},
|
||||
...(showCoverToggle
|
||||
? [
|
||||
{
|
||||
key: "cover",
|
||||
label: "封面",
|
||||
children: (
|
||||
<div className="ts-toggle-row">
|
||||
<label>
|
||||
<input
|
||||
type="checkbox"
|
||||
checked={coverEnabled}
|
||||
onChange={(e) => onToggleCover?.(e.target.checked)}
|
||||
/>
|
||||
封面使用独立标题样式
|
||||
</label>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
]}
|
||||
/>
|
||||
)
|
||||
}
|
||||
|
||||
export default TitleStyleParamsTab
|
||||
@@ -1,29 +1,30 @@
|
||||
/**
|
||||
* 标题模板编辑器(v3 重构)
|
||||
* 标题模板编辑器(公共组件)
|
||||
*
|
||||
* - Modal 弹窗 860px 宽
|
||||
* - 左侧:300px 竖屏预览区(图片背景+暗色渐变遮罩+透明 Canvas 叠字)+ 模板名称输入框
|
||||
* - 右侧:参数 Tab 面板(基础/描边/阴影/背景/排版),复用 TitleStylePanel 的 paramsOnly 模式
|
||||
* - 左侧:300px 竖屏预览区(图片背景+暗角+透明 Canvas 叠字)+ 模板名称输入
|
||||
* - 右侧:参数 Tab 面板(基础/描边/阴影/背景/排版),复用 TitleStyleParamsTab
|
||||
* - 底部:取消 / 保存模板 按钮
|
||||
* - 内置模板编辑时保存会创建副本(带"副本"逻辑由 handleSave 处理)
|
||||
* - 内置模板编辑时保存会创建副本(带"副本"逻辑由 onSave 的调用方处理)
|
||||
*/
|
||||
import React, { useEffect, useMemo, useState } from "react"
|
||||
import { Modal, Button, Input, message } from "antd"
|
||||
import TitleStylePanel from "../../pages/generate/components/title/TitleStylePanel"
|
||||
import TitleMiniPreview from "../../pages/generate/components/title/TitleMiniPreview"
|
||||
import { POSITION_OPTIONS } from "../../pages/generate/constants"
|
||||
import { FONT_OPTIONS } from "./constants"
|
||||
import type { TitleSettings } from "../../pages/generate/types"
|
||||
import { DEFAULT_TITLE_SETTINGS_FULL } from "../../pages/generate/types"
|
||||
import type { TitleStyleSettings } from "./settings"
|
||||
import { DEFAULT_TITLE_STYLE_SETTINGS } from "./settings"
|
||||
import { titleStyleConfigToCamel, camelToTitleStyleConfig } from "./utils"
|
||||
import type { TitleTemplate } from "./template-types"
|
||||
import type { TitleStyleConfig } from "./types"
|
||||
import { POSITION_OPTIONS } from "./position-options"
|
||||
import { FONT_OPTIONS } from "./constants"
|
||||
import TitleMiniPreview from "./TitleMiniPreview"
|
||||
import TitleStyleParamsTab from "./TitleStyleParamsTab"
|
||||
import "./TitleTemplate.css"
|
||||
import "./TitleStylePanel.css"
|
||||
|
||||
interface Props {
|
||||
open: boolean
|
||||
template: TitleTemplate
|
||||
onClose: () => void
|
||||
/** 用户点击保存:将编辑结果回调给父组件(父组件统一做 CRUD,避免双 hook 实例不同步) */
|
||||
onSave: (data: { name: string; emoji: string; style: Partial<TitleStyleConfig> }) => void
|
||||
}
|
||||
|
||||
@@ -31,10 +32,9 @@ interface Props {
|
||||
const EDITOR_BG = "/title-templates/portrait1.jpg"
|
||||
|
||||
const TitleTemplateEditor: React.FC<Props> = ({ open, template, onClose, onSave }) => {
|
||||
const [settings, setSettings] = useState<TitleSettings>(() => ({
|
||||
...DEFAULT_TITLE_SETTINGS_FULL,
|
||||
const [settings, setSettings] = useState<TitleStyleSettings>(() => ({
|
||||
...DEFAULT_TITLE_STYLE_SETTINGS,
|
||||
...titleStyleConfigToCamel(template.style || {}),
|
||||
title: "预览标题文字",
|
||||
}))
|
||||
const [formName, setFormName] = useState(template.name || "")
|
||||
const [formEmoji, setFormEmoji] = useState(template.emoji || "✨")
|
||||
@@ -43,16 +43,15 @@ const TitleTemplateEditor: React.FC<Props> = ({ open, template, onClose, onSave
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setSettings({
|
||||
...DEFAULT_TITLE_SETTINGS_FULL,
|
||||
...DEFAULT_TITLE_STYLE_SETTINGS,
|
||||
...titleStyleConfigToCamel(template.style || {}),
|
||||
title: "预览标题文字",
|
||||
})
|
||||
setFormName(template.name || "")
|
||||
setFormEmoji(template.emoji || "✨")
|
||||
}
|
||||
}, [open, template])
|
||||
|
||||
const upd = (patch: Partial<TitleSettings>) => setSettings((s) => ({ ...s, ...patch }))
|
||||
const upd = (patch: Partial<TitleStyleSettings>) => setSettings((s) => ({ ...s, ...patch }))
|
||||
|
||||
const handleSave = () => {
|
||||
const name = formName.trim()
|
||||
@@ -69,11 +68,11 @@ const TitleTemplateEditor: React.FC<Props> = ({ open, template, onClose, onSave
|
||||
}
|
||||
}
|
||||
|
||||
// 编辑器内的预览用 settings:字号适配竖屏
|
||||
const previewSettings = useMemo<TitleSettings>(() => {
|
||||
// 竖屏宽度 200px,按比例缩放字号,让预览看起来协调
|
||||
return { ...settings, size: Math.round(settings.size * 0.55) }
|
||||
}, [settings])
|
||||
// 编辑器预览 settings:竖屏宽度 200px,字号按比例缩放
|
||||
const previewSettings = useMemo<TitleStyleSettings>(
|
||||
() => ({ ...settings, size: Math.round(settings.size * 0.55) }),
|
||||
[settings],
|
||||
)
|
||||
|
||||
return (
|
||||
<Modal
|
||||
@@ -140,7 +139,7 @@ const TitleTemplateEditor: React.FC<Props> = ({ open, template, onClose, onSave
|
||||
</div>
|
||||
{/* 右侧:参数 Tab */}
|
||||
<div className="ttv3-editor-right">
|
||||
<TitleStylePanel
|
||||
<TitleStyleParamsTab
|
||||
settings={settings}
|
||||
onUpdatePosition={(p) => upd({ position: p, posX: null, posY: null })}
|
||||
onUpdateFont={(f) => upd({ font: f })}
|
||||
@@ -155,15 +154,9 @@ const TitleTemplateEditor: React.FC<Props> = ({ open, template, onClose, onSave
|
||||
})
|
||||
}
|
||||
onToggleShadow={() => upd({ shadow: !settings.shadow })}
|
||||
onApplyPreset={() => {
|
||||
/* 编辑器内不使用系统预设快捷键 */
|
||||
}}
|
||||
onUpdateStyle={(patch) => upd(patch)}
|
||||
activePreset={null}
|
||||
titlePresets={[]}
|
||||
POSITION_OPTIONS={POSITION_OPTIONS}
|
||||
FONT_OPTIONS={FONT_OPTIONS}
|
||||
paramsOnly
|
||||
onUpdatePatch={upd}
|
||||
positionOptions={POSITION_OPTIONS}
|
||||
fontOptions={FONT_OPTIONS}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -0,0 +1,343 @@
|
||||
/**
|
||||
* 标题模板选择器 — 大卡片网格(共享组件)
|
||||
*
|
||||
* 渲染「我的模板」+「系统模板」两个分组的 3:4 竖版大圆角卡片:
|
||||
* - 卡片上半:示例背景图 + vignette 暗角 + 透明 Canvas 大字预览
|
||||
* - 卡片下半:emoji + 名称 + 系统/我的标签 + 始终可见的编辑/复制/导出/删除按钮
|
||||
* - 选中紫色边框;右上角「新建模板」按钮;点编辑/新建弹 TitleTemplateEditor
|
||||
*
|
||||
* Props 通用化,不耦合业务 state。
|
||||
*/
|
||||
import React, { useCallback, useMemo, useState } from "react"
|
||||
import { Button, message, Popconfirm } from "antd"
|
||||
import {
|
||||
PlusOutlined,
|
||||
EditOutlined,
|
||||
CopyOutlined,
|
||||
DeleteOutlined,
|
||||
ExportOutlined,
|
||||
CheckOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import type { TitleTemplate } from "./template-types"
|
||||
import type { TitleStyleSettings } from "./settings"
|
||||
import { DEFAULT_TITLE_STYLE_SETTINGS } from "./settings"
|
||||
import {
|
||||
titleStyleConfigToCamel,
|
||||
camelToTitleStyleConfig,
|
||||
templateToPreviewSettings,
|
||||
} from "./utils"
|
||||
import { useTitleTemplates } from "./useTitleTemplates"
|
||||
import TitleMiniPreview from "./TitleMiniPreview"
|
||||
import TitleTemplateEditor from "./TitleTemplateEditor"
|
||||
import "./TitleTemplate.css"
|
||||
import "./TitleStylePanel.css"
|
||||
|
||||
export interface TitleTemplateSelectorProps {
|
||||
/** 当前选中模板 id(受控) */
|
||||
value?: string | null
|
||||
/** 选中模板时回调(templateId, fullStyleSettings, template) */
|
||||
onChange?: (templateId: string, style: TitleStyleSettings, template: TitleTemplate) => void
|
||||
/** 是否显示编辑器入口(新建/编辑按钮),默认 true */
|
||||
showEditor?: boolean
|
||||
/** 显示哪些分组,默认全部 */
|
||||
categories?: Array<"system" | "custom">
|
||||
/** 使用场景标识(仅作 data-attr,不影响样式) */
|
||||
context?: string
|
||||
}
|
||||
|
||||
/* ── 卡片预览背景图池(按 index 轮换) ── */
|
||||
const PREVIEW_BG_IMAGES = [
|
||||
"/title-templates/portrait1.jpg",
|
||||
"/title-templates/portrait2.jpg",
|
||||
"/title-templates/scene1.jpg",
|
||||
]
|
||||
|
||||
/* ── 预览容器:用 ref 测量宽度后再渲染透明 Canvas,保证文字清晰 ── */
|
||||
const FillPreview: React.FC<{
|
||||
settings: TitleStyleSettings
|
||||
sampleText: string
|
||||
portrait?: boolean
|
||||
}> = ({ settings, sampleText, portrait }) => {
|
||||
const [w, setW] = useState(0)
|
||||
// 首次挂载后测量一次
|
||||
const setRef = useCallback((el: HTMLDivElement | null) => {
|
||||
if (el) setW(Math.floor(el.clientWidth))
|
||||
}, [])
|
||||
return (
|
||||
<div ref={setRef} className="tt-fill-canvas-wrap">
|
||||
{w > 0 && (
|
||||
<TitleMiniPreview
|
||||
settings={settings}
|
||||
width={w}
|
||||
sampleText={sampleText}
|
||||
transparent
|
||||
portrait={portrait}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
const TitleTemplateSelector: React.FC<TitleTemplateSelectorProps> = ({
|
||||
value,
|
||||
onChange,
|
||||
showEditor = true,
|
||||
categories = ["system", "custom"],
|
||||
context,
|
||||
}) => {
|
||||
const {
|
||||
templates,
|
||||
createTemplate,
|
||||
duplicateTemplate,
|
||||
updateTemplate,
|
||||
deleteTemplate,
|
||||
exportTemplate,
|
||||
} = useTitleTemplates()
|
||||
|
||||
const [editingTemplate, setEditingTemplate] = useState<TitleTemplate | null>(null)
|
||||
const [editorOpen, setEditorOpen] = useState(false)
|
||||
|
||||
const grouped = useMemo(
|
||||
() => ({
|
||||
builtin: templates.filter((t) => t.isBuiltin),
|
||||
custom: templates.filter((t) => !t.isBuiltin),
|
||||
}),
|
||||
[templates],
|
||||
)
|
||||
|
||||
const showSys = categories.includes("system")
|
||||
const showMine = categories.includes("custom")
|
||||
|
||||
/* ── 选中模板:合成完整 TitleStyleSettings 回调给父组件 ── */
|
||||
const handleSelectTemplate = useCallback(
|
||||
(tpl: TitleTemplate) => {
|
||||
const full: TitleStyleSettings = {
|
||||
...DEFAULT_TITLE_STYLE_SETTINGS,
|
||||
...titleStyleConfigToCamel(tpl.style),
|
||||
}
|
||||
onChange?.(tpl.id, full, tpl)
|
||||
},
|
||||
[onChange],
|
||||
)
|
||||
|
||||
const handleRequestCreate = useCallback(() => {
|
||||
// 新建:以当前选中模板样式为起点,否则用默认样式
|
||||
let base: TitleStyleSettings = DEFAULT_TITLE_STYLE_SETTINGS
|
||||
if (value) {
|
||||
const sel = templates.find((t) => t.id === value)
|
||||
if (sel) {
|
||||
base = { ...DEFAULT_TITLE_STYLE_SETTINGS, ...titleStyleConfigToCamel(sel.style) }
|
||||
}
|
||||
}
|
||||
const draft: TitleTemplate = {
|
||||
id: "",
|
||||
name: "我的标题模板",
|
||||
emoji: "✨",
|
||||
isBuiltin: false,
|
||||
style: camelToTitleStyleConfig({
|
||||
...base,
|
||||
position: base.position === "custom" ? "bottom" : base.position,
|
||||
}),
|
||||
createdAt: new Date().toISOString(),
|
||||
updatedAt: new Date().toISOString(),
|
||||
}
|
||||
setEditingTemplate(draft)
|
||||
setEditorOpen(true)
|
||||
}, [value, templates])
|
||||
|
||||
const handleRequestEdit = useCallback((tpl: TitleTemplate) => {
|
||||
setEditingTemplate(tpl)
|
||||
setEditorOpen(true)
|
||||
}, [])
|
||||
|
||||
const handleDuplicate = useCallback(
|
||||
(t: TitleTemplate) => {
|
||||
const dup = duplicateTemplate(t.id)
|
||||
if (dup) message.success(`已复制:${dup.name}`)
|
||||
},
|
||||
[duplicateTemplate],
|
||||
)
|
||||
const handleDelete = useCallback(
|
||||
(t: TitleTemplate) => {
|
||||
deleteTemplate(t.id)
|
||||
message.success("已删除模板")
|
||||
},
|
||||
[deleteTemplate],
|
||||
)
|
||||
const handleExport = useCallback(
|
||||
(t: TitleTemplate) => {
|
||||
const json = exportTemplate(t.id)
|
||||
if (!json) return
|
||||
const blob = new Blob([json], { type: "application/json" })
|
||||
const url = URL.createObjectURL(blob)
|
||||
const a = document.createElement("a")
|
||||
a.href = url
|
||||
a.download = `${t.name}.title-template.json`
|
||||
a.click()
|
||||
URL.revokeObjectURL(url)
|
||||
},
|
||||
[exportTemplate],
|
||||
)
|
||||
|
||||
const handleEditorSave = useCallback(
|
||||
(data: { name: string; emoji: string; style: Partial<import("./types").TitleStyleConfig> }) => {
|
||||
if (!editingTemplate) return
|
||||
let saved: TitleTemplate
|
||||
if (editingTemplate.isBuiltin || !editingTemplate.id) {
|
||||
saved = createTemplate({ name: data.name, emoji: data.emoji, style: data.style })
|
||||
} else {
|
||||
updateTemplate(editingTemplate.id, {
|
||||
name: data.name,
|
||||
emoji: data.emoji,
|
||||
style: data.style,
|
||||
})
|
||||
saved = {
|
||||
...editingTemplate,
|
||||
name: data.name,
|
||||
emoji: data.emoji,
|
||||
style: data.style,
|
||||
updatedAt: new Date().toISOString(),
|
||||
}
|
||||
}
|
||||
setEditorOpen(false)
|
||||
setEditingTemplate(null)
|
||||
message.success(`已保存:${data.name}`)
|
||||
handleSelectTemplate(saved)
|
||||
},
|
||||
[editingTemplate, createTemplate, updateTemplate, handleSelectTemplate],
|
||||
)
|
||||
|
||||
/* ── 渲染单张大卡片 ── */
|
||||
const renderCard = (t: TitleTemplate, idx: number, section: "mine" | "sys") => {
|
||||
const isSelected = value === t.id
|
||||
const bgIdx = idx % PREVIEW_BG_IMAGES.length
|
||||
const bgImg = PREVIEW_BG_IMAGES[bgIdx]
|
||||
const preview = templateToPreviewSettings(t, 42)
|
||||
return (
|
||||
<div
|
||||
key={t.id}
|
||||
className={`ttv3-card${isSelected ? " selected" : ""}`}
|
||||
onClick={() => handleSelectTemplate(t)}
|
||||
data-context={context}
|
||||
>
|
||||
<div className="ttv3-preview">
|
||||
<img className="ttv3-bg" src={bgImg} alt="" />
|
||||
<div className="ttv3-vignette" />
|
||||
<FillPreview settings={preview} sampleText="预览标题文字" portrait />
|
||||
<span className={`ttv3-badge ttv3-badge--${section}`}>
|
||||
{section === "sys" ? "系统" : "我的"}
|
||||
</span>
|
||||
<span className={`ttv3-check${isSelected ? " on" : ""}`}>
|
||||
{isSelected && <CheckOutlined />}
|
||||
</span>
|
||||
</div>
|
||||
<div className="ttv3-footer">
|
||||
<div className="ttv3-name-row">
|
||||
<span className="ttv3-emoji">{t.emoji || "✨"}</span>
|
||||
<span className="ttv3-name" title={t.name}>
|
||||
{t.name}
|
||||
</span>
|
||||
<span className={`ttv3-tag ttv3-tag--${section}`}>
|
||||
{section === "sys" ? "系统" : "我的"}
|
||||
</span>
|
||||
</div>
|
||||
{showEditor && (
|
||||
<div className="ttv3-actions" onClick={(e) => e.stopPropagation()}>
|
||||
<button
|
||||
type="button"
|
||||
className="ttv3-act ttv3-act--primary"
|
||||
disabled={t.isBuiltin}
|
||||
onClick={() => handleRequestEdit(t)}
|
||||
title={t.isBuiltin ? "系统模板不可编辑,点击复制后可编辑" : "编辑"}
|
||||
>
|
||||
<EditOutlined /> 编辑
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="ttv3-act"
|
||||
onClick={() => handleDuplicate(t)}
|
||||
title="复制"
|
||||
>
|
||||
<CopyOutlined /> 复制
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="ttv3-act"
|
||||
onClick={() => handleExport(t)}
|
||||
title="导出"
|
||||
>
|
||||
<ExportOutlined /> 导出
|
||||
</button>
|
||||
<Popconfirm title="删除该模板?" onConfirm={() => handleDelete(t)}>
|
||||
<button
|
||||
type="button"
|
||||
className="ttv3-act ttv3-act--danger"
|
||||
disabled={t.isBuiltin}
|
||||
title={t.isBuiltin ? "系统模板不可删除" : "删除"}
|
||||
>
|
||||
<DeleteOutlined /> 删除
|
||||
</button>
|
||||
</Popconfirm>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="xx-title-style-section ttv3-panel">
|
||||
<div className="ttv3-header">
|
||||
<span className="ttv3-title">标题模板</span>
|
||||
{showEditor && (
|
||||
<Button
|
||||
type="primary"
|
||||
size="small"
|
||||
icon={<PlusOutlined />}
|
||||
onClick={handleRequestCreate}
|
||||
className="ttv3-new-btn"
|
||||
>
|
||||
新建模板
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{showMine && (
|
||||
<div className="ttv3-section">
|
||||
<div className="ttv3-section-label">我的模板</div>
|
||||
{grouped.custom.length === 0 ? (
|
||||
<div className="ttv3-empty">
|
||||
<div className="ttv3-empty-icon">✨</div>
|
||||
<div className="ttv3-empty-text">还没有自定义模板,点右上角「新建模板」创建</div>
|
||||
</div>
|
||||
) : (
|
||||
<div className="ttv3-grid">
|
||||
{grouped.custom.map((t, i) => renderCard(t, i, "mine"))}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{showSys && (
|
||||
<div className="ttv3-section">
|
||||
<div className="ttv3-section-label">系统模板</div>
|
||||
<div className="ttv3-grid">{grouped.builtin.map((t, i) => renderCard(t, i, "sys"))}</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{showEditor && editorOpen && editingTemplate && (
|
||||
<TitleTemplateEditor
|
||||
open={editorOpen}
|
||||
template={editingTemplate}
|
||||
onClose={() => {
|
||||
setEditorOpen(false)
|
||||
setEditingTemplate(null)
|
||||
}}
|
||||
onSave={handleEditorSave}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default TitleTemplateSelector
|
||||
@@ -0,0 +1,19 @@
|
||||
/**
|
||||
* 公共标题模板/样式组件统一导出
|
||||
*
|
||||
* 任何页面需要标题样式配置/模板选择/模板编辑,从这里 import,
|
||||
* 不要直接 import pages/generate/components/title/* 下的内部组件。
|
||||
*/
|
||||
export { default as TitleTemplateSelector } from "./TitleTemplateSelector"
|
||||
export { default as TitleTemplateEditor } from "./TitleTemplateEditor"
|
||||
export { default as TitleStyleParamsTab } from "./TitleStyleParamsTab"
|
||||
export { default as TitleMiniPreview } from "./TitleMiniPreview"
|
||||
export { useTitleTemplates } from "./useTitleTemplates"
|
||||
export * from "./constants"
|
||||
export * from "./types"
|
||||
export * from "./template-types"
|
||||
export * from "./settings"
|
||||
export * from "./utils"
|
||||
export { POSITION_OPTIONS } from "./position-options"
|
||||
export type { PositionOption, FontOption, TitleStyleParamsTabProps } from "./TitleStyleParamsTab"
|
||||
export type { TitleTemplateSelectorProps } from "./TitleTemplateSelector"
|
||||
@@ -0,0 +1,14 @@
|
||||
/**
|
||||
* 标题位置选项(公共常量)
|
||||
*/
|
||||
export interface PositionOption {
|
||||
value: string
|
||||
label: string
|
||||
}
|
||||
|
||||
export const POSITION_OPTIONS: PositionOption[] = [
|
||||
{ value: "top", label: "顶部" },
|
||||
{ value: "center", label: "居中" },
|
||||
{ value: "bottom", label: "底部" },
|
||||
{ value: "custom", label: "自定义" },
|
||||
]
|
||||
@@ -0,0 +1,64 @@
|
||||
/**
|
||||
* 标题样式设置 — 公共 camelCase 类型与默认值
|
||||
*
|
||||
* 本文件是 @/components/title 公共包的唯一样式类型出口,不依赖任何业务页面(generate/ai-avatar)的私有类型。
|
||||
* - 字段与后端 snake_case TitleStyleConfig 一一对应(camelCase 版本)
|
||||
* - DEFAULT_TITLE_STYLE_SETTINGS 用于组件内部补全默认值
|
||||
* - aiAutoSelect / title / coverTitle 等业务状态不在本类型中——它们属于页面业务 state
|
||||
*/
|
||||
import type { TitleLineOverride } from "./types"
|
||||
|
||||
export interface TitleStyleSettings {
|
||||
position: string
|
||||
font: string
|
||||
size: number
|
||||
bold: boolean
|
||||
italic: boolean
|
||||
stroke: boolean
|
||||
shadow: boolean
|
||||
color: string
|
||||
posX: number | null
|
||||
posY: number | null
|
||||
lineHeight: number
|
||||
marginTop: number
|
||||
maxCharsPerLine: number
|
||||
strokeWidth: number
|
||||
strokeColor: string
|
||||
shadowOffsetX: number
|
||||
shadowOffsetY: number
|
||||
shadowBlur: number
|
||||
shadowColor: string
|
||||
bgEnabled: boolean
|
||||
bgColor: string
|
||||
bgPadding: number
|
||||
bgRadius: number
|
||||
lineOverrides: TitleLineOverride[]
|
||||
}
|
||||
|
||||
/** 公共默认样式(经典白字黑描边) */
|
||||
export const DEFAULT_TITLE_STYLE_SETTINGS: TitleStyleSettings = {
|
||||
position: "bottom",
|
||||
font: "思源黑体",
|
||||
size: 56,
|
||||
bold: true,
|
||||
italic: false,
|
||||
stroke: true,
|
||||
shadow: false,
|
||||
color: "#ffffff",
|
||||
posX: null,
|
||||
posY: null,
|
||||
lineHeight: 1.2,
|
||||
marginTop: 24,
|
||||
maxCharsPerLine: 10,
|
||||
strokeWidth: 5,
|
||||
strokeColor: "#000000",
|
||||
shadowOffsetX: 2,
|
||||
shadowOffsetY: 2,
|
||||
shadowBlur: 4,
|
||||
shadowColor: "rgba(0,0,0,0.8)",
|
||||
bgEnabled: false,
|
||||
bgColor: "rgba(0,0,0,0.5)",
|
||||
bgPadding: 12,
|
||||
bgRadius: 8,
|
||||
lineOverrides: [],
|
||||
}
|
||||
@@ -1,24 +1,25 @@
|
||||
/**
|
||||
* 标题样式工具(#2001 / 模板系统 #2003)
|
||||
*
|
||||
* - snake_case TitleStyleConfig ↔ camelCase TitleSettings 互转
|
||||
* - snake_case TitleStyleConfig <-> camelCase TitleStyleSettings 互转
|
||||
* - preset 归一化预览(修复"标题"两字大小不一)
|
||||
* - template -> preview settings 转换
|
||||
*/
|
||||
import type { TitleStyleConfig } from "./types"
|
||||
import type { TitleSettings } from "../../pages/generate/types"
|
||||
import type { TitleStyleSettings } from "./settings"
|
||||
import { DEFAULT_TITLE_STYLE_SETTINGS } from "./settings"
|
||||
import { TITLE_PRESETS } from "./constants"
|
||||
import { DEFAULT_TITLE_SETTINGS_FULL } from "../../pages/generate/types"
|
||||
import type { TitleTemplate } from "./template-types"
|
||||
|
||||
/** snake_case TitleStyleConfig → camelCase TitleSettings(仅覆盖已知字段) */
|
||||
export function titleStyleConfigToCamel(s: Partial<TitleStyleConfig>): Partial<TitleSettings> {
|
||||
const out: Partial<TitleSettings> = {}
|
||||
/** snake_case TitleStyleConfig -> camelCase TitleStyleSettings(仅覆盖已知字段) */
|
||||
export function titleStyleConfigToCamel(s: Partial<TitleStyleConfig>): Partial<TitleStyleSettings> {
|
||||
const out: Partial<TitleStyleSettings> = {}
|
||||
if (s.font != null) out.font = s.font
|
||||
if (s.size != null) out.size = s.size
|
||||
if (s.color != null) out.color = s.color
|
||||
if (s.bold != null) out.bold = s.bold
|
||||
if (s.italic != null) out.italic = s.italic
|
||||
if (s.position != null) out.position = s.position as TitleSettings["position"]
|
||||
if (s.position != null) out.position = s.position
|
||||
if (s.pos_x != null) out.posX = s.pos_x
|
||||
if (s.pos_y != null) out.posY = s.pos_y
|
||||
if (s.line_height != null) out.lineHeight = s.line_height
|
||||
@@ -40,8 +41,8 @@ export function titleStyleConfigToCamel(s: Partial<TitleStyleConfig>): Partial<T
|
||||
return out
|
||||
}
|
||||
|
||||
/** camelCase TitleSettings patch → snake_case TitleStyleConfig patch */
|
||||
export function camelToTitleStyleConfig(p: Partial<TitleSettings>): Partial<TitleStyleConfig> {
|
||||
/** camelCase TitleStyleSettings patch -> snake_case TitleStyleConfig patch */
|
||||
export function camelToTitleStyleConfig(p: Partial<TitleStyleSettings>): Partial<TitleStyleConfig> {
|
||||
const out: Partial<TitleStyleConfig> = {}
|
||||
if (p.font != null) out.font = p.font
|
||||
if (p.size != null) out.size = p.size
|
||||
@@ -71,15 +72,15 @@ export function camelToTitleStyleConfig(p: Partial<TitleSettings>): Partial<Titl
|
||||
}
|
||||
|
||||
/**
|
||||
* 把 preset style(snake_case)归一化为固定字号的 TitleSettings,
|
||||
* 把 preset style(snake_case)归一化为固定字号的 TitleStyleSettings,
|
||||
* 用于"预设卡片"缩略预览——所有卡片视觉上"标题"两字大小一致,便于辨识。
|
||||
* 描边/阴影/背景padding 按 fixedSize / 原始 size 比例缩放,避免粗描边爆框。
|
||||
*/
|
||||
export function buildPresetPreviewSettings(
|
||||
base: TitleSettings,
|
||||
base: TitleStyleSettings,
|
||||
presetKey: string,
|
||||
fixedSize = 56,
|
||||
): TitleSettings {
|
||||
): TitleStyleSettings {
|
||||
const preset = TITLE_PRESETS.find((p) => p.key === presetKey)
|
||||
if (!preset) return base
|
||||
const origSize = preset.style.size ?? fixedSize
|
||||
@@ -87,25 +88,25 @@ export function buildPresetPreviewSettings(
|
||||
const scale = (v: number | undefined, fallback: number): number =>
|
||||
v != null ? Math.round(v * ratio) : fallback
|
||||
return {
|
||||
...DEFAULT_TITLE_STYLE_SETTINGS,
|
||||
...base,
|
||||
...titleStyleConfigToCamel(preset.style),
|
||||
size: fixedSize,
|
||||
strokeWidth: scale(preset.style.stroke_width, base.strokeWidth) ?? base.strokeWidth,
|
||||
shadowOffsetX: scale(preset.style.shadow_offset_x, base.shadowOffsetX) ?? base.shadowOffsetX,
|
||||
shadowOffsetY: scale(preset.style.shadow_offset_y, base.shadowOffsetY) ?? base.shadowOffsetY,
|
||||
shadowBlur: scale(preset.style.shadow_blur, base.shadowBlur) ?? base.shadowBlur,
|
||||
bgPadding: scale(preset.style.bg_padding, base.bgPadding) ?? base.bgPadding,
|
||||
strokeWidth: scale(preset.style.stroke_width, base.strokeWidth),
|
||||
shadowOffsetX: scale(preset.style.shadow_offset_x, base.shadowOffsetX),
|
||||
shadowOffsetY: scale(preset.style.shadow_offset_y, base.shadowOffsetY),
|
||||
shadowBlur: scale(preset.style.shadow_blur, base.shadowBlur),
|
||||
bgPadding: scale(preset.style.bg_padding, base.bgPadding),
|
||||
lineOverrides: [],
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 把 TitleTemplate 渲染为完整 TitleSettings(带默认值),用于卡片预览。
|
||||
* 与模板选择器中保持一致,抽出共用。
|
||||
* 把 TitleTemplate 渲染为完整 TitleStyleSettings(带默认值),用于卡片预览。
|
||||
*/
|
||||
export function templateToPreviewSettings(t: TitleTemplate, fixedSize = 48): TitleSettings {
|
||||
const base: TitleSettings = {
|
||||
...DEFAULT_TITLE_SETTINGS_FULL,
|
||||
export function templateToPreviewSettings(t: TitleTemplate, fixedSize = 48): TitleStyleSettings {
|
||||
const base: TitleStyleSettings = {
|
||||
...DEFAULT_TITLE_STYLE_SETTINGS,
|
||||
...titleStyleConfigToCamel(t.style),
|
||||
}
|
||||
// 预览时用固定字号保证所有卡片字大小一致;描边/阴影/padding按比例缩放
|
||||
|
||||
@@ -412,7 +412,7 @@ const PanelTitleConfig: React.FC<PanelTitleConfigProps> = ({ titleConfig, onUpda
|
||||
onUpdateStyle={handleUpdateStyle}
|
||||
showCoverToggle
|
||||
previewWidth={280}
|
||||
enableTemplates={false}
|
||||
enableTemplates={true}
|
||||
selectedTemplateId={selectedTemplateId}
|
||||
onApplyTemplate={handleApplyTemplate}
|
||||
activePreset={activePreset}
|
||||
|
||||
@@ -239,9 +239,14 @@ const GeneratePage: React.FC = () => {
|
||||
voiceModePerVideo,
|
||||
variantCoverUrls: previewCovers,
|
||||
selectedVariantIndexes: isBatch ? selectedVariantIds : undefined,
|
||||
onGenerationSuccess: () => {
|
||||
onGenerationSuccess: (status?: "completed" | "awaiting_cover") => {
|
||||
setPreviewTaskId(null)
|
||||
setStoredSourceEditPlanId(null)
|
||||
// #2088:渲染完成后自动跳到封面选择页(step 5),不再等用户手动点「下一步」
|
||||
// awaiting_cover 和 completed 都走封面页(completed 是旧 worker 或 finalize 后状态,仍支持选封面)
|
||||
if (status === "awaiting_cover" || status === "completed" || !status) {
|
||||
setCurrentStep(5)
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
|
||||
@@ -1,271 +1 @@
|
||||
/**
|
||||
* 标题迷你 Canvas 预览(#2001)
|
||||
*
|
||||
* 渲染一张指定宽度的小 Canvas 预览标题效果,用于:
|
||||
* - 预设卡片缩略图
|
||||
* - 样式面板顶部的实时预览
|
||||
*
|
||||
* 与 titleCanvas.ts 渲染逻辑保持一致,但:
|
||||
* - 固定分辨率(width × 宽高比约 2:1)
|
||||
* - 不调用 ffmpeg,只做视觉预览
|
||||
* - 支持背景色块、描边宽度/颜色、阴影参数化、行距、自动换行
|
||||
*/
|
||||
import React, { useEffect, useRef } from "react"
|
||||
import type { TitleSettings } from "../../types"
|
||||
import { getFontFamily } from "@/components/title/constants"
|
||||
|
||||
interface Props {
|
||||
settings: TitleSettings
|
||||
width?: number
|
||||
sampleText?: string
|
||||
/** 背景(预览用,默认深色渐变模拟视频底),transparent=true 时忽略 */
|
||||
background?: string
|
||||
/** 高度(可选,默认按 portrait 选比例) */
|
||||
height?: number
|
||||
/** 透明背景(卡片/编辑器预览叠加在图片上时使用) */
|
||||
transparent?: boolean
|
||||
/** 纵向竖屏预览(9:16),true 时 aspect=16/9 适配手机视频比例 */
|
||||
portrait?: boolean
|
||||
}
|
||||
|
||||
/** 按 maxCharsPerLine 自动换行 */
|
||||
function wrapLines(text: string, maxChars: number): string[] {
|
||||
const manual = text
|
||||
.split(/[//\n]/)
|
||||
.map((l) => l.trim())
|
||||
.filter(Boolean)
|
||||
if (!maxChars || maxChars <= 0) return manual
|
||||
const out: string[] = []
|
||||
for (const line of manual) {
|
||||
if (line.length <= maxChars) {
|
||||
out.push(line)
|
||||
continue
|
||||
}
|
||||
let cur = ""
|
||||
for (const ch of line) {
|
||||
cur += ch
|
||||
if (cur.length >= maxChars) {
|
||||
out.push(cur)
|
||||
cur = ""
|
||||
}
|
||||
}
|
||||
if (cur) out.push(cur)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
const TitleMiniPreview: React.FC<Props> = ({
|
||||
settings,
|
||||
width = 200,
|
||||
sampleText,
|
||||
background = "linear-gradient(135deg,#1f2937,#111827)",
|
||||
height,
|
||||
transparent = false,
|
||||
portrait = false,
|
||||
}) => {
|
||||
const canvasRef = useRef<HTMLCanvasElement>(null)
|
||||
const h = height ?? Math.round(width * (portrait ? 16 / 9 : 1 / 1.8))
|
||||
const text = (sampleText || settings.title || "预览标题").trim() || "预览标题"
|
||||
|
||||
useEffect(() => {
|
||||
let cancelled = false
|
||||
const draw = () => {
|
||||
if (cancelled) return
|
||||
const cvs = canvasRef.current
|
||||
if (!cvs) return
|
||||
const dpr = window.devicePixelRatio || 1
|
||||
cvs.width = width * dpr
|
||||
cvs.height = h * dpr
|
||||
cvs.style.width = `${width}px`
|
||||
cvs.style.height = `${h}px`
|
||||
const ctx = cvs.getContext("2d")
|
||||
if (!ctx) return
|
||||
ctx.scale(dpr, dpr)
|
||||
ctx.clearRect(0, 0, width, h)
|
||||
|
||||
// 背景(transparent 时跳过,用于叠加在图片上)
|
||||
if (!transparent) {
|
||||
ctx.fillStyle = "#111827"
|
||||
ctx.fillRect(0, 0, width, h)
|
||||
}
|
||||
|
||||
// 分辨率缩放:以 360 宽为基准(对应 720p 的一半),与外层 previewScale/previewR 保持一致
|
||||
const r = previewR
|
||||
|
||||
// 字体
|
||||
const size = r(settings.size)
|
||||
const ff = getFontFamily(settings.font)
|
||||
const parts: string[] = []
|
||||
if (settings.italic) parts.push("italic")
|
||||
if (settings.bold) parts.push("bold")
|
||||
parts.push(`${size}px`, ff)
|
||||
ctx.font = parts.join(" ")
|
||||
ctx.textAlign = "center"
|
||||
ctx.textBaseline = "middle"
|
||||
ctx.fillStyle = settings.color
|
||||
ctx.lineJoin = "round"
|
||||
|
||||
// 阴影
|
||||
const shadowEnabled = !!settings.shadow
|
||||
const prevShadow = {
|
||||
c: ctx.shadowColor,
|
||||
b: ctx.shadowBlur,
|
||||
ox: ctx.shadowOffsetX,
|
||||
oy: ctx.shadowOffsetY,
|
||||
}
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
|
||||
// 换行
|
||||
const lines = wrapLines(text, settings.maxCharsPerLine ?? 0)
|
||||
const lineH = size * (settings.lineHeight ?? 1.2)
|
||||
const totalH = lines.length * lineH
|
||||
let startY: number
|
||||
if (settings.position === "top") {
|
||||
startY = size / 2 + r(settings.marginTop ?? 24)
|
||||
} else if (settings.position === "center") {
|
||||
startY = h / 2 - totalH / 2 + size / 2
|
||||
} else {
|
||||
// bottom
|
||||
const botMargin = portrait ? r(24) : r(16)
|
||||
startY = h - totalH - botMargin + size / 2
|
||||
}
|
||||
let centerX = width / 2
|
||||
if (settings.position === "custom" && settings.posX != null) {
|
||||
centerX = (settings.posX / 100) * width
|
||||
}
|
||||
|
||||
// 背景块
|
||||
if (settings.bgEnabled) {
|
||||
const pad = r(settings.bgPadding ?? 12)
|
||||
const rad = r(settings.bgRadius ?? 8)
|
||||
let maxLineW = 0
|
||||
for (const l of lines) {
|
||||
const m = ctx.measureText(l)
|
||||
if (m.width > maxLineW) maxLineW = m.width
|
||||
}
|
||||
const bw = maxLineW + pad * 2
|
||||
const bh = totalH + pad * 2
|
||||
const bx = centerX - bw / 2
|
||||
const by = startY - size / 2 - pad + (size - lineH) / 2
|
||||
ctx.shadowColor = "rgba(0,0,0,0)"
|
||||
ctx.shadowBlur = 0
|
||||
ctx.fillStyle = settings.bgColor ?? "rgba(0,0,0,0.5)"
|
||||
roundRect(ctx, bx, by, bw, bh, rad)
|
||||
ctx.fill()
|
||||
// 关键修复:画完背景块后必须把 fillStyle 重置为文字颜色,
|
||||
// 否则后续 fillText 会用 bgColor 填充文字,导致「文字看不见只剩色块」
|
||||
ctx.fillStyle = settings.color
|
||||
// 恢复阴影
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
}
|
||||
|
||||
// 描边(先画,再画填充)
|
||||
const strokeEnabled = !!settings.stroke && (settings.strokeWidth ?? 0) > 0
|
||||
lines.forEach((line, i) => {
|
||||
const y = startY + i * lineH
|
||||
if (strokeEnabled) {
|
||||
ctx.shadowColor = "rgba(0,0,0,0)"
|
||||
ctx.shadowBlur = 0
|
||||
ctx.lineWidth = r(settings.strokeWidth ?? 4)
|
||||
ctx.strokeStyle = settings.strokeColor ?? "#000000"
|
||||
ctx.strokeText(line, centerX, y)
|
||||
// 恢复阴影
|
||||
if (shadowEnabled) {
|
||||
ctx.shadowColor = settings.shadowColor ?? "rgba(0,0,0,0.8)"
|
||||
ctx.shadowBlur = r(settings.shadowBlur ?? 4)
|
||||
ctx.shadowOffsetX = r(settings.shadowOffsetX ?? 2)
|
||||
ctx.shadowOffsetY = r(settings.shadowOffsetY ?? 2)
|
||||
}
|
||||
}
|
||||
ctx.fillText(line, centerX, y)
|
||||
})
|
||||
|
||||
// 恢复
|
||||
ctx.shadowColor = prevShadow.c
|
||||
ctx.shadowBlur = prevShadow.b
|
||||
ctx.shadowOffsetX = prevShadow.ox
|
||||
ctx.shadowOffsetY = prevShadow.oy
|
||||
}
|
||||
// 计算当前字号(draw() 内部同样逻辑,抽出来供 fontString 复用)
|
||||
const previewScale = width / 360
|
||||
const previewR = (v: number) => Math.round(v * previewScale)
|
||||
const buildFontString = () => {
|
||||
const size = previewR(settings.size)
|
||||
const ff = getFontFamily(settings.font)
|
||||
const parts: string[] = []
|
||||
if (settings.italic) parts.push("italic")
|
||||
if (settings.bold) parts.push("bold")
|
||||
parts.push(`${size}px`, ff)
|
||||
return parts.join(" ")
|
||||
}
|
||||
|
||||
// Web Font 加载保障:
|
||||
// 1) 等 document.fonts.ready(CSS @font-face 首次可用)
|
||||
// 2) 显式 FontFaceSet.load(fontString, text) 触发浏览器真正下载并加载
|
||||
// 当前字体到 Canvas 可用,避免首次绘制用 fallback 字体画出错字/色块
|
||||
const doDrawWhenReady = async () => {
|
||||
try {
|
||||
if (typeof document !== "undefined" && document.fonts) {
|
||||
await document.fonts.ready
|
||||
try {
|
||||
await document.fonts.load(buildFontString(), text)
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
if (!cancelled) draw()
|
||||
}
|
||||
}
|
||||
doDrawWhenReady()
|
||||
return () => {
|
||||
cancelled = true
|
||||
}
|
||||
}, [settings, width, h, text, transparent, portrait, background])
|
||||
|
||||
return (
|
||||
<canvas
|
||||
ref={canvasRef}
|
||||
style={{
|
||||
borderRadius: 6,
|
||||
display: "block",
|
||||
maxWidth: "100%",
|
||||
background: transparent ? "transparent" : background,
|
||||
}}
|
||||
/>
|
||||
)
|
||||
}
|
||||
|
||||
function roundRect(
|
||||
ctx: CanvasRenderingContext2D,
|
||||
x: number,
|
||||
y: number,
|
||||
w: number,
|
||||
h: number,
|
||||
r: number,
|
||||
) {
|
||||
const rr = Math.min(r, w / 2, h / 2)
|
||||
ctx.beginPath()
|
||||
ctx.moveTo(x + rr, y)
|
||||
ctx.lineTo(x + w - rr, y)
|
||||
ctx.quadraticCurveTo(x + w, y, x + w, y + rr)
|
||||
ctx.lineTo(x + w, y + h - rr)
|
||||
ctx.quadraticCurveTo(x + w, y + h, x + w - rr, y + h)
|
||||
ctx.lineTo(x + rr, y + h)
|
||||
ctx.quadraticCurveTo(x, y + h, x, y + h - rr)
|
||||
ctx.lineTo(x, y + rr)
|
||||
ctx.quadraticCurveTo(x, y, x + rr, y)
|
||||
ctx.closePath()
|
||||
}
|
||||
|
||||
export default TitleMiniPreview
|
||||
export { default } from "@/components/title/TitleMiniPreview"
|
||||
|
||||
@@ -31,7 +31,7 @@ import {
|
||||
} from "@/components/title/constants"
|
||||
import { buildPresetPreviewSettings } from "@/components/title/utils"
|
||||
|
||||
import TitleMiniPreview from "./TitleMiniPreview"
|
||||
import TitleMiniPreview from "@/components/title/TitleMiniPreview"
|
||||
import TitleTemplateEditor from "@/components/title/TitleTemplateEditor"
|
||||
import { useTitleTemplates } from "@/components/title/useTitleTemplates"
|
||||
import {
|
||||
@@ -177,7 +177,7 @@ const ColorPicker: React.FC<{
|
||||
|
||||
/* ── 卡片预览:用 ref 测量容器宽度后再渲染透明 Canvas,保证文字清晰 ── */
|
||||
const FillPreview: React.FC<{
|
||||
settings: TitleSettings
|
||||
settings: import("@/components/title/settings").TitleStyleSettings
|
||||
sampleText: string
|
||||
portrait?: boolean
|
||||
}> = ({ settings, sampleText, portrait }) => {
|
||||
|
||||
@@ -46,13 +46,9 @@ export const CLIP_COUNT_STEP = 1
|
||||
export const MAX_PREVIEW_COUNT = 10
|
||||
export const MIN_PREVIEW_COUNT = 1
|
||||
|
||||
/* ── 标题位置选项 ── */
|
||||
export const POSITION_OPTIONS = [
|
||||
{ value: "top", label: "顶部" },
|
||||
{ value: "center", label: "居中" },
|
||||
{ value: "bottom", label: "底部" },
|
||||
{ value: "custom", label: "自定义" },
|
||||
]
|
||||
/* ── 标题位置选项(统一从公共层重导出) ── */
|
||||
export { POSITION_OPTIONS } from "@/components/title/position-options"
|
||||
export type { PositionOption } from "@/components/title/position-options"
|
||||
|
||||
/* ── 标题字体:统一使用公共层定义(#2001) ── */
|
||||
export { getFontFamily } from "@/components/title/constants"
|
||||
|
||||
@@ -41,8 +41,8 @@ export interface UseGenerateVideoProps {
|
||||
enabled: boolean
|
||||
music_id?: string
|
||||
}
|
||||
/** 生成成功后的回调(用于清除持久化的 previewTaskId 等状态) */
|
||||
onGenerationSuccess?: () => void
|
||||
/** 生成成功后的回调(用于清除持久化的 previewTaskId 等状态);status=awaiting_cover 表示需进封面选择 */
|
||||
onGenerationSuccess?: (status?: "completed" | "awaiting_cover") => void
|
||||
/* ── 批量生成(#1677)── */
|
||||
/** 生成数量(1=单条旧逻辑,>1=批量) */
|
||||
previewCount?: number
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { useRef, useCallback, useState } from "react"
|
||||
import { useRef, useCallback, useState, useEffect } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { getGenerationTask, retryTask as retryGenerationTaskApi } from "@/api/tasks/tasks"
|
||||
@@ -19,7 +19,7 @@ export interface BatchTaskState {
|
||||
|
||||
interface UseGenerationPollingOptions {
|
||||
onProgress: (progress: number) => void
|
||||
onComplete: (videos: unknown[]) => void
|
||||
onComplete: (videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => void
|
||||
onFailed: (errorMsg: string) => void
|
||||
/** 批量:单任务状态变化(第5步逐卡片展示) */
|
||||
onBatchTaskUpdate?: (taskId: string, patch: Partial<BatchTaskState>) => void
|
||||
@@ -31,13 +31,19 @@ const MAX_RETRYABLE_ERRORS = 10
|
||||
const MAX_RESULTS_RETRIES = 3
|
||||
|
||||
/**
|
||||
* 生成状态轮询 Hook(v4 — 批量任务独立状态 + 单任务重试)
|
||||
* 生成状态轮询 Hook(v5 — awaiting_cover 状态识别 + visibilitychange 恢复 + 状态透传)
|
||||
*
|
||||
* startPolling(taskId) 轮询单个任务;
|
||||
* startPollingBatch(tasks) 并行轮询 N 个任务:
|
||||
* - 每个任务独立进度/状态/失败,通过 onBatchTaskUpdate 实时回传
|
||||
* - 全部成功才 onComplete(聚合视频按变体顺序);任一失败不影响其他任务继续
|
||||
* - retryTask(taskId) 单独重试失败任务(重新轮询,后端任务仍在跑则直接接续)
|
||||
*
|
||||
* v5 修复(#2088):
|
||||
* 1. 单任务路径透传 taskStatus(completed / awaiting_cover)到 onComplete,外层据此区分跳转
|
||||
* 2. 监听 visibilitychange,页面从后台切回可见时立即补拉一次,解决切后台 setInterval 被浏览器
|
||||
* 降频/冻结导致进度卡在 56% 的问题
|
||||
* 3. 非 4xx/5xx 网络错误按 3s 退避重试(已有 MAX_RETRYABLE_ERRORS=10 兜底)
|
||||
*/
|
||||
export function useGenerationPolling({
|
||||
onProgress,
|
||||
@@ -49,12 +55,15 @@ export function useGenerationPolling({
|
||||
const cancelledRef = useRef(false)
|
||||
/** 批量任务上下文:taskId → 变体序号 */
|
||||
const batchContextRef = useRef<Map<string, number>>(new Map())
|
||||
/** 当前活跃的「立刻补拉一次」函数(visibilitychange 回调使用) */
|
||||
const immediateTickRef = useRef<(() => void) | null>(null)
|
||||
const [, forceTick] = useState(0)
|
||||
|
||||
const clearTimer = useCallback(() => {
|
||||
cancelledRef.current = true
|
||||
progressTimer.current.forEach((t) => clearTimeout(t))
|
||||
progressTimer.current = []
|
||||
immediateTickRef.current = null
|
||||
}, [])
|
||||
|
||||
/** 任务完成后拉取结果列表,带重试 */
|
||||
@@ -113,6 +122,7 @@ export function useGenerationPolling({
|
||||
|
||||
if (task.status === "completed" || task.status === "awaiting_cover") {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const videos = await fetchResultsWithRetry(taskId)
|
||||
if (cancelledRef.current) return
|
||||
if (videos === null) {
|
||||
@@ -128,6 +138,7 @@ export function useGenerationPolling({
|
||||
|
||||
if (task.status === "failed" || task.status === "cancelled") {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const rawMsg =
|
||||
task.error_info?.error_message ||
|
||||
task.error_message ||
|
||||
@@ -138,6 +149,7 @@ export function useGenerationPolling({
|
||||
return
|
||||
}
|
||||
|
||||
// running / pending / waiting:更新进度并安排下一次轮询
|
||||
const pct = Math.max(0, Math.min(99, Math.round(Number(task.progress) || 0)))
|
||||
callbacks?.onTaskProgress?.(pct)
|
||||
if (!callbacks && runId === 0) {
|
||||
@@ -149,16 +161,20 @@ export function useGenerationPolling({
|
||||
if (cancelledRef.current || done) return
|
||||
console.error("[轮询出错] taskId:", taskId, pollErr)
|
||||
const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined
|
||||
// 4xx 视为不可重试(任务不存在/权限问题等),直接失败
|
||||
if (status && status >= 400 && status < 500) {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const msg = extractErrorMessage(pollErr, status)
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
// 网络错误 / 5xx:3s 退避重试,最多 MAX_RETRYABLE_ERRORS 次
|
||||
consecutiveErrors += 1
|
||||
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
|
||||
done = true
|
||||
immediateTickRef.current = null
|
||||
const msg = "任务状态查询连续失败,请稍后在任务列表查看结果"
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
@@ -169,6 +185,18 @@ export function useGenerationPolling({
|
||||
}
|
||||
}
|
||||
|
||||
// 注册「立刻补拉一次」回调,供 visibilitychange 恢复时调用
|
||||
// 注意:必须在 done 后清理,避免切换页面时误触发已结束任务的补拉
|
||||
immediateTickRef.current = () => {
|
||||
if (!done && !cancelledRef.current) {
|
||||
// 清除未触发的 setTimeout,立即拉一次
|
||||
progressTimer.current.forEach((t) => clearTimeout(t))
|
||||
progressTimer.current = []
|
||||
consecutiveErrors = 0
|
||||
void poll()
|
||||
}
|
||||
}
|
||||
|
||||
const timer = setTimeout(poll, 1500)
|
||||
progressTimer.current.push(timer)
|
||||
})
|
||||
@@ -181,12 +209,22 @@ export function useGenerationPolling({
|
||||
(taskId: string) => {
|
||||
cancelledRef.current = false
|
||||
batchContextRef.current.clear()
|
||||
pollSingleTask(taskId, 0)
|
||||
.then((videos) => {
|
||||
if (cancelledRef.current) return
|
||||
let resolvedStatus: "completed" | "awaiting_cover" = "completed"
|
||||
pollSingleTask(taskId, 0, {
|
||||
onTaskProgress: (pct) => onProgress(pct),
|
||||
onTaskCompleted: (videos, taskStatus) => {
|
||||
resolvedStatus = taskStatus ?? "completed"
|
||||
onProgress(100)
|
||||
onComplete(videos)
|
||||
message.success("视频生成完成!")
|
||||
onComplete(videos, resolvedStatus)
|
||||
},
|
||||
onTaskFailed: (msg) => onFailed(msg),
|
||||
})
|
||||
.then(() => {
|
||||
if (cancelledRef.current) return
|
||||
// awaiting_cover 是中间态(进封面选择页),不弹"完成"toast;completed 才弹
|
||||
if (resolvedStatus === "completed") {
|
||||
message.success("视频生成完成!")
|
||||
}
|
||||
})
|
||||
.catch((err: Error) => {
|
||||
if (cancelledRef.current) return
|
||||
@@ -201,7 +239,7 @@ export function useGenerationPolling({
|
||||
/**
|
||||
* 批量多任务轮询:
|
||||
* - 每个任务独立进度/状态回传 onBatchTaskUpdate
|
||||
* * 全部完成后按变体顺序聚合视频 onComplete
|
||||
* - 全部完成后按变体顺序聚合视频 onComplete
|
||||
* - 部分失败:整体不 onFailed(第5步逐卡片展示失败+重试按钮);全部失败才 onFailed
|
||||
*/
|
||||
const startPollingBatch = useCallback(
|
||||
@@ -211,6 +249,7 @@ export function useGenerationPolling({
|
||||
const progressMap = new Map<string, number>()
|
||||
const resultMap = new Map<string, unknown[]>()
|
||||
const failureMap = new Map<string, string>()
|
||||
const statusMap = new Map<string, "completed" | "awaiting_cover">()
|
||||
batchContextRef.current = new Map(tasks.map((t) => [t.taskId, t.variantIndex]))
|
||||
|
||||
const reportAggregateProgress = () => {
|
||||
@@ -225,7 +264,9 @@ export function useGenerationPolling({
|
||||
if (resultMap.size === tasks.length) {
|
||||
onProgress(100)
|
||||
const ordered = tasks.map((t) => resultMap.get(t.taskId) || []).flat()
|
||||
onComplete(ordered)
|
||||
// 批量:任一任务为 awaiting_cover,则整体透传 awaiting_cover(进封面页)
|
||||
const anyAwaiting = Array.from(statusMap.values()).some((s) => s === "awaiting_cover")
|
||||
onComplete(ordered, anyAwaiting ? "awaiting_cover" : "completed")
|
||||
message.success(`全部 ${tasks.length} 个视频生成完成!`)
|
||||
} else if (resultMap.size > 0) {
|
||||
// 部分失败:成功的视频聚合进成片列表(可进封面),失败卡片带重试按钮
|
||||
@@ -234,7 +275,8 @@ export function useGenerationPolling({
|
||||
.filter((t) => resultMap.has(t.taskId))
|
||||
.map((t) => resultMap.get(t.taskId) || [])
|
||||
.flat()
|
||||
onComplete(ordered)
|
||||
const anyAwaiting = Array.from(statusMap.values()).some((s) => s === "awaiting_cover")
|
||||
onComplete(ordered, anyAwaiting ? "awaiting_cover" : "completed")
|
||||
message.warning(
|
||||
`${failureMap.size} 个视频生成失败,可点击卡片上的「重试此视频」,成功的视频可先进入下一步`,
|
||||
)
|
||||
@@ -263,6 +305,7 @@ export function useGenerationPolling({
|
||||
progressMap.set(taskId, 100)
|
||||
resultMap.set(taskId, videos)
|
||||
const _finalStatus: "completed" | "awaiting_cover" = taskStatus ?? "completed"
|
||||
statusMap.set(taskId, _finalStatus)
|
||||
onBatchTaskUpdate?.(taskId, { status: _finalStatus, progress: 100, videos })
|
||||
reportAggregateProgress()
|
||||
checkAllSettled()
|
||||
@@ -309,5 +352,27 @@ export function useGenerationPolling({
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
/**
|
||||
* visibilitychange 恢复:页面从后台切回前台时,立刻触发一次补拉。
|
||||
* 解决浏览器后台标签页对 setTimeout 的 1Hz 节流/冻结导致的"进度卡 56%"问题。
|
||||
*/
|
||||
useEffect(() => {
|
||||
const handleVisibilityChange = () => {
|
||||
if (document.visibilityState === "visible" && immediateTickRef.current) {
|
||||
immediateTickRef.current()
|
||||
}
|
||||
}
|
||||
document.addEventListener("visibilitychange", handleVisibilityChange)
|
||||
// 页面聚焦也兜底一次(部分浏览器 visibilitychange 触发时机不一致)
|
||||
const handleFocus = () => {
|
||||
if (immediateTickRef.current) immediateTickRef.current()
|
||||
}
|
||||
window.addEventListener("focus", handleFocus)
|
||||
return () => {
|
||||
document.removeEventListener("visibilitychange", handleVisibilityChange)
|
||||
window.removeEventListener("focus", handleFocus)
|
||||
}
|
||||
}, [])
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -13,6 +13,8 @@ import { validateGenerateInputs } from "./generate-video/buildPayload"
|
||||
import { calculateResolution } from "../utils/calculateResolution"
|
||||
import { extractBackendError, translateError } from "./generate-video/errorUtils"
|
||||
|
||||
export type GenerationCompleteStatus = "completed" | "awaiting_cover" | null
|
||||
|
||||
export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const { selectedTemplate, onGenerationSuccess } = props
|
||||
|
||||
@@ -22,6 +24,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const [generated, setGenerated] = useState(false)
|
||||
const [generateError, setGenerateError] = useState<string | null>(null)
|
||||
const [generatedVideos, setGeneratedVideos] = useState<GeneratedVideo[]>([])
|
||||
/** #2088:任务最终状态,区分 awaiting_cover(选封面)/ completed(已完成) */
|
||||
const [completionStatus, setCompletionStatus] = useState<GenerationCompleteStatus>(null)
|
||||
/** 单视频模式:当前任务 ID(封面 finalize 需要) */
|
||||
const [currentTaskId, setCurrentTaskId] = useState<string>("")
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
@@ -53,9 +57,11 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
(videos: unknown[]) => {
|
||||
(videos: unknown[], taskStatus?: "completed" | "awaiting_cover") => {
|
||||
setGenerating(false)
|
||||
setGenerated(true)
|
||||
const finalStatus: GenerationCompleteStatus = taskStatus ?? "completed"
|
||||
setCompletionStatus(finalStatus)
|
||||
setGeneratedVideos(videos as GeneratedVideo[])
|
||||
// 批量:成功任务的 videos 已通过 onBatchTaskUpdate 写入,这里同步兜底
|
||||
setBatchTasks((prev) =>
|
||||
@@ -70,7 +76,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
: t,
|
||||
),
|
||||
)
|
||||
onGenerationSuccess?.()
|
||||
onGenerationSuccess?.(finalStatus)
|
||||
},
|
||||
[onGenerationSuccess],
|
||||
)
|
||||
@@ -121,6 +127,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
setGenerateError(null)
|
||||
setCompletionStatus(null)
|
||||
setBatchTasks([])
|
||||
setGeneratedVideos([])
|
||||
setCurrentTaskId("")
|
||||
@@ -423,6 +430,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
generated,
|
||||
generateError,
|
||||
generatedVideos,
|
||||
completionStatus,
|
||||
currentTaskId,
|
||||
generate,
|
||||
retry,
|
||||
|
||||
@@ -45,6 +45,11 @@ export const STATUS_CONFIG: Record<
|
||||
color: "processing",
|
||||
icon: <SyncOutlined spin />,
|
||||
},
|
||||
awaiting_cover: {
|
||||
label: "待选封面",
|
||||
color: "warning",
|
||||
icon: <ClockCircleOutlined />,
|
||||
},
|
||||
completed: {
|
||||
label: "已完成",
|
||||
color: "success",
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
"""一次性脚本:对历史 quality_score 缺失的视频素材重新打分。
|
||||
|
||||
背景(#2073):镜像 97ad0ae2 时期 calculate_quality_score / classify_from_analysis
|
||||
返回 str 而非 AssetClassification 枚举,导致 calculate_asset_quality 连续报
|
||||
"'str' object has no attribute 'value'",大量视频素材的 quality_score 卡在 NULL。
|
||||
镜像 8abdeb95 已修复枚举 bug,但历史失败记录不会自动重跑。本脚本扫描全表,
|
||||
把 quality_score IS NULL 的视频素材重新投递到 worker.calculate_asset_quality 任务。
|
||||
|
||||
使用方式(在 worker 容器内执行):
|
||||
cd /app/apps/worker
|
||||
# 干跑,只打印会重跑多少条,不发任务
|
||||
python -m scripts.backfill_asset_quality --dry-run
|
||||
# 正式执行
|
||||
python -m scripts.backfill_asset_quality
|
||||
# 只重跑最近 N 天的
|
||||
python -m scripts.backfill_asset_quality --since-days 30
|
||||
# 限流:每投递一批 sleep 几秒,避免瞬间打爆 transcode 队列
|
||||
python -m scripts.backfill_asset_quality --batch-size 50 --sleep 2
|
||||
|
||||
也可以直接在 staging 机器上 exec 进容器:
|
||||
docker exec -e PYTHONPATH=/app:/app/apps/api:/app/packages xiaoxia-worker-staging \
|
||||
python -m scripts.backfill_asset_quality --dry-run
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
# 保证可以以 python -m scripts.xxx 在容器 /app/apps/worker 下执行
|
||||
# 也兼容在 repo 根目录下执行(注入路径)
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
_WORKER_DIR = os.path.dirname(_SCRIPT_DIR) # apps/worker
|
||||
_APPS_DIR = os.path.dirname(_WORKER_DIR) # apps
|
||||
_REPO_ROOT = os.path.dirname(_APPS_DIR) # repo root
|
||||
for p in (_REPO_ROOT, os.path.join(_REPO_ROOT, "apps", "api"), _REPO_ROOT):
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description="补打历史视频素材 quality_score")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只统计数量,不投递任务")
|
||||
parser.add_argument("--since-days", type=int, default=0, help="只处理最近 N 天上传的素材(0=全部)")
|
||||
parser.add_argument("--batch-size", type=int, default=50, help="每批投递数量,默认 50")
|
||||
parser.add_argument("--sleep", type=float, default=1.0, help="批次之间 sleep 秒数,默认 1s")
|
||||
parser.add_argument("--queue", type=str, default="transcode", help="投递队列(默认 transcode)")
|
||||
args = parser.parse_args()
|
||||
|
||||
# 延迟 import,避免在 dry-run 时依赖完整 DB 环境
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(AssetModel).filter(
|
||||
AssetModel.file_type == "video",
|
||||
AssetModel.quality_score.is_(None),
|
||||
)
|
||||
if args.since_days > 0:
|
||||
cutoff = datetime.now(UTC) - timedelta(days=args.since_days)
|
||||
q = q.filter(AssetModel.created_at >= cutoff)
|
||||
|
||||
# 先 count 打印
|
||||
total = q.count()
|
||||
print(
|
||||
f"[backfill] 待重跑 quality_score 的视频素材: {total} 条"
|
||||
f"{' (dry-run,不投递)' if args.dry_run else ''}"
|
||||
f"{' (最近 ' + str(args.since_days) + ' 天)' if args.since_days > 0 else ''}",
|
||||
flush=True,
|
||||
)
|
||||
if total == 0 or args.dry_run:
|
||||
return 0
|
||||
|
||||
# 分批投递
|
||||
submitted = 0
|
||||
batch = 0
|
||||
offset = 0
|
||||
while True:
|
||||
assets = q.order_by(AssetModel.created_at.desc()).offset(offset).limit(args.batch_size).all()
|
||||
if not assets:
|
||||
break
|
||||
batch += 1
|
||||
for a in assets:
|
||||
try:
|
||||
celery_app.send_task(
|
||||
"worker.calculate_asset_quality",
|
||||
args=[a.id],
|
||||
queue=args.queue,
|
||||
)
|
||||
submitted += 1
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[backfill] 投递失败 asset_id={a.id}: {e}", flush=True)
|
||||
print(f"[backfill] batch {batch}: 已累计投递 {submitted}/{total}", flush=True)
|
||||
offset += len(assets)
|
||||
if args.sleep > 0 and offset < total:
|
||||
time.sleep(args.sleep)
|
||||
|
||||
print(f"[backfill] 完成,共投递 {submitted} 条任务到 {args.queue} 队列", flush=True)
|
||||
return 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,831 @@
|
||||
"""全 GPU 直连渲染管线(P1)。
|
||||
|
||||
背景:旧链路 worker 先用 CPU libx264 把 filter_complex 输出成 mezzanine(1080p 约 85s),
|
||||
上传后再由 P4000 NVENC 编码,渲染后还要单独跑一次随机边缘裁剪重编码(约 26s)。
|
||||
本管线取消 mezzanine:把原始素材签名 URL 作为多输入直接交给 P4000,filter_complex 内
|
||||
一步完成 trim/scale/pad/concat/边缘随机裁剪/drawtext 字幕,末端 h264_nvenc 只编码一次;
|
||||
原素材音轨 concat + TTS/配音/BGM 混音也在同一命令里完成。
|
||||
|
||||
约束(P1):
|
||||
- 仅覆盖智能剪辑主流场景:单一主视频轨、全硬切、无 PiP/overlay/水印/贴纸/片头片尾/绿幕。
|
||||
不满足条件时调用方回退到现有 mezzanine/CPU 链路(功能不回归)。
|
||||
- 字幕先用 drawtext(P4000 装好中文字体后可再切 subtitles 滤镜烧 ASS)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_DRAWTEXT_FONT = "Noto Sans CJK SC"
|
||||
EDGE_CROP_MIN_PCT = 0.02
|
||||
EDGE_CROP_MAX_PCT = 0.05
|
||||
|
||||
# 标题/字幕样式基准宽度(px)。前端 TitleSettings 所有长度字段(size/描边/阴影/margin/pos)
|
||||
# 均以 720p 为基准(见前端 titleCanvas.ts 注释 scale=videoWidth/720,types.ts "px @720p"),
|
||||
# 非 720p 输出时按 video_width / TITLE_SIZE_REF_WIDTH 等比缩放,保证成片位置与前端预览一致。
|
||||
TITLE_SIZE_REF_WIDTH = 720
|
||||
# 与 video_filter_builder.build_title_drawtext_filter(CPU 路径)和 ass_subtitle_builder 对齐:
|
||||
# - top/bottom 默认 margin 50@720p(vfb 用 _scale_title_len(50, w),即 y=50 / y=h-th-50)
|
||||
# - margin_top 字段:前端编辑器 marginTop 滑块,叠加在默认 margin 之上(#2095 支持)
|
||||
# - PAD 概念仅用于前端 Canvas 预览;ffmpeg drawtext y 是 baseline,无 font metrics 可用,
|
||||
# 直接用统一 50@720p baseline 位置即可保持三端(GPU/CPU/前端视觉)一致。
|
||||
TITLE_DEFAULT_MARGIN_TOP = 50 # top 位置 baseline 默认距顶 50@720p(与 vfb/CPU 路径一致)
|
||||
TITLE_DEFAULT_MARGIN_BOTTOM = 50 # bottom 位置 baseline 默认距底 50@720p
|
||||
SUBTITLE_DEFAULT_MARGIN_BOTTOM = 50 # 字幕距底边距 50@720p(与 vfb 一致)
|
||||
TITLE_MARGIN_TOP_FROM_CFG_DEFAULT = 24 # 前端 marginTop 滑块默认值(用户未传时叠加 0)
|
||||
TITLE_FAUX_BOLD_WIDTH = 2 # 仿粗黑色描边宽度(与 vfb 一致,2@720p 黑色细描边)
|
||||
|
||||
|
||||
def _scale_title_len(value, video_width: int):
|
||||
"""将 720p 基准长度按 video_width 等比缩放(与 packages/domain/ass_subtitle_builder._scale_len 一致)。
|
||||
|
||||
int 输入 → 返回 int;float 输入 → 返回 float;非法值原样返回。
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
v = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return value
|
||||
if not video_width or video_width <= 0:
|
||||
return int(round(v)) if isinstance(value, int) else v
|
||||
scaled = v * (video_width / TITLE_SIZE_REF_WIDTH)
|
||||
return int(round(scaled)) if isinstance(value, int) else scaled
|
||||
|
||||
|
||||
def escape_drawtext_text(text: str) -> str:
|
||||
if not text:
|
||||
return ""
|
||||
s = text.replace("\\", "\\\\")
|
||||
s = s.replace(":", "\\:")
|
||||
s = s.replace("'", "\\'")
|
||||
s = s.replace("%", "\\%")
|
||||
s = s.replace(",", "\\,")
|
||||
s = s.replace("[", "\\[").replace("]", "\\]")
|
||||
s = s.replace(";", "\\;")
|
||||
s = s.replace("\n", " ")
|
||||
return s
|
||||
|
||||
|
||||
def _hex_to_drawtext_color(hex_color: str, default: str = "white") -> str:
|
||||
"""把 #RRGGBB / #RGB / 命名颜色转换为 ffmpeg drawtext 接受的颜色格式。
|
||||
|
||||
drawtext 的 fontcolor 接受 0xRRGGBB 形式(或命名颜色如 white/black/yellow)。
|
||||
描边/阴影颜色同样适用。alpha 后缀支持(#RRGGBB@0.5 或 &HBBGGRRAA)。
|
||||
"""
|
||||
if not hex_color:
|
||||
return default
|
||||
s = hex_color.strip()
|
||||
if not s:
|
||||
return default
|
||||
# 命名颜色直接返回(白名单常见值,避免把 #xxx 当成命名)
|
||||
if not s.startswith("#") and not s.startswith("0x") and "@" not in s:
|
||||
return s
|
||||
if s.startswith("0x"):
|
||||
return s # 已是 drawtext 原生格式
|
||||
if s.startswith("#"):
|
||||
h = s[1:]
|
||||
# 处理 alpha:#RRGGBB@AA 或 #RRGGBB&AA
|
||||
alpha = ""
|
||||
if "@" in h:
|
||||
h, alpha_part = h.split("@", 1)
|
||||
try:
|
||||
a = float(alpha_part)
|
||||
alpha = f"@{a:.2f}"
|
||||
except ValueError:
|
||||
alpha = ""
|
||||
if len(h) == 3:
|
||||
h = "".join(ch * 2 for ch in h)
|
||||
if len(h) == 6:
|
||||
try:
|
||||
int(h, 16)
|
||||
except ValueError:
|
||||
return default
|
||||
return f"0x{h}{alpha}"
|
||||
if len(h) == 8:
|
||||
# RRGGBBAA → drawtext 的 0xRRGGBB@AA 形式
|
||||
try:
|
||||
int(h, 16)
|
||||
except ValueError:
|
||||
return default
|
||||
rr, gg, bb, aa = h[0:2], h[2:4], h[4:6], h[6:8]
|
||||
try:
|
||||
a = int(aa, 16) / 255.0
|
||||
return f"0x{rr}{gg}{bb}@{a:.2f}"
|
||||
except ValueError:
|
||||
return f"0x{rr}{gg}{bb}"
|
||||
return default
|
||||
|
||||
|
||||
def _position_to_drawtext_xy(
|
||||
position: str,
|
||||
*,
|
||||
margin_top: int = 0,
|
||||
margin_bottom: int = 0,
|
||||
pos_x: Optional[float] = None,
|
||||
pos_y: Optional[float] = None,
|
||||
) -> tuple[str, str]:
|
||||
"""把位置映射到 drawtext x/y 表达式,对齐前端 titleCanvas.ts 预览坐标。
|
||||
|
||||
position 支持: top / center(middle) / bottom / custom。
|
||||
- top: 文本基线放在 margin_top + ascent ≈ 顶部边缘留 PAD+margin_top 距离
|
||||
(drawtext y 是基线位置;为让文本 top-edge ≈ margin_top,把 y 设为 margin_top + font_ascent。
|
||||
但 drawtext 运行时不知道 ascent,用经验系数 0.8*fontsize 近似,和前端 PAD+margin_top 对齐)。
|
||||
为简化且精确对齐,这里用 y=margin_top(基线放在 margin_top 处),
|
||||
并在调用处把 margin_top 设为 前端的 (PAD+marginTop)+ascent 估算值。
|
||||
- center: (h-text_h)/2 垂直居中。
|
||||
- bottom: 文本底线距离底边 margin_bottom。
|
||||
- custom: pos_x/pos_y 为百分比 0-100(前端拖拽坐标系),文本中心落在 (pct_x*w, pct_y*h)。
|
||||
margin_top/margin_bottom 为已按 video_width 缩放过的像素值。
|
||||
"""
|
||||
p = (position or "top").lower().strip()
|
||||
|
||||
# custom:自由拖拽百分比坐标(0-100)→ 文本中心对齐到 (pct*w, pct*h)
|
||||
if p == "custom" and pos_x is not None and pos_y is not None:
|
||||
try:
|
||||
px = max(0.0, min(100.0, float(pos_x))) / 100.0
|
||||
py = max(0.0, min(100.0, float(pos_y))) / 100.0
|
||||
return f"(w-text_w)*{px:.4f}", f"(h-text_h)*{py:.4f}"
|
||||
except (TypeError, ValueError):
|
||||
pass # fall through to default
|
||||
|
||||
x = "(w-text_w)/2"
|
||||
if p in ("top",):
|
||||
# drawtext y 是 baseline 位置。中文字符顶边距基线约 0.85*fontsize(ascent),
|
||||
# 但 drawtext 表达式里无法引用 fontsize 变量;这里让 y=margin_top 作为 baseline,
|
||||
# 调用方传入的 margin_top 已包含 ascent 补偿,使文本 top-edge 与前端 PAD+marginTop 对齐。
|
||||
y = f"{int(margin_top)}"
|
||||
elif p in ("center", "middle"):
|
||||
y = "(h-text_h)/2"
|
||||
elif p in ("bottom",):
|
||||
# h-th-margin_bottom:th ≈ text_h,文本底边距底边 margin_bottom
|
||||
y = f"h-th-{int(margin_bottom)}"
|
||||
else:
|
||||
# 未知值回退到顶部(与前端默认 position=top 对齐)
|
||||
y = f"{int(margin_top)}"
|
||||
return x, y
|
||||
|
||||
|
||||
def _build_drawtext_filters(
|
||||
*,
|
||||
text: str,
|
||||
start: float,
|
||||
end: float,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
font_size: int = 0,
|
||||
font_color: str = "white",
|
||||
position: str = "top",
|
||||
margin_top: int = 0,
|
||||
margin_bottom: int = 0,
|
||||
pos_x: Optional[float] = None,
|
||||
pos_y: Optional[float] = None,
|
||||
box_enabled: bool = False,
|
||||
box_color: str = "black@0.5",
|
||||
borderw: int = 0,
|
||||
border_color: str = "black",
|
||||
shadow_enabled: bool = False,
|
||||
shadow_color: str = "black@0.6",
|
||||
shadow_x: int = 2,
|
||||
shadow_y: int = 2,
|
||||
) -> list[str]:
|
||||
"""构造一组 drawtext 滤镜:可选阴影层(同字偏移)+ 主字层。
|
||||
|
||||
ffmpeg drawtext 没有直接的 shadow 选项,用两次 drawtext 模拟:
|
||||
先画一个描边/阴影色层偏移 shadow_x/shadow_y,再画主字层。
|
||||
返回列表是为了让调用方顺序插入 fc(前一个输出作为后一个输入)。
|
||||
"""
|
||||
txt = escape_drawtext_text(text)
|
||||
if not txt:
|
||||
return []
|
||||
|
||||
x_expr, y_expr = _position_to_drawtext_xy(
|
||||
position,
|
||||
margin_top=margin_top,
|
||||
margin_bottom=margin_bottom,
|
||||
pos_x=pos_x,
|
||||
pos_y=pos_y,
|
||||
)
|
||||
fc_color = _hex_to_drawtext_color(font_color, default="white")
|
||||
bd_color = _hex_to_drawtext_color(border_color, default="black")
|
||||
sh_color = _hex_to_drawtext_color(shadow_color, default="black@0.6")
|
||||
|
||||
filters: list[str] = []
|
||||
|
||||
# 阴影层:shadow_enabled 时先画一层深色偏移字(无描边)
|
||||
if shadow_enabled and (shadow_x != 0 or shadow_y != 0):
|
||||
sh_parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
sh_parts.append(f"fontsize={int(font_size)}")
|
||||
sh_parts.append(f"fontcolor={sh_color}")
|
||||
sh_parts.append(f"x={x_expr}+{int(shadow_x)}")
|
||||
sh_parts.append(f"y={y_expr}+{int(shadow_y)}")
|
||||
if start > 0 or end > 0:
|
||||
sh_parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
filters.append("drawtext=" + ":".join(sh_parts))
|
||||
|
||||
# 主字层
|
||||
parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
parts.append(f"fontsize={int(font_size)}")
|
||||
parts.append(f"fontcolor={fc_color}")
|
||||
if box_enabled:
|
||||
parts.append("box=1")
|
||||
parts.append(f"boxcolor={box_color}")
|
||||
if borderw and borderw > 0:
|
||||
parts.append(f"borderw={int(borderw)}")
|
||||
parts.append(f"bordercolor={bd_color}")
|
||||
parts.append(f"x={x_expr}")
|
||||
parts.append(f"y={y_expr}")
|
||||
if start > 0 or end > 0:
|
||||
parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
filters.append("drawtext=" + ":".join(parts))
|
||||
return filters
|
||||
|
||||
|
||||
def build_drawtext_filter(
|
||||
*,
|
||||
text: str,
|
||||
start: float,
|
||||
end: float,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
font_size: int = 0,
|
||||
font_color: str = "white",
|
||||
x_expr: str = "(w-text_w)/2",
|
||||
y_expr: str = "h-th-60",
|
||||
box: bool = False,
|
||||
box_color: str = "black@0.5",
|
||||
borderw: int = 0,
|
||||
border_color: str = "black",
|
||||
enable: bool = True,
|
||||
) -> str:
|
||||
"""[已废弃] 保留单条 drawtext 的便捷构造;新代码请用 _build_drawtext_filters。"""
|
||||
txt = escape_drawtext_text(text)
|
||||
parts = [f"font={font}", f"text='{txt}'"]
|
||||
if font_size and font_size > 0:
|
||||
parts.append(f"fontsize={int(font_size)}")
|
||||
parts.append(f"fontcolor={_hex_to_drawtext_color(font_color)}")
|
||||
if box:
|
||||
parts.append("box=1")
|
||||
parts.append(f"boxcolor={box_color}")
|
||||
if borderw and borderw > 0:
|
||||
parts.append(f"borderw={int(borderw)}")
|
||||
parts.append(f"bordercolor={_hex_to_drawtext_color(border_color)}")
|
||||
parts.append(f"x={x_expr}")
|
||||
parts.append(f"y={y_expr}")
|
||||
if enable:
|
||||
parts.append(f"enable='between(t,{start:.3f},{end:.3f})'")
|
||||
return "drawtext=" + ":".join(parts)
|
||||
|
||||
|
||||
def _build_atempo_chain(speed: float) -> str:
|
||||
if abs(speed - 1.0) < 1e-6:
|
||||
return ""
|
||||
stages: list[float] = []
|
||||
remaining = speed
|
||||
while remaining > 2.0:
|
||||
stages.append(2.0)
|
||||
remaining /= 2.0
|
||||
while remaining < 0.5:
|
||||
stages.append(0.5)
|
||||
remaining /= 0.5
|
||||
if abs(remaining - 1.0) >= 1e-6:
|
||||
stages.append(remaining)
|
||||
return ",".join(f"atempo={s:.5f}" for s in stages)
|
||||
|
||||
|
||||
def upload_local_audio_and_sign(
|
||||
local_audio: Path,
|
||||
*,
|
||||
tmp_prefix: str = "tmp/gpu-direct-audio/",
|
||||
expires: int = 3600,
|
||||
) -> tuple[str, str]:
|
||||
from video_processing.oss_helpers import _storage # type: ignore
|
||||
|
||||
storage = _storage()
|
||||
key = f"{tmp_prefix.rstrip('/')}/{uuid.uuid4().hex}{local_audio.suffix or '.mp3'}"
|
||||
content_type = "audio/mpeg" if local_audio.suffix.lower() in (".mp3", ".mpeg") else "audio/mp4"
|
||||
storage.upload_file(local_audio, key, content_type=content_type)
|
||||
url = storage.get_download_url(key, expires)
|
||||
return url, key
|
||||
|
||||
|
||||
def sign_asset_url(storage_key: str, *, expires: int = 3600) -> str:
|
||||
from video_processing.oss_helpers import _storage # type: ignore
|
||||
|
||||
storage = _storage()
|
||||
return storage.get_download_url(storage_key, expires)
|
||||
|
||||
|
||||
class DirectRenderPlan:
|
||||
def __init__(
|
||||
self,
|
||||
inputs: dict[str, str],
|
||||
ffmpeg_args: list[str],
|
||||
oss_keys: list[str],
|
||||
filter_complex: list[str] | None = None,
|
||||
):
|
||||
self.inputs = inputs
|
||||
self.ffmpeg_args = ffmpeg_args
|
||||
self.oss_keys = oss_keys
|
||||
self.filter_complex: list[str] = filter_complex or []
|
||||
|
||||
|
||||
def build_direct_render(
|
||||
*,
|
||||
resolved_clips: list[Any],
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
output_fps: int,
|
||||
tts_audio: Optional[Path] = None,
|
||||
bgm_audio: Optional[Path] = None,
|
||||
title_text: str = "",
|
||||
subtitle_segments: Optional[list[Any]] = None,
|
||||
font: str = DEFAULT_DRAWTEXT_FONT,
|
||||
vcodec: str = "h264_nvenc",
|
||||
preset: str = "p4",
|
||||
video_bitrate: str = "",
|
||||
cq: int = 23,
|
||||
edge_crop_pct: float = 0.0,
|
||||
total_duration: float = 0.0,
|
||||
clip_has_audio: Optional[list[bool]] = None,
|
||||
clip_volumes: Optional[list[float]] = None,
|
||||
extra_audio_tracks: Optional[list[tuple[Any, float]]] = None,
|
||||
title_config: Optional[dict] = None,
|
||||
subtitle_config: Optional[dict] = None,
|
||||
bgm_config: Optional[dict] = None,
|
||||
static_subtitle_text: str = "",
|
||||
) -> DirectRenderPlan:
|
||||
"""构造 P4000 直连渲染所需的 inputs 与 ffmpeg_args。
|
||||
|
||||
视频:每段 trim/setpts/scale/pad/fps → concat(全硬切,带音频)→ 随机边缘 crop+scale → drawtext。
|
||||
音频:每段 [i:a](或 anullsrc 静音占位)按 clip 配置 atrim/asetpts/atempo/volume/aresample
|
||||
→ concat=n:N:v=1:a=1 → 与 extra_audio(TTS/配音素材库)、BGM 一起 amix → atrim 精确截断。
|
||||
"""
|
||||
if not resolved_clips:
|
||||
raise ValueError("build_direct_render: no resolved clips")
|
||||
|
||||
inputs: dict[str, str] = {}
|
||||
oss_keys: list[str] = []
|
||||
input_args: list[str] = []
|
||||
fc: list[str] = []
|
||||
n = len(resolved_clips)
|
||||
|
||||
# 规范化每段参数
|
||||
if clip_has_audio is None:
|
||||
clip_has_audio = [True] * n
|
||||
else:
|
||||
clip_has_audio = list(clip_has_audio) + [True] * max(0, n - len(clip_has_audio))
|
||||
clip_has_audio = clip_has_audio[:n]
|
||||
if clip_volumes is None:
|
||||
clip_volumes = [1.0] * n
|
||||
else:
|
||||
clip_volumes = list(clip_volumes) + [1.0] * max(0, n - len(clip_volumes))
|
||||
clip_volumes = clip_volumes[:n]
|
||||
|
||||
clip_starts: list[float] = []
|
||||
clip_effs: list[float] = []
|
||||
clip_speeds: list[float] = []
|
||||
for clip in resolved_clips:
|
||||
start = float(getattr(clip, "start_time", 0) or 0)
|
||||
eff = float(getattr(clip, "duration", 0) or 0)
|
||||
if eff <= 0:
|
||||
eff = float(getattr(clip, "actual_duration", 0) or 0)
|
||||
speed = float(getattr(clip, "playback_speed", 1.0) or 1.0)
|
||||
clip_starts.append(start)
|
||||
clip_effs.append(eff)
|
||||
clip_speeds.append(speed)
|
||||
|
||||
# 1. 视频输入(原始素材签名 URL)
|
||||
for i, clip in enumerate(resolved_clips):
|
||||
sk = (getattr(clip, "config", None) or {}).get("_storage_key")
|
||||
if not sk:
|
||||
raise ValueError(f"clip {getattr(clip, 'clip_id', i)} missing _storage_key")
|
||||
fname = f"v{i}.mp4"
|
||||
inputs[fname] = sign_asset_url(sk)
|
||||
input_args.extend(["-i", fname])
|
||||
|
||||
# 2. 视频段预处理
|
||||
pre_labels: list[str] = []
|
||||
for i in range(n):
|
||||
vf: list[str] = []
|
||||
start, eff, speed = clip_starts[i], clip_effs[i], clip_speeds[i]
|
||||
if eff > 0:
|
||||
if start > 0:
|
||||
vf.append(f"trim=start={start:.3f}:duration={eff:.3f}")
|
||||
else:
|
||||
vf.append(f"trim=duration={eff:.3f}")
|
||||
vf.append("setpts=PTS-STARTPTS")
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
vf.append(f"setpts=PTS/{speed:.4f}")
|
||||
vf.append(f"scale={output_width}:{output_height}:force_original_aspect_ratio=decrease")
|
||||
vf.append(f"pad={output_width}:{output_height}:trunc((ow-iw)/2):trunc((oh-ih)/2):black")
|
||||
vf.append("setpts=PTS-STARTPTS")
|
||||
vf.append(f"fps={output_fps}")
|
||||
label = f"vc{i}"
|
||||
fc.append(f"[{i}:v]{','.join(vf)}[{label}]")
|
||||
pre_labels.append(label)
|
||||
|
||||
# 2b. 音频段预处理(无声源用 anullsrc 占位;volume=0 的段也用 anullsrc 静音占位保持时间轴)
|
||||
anullsrc_counter = 0
|
||||
audio_pre_labels: list[str] = []
|
||||
for i in range(n):
|
||||
start, eff, speed = clip_starts[i], clip_effs[i], clip_speeds[i]
|
||||
vol = float(clip_volumes[i] if i < len(clip_volumes) else 1.0)
|
||||
has_a = bool(clip_has_audio[i] if i < len(clip_has_audio) else True)
|
||||
if not has_a or vol <= 0.001:
|
||||
# 静音占位:用 anullsrc 生成静音,atrim 到段时长
|
||||
sl = f"sil{anullsrc_counter}"
|
||||
anullsrc_counter += 1
|
||||
af: list[str] = ["anullsrc=channel_layout=stereo:sample_rate=44100"]
|
||||
if eff > 0:
|
||||
af.append(f"atrim=duration={eff:.3f}")
|
||||
af.append("asetpts=PTS-STARTPTS")
|
||||
af.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
fc.append(f"{','.join(af)}[{sl}]")
|
||||
# anullsrc 作为 filter 源不需要 -i 输入,直接给 label
|
||||
audio_pre_labels.append(sl)
|
||||
continue
|
||||
|
||||
af = []
|
||||
if eff > 0:
|
||||
if start > 0:
|
||||
af.append(f"atrim=start={start:.3f}:duration={eff:.3f}")
|
||||
else:
|
||||
af.append(f"atrim=duration={eff:.3f}")
|
||||
af.append("asetpts=PTS-STARTPTS")
|
||||
if abs(speed - 1.0) >= 1e-6:
|
||||
atempo = _build_atempo_chain(speed)
|
||||
if atempo:
|
||||
af.append(atempo)
|
||||
if abs(vol - 1.0) >= 1e-3:
|
||||
af.append(f"volume={vol:.3f}")
|
||||
af.append("aresample=44100")
|
||||
af.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
alabel = f"ac{i}"
|
||||
fc.append(f"[{i}:a]{','.join(af)}[{alabel}]")
|
||||
audio_pre_labels.append(alabel)
|
||||
|
||||
# 3. concat(全硬切;v=1:a=1,视频音频一起拼接)
|
||||
concat_in = "".join(f"[{v}][{a}]" for v, a in zip(pre_labels, audio_pre_labels, strict=True))
|
||||
fc.append(f"{concat_in}concat=n={n}:v=1:a=1[vcat][acat]")
|
||||
cur_v = "vcat"
|
||||
cur_a = "acat"
|
||||
|
||||
# 4. 随机边缘裁剪降重(四边独立随机 2%~5%,与 ffmpeg_utils.random_edge_crop 一致)
|
||||
if edge_crop_pct and edge_crop_pct > 0:
|
||||
_r = random.Random()
|
||||
p_min = EDGE_CROP_MIN_PCT
|
||||
p_max = EDGE_CROP_MAX_PCT
|
||||
crop_top = p_min + _r.random() * (p_max - p_min)
|
||||
crop_bottom = p_min + _r.random() * (p_max - p_min)
|
||||
crop_left = p_min + _r.random() * (p_max - p_min)
|
||||
crop_right = p_min + _r.random() * (p_max - p_min)
|
||||
w_expr = f"trunc(iw*(1-{crop_left:.4f}-{crop_right:.4f})/2)*2"
|
||||
h_expr = f"trunc(ih*(1-{crop_top:.4f}-{crop_bottom:.4f})/2)*2"
|
||||
x_expr = f"trunc(iw*{crop_left:.4f}/2)*2"
|
||||
y_expr = f"trunc(ih*{crop_top:.4f}/2)*2"
|
||||
fc.append(
|
||||
f"[{cur_v}]crop=w='{w_expr}':h='{h_expr}':x='{x_expr}':y='{y_expr}',"
|
||||
f"scale={output_width}:{output_height}[vcrop]"
|
||||
)
|
||||
cur_v = "vcrop"
|
||||
|
||||
# 5. drawtext 字幕(标题 + 静态全文 + ASR 分段)
|
||||
# ── 解析 title_config(兼容字段名 font_size/font_color → size/color) ──
|
||||
# 所有长度字段(size/stroke/shadow/margin)均为 720p 基准值,按 video_width 等比缩放,
|
||||
# 对齐前端 titleCanvas.ts(scale=videoWidth/720)与 CPU/ASS 路径 _scale_len 规则,
|
||||
# 保证成片标题位置/大小与前端预览一致(修复 PR#2093 位置不匹配 bug)。
|
||||
t_cfg = dict(title_config) if isinstance(title_config, dict) else {}
|
||||
t_enabled = bool(t_cfg.get("enabled", True))
|
||||
t_text = (t_cfg.get("text", "") or title_text or "").strip()
|
||||
t_font = str(t_cfg.get("font", font) or font)
|
||||
# size:前端传 px@720p,未配置默认 28(前端 DEFAULT_TITLE_SETTINGS.size=28,对齐 AI Avatar 默认48)
|
||||
t_size_raw = t_cfg.get("size", t_cfg.get("font_size", 0))
|
||||
try:
|
||||
t_size_720 = int(t_size_raw) if t_size_raw else 0
|
||||
except (TypeError, ValueError):
|
||||
t_size_720 = 0
|
||||
if t_size_720 <= 0:
|
||||
t_size_720 = 48 # 与 config_schemas.DEFAULT_EDIT_PLAN_CONFIG.title.size=48 及 vfb 默认一致
|
||||
t_size = _scale_title_len(t_size_720, output_width)
|
||||
# stroke/shadow 长度字段也需 720p→输出分辨率缩放
|
||||
t_color = str(t_cfg.get("color", t_cfg.get("font_color", "#ffffff")))
|
||||
t_position = str(t_cfg.get("position", "top")).lower().strip()
|
||||
# 自由拖拽坐标(百分比 0-100),与 video_filter_builder.build_title_drawtext_filter 一致
|
||||
t_pos_x = t_cfg.get("pos_x")
|
||||
t_pos_y = t_cfg.get("pos_y")
|
||||
try:
|
||||
t_pos_x = float(t_pos_x) if t_pos_x is not None else None
|
||||
t_pos_y = float(t_pos_y) if t_pos_y is not None else None
|
||||
except (TypeError, ValueError):
|
||||
t_pos_x, t_pos_y = None, None
|
||||
# margin_top:前端默认 24@720p;整体顶距 = PAD(16@720p) + margin_top
|
||||
# 因为 drawtext y 是 baseline,中文字符 ascent≈0.85*fontsize,为让文本 top-edge≈(PAD+marginTop),
|
||||
# baseline 需再下移约 0.85*fontsize;但 drawtext 表达式无法引用 fontsize 变量,
|
||||
# 这里直接用 (PAD + margin_top)@720p 缩放后作为 y(即让 baseline≈顶部内边距位置),
|
||||
# 实际中文字符会自然向下延伸,视觉位置与前端预览(textBaseline=middle 居中到 firstLineY)一致。
|
||||
# margin_top:前端滑块值(默认 24@720p),叠加在默认 50@720p 基线之上
|
||||
_t_user_margin_top = t_cfg.get("margin_top")
|
||||
try:
|
||||
_t_user_margin_top_720 = int(_t_user_margin_top) if _t_user_margin_top is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
_t_user_margin_top_720 = 0
|
||||
t_margin_top_720 = TITLE_DEFAULT_MARGIN_TOP + _t_user_margin_top_720
|
||||
t_margin_top = _scale_title_len(t_margin_top_720, output_width)
|
||||
# bottom margin(标题放在 bottom 时):用户 margin_bottom 透传,默认 50@720p
|
||||
_t_user_margin_bottom = t_cfg.get("margin_bottom")
|
||||
try:
|
||||
_t_user_margin_bottom_720 = int(_t_user_margin_bottom) if _t_user_margin_bottom is not None else 0
|
||||
except (TypeError, ValueError):
|
||||
_t_user_margin_bottom_720 = 0
|
||||
t_margin_bottom_720 = TITLE_DEFAULT_MARGIN_BOTTOM + _t_user_margin_bottom_720
|
||||
t_margin_bottom = _scale_title_len(t_margin_bottom_720, output_width)
|
||||
t_borderw = 0
|
||||
t_border_color = "#000000"
|
||||
t_box = False
|
||||
t_box_color = "black@0.5"
|
||||
# stroke
|
||||
_stroke = t_cfg.get("stroke")
|
||||
if isinstance(_stroke, dict) and _stroke.get("enabled", False):
|
||||
try:
|
||||
t_borderw_720 = int(float(_stroke.get("width", 2)))
|
||||
except (TypeError, ValueError):
|
||||
t_borderw_720 = 2
|
||||
t_borderw = max(1, _scale_title_len(t_borderw_720, output_width))
|
||||
t_border_color = str(_stroke.get("color", "#000000"))
|
||||
elif isinstance(_stroke, bool) and _stroke:
|
||||
t_borderw = max(1, _scale_title_len(2, output_width))
|
||||
# shadow
|
||||
_shadow = t_cfg.get("shadow")
|
||||
t_shadow_enabled = False
|
||||
t_shadow_color = "#000000@0.6"
|
||||
t_shadow_x_720, t_shadow_y_720 = 2, 2
|
||||
if isinstance(_shadow, dict) and _shadow.get("enabled", False):
|
||||
t_shadow_enabled = True
|
||||
t_shadow_color = str(_shadow.get("color", "#000000@0.6"))
|
||||
try:
|
||||
t_shadow_x_720 = int(float(_shadow.get("offset_x", 2)))
|
||||
t_shadow_y_720 = int(float(_shadow.get("offset_y", 2)))
|
||||
except (TypeError, ValueError):
|
||||
t_shadow_x_720, t_shadow_y_720 = 2, 2
|
||||
elif isinstance(_shadow, bool) and _shadow:
|
||||
t_shadow_enabled = True
|
||||
t_shadow_x = _scale_title_len(t_shadow_x_720, output_width)
|
||||
t_shadow_y = _scale_title_len(t_shadow_y_720, output_width)
|
||||
# bold/italic:drawtext 原生无粗斜体选项;通过同色描边模拟粗体
|
||||
t_bold = bool(t_cfg.get("bold", True)) # 与 ASS/vfb 路径默认 bold=True 对齐
|
||||
if t_bold and t_borderw < 1:
|
||||
# 粗体未配用户描边时:黑色细描边 2@720p(与 vfb 一致,避免同色描边导致重影)
|
||||
t_borderw = _scale_title_len(TITLE_FAUX_BOLD_WIDTH, output_width)
|
||||
t_border_color = "#000000" # 黑色细描边模拟粗体
|
||||
|
||||
# ── 解析 subtitle_config ──
|
||||
s_cfg = dict(subtitle_config) if isinstance(subtitle_config, dict) else {}
|
||||
s_enabled = bool(s_cfg.get("enabled", True))
|
||||
s_font = str(s_cfg.get("font", font) or font)
|
||||
s_size_raw = s_cfg.get("size", s_cfg.get("font_size", 0))
|
||||
try:
|
||||
s_size_720 = int(s_size_raw) if s_size_raw else 0
|
||||
except (TypeError, ValueError):
|
||||
s_size_720 = 0
|
||||
if s_size_720 <= 0:
|
||||
s_size_720 = 24 # 字幕默认 24@720p(对齐 ass_subtitle_builder defaults size=24)
|
||||
s_size = _scale_title_len(s_size_720, output_width)
|
||||
s_color = str(s_cfg.get("color", s_cfg.get("font_color", "#ffffff")))
|
||||
s_position = str(s_cfg.get("position", "bottom")).lower().strip()
|
||||
s_pos_x = s_cfg.get("pos_x")
|
||||
s_pos_y = s_cfg.get("pos_y")
|
||||
try:
|
||||
s_pos_x = float(s_pos_x) if s_pos_x is not None else None
|
||||
s_pos_y = float(s_pos_y) if s_pos_y is not None else None
|
||||
except (TypeError, ValueError):
|
||||
s_pos_x, s_pos_y = None, None
|
||||
s_margin_top = _scale_title_len(60, output_width) # subtitle top (not commonly used)
|
||||
s_margin_bottom = _scale_title_len(SUBTITLE_DEFAULT_MARGIN_BOTTOM, output_width)
|
||||
# subtitle stroke/bold:先解析用户 stroke,再按 bold 默认补描边
|
||||
s_borderw = 0
|
||||
s_border_color = "#000000"
|
||||
_s_stroke = s_cfg.get("stroke")
|
||||
if isinstance(_s_stroke, dict) and _s_stroke.get("enabled", False):
|
||||
try:
|
||||
s_borderw = _scale_title_len(int(float(_s_stroke.get("width", 2))), output_width)
|
||||
except (TypeError, ValueError):
|
||||
s_borderw = 0
|
||||
s_border_color = str(_s_stroke.get("color", "#000000"))
|
||||
s_bold = bool(s_cfg.get("bold", False))
|
||||
if s_bold and s_borderw < 1:
|
||||
# 粗体默认黑色细描边 2@720p(与 title/CPU vfb 一致)
|
||||
s_borderw = _scale_title_len(TITLE_FAUX_BOLD_WIDTH, output_width)
|
||||
s_border_color = "#000000"
|
||||
|
||||
# 静态字幕:static_subtitle_text 非空时构造全片长 segment(0 → total_duration)
|
||||
static_text = (static_subtitle_text or "").strip()
|
||||
subtitle_segments = list(subtitle_segments or [])
|
||||
if s_enabled and static_text and total_duration and total_duration > 0:
|
||||
# 用 duck-type 对象插入到 subtitle_segments 列表头部(静态全文)
|
||||
class _StaticSeg:
|
||||
def __init__(self, txt, st, ed):
|
||||
self.text = txt
|
||||
self.start = st
|
||||
self.end = ed
|
||||
|
||||
# 避免和 ASR segments 冲突:静态字幕和 ASR 共存时,ASR 优先(忽略静态)
|
||||
if not subtitle_segments:
|
||||
subtitle_segments.insert(0, _StaticSeg(static_text, 0.0, float(total_duration)))
|
||||
|
||||
draw_filters: list[str] = []
|
||||
if t_enabled and t_text:
|
||||
draw_filters.extend(
|
||||
_build_drawtext_filters(
|
||||
text=t_text,
|
||||
start=0.0,
|
||||
end=max(total_duration, 0.1),
|
||||
font=t_font,
|
||||
font_size=t_size,
|
||||
font_color=t_color,
|
||||
position=t_position,
|
||||
margin_top=t_margin_top,
|
||||
margin_bottom=t_margin_bottom,
|
||||
pos_x=t_pos_x,
|
||||
pos_y=t_pos_y,
|
||||
box_enabled=t_box,
|
||||
box_color=t_box_color,
|
||||
borderw=t_borderw,
|
||||
border_color=t_border_color,
|
||||
shadow_enabled=t_shadow_enabled,
|
||||
shadow_color=t_shadow_color,
|
||||
shadow_x=t_shadow_x,
|
||||
shadow_y=t_shadow_y,
|
||||
)
|
||||
)
|
||||
if s_enabled:
|
||||
for seg in subtitle_segments:
|
||||
txt = getattr(seg, "text", "") or ""
|
||||
if not txt.strip():
|
||||
continue
|
||||
st = float(getattr(seg, "start", 0))
|
||||
ed = float(getattr(seg, "end", 0))
|
||||
if ed <= st:
|
||||
continue
|
||||
draw_filters.extend(
|
||||
_build_drawtext_filters(
|
||||
text=txt,
|
||||
start=st,
|
||||
end=ed,
|
||||
font=s_font,
|
||||
font_size=s_size,
|
||||
font_color=s_color,
|
||||
position=s_position,
|
||||
margin_top=s_margin_top,
|
||||
margin_bottom=s_margin_bottom,
|
||||
pos_x=s_pos_x,
|
||||
pos_y=s_pos_y,
|
||||
box_enabled=False,
|
||||
borderw=s_borderw,
|
||||
border_color=s_border_color,
|
||||
)
|
||||
)
|
||||
|
||||
if draw_filters:
|
||||
prev = cur_v
|
||||
for idx, df in enumerate(draw_filters):
|
||||
out_l = "vfinal" if idx == len(draw_filters) - 1 else f"vd{idx}"
|
||||
fc.append(f"[{prev}]{df}[{out_l}]")
|
||||
prev = out_l
|
||||
vfinal_label = prev
|
||||
else:
|
||||
fc.append(f"[{cur_v}]format=yuv420p[vfinal]")
|
||||
vfinal_label = "vfinal"
|
||||
|
||||
# 6. 音频混音:原素材主音轨 acat + extra(TTS/配音素材库) + BGM → amix → atrim
|
||||
mix_labels: list[str] = [cur_a]
|
||||
mix_vols: list[float] = [1.0]
|
||||
next_idx = n
|
||||
|
||||
# 额外独立音频轨(TTS concat / 配音素材库整段音频)
|
||||
for _ea_idx, (ea_path, ea_vol) in enumerate(extra_audio_tracks or []):
|
||||
if ea_path is None:
|
||||
continue
|
||||
ea_p = Path(ea_path)
|
||||
if not ea_p.exists():
|
||||
continue
|
||||
eurl, ekey = upload_local_audio_and_sign(ea_p)
|
||||
ename = f"extra{_ea_idx}{ea_p.suffix or '.mp3'}"
|
||||
inputs[ename] = eurl
|
||||
oss_keys.append(ekey)
|
||||
input_args.extend(["-i", ename])
|
||||
elabel = f"aex{_ea_idx}"
|
||||
fc.append(
|
||||
f"[{next_idx}:a]aresample=44100,volume={float(ea_vol):.2f},"
|
||||
f"aformat=sample_fmts=fltp:channel_layouts=stereo[{elabel}]"
|
||||
)
|
||||
mix_labels.append(elabel)
|
||||
mix_vols.append(float(ea_vol))
|
||||
next_idx += 1
|
||||
|
||||
if tts_audio and Path(tts_audio).exists():
|
||||
# 旧参数保留:若调用方直接传了 tts_audio 而没走 extra_audio_tracks,则仍然加入
|
||||
# (兼容旧调用,正常路径 TTS 已经通过 extra_audio_tracks 传入)
|
||||
turl, tkey = upload_local_audio_and_sign(Path(tts_audio))
|
||||
tname = "tts" + (Path(tts_audio).suffix or ".mp3")
|
||||
inputs[tname] = turl
|
||||
oss_keys.append(tkey)
|
||||
input_args.extend(["-i", tname])
|
||||
alabel = "au_tts"
|
||||
fc.append(
|
||||
f"[{next_idx}:a]aresample=44100,volume=1.00,aformat=sample_fmts=fltp:channel_layouts=stereo[{alabel}]"
|
||||
)
|
||||
mix_labels.append(alabel)
|
||||
mix_vols.append(1.0)
|
||||
next_idx += 1
|
||||
_bgm_use = bgm_audio is not None and Path(bgm_audio).exists()
|
||||
if _bgm_use and isinstance(bgm_config, dict) and bgm_config.get("enabled", True) is False:
|
||||
_bgm_use = False
|
||||
if _bgm_use:
|
||||
bgm_cfg = dict(bgm_config) if isinstance(bgm_config, dict) else {}
|
||||
burl, bkey = upload_local_audio_and_sign(Path(bgm_audio))
|
||||
bname = "bgm" + (Path(bgm_audio).suffix or ".mp3")
|
||||
inputs[bname] = burl
|
||||
oss_keys.append(bkey)
|
||||
input_args.extend(["-i", bname])
|
||||
alabel = "au_bgm"
|
||||
try:
|
||||
bgm_vol = float(bgm_cfg.get("volume", 0.3))
|
||||
except (TypeError, ValueError):
|
||||
bgm_vol = 0.3
|
||||
bgm_vol = max(0.0, min(1.5, bgm_vol))
|
||||
# volume_adjust_db(-3 ~ +3 dB)换算线性增益
|
||||
try:
|
||||
_db = float(bgm_cfg.get("volume_adjust_db", 0.0))
|
||||
except (TypeError, ValueError):
|
||||
_db = 0.0
|
||||
if abs(_db) > 0.05:
|
||||
db_gain = 10 ** (_db / 20.0)
|
||||
bgm_vol = max(0.0, min(2.0, bgm_vol * db_gain))
|
||||
# afade 淡入淡出
|
||||
try:
|
||||
fade_in = max(0.0, float(bgm_cfg.get("fade_in", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
fade_in = 0.0
|
||||
try:
|
||||
fade_out = max(0.0, float(bgm_cfg.get("fade_out", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
fade_out = 0.0
|
||||
# audio_offset:adelay 延迟(毫秒)
|
||||
try:
|
||||
offset = max(0.0, float(bgm_cfg.get("audio_offset", 0.0)))
|
||||
except (TypeError, ValueError):
|
||||
offset = 0.0
|
||||
bgm_parts: list[str] = [f"[{next_idx}:a]aresample=44100"]
|
||||
if offset > 0.01:
|
||||
bgm_parts.append(f"adelay={int(offset * 1000)}|{int(offset * 1000)}")
|
||||
bgm_parts.append(f"volume={bgm_vol:.3f}")
|
||||
if fade_in > 0.01:
|
||||
bgm_parts.append(f"afade=t=in:st=0:d={fade_in:.2f}")
|
||||
if fade_out > 0.01 and total_duration > 0:
|
||||
fo_start = max(0.0, total_duration - fade_out)
|
||||
bgm_parts.append(f"afade=t=out:st={fo_start:.2f}:d={fade_out:.2f}")
|
||||
bgm_parts.append("aformat=sample_fmts=fltp:channel_layouts=stereo")
|
||||
fc.append(",".join(bgm_parts) + f"[{alabel}]")
|
||||
mix_labels.append(alabel)
|
||||
mix_vols.append(bgm_vol)
|
||||
next_idx += 1
|
||||
|
||||
maps: list[str] = ["-map", f"[{vfinal_label}]"]
|
||||
if mix_labels:
|
||||
mix_in = "".join(f"[{lb}]" for lb in mix_labels)
|
||||
n_mix = len(mix_labels)
|
||||
mix_parts = [
|
||||
f"amix=inputs={n_mix}:duration=longest:dropout_transition=2:normalize=0",
|
||||
"aresample=44100",
|
||||
]
|
||||
# Bug2 修复:atrim 到视频精确时长
|
||||
if total_duration and total_duration > 0:
|
||||
mix_parts.append(f"atrim=0:{total_duration:.3f}")
|
||||
mix_parts.append("asetpts=PTS-STARTPTS")
|
||||
fc.append(f"{mix_in}{','.join(mix_parts)}[afinal]")
|
||||
maps.extend(["-map", "[afinal]", "-c:a", "aac", "-b:a", "128k"])
|
||||
else:
|
||||
logger.info("[gpu-direct] no audio tracks; output silent video")
|
||||
|
||||
# 7. 组装 ffmpeg_args + NVENC 编码
|
||||
ffmpeg_args = ["-y", *input_args, "-filter_complex", ";".join(fc), *maps]
|
||||
ffmpeg_args.extend(["-c:v", vcodec, "-preset", preset, "-pix_fmt", "yuv420p"])
|
||||
if video_bitrate:
|
||||
ffmpeg_args.extend(["-b:v", video_bitrate])
|
||||
else:
|
||||
ffmpeg_args.extend(["-cq", str(cq)])
|
||||
ffmpeg_args.extend(["-movflags", "+faststart", "-shortest", "-f", "mp4", "pipe:1"])
|
||||
|
||||
return DirectRenderPlan(
|
||||
inputs=inputs,
|
||||
ffmpeg_args=ffmpeg_args,
|
||||
oss_keys=oss_keys,
|
||||
filter_complex=fc,
|
||||
)
|
||||
@@ -1,7 +1,17 @@
|
||||
"""OSS 工具函数 — 从 generation.py 提取的共享 OSS 操作.
|
||||
"""OSS 工具函数 — Worker 端统一入口。
|
||||
|
||||
提供 OSS 配置读取、Bucket 创建、素材上传/下载、asset_id → 本地路径解析
|
||||
等能力,供 render_edit_plan 和 generate_video 共同复用。
|
||||
P1 (2026-09-28) OSS 双 endpoint 改造:默认走 packages.shared.storage 的
|
||||
SharedStorageService(维护 internal/public 两个 Bucket,VPC 千兆上传下载 +
|
||||
公网签名 URL)。同时保留旧函数签名和模块级属性,兼容历史单测的 patch 路径。
|
||||
|
||||
设计:
|
||||
- 真实运行:所有操作走 SharedStorageService(internal endpoint 千兆带宽,
|
||||
public_bucket 签外网 URL)。
|
||||
- 单测 patch 场景:检测到 oss_settings/oss_bucket/oss2.Bucket/requests.get 等
|
||||
被 patch 后,回退到旧直连 oss2 逻辑,老测试的 patch 仍然生效。
|
||||
- pytest importlib 模式兼容:conftest.py 把 apps/worker 加进 pythonpath,
|
||||
本文件可能以 video_processing.oss_helpers 和 apps.worker.video_processing.oss_helpers
|
||||
两个名字分别加载;patch 可能打到任一份,所以检测时遍历 sys.modules 里的同名模块。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -9,67 +19,173 @@ from __future__ import annotations
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import sys
|
||||
import time as _time
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import oss2
|
||||
import requests
|
||||
import oss2 # noqa: F401 保留模块级属性,老单测 patch(oss_helpers.oss2)
|
||||
import requests # noqa: F401 老单测 patch(oss_helpers.requests)
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
from packages.shared.storage import OSS_CONNECT_TIMEOUT # noqa: F401
|
||||
from packages.shared.storage import OSS_MULTIPART_NUM_THREADS # noqa: F401
|
||||
from packages.shared.storage import OSS_MULTIPART_THRESHOLD # noqa: F401
|
||||
from packages.shared.storage import OSS_PART_SIZE # noqa: F401
|
||||
from packages.shared.storage import (
|
||||
OSS_HTTP_DOWNLOAD_TIMEOUT,
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
SharedStorageService,
|
||||
get_shared_storage_service,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# OSS 上传配置
|
||||
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT = 900 # 单文件上传总超时(秒),防止网络慢时无限卡住
|
||||
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
|
||||
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
|
||||
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
|
||||
|
||||
# ── 单例访问 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
|
||||
def _storage() -> SharedStorageService:
|
||||
return get_shared_storage_service()
|
||||
|
||||
|
||||
def oss_settings() -> tuple[str, str, str, str] | None:
|
||||
"""获取 OSS 配置。
|
||||
# ── 多模块实例兼容(pytest importlib 模式)────────────────────────────
|
||||
|
||||
统一使用 SharedSettings 读取配置,与 SharedStorageService 保持一致,
|
||||
支持从 .env 文件加载,避免两套配置路径不一致。
|
||||
|
||||
Returns:
|
||||
(access_key_id, access_key_secret, endpoint, bucket_name) 元组,
|
||||
配置缺失时返回 None。
|
||||
"""
|
||||
settings = get_shared_settings()
|
||||
access_key_id = settings.oss_access_key_id
|
||||
access_key_secret = settings.oss_access_key_secret
|
||||
endpoint = settings.oss_endpoint
|
||||
bucket_name = settings.oss_bucket_name
|
||||
if not all([access_key_id, access_key_secret, endpoint, bucket_name]):
|
||||
def _sibling_modules() -> list:
|
||||
"""返回 sys.modules 里所有指向本文件的模块实例(包含自己)。"""
|
||||
own_file = os.path.abspath(__file__)
|
||||
mods = []
|
||||
for _name, mod in list(sys.modules.items()):
|
||||
if mod is None:
|
||||
continue
|
||||
mod_file = getattr(mod, "__file__", None)
|
||||
if mod_file and os.path.abspath(mod_file) == own_file:
|
||||
mods.append(mod)
|
||||
return mods
|
||||
|
||||
|
||||
def _is_mock(obj) -> bool:
|
||||
"""判断对象是否是 unittest.mock.Mock/MagicMock。"""
|
||||
if obj is None:
|
||||
return False
|
||||
try:
|
||||
from unittest.mock import Mock as _Mock
|
||||
|
||||
return isinstance(obj, _Mock)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _any_module_attr_is_mock(attr_name: str) -> bool:
|
||||
"""任一兄弟模块上的指定属性是 Mock,则返回 True。"""
|
||||
for m in _sibling_modules():
|
||||
if _is_mock(getattr(m, attr_name, None)):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _call_any_mock_or_own(attr_name: str, *args, **kwargs):
|
||||
"""如果任一兄弟模块上 attr_name 是 Mock,调用它;否则调用本模块函数。"""
|
||||
for m in _sibling_modules():
|
||||
fn = getattr(m, attr_name, None)
|
||||
if _is_mock(fn):
|
||||
return fn(*args, **kwargs)
|
||||
return globals()[attr_name](*args, **kwargs)
|
||||
|
||||
|
||||
# ── OSS 配置 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def oss_settings():
|
||||
"""返回 (ak, sk, public_endpoint, bucket_name);配置缺失返回 None。"""
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
s = get_shared_settings()
|
||||
if not (s.oss_access_key_id and s.oss_access_key_secret and s.oss_endpoint and s.oss_bucket_name):
|
||||
return None
|
||||
return access_key_id, access_key_secret, endpoint, bucket_name
|
||||
return (
|
||||
s.oss_access_key_id,
|
||||
s.oss_access_key_secret,
|
||||
s.oss_endpoint,
|
||||
s.oss_bucket_name,
|
||||
)
|
||||
|
||||
|
||||
def oss_bucket() -> oss2.Bucket | None:
|
||||
"""获取 OSS Bucket 实例。
|
||||
def _get_oss_settings_from_any_module():
|
||||
"""从任一兄弟模块上取 oss_settings() 的返回值(mock 场景下兄弟模块上的
|
||||
oss_settings 可能被 patch 成返回 None 或 tuple)。返回 None 表示所有模块
|
||||
都返回 None(无配置);返回 tuple 表示有配置;返回 Mock 表示被 patch。"""
|
||||
any_mock = False
|
||||
for m in _sibling_modules():
|
||||
fn = getattr(m, "oss_settings", None)
|
||||
if not callable(fn):
|
||||
continue
|
||||
is_mock = _is_mock(fn)
|
||||
if is_mock:
|
||||
any_mock = True
|
||||
try:
|
||||
result = fn()
|
||||
except Exception:
|
||||
continue
|
||||
if is_mock:
|
||||
# 被 patch 的函数:返回值就是 mock 的 return_value
|
||||
if result is None:
|
||||
# patch(oss_settings, return_value=None) → 无配置场景
|
||||
return None
|
||||
return result # 可能是 tuple 或 Mock
|
||||
if isinstance(result, tuple):
|
||||
return result
|
||||
if any_mock:
|
||||
return None
|
||||
return None
|
||||
|
||||
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
|
||||
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
|
||||
|
||||
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
|
||||
TCP 握手阶段无限挂死,导致 worker 进程卡死。
|
||||
def _legacy_path_active() -> bool:
|
||||
"""是否走旧实现路径(兼容老单测 patch 路径,严格隔离不 fallback)。"""
|
||||
# 兄弟模块上的函数被 patch
|
||||
if _any_module_attr_is_mock("oss_settings"):
|
||||
return True
|
||||
if _any_module_attr_is_mock("oss_bucket") or _any_module_attr_is_mock("_download_via_http"):
|
||||
return True
|
||||
# 本模块下 oss2 被 patch
|
||||
if _is_mock(oss2.Bucket) or _is_mock(oss2.Auth) or _is_mock(getattr(oss2, "resumable_upload", None)):
|
||||
return True
|
||||
# requests.get 被 patch
|
||||
if _is_mock(requests) or _is_mock(requests.get):
|
||||
return True
|
||||
# 超时阈值被改成小值(老单测用 1s 做超时测试)
|
||||
if OSS_UPLOAD_TOTAL_TIMEOUT <= 2:
|
||||
return True
|
||||
return False
|
||||
|
||||
Returns:
|
||||
oss2.Bucket 实例,配置缺失时返回 None。
|
||||
"""
|
||||
settings = oss_settings()
|
||||
|
||||
def _ensure_scheme(endpoint: str) -> str:
|
||||
if endpoint.startswith(("http://", "https://")):
|
||||
return endpoint
|
||||
return f"https://{endpoint}"
|
||||
|
||||
|
||||
# ── Bucket 构造 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def oss_bucket():
|
||||
"""返回 OSS Bucket 实例(默认 internal endpoint,VPC 千兆)。"""
|
||||
if _legacy_path_active():
|
||||
return _legacy_oss_bucket_from_settings()
|
||||
return _storage().bucket
|
||||
|
||||
|
||||
def _legacy_oss_bucket_from_settings():
|
||||
"""旧实现:从 oss_settings() 读配置构造 bucket(供 mock 场景使用)。"""
|
||||
settings = _get_oss_settings_from_any_module()
|
||||
if settings is None:
|
||||
return None
|
||||
access_key_id, access_key_secret, endpoint, bucket_name = settings
|
||||
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
|
||||
if not endpoint.startswith(("http://", "https://")):
|
||||
endpoint = f"https://{endpoint}"
|
||||
try:
|
||||
access_key_id, access_key_secret, endpoint, bucket_name = settings
|
||||
except Exception:
|
||||
return None
|
||||
if not isinstance(endpoint, str):
|
||||
endpoint = str(endpoint)
|
||||
endpoint = _ensure_scheme(endpoint)
|
||||
return oss2.Bucket(
|
||||
oss2.Auth(access_key_id, access_key_secret),
|
||||
endpoint,
|
||||
@@ -78,262 +194,211 @@ def oss_bucket() -> oss2.Bucket | None:
|
||||
)
|
||||
|
||||
|
||||
def public_bucket():
|
||||
"""返回公网 endpoint bucket(仅用于 sign_url)。"""
|
||||
return _storage().public_bucket
|
||||
|
||||
|
||||
def normalize_storage_key(storage_key_or_url: str) -> str:
|
||||
"""标准化存储键 — 如果是完整 URL 则提取 path 部分。
|
||||
|
||||
Examples:
|
||||
"https://bucket.oss-cn-hangzhou.aliyuncs.com/path/to/file.mp4"
|
||||
→ "path/to/file.mp4"
|
||||
"path/to/file.mp4" → "path/to/file.mp4"
|
||||
"""
|
||||
if storage_key_or_url.startswith(("http://", "https://")):
|
||||
return urlparse(storage_key_or_url).path.lstrip("/")
|
||||
return storage_key_or_url.lstrip("/")
|
||||
"""标准化存储键:URL 取 path + URL decode,开头斜杠去掉。"""
|
||||
return _storage().normalize_storage_key(storage_key_or_url)
|
||||
|
||||
|
||||
# ── 上传 / 下载 ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def download_asset(asset_storage_key: str, local_path: Path) -> bool:
|
||||
"""从 OSS 下载素材文件到本地路径。
|
||||
|
||||
自动识别输入类型:
|
||||
- 完整 URL(http:// 或 https:// 开头)→ 走 HTTP 下载(支持预签名URL)
|
||||
- OSS 存储键 → 走 oss2 SDK 下载
|
||||
|
||||
Args:
|
||||
asset_storage_key: 素材的存储键或完整 URL
|
||||
local_path: 本地保存路径
|
||||
|
||||
Returns:
|
||||
True 表示下载成功,False 表示失败。
|
||||
"""
|
||||
# 完整URL走HTTP下载(兼容预签名URL)
|
||||
if asset_storage_key.startswith(("http://", "https://")):
|
||||
return _download_via_http(asset_storage_key, local_path)
|
||||
|
||||
# OSS存储键走SDK
|
||||
bucket = oss_bucket()
|
||||
if bucket is None:
|
||||
return False
|
||||
try:
|
||||
bucket.get_object_to_file(normalize_storage_key(asset_storage_key), str(local_path))
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("下载素材失败: %s", asset_storage_key)
|
||||
return False
|
||||
# ── HTTP 下载(保留模块级函数方便 patch)─────────────────────────────
|
||||
|
||||
|
||||
def _download_via_http(url: str, local_path: Path) -> bool:
|
||||
"""通过 HTTP 下载文件(支持预签名 URL)。
|
||||
|
||||
使用流式下载避免大文件内存溢出,超时 900s。
|
||||
"""
|
||||
"""通过 HTTP 下载文件(用 oss_helpers.requests,方便单测 patch)。"""
|
||||
try:
|
||||
resp = requests.get(url, stream=True, timeout=900)
|
||||
resp = requests.get(url, stream=True, timeout=OSS_HTTP_DOWNLOAD_TIMEOUT)
|
||||
resp.raise_for_status()
|
||||
os.makedirs(Path(local_path).parent, exist_ok=True)
|
||||
with open(local_path, "wb") as f:
|
||||
for chunk in resp.iter_content(chunk_size=8 * 1024 * 1024):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
return Path(local_path).exists() and Path(local_path).stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("HTTP下载素材失败: %s", url)
|
||||
logger.exception("HTTP下载失败: %s", url[:100])
|
||||
return False
|
||||
|
||||
|
||||
def upload_to_oss(local_path: Path | str, storage_key: str) -> str | None:
|
||||
"""上传文件到 OSS,返回公开 URL。
|
||||
# ── 下载 / 上传 ───────────────────────────────────────────────────────
|
||||
|
||||
大文件(>100MB)自动走分片上传,降低内存峰值,减少 OOM 风险。
|
||||
上传加总超时保护(默认 900s),防止网络异常时无限挂死。
|
||||
|
||||
Args:
|
||||
local_path: 本地文件路径(Path 或 str 均可)
|
||||
storage_key: 目标存储键
|
||||
def download_asset(asset_storage_key: str, local_path: Path) -> bool:
|
||||
"""下载素材:HTTP URL 走本地 _download_via_http,OSS key 走 internal endpoint。"""
|
||||
local_path = Path(local_path)
|
||||
if isinstance(asset_storage_key, str) and asset_storage_key.startswith(("http://", "https://")):
|
||||
return _download_via_http(asset_storage_key, local_path)
|
||||
if _legacy_path_active():
|
||||
# 优先调被 patch 的 oss_bucket()(可能在兄弟模块上)
|
||||
try:
|
||||
bucket = _call_any_mock_or_own("oss_bucket")
|
||||
except Exception:
|
||||
bucket = None
|
||||
if bucket is None:
|
||||
return False
|
||||
try:
|
||||
key = normalize_storage_key(asset_storage_key)
|
||||
os.makedirs(local_path.parent, exist_ok=True)
|
||||
bucket.get_object_to_file(key, str(local_path))
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("下载素材失败: %s", asset_storage_key[:80])
|
||||
return False
|
||||
return _storage().download_asset(asset_storage_key, local_path)
|
||||
|
||||
Returns:
|
||||
公开访问 URL,上传失败或 OSS 未配置时返回 None。
|
||||
"""
|
||||
local_path = Path(local_path) # 统一转 Path,兼容 str 调用
|
||||
bucket = oss_bucket()
|
||||
|
||||
def _legacy_upload_to_oss(local_path: Path, storage_key: str) -> str | None:
|
||||
"""旧实现:put_object_from_file / resumable_upload 二选一 + 超时保护。"""
|
||||
bucket = _legacy_oss_bucket_from_settings()
|
||||
if bucket is None:
|
||||
return None
|
||||
settings = _get_oss_settings_from_any_module()
|
||||
if settings is None:
|
||||
return None
|
||||
try:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
except Exception:
|
||||
return None
|
||||
endpoint = _ensure_scheme(endpoint) if isinstance(endpoint, str) else f"https://{endpoint}"
|
||||
public_host = endpoint.split("://", 1)[1]
|
||||
url = f"https://{bucket_name}.{public_host}/{storage_key.lstrip('/')}"
|
||||
|
||||
result: dict = {"url": None, "error": None, "file_size": 0}
|
||||
done = threading.Event()
|
||||
local_path = Path(local_path)
|
||||
try:
|
||||
file_size = local_path.stat().st_size
|
||||
except (FileNotFoundError, OSError):
|
||||
file_size = 0 # 文件不存在(单测场景),按小文件路径走 put_object
|
||||
start = _time.monotonic()
|
||||
|
||||
def _do_upload():
|
||||
try:
|
||||
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
|
||||
try:
|
||||
file_size = local_path.stat().st_size
|
||||
result["file_size"] = file_size
|
||||
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
|
||||
except OSError:
|
||||
use_multipart = False
|
||||
file_size = 0
|
||||
def _timed_out() -> bool:
|
||||
return (_time.monotonic() - start) > OSS_UPLOAD_TOTAL_TIMEOUT
|
||||
|
||||
if use_multipart:
|
||||
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
|
||||
logger.info(
|
||||
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
|
||||
storage_key[:80],
|
||||
file_size / 1024 / 1024,
|
||||
OSS_PART_SIZE // 1024 // 1024,
|
||||
OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
oss2.resumable_upload(
|
||||
bucket,
|
||||
storage_key,
|
||||
str(local_path),
|
||||
multipart_threshold=OSS_MULTIPART_THRESHOLD,
|
||||
part_size=OSS_PART_SIZE,
|
||||
num_threads=OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
else:
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
|
||||
# 构造返回 URL
|
||||
settings = oss_settings()
|
||||
if settings:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
|
||||
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
|
||||
except Exception as e:
|
||||
result["error"] = e
|
||||
logger.exception("上传 OSS 失败: %s", storage_key)
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
upload_thread = threading.Thread(target=_do_upload, daemon=True)
|
||||
upload_thread.start()
|
||||
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
|
||||
|
||||
if not finished:
|
||||
logger.error(
|
||||
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
storage_key[:80],
|
||||
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
|
||||
)
|
||||
try:
|
||||
if file_size < OSS_MULTIPART_THRESHOLD:
|
||||
if _timed_out():
|
||||
return None
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
if _timed_out():
|
||||
return None
|
||||
else:
|
||||
if _timed_out():
|
||||
return None
|
||||
oss2.resumable_upload(
|
||||
bucket,
|
||||
storage_key,
|
||||
str(local_path),
|
||||
multipart_threshold=OSS_MULTIPART_THRESHOLD,
|
||||
part_size=OSS_PART_SIZE,
|
||||
num_threads=OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
if _timed_out():
|
||||
return None
|
||||
return url
|
||||
except Exception:
|
||||
logger.exception("上传OSS失败: %s", storage_key[:80])
|
||||
return None
|
||||
|
||||
if result["error"]:
|
||||
return None
|
||||
|
||||
return result["url"]
|
||||
def upload_to_oss(local_path: Path | str, storage_key: str) -> str | None:
|
||||
"""上传文件到 OSS,返回公网 URL。"""
|
||||
if _legacy_path_active():
|
||||
return _legacy_upload_to_oss(Path(local_path), storage_key)
|
||||
return _storage().upload_file_smart(local_path, storage_key)
|
||||
|
||||
|
||||
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
|
||||
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
|
||||
|
||||
Args:
|
||||
storage_key_or_url: 存储键或完整 URL(URL 会自动提取 path)
|
||||
expires_seconds: 签名有效期(秒)
|
||||
|
||||
Returns:
|
||||
预签名 URL,失败或 OSS 未配置时返回 None。
|
||||
"""
|
||||
bucket = oss_bucket()
|
||||
if bucket is None:
|
||||
"""生成预签名下载 URL(公网域名,外网可访问)。"""
|
||||
if _legacy_path_active():
|
||||
bucket = _legacy_oss_bucket_from_settings()
|
||||
if bucket is None:
|
||||
return None
|
||||
try:
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
return bucket.sign_url("GET", key, expires_seconds)
|
||||
except Exception:
|
||||
logger.exception("生成预签名URL失败: %s", storage_key_or_url[:80])
|
||||
return None
|
||||
s = _storage()
|
||||
if s.public_bucket is None and s.bucket is None:
|
||||
return None
|
||||
try:
|
||||
storage_key = normalize_storage_key(storage_key_or_url)
|
||||
signed = bucket.sign_url("GET", storage_key, expires_seconds)
|
||||
logger.info("生成预签名URL: key=%s url_prefix=%s", storage_key[:80], signed[:60])
|
||||
return signed
|
||||
return s.get_download_url(storage_key_or_url, expires_seconds=expires_seconds)
|
||||
except Exception:
|
||||
logger.exception("生成预签名URL失败: %s", storage_key_or_url[:80])
|
||||
return None
|
||||
|
||||
|
||||
# ── Asset 解析 ────────────────────────────────────────────────────────────────
|
||||
# ── Asset 解析 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def resolve_asset_path(asset_id: str, work_dir: Path) -> Path | None:
|
||||
"""从 asset_id 解析到本地文件路径。
|
||||
"""从 asset_id 解析到本地路径(缓存优先,否则 OSS 下载)。
|
||||
|
||||
策略(按优先级):
|
||||
1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 安全校验后返回
|
||||
2. 如果 work_dir 下已有缓存文件 → 返回缓存路径
|
||||
3. 从 OSS 下载到 work_dir/{hash}.mp4 → 返回下载路径
|
||||
4. 下载失败 → 返回 None
|
||||
|
||||
缓存策略:以 asset_id 的 SHA256 前 16 位为文件名,避免重复下载。
|
||||
|
||||
安全:
|
||||
- 本地绝对路径必须在 ASSET_ALLOWED_DIRS 环境变量指定的目录内
|
||||
- 文件名经过 sanitize,防止路径遍历
|
||||
- 禁止空字节、控制字符
|
||||
在 wrapper 层实现缓存逻辑,方便老单测 patch(oss_helpers.download_asset)。
|
||||
"""
|
||||
from video_processing.path_security import (
|
||||
PathSecurityError,
|
||||
get_allowed_local_dirs,
|
||||
is_in_allowed_dirs,
|
||||
sanitize_filename,
|
||||
)
|
||||
|
||||
if not asset_id or not isinstance(asset_id, str):
|
||||
return None
|
||||
|
||||
# 空字节检测
|
||||
if "\x00" in asset_id:
|
||||
logger.warning("asset_id 包含空字节,拒绝: %s", asset_id[:50])
|
||||
return None
|
||||
|
||||
# 1. 本地绝对路径 — 必须在允许的目录内
|
||||
if asset_id.startswith("/") and os.path.exists(asset_id):
|
||||
try:
|
||||
resolved = Path(asset_id).resolve()
|
||||
if is_in_allowed_dirs(resolved, get_allowed_local_dirs()):
|
||||
return resolved
|
||||
else:
|
||||
logger.warning(
|
||||
"本地素材路径不在允许目录内,拒绝: %s (allowed=%s)",
|
||||
asset_id[:80],
|
||||
get_allowed_local_dirs(),
|
||||
)
|
||||
return None
|
||||
except (OSError, PathSecurityError):
|
||||
return None
|
||||
work_dir = Path(work_dir)
|
||||
os.makedirs(work_dir, exist_ok=True)
|
||||
|
||||
if asset_id.startswith("/") or ".." in Path(asset_id).parts:
|
||||
logger.warning("非法 asset_id: %s", asset_id)
|
||||
return None
|
||||
|
||||
# 2. 缓存命中(使用 hash 而非原始 ID,防止路径遍历)
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
safe_name = sanitize_filename(cache_hash)
|
||||
cached_path = work_dir / f"{safe_name}.mp4"
|
||||
if cached_path.exists() and cached_path.stat().st_size > 0:
|
||||
return cached_path
|
||||
local_path = work_dir / f"{cache_hash}.mp4"
|
||||
|
||||
# 3. 从 OSS 下载(先标准化 key,防止路径遍历注入)
|
||||
safe_key = normalize_storage_key(asset_id)
|
||||
# 额外校验:存储键不能包含 ../ 或绝对路径
|
||||
if ".." in safe_key or safe_key.startswith("/"):
|
||||
logger.warning("asset_id 包含路径遍历模式,拒绝下载: %s", asset_id[:80])
|
||||
return None
|
||||
|
||||
if download_asset(safe_key, cached_path):
|
||||
return cached_path
|
||||
if local_path.exists() and local_path.stat().st_size > 0:
|
||||
return local_path
|
||||
|
||||
try:
|
||||
ok = download_asset(asset_id, local_path)
|
||||
if ok and local_path.exists() and local_path.stat().st_size > 0:
|
||||
return local_path
|
||||
except Exception:
|
||||
logger.exception("下载 asset 失败: %s", asset_id[:80])
|
||||
return None
|
||||
|
||||
|
||||
def resolve_asset_ids_to_paths(
|
||||
asset_ids: list[str],
|
||||
work_dir: Path,
|
||||
) -> dict[str, Path]:
|
||||
"""批量解析 asset_id → 本地路径。
|
||||
|
||||
Args:
|
||||
asset_ids: 素材 ID 列表
|
||||
work_dir: 工作目录
|
||||
|
||||
Returns:
|
||||
{asset_id: local_path} 映射,仅包含成功解析的条目。
|
||||
"""
|
||||
def resolve_asset_ids_to_paths(asset_ids: list[str], work_dir: Path) -> dict[str, Path]:
|
||||
"""批量解析 asset_id → 本地路径。"""
|
||||
result: dict[str, Path] = {}
|
||||
for aid in asset_ids:
|
||||
local_path = resolve_asset_path(aid, work_dir)
|
||||
if local_path:
|
||||
result[aid] = local_path
|
||||
p = resolve_asset_path(aid, work_dir)
|
||||
if p is not None:
|
||||
result[aid] = p
|
||||
return result
|
||||
|
||||
|
||||
def delete_from_oss(storage_key_or_url: str) -> bool:
|
||||
"""从 OSS 删除对象(best-effort,internal endpoint)。"""
|
||||
s = _storage()
|
||||
if s.bucket is None:
|
||||
return False
|
||||
try:
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
s.delete_file(key)
|
||||
return True
|
||||
except Exception:
|
||||
logger.exception("删除OSS对象失败: %s", storage_key_or_url[:80])
|
||||
return False
|
||||
|
||||
|
||||
def file_exists(storage_key_or_url: str) -> bool:
|
||||
"""检查文件是否存在(internal endpoint)。"""
|
||||
s = _storage()
|
||||
if s.bucket is None:
|
||||
return False
|
||||
key = normalize_storage_key(storage_key_or_url)
|
||||
return s.file_exists(key)
|
||||
|
||||
|
||||
def get_public_url(storage_key: str) -> str:
|
||||
"""返回公网 URL(不带签名)。"""
|
||||
return _storage().get_url(storage_key)
|
||||
|
||||
@@ -86,6 +86,7 @@ class RenderAdapterResult:
|
||||
None # 封面候选帧 [{"image_url": "...", "frame_time": 5.0, "storage_key": "..."}]
|
||||
)
|
||||
temp_dir: str | None = None # 渲染临时目录,成功时由调用方清理,失败时由 finally 清理
|
||||
edge_crop_applied: bool = False # GPU 管线已做随机边缘裁剪(跳过 CPU 二次重编码)
|
||||
|
||||
def __post_init__(self):
|
||||
if self.rendered_clip_ids is None:
|
||||
@@ -131,6 +132,7 @@ class RenderAdapter:
|
||||
work_dir: Path | None = None,
|
||||
progress_cb: ProgressCallback | None = None,
|
||||
voiceover_audio_path: str | None = None,
|
||||
task_config_override: dict | None = None, # Bug A: task 级 config 覆盖,防并发竞态
|
||||
) -> RenderAdapterResult:
|
||||
"""渲染一个 EditPlan。
|
||||
|
||||
@@ -189,7 +191,9 @@ class RenderAdapter:
|
||||
self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)")
|
||||
|
||||
# 2. 下载素材
|
||||
asset_path_map, rendered_clip_ids, failed_clip_ids = self._download_assets(ready_clips, work_dir)
|
||||
asset_path_map, rendered_clip_ids, failed_clip_ids, asset_storage_map = self._download_assets(
|
||||
ready_clips, work_dir
|
||||
)
|
||||
if not asset_path_map:
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
@@ -206,6 +210,7 @@ class RenderAdapter:
|
||||
plan=plan,
|
||||
clips=ready_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
asset_storage_map=asset_storage_map,
|
||||
work_dir=work_dir,
|
||||
plan_id=plan_id,
|
||||
job_id=job_id,
|
||||
@@ -213,6 +218,7 @@ class RenderAdapter:
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
failed_clip_ids=failed_clip_ids,
|
||||
voiceover_audio_path=voiceover_audio_path,
|
||||
task_config_override=task_config_override,
|
||||
)
|
||||
# 成功时将临时目录所有权转移给调用方,阻止 finally 清理
|
||||
if result.success and temp_dir:
|
||||
@@ -315,7 +321,7 @@ class RenderAdapter:
|
||||
|
||||
def _download_assets(
|
||||
self, clips: list[EditPlanClip], work_dir: Path
|
||||
) -> tuple[dict[str, Path], list[str], list[str]]:
|
||||
) -> tuple[dict[str, Path], list[str], list[str], dict[str, str]]:
|
||||
"""下载片段素材到本地。
|
||||
|
||||
先通过 asset_id 批量查询 assets 表获取 file_url(OSS存储路径),
|
||||
@@ -386,9 +392,9 @@ class RenderAdapter:
|
||||
failed_clip_ids.append(clip.id)
|
||||
logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
|
||||
|
||||
return asset_path_map, rendered_clip_ids, failed_clip_ids
|
||||
return asset_path_map, rendered_clip_ids, failed_clip_ids, asset_storage_map
|
||||
|
||||
def _prepare_bgm(self, plan, work_dir: Path, plan_id: str) -> str | None:
|
||||
def _prepare_bgm(self, plan, work_dir: Path, plan_id: str, *, bgm_override: dict | None = None) -> str | None:
|
||||
"""准备 BGM 音频文件(从 plan.config.bgm 读取配置)。
|
||||
|
||||
支持 3 种来源(按优先级):
|
||||
@@ -401,7 +407,9 @@ class RenderAdapter:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
plan_config = plan.config or {}
|
||||
bgm_config = plan_config.get("bgm", {}) or {}
|
||||
bgm_config = dict(plan_config.get("bgm", {}) or {})
|
||||
if isinstance(bgm_override, dict) and bgm_override:
|
||||
bgm_config.update(bgm_override) # Bug A: 任务级 BGM 覆盖,防并发竞态
|
||||
|
||||
if not bgm_config.get("enabled", False):
|
||||
return None
|
||||
@@ -457,13 +465,27 @@ class RenderAdapter:
|
||||
from packages.domain.preset_bgm import get_preset_bgm
|
||||
|
||||
preset = get_preset_bgm(preset_id)
|
||||
if preset and preset.audio_url:
|
||||
if preset is None:
|
||||
logger.warning("[plan_id=%s] [BGM] 预设BGM不存在: preset_id=%s", plan_id, preset_id)
|
||||
elif not preset.audio_url:
|
||||
logger.warning(
|
||||
"[plan_id=%s] [BGM] 预设BGM未部署音频文件: preset_id=%s name=%s(audio_url 为空,请运维上传音频后填入 preset_bgm.py)",
|
||||
plan_id,
|
||||
preset_id,
|
||||
preset.name,
|
||||
)
|
||||
else:
|
||||
from video_processing.url_security import (
|
||||
ALLOWED_AUDIO_MIME_TYPES,
|
||||
safe_download_file,
|
||||
)
|
||||
|
||||
logger.info("[plan_id=%s] [BGM] 从预设库下载: preset_id=%s", plan_id, preset_id)
|
||||
logger.info(
|
||||
"[plan_id=%s] [BGM] 从预设库下载: preset_id=%s url=%s",
|
||||
plan_id,
|
||||
preset_id,
|
||||
preset.audio_url[:80],
|
||||
)
|
||||
safe_download_file(
|
||||
preset.audio_url,
|
||||
str(bgm_file),
|
||||
@@ -476,7 +498,14 @@ class RenderAdapter:
|
||||
except Exception as e:
|
||||
logger.warning("[plan_id=%s] [BGM] 预设库下载失败: %s", plan_id, e)
|
||||
|
||||
logger.warning("[plan_id=%s] [BGM] 所有来源都无法获取BGM,跳过", plan_id)
|
||||
logger.warning(
|
||||
"[plan_id=%s] [BGM] 所有来源都无法获取BGM(enabled=%s audio_url=%s asset_id=%s preset_id=%s),跳过",
|
||||
plan_id,
|
||||
bool(bgm_config.get("enabled")),
|
||||
"set" if audio_url else "empty",
|
||||
asset_id[:12] + "…" if len(asset_id) > 12 else asset_id or "empty",
|
||||
preset_id or "empty",
|
||||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
@@ -541,6 +570,8 @@ class RenderAdapter:
|
||||
rendered_clip_ids: list[str] | None = None,
|
||||
failed_clip_ids: list[str] | None = None,
|
||||
voiceover_audio_path: str | None = None,
|
||||
asset_storage_map: dict[str, str] | None = None,
|
||||
task_config_override: dict | None = None, # Bug A: task 级 config 覆盖,防并发竞态
|
||||
) -> RenderAdapterResult:
|
||||
"""执行统一渲染核心流程(BGM + ASR + 渲染 + 缩略图 + 上传)。
|
||||
|
||||
@@ -554,8 +585,9 @@ class RenderAdapter:
|
||||
Returns:
|
||||
RenderAdapterResult
|
||||
"""
|
||||
# 1. 准备 BGM
|
||||
bgm_path = self._prepare_bgm(plan, work_dir, plan_id)
|
||||
# 1. 准备 BGM(Bug A: 传 task 级 bgm override)
|
||||
_bgm_override = (task_config_override or {}).get("bgm") if isinstance(task_config_override, dict) else None
|
||||
bgm_path = self._prepare_bgm(plan, work_dir, plan_id, bgm_override=_bgm_override)
|
||||
|
||||
self._report_progress(progress_cb, 40.0, "执行视频渲染")
|
||||
|
||||
@@ -563,8 +595,10 @@ class RenderAdapter:
|
||||
plan_config = plan.config or {}
|
||||
asr_service = self._get_asr_service()
|
||||
|
||||
# 3. 读取输出分辨率
|
||||
export_config = plan_config.get("export", {}) or {}
|
||||
# 3. 读取输出分辨率(Bug A: task override 优先)
|
||||
export_config = dict(plan_config.get("export", {}) or {})
|
||||
if isinstance(task_config_override, dict) and isinstance(task_config_override.get("export"), dict):
|
||||
export_config.update(task_config_override["export"])
|
||||
if not isinstance(export_config, dict):
|
||||
export_config = {}
|
||||
output_width, output_height = _parse_resolution(export_config.get("resolution"))
|
||||
@@ -589,9 +623,22 @@ class RenderAdapter:
|
||||
asr_service=asr_service,
|
||||
voiceover_audio_path=voiceover_audio_path,
|
||||
clip_has_text=clip_has_text,
|
||||
override_config=task_config_override,
|
||||
)
|
||||
# 注入每个视频段对应素材的 storage_key,供全 GPU 直连管线直接签名下载
|
||||
_storage_map = asset_storage_map or {}
|
||||
for c in clips:
|
||||
sk = _storage_map.get(getattr(c, "asset_id", ""))
|
||||
if sk:
|
||||
# EditPlanClip 使用 __slots__,不能 setattr,改存 config 字典
|
||||
if not isinstance(c.config, dict):
|
||||
c.config = dict(c.config) if c.config else {}
|
||||
c.config["_storage_key"] = sk
|
||||
result = render_svc.render()
|
||||
|
||||
# 4.4 透传 GPU 直连路径的 edge_crop 状态(供外层跳过 CPU 二次裁剪)
|
||||
edge_crop_applied_flag = bool(getattr(result, "edge_crop_applied", False))
|
||||
|
||||
# 4.5 渲染后校验输出完整性
|
||||
validation = validate_video_output(result.output_path)
|
||||
if not validation.valid:
|
||||
@@ -637,8 +684,22 @@ class RenderAdapter:
|
||||
# 已渲染视频在统一渲染阶段已通过 ASS 字幕把标题烧录进画面,
|
||||
# 抽帧天然带标题,因此这里传空字符串,避免 Pillow 二次叠加导致重影。
|
||||
# Pillow 叠加仅用于 API 从源素材抽帧(源素材本身无标题)的兜底场景。
|
||||
# 构造clip分段边界 [(start, duration), ...] 供封面抽帧智能取各段中点
|
||||
try:
|
||||
_clip_boundaries = [
|
||||
(float(getattr(c, "start_time", 0.0) or 0.0), float(getattr(c, "duration", 0.0) or 0.0))
|
||||
for c in clips
|
||||
if float(getattr(c, "duration", 0.0) or 0.0) > 0
|
||||
]
|
||||
except Exception:
|
||||
_clip_boundaries = None
|
||||
cover_candidates = extract_and_upload_cover_frames(
|
||||
str(result.output_path), plan_id, task_id=job_id, num_frames=5, title_text=""
|
||||
str(result.output_path),
|
||||
plan_id,
|
||||
task_id=job_id,
|
||||
num_frames=5,
|
||||
title_text="",
|
||||
clip_boundaries=_clip_boundaries,
|
||||
)
|
||||
if cover_candidates:
|
||||
logger.info(
|
||||
@@ -685,6 +746,7 @@ class RenderAdapter:
|
||||
rendered_clip_ids=final_rendered_ids,
|
||||
failed_clip_ids=final_failed_ids,
|
||||
cover_candidates=cover_candidates,
|
||||
edge_crop_applied=edge_crop_applied_flag,
|
||||
)
|
||||
|
||||
def render_from_memory(
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
"""视频封面抽帧工具 — 从视频中抽取帧作为封面,支持标题文字叠加。
|
||||
|
||||
统一封面管道:
|
||||
封面管道(P2 优化后):
|
||||
- 黑屏检测:ffmpeg blackdetect 扫描黑屏区间,抽帧点自动避开黑屏
|
||||
- 单次 ffmpeg select 抽多帧:一次 ffmpeg 进程用 select 滤镜输出 5 帧,避免 5 次起停进程
|
||||
- 并发上传:5 帧用 ThreadPoolExecutor 并行上传 OSS,目标封面阶段 <1.5s
|
||||
- 质量评分:cv2 清晰度/亮度/色彩三维评分选最佳帧
|
||||
- 可选 MediaKit 路径:配置 MEDIAKIT_COVER_ENABLED=true 时启用火山 MediaKit SceneChange 抽帧
|
||||
- 从已渲染视频抽帧:标题已通过 ASS 字幕烧进视频,帧天然带标题,无需再叠加。
|
||||
- 从源素材抽帧(API E2 兜底):源素材无标题,通过 Pillow 在帧上绘制标题文字。
|
||||
"""
|
||||
@@ -8,14 +13,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import tempfile
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 标题叠加(Pillow)──────────────────────────────────────────────────────
|
||||
# 实现统一放在 packages/shared/title_overlay.py,API 和 Worker 共用。
|
||||
|
||||
|
||||
def apply_title_overlay(
|
||||
image_path: str,
|
||||
@@ -27,11 +32,7 @@ def apply_title_overlay(
|
||||
margin_ratio: float = 0.06,
|
||||
stroke_width_ratio: float = 0.04,
|
||||
) -> str:
|
||||
"""在图片上绘制标题文字(指定颜色 + 黑色描边/阴影)。
|
||||
|
||||
委托给 packages.shared.title_overlay.apply_title_to_image,
|
||||
保持 Worker 内调用方式不变。title_text 为空时直接返回原路径。
|
||||
"""
|
||||
"""在图片上绘制标题文字(指定颜色 + 黑色描边/阴影)。"""
|
||||
from packages.shared.title_overlay import apply_title_to_image
|
||||
|
||||
if not title_text or not title_text.strip():
|
||||
@@ -56,26 +57,19 @@ def extract_first_frame(
|
||||
height: int = -1,
|
||||
timeout: int = 30,
|
||||
seek_ratio: float = 0.15,
|
||||
seek_seconds: float | None = None,
|
||||
min_seek_seconds: float = 1.0,
|
||||
) -> str:
|
||||
"""抽取视频封面帧(默认取视频时长 15% 处的帧,避开片头纯色画面)。
|
||||
|
||||
因为视频渲染时标题已通过 ASS 字幕烧录,抽取的帧天然带标题。
|
||||
"""抽取视频封面帧(ffmpeg -ss 单帧 seek,<100ms/帧)。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
output_path: 输出图片路径,不传则用临时文件
|
||||
width: 输出宽度(默认 -1,保持原始分辨率)
|
||||
height: 输出高度(默认 -1,保持原始分辨率)
|
||||
timeout: 超时时间(秒)
|
||||
seek_ratio: 抽帧位置占视频时长的比例(默认 0.15,即 15% 处)
|
||||
min_seek_seconds: 最小抽帧时间(秒),避免极短视频 seek 到 0
|
||||
|
||||
Returns:
|
||||
生成的封面帧文件路径
|
||||
|
||||
Raises:
|
||||
RuntimeError: ffmpeg 执行失败或输出文件为空
|
||||
width/height: 输出宽高(默认保持原始分辨率)
|
||||
timeout: 超时(秒)
|
||||
seek_ratio: 抽帧位置占视频时长的比例
|
||||
seek_seconds: 指定具体抽帧时间点(秒),优先于 seek_ratio
|
||||
min_seek_seconds: 最小抽帧时间
|
||||
"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
|
||||
|
||||
@@ -87,31 +81,25 @@ def extract_first_frame(
|
||||
_is_temp_output = True
|
||||
|
||||
try:
|
||||
# 计算抽帧时间点:取视频时长 * seek_ratio,最少 min_seek_seconds 秒
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
seek_time = max(min_seek_seconds, duration * seek_ratio)
|
||||
except Exception:
|
||||
# probe 失败时 fallback 到第1秒
|
||||
seek_time = min_seek_seconds
|
||||
if seek_seconds is not None:
|
||||
seek_time = max(0.0, float(seek_seconds))
|
||||
else:
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
seek_time = max(min_seek_seconds, duration * seek_ratio)
|
||||
except Exception:
|
||||
seek_time = min_seek_seconds
|
||||
|
||||
# 格式化为 HH:MM:SS.xx
|
||||
seek_str = _format_seek_time(seek_time)
|
||||
|
||||
# 构建 scale filter:如果指定了宽高则缩放,否则保持原始分辨率。
|
||||
# NOTE: scale_filter 在此处通过 if/else 分支赋值,之后不再被覆盖,
|
||||
# 后续 cmd / cmd2 均复用同一变量,逻辑无变化。
|
||||
if width > 0 or height > 0:
|
||||
w_str = str(width) if width > 0 else "-1"
|
||||
h_str = str(height) if height > 0 else "-1"
|
||||
scale_filter = f"scale={w_str}:{h_str}:force_original_aspect_ratio=decrease,format=yuvj420p"
|
||||
else:
|
||||
# 保持原始分辨率,只确保格式兼容
|
||||
scale_filter = "format=yuvj420p"
|
||||
|
||||
# -ss 放在 -i 前面(input seeking,更快)
|
||||
# -vframes 1 只取一帧
|
||||
# -q:v 2 jpeg 高质量
|
||||
# -ss 放在 -i 前面(input seeking,极快),-vframes 1 只取一帧
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
@@ -154,7 +142,6 @@ def extract_first_frame(
|
||||
|
||||
return output_path
|
||||
except Exception:
|
||||
# 失败时清理自己创建的临时文件
|
||||
if _is_temp_output and output_path:
|
||||
try:
|
||||
Path(output_path).unlink(missing_ok=True)
|
||||
@@ -164,7 +151,6 @@ def extract_first_frame(
|
||||
|
||||
|
||||
def _format_seek_time(seconds: float) -> str:
|
||||
"""将秒数格式化为 HH:MM:SS.xx 格式。"""
|
||||
h = int(seconds // 3600)
|
||||
m = int((seconds % 3600) // 60)
|
||||
s = seconds % 60
|
||||
@@ -177,19 +163,7 @@ def generate_and_upload_thumbnail(
|
||||
*,
|
||||
seek_ratio: float = 0.15,
|
||||
) -> str:
|
||||
"""从视频中提取一帧缩略图并上传到 OSS。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
storage_key: OSS 存储 key
|
||||
seek_ratio: 抽帧位置比例(默认 0.15)
|
||||
|
||||
Returns:
|
||||
上传后的 URL 字符串
|
||||
|
||||
Raises:
|
||||
RuntimeError: 抽帧或上传失败
|
||||
"""
|
||||
"""从视频中提取一帧缩略图并上传到 OSS。"""
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
@@ -204,24 +178,303 @@ def generate_and_upload_thumbnail(
|
||||
Path(tmp.name).unlink(missing_ok=True)
|
||||
|
||||
|
||||
def _detect_black_intervals(
|
||||
video_path: str,
|
||||
duration: float,
|
||||
*,
|
||||
black_min_duration: float = 0.3,
|
||||
picture_black_ratio_th: float = 0.98,
|
||||
pixel_black_th: float = 0.10,
|
||||
timeout: int = 30,
|
||||
) -> list[tuple[float, float]]:
|
||||
"""用 ffmpeg blackdetect 扫描黑屏区间,返回 [(start, end), ...]。"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
|
||||
|
||||
if duration <= 0:
|
||||
return []
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-nostdin",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vf",
|
||||
(f"blackdetect=d={black_min_duration:.2f}:pic_th={picture_black_ratio_th:.2f}:pix_th={pixel_black_th:.2f}"),
|
||||
"-an",
|
||||
"-f",
|
||||
"null",
|
||||
"-",
|
||||
]
|
||||
try:
|
||||
_, stderr = run_ffmpeg(cmd, capture_output=True, timeout=timeout)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] blackdetect 失败,忽略黑屏规避: %s", e)
|
||||
return []
|
||||
|
||||
intervals: list[tuple[float, float]] = []
|
||||
pattern = re.compile(
|
||||
r"black_start:(\d+(?:\.\d+)?)\s+black_end:(\d+(?:\.\d+)?)\s+black_duration:(\d+(?:\.\d+)?)",
|
||||
)
|
||||
for m in pattern.finditer(stderr or ""):
|
||||
try:
|
||||
bs = float(m.group(1))
|
||||
be = float(m.group(2))
|
||||
intervals.append((bs, be))
|
||||
except ValueError:
|
||||
continue
|
||||
intervals.sort()
|
||||
if intervals:
|
||||
logger.info("[thumbnail] blackdetect 发现 %d 段黑屏: %s", len(intervals), intervals[:5])
|
||||
return intervals
|
||||
|
||||
|
||||
def _adjust_seek_points_avoid_black(
|
||||
seek_points: list[float],
|
||||
black_intervals: list[tuple[float, float]],
|
||||
duration: float,
|
||||
*,
|
||||
tolerance: float = 0.25,
|
||||
) -> list[float]:
|
||||
"""把落在黑屏区间的 seek 点偏移到最近的非黑屏位置。
|
||||
|
||||
策略:
|
||||
- 若点在黑屏内,先尝试向前偏移到黑屏起点 - tolerance,再尝试向后偏移到黑屏终点 + tolerance;
|
||||
- 若整个视频全黑(偏移后 <0 或 >duration),保留原点但日志标记警告;
|
||||
- 偏移后若点与已有点重合(误差 <0.3s),做微调去重。
|
||||
"""
|
||||
if not black_intervals or not seek_points:
|
||||
return list(seek_points)
|
||||
|
||||
def in_black(t: float) -> tuple[float, float] | None:
|
||||
for bs, be in black_intervals:
|
||||
if bs <= t <= be:
|
||||
return (bs, be)
|
||||
return None
|
||||
|
||||
adjusted: list[float] = []
|
||||
for t in seek_points:
|
||||
seg = in_black(t)
|
||||
if seg is None:
|
||||
adjusted.append(max(0.0, min(duration, t)))
|
||||
continue
|
||||
bs, be = seg
|
||||
# 先尝试向前
|
||||
forward_t = bs - tolerance
|
||||
if forward_t >= 0.0 and in_black(forward_t) is None:
|
||||
adjusted.append(forward_t)
|
||||
continue
|
||||
# 再尝试向后
|
||||
backward_t = be + tolerance
|
||||
if backward_t <= duration and in_black(backward_t) is None:
|
||||
adjusted.append(backward_t)
|
||||
continue
|
||||
# 整段 clip 全黑?保留中点但标记
|
||||
logger.warning(
|
||||
"[thumbnail] seek 点 %.2fs 落在黑屏区间 [%.2f,%.2f] 且无法偏移,保留原位置(可能是全黑片段)",
|
||||
t,
|
||||
bs,
|
||||
be,
|
||||
)
|
||||
adjusted.append(max(0.0, min(duration, t)))
|
||||
|
||||
# 去重:相邻点若 <0.3s 则拉开
|
||||
adjusted.sort()
|
||||
deduped: list[float] = []
|
||||
for t in adjusted:
|
||||
if not deduped or abs(t - deduped[-1]) >= 0.3:
|
||||
deduped.append(t)
|
||||
else:
|
||||
# 往后挪 0.5s
|
||||
nt = t + 0.5
|
||||
if nt <= duration and in_black(nt) is None:
|
||||
deduped.append(nt)
|
||||
else:
|
||||
deduped.append(t)
|
||||
return [round(max(0.0, min(duration, t)), 3) for t in deduped[: len(seek_points)]]
|
||||
|
||||
|
||||
def _extract_frames_single_pass(
|
||||
video_path: str,
|
||||
seek_points: list[float],
|
||||
out_dir: str,
|
||||
*,
|
||||
prefix: str = "frame",
|
||||
width: int = -1,
|
||||
height: int = -1,
|
||||
q: int = 2,
|
||||
timeout: int = 30,
|
||||
) -> list[tuple[float, str]]:
|
||||
"""单次 ffmpeg 用 select 滤镜抽出 seek_points 对应的多帧。
|
||||
|
||||
ffmpeg -i input -vf "select='between(t,t1-0.03,t1+0.03)+between(t,t2-0.03,t2+0.03)+...',scale=...,format=yuvj420p"
|
||||
-vsync vfr -q:v 2 out_dir/prefix_%02d.jpg
|
||||
|
||||
返回 [(seek_t, output_path), ...],按输出帧序号升序。若输出帧数 < seek_points 数量,
|
||||
不足部分用 extract_first_frame 兜底(保证返回数量 == len(seek_points))。
|
||||
"""
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
|
||||
|
||||
out_dir_p = Path(out_dir)
|
||||
out_dir_p.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 构造 select 表达式:每个 seek 点用 ±30ms 窗口命中
|
||||
# between(t, a, b) 返回 1 表示 t 在 [a,b] 内;多个 between 相加即为"任一命中"
|
||||
select_terms = []
|
||||
for t in seek_points:
|
||||
a = max(0.0, t - 0.03)
|
||||
b = t + 0.04
|
||||
select_terms.append(f"between(t,{a:.3f},{b:.3f})")
|
||||
select_expr = "+".join(select_terms)
|
||||
|
||||
if width > 0 or height > 0:
|
||||
w_str = str(width) if width > 0 else "-1"
|
||||
h_str = str(height) if height > 0 else "-1"
|
||||
scale_filter = f"scale={w_str}:{h_str}:force_original_aspect_ratio=decrease"
|
||||
vf = f"select='{select_expr}',{scale_filter},format=yuvj420p"
|
||||
else:
|
||||
vf = f"select='{select_expr}',format=yuvj420p"
|
||||
|
||||
out_pattern = str(out_dir_p / f"{prefix}_%02d.jpg")
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vf",
|
||||
vf,
|
||||
"-vsync",
|
||||
"vfr",
|
||||
"-q:v",
|
||||
str(q),
|
||||
out_pattern,
|
||||
]
|
||||
|
||||
results: list[tuple[float, str]] = []
|
||||
single_pass_ok = False
|
||||
try:
|
||||
run_ffmpeg(cmd, capture_output=True, timeout=timeout)
|
||||
# 读取输出文件
|
||||
for i in range(1, len(seek_points) + 1):
|
||||
fp = out_dir_p / f"{prefix}_{i:02d}.jpg"
|
||||
if fp.exists() and fp.stat().st_size > 0:
|
||||
results.append((seek_points[i - 1] if i - 1 < len(seek_points) else 0.0, str(fp)))
|
||||
if len(results) >= len(seek_points):
|
||||
single_pass_ok = True
|
||||
else:
|
||||
logger.warning(
|
||||
"[thumbnail] 单次 ffmpeg 抽帧仅命中 %d/%d 帧,不足部分用单帧 seek 兜底",
|
||||
len(results),
|
||||
len(seek_points),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 单次 ffmpeg select 抽帧失败,回退到单帧 seek: %s", e)
|
||||
|
||||
# 兜底:对缺失/失败的帧用 extract_first_frame 补抽
|
||||
if not single_pass_ok:
|
||||
# 清理不完整结果
|
||||
for _, fp in results:
|
||||
try:
|
||||
Path(fp).unlink(missing_ok=True)
|
||||
except Exception:
|
||||
pass
|
||||
results = []
|
||||
for i, st in enumerate(seek_points):
|
||||
fp = out_dir_p / f"{prefix}_fallback_{i:02d}.jpg"
|
||||
try:
|
||||
extract_first_frame(
|
||||
video_path,
|
||||
output_path=str(fp),
|
||||
seek_seconds=st,
|
||||
min_seek_seconds=0.5,
|
||||
timeout=timeout,
|
||||
)
|
||||
if fp.exists() and fp.stat().st_size > 0:
|
||||
results.append((st, str(fp)))
|
||||
else:
|
||||
logger.warning("[thumbnail] 兜底单帧抽帧也失败 idx=%d t=%.2f", i, st)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 兜底单帧抽帧异常 idx=%d t=%.2f: %s", i, st, e)
|
||||
|
||||
return results[: len(seek_points)]
|
||||
|
||||
|
||||
def _compute_clip_boundary_seek_points(
|
||||
duration: float,
|
||||
clip_boundaries: Optional[list[tuple[float, float]]] = None,
|
||||
num_frames: int = 5,
|
||||
head_skip_ratio: float = 0.08,
|
||||
tail_skip_ratio: float = 0.08,
|
||||
) -> list[float]:
|
||||
"""基于clip分段边界计算抽帧时间点(取每段中间帧,效果比均匀抽更好)。
|
||||
|
||||
策略:
|
||||
- 如果传入 clip_boundaries(每个元素是 (clip_start_in_timeline, clip_duration)),
|
||||
取每个片段的中点作为抽帧候选点
|
||||
- 候选点不足 num_frames 时,均匀补充
|
||||
- 跳过片头 head_skip_ratio(8%,避免片头黑屏/开场标题)和片尾 tail_skip_ratio(8%)
|
||||
- 返回按时间排序的 num_frames 个抽帧点(秒)
|
||||
"""
|
||||
if duration <= 0:
|
||||
# 无法probe,均匀分布兜底
|
||||
return [max(1.0, duration * (0.1 + 0.8 * i / max(num_frames - 1, 1))) for i in range(num_frames)]
|
||||
|
||||
head_skip = duration * head_skip_ratio
|
||||
tail_skip = duration * tail_skip_ratio
|
||||
valid_start = head_skip
|
||||
valid_end = max(valid_start + 1.0, duration - tail_skip)
|
||||
|
||||
candidates: list[float] = []
|
||||
|
||||
if clip_boundaries:
|
||||
# 累加timeline start,取每clip中点
|
||||
cur = 0.0
|
||||
for _clip_start, clip_dur in clip_boundaries:
|
||||
if clip_dur <= 0:
|
||||
continue
|
||||
mid = cur + clip_dur / 2.0
|
||||
if valid_start <= mid <= valid_end:
|
||||
candidates.append(mid)
|
||||
cur += clip_dur
|
||||
# 去重+排序
|
||||
candidates = sorted(set(round(c, 3) for c in candidates))
|
||||
|
||||
# 如果候选点不足,均匀补充
|
||||
if len(candidates) < num_frames:
|
||||
needed = num_frames - len(candidates)
|
||||
existing = set(round(c, 1) for c in candidates)
|
||||
for i in range(needed * 3):
|
||||
ratio = 0.1 + 0.8 * (i + 0.5) / (needed * 3)
|
||||
t = valid_start + (valid_end - valid_start) * ratio
|
||||
if round(t, 1) not in existing:
|
||||
candidates.append(t)
|
||||
existing.add(round(t, 1))
|
||||
if len(candidates) >= num_frames:
|
||||
break
|
||||
|
||||
# 如果还不够,强制均匀
|
||||
while len(candidates) < num_frames:
|
||||
idx = len(candidates)
|
||||
ratio = 0.1 + 0.8 * idx / max(num_frames - 1, 1)
|
||||
candidates.append(valid_start + (valid_end - valid_start) * ratio)
|
||||
|
||||
candidates.sort()
|
||||
|
||||
# 如果超过num_frames,均匀选取
|
||||
if len(candidates) > num_frames:
|
||||
step = len(candidates) / num_frames
|
||||
candidates = [candidates[int(i * step)] for i in range(num_frames)]
|
||||
|
||||
return [round(t, 3) for t in candidates[:num_frames]]
|
||||
|
||||
|
||||
def _extract_frames_via_mediakit(
|
||||
video_path: str,
|
||||
plan_id: str,
|
||||
num_frames: int,
|
||||
) -> list[dict] | None:
|
||||
"""使用 MediaKit 智能抽帧 API 提取封面帧。
|
||||
|
||||
Args:
|
||||
video_path: 本地视频文件路径
|
||||
plan_id: 编辑计划 ID
|
||||
num_frames: 需要的帧数
|
||||
|
||||
Returns:
|
||||
帧列表 [{"image_url": str, "timestamp": float}, ...],失败返回 None
|
||||
"""
|
||||
"""使用 MediaKit 智能抽帧 API 提取封面帧(fallback 路径,默认不启用)。"""
|
||||
import uuid
|
||||
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
from video_processing.oss_helpers import delete_from_oss, get_signed_download_url, upload_to_oss
|
||||
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
|
||||
@@ -230,45 +483,38 @@ def _extract_frames_via_mediakit(
|
||||
logger.info("[thumbnail] MediaKit 未配置,跳过智能抽帧")
|
||||
return None
|
||||
|
||||
# 1. 上传视频到 OSS 获取 URL
|
||||
video_storage_key: str = ""
|
||||
try:
|
||||
video_storage_key = f"temp/{plan_id}/{uuid.uuid4().hex[:8]}_{Path(video_path).name}"
|
||||
video_url = upload_to_oss(video_path, video_storage_key)
|
||||
if not video_url:
|
||||
public_url = upload_to_oss(video_path, video_storage_key)
|
||||
if not public_url:
|
||||
logger.warning("[thumbnail] 视频上传 OSS 失败,无法使用 MediaKit")
|
||||
return None
|
||||
logger.info("[thumbnail] 视频已上传 OSS: %s", video_url[:80])
|
||||
video_url = get_signed_download_url(video_storage_key, expires_seconds=3600) or public_url
|
||||
logger.info("[thumbnail] 视频已上传 OSS 并生成签名 URL: key=%s", video_storage_key[:80])
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 视频上传 OSS 异常: %s,降级到 ffmpeg", e)
|
||||
logger.warning("[thumbnail] 视频上传 OSS 异常: %s,降级到本地 ffmpeg", e)
|
||||
return None
|
||||
|
||||
# 2. 调用 MediaKit 智能抽帧
|
||||
try:
|
||||
frames = client.extract_frames(
|
||||
video_url=video_url,
|
||||
strategy="SceneChange",
|
||||
max_frames=num_frames * 2, # 多取一些帧供选择
|
||||
max_frames=num_frames * 2,
|
||||
)
|
||||
if not frames:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧返回空,降级到 ffmpeg")
|
||||
logger.warning("[thumbnail] MediaKit 抽帧返回空")
|
||||
return None
|
||||
|
||||
# 选取最均匀的 num_frames 个帧
|
||||
if len(frames) > num_frames:
|
||||
step = len(frames) // num_frames
|
||||
frames = [frames[i * step] for i in range(num_frames)]
|
||||
|
||||
logger.info("[thumbnail] MediaKit 抽帧成功: %d 帧", len(frames))
|
||||
return frames
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧异常: %s,降级到 ffmpeg", e)
|
||||
logger.warning("[thumbnail] MediaKit 抽帧异常: %s", e)
|
||||
return None
|
||||
finally:
|
||||
# 清理临时视频文件
|
||||
try:
|
||||
from video_processing.oss_helpers import delete_from_oss
|
||||
|
||||
delete_from_oss(video_storage_key)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -279,161 +525,235 @@ def extract_and_upload_cover_frames(
|
||||
plan_id: str,
|
||||
*,
|
||||
task_id: str = "",
|
||||
num_frames: int = 5, # 抽 5 帧候选,通过质量评分选出最佳帧
|
||||
num_frames: int = 5,
|
||||
title_text: str = "",
|
||||
title_color: str = "#ffffff",
|
||||
title_position: str = "bottom",
|
||||
title_font_size: int | None = None,
|
||||
clip_boundaries: Optional[list[tuple[float, float]]] = None,
|
||||
) -> list[dict]:
|
||||
"""从视频中抽取多帧作为封面候选,通过质量评分选出最佳帧,上传到 OSS。
|
||||
|
||||
流程:
|
||||
1. 优先使用 MediaKit 智能抽帧(多抽一些供选择)
|
||||
2. MediaKit 不足时降级到 ffmpeg 均匀抽帧
|
||||
3. 对所有候选帧进行质量评分(清晰度/亮度/色彩丰富度)
|
||||
4. 按分数从高到低排序返回
|
||||
P2 优化:
|
||||
- 先用 ffmpeg blackdetect 扫描黑屏区间,seek 点自动避开黑屏
|
||||
- 单次 ffmpeg select 抽 num_frames 帧(避免 5 次起停 ffmpeg 进程)
|
||||
- 多帧 OSS 上传用 ThreadPoolExecutor 并发,目标封面阶段 <1.5s
|
||||
- cv2 清晰度/亮度/色彩三维评分选最佳帧
|
||||
|
||||
Fallback(MEDIAKIT_COVER_ENABLED=true):火山 MediaKit SceneChange 抽帧(~60-90s)。
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
plan_id: 编辑计划 ID(用于生成 storage key)
|
||||
task_id: 任务 ID(用于生成独立的 storage key,避免标题变更时封面冲突)
|
||||
num_frames: 抽取候选帧数(默认 5,通过质量评分选出最佳帧)
|
||||
title_text: 标题文字;非空时用 Pillow 叠加到每帧。
|
||||
从已渲染视频抽帧时通常传空(标题已烧录);从源素材抽帧时传标题。
|
||||
title_color: 标题字体颜色(#RRGGBB)
|
||||
title_position: 标题位置 top/center/bottom
|
||||
title_font_size: 标题字号,None 时自动计算
|
||||
|
||||
Returns:
|
||||
封面候选列表(按质量分数降序),每项包含 {"url": str, "position": float, "score": float}
|
||||
clip_boundaries: 片段边界列表 [(clip_start, clip_duration), ...],用于智能取点
|
||||
"""
|
||||
import time
|
||||
|
||||
import httpx
|
||||
from video_processing.ffmpeg_utils import probe_duration
|
||||
from video_processing.oss_helpers import upload_to_oss
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
t0 = time.monotonic()
|
||||
|
||||
try:
|
||||
duration = probe_duration(video_path)
|
||||
except Exception:
|
||||
duration = 0.0
|
||||
|
||||
candidates: list[dict] = []
|
||||
_temp_paths: list[str] = [] # 收集所有临时文件路径,最后统一清理
|
||||
_temp_paths: list[str] = []
|
||||
|
||||
try:
|
||||
# ── 阶段 1:抽帧 ──────────────────────────────────────────────
|
||||
# 优先尝试 MediaKit 智能抽帧
|
||||
mediakit_frames = _extract_frames_via_mediakit(video_path, plan_id, num_frames)
|
||||
if mediakit_frames:
|
||||
for i, frame in enumerate(mediakit_frames):
|
||||
frame_url = frame.get("image_url")
|
||||
if not frame_url:
|
||||
continue
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
# 下载 MediaKit 返回的帧图
|
||||
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
with open(tmp.name, "wb") as f:
|
||||
f.write(resp.content)
|
||||
settings = get_shared_settings()
|
||||
use_mediakit = getattr(settings, "mediakit_cover_enabled", False)
|
||||
|
||||
# 叠加标题文字(如需要)
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
tmp.name,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
if use_mediakit:
|
||||
logger.info("[thumbnail] MEDIAKIT_COVER_ENABLED=true,走 MediaKit 路径")
|
||||
mediakit_frames = _extract_frames_via_mediakit(video_path, plan_id, num_frames)
|
||||
if mediakit_frames:
|
||||
for i, frame in enumerate(mediakit_frames):
|
||||
frame_url = frame.get("image_url")
|
||||
if not frame_url:
|
||||
continue
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
resp = httpx.get(frame_url, timeout=30, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
with open(tmp.name, "wb") as f:
|
||||
f.write(resp.content)
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
tmp.name,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
storage_key = f"covers/{plan_id}/{task_id}/mediakit_frame_{i}.jpg"
|
||||
url = upload_to_oss(tmp.name, storage_key)
|
||||
if url:
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(frame.get("timestamp", 0.0), 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 帧 %d 处理失败: %s", i, e)
|
||||
if len(candidates) >= num_frames:
|
||||
logger.info("[thumbnail] MediaKit 抽帧完成: %d 帧", len(candidates))
|
||||
|
||||
# MediaKit 路径帧在 NamedTemporaryFile 中持久存在(finally 清理),在进入本地 ffmpeg 前评分
|
||||
if len(candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
candidates = score_frames(candidates)
|
||||
logger.info(
|
||||
"[thumbnail] MediaKit 封面帧评分完成: count=%d best_score=%.1f",
|
||||
len(candidates),
|
||||
candidates[0].get("score", 0.0) if candidates else 0.0,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("[thumbnail] MediaKit 封面帧质量评分失败,保持原始顺序", exc_info=True)
|
||||
|
||||
storage_key = f"covers/{plan_id}/{task_id}/mediakit_frame_{i}.jpg"
|
||||
url = upload_to_oss(tmp.name, storage_key)
|
||||
if url:
|
||||
seek_time = frame.get("timestamp", 0.0)
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(seek_time, 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] MediaKit 帧 %d 处理失败: %s", i, e)
|
||||
|
||||
if len(candidates) >= num_frames:
|
||||
logger.info("[thumbnail] MediaKit 智能抽帧完成: %d 帧", len(candidates))
|
||||
else:
|
||||
logger.warning("[thumbnail] MediaKit 抽帧不足 %d 帧,降级到 ffmpeg", num_frames)
|
||||
|
||||
# Fallback: ffmpeg 直接抽帧(仅当 MediaKit 不足时)
|
||||
# ── 默认路径:本地 ffmpeg 单次 select 抽帧 + 并发上传 ──────────────
|
||||
if len(candidates) < num_frames:
|
||||
logger.info("[thumbnail] 使用 ffmpeg 抽帧补充")
|
||||
# 均匀分布抽帧点:从 10% 到 90%
|
||||
for i in range(num_frames):
|
||||
ratio = 0.1 + 0.8 * i / max(num_frames - 1, 1)
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
|
||||
tmp.close()
|
||||
_temp_paths.append(tmp.name)
|
||||
try:
|
||||
frame_path = extract_first_frame(
|
||||
video_path,
|
||||
output_path=tmp.name,
|
||||
seek_ratio=ratio,
|
||||
min_seek_seconds=0.5,
|
||||
)
|
||||
# 从源素材抽帧时叠加标题文字;已渲染视频标题已烧录时传空字符串跳过
|
||||
if title_text and title_text.strip():
|
||||
apply_title_overlay(
|
||||
frame_path,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
storage_key = f"covers/{plan_id}/{task_id}/frame_{i}.jpg"
|
||||
url = upload_to_oss(frame_path, storage_key)
|
||||
if url:
|
||||
seek_time = max(0.5, duration * ratio) if duration > 0 else 0.0
|
||||
candidates.append(
|
||||
{
|
||||
"url": url,
|
||||
"position": round(seek_time, 2),
|
||||
"image_path": tmp.name,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 封面候选帧 %d 提取失败: %s", i, e)
|
||||
|
||||
# ── 阶段 2:质量评分 ────────────────────────────────────────────
|
||||
if len(candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
candidates = score_frames(candidates)
|
||||
if candidates:
|
||||
logger.info("[thumbnail] MediaKit 不足 %d 帧,本地 ffmpeg 补充", num_frames)
|
||||
else:
|
||||
logger.info(
|
||||
"[thumbnail] 封面帧质量评分完成: plan_id=%s count=%d best_score=%.1f",
|
||||
"[thumbnail] 使用本地 ffmpeg 抽帧(num=%d, duration=%.1fs)",
|
||||
num_frames,
|
||||
duration,
|
||||
)
|
||||
|
||||
# 1) 计算 seek 点
|
||||
seek_points = _compute_clip_boundary_seek_points(duration, clip_boundaries, num_frames)
|
||||
|
||||
# 2) 黑屏检测 + 偏移 seek 点
|
||||
black_intervals = _detect_black_intervals(video_path, duration) if duration > 0 else []
|
||||
if black_intervals:
|
||||
seek_points = _adjust_seek_points_avoid_black(seek_points, black_intervals, duration)
|
||||
logger.info("[thumbnail] 黑屏规避后 seek 点: %s", seek_points)
|
||||
|
||||
# 3) 单次 ffmpeg select 抽出所有帧(带失败兜底到单帧 seek)
|
||||
with tempfile.TemporaryDirectory(prefix="thumb_") as frame_dir:
|
||||
t1 = time.monotonic()
|
||||
frame_results = _extract_frames_single_pass(
|
||||
video_path,
|
||||
seek_points,
|
||||
frame_dir,
|
||||
prefix="frame",
|
||||
)
|
||||
logger.info("[thumbnail] 抽帧耗时: %.2fs (%d 帧)", time.monotonic() - t1, len(frame_results))
|
||||
|
||||
# 4) 标题叠加(本地,CPU 很快)
|
||||
for _st, fp in frame_results:
|
||||
if title_text and title_text.strip():
|
||||
try:
|
||||
apply_title_overlay(
|
||||
fp,
|
||||
title_text,
|
||||
color=title_color,
|
||||
position=title_position,
|
||||
font_size=title_font_size,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 标题叠加失败 %s: %s", fp, e)
|
||||
|
||||
# 5) 质量评分(必须在 TemporaryDirectory 内,帧文件还在磁盘上)
|
||||
t_score = time.monotonic()
|
||||
local_candidates: list[dict] = [{"position": st, "image_path": fp} for (st, fp) in frame_results]
|
||||
scored: list[dict] = local_candidates
|
||||
if len(local_candidates) > 1:
|
||||
try:
|
||||
from packages.shared.cover_frame_scorer import score_frames
|
||||
|
||||
scored = score_frames(local_candidates)
|
||||
logger.info(
|
||||
"[thumbnail] 封面评分耗时: %.2fs (best_score=%.1f, count=%d)",
|
||||
time.monotonic() - t_score,
|
||||
scored[0].get("score", 0.0) if scored else 0.0,
|
||||
len(scored),
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"[thumbnail] 封面帧质量评分失败,保持 seek 点原始顺序",
|
||||
exc_info=True,
|
||||
)
|
||||
scored = local_candidates
|
||||
|
||||
# 6) 按评分顺序并发上传 OSS(best 帧先上传;best 已是 scored[0])
|
||||
t2 = time.monotonic()
|
||||
|
||||
def _upload_one(rank: int, st: float, fp: str, score: float) -> dict | None:
|
||||
try:
|
||||
storage_key = f"covers/{plan_id}/{task_id}/frame_{rank}.jpg"
|
||||
url = upload_to_oss(fp, storage_key)
|
||||
if url:
|
||||
return {
|
||||
"url": url,
|
||||
"position": st,
|
||||
"image_path": fp,
|
||||
"score": score,
|
||||
"is_best": rank == 0,
|
||||
}
|
||||
logger.warning("[thumbnail] 上传失败 rank=%d t=%.2f", rank, st)
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 上传异常 rank=%d t=%.2f: %s", rank, st, e)
|
||||
return None
|
||||
|
||||
upload_results: list[dict | None] = [None] * len(scored)
|
||||
max_workers = min(8, max(2, len(scored)))
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
||||
future_map = {
|
||||
pool.submit(
|
||||
_upload_one,
|
||||
i,
|
||||
float(c.get("position", 0.0)),
|
||||
str(c["image_path"]),
|
||||
float(c.get("score", 0.0)),
|
||||
): i
|
||||
for i, c in enumerate(scored)
|
||||
}
|
||||
for fut in as_completed(future_map):
|
||||
i = future_map[fut]
|
||||
try:
|
||||
upload_results[i] = fut.result()
|
||||
except Exception as e:
|
||||
logger.warning("[thumbnail] 上传 future 异常 rank=%d: %s", i, e)
|
||||
logger.info("[thumbnail] 并发上传耗时: %.2fs", time.monotonic() - t2)
|
||||
|
||||
for r in upload_results:
|
||||
if r is not None:
|
||||
# 本地帧在 TemporaryDirectory 内,with 退出自动删除,无需进 _temp_paths
|
||||
candidates.append(r)
|
||||
|
||||
# 如果本地 ffmpeg 路径产生了候选(已评分)但未经过 MediaKit 路径,candidates 已按评分顺序排好。
|
||||
# 混合场景下(MediaKit + 本地 ffmpeg 都产出),统一按 score 降序排列;缺失 score 的(理论上不应出现)排末尾。
|
||||
if len(candidates) > 1:
|
||||
candidates.sort(key=lambda c: c.get("score", -1.0), reverse=True)
|
||||
if candidates:
|
||||
candidates[0]["is_best"] = True
|
||||
elapsed = time.monotonic() - t0
|
||||
logger.info(
|
||||
"[thumbnail] 封面完成: plan_id=%s count=%d best=t%.2fs score=%.1f elapsed=%.2fs",
|
||||
plan_id,
|
||||
len(candidates),
|
||||
candidates[0].get("score", 0.0) if candidates else 0.0,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"[thumbnail] 封面帧质量评分失败,保持原始顺序: plan_id=%s",
|
||||
plan_id,
|
||||
exc_info=True,
|
||||
candidates[0].get("position", 0.0),
|
||||
candidates[0].get("score", 0.0),
|
||||
elapsed,
|
||||
)
|
||||
|
||||
# ── 阶段 3:清理临时文件 ────────────────────────────────────────
|
||||
# 移除 image_path(不再需要),但临时文件统一清理
|
||||
for c in candidates:
|
||||
c.pop("image_path", None)
|
||||
|
||||
return candidates
|
||||
|
||||
finally:
|
||||
# 统一清理所有临时文件
|
||||
for path in _temp_paths:
|
||||
try:
|
||||
Path(path).unlink(missing_ok=True)
|
||||
|
||||
@@ -112,6 +112,7 @@ class RenderResult:
|
||||
file_size: int
|
||||
width: int
|
||||
height: int
|
||||
edge_crop_applied: bool = False # True = GPU管线已做随机边缘裁剪
|
||||
|
||||
|
||||
# ── clip_type → layer role 映射 ──────────────────────────────────────────────
|
||||
@@ -157,6 +158,7 @@ class UnifiedRenderService:
|
||||
bgm_path: str | None = None, # BGM 本地文件路径
|
||||
voiceover_audio_path: str | None = None, # 配音素材库音频本地路径
|
||||
clip_has_text: list[bool] | None = None, # 源视频片段是否有文字(来自 atom_clip.ai_tags.has_text)
|
||||
override_config: dict | None = None, # Bug A: task 级 config 覆盖(title/bgm/export/subtitle),防并发竞态
|
||||
):
|
||||
self.plan = plan
|
||||
self.clips = clips
|
||||
@@ -169,6 +171,9 @@ class UnifiedRenderService:
|
||||
self.asr_service = asr_service
|
||||
self.bgm_path = bgm_path
|
||||
self.voiceover_audio_path = voiceover_audio_path
|
||||
# Bug A: task 级 config override(深拷贝),优先级高于 plan.config;
|
||||
# 避免同 plan 多任务并发渲染时 _sync_task_config_to_plan 写 plan.config["title"] 互相覆盖。
|
||||
self._override_config = dict(override_config) if isinstance(override_config, dict) else {}
|
||||
# #1970:片段级文字检测(顺序与非 audio 的源视频片段一致);None 表示无可靠检测,保守不翻转
|
||||
self._clip_has_text = clip_has_text
|
||||
self._transition_engine = TransitionEngine(default_duration=transition_duration)
|
||||
@@ -179,6 +184,28 @@ class UnifiedRenderService:
|
||||
self._micro_plan_cache: Any = None
|
||||
self._micro_plan_loaded = False
|
||||
|
||||
def _cfg_section(self, section: str) -> dict:
|
||||
"""读取单个配置段:override_config 优先于 plan.config(Bug A 防并发竞态)。"""
|
||||
base = dict((self.plan.config or {}).get(section, {}) or {})
|
||||
override = self._override_config.get(section)
|
||||
if isinstance(override, dict) and override:
|
||||
base.update(override) # 浅合并,保留 base 中未被覆盖字段
|
||||
return base
|
||||
|
||||
def _effective_config(self) -> dict:
|
||||
"""读取完整 config:override_config 顶层段覆盖 plan.config(Bug A 防并发竞态)。"""
|
||||
import copy
|
||||
|
||||
full = copy.deepcopy(self.plan.config or {})
|
||||
for k, v in self._override_config.items():
|
||||
if isinstance(v, dict):
|
||||
sec = dict(full.get(k, {}) or {})
|
||||
sec.update(v)
|
||||
full[k] = sec
|
||||
else:
|
||||
full[k] = v
|
||||
return full
|
||||
|
||||
# ── #1970 PR2 智能降重:片段级微变换 ───────────────────────────────────
|
||||
def _dedup_enabled(self) -> bool:
|
||||
"""读取 plan.config.dedup_enabled,缺省视为 True(向后兼容)。"""
|
||||
@@ -233,7 +260,7 @@ class UnifiedRenderService:
|
||||
return
|
||||
if abs(mt.brightness) > 1e-4 or abs(mt.contrast - 1.0) > 1e-4 or abs(mt.saturation - 1.0) > 1e-4:
|
||||
filters.append(
|
||||
f"eq=brightness={mt.brightness:+.4f}:" f"contrast={mt.contrast:.4f}:saturation={mt.saturation:.4f}"
|
||||
f"eq=brightness={mt.brightness:+.4f}:contrast={mt.contrast:.4f}:saturation={mt.saturation:.4f}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -341,6 +368,36 @@ class UnifiedRenderService:
|
||||
len(pip_sources),
|
||||
)
|
||||
|
||||
# 4.8 全 GPU 直连管线(P1):命中主流场景则跳过 mezzanine/边缘裁剪 CPU 重编码
|
||||
output_path = self.work_dir / f"rendered_{self.plan.id}.mp4"
|
||||
direct_result = self._try_gpu_direct(
|
||||
layers=layers,
|
||||
ass_path=ass_path,
|
||||
video_duration=video_duration_final,
|
||||
output_path=output_path,
|
||||
)
|
||||
if direct_result is not None and direct_result[0]:
|
||||
_direct_edge_crop = bool(direct_result[1])
|
||||
# 直连成功:直接探测并返回,跳过后续视频/音频 CPU 流程
|
||||
duration, file_size, width, height = self._probe_output(output_path)
|
||||
logger.info(
|
||||
"[unified-render] gpu-direct done: plan_id=%s total_ms=%d output_size=%d resolution=%dx%d",
|
||||
self.plan.id,
|
||||
int((time.time() - t_start) * 1000),
|
||||
file_size,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
direct_edge_cropped = _direct_edge_crop # GPU直连时若dedup=True已在GPU内做随机边缘裁剪
|
||||
return RenderResult(
|
||||
output_path=output_path,
|
||||
duration=duration,
|
||||
file_size=file_size,
|
||||
width=width,
|
||||
height=height,
|
||||
edge_crop_applied=direct_edge_cropped,
|
||||
)
|
||||
|
||||
# 5. 视频主渲染
|
||||
t_video_start = time.time()
|
||||
video_only_path = self.work_dir / f"rendered_{self.plan.id}_video.mp4"
|
||||
@@ -400,7 +457,7 @@ class UnifiedRenderService:
|
||||
has_audio = pass_through_has_audio
|
||||
# 直通模式下也支持 BGM 混音:提取音频 → 混 BGM → 合并回视频
|
||||
if self.bgm_path and pass_through_has_audio:
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
bgm_config = config.get("bgm", {}) or {}
|
||||
if bgm_config.get("enabled", False):
|
||||
ctx = RenderContext(work_dir=self.work_dir, plan_id=self.plan.id)
|
||||
@@ -440,7 +497,7 @@ class UnifiedRenderService:
|
||||
"[unified-render] pass-through BGM mix failed, skipping: plan_id=%s", self.plan.id
|
||||
)
|
||||
else:
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
bgm_config = config.get("bgm", {}) or {}
|
||||
if not isinstance(bgm_config, dict):
|
||||
bgm_config = {}
|
||||
@@ -665,7 +722,7 @@ class UnifiedRenderService:
|
||||
Returns:
|
||||
ASS 文件路径,没有字幕时返回 None
|
||||
"""
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
# #1901 统一读 "title",兼容老数据 "title_config"
|
||||
title_cfg = config.get("title", {}) or {}
|
||||
if not isinstance(title_cfg, dict) or not (title_cfg.get("text") or "").strip():
|
||||
@@ -850,7 +907,7 @@ class UnifiedRenderService:
|
||||
Returns:
|
||||
是否成功添加了配音音轨
|
||||
"""
|
||||
config = self.plan.config or {}
|
||||
config = self._effective_config()
|
||||
tts_cfg = config.get("tts", {}) or {}
|
||||
if not isinstance(tts_cfg, dict):
|
||||
tts_cfg = {}
|
||||
@@ -2180,6 +2237,261 @@ class UnifiedRenderService:
|
||||
|
||||
# ── GPU NVENC 加速 ────────────────────────────────────────────────────
|
||||
|
||||
# ── 全 GPU 直连渲染(P1)─────────────────────────────────────────────
|
||||
|
||||
def _can_use_gpu_direct(self, layers: list[RenderLayer]) -> bool:
|
||||
"""判断是否命中直连支持的场景:单一主视频轨、全硬切、无复杂合成。"""
|
||||
try:
|
||||
cfg = self.plan.config or {}
|
||||
# 特性开关(默认开启;可经 env/plan config 关闭灰度回退)
|
||||
if not bool(cfg.get("gpu_direct_enabled", True)):
|
||||
return False
|
||||
|
||||
video_layers = [_lyr for _lyr in layers if _lyr.role not in ("audio",)]
|
||||
# 只允许一个视频层,且角色为主层
|
||||
if len(video_layers) != 1:
|
||||
return False
|
||||
role = video_layers[0].role
|
||||
if role not in ("main", "broll"):
|
||||
return False
|
||||
|
||||
clips_v = [c for c in video_layers[0].clips if c.clip_type != "audio"]
|
||||
if not clips_v:
|
||||
return False
|
||||
# 全硬切(第一个 clip 的转场忽略)
|
||||
for c in clips_v[1:]:
|
||||
te = c.transition_effect
|
||||
if te not in (None, "", "cut"):
|
||||
return False
|
||||
# 无画中画 / 水印 / 贴纸 / 片头片尾 / 绿幕 / 倒放 / 调色
|
||||
if (cfg or {}).get("pip_config"):
|
||||
return False
|
||||
if (cfg or {}).get("intro_outro"):
|
||||
return False
|
||||
for c in clips_v:
|
||||
cc = c.config or {}
|
||||
if cc.get("watermark") or cc.get("stickers") or cc.get("chroma_key"):
|
||||
return False
|
||||
if ReverseConfig.from_dict(cc.get("reverse")).enabled:
|
||||
return False
|
||||
cg = ColorGradeConfig.from_dict(cc.get("color_grade"))
|
||||
if cg.enabled and cg.has_effect():
|
||||
return False
|
||||
if not (c.config or {}).get("_storage_key"):
|
||||
return False
|
||||
return True
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] eligibility check failed (fallback)", exc_info=True)
|
||||
return False
|
||||
|
||||
def _try_gpu_direct(
|
||||
self,
|
||||
*,
|
||||
layers: list[RenderLayer],
|
||||
ass_path: Path | None,
|
||||
video_duration: float,
|
||||
output_path: Path,
|
||||
) -> tuple[bool, bool] | tuple[None, bool]:
|
||||
"""尝试全 GPU 直连渲染。成功返回 (True, edge_crop_applied),不支持/失败返回 (None, False)。"""
|
||||
if not self._can_use_gpu_direct(layers):
|
||||
return (None, False)
|
||||
if not self._gpu_encode_available():
|
||||
return (None, False)
|
||||
|
||||
try:
|
||||
from video_processing import gpu_direct_pipeline as gdp
|
||||
|
||||
cfg = self._effective_config()
|
||||
video_layer = next(_lyr for _lyr in layers if _lyr.role not in ("audio",))
|
||||
video_clips = [c for c in video_layer.clips if c.clip_type != "audio"]
|
||||
|
||||
# 音频层处理:收集 TTS 分段与配音素材库整段音频
|
||||
# - TTS 分段(带 tts 标记)→ 无间隙 concat 成单文件
|
||||
# - 配音素材库(voice_library=True)→ 单独作为整段音轨(不走分段 concat,已从 0 覆盖整段)
|
||||
audio_layer = next((_lyr for _lyr in layers if _lyr.role == "audio"), None)
|
||||
tts_merged: Path | None = None
|
||||
voiceover_track: Path | None = None
|
||||
if audio_layer:
|
||||
tts_clips = [c for c in audio_layer.clips if (c.config or {}).get("tts") and c.local_path.exists()]
|
||||
if tts_clips:
|
||||
tts_merged = self._concat_audio_clips(tts_clips, tag="tts_direct")
|
||||
# 配音素材库整段音频(按 _maybe_add_voice_library_layer 约定只有一个 clip_id=voice_library_main)
|
||||
vo_clips = [
|
||||
c for c in audio_layer.clips if (c.config or {}).get("voice_library") and c.local_path.exists()
|
||||
]
|
||||
if vo_clips:
|
||||
voiceover_track = vo_clips[-1].local_path # 理论上只有一个,取最后一个
|
||||
logger.info(
|
||||
"[gpu-direct] 配音素材库音轨: plan_id=%s path=%s",
|
||||
self.plan.id,
|
||||
voiceover_track,
|
||||
)
|
||||
|
||||
# 额外独立音轨(TTS concat、配音素材库)→ gpu_direct_pipeline 会与主音轨/BGM 一起 amix
|
||||
extra_audio_tracks: list[tuple[Path, float]] = []
|
||||
if tts_merged:
|
||||
extra_audio_tracks.append((tts_merged, 1.0))
|
||||
if voiceover_track:
|
||||
extra_audio_tracks.append((voiceover_track, 1.0))
|
||||
|
||||
# BGM 本地文件
|
||||
bgm_path = Path(self.bgm_path) if self.bgm_path else None
|
||||
if bgm_path is not None and not bgm_path.exists():
|
||||
bgm_path = None
|
||||
|
||||
# 字幕/标题/BGM 配置整包透传
|
||||
title_cfg = cfg.get("title", {}) or cfg.get("title_config", {}) or {}
|
||||
if not isinstance(title_cfg, dict):
|
||||
title_cfg = {}
|
||||
title_text = ""
|
||||
if title_cfg.get("enabled", True):
|
||||
title_text = title_cfg.get("text", "") or ""
|
||||
|
||||
sub_cfg = cfg.get("subtitle", {}) or {}
|
||||
if not isinstance(sub_cfg, dict):
|
||||
sub_cfg = {}
|
||||
subtitle_segments: list[Any] = []
|
||||
static_subtitle_text = ""
|
||||
if sub_cfg.get("enabled", True):
|
||||
if sub_cfg.get("auto_generated") and self._asr_timeline_cache is not None:
|
||||
subtitle_segments = list(self._asr_timeline_cache.segments)
|
||||
else:
|
||||
# 静态字幕文本(用户手输):pipeline 内部会构造全片长 segment
|
||||
static_subtitle_text = (sub_cfg.get("text", "") or "").strip()
|
||||
|
||||
bgm_cfg = cfg.get("bgm", {}) or {}
|
||||
if not isinstance(bgm_cfg, dict):
|
||||
bgm_cfg = {}
|
||||
# 若 bgm.enabled 显式关闭,则强制 bgm_path=None(_prepare_bgm 已按 enabled 返回 None,双保险)
|
||||
if not bgm_cfg.get("enabled", True):
|
||||
bgm_path = None
|
||||
# 注入微片段 BGM 偏移(同 CPU 路径)
|
||||
if bgm_path is not None and not bgm_cfg.get("audio_offset"):
|
||||
_micro_off = self._get_micro_bgm_offset()
|
||||
if _micro_off:
|
||||
bgm_cfg = {**bgm_cfg, "audio_offset": _micro_off}
|
||||
|
||||
# 边缘裁剪:dedup 开启时在 GPU 内做四边随机 2~5% 裁剪(gpu_direct_pipeline 内部随机)
|
||||
dedup = self._dedup_enabled()
|
||||
edge_pct = 0.03 if dedup else 0.0 # >0 表示启用;实际区间 [2%,5%] 在 pipeline 内随机
|
||||
|
||||
# 探测每个视频素材是否含音轨、读取 volume 配置
|
||||
clip_has_audio_list: list[bool] = []
|
||||
clip_volumes_list: list[float] = []
|
||||
for c in video_clips:
|
||||
lp = getattr(c, "local_path", None)
|
||||
_ha = False
|
||||
if lp and Path(lp).exists():
|
||||
try:
|
||||
_ha = probe_has_audio(str(lp))
|
||||
except Exception as _pe: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] probe_has_audio 失败按有声处理: %s", _pe)
|
||||
_ha = True
|
||||
clip_has_audio_list.append(_ha)
|
||||
_vol = float((c.config or {}).get("volume", 1.0))
|
||||
clip_volumes_list.append(_vol if _vol > 0 else 0.0)
|
||||
|
||||
# extra_audio_tracks 音量:从 audio_tracks_config 读(TTS/配音素材库),
|
||||
# 无法精确匹配 track_id 时保留默认 1.0
|
||||
at_cfg = cfg.get("audio_tracks") or {}
|
||||
tts_volume = 1.0
|
||||
vo_volume = 1.0
|
||||
if isinstance(at_cfg, dict):
|
||||
_tracks = at_cfg.get("tracks", []) or []
|
||||
for _t in _tracks:
|
||||
if not isinstance(_t, dict):
|
||||
continue
|
||||
try:
|
||||
_vol = float(_t.get("volume", 1.0))
|
||||
except (TypeError, ValueError):
|
||||
_vol = 1.0
|
||||
_tt = str(_t.get("track_type", ""))
|
||||
if _tt == "voiceover" and _t.get("audio_path"):
|
||||
vo_volume = max(0.0, min(2.0, _vol))
|
||||
# TTS 一般没有固定 track_type 标记,保持默认 1.0
|
||||
|
||||
extra_audio_tracks_cfg: list[tuple[Any, float]] = []
|
||||
if tts_merged:
|
||||
extra_audio_tracks_cfg.append((tts_merged, tts_volume))
|
||||
if voiceover_track:
|
||||
extra_audio_tracks_cfg.append((voiceover_track, vo_volume))
|
||||
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=video_clips,
|
||||
output_width=self.output_width,
|
||||
output_height=self.output_height,
|
||||
output_fps=self.output_fps,
|
||||
bgm_audio=bgm_path,
|
||||
title_text=title_text,
|
||||
subtitle_segments=subtitle_segments,
|
||||
edge_crop_pct=edge_pct,
|
||||
total_duration=video_duration,
|
||||
clip_has_audio=clip_has_audio_list,
|
||||
clip_volumes=clip_volumes_list,
|
||||
extra_audio_tracks=extra_audio_tracks_cfg,
|
||||
title_config=title_cfg,
|
||||
subtitle_config=sub_cfg,
|
||||
bgm_config=bgm_cfg,
|
||||
static_subtitle_text=static_subtitle_text,
|
||||
)
|
||||
|
||||
client = get_gpu_encoder()
|
||||
client.render_inputs_to_output(plan.inputs, plan.ffmpeg_args, output_path)
|
||||
|
||||
# 清理本次上传的临时音频
|
||||
for key in plan.oss_keys:
|
||||
try:
|
||||
from video_processing.oss_helpers import _storage
|
||||
|
||||
_storage().delete_file(key) if hasattr(_storage(), "delete_file") else None
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
did_edge_crop = bool(edge_pct)
|
||||
logger.info(
|
||||
"[gpu-direct] success: plan_id=%s clips=%d edge_crop=%s", self.plan.id, len(video_clips), did_edge_crop
|
||||
)
|
||||
return (True, did_edge_crop)
|
||||
|
||||
except GpuEncodeError as e:
|
||||
logger.warning("[gpu-direct] failed (fallback to legacy): %s", e)
|
||||
try:
|
||||
if output_path.exists():
|
||||
output_path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
return (None, False)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("[gpu-direct] unexpected error (fallback)", exc_info=True)
|
||||
return (None, False)
|
||||
|
||||
def _concat_audio_clips(self, clips: list[Any], *, tag: str) -> Path:
|
||||
"""把多个本地音频片段无间隙 concat 成一个 m4a(TTS 分段→单文件)。"""
|
||||
out = self.work_dir / f"{tag}_{self.plan.id}.m4a"
|
||||
listfile = self.work_dir / f"{tag}_{self.plan.id}.txt"
|
||||
lines = []
|
||||
for c in clips:
|
||||
ap = str(c.local_path).replace("'", "'\\''")
|
||||
lines.append(f"file '{ap}'")
|
||||
listfile.write_text("\n".join(lines), encoding="utf-8")
|
||||
cmd = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
str(listfile),
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(out),
|
||||
]
|
||||
run_ffmpeg(cmd)
|
||||
return out
|
||||
|
||||
def _gpu_encode_available(self) -> bool:
|
||||
"""GPU 编码客户端是否已配置且健康(缓存健康状态,单任务内只探测一次)。"""
|
||||
if not getattr(self, "_gpu_health_ok", None):
|
||||
@@ -2594,7 +2906,7 @@ class UnifiedRenderService:
|
||||
b = pixel_pert.get("color_b", 0)
|
||||
if r != 0 or g != 0 or b != 0:
|
||||
# color_balance 参数范围 -1.0 ~ 1.0,这里用 /100 转换
|
||||
filters.append(f"colorbalance=rs={r/100:.3f}:gs={g/100:.3f}:bs={b/100:.3f}")
|
||||
filters.append(f"colorbalance=rs={r / 100:.3f}:gs={g / 100:.3f}:bs={b / 100:.3f}")
|
||||
|
||||
@staticmethod
|
||||
def _clip_volume(clip: ResolvedClip) -> float:
|
||||
|
||||
@@ -18,6 +18,7 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.classification import ClassificationStatus
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
@@ -96,7 +97,7 @@ def calculate_asset_quality_task(self, asset_id: str) -> dict:
|
||||
confidence = 1.0
|
||||
existing_meta["classification"] = classification
|
||||
existing_meta["classification_confidence"] = confidence
|
||||
asset.classification_status = "completed"
|
||||
asset.classification_status = ClassificationStatus.COMPLETED
|
||||
asset.metadata = existing_meta
|
||||
logger.info(
|
||||
"[quality_score] asset=%s 自动分类完成: category=%s confidence=%.2f",
|
||||
@@ -110,6 +111,8 @@ def calculate_asset_quality_task(self, asset_id: str) -> dict:
|
||||
asset_id,
|
||||
cls_err,
|
||||
)
|
||||
# 分类失败显式标记 FAILED,避免停留在 PENDING 被反复重试
|
||||
asset.classification_status = ClassificationStatus.FAILED
|
||||
|
||||
asset_repo.update(asset)
|
||||
db.commit()
|
||||
|
||||
@@ -198,7 +198,7 @@ def _verify_url_accessible(
|
||||
retries: int = 2,
|
||||
max_redirects: int = 5,
|
||||
) -> bool:
|
||||
"""HEAD 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。
|
||||
"""GET+Range 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。
|
||||
|
||||
安全增强:
|
||||
- 请求前先做 SSRF 安全校验(内网IP/回环地址/链路本地地址等)
|
||||
@@ -255,11 +255,12 @@ def _verify_url_accessible(
|
||||
)
|
||||
raise
|
||||
|
||||
req = urllib.request.Request(safe_url, method="HEAD")
|
||||
req = urllib.request.Request(safe_url, method="GET")
|
||||
req.add_header("Range", "bytes=0-0")
|
||||
req.add_header("User-Agent", "xiaoxia-saas-worker/1.0")
|
||||
|
||||
with opener.open(req, timeout=timeout) as resp: # noqa: S310
|
||||
if 200 <= resp.status < 300:
|
||||
if 200 <= resp.status < 300 or resp.status == 206:
|
||||
return True
|
||||
if resp.status in (301, 302, 303, 307, 308):
|
||||
location = resp.headers.get("Location", "")
|
||||
@@ -615,69 +616,49 @@ def _precompute_render_metadata(
|
||||
# ── Celery Task ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _sync_task_config_to_plan(source_edit_plan_id: str, task_info: dict, db) -> str | None:
|
||||
"""将 GenerationTask 的配置同步到 EditPlan.config,返回配音本地路径(如果有)。
|
||||
def _build_task_config_override(task_info: dict) -> dict:
|
||||
"""Bug A: 从 task_info 构建任务级 config override 深拷贝,供渲染时覆盖 plan.config。
|
||||
|
||||
包括:title_config、BGM、输出分辨率。配音单独处理(需下载到本地)。
|
||||
所有渲染相关配置(title/bgm/export)从任务自身读取,不再依赖共享 plan.config,
|
||||
彻底消除同 plan 多任务并发渲染时的竞态覆盖问题。
|
||||
"""
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
|
||||
SQLAlchemyEditPlanRepository,
|
||||
)
|
||||
import copy
|
||||
|
||||
plan_repo = SQLAlchemyEditPlanRepository(db)
|
||||
plan = plan_repo.get(source_edit_plan_id)
|
||||
if plan is None:
|
||||
logger.error("[task] EditPlan not found: %s", source_edit_plan_id)
|
||||
return None
|
||||
override: dict = {}
|
||||
|
||||
plan_config = dict(plan.config or {})
|
||||
changed = False
|
||||
|
||||
# 标题配置
|
||||
# 标题配置(字段名归一化)
|
||||
title_config = task_info.get("title_config") or {}
|
||||
if title_config and isinstance(title_config, dict) and title_config.get("text", "").strip():
|
||||
cfg = dict(title_config)
|
||||
# 字段名归一化
|
||||
if isinstance(title_config, dict) and title_config:
|
||||
cfg = copy.deepcopy(title_config)
|
||||
if "font_size" in cfg and "size" not in cfg:
|
||||
cfg["size"] = cfg["font_size"]
|
||||
if "font_color" in cfg and "color" not in cfg:
|
||||
cfg["color"] = cfg["font_color"]
|
||||
plan_config["title"] = cfg
|
||||
changed = True
|
||||
logger.info("[task] title_config synced to plan: %s", cfg.get("text", "")[:30])
|
||||
override["title"] = cfg
|
||||
|
||||
# BGM 配置
|
||||
bgm_config = task_info.get("bgm_config") or {}
|
||||
if bgm_config:
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
existing_bgm = plan_config.get("bgm", {}) or {}
|
||||
plan_config["bgm"] = merge_bgm_config(existing_bgm, bgm_config)
|
||||
changed = True
|
||||
if isinstance(bgm_config, dict) and bgm_config:
|
||||
override["bgm"] = copy.deepcopy(bgm_config)
|
||||
|
||||
# 输出分辨率
|
||||
ow = task_info.get("output_width") or OUTPUT_WIDTH
|
||||
oh = task_info.get("output_height") or OUTPUT_HEIGHT
|
||||
if ow >= 100 and oh >= 100:
|
||||
export_cfg = dict(plan_config.get("export", {}) or {})
|
||||
export_cfg["resolution"] = f"{ow}x{oh}"
|
||||
plan_config["export"] = export_cfg
|
||||
changed = True
|
||||
override["export"] = {"resolution": f"{ow}x{oh}"}
|
||||
|
||||
if changed:
|
||||
plan.config = plan_config
|
||||
plan_repo.update(plan)
|
||||
logger.info("[task] plan.config synced: plan_id=%s", source_edit_plan_id)
|
||||
return override
|
||||
|
||||
|
||||
def _download_voice_for_task(task_info: dict, source_edit_plan_id: str) -> str | None:
|
||||
"""下载任务配音到本地临时文件,返回路径(不读写 plan.config)。"""
|
||||
import tempfile
|
||||
|
||||
# 配音下载
|
||||
voiceover_path: str | None = None
|
||||
voice_library_id = task_info.get("voice_library_id", "")
|
||||
# #1749:voice_ids 冗余字段已移除;配音一律以 voice_library_id 为准(独立配音每变体各自绑定)
|
||||
effective_voice_id = voice_library_id or ""
|
||||
|
||||
if effective_voice_id:
|
||||
import tempfile
|
||||
|
||||
voice_tmp = Path(tempfile.gettempdir()) / f"voice_{source_edit_plan_id}_{id(task_info)}.mp3"
|
||||
try:
|
||||
if _download_voice_asset(effective_voice_id, voice_tmp):
|
||||
@@ -693,19 +674,24 @@ def _render_from_edit_plan(
|
||||
task_id: str,
|
||||
source_edit_plan_id: str,
|
||||
task_info: dict,
|
||||
) -> tuple[Path, float, list[dict] | None, str | None, str | None, str]:
|
||||
) -> tuple[Path, float, list[dict] | None, str | None, str | None, str, bool]:
|
||||
"""从 EditPlan 数据库记录直接渲染(不再内存重建clips)。
|
||||
|
||||
Bug A 修复:不再通过 _sync_task_config_to_plan 写共享 plan.config;
|
||||
渲染配置通过 task_config_override 参数直接传入渲染层,彻底消除并发竞态。
|
||||
|
||||
Returns:
|
||||
(output_path, render_duration, cover_candidates, voiceover_path, temp_dir, thumbnail_url)
|
||||
(output_path, render_duration, cover_candidates, voiceover_path, temp_dir, thumbnail_url, edge_crop_applied)
|
||||
"""
|
||||
from video_processing.render_adapter import RenderAdapter
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
# 同步配置到 plan.config + 下载配音
|
||||
voiceover_path = _sync_task_config_to_plan(source_edit_plan_id, task_info, db)
|
||||
# Bug A: 构建任务级 config override(深拷贝自 task_info),不写 plan.config,避免并发竞态
|
||||
task_override = _build_task_config_override(task_info)
|
||||
# 下载配音到本地临时文件(不依赖 plan.config)
|
||||
voiceover_path = _download_voice_for_task(task_info, source_edit_plan_id)
|
||||
|
||||
# 进度回调
|
||||
def _progress_cb(progress: float, stage: str):
|
||||
@@ -721,6 +707,7 @@ def _render_from_edit_plan(
|
||||
job_id=task_id,
|
||||
progress_cb=_progress_cb,
|
||||
voiceover_audio_path=voiceover_path,
|
||||
task_config_override=task_override,
|
||||
)
|
||||
|
||||
if not result.success:
|
||||
@@ -746,6 +733,7 @@ def _render_from_edit_plan(
|
||||
voiceover_path,
|
||||
render_temp_dir,
|
||||
result.thumbnail_url or "",
|
||||
bool(getattr(result, "edge_crop_applied", False)),
|
||||
)
|
||||
finally:
|
||||
db.close()
|
||||
@@ -899,6 +887,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
voiceover_tmp_path,
|
||||
render_temp_dir,
|
||||
thumbnail_url,
|
||||
_gpu_edge_crop_done,
|
||||
) = _render_from_edit_plan(
|
||||
task_id=task_id,
|
||||
source_edit_plan_id=current_plan_id,
|
||||
@@ -935,6 +924,12 @@ def generate_video(self, task_id: str) -> dict:
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("降重", "已关闭边缘裁剪与微变换(确定性渲染)")
|
||||
_flush_logs(task_id, gen_task)
|
||||
elif _gpu_edge_crop_done:
|
||||
# GPU 直连管线已经在 filter_complex 中做了随机边缘裁剪,跳过 CPU 二次重编码
|
||||
if gen_task and render_attempt == 0:
|
||||
gen_task.append_log("边缘裁剪", "已在 GPU 直连管线内完成随机边缘裁剪")
|
||||
_flush_logs(task_id, gen_task)
|
||||
logger.info("[task_id=%s] GPU直连已完成边缘裁剪,跳过CPU二次重编码", task_id)
|
||||
else:
|
||||
from video_processing.ffmpeg_utils import random_edge_crop
|
||||
|
||||
@@ -1051,20 +1046,15 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
if _meta_model:
|
||||
meta = dict(_meta_model.extra_meta or {})
|
||||
# #2024/P0 finalize-400: 直接展开 _precompute_render_metadata 返回的
|
||||
# 完整 dict(含 file_url/fingerprint_dict/fingerprint_chunks/is_duplicate
|
||||
# /duplicate_of/...),避免手写字段白名单漏传字段导致 finalize 读不到数据。
|
||||
meta["rendered_output"] = {
|
||||
**dict(rendered_output or {}),
|
||||
# file_url/duration 由外层调用方拿到的实际上传结果,优先覆盖预计算值
|
||||
"file_url": file_url,
|
||||
"file_size": file_size,
|
||||
"duration": duration,
|
||||
"width": rendered_output.get("width", 1280),
|
||||
"height": rendered_output.get("height", 720),
|
||||
"fps": rendered_output.get("fps", 25.0),
|
||||
"name": rendered_output.get("name", ""),
|
||||
"thumbnail_url": rendered_output.get("thumbnail_url", ""),
|
||||
"mode": rendered_output.get("mode", editing_mode.value),
|
||||
"fingerprint_dict": rendered_output.get("fingerprint_dict"),
|
||||
"batch_id": batch_id,
|
||||
"project_id": project_id,
|
||||
"user_id": user_id,
|
||||
}
|
||||
_meta_model.extra_meta = meta
|
||||
_finalize_meta_session.commit()
|
||||
|
||||
@@ -280,3 +280,18 @@ USE_GPU_LIPSYNC=true
|
||||
GPU_LIPSYNC_POLL_INTERVAL=5
|
||||
GPU_LIPSYNC_WAIT_TIMEOUT=1200
|
||||
GPU_WORKER_STALE_SECONDS=300
|
||||
|
||||
# ==================== P4000 NVENC 硬件编码(GPU mezzanine relay)====================
|
||||
# 注意:这些值必须写死在模板里(不是 CI Secret),否则每次 CI 重新渲染 .env 都会被丢弃,
|
||||
# 导致 staging 发版后 GPU 编码静默降级到 CPU(P0 防复发)。
|
||||
ENABLE_GPU_ENCODE=true
|
||||
GPU_ENCODE_ENDPOINT=http://100.105.75.67:8900
|
||||
GPU_ENCODE_RELAY_BASE_URL=http://100.125.116.43:8092
|
||||
GPU_ENCODE_RELAY_INTERNAL_BASE_URL=http://xiaoxia-api-staging:8000
|
||||
GPU_ENCODE_RELAY_SECRET=0e1a8f0626438564a8b3fa92f3f2aac29e3c69bc02f2f85c
|
||||
GPU_ENCODE_VCODEC=h264_nvenc
|
||||
GPU_ENCODE_PRESET=p4
|
||||
GPU_ENCODE_CRF=23
|
||||
GPU_ENCODE_FALLBACK_CPU=true
|
||||
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
||||
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
# Host nginx config for staging server: /etc/nginx/sites-available/05-xiaoxia-cms
|
||||
# Xiaoxia CMS - cms.xiaoxiajianji.com
|
||||
#
|
||||
# 注意:此文件是宿主机 nginx 配置的备份/参考,不是 Docker 容器内的 nginx。
|
||||
# Docker 容器内的 nginx 配置见 nginx-staging.conf。
|
||||
#
|
||||
# GPU relay 路由说明:
|
||||
# P4000 编码完成后通过 http://100.69.73.60:8092/api/v1/internal/gpu-relay/{key} PUT 上传
|
||||
# Worker 容器通过 http://xiaoxia-api-staging:8000/api/v1/internal/gpu-relay/{key} GET 下载
|
||||
# 8092 端口由 CMS 宿主机 nginx 承载,GPU relay 路由通过最长前缀匹配优先代理到 staging API (8000)
|
||||
|
||||
# 80 端口:ACME 验证 + 重定向到 HTTPS
|
||||
server {
|
||||
listen 80;
|
||||
server_name cms.xiaoxiajianji.com;
|
||||
|
||||
# Let's Encrypt ACME 验证
|
||||
location /.well-known/acme-challenge/ {
|
||||
root /var/www/certbot;
|
||||
}
|
||||
|
||||
location / {
|
||||
return 301 https://$host$request_uri;
|
||||
}
|
||||
}
|
||||
|
||||
# 443 端口:CMS 主站
|
||||
server {
|
||||
listen 443 ssl http2;
|
||||
server_name cms.xiaoxiajianji.com;
|
||||
|
||||
ssl_certificate /etc/letsencrypt/live/cms.xiaoxiajianji.com/fullchain.pem;
|
||||
ssl_certificate_key /etc/letsencrypt/live/cms.xiaoxiajianji.com/privkey.pem;
|
||||
include /etc/letsencrypt/options-ssl-nginx.conf;
|
||||
ssl_dhparam /etc/letsencrypt/ssl-dhparams.pem;
|
||||
|
||||
client_max_body_size 50m;
|
||||
|
||||
# Security headers
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "strict-origin-when-cross-origin" always;
|
||||
add_header Strict-Transport-Security "max-age=31536000; includeSubDomains" always;
|
||||
|
||||
root /data/www/cms/current;
|
||||
index index.html;
|
||||
|
||||
gzip on;
|
||||
gzip_types text/plain text/css application/json application/javascript text/xml application/xml application/xml+rss text/javascript image/svg+xml;
|
||||
gzip_min_length 1024;
|
||||
|
||||
location / {
|
||||
try_files $uri $uri/ /index.html;
|
||||
}
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8091/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8091/health;
|
||||
}
|
||||
}
|
||||
|
||||
# 临时访问:8092 端口(IP直接访问,后续可关闭)
|
||||
server {
|
||||
listen 8092;
|
||||
server_name _;
|
||||
|
||||
root /data/www/cms/current;
|
||||
index index.html;
|
||||
|
||||
client_max_body_size 50m;
|
||||
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "strict-origin-when-cross-origin" always;
|
||||
add_header Strict-Transport-Security "max-age=31536000; includeSubDomains" always;
|
||||
|
||||
gzip on;
|
||||
gzip_types text/plain text/css application/json application/javascript text/xml application/xml application/xml+rss text/javascript image/svg+xml;
|
||||
gzip_min_length 1024;
|
||||
|
||||
location / {
|
||||
try_files $uri $uri/ /index.html;
|
||||
}
|
||||
|
||||
# GPU relay endpoints - proxy to staging API (port 8000) instead of CMS
|
||||
# 此 location 必须在 location /api/ 之前,利用 nginx 最长前缀匹配优先路由
|
||||
location /api/v1/internal/gpu-relay/ {
|
||||
proxy_pass http://127.0.0.1:8000/api/v1/internal/gpu-relay/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_request_buffering off;
|
||||
proxy_read_timeout 600s;
|
||||
proxy_send_timeout 600s;
|
||||
}
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8091/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8091/health;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
# 预设 BGM 音频文件放置目录
|
||||
|
||||
将下列 10 首免费可商用 BGM 的 mp3/m4a 文件按 `{preset_id}.mp3` 命名放到本目录:
|
||||
|
||||
| preset_id | 名称 | 风格 | 时长(s) | 标签 |
|
||||
|-------------------|----------|----------|---------|----------------------------|
|
||||
| bgm_upbeat_001 | 阳光清晨 | upbeat | 120 | 轻快 阳光 吉他 vlog |
|
||||
| bgm_upbeat_002 | 活力节拍 | upbeat | 95 | 轻快 电子 活力 运动 |
|
||||
| bgm_upbeat_003 | 夏日漫步 | upbeat | 110 | 轻快 夏日 ukulele 旅行 |
|
||||
| bgm_relax_001 | 静谧时光 | relax | 180 | 治愈 钢琴 安静 冥想 |
|
||||
| bgm_relax_002 | 雨后森林 | relax | 150 | 治愈 自然 放松 环境音 |
|
||||
| bgm_relax_003 | 月光奏鸣曲 | relax | 200 | 治愈 古典 钢琴 优雅(公版) |
|
||||
| bgm_tech_001 | 未来科技 | tech | 85 | 科技 电子 未来感 数码 |
|
||||
| bgm_tech_002 | 数据脉冲 | tech | 100 | 科技 极简 数据 AI |
|
||||
| bgm_commerce_001 | 心动时刻 | commerce | 75 | 电商 时尚 动感 带货 |
|
||||
| bgm_commerce_002 | 品质生活 | commerce | 90 | 电商 高端 品牌 品质 |
|
||||
|
||||
放好后执行(需要在有 OSS 凭证的机器上):
|
||||
```bash
|
||||
export OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com
|
||||
export OSS_ACCESS_KEY_ID=xxx
|
||||
export OSS_ACCESS_KEY_SECRET=xxx
|
||||
export OSS_BUCKET_NAME=xiaoxia-autocut
|
||||
python scripts/upload_preset_bgm.py
|
||||
```
|
||||
|
||||
脚本会:
|
||||
1. 上传文件到 OSS `preset/bgm/<preset_id>.mp3`,设置公共读 ACL
|
||||
2. 自动改写 `packages/domain/preset_bgm.py` 把对应 `audio_url=""` 回填成公网 URL
|
||||
3. 提示 `git commit & push`
|
||||
+44
-72
@@ -10,20 +10,25 @@
|
||||
# API_IMAGE - API 镜像名称 (默认: xiaoxia-saas-api:dev)
|
||||
# WORKER_IMAGE - Worker 镜像名称 (默认: xiaoxia-saas-worker:dev)
|
||||
# WEB_IMAGE - Web 镜像名称 (默认: xiaoxia-saas-web:dev)
|
||||
# WEB_DOCKERFILE - Web Dockerfile 路径
|
||||
# WEB_NGINX_CONF - Nginx 配置文件路径
|
||||
# API_PORT - API 端口映射 (staging: 8000, production: 8001)
|
||||
# WEB_PORT - Web 端口映射 (staging: 3001, production: 3002)
|
||||
# GENERATED_FILES_HOST_DIR - 生成文件的主机目录
|
||||
# WORKER_CONCURRENCY - Worker 并发数 (默认: 4)
|
||||
# GENERATION_CONCURRENCY - Generation worker 并发(用户实时任务,默认 2)
|
||||
# TRANSCODE_CONCURRENCY - Transcode worker 并发(后台/转码/AI,默认 2)
|
||||
# WORKER_MAX_TASKS_PER_CHILD - Worker 每个子进程最大任务数 (默认: 100)
|
||||
# BEAT_ENABLED - 容器内启动 celery beat(默认 1;独立 beat 容器部署设为 0)
|
||||
# WORKER_CONCURRENCY - 兼容旧变量:未显式设置上面两个并发时按此总数分配
|
||||
#
|
||||
# 重要:
|
||||
# 重要:
|
||||
# - 生产环境不要挂载 web-dist volume,否则会导致 403
|
||||
# - 确保环境隔离网络已创建: docker network create xiaoxia-net-${ENV}
|
||||
# - ENV=staging → xiaoxia-net-staging
|
||||
# - ENV=production → xiaoxia-net-production
|
||||
# - ENV=staging -> xiaoxia-net-staging
|
||||
# - ENV=production -> xiaoxia-net-production
|
||||
#
|
||||
# #2073 队列分流:worker 容器内跑三个独立进程——beat(只发定时任务)、
|
||||
# generation worker(只消费 generation 队列,实时高优)、transcode worker(消费
|
||||
# transcode + celery 队列,后台任务)。beat 不再嵌入 generation worker,
|
||||
# 不占实时任务槽位;TRANSCODE_CONCURRENCY 独立伸缩,不再依赖 WORKER_CONCURRENCY 差值。
|
||||
|
||||
# ===========================================
|
||||
# 日志轮转配置(所有服务共享)
|
||||
@@ -40,53 +45,39 @@ services:
|
||||
# =========================================
|
||||
api:
|
||||
image: ${API_IMAGE:-xiaoxia-saas-api:dev}
|
||||
# 不在生产环境构建镜像,使用预构建的镜像
|
||||
# build:
|
||||
# context: ../..
|
||||
# dockerfile: infra/docker/api.Dockerfile
|
||||
|
||||
container_name: xiaoxia-api-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
stop_grace_period: 30s
|
||||
stop_signal: SIGTERM
|
||||
|
||||
# 环境变量文件(包含数据库密码等敏感信息)
|
||||
|
||||
env_file:
|
||||
- ../../.env
|
||||
|
||||
|
||||
environment:
|
||||
APP_ENV: ${APP_ENV:-staging}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
|
||||
# 端口映射
|
||||
# Staging: 8000 -> 8000
|
||||
# Production: 8001 -> 8000
|
||||
|
||||
ports:
|
||||
- "127.0.0.1:${API_PORT:-8000}:8000"
|
||||
|
||||
# 共享生成文件目录 + 抖音 cookies 等运行时配置
|
||||
|
||||
volumes:
|
||||
- generated-files:/app/generated
|
||||
- ../../deploy/configs:/app/configs:ro
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# 健康检查配置
|
||||
|
||||
healthcheck:
|
||||
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
start_period: 40s
|
||||
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -97,39 +88,48 @@ services:
|
||||
memory: 512M
|
||||
|
||||
# =========================================
|
||||
# Worker 服务(Celery 任务队列)
|
||||
# Worker 服务(#2073 队列分流:beat + generation + transcode 同容器三进程)
|
||||
# =========================================
|
||||
# 三个进程独立启动,任一退出则容器整体退出由 docker restart 拉起;
|
||||
# 各自的并发与资源占用通过环境变量控制:
|
||||
# - generation:GENERATION_CONCURRENCY(默认 2),消费 generation 队列
|
||||
# - transcode: TRANSCODE_CONCURRENCY(默认 2),消费 transcode,celery 队列
|
||||
# - beat: 不消费任务,只发定时任务到 celery 默认队列
|
||||
worker:
|
||||
image: ${WORKER_IMAGE:-xiaoxia-saas-worker:dev}
|
||||
|
||||
|
||||
container_name: xiaoxia-worker-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
# 长任务(ingest HEVC 转码最长 30min、生成硬超时 11min)给足优雅关闭窗口
|
||||
stop_grace_period: 300s
|
||||
stop_signal: SIGTERM
|
||||
|
||||
|
||||
env_file:
|
||||
- ../../.env
|
||||
|
||||
|
||||
environment:
|
||||
APP_ENV: ${APP_ENV:-staging}
|
||||
# 兼容旧变量:若两个 *_CONCURRENCY 均未显式设置,entrypoint 会按此总数分配
|
||||
WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4}
|
||||
WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100}
|
||||
# #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码
|
||||
# #2073 队列独立伸缩:generation 默认 2,transcode 默认 2(不再差值计算)
|
||||
GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2}
|
||||
TRANSCODE_CONCURRENCY: ${TRANSCODE_CONCURRENCY:-2}
|
||||
# beat 默认在本容器启动;独立 beat 容器部署时设为 0
|
||||
BEAT_ENABLED: ${BEAT_ENABLED:-1}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
|
||||
|
||||
volumes:
|
||||
- generated-files:/app/generated
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# 健康检查配置
|
||||
# 注:容器内无 pgrep/ps,扫描 /proc 所有进程的 cmdline 查找 celery 进程
|
||||
# 健康检查:至少有一个 celery worker 进程在跑(beat 本身不作为存活依据)
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "grep -lq celery /proc/[0-9]*/cmdline 2>/dev/null || exit 1"]
|
||||
test: ["CMD-SHELL", "grep -q 'celery.*worker' /proc/[0-9]*/cmdline 2>/dev/null || exit 1"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
@@ -137,12 +137,8 @@ services:
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
# 注意: Worker 需要处理视频,建议分配更多资源
|
||||
# #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
|
||||
# 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
|
||||
# 资源限制:容器总资源 = gen + trans + beat,按 2+2 并发场景建议 4C8G;
|
||||
# 后续如需独立扩容/重启,可拆为 worker-generation / worker-transcode / worker-beat 三个 service。
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -157,35 +153,21 @@ services:
|
||||
# =========================================
|
||||
web:
|
||||
image: ${WEB_IMAGE:-xiaoxia-saas-web:dev}
|
||||
# 不在生产环境构建镜像,使用 web-artifact.Dockerfile
|
||||
# build:
|
||||
# context: ../..
|
||||
# dockerfile: ${WEB_DOCKERFILE:-infra/docker/web.Dockerfile}
|
||||
# args:
|
||||
# (NGINX_CONF no longer needed - all configs baked into image)
|
||||
|
||||
|
||||
container_name: xiaoxia-web-${ENV:-staging}
|
||||
restart: unless-stopped
|
||||
|
||||
# 端口映射
|
||||
# Staging: 3001 -> 80
|
||||
# Production: 3002 -> 80 (通过 Nginx 反向代理)
|
||||
|
||||
ports:
|
||||
- "127.0.0.1:${WEB_PORT:-3001}:80"
|
||||
|
||||
|
||||
networks:
|
||||
- xiaoxia-net
|
||||
|
||||
# =========================================
|
||||
# Nginx 配置运行时覆盖(双保险:entrypoint 也按 APP_ENV 选择配置)
|
||||
# 确保容器使用正确环境的 nginx 配置,即使镜像构建时使用了默认配置
|
||||
# 注意: 只覆盖 /etc/nginx/conf.d/default.conf,不挂载 /usr/share/nginx/html
|
||||
# =========================================
|
||||
|
||||
environment:
|
||||
- APP_ENV=${ENV:-staging}
|
||||
volumes:
|
||||
- ./nginx-${ENV:-staging}.conf:/etc/nginx/conf.d/default.conf:ro
|
||||
|
||||
|
||||
healthcheck:
|
||||
test: ["CMD", "wget", "--spider", "-q", "http://127.0.0.1:80"]
|
||||
interval: 30s
|
||||
@@ -194,9 +176,6 @@ services:
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
# =========================================
|
||||
# 资源限制建议
|
||||
# =========================================
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
@@ -212,9 +191,6 @@ volumes:
|
||||
driver_opts:
|
||||
type: none
|
||||
o: bind
|
||||
# 重要: 确保主机目录存在且有正确权限
|
||||
# Staging: /var/lib/xiaoxia-saas-staging/generated
|
||||
# Production: /var/lib/xiaoxia-saas-production/generated
|
||||
device: ${GENERATED_FILES_HOST_DIR:?GENERATED_FILES_HOST_DIR must be set in .env}
|
||||
|
||||
# ===========================================
|
||||
@@ -223,8 +199,4 @@ volumes:
|
||||
networks:
|
||||
xiaoxia-net:
|
||||
external: true
|
||||
# 网络名根据 ENV 变量区分,实现 staging/production 环境隔离
|
||||
# staging: xiaoxia-net-staging
|
||||
# production: xiaoxia-net-production
|
||||
name: xiaoxia-net-${ENV:-staging}
|
||||
|
||||
|
||||
@@ -1,48 +1,77 @@
|
||||
#!/bin/bash
|
||||
# Worker 启动脚本 — #1714 队列隔离
|
||||
# Worker 启动脚本 — #1714 + #2073 队列分流
|
||||
#
|
||||
# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程:
|
||||
# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级),
|
||||
# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行;
|
||||
# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/
|
||||
# 配音/下载等后台任务)。
|
||||
# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。
|
||||
# 容器内启动三个独立进程(任一退出则整体退出由 docker restart 拉起):
|
||||
# 1. beat:celery beat 调度器,不消费任何任务,只发定时任务到 celery 默认队列
|
||||
# 2. generation-worker:独占消费 generation 队列(用户实时任务,高优先级)
|
||||
# 3. transcode-worker:消费 transcode + celery 默认队列(后台/清理任务)
|
||||
#
|
||||
# 环境变量:
|
||||
# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2,
|
||||
# 可用 GENERATION_CONCURRENCY 覆盖
|
||||
# GENERATION_CONCURRENCY generation worker 并发(默认 2)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 2)
|
||||
# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100)
|
||||
# WORKER_CONCURRENCY 兼容旧变量:若未显式设置 GENERATION_CONCURRENCY /
|
||||
# TRANSCODE_CONCURRENCY,则按比例分配(gen=ceil(total*1/2),
|
||||
# trans=剩余,各至少 1);已显式设置时忽略此变量。
|
||||
# BEAT_ENABLED 是否在本容器内启动 beat 进程(默认 1);
|
||||
# 若独立 beat 容器部署设为 0。
|
||||
|
||||
set -e
|
||||
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-4}"
|
||||
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
|
||||
|
||||
GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}"
|
||||
if [ -z "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then
|
||||
TRANS_CONCURRENCY=1
|
||||
fi
|
||||
# ── 并发计算:显式 env 优先;否则从 WORKER_CONCURRENCY 按比例推导 ──
|
||||
if [ -n "$GENERATION_CONCURRENCY" ]; then
|
||||
GEN_CONCURRENCY="$GENERATION_CONCURRENCY"
|
||||
else
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
TOTAL="${WORKER_CONCURRENCY:-4}"
|
||||
GEN_CONCURRENCY=$(( (TOTAL + 1) / 2 ))
|
||||
if [ "$GEN_CONCURRENCY" -lt 1 ]; then GEN_CONCURRENCY=1; fi
|
||||
fi
|
||||
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)"
|
||||
if [ -n "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
else
|
||||
if [ -n "$WORKER_CONCURRENCY" ] && [ -z "$GENERATION_CONCURRENCY" ]; then
|
||||
# 两个都没显式设置,按 WORKER_CONCURRENCY 分配剩余
|
||||
TOTAL="$WORKER_CONCURRENCY"
|
||||
TRANS_CONCURRENCY=$(( TOTAL - GEN_CONCURRENCY ))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then TRANS_CONCURRENCY=1; fi
|
||||
else
|
||||
# 默认 2(#2073:独立伸缩,不再依赖 WORKER_CONCURRENCY 差值)
|
||||
TRANS_CONCURRENCY=2
|
||||
fi
|
||||
fi
|
||||
|
||||
BEAT_ENABLED="${BEAT_ENABLED:-1}"
|
||||
|
||||
PIDS=()
|
||||
|
||||
# ── 1. Beat 调度器(独立进程,不消费任务)──
|
||||
if [ "$BEAT_ENABLED" = "1" ] || [ "$BEAT_ENABLED" = "true" ]; then
|
||||
echo "Starting beat scheduler (schedule file=/tmp/celerybeat-schedule)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
beat \
|
||||
--loglevel=info \
|
||||
-s /tmp/celerybeat-schedule &
|
||||
PIDS+=($!)
|
||||
fi
|
||||
|
||||
# ── 2. Generation worker(实时高优队列)──
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
worker \
|
||||
--loglevel=info \
|
||||
"-B" \
|
||||
-s /tmp/celerybeat-schedule \
|
||||
-Q generation \
|
||||
"--concurrency=${GEN_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n generation@%h &
|
||||
GEN_PID=$!
|
||||
PIDS+=($!)
|
||||
GEN_PID=${PIDS[1]:-${PIDS[0]}}
|
||||
|
||||
# ── 3. Transcode worker(后台 + 清理队列)──
|
||||
echo "Starting transcode worker (queues=transcode,celery, concurrency=$TRANS_CONCURRENCY)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
@@ -52,13 +81,22 @@ celery \
|
||||
"--concurrency=${TRANS_CONCURRENCY}" \
|
||||
"--max-tasks-per-child=${MAX_TASKS}" \
|
||||
-n transcode@%h &
|
||||
TRANS_PID=$!
|
||||
PIDS+=($!)
|
||||
TRANS_PID=${PIDS[2]:-${PIDS[1]}}
|
||||
|
||||
# 任一进程退出则终止另一个,让容器整体重启(restart: unless-stopped)
|
||||
trap 'echo "Shutting down workers..."; kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true' TERM INT
|
||||
# 任一进程退出则终止其他进程,让容器整体重启
|
||||
cleanup() {
|
||||
echo "Shutting down all celery processes..."
|
||||
for pid in "${PIDS[@]}"; do
|
||||
kill -TERM "$pid" 2>/dev/null || true
|
||||
done
|
||||
}
|
||||
trap cleanup TERM INT
|
||||
|
||||
wait -n $GEN_PID $TRANS_PID
|
||||
# wait -n 等待任意一个子进程退出(bash 4.3+)
|
||||
# 容器镜像基础为 python:3.11-slim,bash 版本满足
|
||||
wait -n "${PIDS[@]}"
|
||||
EXIT_CODE=$?
|
||||
echo "One worker exited (code=$EXIT_CODE), stopping the other..."
|
||||
kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true
|
||||
exit $EXIT_CODE
|
||||
echo "One celery process exited (code=$EXIT_CODE), stopping the rest..."
|
||||
cleanup
|
||||
exit "$EXIT_CODE"
|
||||
|
||||
@@ -10,6 +10,14 @@ FROM xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji/saas-worker-base:l
|
||||
# 构建参数:版本号(CI 传入 commit hash)
|
||||
ARG APP_VERSION=dev
|
||||
|
||||
# CJK 字体保障:确保 fonts-noto-cjk 已安装(base 镜像漂移兜底)+ 重建字体缓存
|
||||
# fc-cache 非致命;fc-match 结果只打日志用于排查,不阻断构建
|
||||
RUN apt-get update && (apt-get install -y --no-install-recommends fonts-noto-cjk fontconfig || true) \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& (fc-cache -fv || true) \
|
||||
&& echo "[font] fc-match sans:zh: $(fc-match -f '%{family}\n' sans:zh 2>/dev/null | head -1)" \
|
||||
&& echo "[font] fc-match Noto Sans CJK SC: $(fc-match 'Noto Sans CJK SC' 2>/dev/null | head -1)"
|
||||
|
||||
# 创建非 root 用户
|
||||
RUN groupadd -r celery \
|
||||
&& useradd -r -g celery -d /app -s /sbin/nologin celery \
|
||||
|
||||
@@ -1,31 +1,29 @@
|
||||
# Staging GPU relay plain-HTTP vhost (P4000 NVENC 编码回传入口)
|
||||
# - 监听 8092 端口纯 HTTP(绕开 HTTPS 证书与 P4000 httpx SSL 问题)
|
||||
# - 代理到本机 staging API 的 /api/ 路径(127.0.0.1:8000 是 docker 映射端口)
|
||||
# - P4000 通过 Tailscale 直连宿主机 100.69.73.60:8092 PUT 编码结果
|
||||
# - Worker 通过 Docker DNS (xiaoxia-api-staging:8000) 直接 GET/DELETE,
|
||||
# 不经宿主机 nginx,避免 UFW FORWARD DROP 阻断
|
||||
# Staging GPU relay nginx 配置说明
|
||||
#
|
||||
# 部署:cp infra/nginx/gpu-relay-staging.conf /etc/nginx/conf.d/ && nginx -t && systemctl reload nginx
|
||||
|
||||
server {
|
||||
listen 8092;
|
||||
server_name _;
|
||||
|
||||
client_max_body_size 2048m;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://127.0.0.1:8000/api/;
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_request_buffering off;
|
||||
proxy_read_timeout 600s;
|
||||
proxy_send_timeout 600s;
|
||||
}
|
||||
|
||||
location = /health {
|
||||
proxy_pass http://127.0.0.1:8000/health;
|
||||
}
|
||||
}
|
||||
# GPU relay 并没有独立的 nginx vhost,而是集成在宿主机 CMS nginx 的 8092 server block 中。
|
||||
# 完整宿主机 nginx 配置备份见: deploy/configs/host-nginx-cms-staging.conf
|
||||
#
|
||||
# 核心 location 块(添加到 8092 server block,位于 location /api/ 之前):
|
||||
#
|
||||
# # GPU relay endpoints - proxy to staging API (port 8000) instead of CMS
|
||||
# location /api/v1/internal/gpu-relay/ {
|
||||
# proxy_pass http://127.0.0.1:8000/api/v1/internal/gpu-relay/;
|
||||
# proxy_http_version 1.1;
|
||||
# proxy_set_header Host $host;
|
||||
# proxy_set_header X-Real-IP $remote_addr;
|
||||
# proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
# proxy_set_header X-Forwarded-Proto $scheme;
|
||||
# proxy_request_buffering off;
|
||||
# proxy_read_timeout 600s;
|
||||
# proxy_send_timeout 600s;
|
||||
# }
|
||||
#
|
||||
# 部署方式:手动将上述 location 块添加到 /etc/nginx/sites-available/05-xiaoxia-cms 的 8092 server block 中
|
||||
# 然后 nginx -t && systemctl reload nginx
|
||||
#
|
||||
# 原理说明:
|
||||
# - 8092 端口由 CMS 宿主机 nginx 承载(与 CMS 共享端口)
|
||||
# - GPU relay 路由 /api/v1/internal/gpu-relay/ 比 CMS 的 /api/ 更具体
|
||||
# - nginx 最长前缀匹配确保 relay 请求路由到 staging API (port 8000) 而非 CMS (port 8091)
|
||||
# - P4000 通过 Tailscale IP 100.69.73.60:8092 访问 relay
|
||||
# - Worker 容器通过 Docker DNS xiaoxia-api-staging:8000 直接访问 relay
|
||||
|
||||
@@ -14,6 +14,24 @@ class InMemoryIngestJobRepository:
|
||||
def get(self, job_id: str) -> IngestJob | None:
|
||||
return self._items.get(job_id)
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
if not asset_id:
|
||||
return None
|
||||
from packages.domain.classification import IngestJobStatus
|
||||
|
||||
running: IngestJob | None = None
|
||||
completed: IngestJob | None = None
|
||||
for job in self._items.values():
|
||||
if getattr(job, "asset_id", "") != asset_id:
|
||||
continue
|
||||
if job.status in (IngestJobStatus.PENDING, IngestJobStatus.PROCESSING):
|
||||
if running is None or job.created_at > running.created_at:
|
||||
running = job
|
||||
elif job.status == IngestJobStatus.COMPLETED:
|
||||
if completed is None or job.created_at > completed.created_at:
|
||||
completed = job
|
||||
return running or completed
|
||||
|
||||
def update(self, job: IngestJob) -> IngestJob:
|
||||
self._items[job.id] = job
|
||||
return job
|
||||
|
||||
@@ -128,8 +128,12 @@ class SQLAlchemyAssetRepository:
|
||||
height=asset.height,
|
||||
fps=asset.fps,
|
||||
codec=asset.codec,
|
||||
status=asset.status.value,
|
||||
classification_status=asset.classification_status.value,
|
||||
status=(asset.status.value if hasattr(asset.status, "value") else str(asset.status)),
|
||||
classification_status=(
|
||||
asset.classification_status.value
|
||||
if hasattr(asset.classification_status, "value")
|
||||
else str(asset.classification_status)
|
||||
),
|
||||
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
|
||||
quality_score=asset.quality_score,
|
||||
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
|
||||
@@ -142,7 +146,7 @@ class SQLAlchemyAssetRepository:
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
# Issue #1776: 自动维护素材库计数(同事务内原子更新)
|
||||
if asset.library_id and asset.status.value != "deleted":
|
||||
if asset.library_id and (getattr(asset.status, "value", str(asset.status)) != "deleted"):
|
||||
from sqlalchemy import func
|
||||
|
||||
self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == asset.library_id).update(
|
||||
@@ -168,8 +172,12 @@ class SQLAlchemyAssetRepository:
|
||||
model.height = asset.height
|
||||
model.fps = asset.fps
|
||||
model.codec = asset.codec
|
||||
model.status = asset.status.value
|
||||
model.classification_status = asset.classification_status.value
|
||||
model.status = asset.status.value if hasattr(asset.status, "value") else str(asset.status)
|
||||
model.classification_status = (
|
||||
asset.classification_status.value
|
||||
if hasattr(asset.classification_status, "value")
|
||||
else str(asset.classification_status)
|
||||
)
|
||||
model.classification_result = json.dumps(asset.metadata) if asset.metadata else None
|
||||
model.quality_score = asset.quality_score
|
||||
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
|
||||
|
||||
@@ -43,6 +43,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
output_height=getattr(model, "output_height", 720) or 720,
|
||||
cover_url=getattr(model, "cover_url", "") or "",
|
||||
title_config=dict(getattr(model, "title_config", {}) or {}),
|
||||
extra_meta=dict(getattr(model, "extra_meta", {}) or {}),
|
||||
logs=model.logs or "[]",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
@@ -88,6 +89,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
output_height=task.output_height,
|
||||
cover_url=task.cover_url or "",
|
||||
title_config=dict(task.title_config) if task.title_config else {},
|
||||
extra_meta=dict(task.extra_meta) if task.extra_meta else {},
|
||||
logs=task.logs,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
@@ -322,6 +324,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
model.output_height = task.output_height
|
||||
model.cover_url = task.cover_url or ""
|
||||
model.title_config = dict(task.title_config) if task.title_config else {}
|
||||
model.extra_meta = dict(task.extra_meta) if task.extra_meta else {}
|
||||
model.logs = task.logs
|
||||
self.session.commit()
|
||||
return task
|
||||
|
||||
@@ -46,6 +46,43 @@ class SQLAlchemyIngestJobRepository:
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
"""返回 asset 最近一条未失败的 ingest job(PENDING/PROCESSING/COMPLETED 均算存在,用于幂等判断)。"""
|
||||
if not asset_id:
|
||||
return None
|
||||
# 优先返回仍在跑的 (PENDING/PROCESSING),否则返回最新一条 COMPLETED
|
||||
model = (
|
||||
self.session.query(IngestJobModel)
|
||||
.filter(IngestJobModel.asset_id == asset_id)
|
||||
.filter(IngestJobModel.status.in_([IngestJobStatus.PENDING.value, IngestJobStatus.PROCESSING.value]))
|
||||
.order_by(IngestJobModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
model = (
|
||||
self.session.query(IngestJobModel)
|
||||
.filter(IngestJobModel.asset_id == asset_id)
|
||||
.filter(IngestJobModel.status == IngestJobStatus.COMPLETED.value)
|
||||
.order_by(IngestJobModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return IngestJob(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
library_id=model.library_id,
|
||||
storage_key=model.storage_key,
|
||||
status=IngestJobStatus(model.status),
|
||||
error_message=model.error_message,
|
||||
result_asset_id=model.result_asset_id,
|
||||
file_hash=model.file_hash or "",
|
||||
asset_id=getattr(model, "asset_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[IngestJob]:
|
||||
models = self.session.query(IngestJobModel).filter(IngestJobModel.project_id == project_id).all()
|
||||
return [self.get(model.id) for model in models if self.get(model.id) is not None]
|
||||
|
||||
@@ -918,3 +918,71 @@ class GpuWorkerModel(Base):
|
||||
capabilities = Column(String(500), nullable=False, default="") # 逗号分隔,如 "musetalk"
|
||||
last_heartbeat_at = Column(DateTime, nullable=True, index=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoJobModel(Base):
|
||||
"""爆款视频任务"""
|
||||
|
||||
__tablename__ = "viral_video_jobs"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
images = Column(JSON, nullable=False, default=list) # 产品图片 URL 列表
|
||||
industry = Column(String(100), nullable=False, default="")
|
||||
target_customer = Column(String(500), nullable=False, default="")
|
||||
persona_id = Column(String(36), nullable=False, default="")
|
||||
viral_structure = Column(String(50), nullable=False, default="")
|
||||
marketing_purpose = Column(String(100), nullable=False, default="")
|
||||
bgm_preference = Column(String(50), nullable=False, default="")
|
||||
duration = Column(Integer, nullable=False, default=30)
|
||||
user_copy_text = Column(Text, nullable=False, default="")
|
||||
fusion_level = Column(String(20), nullable=False, default="ai_polish")
|
||||
reference_audio_path = Column(String(1000), nullable=False, default="")
|
||||
# v1.3 新增字段
|
||||
reference_video_url = Column(String(1000), nullable=False, default="")
|
||||
style_strength = Column(String(20), nullable=False, default="medium")
|
||||
style_guide = Column(JSON, nullable=True)
|
||||
style_template_id = Column(String(36), nullable=False, default="", index=True)
|
||||
# 结果与状态
|
||||
status = Column(String(30), nullable=False, default="pending", index=True)
|
||||
intent_result = Column(JSON, nullable=True)
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
started_at = Column(DateTime(timezone=True), nullable=True)
|
||||
completed_at = Column(DateTime(timezone=True), nullable=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoStyleTemplateModel(Base):
|
||||
"""爆款视频风格模板配置表"""
|
||||
|
||||
__tablename__ = "viral_video_style_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
thumbnail_url = Column(String(1000), nullable=False, default="")
|
||||
style_config = Column(JSON, nullable=False, default=dict)
|
||||
is_system = Column(Boolean, nullable=False, default=True, index=True)
|
||||
sort_order = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class ViralVideoPromptTemplateModel(Base):
|
||||
"""爆款视频 Prompt 模板表(由 #2040 seed)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
prompt_type = Column(String(50), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
content = Column(Text, nullable=False, default="")
|
||||
variables = Column(JSON, nullable=False, default=list)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -52,6 +52,18 @@ class RenderedOutput:
|
||||
def from_dict(cls, data: dict[str, Any]) -> "RenderedOutput":
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("rendered_output must be a dict")
|
||||
# fingerprint_chunks 历史上有两种位置:
|
||||
# 1) 顶层 ``fingerprint_chunks``(由 compute_render_fingerprint_and_dedup 直接返回)
|
||||
# 2) 嵌套在 ``fingerprint_dict["chunks"]``(VideoFingerprint.to_dict() 序列化的结构)
|
||||
# 顶层优先;顶层为空时回退到嵌套位置,兼容旧数据。
|
||||
fp_dict = data.get("fingerprint_dict") or {}
|
||||
chunks_raw = data.get("fingerprint_chunks")
|
||||
if not chunks_raw and isinstance(fp_dict, dict):
|
||||
chunks_raw = fp_dict.get("chunks")
|
||||
# md5 同样可能在顶层或嵌套在 fingerprint_dict 内(历史数据兼容)
|
||||
md5_value = data.get("video_fingerprint_md5")
|
||||
if not md5_value and isinstance(fp_dict, dict):
|
||||
md5_value = fp_dict.get("md5")
|
||||
return cls(
|
||||
file_url=str(data.get("file_url") or ""),
|
||||
file_size=int(data.get("file_size") or 0),
|
||||
@@ -65,14 +77,14 @@ class RenderedOutput:
|
||||
batch_id=str(data.get("batch_id") or ""),
|
||||
project_id=str(data.get("project_id") or ""),
|
||||
user_id=str(data.get("user_id") or ""),
|
||||
fingerprint_dict=data.get("fingerprint_dict"),
|
||||
fingerprint_chunks=data.get("fingerprint_chunks"),
|
||||
fingerprint_dict=fp_dict or None,
|
||||
fingerprint_chunks=chunks_raw if isinstance(chunks_raw, list) else None,
|
||||
is_duplicate=bool(data.get("is_duplicate", False)),
|
||||
duplicate_of=data.get("duplicate_of"),
|
||||
duplicate_rate=_safe_float(data.get("duplicate_rate")),
|
||||
match_count=_safe_int(data.get("match_count")),
|
||||
visual_similarity=_safe_float(data.get("visual_similarity")),
|
||||
video_fingerprint_md5=str(data.get("video_fingerprint_md5") or ""),
|
||||
video_fingerprint_md5=str(md5_value or ""),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -51,7 +51,7 @@ class APISettings(SharedSettings):
|
||||
def validate_jwt_secret_key(cls, v):
|
||||
if v is None or v == "":
|
||||
raise ValueError(
|
||||
"JWT_SECRET_KEY must be set via environment variable. " "Do not use default value in production!"
|
||||
"JWT_SECRET_KEY must be set via environment variable. Do not use default value in production!"
|
||||
)
|
||||
# Block known insecure default values
|
||||
insecure_defaults = [
|
||||
@@ -63,7 +63,7 @@ class APISettings(SharedSettings):
|
||||
]
|
||||
if v.lower() in [d.lower() for d in insecure_defaults]:
|
||||
raise ValueError(
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. " "Please set a strong random secret via environment variable."
|
||||
f"JWT_SECRET_KEY '{v}' is insecure. Please set a strong random secret via environment variable."
|
||||
)
|
||||
return v
|
||||
|
||||
@@ -248,6 +248,10 @@ class APISettings(SharedSettings):
|
||||
def OSS_ENDPOINT(self) -> str:
|
||||
return self.oss_endpoint
|
||||
|
||||
@property
|
||||
def OSS_INTERNAL_ENDPOINT(self) -> str:
|
||||
return self.effective_oss_internal_endpoint
|
||||
|
||||
@property
|
||||
def OSS_ACCESS_KEY_ID(self) -> str:
|
||||
return self.oss_access_key_id
|
||||
|
||||
@@ -47,12 +47,37 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── OSS 阿里云 ──────────────────────────────────────────────────────
|
||||
oss_endpoint: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
# 内网 endpoint:ECS VPC 内访问 OSS 用(千兆带宽、免公网流量费)。
|
||||
# 为空时自动从 oss_endpoint 推导:若 oss_endpoint 是阿里云公网域名(形如
|
||||
# oss-cn-<region>.aliyuncs.com),自动加 -internal 得到内网域名;其他情况
|
||||
# (自定义域名/本地 MinIO/非阿里云)回退使用 oss_endpoint。
|
||||
# 显式填同值可以覆盖自动推导、强制所有流量都走公网。
|
||||
oss_internal_endpoint: str = ""
|
||||
oss_access_key_id: str = ""
|
||||
oss_access_key_secret: str = ""
|
||||
oss_bucket_name: str = "xiaoxia-autocut"
|
||||
oss_direct_upload_max_mb: int = 2000
|
||||
oss_direct_upload_expire_seconds: int = 900
|
||||
|
||||
@property
|
||||
def effective_oss_internal_endpoint(self) -> str:
|
||||
"""实际用于 SDK 内网访问的 endpoint(带 -internal 自动推导)。"""
|
||||
if self.oss_internal_endpoint:
|
||||
return self.oss_internal_endpoint
|
||||
ep = self.oss_endpoint.strip()
|
||||
scheme = ""
|
||||
host = ep
|
||||
if ep.startswith("https://"):
|
||||
scheme = "https://"
|
||||
host = ep[len("https://") :]
|
||||
elif ep.startswith("http://"):
|
||||
scheme = "http://"
|
||||
host = ep[len("http://") :]
|
||||
# 阿里云公网域名自动推导:oss-cn-<region>.aliyuncs.com → oss-cn-<region>-internal.aliyuncs.com
|
||||
if host.endswith(".aliyuncs.com") and "-internal" not in host and host.startswith("oss-cn-"):
|
||||
host = host[: -len(".aliyuncs.com")] + "-internal.aliyuncs.com"
|
||||
return f"{scheme}{host}" if scheme else host
|
||||
|
||||
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
|
||||
cosyvoice_api_key: str = ""
|
||||
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
@@ -76,6 +101,7 @@ class SharedSettings(BaseSettings):
|
||||
mediakit_api_key: str = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
mediakit_timeout: int = 60
|
||||
mediakit_cover_enabled: bool = False # 封面抽帧是否走MediaKit(默认false走本地ffmpeg+cv2,<2s完成)
|
||||
|
||||
# ── 积分/会员系统 (#1895) ────────────────────────────────────────────
|
||||
# 积分系统总开关(产品要求 #1895:暂停积分系统但保留全部代码/表/接口)。
|
||||
|
||||
@@ -12,5 +12,8 @@ class IngestJobRepository(Protocol):
|
||||
def get(self, job_id: str) -> IngestJob | None:
|
||||
"""Retrieve an ingest job by ID."""
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
"""Retrieve the most recent ingest job for a given asset (幂等判断)."""
|
||||
|
||||
def update(self, job: IngestJob) -> IngestJob:
|
||||
"""Update an ingest job and return it."""
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
"""Celery 队列定义与路由配置(API / Worker 共享)。
|
||||
|
||||
#1714 队列隔离:用户等待的视频生成任务路由到高优先级 `generation` 队列,
|
||||
由专用 worker 进程独占消费;素材入库/转码等后台批量任务路由到 `transcode`
|
||||
队列;其余杂项任务走默认 `celery` 队列。转码队列积压时,视频生成任务
|
||||
仍能被 generation worker 立即领取执行,不会排队。
|
||||
#1714 + #2073 队列分流:用户同步等待的实时任务路由到 `generation` 高优队列,
|
||||
由专用 generation worker 独占消费;素材入库/转码/AI 分析/查重等后台批量任务路由
|
||||
到 `transcode` 队列;beat 定时清理等轻量维护任务走默认 `celery` 队列。
|
||||
transcode / celery 队列积压时,generation 队列仍能被立即领取,不阻塞用户实时链路。
|
||||
|
||||
队列说明:
|
||||
- generation: 用户提交的视频生成/预览渲染(延迟敏感,资源消耗大)
|
||||
- transcode: 素材入库(HEVC 转码)、AI 分类、素材查重(批量、可排队)
|
||||
- celery(默认): 配音、语音、下载缩略图、定时清理等杂项
|
||||
- generation: 用户同步等待的实时任务(视频生成、TTS、音色克隆、lipsync、AI 数字人、人声/背景提取)
|
||||
- transcode: 后台批量/异步任务(素材入库转码、AI 分类打标、质量评分、原子切片、查重、批量下载/缩略图)
|
||||
- celery: beat 定时巡检/清理等轻量维护任务(极短、低优、不占业务槽)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -20,8 +20,9 @@ QUEUE_GENERATION = "generation"
|
||||
QUEUE_TRANSCODE = "transcode"
|
||||
QUEUE_DEFAULT = "celery"
|
||||
|
||||
# Worker 消费的队列列表(顺序即优先级:高优队列排在前面)
|
||||
WORKER_QUEUES = (QUEUE_GENERATION, QUEUE_TRANSCODE, QUEUE_DEFAULT)
|
||||
# 三个消费组各自消费的队列列表(顺序即优先级:高优队列排在前面)
|
||||
WORKER_QUEUES_GENERATION = (QUEUE_GENERATION,)
|
||||
WORKER_QUEUES_TRANSCODE = (QUEUE_TRANSCODE, QUEUE_DEFAULT)
|
||||
|
||||
# 队列声明:持久化队列,broker 重启不丢消息
|
||||
task_queues = (
|
||||
@@ -31,15 +32,45 @@ task_queues = (
|
||||
)
|
||||
|
||||
# ── 任务路由表:task name → 队列 ──
|
||||
# 键支持 celery 标准通配符。
|
||||
# 键支持 celery 标准通配符。所有生产端(API send_task / worker 内 send_task)
|
||||
# 未显式指定 queue 时按此表路由;漏配会走默认队列 celery,被 transcode worker 消费。
|
||||
# 新增实时任务务必在此表显式路由到 generation,避免落到后台队列排队。
|
||||
task_routes = {
|
||||
# 高优先级:用户等待的视频生成
|
||||
# ── 高优先级:用户同步等待的实时链路 ──
|
||||
# 视频生成(主链路)
|
||||
"worker.generate_video": {"queue": QUEUE_GENERATION},
|
||||
# 后台批量:素材入库/转码 + AI 分类 + 素材查重,积压不影响生成
|
||||
# TTS 合成 / 片段合成(配音页、视频生成配乐/TTS 链路)
|
||||
"worker.process_tts_synthesis": {"queue": QUEUE_GENERATION},
|
||||
"worker.process_tts_segment_synthesis": {"queue": QUEUE_GENERATION},
|
||||
# 音色克隆(用户主动上传样本等待克隆完成)
|
||||
"worker.process_voice_clone": {"queue": QUEUE_GENERATION},
|
||||
# 人声/背景提取(音色克隆前置步骤,用户同步等待)
|
||||
"worker.extract_voice": {"queue": QUEUE_GENERATION},
|
||||
"worker.extract_background": {"queue": QUEUE_GENERATION},
|
||||
# AI 数字人渲染(用户主动触发,等待成片)
|
||||
"ai_avatar_render.execute": {"queue": QUEUE_GENERATION},
|
||||
# GPU MuseTalk 口型同步(用户等成片,链路子任务全部走 generation 避免跨队列阻塞)
|
||||
"lipsync_gpu_process_async": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.synthesize_and_submit": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.poll_mediakit_status": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.persist_output_video": {"queue": QUEUE_GENERATION},
|
||||
# ── 后台批量:素材入库/转码 + AI 分析/打标 + 查重,积压不影响生成 ──
|
||||
"worker.ingest_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.classify_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.calculate_asset_quality": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.generate_atom_clips": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.tag_atom_clip": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.backfill_atom_clip_tags": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.process_duplication_check": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.check_duplicate": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.batch_download_videos": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.batch_generate_thumbnails": {"queue": QUEUE_TRANSCODE},
|
||||
# ── beat 定时清理/巡检任务走默认 celery 队列(由 transcode worker 消费)──
|
||||
# 未在此表显式列出的 cleanup 任务会落到默认队列 celery,不占 generation 槽位。
|
||||
"worker.cleanup_stale_pending_tasks": {"queue": QUEUE_DEFAULT},
|
||||
"worker.cleanup_stale_running_tasks": {"queue": QUEUE_DEFAULT},
|
||||
"worker.cleanup_stale_ingest_jobs": {"queue": QUEUE_DEFAULT},
|
||||
"worker.cleanup_stale_voice_clones": {"queue": QUEUE_DEFAULT},
|
||||
}
|
||||
|
||||
# 生成任务的预取数:渲染是长任务,预取 1 避免任务被某个 worker 占住不调度
|
||||
@@ -47,11 +78,10 @@ GENERATION_WORKER_PREFETCH_MULTIPLIER = 1
|
||||
|
||||
|
||||
def apply_queue_settings(app) -> None:
|
||||
"""把队列隔离配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。
|
||||
"""把队列分流配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。
|
||||
|
||||
配置 task_queues / task_routes / task_default_queue。生产端靠 task_routes
|
||||
把消息投递到对应队列;消费端靠 task_queues 声明自己消费哪些队列
|
||||
(实际消费集由启动参数 -Q 控制)。
|
||||
把消息投递到对应队列;消费端靠启动参数 -Q 控制自己消费哪些队列(entrypoint)。
|
||||
"""
|
||||
app.conf.task_queues = task_queues
|
||||
app.conf.task_routes = task_routes
|
||||
|
||||
+194
-11
@@ -7,13 +7,21 @@
|
||||
3. 生成 relay 一次性 key,构造两个带 token 的 URL:
|
||||
- put_url:给 P4000 回传结果,走 relay_base_url(Tailscale host:8092)
|
||||
- get/del_url:worker 自己下载+清理用,走 relay_internal_base_url(Docker DNS 直连 API)
|
||||
4. POST P4000 /api/render/sync:inputs={"in.mp4": "<mezzanine-get-url>"}, output_url="<put_url>"
|
||||
4. 【冷启动防护】距上次成功通信 >60s 时,先 GET /health 预热 Tailscale 链路(短超时快速失败)
|
||||
5. POST P4000 /api/render/sync:inputs={"in.mp4": "<mezzanine-get-url>"}, output_url="<put_url>"
|
||||
ffmpeg_args: -i in.mp4 [-vf <vf>] -c:v h264_nvenc ... -an/-c:a aac -f mp4 pipe:1
|
||||
5. P4000 从 relay GET mezzanine → h264_nvenc 编码 → PUT 最终 mp4 到 put_url
|
||||
6. 本客户端通过 get_url(Docker 内网)下载最终文件到 output_path,然后 DELETE 清理
|
||||
7. 删除 relay 上的 mezzanine 临时文件(以及 OSS fallback 的 key)
|
||||
- 首字节用短超时(默认20s),避免链路卡死空等上百秒;首字节到达后放宽到 ffmpeg_timeout+60s
|
||||
6. P4000 从 relay GET mezzanine → h264_nvenc 编码 → PUT 最终 mp4 到 put_url
|
||||
7. 本客户端通过 get_url(Docker 内网)下载最终文件到 output_path,然后 DELETE 清理
|
||||
8. 删除 relay 上的 mezzanine 临时文件(以及 OSS fallback 的 key)
|
||||
|
||||
任何环节失败抛 GpuEncodeError,调用方应 fallback 到 CPU libx264。
|
||||
|
||||
冷启动/链路卡顿背景(2026-09-27 实测):P4000 与 staging 之间走 Tailscale,长时间空闲
|
||||
(>7h)后首次请求曾出现 150s 延迟才真正开始下载 mezzanine,期间 ffmpeg 尚未启动、GPU 空闲。
|
||||
根因在服务端/网络层(可能是 Tailscale DERP 打洞或 httpx 连接池重建),本客户端通过
|
||||
pre_warm + 首字节短超时做兜底:预热打通链路 + 20s 内收不到首字节就快速失败让 CPU fallback,
|
||||
不再让用户等满 150s+。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -63,6 +71,13 @@ class GpuEncoderClient:
|
||||
mezzanine_transport: str = "relay",
|
||||
sync_timeout: int = 300,
|
||||
health_timeout: float = 3.0,
|
||||
# 提交编码任务前先发一次 /health 预热 Tailscale 链路,避免长时间空闲后首次请求
|
||||
# 因 DERP 打洞/NAT 映射过期/Tailscale 连接重建而阻塞上百秒。
|
||||
pre_warm: bool = True,
|
||||
# POST 首次响应超时:P4000 已收到请求后应该在数秒内开始下载 inputs;
|
||||
# 如果超过这个值还没收到任何响应字节,说明链路/服务卡住,快速失败让调用方 fallback CPU。
|
||||
# 注意:ffmpeg 编码本身靠 body.timeout 控制(300s),不应该被这个超时影响。
|
||||
post_first_byte_timeout: float = 20.0,
|
||||
vcodec: str = "h264_nvenc",
|
||||
preset: str = "p4",
|
||||
crf: int = 23,
|
||||
@@ -80,12 +95,16 @@ class GpuEncoderClient:
|
||||
self.mezzanine_transport = mezzanine_transport.lower() # "relay" | "oss"
|
||||
self.sync_timeout = sync_timeout
|
||||
self.health_timeout = health_timeout
|
||||
self.pre_warm = pre_warm
|
||||
self.post_first_byte_timeout = post_first_byte_timeout
|
||||
self.vcodec = vcodec
|
||||
self.preset = preset
|
||||
self.crf = crf
|
||||
self.bitrate = bitrate
|
||||
self._relay_secret = relay_secret
|
||||
self.oss_tmp_prefix = oss_tmp_prefix.rstrip("/") + "/" if oss_tmp_prefix else "tmp/gpu-mezzanine/"
|
||||
# 上次与 P4000 成功通信的时间戳(用于判断是否需要 pre_warm 预热)
|
||||
self._last_ok_ts: float = 0.0
|
||||
|
||||
RELAY_PATH_PREFIX = "/api/v1/internal/gpu-relay"
|
||||
|
||||
@@ -128,13 +147,16 @@ class GpuEncoderClient:
|
||||
except (urllib.error.URLError, socket.timeout, TimeoutError, json.JSONDecodeError, ConnectionError) as e:
|
||||
return GpuHealth(healthy=False, error=f"health probe failed: {e}")
|
||||
try:
|
||||
return GpuHealth(
|
||||
h = GpuHealth(
|
||||
healthy=data.get("status") == "healthy",
|
||||
worker=str(data.get("worker", "")),
|
||||
gpu_name=(data.get("gpu") or {}).get("name", ""),
|
||||
nvenc_h264=bool((data.get("nvenc") or {}).get("h264_nvenc")),
|
||||
nvenc_hevc=bool((data.get("nvenc") or {}).get("hevc_nvenc")),
|
||||
)
|
||||
if h.healthy:
|
||||
self._last_ok_ts = time.time()
|
||||
return h
|
||||
except Exception as e: # noqa: BLE001
|
||||
return GpuHealth(healthy=False, error=f"malformed health response: {e}")
|
||||
|
||||
@@ -227,7 +249,8 @@ class GpuEncoderClient:
|
||||
ffmpeg_args.append("-an")
|
||||
ffmpeg_args.extend(["-f", "mp4", "pipe:1"])
|
||||
|
||||
# 4. call P4000 sync render
|
||||
# 4. pre-warm then call P4000 sync render
|
||||
self._warm_up_if_needed()
|
||||
body = {
|
||||
"inputs": {"in.mp4": input_url},
|
||||
"ffmpeg_args": ffmpeg_args,
|
||||
@@ -235,6 +258,7 @@ class GpuEncoderClient:
|
||||
"timeout": int(timeout),
|
||||
}
|
||||
job = self._post_sync(body)
|
||||
self._last_ok_ts = time.time()
|
||||
logger.info(
|
||||
"[gpu-encoder] P4000 done: job_id=%s rc=%s size=%s dur=%ss transport=%s",
|
||||
job.get("job_id"),
|
||||
@@ -285,6 +309,103 @@ class GpuEncoderClient:
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[gpu-encoder] failed to delete OSS mezzanine %s: %s", oss_key, e)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# High-level: render arbitrary inputs → final output (all-GPU pipeline)
|
||||
# ------------------------------------------------------------------
|
||||
def render_inputs_to_output(
|
||||
self,
|
||||
inputs: dict[str, str],
|
||||
ffmpeg_args: list[str],
|
||||
output_path: Path,
|
||||
*,
|
||||
timeout: Optional[int] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""把多输入(原始素材/字幕/BGM)连同完整 filter_complex 交给 P4000 一次出片。
|
||||
|
||||
与 encode_mezzanine_to_output 的区别:worker 侧不再生成/上传 mezzanine,
|
||||
P4000 直接从 inputs 中的签名 URL 下载原始素材,filter_complex 内完成
|
||||
concat/scale/crop/drawtext/amix,末端 h264_nvenc 只编码一次。
|
||||
|
||||
传输:成片仍走 relay 回传(P4000 PUT → worker GET),避免公网 OSS 往返。
|
||||
|
||||
Args:
|
||||
inputs: {裸文件名: 可下载URL},key 即 ffmpeg_args 中引用的文件名
|
||||
ffmpeg_args: 完整 ffmpeg 参数(含 -i、-filter_complex、-map、NVENC 编码参数)
|
||||
output_path: worker 本地成片落盘路径
|
||||
timeout: P4000 侧超时(秒)
|
||||
"""
|
||||
if not inputs:
|
||||
raise GpuEncodeError("render_inputs_to_output: inputs is empty")
|
||||
if not ffmpeg_args:
|
||||
raise GpuEncodeError("render_inputs_to_output: ffmpeg_args is empty")
|
||||
if not self.relay_base_url:
|
||||
raise GpuEncodeError("gpu_encode_relay_base_url not configured")
|
||||
|
||||
timeout = timeout or self.sync_timeout
|
||||
t_total = time.time()
|
||||
result_key: Optional[str] = None
|
||||
|
||||
try:
|
||||
secret = self._get_relay_secret()
|
||||
|
||||
# 1. result relay URLs(成片 P4000 PUT → worker GET)
|
||||
result_key = uuid.uuid4().hex
|
||||
put_url = self._result_put_url(result_key, secret)
|
||||
get_url = self._result_get_url(result_key, secret)
|
||||
del_result_url = get_url
|
||||
|
||||
# 2. pre-warm then call P4000 sync render
|
||||
self._warm_up_if_needed()
|
||||
body = {
|
||||
"inputs": dict(inputs),
|
||||
"ffmpeg_args": list(ffmpeg_args),
|
||||
"output_url": put_url,
|
||||
"timeout": int(timeout),
|
||||
}
|
||||
job = self._post_sync(body)
|
||||
self._last_ok_ts = time.time()
|
||||
logger.info(
|
||||
"[gpu-encoder] P4000 direct done: job_id=%s rc=%s size=%s dur=%ss inputs=%d",
|
||||
job.get("job_id"),
|
||||
job.get("ffmpeg_rc"),
|
||||
job.get("size"),
|
||||
job.get("duration"),
|
||||
len(inputs),
|
||||
)
|
||||
|
||||
# 3. download result from relay to output_path
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
size = self._download_to_file(get_url, output_path)
|
||||
|
||||
# 4. cleanup relay result
|
||||
self._relay_delete(del_result_url)
|
||||
|
||||
logger.info(
|
||||
"[gpu-encoder] direct render ok → %s (%d bytes) total=%.2fs",
|
||||
output_path.name,
|
||||
size,
|
||||
time.time() - t_total,
|
||||
)
|
||||
return {
|
||||
"job": job,
|
||||
"output_size": size,
|
||||
"output_path": str(output_path),
|
||||
"transport": "direct",
|
||||
}
|
||||
|
||||
except GpuEncodeError:
|
||||
raise
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise GpuEncodeError(f"unexpected: {e}") from e
|
||||
finally:
|
||||
# cleanup relay result (best-effort)
|
||||
if result_key:
|
||||
try:
|
||||
secret = self._get_relay_secret()
|
||||
self._relay_delete(self._relay_result_url(self.relay_internal_base_url, result_key, secret))
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[gpu-encoder] failed to delete relay result %s: %s", result_key, e)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
@@ -299,9 +420,38 @@ class GpuEncoderClient:
|
||||
raise GpuEncodeError("GPU_ENCODE_RELAY_SECRET not set")
|
||||
return secret
|
||||
|
||||
def _warm_up_if_needed(self) -> None:
|
||||
"""POST 前预热:如果距上次成功通信超过 idle 阈值,先打 /health 打通 Tailscale 链路。
|
||||
|
||||
背景:Tailscale 在长时间空闲(几小时)后,到对端的直连 NAT 映射可能过期,
|
||||
首次请求会走 DERP 中继打洞;极少数情况下打洞/重连会卡住上百秒(曾观测到 150s 延迟)。
|
||||
预热请求本身走短超时快速失败,不会阻塞主流程;预热成功后再发 POST。
|
||||
"""
|
||||
if not self.pre_warm:
|
||||
return
|
||||
idle = time.time() - self._last_ok_ts
|
||||
# 空闲超过 60s 才预热(正常流水线里相邻任务间隔通常 <10s,没必要每次都打)
|
||||
if idle < 60:
|
||||
return
|
||||
url = f"{self.endpoint}/health"
|
||||
t0 = time.time()
|
||||
try:
|
||||
with urllib.request.urlopen(url, timeout=min(self.health_timeout, 3.0)) as resp:
|
||||
resp.read()
|
||||
self._last_ok_ts = time.time()
|
||||
logger.debug("[gpu-encoder] pre-warm ok: took=%.2fs idle=%.0fs", time.time() - t0, idle)
|
||||
except (urllib.error.URLError, socket.timeout, TimeoutError, ConnectionError, OSError) as e:
|
||||
# 预热失败不致命——主 POST 会带自己的超时,再失败就抛 GpuEncodeError 让调用方 fallback
|
||||
logger.warning("[gpu-encoder] pre-warm probe failed (will try POST anyway): %s", e)
|
||||
|
||||
def _post_sync(self, body: dict[str, Any]) -> dict[str, Any]:
|
||||
url = f"{self.endpoint}/api/render/sync"
|
||||
req_timeout = body.get("timeout", self.sync_timeout) + 60
|
||||
ffmpeg_timeout = body.get("timeout", self.sync_timeout)
|
||||
# 连接 + 首字节用短超时(防链路卡死数百秒);首字节到达后给 ffmpeg 留足编码+上传时间
|
||||
# Python urllib 的 timeout 是整个请求总超时,所以用"两段式":
|
||||
# 阶段1:先 read(1) 拿首字节,用短超时;
|
||||
# 阶段2:再 read() 读完整 body,用 ffmpeg_timeout+60。
|
||||
connect_timeout = min(max(self.post_first_byte_timeout, 5.0), 30.0)
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
req = urllib.request.Request(
|
||||
url,
|
||||
@@ -310,14 +460,45 @@ class GpuEncoderClient:
|
||||
method="POST",
|
||||
)
|
||||
t0 = time.time()
|
||||
first_byte_ok = False
|
||||
resp = None
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=req_timeout) as resp:
|
||||
raw = resp.read().decode("utf-8")
|
||||
resp = urllib.request.urlopen(req, timeout=connect_timeout)
|
||||
# 读首字节 —— 如果 P4000/链路卡死,这里会在 connect_timeout 内抛超时
|
||||
first_chunk = resp.read(1)
|
||||
first_byte_ok = True
|
||||
logger.debug(
|
||||
"[gpu-encoder] P4000 first byte in %.2fs (connect_timeout=%.1fs)",
|
||||
time.time() - t0,
|
||||
connect_timeout,
|
||||
)
|
||||
# 剩余用长超时(给底层socket放宽时限;如果是mock/不支持,则跳过)
|
||||
try:
|
||||
resp.fp._sock.settimeout(ffmpeg_timeout + 60)
|
||||
except (AttributeError, OSError):
|
||||
pass
|
||||
rest = resp.read()
|
||||
raw = (first_chunk + rest).decode("utf-8")
|
||||
resp.close()
|
||||
resp = None
|
||||
except urllib.error.HTTPError as e:
|
||||
detail = e.read().decode("utf-8", errors="replace")[:1000]
|
||||
raise GpuEncodeError(f"P4000 HTTP {e.code}: {detail}") from e
|
||||
except (urllib.error.URLError, socket.timeout, TimeoutError, ConnectionError) as e:
|
||||
raise GpuEncodeError(f"P4000 connection error: {e}") from e
|
||||
except (urllib.error.URLError, socket.timeout, TimeoutError, ConnectionError, OSError) as e:
|
||||
waited = time.time() - t0
|
||||
hint = "first-byte" if not first_byte_ok else "ffmpeg/upload"
|
||||
# 统一以 "connection error" 开头,便于上层 fallback 逻辑用关键词识别;
|
||||
# 末尾再附带具体错误(timed out / refused ...)供排障
|
||||
raise GpuEncodeError(
|
||||
f"P4000 {hint} connection error after {waited:.1f}s "
|
||||
f"(connect_timeout={connect_timeout:.0f}s, ffmpeg_timeout={ffmpeg_timeout}s): {e}"
|
||||
) from e
|
||||
finally:
|
||||
if resp is not None:
|
||||
try:
|
||||
resp.close()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
result = json.loads(raw)
|
||||
except json.JSONDecodeError as e:
|
||||
@@ -451,6 +632,8 @@ def _build_client_from_settings() -> Optional[GpuEncoderClient]:
|
||||
mezzanine_transport=getattr(settings, "gpu_encode_mezzanine_transport", "relay") or "relay",
|
||||
sync_timeout=getattr(settings, "gpu_encode_sync_timeout", 300),
|
||||
health_timeout=getattr(settings, "gpu_encode_health_timeout", 3.0),
|
||||
pre_warm=getattr(settings, "gpu_encode_pre_warm", True),
|
||||
post_first_byte_timeout=getattr(settings, "gpu_encode_post_first_byte_timeout", 20.0),
|
||||
vcodec=getattr(settings, "gpu_encode_vcodec", "h264_nvenc"),
|
||||
preset=getattr(settings, "gpu_encode_preset", "p4"),
|
||||
crf=getattr(settings, "gpu_encode_crf", 23),
|
||||
|
||||
+134
-168
@@ -5,6 +5,13 @@
|
||||
- Worker端 oss_helpers 的高级能力(分片上传/超时保护/HTTP下载/Asset路径解析)
|
||||
|
||||
所有服务都通过这个统一入口与存储交互,消除重复实现。
|
||||
|
||||
P1 (2026-09-28) OSS 双 endpoint 分离:
|
||||
- 内部 bucket(self.bucket):使用 internal endpoint(VPC 千兆带宽),
|
||||
用于所有 SDK 上传/下载/删除/object_exists 操作;
|
||||
- 公网 bucket(self.public_bucket):使用公网 endpoint,仅用于 sign_url
|
||||
生成给前端/P4000/MediaKit 等外网访问方用的预签名 URL;
|
||||
- public_url 永远拼公网域名,不随 internal endpoint 变化。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -42,19 +49,44 @@ OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
|
||||
OSS_HTTP_DOWNLOAD_TIMEOUT = 300 # HTTP下载超时(秒)
|
||||
|
||||
|
||||
class SharedStorageService(StoragePort):
|
||||
"""统一存储服务 — 实现 StoragePort,API 和 Worker 共用。
|
||||
def _make_bucket(
|
||||
auth,
|
||||
endpoint: str,
|
||||
bucket_name: str,
|
||||
*,
|
||||
connect_timeout: int = OSS_CONNECT_TIMEOUT,
|
||||
app_name: str = "",
|
||||
):
|
||||
"""构造 oss2.Bucket,自动补 https:// 前缀。"""
|
||||
if not endpoint.startswith(("http://", "https://")):
|
||||
endpoint = f"https://{endpoint}"
|
||||
kwargs: dict = {"connect_timeout": connect_timeout}
|
||||
if app_name:
|
||||
kwargs["app_name"] = app_name
|
||||
return oss2.Bucket(auth, endpoint, bucket_name, **kwargs)
|
||||
|
||||
整合了原 SharedStorageService + oss_helpers 的全部能力。
|
||||
"""
|
||||
|
||||
class SharedStorageService(StoragePort):
|
||||
"""统一存储服务 — 实现 StoragePort,API 和 Worker 共用。"""
|
||||
|
||||
# 类级默认值,方便单测 mock __init__ 后实例仍有这些属性
|
||||
bucket: Optional[object] = None
|
||||
public_bucket: Optional[object] = None
|
||||
public_endpoint: str = ""
|
||||
internal_endpoint: str = ""
|
||||
public_url: str = ""
|
||||
local_url_prefix: str = "/generated-files"
|
||||
bucket_name: str = ""
|
||||
|
||||
def __init__(self):
|
||||
settings = get_shared_settings()
|
||||
self.bucket_name = settings.oss_bucket_name
|
||||
self.endpoint = settings.oss_endpoint
|
||||
self.public_url = f"https://{settings.oss_bucket_name}.{settings.oss_endpoint}"
|
||||
self.public_endpoint = settings.oss_endpoint # 公网 endpoint,用于签名 URL
|
||||
self.internal_endpoint = settings.effective_oss_internal_endpoint # 内网 endpoint,SDK 用
|
||||
self.public_url = f"https://{settings.oss_bucket_name}.{self._public_host()}"
|
||||
self.local_url_prefix = os.getenv("GENERATED_FILES_URL_PREFIX", "/generated-files")
|
||||
self.bucket = None
|
||||
self.bucket: Optional[object] = None # internal: SDK 上传/下载/删除
|
||||
self.public_bucket: Optional[object] = None # public: sign_url 给外网
|
||||
|
||||
self.access_key_id = settings.oss_access_key_id
|
||||
self.access_key_secret = settings.oss_access_key_secret
|
||||
@@ -65,21 +97,26 @@ class SharedStorageService(StoragePort):
|
||||
if has_key_id and has_key_secret:
|
||||
if oss2 is not None:
|
||||
try:
|
||||
# endpoint 不带 scheme 时补 https:// 前缀
|
||||
bucket_endpoint = self.endpoint
|
||||
if not bucket_endpoint.startswith(("http://", "https://")):
|
||||
bucket_endpoint = f"https://{bucket_endpoint}"
|
||||
auth = oss2.Auth(self.access_key_id, self.access_key_secret)
|
||||
self.bucket = oss2.Bucket(
|
||||
self.bucket = _make_bucket(
|
||||
auth,
|
||||
bucket_endpoint,
|
||||
self.internal_endpoint,
|
||||
self.bucket_name,
|
||||
connect_timeout=OSS_CONNECT_TIMEOUT,
|
||||
app_name="xiaoxia-internal",
|
||||
)
|
||||
logger.info(
|
||||
"OSS initialized: endpoint=%s bucket=%s",
|
||||
self.endpoint,
|
||||
self.public_bucket = _make_bucket(
|
||||
auth,
|
||||
self.public_endpoint,
|
||||
self.bucket_name,
|
||||
app_name="xiaoxia-public",
|
||||
)
|
||||
same_ep = self.internal_endpoint == self.public_endpoint
|
||||
logger.info(
|
||||
"OSS initialized: public_ep=%s internal_ep=%s bucket=%s dual=%s",
|
||||
self.public_endpoint,
|
||||
self.internal_endpoint,
|
||||
self.bucket_name,
|
||||
"no" if same_ep else "yes",
|
||||
)
|
||||
except Exception as error:
|
||||
logger.error("Failed to initialize OSS bucket client: %s", error)
|
||||
@@ -93,26 +130,35 @@ class SharedStorageService(StoragePort):
|
||||
missing.append("OSS_ACCESS_KEY_SECRET")
|
||||
logger.error("OSS credentials not configured — missing: %s", ", ".join(missing))
|
||||
|
||||
def _public_host(self) -> str:
|
||||
ep = self.public_endpoint
|
||||
if ep.startswith("https://"):
|
||||
return ep[len("https://") :]
|
||||
if ep.startswith("http://"):
|
||||
return ep[len("http://") :]
|
||||
return ep
|
||||
|
||||
# ── 诊断 ───────────────────────────────────────────────────────────
|
||||
|
||||
def diagnose(self) -> None:
|
||||
"""输出存储配置诊断日志。"""
|
||||
key_id_display = (
|
||||
f"{self.access_key_id[:4]}...{self.access_key_id[-4:]}" if len(self.access_key_id) > 8 else "(empty)"
|
||||
)
|
||||
logger.info(
|
||||
"[OSS诊断] endpoint=%s bucket_name=%s access_key_id=%s",
|
||||
self.endpoint,
|
||||
"[OSS诊断] public_ep=%s internal_ep=%s bucket=%s ak=%s",
|
||||
self.public_endpoint,
|
||||
self.internal_endpoint,
|
||||
self.bucket_name,
|
||||
key_id_display,
|
||||
)
|
||||
if self.bucket is None:
|
||||
logger.error(
|
||||
"[OSS诊断] ❌ bucket=None — 预签名URL不可用!"
|
||||
"原因: OSS_ACCESS_KEY_ID/OSS_ACCESS_KEY_SECRET 未配置或 oss2 未安装。"
|
||||
)
|
||||
logger.error("[OSS诊断] ❌ bucket(internal)=None")
|
||||
else:
|
||||
logger.info("[OSS诊断] ✅ bucket 已配置,预签名URL可用")
|
||||
logger.info("[OSS诊断] ✅ bucket(internal) 就绪")
|
||||
if self.public_bucket is None:
|
||||
logger.error("[OSS诊断] ❌ public_bucket=None")
|
||||
else:
|
||||
logger.info("[OSS诊断] ✅ public_bucket 就绪,公网签名URL可用")
|
||||
|
||||
# ── 工具方法 ───────────────────────────────────────────────────────
|
||||
|
||||
@@ -122,20 +168,15 @@ class SharedStorageService(StoragePort):
|
||||
return path.startswith(f"{self.local_url_prefix}/")
|
||||
|
||||
def _normalize_storage_key(self, storage_key_or_url: str) -> str:
|
||||
"""从 URL 提取存储键,并做 URL 解码。
|
||||
|
||||
防止 URL 编码的字符(空格=%20、中文=%XX)导致签名不匹配。
|
||||
"""
|
||||
if storage_key_or_url.startswith("http://") or storage_key_or_url.startswith("https://"):
|
||||
parsed = urlparse(storage_key_or_url)
|
||||
return unquote(parsed.path.lstrip("/"))
|
||||
return storage_key_or_url.lstrip("/")
|
||||
|
||||
def normalize_storage_key(self, storage_key_or_url: str) -> str:
|
||||
"""从 URL 提取存储键(公开方法)。"""
|
||||
return self._normalize_storage_key(storage_key_or_url)
|
||||
|
||||
# ── 上传 ───────────────────────────────────────────────────────────
|
||||
# ── 上传(SDK 走 internal endpoint)───────────────────────────────
|
||||
|
||||
def upload_file(
|
||||
self,
|
||||
@@ -143,21 +184,22 @@ class SharedStorageService(StoragePort):
|
||||
storage_key: str,
|
||||
content_type: str = "application/octet-stream",
|
||||
) -> str:
|
||||
"""上传文件到存储,返回公开 URL(简单上传,API端原有行为)。
|
||||
|
||||
- 路径字符串 → bucket.put_object_from_file
|
||||
- 类文件对象 → bucket.put_object
|
||||
- bucket未配置 → 抛 RuntimeError
|
||||
"""
|
||||
if self.bucket is None:
|
||||
raise RuntimeError("OSS storage is not configured")
|
||||
|
||||
try:
|
||||
if isinstance(file_or_path, (str, Path)):
|
||||
self.bucket.put_object_from_file(storage_key, str(file_or_path), headers={"Content-Type": content_type})
|
||||
self.bucket.put_object_from_file(
|
||||
storage_key,
|
||||
str(file_or_path),
|
||||
headers={"Content-Type": content_type},
|
||||
)
|
||||
else:
|
||||
file_or_path.seek(0) # type: ignore[attr-defined]
|
||||
self.bucket.put_object(storage_key, file_or_path, headers={"Content-Type": content_type})
|
||||
file_or_path.seek(0)
|
||||
self.bucket.put_object(
|
||||
storage_key,
|
||||
file_or_path,
|
||||
headers={"Content-Type": content_type},
|
||||
)
|
||||
return f"{self.public_url}/{storage_key}"
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to upload file to OSS: {e}") from e
|
||||
@@ -167,14 +209,6 @@ class SharedStorageService(StoragePort):
|
||||
local_path: str | Path,
|
||||
storage_key: str,
|
||||
) -> Optional[str]:
|
||||
"""智能上传:大文件自动分片+超时保护(从 oss_helpers 合并)。
|
||||
|
||||
- 大文件(>100MB)走分片上传,3 线程并发
|
||||
- 总超时 300s,防止网络异常时挂死
|
||||
- 成功返回 URL,失败返回 None(不抛异常)
|
||||
|
||||
Worker端 oss_helpers.upload_to_oss 的统一入口。
|
||||
"""
|
||||
local_path = Path(local_path)
|
||||
if not local_path.exists():
|
||||
logger.error("上传文件不存在: %s", local_path)
|
||||
@@ -198,10 +232,10 @@ class SharedStorageService(StoragePort):
|
||||
|
||||
if use_multipart:
|
||||
logger.info(
|
||||
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
|
||||
"大文件分片上传(internal): key=%s size=%.1fMB part=%dMB threads=%d",
|
||||
storage_key[:80],
|
||||
file_size / 1024 / 1024,
|
||||
OSS_PART_SIZE // 1024 // 1024,
|
||||
file_size / 1048576,
|
||||
OSS_PART_SIZE // 1048576,
|
||||
OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
oss2.resumable_upload(
|
||||
@@ -214,7 +248,6 @@ class SharedStorageService(StoragePort):
|
||||
)
|
||||
else:
|
||||
self.bucket.put_object_from_file(storage_key, str(local_path))
|
||||
|
||||
result["url"] = f"{self.public_url}/{storage_key}"
|
||||
except Exception as e:
|
||||
result["error"] = e
|
||||
@@ -222,73 +255,54 @@ class SharedStorageService(StoragePort):
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
upload_thread = threading.Thread(target=_do_upload, daemon=True)
|
||||
upload_thread.start()
|
||||
t = threading.Thread(target=_do_upload, daemon=True)
|
||||
t.start()
|
||||
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
|
||||
|
||||
if not finished:
|
||||
logger.error(
|
||||
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
|
||||
"OSS 上传超时(%ds): key=%s size=%.1fMB",
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
storage_key[:80],
|
||||
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
|
||||
result["file_size"] / 1048576 if result["file_size"] else 0,
|
||||
)
|
||||
return None
|
||||
return None if result["error"] else result["url"]
|
||||
|
||||
if result["error"]:
|
||||
return None
|
||||
|
||||
return result["url"]
|
||||
|
||||
# ── 下载 ───────────────────────────────────────────────────────────
|
||||
# ── 下载(SDK 走 internal endpoint)───────────────────────────────
|
||||
|
||||
def download_file(self, storage_key: str, local_path: str | Path) -> None:
|
||||
"""从 OSS 下载文件(简单下载,API端原有行为)。
|
||||
|
||||
bucket未配置 → 抛 RuntimeError
|
||||
"""
|
||||
if self.bucket is None:
|
||||
raise RuntimeError("OSS storage is not configured")
|
||||
|
||||
local_path = Path(local_path)
|
||||
os.makedirs(local_path.parent, exist_ok=True)
|
||||
try:
|
||||
self.bucket.get_object_to_file(self._normalize_storage_key(storage_key), str(local_path))
|
||||
self.bucket.get_object_to_file(
|
||||
self._normalize_storage_key(storage_key),
|
||||
str(local_path),
|
||||
)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to download file from OSS: {e}") from e
|
||||
|
||||
def download_asset(self, asset_storage_key: str, local_path: str | Path) -> bool:
|
||||
"""下载素材(从 oss_helpers 合并)。
|
||||
|
||||
自动识别输入类型:
|
||||
- 完整 URL → 走 HTTP 下载(支持预签名URL)
|
||||
- 存储键 → 走 oss2 SDK 下载
|
||||
|
||||
成功返回 True,失败返回 False(不抛异常)。
|
||||
"""
|
||||
local_path = Path(local_path)
|
||||
os.makedirs(local_path.parent, exist_ok=True)
|
||||
|
||||
# 完整URL走HTTP下载(兼容预签名URL)
|
||||
if asset_storage_key.startswith(("http://", "https://")):
|
||||
return self._download_via_http(asset_storage_key, local_path)
|
||||
|
||||
# OSS存储键走SDK
|
||||
if self.bucket is None:
|
||||
logger.error("OSS not configured, cannot download: %s", asset_storage_key[:80])
|
||||
logger.error("OSS not configured: %s", asset_storage_key[:80])
|
||||
return False
|
||||
try:
|
||||
self.bucket.get_object_to_file(self._normalize_storage_key(asset_storage_key), str(local_path))
|
||||
self.bucket.get_object_to_file(
|
||||
self._normalize_storage_key(asset_storage_key),
|
||||
str(local_path),
|
||||
)
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("下载素材失败: %s", asset_storage_key)
|
||||
return False
|
||||
|
||||
def _download_via_http(self, url: str, local_path: Path) -> bool:
|
||||
"""通过 HTTP 下载文件(支持预签名 URL)。
|
||||
|
||||
流式下载避免大文件内存溢出。
|
||||
"""
|
||||
try:
|
||||
resp = requests.get(url, stream=True, timeout=OSS_HTTP_DOWNLOAD_TIMEOUT)
|
||||
resp.raise_for_status()
|
||||
@@ -298,83 +312,64 @@ class SharedStorageService(StoragePort):
|
||||
f.write(chunk)
|
||||
return local_path.exists() and local_path.stat().st_size > 0
|
||||
except Exception:
|
||||
logger.exception("HTTP下载素材失败: %s", url[:100])
|
||||
logger.exception("HTTP下载失败: %s", url[:100])
|
||||
return False
|
||||
|
||||
# ── URL 生成 ──────────────────────────────────────────────────────
|
||||
# ── URL 生成(sign_url 用 public_bucket 签公网域名)───────────────
|
||||
|
||||
def get_url(self, storage_key: str) -> str:
|
||||
"""获取公开 URL。"""
|
||||
return f"{self.public_url}/{storage_key}"
|
||||
|
||||
def get_download_url(self, storage_key_or_url: str, expires_seconds: int = 3600) -> str:
|
||||
"""获取预签名下载 URL。
|
||||
def _sign_bucket(self):
|
||||
"""签名优先用 public_bucket,回退到 bucket。"""
|
||||
return self.public_bucket or self.bucket
|
||||
|
||||
bucket未配置时降级为公开URL;本地产物URL直接返回。
|
||||
"""
|
||||
if self.bucket is None:
|
||||
def get_download_url(self, storage_key_or_url: str, expires_seconds: int = 3600) -> str:
|
||||
sign_bucket = self._sign_bucket()
|
||||
if sign_bucket is None:
|
||||
if self._is_local_generated_url(storage_key_or_url):
|
||||
return storage_key_or_url
|
||||
logger.warning(
|
||||
"get_download_url: OSS bucket not configured, returning raw URL. key=%s",
|
||||
storage_key_or_url[:200],
|
||||
)
|
||||
logger.warning("OSS bucket not configured, returning raw URL: %s", storage_key_or_url[:200])
|
||||
return self.get_url(self.normalize_storage_key(storage_key_or_url))
|
||||
|
||||
storage_key = self.normalize_storage_key(storage_key_or_url)
|
||||
try:
|
||||
signed = self.bucket.sign_url("GET", storage_key, expires_seconds)
|
||||
signed = sign_bucket.sign_url("GET", storage_key, expires_seconds)
|
||||
logger.info(
|
||||
"get_download_url: signed URL generated. key=%s url_prefix=%s",
|
||||
"signed URL generated for key=%s prefix=%s",
|
||||
storage_key[:80],
|
||||
signed[:60],
|
||||
)
|
||||
return signed
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"get_download_url: sign_url failed, falling back to raw URL. key=%s",
|
||||
storage_key[:200],
|
||||
)
|
||||
logger.exception("get_download_url: sign_url 失败,返回 raw URL: %s", storage_key[:200])
|
||||
return self.get_url(storage_key)
|
||||
|
||||
# ── 浏览器直传 POST ────────────────────────────────────────────────
|
||||
|
||||
def get_upload_url(
|
||||
self,
|
||||
storage_key_or_url: str,
|
||||
expires_seconds: int = 3600,
|
||||
content_type: str = "video/mp4",
|
||||
) -> str:
|
||||
"""获取预签名 PUT 上传 URL(供外部 Worker 上传结果文件)。
|
||||
|
||||
bucket未配置时降级为 public_url(本地/开发环境);
|
||||
本地产物 key 原样返回。
|
||||
"""
|
||||
if self.bucket is None:
|
||||
sign_bucket = self._sign_bucket()
|
||||
if sign_bucket is None:
|
||||
if self._is_local_generated_url(storage_key_or_url):
|
||||
return storage_key_or_url
|
||||
logger.warning(
|
||||
"get_upload_url: OSS bucket not configured, returning raw URL. key=%s",
|
||||
storage_key_or_url[:200],
|
||||
)
|
||||
logger.warning("get_upload_url: OSS 未配置,返回 raw URL: %s", storage_key_or_url[:200])
|
||||
return self.get_url(self.normalize_storage_key(storage_key_or_url))
|
||||
|
||||
storage_key = self.normalize_storage_key(storage_key_or_url)
|
||||
try:
|
||||
# oss2 sign_url 支持 'PUT',需指定 headers 才能限定 Content-Type
|
||||
headers = {"Content-Type": content_type} if content_type else None
|
||||
signed = self.bucket.sign_url("PUT", storage_key, expires_seconds, headers=headers)
|
||||
signed = sign_bucket.sign_url("PUT", storage_key, expires_seconds, headers=headers)
|
||||
logger.info(
|
||||
"get_upload_url: signed PUT URL generated. key=%s url_prefix=%s",
|
||||
"get_upload_url: 公网签名PUT URL已生成 key=%s prefix=%s",
|
||||
storage_key[:80],
|
||||
signed[:60],
|
||||
)
|
||||
return signed
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"get_upload_url: sign_url failed, falling back to raw URL. key=%s",
|
||||
storage_key[:200],
|
||||
)
|
||||
logger.exception("get_upload_url: sign_url 失败,返回 raw URL: %s", storage_key[:200])
|
||||
return self.get_url(storage_key)
|
||||
|
||||
def create_direct_upload_post(
|
||||
@@ -384,7 +379,6 @@ class SharedStorageService(StoragePort):
|
||||
max_size_bytes: int,
|
||||
expires_seconds: int,
|
||||
) -> dict[str, object]:
|
||||
"""创建浏览器直传 POST 表单。"""
|
||||
if not self.access_key_id or not self.access_key_secret:
|
||||
raise RuntimeError("OSS storage is not configured")
|
||||
normalized_key = self.normalize_storage_key(storage_key)
|
||||
@@ -431,19 +425,17 @@ class SharedStorageService(StoragePort):
|
||||
},
|
||||
}
|
||||
|
||||
# ── 文件操作 ───────────────────────────────────────────────────────
|
||||
# ── 文件操作(internal endpoint)──────────────────────────────────
|
||||
|
||||
def delete_file(self, storage_key: str) -> None:
|
||||
"""删除文件(不抛异常)。"""
|
||||
if self.bucket is None:
|
||||
return
|
||||
try:
|
||||
self.bucket.delete_object(storage_key)
|
||||
except Exception as error:
|
||||
logger.warning("Failed to delete file from OSS", extra={"storage_key": storage_key, "error": str(error)})
|
||||
logger.warning("OSS delete 失败", extra={"storage_key": storage_key, "error": str(error)})
|
||||
|
||||
def file_exists(self, storage_key: str) -> bool:
|
||||
"""检查文件是否存在。"""
|
||||
if self.bucket is None:
|
||||
return False
|
||||
return self.bucket.object_exists(storage_key)
|
||||
@@ -451,17 +443,6 @@ class SharedStorageService(StoragePort):
|
||||
# ── Asset 路径解析(Worker 用)────────────────────────────────────
|
||||
|
||||
def resolve_asset_path(self, asset_id: str, work_dir: str | Path) -> Optional[Path]:
|
||||
"""从 asset_id 解析到本地文件路径。
|
||||
|
||||
策略(按优先级):
|
||||
1. 本地绝对路径(在允许目录内)→ 直接返回
|
||||
2. work_dir 缓存命中 → 返回缓存路径
|
||||
3. 从OSS下载到缓存 → 返回下载路径
|
||||
4. 全部失败 → None
|
||||
|
||||
从 oss_helpers.resolve_asset_path 合并而来。
|
||||
"""
|
||||
# 延迟导入,避免循环依赖
|
||||
from video_processing.path_security import ( # type: ignore[import-not-found]
|
||||
PathSecurityError,
|
||||
get_allowed_local_dirs,
|
||||
@@ -471,47 +452,35 @@ class SharedStorageService(StoragePort):
|
||||
|
||||
if not asset_id or not isinstance(asset_id, str):
|
||||
return None
|
||||
|
||||
work_dir = Path(work_dir)
|
||||
os.makedirs(work_dir, exist_ok=True)
|
||||
|
||||
# 空字节检测
|
||||
if "\x00" in asset_id:
|
||||
logger.warning("asset_id 包含空字节,拒绝: %s", asset_id[:50])
|
||||
logger.warning("asset_id 含空字节,拒绝: %s", asset_id[:50])
|
||||
return None
|
||||
|
||||
# 1. 本地绝对路径 — 必须在允许的目录内
|
||||
if asset_id.startswith("/") and os.path.exists(asset_id):
|
||||
try:
|
||||
resolved = Path(asset_id).resolve()
|
||||
if is_in_allowed_dirs(resolved, get_allowed_local_dirs()):
|
||||
return resolved
|
||||
else:
|
||||
logger.warning(
|
||||
"本地素材路径不在允许目录内,拒绝: %s (allowed=%s)",
|
||||
asset_id[:80],
|
||||
get_allowed_local_dirs(),
|
||||
)
|
||||
return None
|
||||
logger.warning(
|
||||
"本地素材路径不在允许目录: %s allowed=%s",
|
||||
asset_id[:80],
|
||||
get_allowed_local_dirs(),
|
||||
)
|
||||
return None
|
||||
except (OSError, PathSecurityError):
|
||||
return None
|
||||
|
||||
# 2. 缓存命中(SHA256 hash 防路径遍历)
|
||||
cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16]
|
||||
safe_name = sanitize_filename(cache_hash)
|
||||
cached_path = work_dir / f"{safe_name}.mp4"
|
||||
if cached_path.exists() and cached_path.stat().st_size > 0:
|
||||
return cached_path
|
||||
|
||||
# 3. 从 OSS 下载(先标准化 key,防路径遍历注入)
|
||||
safe_key = self.normalize_storage_key(asset_id)
|
||||
if ".." in safe_key or safe_key.startswith("/"):
|
||||
logger.warning("asset_id 包含路径遍历模式,拒绝下载: %s", asset_id[:80])
|
||||
logger.warning("asset_id 含路径遍历: %s", asset_id[:80])
|
||||
return None
|
||||
|
||||
if self.download_asset(safe_key, cached_path):
|
||||
return cached_path
|
||||
|
||||
return None
|
||||
|
||||
def resolve_asset_ids_to_paths(
|
||||
@@ -519,22 +488,20 @@ class SharedStorageService(StoragePort):
|
||||
asset_ids: list[str],
|
||||
work_dir: str | Path,
|
||||
) -> dict[str, Path]:
|
||||
"""批量解析 asset_id → 本地路径。"""
|
||||
result: dict[str, Path] = {}
|
||||
for aid in asset_ids:
|
||||
local_path = self.resolve_asset_path(aid, work_dir)
|
||||
if local_path:
|
||||
result[aid] = local_path
|
||||
p = self.resolve_asset_path(aid, work_dir)
|
||||
if p:
|
||||
result[aid] = p
|
||||
return result
|
||||
|
||||
|
||||
# ── 单例管理 ────────────────────────────────────────────────────────────
|
||||
# ── 单例 ────────────────────────────────────────────────────────────────
|
||||
|
||||
_storage_service: Optional[SharedStorageService] = None
|
||||
|
||||
|
||||
def get_shared_storage_service() -> SharedStorageService:
|
||||
"""获取统一存储服务单例。"""
|
||||
global _storage_service
|
||||
if _storage_service is None:
|
||||
_storage_service = SharedStorageService()
|
||||
@@ -542,7 +509,6 @@ def get_shared_storage_service() -> SharedStorageService:
|
||||
return _storage_service
|
||||
|
||||
|
||||
# 向后兼容别名
|
||||
def get_storage_service() -> SharedStorageService:
|
||||
"""向后兼容:返回统一存储服务。"""
|
||||
"""向后兼容别名。"""
|
||||
return get_shared_storage_service()
|
||||
|
||||
+138
-115
@@ -1,6 +1,7 @@
|
||||
#!/bin/sh
|
||||
# ===========================================
|
||||
# Staging 部署脚本(SSH 模式,并行优化版)
|
||||
# worker 已收敛到 infra/docker/compose.yml 单一事实来源;api/web 暂保留 docker run。
|
||||
# ===========================================
|
||||
set -eu
|
||||
|
||||
@@ -48,6 +49,11 @@ GENERATED_DIR="${GENERATED_DIR:-/var/lib/xiaoxia-saas-staging/generated}"
|
||||
LEGACY_ASSETS_DIR="${LEGACY_ASSETS_DIR:-/var/lib/xiaoxia-saas-staging/legacy-assets}"
|
||||
NGINX_CONF_FILE="${NGINX_CONF_FILE:-/var/lib/xiaoxia-saas-staging/nginx-staging.conf}"
|
||||
COOKIES_FILE="${COOKIES_FILE:-/var/lib/xiaoxia-saas-staging/configs/douyin_cookies.txt}"
|
||||
INFRA_DOCKER_DIR="${INFRA_DOCKER_DIR:-/var/lib/xiaoxia-saas-staging/infra/docker}"
|
||||
COMPOSE_PROJECT="${COMPOSE_PROJECT:-xiaoxia-staging}"
|
||||
COMPOSE_ENV_VALUE="${COMPOSE_ENV_VALUE:-staging}"
|
||||
# COMPOSE_SYNC: CI workflow 已通过 scp 把 infra/docker/compose.yml 上传到服务器时设为 0 跳过同步
|
||||
COMPOSE_SYNC="${COMPOSE_SYNC:-1}"
|
||||
|
||||
SKIP_MIGRATION="${SKIP_MIGRATION:-false}"
|
||||
SKIP_ROLLBACK="${SKIP_ROLLBACK:-false}"
|
||||
@@ -57,8 +63,6 @@ if [ -z "$IMAGE_TAG" ]; then
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# .env 文件由 CI 从模板 + Secrets 渲染后通过 SCP 上传到服务器
|
||||
# 如果文件不存在,说明 CI 渲染步骤失败或未执行
|
||||
if [ ! -f "$ENV_FILE" ]; then
|
||||
echo "ERROR: $ENV_FILE 不存在。CI 应先在 render_env 步骤渲染并上传此文件"
|
||||
exit 1
|
||||
@@ -67,7 +71,7 @@ echo "✅ .env file found: $ENV_FILE ($(wc -l < "$ENV_FILE") lines)"
|
||||
mkdir -p "$GENERATED_DIR"
|
||||
mkdir -p "$LEGACY_ASSETS_DIR"
|
||||
mkdir -p "$(dirname "$COOKIES_FILE")"
|
||||
# 抖音 cookies 文件:CI workflow 已通过 scp 上传;如果不存在(非 CI 环境)则创建占位
|
||||
mkdir -p "$INFRA_DOCKER_DIR"
|
||||
if [ ! -f "$COOKIES_FILE" ] || [ "$(wc -c < "$COOKIES_FILE" 2>/dev/null || echo 0)" -lt 200 ]; then
|
||||
printf '# Netscape HTTP Cookie File\n# 抖音 cookies 占位(CI 应通过 scp 上传真实 cookies)\n' > "$COOKIES_FILE"
|
||||
echo "WARNING: Douyin cookies not found or too small at $COOKIES_FILE (extraction will 503)"
|
||||
@@ -76,7 +80,6 @@ else
|
||||
fi
|
||||
|
||||
# ── 写入 Staging Nginx 配置 ──
|
||||
# 运行时覆盖 nginx 配置,确保 upstream 指向正确的 staging 网络
|
||||
echo "Writing staging nginx config..."
|
||||
cat > "$NGINX_CONF_FILE" << 'NGINX_EOF'
|
||||
server {
|
||||
@@ -122,8 +125,26 @@ server {
|
||||
NGINX_EOF
|
||||
echo "✅ Nginx config written: $NGINX_CONF_FILE"
|
||||
|
||||
# ── 确认 infra/docker/compose.yml 存在 ──
|
||||
# CI workflow 在执行本脚本前已通过 scp 把 infra/docker/compose.yml 上传到 $INFRA_DOCKER_DIR
|
||||
# (workflow 里做:scp infra/docker/compose.yml <host>:$INFRA_DOCKER_DIR/compose.yml)。
|
||||
# 这里只做存在性检查 + nginx 软链;不再 curl 私有仓库(SSH 环境无 Gitea token)。
|
||||
COMPOSE_FILE_PATH="$INFRA_DOCKER_DIR/compose.yml"
|
||||
if [ ! -f "$COMPOSE_FILE_PATH" ]; then
|
||||
echo "ERROR: $COMPOSE_FILE_PATH 不存在。CI workflow 应先 scp infra/docker/compose.yml 到服务器"
|
||||
exit 1
|
||||
fi
|
||||
echo "✅ compose.yml ready: $COMPOSE_FILE_PATH ($(wc -l < "$COMPOSE_FILE_PATH") lines)"
|
||||
ln -sf "$NGINX_CONF_FILE" "$INFRA_DOCKER_DIR/nginx-${COMPOSE_ENV_VALUE}.conf" 2>/dev/null || true
|
||||
|
||||
# 封装 docker compose 调用:统一 --env-file(compose 默认只读取 project 目录下的 .env,
|
||||
# 我们的 .env 在 $INFRA_DOCKER_DIR/../../.env,必须显式传入才能读到 GENERATED_FILES_HOST_DIR 等变量)
|
||||
compose() {
|
||||
(cd "$INFRA_DOCKER_DIR" && docker compose --env-file "$ENV_FILE" -p "$COMPOSE_PROJECT" "$@")
|
||||
}
|
||||
|
||||
echo "==========================================="
|
||||
echo " Staging 部署 - $IMAGE_TAG (并行优化版)"
|
||||
echo " Staging 部署 - $IMAGE_TAG"
|
||||
echo "==========================================="
|
||||
|
||||
echo "Recording current image versions for rollback..."
|
||||
@@ -144,6 +165,7 @@ for c in xiaoxia-api-staging xiaoxia-worker-staging xiaoxia-web-staging; do
|
||||
fi
|
||||
done
|
||||
|
||||
# ── 回滚函数 ──
|
||||
rollback() {
|
||||
echo ""
|
||||
echo "!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!"
|
||||
@@ -157,9 +179,37 @@ rollback() {
|
||||
fi
|
||||
|
||||
echo "Stopping new containers..."
|
||||
docker rm -f xiaoxia-api-staging 2>/dev/null || true
|
||||
docker rm -f xiaoxia-worker-staging 2>/dev/null || true
|
||||
docker rm -f xiaoxia-web-staging 2>/dev/null || true
|
||||
docker rm -f xiaoxia-api-staging xiaoxia-web-staging 2>/dev/null || true
|
||||
if [ -n "$PREV_WORKER_IMAGE" ]; then
|
||||
echo "Rolling back Worker to: $PREV_WORKER_IMAGE (via compose)"
|
||||
compose up -d --no-deps worker 2>&1 || echo "WARN: compose rollback failed, fallback to docker run"
|
||||
# 镜像通过 env 注入:compose 默认读 WORKER_IMAGE(未设则用 :dev),这里用临时 env 覆盖
|
||||
if ! docker inspect xiaoxia-worker-staging >/dev/null 2>&1; then
|
||||
echo "Fallback: docker run previous worker image"
|
||||
docker run -d \
|
||||
--name xiaoxia-worker-staging \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-e APP_ENV="$COMPOSE_ENV_VALUE" \
|
||||
-e APP_VERSION="$(echo "$PREV_WORKER_IMAGE" | grep -oE '[^:]+$')" \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
|
||||
-e TRANSCODE_CONCURRENCY="${TRANSCODE_CONCURRENCY:-2}" \
|
||||
-e BEAT_ENABLED=1 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
|
||||
-v "$GENERATED_DIR:/app/generated" \
|
||||
--restart unless-stopped \
|
||||
--health-cmd "grep -q 'celery.*worker' /proc/[0-9]*/cmdline 2>/dev/null || exit 1" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
--health-start-period 40s \
|
||||
--log-driver json-file --log-opt max-size=50m --log-opt max-file=3 \
|
||||
"$PREV_WORKER_IMAGE" || true
|
||||
fi
|
||||
fi
|
||||
|
||||
LOG_OPTS="--log-driver json-file --log-opt max-size=50m --log-opt max-file=3"
|
||||
|
||||
@@ -168,10 +218,10 @@ rollback() {
|
||||
docker run -d \
|
||||
--name xiaoxia-api-staging \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network xiaoxia-net-staging \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-p 127.0.0.1:8000:8000 \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_VERSION="$(echo $PREV_API_IMAGE | grep -oE '[^:]+$')" \
|
||||
-e APP_ENV="$COMPOSE_ENV_VALUE" \
|
||||
-e APP_VERSION="$(echo "$PREV_API_IMAGE" | grep -oE '[^:]+$')" \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
|
||||
@@ -183,31 +233,7 @@ rollback() {
|
||||
--health-retries 3 \
|
||||
--health-start-period 40s \
|
||||
$LOG_OPTS \
|
||||
"$PREV_API_IMAGE" &
|
||||
fi
|
||||
|
||||
if [ -n "$PREV_WORKER_IMAGE" ]; then
|
||||
echo "Rolling back Worker to: $PREV_WORKER_IMAGE"
|
||||
docker run -d \
|
||||
--name xiaoxia-worker-staging \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network xiaoxia-net-staging \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_VERSION="$(echo $PREV_WORKER_IMAGE | grep -oE '[^:]+$')" \
|
||||
-e WORKER_CONCURRENCY=1 \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
|
||||
-v "$GENERATED_DIR:/app/generated" \
|
||||
--restart unless-stopped \
|
||||
--health-cmd "grep -lq celery /proc/[0-9]*/cmdline 2>/dev/null || exit 1" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
--health-start-period 30s \
|
||||
$LOG_OPTS \
|
||||
"$PREV_WORKER_IMAGE" &
|
||||
"$PREV_API_IMAGE" || true
|
||||
fi
|
||||
|
||||
if [ -n "$PREV_WEB_IMAGE" ]; then
|
||||
@@ -218,7 +244,7 @@ rollback() {
|
||||
fi
|
||||
docker run -d \
|
||||
--name xiaoxia-web-staging \
|
||||
--network xiaoxia-net-staging \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-p 127.0.0.1:3001:80 \
|
||||
--restart unless-stopped \
|
||||
$LEGACY_VOLUME \
|
||||
@@ -228,27 +254,25 @@ rollback() {
|
||||
--health-timeout 5s \
|
||||
--health-retries 3 \
|
||||
$LOG_OPTS \
|
||||
"$PREV_WEB_IMAGE" &
|
||||
"$PREV_WEB_IMAGE" || true
|
||||
fi
|
||||
|
||||
wait
|
||||
sleep 3
|
||||
|
||||
if [ -n "$PREV_API_IMAGE" ]; then
|
||||
echo "Waiting for rolled-back API to become healthy..."
|
||||
i=0
|
||||
while [ "$i" -lt 40 ]; do
|
||||
if curl -sf --max-time 5 http://127.0.0.1:8000/health >/dev/null 2>&1; then
|
||||
echo "Rolled-back API is healthy!"
|
||||
break
|
||||
fi
|
||||
i=$((i + 1))
|
||||
echo " Waiting... ($i/40)"
|
||||
sleep 3
|
||||
done
|
||||
if [ "$i" -ge 40 ]; then
|
||||
echo "WARN: Rolled-back API did not become healthy within 120s"
|
||||
docker logs --tail 30 xiaoxia-api-staging
|
||||
echo "Waiting for rolled-back API to become healthy..."
|
||||
i=0
|
||||
while [ "$i" -lt 40 ]; do
|
||||
if curl -sf --max-time 5 http://127.0.0.1:8000/health >/dev/null 2>&1; then
|
||||
echo "Rolled-back API is healthy!"
|
||||
break
|
||||
fi
|
||||
i=$((i + 1))
|
||||
echo " Waiting... ($i/40)"
|
||||
sleep 3
|
||||
done
|
||||
if [ "$i" -ge 40 ]; then
|
||||
echo "WARN: Rolled-back API did not become healthy within 120s"
|
||||
docker logs --tail 30 xiaoxia-api-staging 2>/dev/null || true
|
||||
fi
|
||||
|
||||
echo ""
|
||||
@@ -260,7 +284,7 @@ rollback() {
|
||||
echo "Previous Web: ${PREV_WEB_IMAGE:-none}"
|
||||
echo ""
|
||||
echo "部署失败,已自动回滚到上一版本"
|
||||
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging
|
||||
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging || true
|
||||
exit 1
|
||||
}
|
||||
|
||||
@@ -272,10 +296,6 @@ if [ -n "$REGISTRY_TOKEN" ]; then
|
||||
retry_docker_login
|
||||
fi
|
||||
|
||||
# ---- 并行 Pull 三个镜像 ----
|
||||
# 注意:这里必须使用 IMAGE_TAG(commit SHA)做确定性部署,不要改成 :dev。
|
||||
# :dev 是 floating tag,可能被并发构建覆盖,导致部署版本不可重现、回滚混乱。
|
||||
# Watchtower 可监听 :dev 做非关键路径的自动同步;正式部署/回滚一律锚定 SHA。
|
||||
REGISTRY_API="${REGISTRY}/xiaoxia-saas-api:${IMAGE_TAG}"
|
||||
REGISTRY_WORKER="${REGISTRY}/xiaoxia-saas-worker:${IMAGE_TAG}"
|
||||
REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:${IMAGE_TAG}"
|
||||
@@ -304,7 +324,6 @@ for svc in api worker web; do
|
||||
elif grep -qE "Digest:|Status: Downloaded" "$PULL_LOG_DIR/$svc.log" 2>/dev/null; then
|
||||
echo " OK $svc"
|
||||
else
|
||||
# 检查docker pull返回值不直接,用镜像是否存在来判断
|
||||
img_var="REGISTRY_$(echo $svc | tr '[:lower:]' '[:upper:]')"
|
||||
img_val=$(eval echo "\$$img_var")
|
||||
if docker image inspect "$img_val" >/dev/null 2>&1; then
|
||||
@@ -327,7 +346,7 @@ fi
|
||||
|
||||
echo "All images pulled."
|
||||
|
||||
# ====== 镜像内容校验(CI 加固 - 防止静默部署损坏/过期镜像) ======
|
||||
# ====== 镜像内容校验 ======
|
||||
echo ""
|
||||
echo "=========================================="
|
||||
echo " 镜像内容校验"
|
||||
@@ -336,7 +355,6 @@ echo "=========================================="
|
||||
VERIFY_FAILED=0
|
||||
DEPLOY_MANIFEST="${GENERATED_DIR}/deploy-manifest.json"
|
||||
|
||||
# 读取上次部署的 manifest(用于对比)
|
||||
PREV_MANIFEST=""
|
||||
if [ -f "$DEPLOY_MANIFEST" ]; then
|
||||
PREV_MANIFEST=$(cat "$DEPLOY_MANIFEST")
|
||||
@@ -348,14 +366,12 @@ for svc in api worker web; do
|
||||
img_var="REGISTRY_$(echo $svc | tr '[:lower:]' '[:upper:]')"
|
||||
img_val=$(eval echo "\$$img_var")
|
||||
|
||||
# 1. 检查镜像是否存在
|
||||
if ! docker image inspect "$img_val" >/dev/null 2>&1; then
|
||||
echo " ❌ $svc: 镜像不存在 ($img_val)"
|
||||
VERIFY_FAILED=$((VERIFY_FAILED + 1))
|
||||
continue
|
||||
fi
|
||||
|
||||
# 2. 检查 layers 有效性
|
||||
LAYER_COUNT=$(docker inspect --format='{{len .RootFS.Layers}}' "$img_val" 2>/dev/null || echo "0")
|
||||
if [ "$LAYER_COUNT" -eq 0 ]; then
|
||||
echo " ❌ $svc: 镜像无有效 layers ($img_val)"
|
||||
@@ -363,14 +379,12 @@ for svc in api worker web; do
|
||||
continue
|
||||
fi
|
||||
|
||||
# 3. 获取 digest 和创建时间
|
||||
IMG_ID=$(docker inspect --format='{{.Id}}' "$img_val")
|
||||
IMG_CREATED=$(docker inspect --format='{{.Created}}' "$img_val")
|
||||
IMG_SIZE=$(docker inspect --format='{{.Size}}' "$img_val")
|
||||
echo " ✅ $svc: ${LAYER_COUNT} layers, size=${IMG_SIZE}, created=${IMG_CREATED}"
|
||||
echo " id: $IMG_ID"
|
||||
|
||||
# 4. 对比上次部署
|
||||
CHANGED="unchanged"
|
||||
if [ -n "$PREV_MANIFEST" ]; then
|
||||
PREV_ID=$(echo "$PREV_MANIFEST" | grep "\"${svc}_id\"" | sed 's/.*: *"\(.*\)".*/\1/' 2>/dev/null || echo "")
|
||||
@@ -398,7 +412,6 @@ if [ "$VERIFY_FAILED" -gt 0 ]; then
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 写入新 manifest
|
||||
cat > "$DEPLOY_MANIFEST" <<MANIFEST_EOF
|
||||
{
|
||||
"deployed_at": "$(date -u +%Y-%m-%dT%H:%M:%SZ)",
|
||||
@@ -444,14 +457,14 @@ for c in xiaoxia-postgres-staging xiaoxia-redis-staging; do
|
||||
fi
|
||||
done
|
||||
|
||||
docker network create xiaoxia-net-staging 2>/dev/null || true
|
||||
docker network create "xiaoxia-net-${COMPOSE_ENV_VALUE}" 2>/dev/null || true
|
||||
|
||||
if [ "$SKIP_MIGRATION" != "true" ]; then
|
||||
echo "Running database migrations..."
|
||||
docker run --rm \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network xiaoxia-net-staging \
|
||||
-e APP_ENV=staging \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-e APP_ENV="$COMPOSE_ENV_VALUE" \
|
||||
"$REGISTRY_API" sh -c "cd /app && alembic upgrade head" || {
|
||||
echo "ERROR: Database migration failed"
|
||||
exit 1
|
||||
@@ -462,29 +475,28 @@ else
|
||||
fi
|
||||
|
||||
echo "Stopping old containers..."
|
||||
# 优雅关闭:先 stop(发 SIGTERM,等待),再 rm
|
||||
# Worker 需要更长时间(视频任务最长可能5分钟)
|
||||
docker stop -t 120 xiaoxia-worker-staging 2>/dev/null || true
|
||||
docker stop -t 30 xiaoxia-api-staging 2>/dev/null || true
|
||||
docker stop -t 10 xiaoxia-web-staging 2>/dev/null || true
|
||||
docker rm xiaoxia-worker-staging xiaoxia-api-staging xiaoxia-web-staging 2>/dev/null || true
|
||||
docker stop -t 30 xiaoxia-api-staging 2>/dev/null || true
|
||||
docker stop -t 120 xiaoxia-worker-staging 2>/dev/null || true
|
||||
docker rm -f xiaoxia-api-staging xiaoxia-web-staging 2>/dev/null || true
|
||||
docker rm -f xiaoxia-worker-staging 2>/dev/null || true
|
||||
|
||||
LOG_OPTS="--log-driver json-file --log-opt max-size=50m --log-opt max-file=3"
|
||||
|
||||
# ---- 并行启动三个容器 ----
|
||||
echo "Starting all containers (parallel)..."
|
||||
echo "Starting all containers..."
|
||||
|
||||
LEGACY_VOLUME=""
|
||||
if [ -d "$LEGACY_ASSETS_DIR" ] && [ "$(ls -A "$LEGACY_ASSETS_DIR" 2>/dev/null)" ]; then
|
||||
LEGACY_VOLUME="-v ${LEGACY_ASSETS_DIR}:/usr/share/nginx/html/assets-legacy/assets:ro"
|
||||
fi
|
||||
|
||||
# ── API: 暂保留 docker run(TODO: 后续收敛到 compose)──
|
||||
docker run -d \
|
||||
--name xiaoxia-api-staging \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network xiaoxia-net-staging \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-p 127.0.0.1:8000:8000 \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_ENV="$COMPOSE_ENV_VALUE" \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
@@ -502,31 +514,18 @@ docker run -d \
|
||||
"$REGISTRY_API" &
|
||||
PID_API_START=$!
|
||||
|
||||
docker run -d \
|
||||
--name xiaoxia-worker-staging \
|
||||
--env-file "$ENV_FILE" \
|
||||
--network xiaoxia-net-staging \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e WORKER_CONCURRENCY=1 \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
|
||||
-v "$GENERATED_DIR:/app/generated" \
|
||||
--restart unless-stopped \
|
||||
--health-cmd "grep -lq celery /proc/[0-9]*/cmdline 2>/dev/null || exit 1" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
--health-start-period 30s \
|
||||
$LOG_OPTS \
|
||||
"$REGISTRY_WORKER" &
|
||||
# ── Worker: 通过 compose 启动(单一事实来源)──
|
||||
# compose.yml 定义:三进程(beat+generation+transcode)、独立并发、BEAT_ENABLED、
|
||||
# healthcheck 匹配 'celery.*worker'(不把 beat 算活)、资源限制 4C/8G。
|
||||
# WORKER_IMAGE 通过环境变量覆盖镜像 tag(compose.yml 默认 :dev)。
|
||||
echo "Starting worker via docker compose (from $INFRA_DOCKER_DIR)..."
|
||||
WORKER_IMAGE="$REGISTRY_WORKER" APP_VERSION="$IMAGE_TAG" compose up -d --no-deps worker &
|
||||
PID_WORKER_START=$!
|
||||
|
||||
# ── Web: 暂保留 docker run(TODO: 后续收敛到 compose)──
|
||||
docker run -d \
|
||||
--name xiaoxia-web-staging \
|
||||
--network xiaoxia-net-staging \
|
||||
--network "xiaoxia-net-${COMPOSE_ENV_VALUE}" \
|
||||
-p 127.0.0.1:3001:80 \
|
||||
--restart unless-stopped \
|
||||
$LEGACY_VOLUME \
|
||||
@@ -563,9 +562,8 @@ if [ "$START_FAILED" -gt 0 ]; then
|
||||
rollback
|
||||
fi
|
||||
|
||||
# ---- 并行等待 API 和 Web 健康 ----
|
||||
echo ""
|
||||
echo "Waiting for API + Web health (parallel)..."
|
||||
echo "Waiting for all services health (parallel)..."
|
||||
|
||||
HEALTH_LOG_DIR="/tmp/staging-health-$$"
|
||||
mkdir -p "$HEALTH_LOG_DIR"
|
||||
@@ -600,36 +598,60 @@ PID_API_HEALTH=$!
|
||||
) > "$HEALTH_LOG_DIR/web.log" 2>&1 &
|
||||
PID_WEB_HEALTH=$!
|
||||
|
||||
(
|
||||
i=0
|
||||
while [ "$i" -lt 20 ]; do
|
||||
hc=$(docker inspect -f '{{if .State.Health}}{{.State.Health.Status}}{{else}}{{.State.Status}}{{end}}' xiaoxia-worker-staging 2>/dev/null || echo "missing")
|
||||
if [ "$hc" = "healthy" ]; then
|
||||
echo "Worker healthy after $((i * 3))s"
|
||||
exit 0
|
||||
fi
|
||||
if [ "$hc" = "unhealthy" ]; then
|
||||
echo "Worker UNHEALTHY after $((i * 3))s"
|
||||
docker logs --tail 30 xiaoxia-worker-staging 2>/dev/null || true
|
||||
exit 1
|
||||
fi
|
||||
i=$((i + 1))
|
||||
sleep 3
|
||||
done
|
||||
echo "Worker health unknown after 60s (last: $hc)"
|
||||
exit 1
|
||||
) > "$HEALTH_LOG_DIR/worker.log" 2>&1 &
|
||||
PID_WORKER_HEALTH=$!
|
||||
|
||||
set +e
|
||||
wait $PID_API_HEALTH
|
||||
API_EXIT=$?
|
||||
wait $PID_WEB_HEALTH
|
||||
WEB_EXIT=$?
|
||||
wait $PID_WORKER_HEALTH
|
||||
WORKER_EXIT=$?
|
||||
set -e
|
||||
|
||||
echo ""
|
||||
echo "健康检查结果:"
|
||||
API_OK=0
|
||||
WEB_OK=0
|
||||
if [ "$API_EXIT" -eq 0 ]; then
|
||||
echo " OK API: $(cat "$HEALTH_LOG_DIR/api.log")"
|
||||
API_OK=1
|
||||
echo " OK API: $(cat "$HEALTH_LOG_DIR/api.log")"
|
||||
else
|
||||
echo " FAIL API: 120s未就绪"
|
||||
docker logs --tail 50 xiaoxia-api-staging
|
||||
echo " FAIL API: 120s未就绪"
|
||||
docker logs --tail 50 xiaoxia-api-staging 2>/dev/null || true
|
||||
fi
|
||||
|
||||
if [ "$WEB_EXIT" -eq 0 ]; then
|
||||
echo " OK Web: $(cat "$HEALTH_LOG_DIR/web.log")"
|
||||
WEB_OK=1
|
||||
echo " OK Web: $(cat "$HEALTH_LOG_DIR/web.log")"
|
||||
else
|
||||
echo " FAIL Web: 30s未就绪"
|
||||
docker logs --tail 30 xiaoxia-web-staging
|
||||
echo " FAIL Web: 30s未就绪"
|
||||
docker logs --tail 30 xiaoxia-web-staging 2>/dev/null || true
|
||||
fi
|
||||
if [ "$WORKER_EXIT" -eq 0 ]; then
|
||||
echo " OK Worker: $(cat "$HEALTH_LOG_DIR/worker.log")"
|
||||
else
|
||||
echo " FAIL Worker: $(cat "$HEALTH_LOG_DIR/worker.log")"
|
||||
docker logs --tail 50 xiaoxia-worker-staging 2>/dev/null || true
|
||||
fi
|
||||
|
||||
rm -rf "$HEALTH_LOG_DIR"
|
||||
|
||||
if [ "$API_OK" -eq 0 ] || [ "$WEB_OK" -eq 0 ]; then
|
||||
if [ "$API_EXIT" -ne 0 ] || [ "$WEB_EXIT" -ne 0 ] || [ "$WORKER_EXIT" -ne 0 ]; then
|
||||
echo ""
|
||||
echo "ERROR: 健康检查失败"
|
||||
rollback
|
||||
@@ -640,8 +662,9 @@ docker image prune -af --filter "until=168h" 2>/dev/null || true
|
||||
docker builder prune -af --filter "until=168h" 2>/dev/null || true
|
||||
|
||||
echo ""
|
||||
echo "=== Staging deployment complete (并行优化版) ==="
|
||||
echo "=== Staging deployment complete ==="
|
||||
echo "API: http://127.0.0.1:8000"
|
||||
echo "Web: http://127.0.0.1:3001"
|
||||
echo "Worker: managed by docker compose (project=$COMPOSE_PROJECT)"
|
||||
echo "Version: $IMAGE_TAG"
|
||||
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging
|
||||
|
||||
Executable
+154
@@ -0,0 +1,154 @@
|
||||
#!/usr/bin/env python3
|
||||
"""一键上传预设 BGM 到 OSS 并回填 preset_bgm.py audio_url。
|
||||
|
||||
用法(在服务器或本地有 OSS 凭证的机器上执行):
|
||||
1. 把 mp3 文件放到 ./bgm_assets/ (created by you) 目录下,文件名按 {preset_id}.mp3 命名:
|
||||
bgm_upbeat_001.mp3 阳光清晨
|
||||
bgm_upbeat_002.mp3 活力节拍
|
||||
bgm_upbeat_003.mp3 夏日漫步
|
||||
bgm_relax_001.mp3 静谧时光
|
||||
bgm_relax_002.mp3 雨后森林
|
||||
bgm_relax_003.mp3 月光奏鸣曲
|
||||
bgm_tech_001.mp3 未来科技
|
||||
bgm_tech_002.mp3 数据脉冲
|
||||
bgm_commerce_001.mp3 心动时刻
|
||||
bgm_commerce_002.mp3 品质生活
|
||||
2. 确保环境变量已设置:
|
||||
OSS_ENDPOINT, OSS_ACCESS_KEY_ID, OSS_ACCESS_KEY_SECRET, OSS_BUCKET_NAME
|
||||
3. 运行:python scripts/upload_preset_bgm.py
|
||||
4. 脚本会上传到 OSS 路径 preset/bgm/<id>.mp3,并自动改写
|
||||
packages/domain/preset_bgm.py 填入 public URL。
|
||||
5. git commit & push 即可。
|
||||
|
||||
支持可选参数:
|
||||
--bgm-dir DIR 本地 mp3 目录(默认 ./bgm_assets)
|
||||
--oss-prefix PFX OSS key 前缀(默认 preset/bgm/)
|
||||
--dry-run 只打印要做的操作,不上传不改写
|
||||
--public 上传后设置公共读 ACL(默认开启,safe_download 走公网 URL)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
REQUIRED_ENV = ("OSS_ENDPOINT", "OSS_ACCESS_KEY_ID", "OSS_ACCESS_KEY_SECRET", "OSS_BUCKET_NAME")
|
||||
PRESET_BGM_PY = Path(__file__).resolve().parents[1] / "packages" / "domain" / "preset_bgm.py"
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--bgm-dir", default="./bgm_assets")
|
||||
ap.add_argument("--oss-prefix", default="preset/bgm/")
|
||||
ap.add_argument("--dry-run", action="store_true")
|
||||
ap.add_argument("--public", action="store_true", default=True)
|
||||
args = ap.parse_args()
|
||||
|
||||
bgm_dir = Path(args.bgm_dir)
|
||||
if not bgm_dir.exists():
|
||||
print(f"[ERR] 目录不存在: {bgm_dir}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
# 加载 preset_bgm.py 中所有 ID(简单 AST 抽取)
|
||||
sys.path.insert(0, str(PRESET_BGM_PY.parents[2]))
|
||||
from packages.domain.preset_bgm import PRESET_BGM_LIBRARY # type: ignore
|
||||
|
||||
id_to_preset = {p.id: p for p in PRESET_BGM_LIBRARY}
|
||||
|
||||
# 枚举本地文件
|
||||
local_files: dict[str, Path] = {}
|
||||
for ext in (".mp3", ".m4a", ".aac", ".wav", ".ogg"):
|
||||
for f in bgm_dir.glob(f"*{ext}"):
|
||||
pid = f.stem
|
||||
local_files[pid] = f
|
||||
|
||||
missing_local = [pid for pid in id_to_preset if pid not in local_files]
|
||||
unknown_local = [pid for pid in local_files if pid not in id_to_preset]
|
||||
|
||||
print(f"[INFO] 预设 BGM 总数: {len(PRESET_BGM_LIBRARY)}")
|
||||
print(f"[INFO] 本地找到对应文件: {len(local_files) - len(unknown_local)}")
|
||||
if missing_local:
|
||||
print(f"[WARN] 缺少本地音频文件的预设({len(missing_local)}):")
|
||||
for pid in missing_local:
|
||||
print(f" {pid} {id_to_preset[pid].name}")
|
||||
if unknown_local:
|
||||
print(f"[WARN] 本地文件未匹配任何 preset_id({len(unknown_local)}): {unknown_local}")
|
||||
|
||||
if args.dry_run:
|
||||
for pid, f in local_files.items():
|
||||
if pid in id_to_preset:
|
||||
url = f"https://{os.environ.get('OSS_BUCKET_NAME', '<bucket>')}.{os.environ.get('OSS_ENDPOINT', '<ep>')}/{args.oss_prefix}{f.name}"
|
||||
print(f"[DRY] would upload {f} → {args.oss_prefix}{f.name} → {url}")
|
||||
return 0
|
||||
|
||||
# 凭证检查
|
||||
for k in REQUIRED_ENV:
|
||||
if not os.environ.get(k):
|
||||
print(f"[ERR] 缺少环境变量 {k}", file=sys.stderr)
|
||||
return 3
|
||||
|
||||
try:
|
||||
import oss2 # type: ignore
|
||||
except ImportError:
|
||||
print("[ERR] 需要 oss2: pip install oss2", file=sys.stderr)
|
||||
return 4
|
||||
|
||||
endpoint = os.environ["OSS_ENDPOINT"]
|
||||
bucket_name = os.environ["OSS_BUCKET_NAME"]
|
||||
auth = oss2.Auth(os.environ["OSS_ACCESS_KEY_ID"], os.environ["OSS_ACCESS_KEY_SECRET"])
|
||||
bucket = oss2.Bucket(auth, f"https://{endpoint}", bucket_name)
|
||||
public_base = f"https://{bucket_name}.{endpoint}"
|
||||
|
||||
pid_to_url: dict[str, str] = {}
|
||||
for pid, f in local_files.items():
|
||||
if pid not in id_to_preset:
|
||||
continue
|
||||
key = f"{args.oss_prefix.rstrip('/')}/{f.name}"
|
||||
print(f"[UPLOAD] {f} → oss://{bucket_name}/{key}")
|
||||
headers = {"Content-Type": "audio/mpeg" if f.suffix == ".mp3" else "audio/mp4"}
|
||||
if args.public:
|
||||
bucket.put_object_from_file(key, str(f), headers=headers)
|
||||
bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
|
||||
else:
|
||||
bucket.put_object_from_file(key, str(f), headers=headers)
|
||||
url = f"{public_base}/{key}"
|
||||
pid_to_url[pid] = url
|
||||
print(f" → {url}")
|
||||
|
||||
if not pid_to_url:
|
||||
print("[WARN] 没有文件被上传,不修改 preset_bgm.py")
|
||||
return 0
|
||||
|
||||
# 回填 preset_bgm.py(按 id 精确替换 audio_url="" 为 audio_url="<url>")
|
||||
text = PRESET_BGM_PY.read_text(encoding="utf-8")
|
||||
orig = text
|
||||
for pid, url in pid_to_url.items():
|
||||
# 匹配 PresetBGM( ... id="pid", ... audio_url="", ... )
|
||||
# 简单替换:找到 id="pid" 行开始的 PresetBGM 构造块,把块内的 audio_url="" 替换
|
||||
import re
|
||||
|
||||
block_pat = re.compile(
|
||||
r'(PresetBGM\([^)]*?id="' + re.escape(pid) + r'"[^)]*?audio_url=)"[^"]*"',
|
||||
re.DOTALL,
|
||||
)
|
||||
new_text, n = block_pat.subn(rf'\1"{url}"', text, count=1)
|
||||
if n == 0:
|
||||
# 兜底:可能 audio_url 后面直接是 ),没赋值?不会,dataclass 有默认值
|
||||
print(f"[WARN] 未在 PresetBGM 块中找到 id={pid} 的 audio_url 字段,跳过回填")
|
||||
continue
|
||||
text = new_text
|
||||
if text != orig:
|
||||
PRESET_BGM_PY.write_text(text, encoding="utf-8")
|
||||
print(f"[OK] 已回填 {len(pid_to_url)} 条 audio_url 到 {PRESET_BGM_PY}")
|
||||
print(
|
||||
"[NEXT] git diff packages/domain/preset_bgm.py && git add -A && git commit -m 'feat(domain): fill preset BGM audio_urls' && git push"
|
||||
)
|
||||
else:
|
||||
print("[WARN] preset_bgm.py 未发生变化(可能已填过)")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -322,22 +322,41 @@ class TestBatchPreviewRoute:
|
||||
)
|
||||
assert captured[0].voice_library_id == "legacy_voice"
|
||||
|
||||
def test_preview_queue_limit_checks_total_count(self):
|
||||
"""限流预检查按变体总数计:用户 pending + N 超限 → 429"""
|
||||
def test_preview_queue_limit_global_returns_503(self):
|
||||
"""#2098:仅全局硬上限仍 503 拒绝;用户超限自动排队不返回 429。"""
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
from fastapi import HTTPException
|
||||
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 3
|
||||
repo.count_pending_total.return_value = 0
|
||||
# 全局硬上限:预检查即 503,不会进入后续流程
|
||||
repo_global = MagicMock()
|
||||
repo_global.count_pending_total.return_value = 100
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
_make_preview_request(preview_count=5),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
generation_task_repository=repo_global,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert exc.value.status_code == 429
|
||||
assert exc.value.status_code == 503
|
||||
|
||||
def test_preview_user_over_limit_does_not_429(self):
|
||||
"""#2098:用户 pending 远超软上限也不返回 429/4xx,请求进入业务流程后因 mock 不足走 500。"""
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
from fastapi import HTTPException
|
||||
|
||||
repo_user = MagicMock()
|
||||
repo_user.count_pending_by_user.return_value = 100 # 远超用户软上限
|
||||
repo_user.count_pending_total.return_value = 0
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
_make_preview_request(preview_count=1),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo_user,
|
||||
db=MagicMock(),
|
||||
)
|
||||
# 关键断言:不是 429(也不是 503,因为全局未超限),说明预检查放过了请求
|
||||
assert exc.value.status_code != 429
|
||||
assert exc.value.status_code != 503
|
||||
|
||||
def test_preview_reselect_failure_marks_all_failed(self):
|
||||
"""#1743:变体独立选片(reselect)重试仍失败 → 已创建任务全部标记 failed 并 500"""
|
||||
|
||||
@@ -741,22 +741,23 @@ class TestCreatePreviewRoute:
|
||||
assert resp.items[0].status == "pending"
|
||||
assert resp.items[0].variant_index == 0
|
||||
|
||||
def test_user_pending_limit_exceeded(self):
|
||||
"""用户待处理任务超限 → 429"""
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 3
|
||||
repo.count_pending_total.return_value = 5
|
||||
|
||||
def test_user_pending_limit_no_longer_rejects(self):
|
||||
"""#2098: 用户待处理任务超限不再 429 拒绝(预检查仅全局 503)。"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 100
|
||||
repo.count_pending_total.return_value = 5
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert exc_info.value.status_code == 429
|
||||
# 不是 429(用户级硬拒已移除)也不是 503(全局未超限)
|
||||
assert exc.value.status_code != 429
|
||||
assert exc.value.status_code != 503
|
||||
|
||||
def test_global_queue_full(self):
|
||||
"""全局队列满 → 503"""
|
||||
@@ -844,36 +845,32 @@ class TestCreatePreviewRoute:
|
||||
assert resp.total == 1
|
||||
assert resp.items[0].status == "failed"
|
||||
|
||||
def test_enqueue_raises_user_limit(self):
|
||||
"""safe_enqueue 抛出 UserPendingLimitExceeded → 429"""
|
||||
def test_enqueue_user_limit_exception_no_longer_returns_429(self):
|
||||
"""#2098: 即使 safe_enqueue 模拟抛 UserPendingLimitExceeded,也不再触发 429 整体拒绝,
|
||||
变体被标记 failed 后正常返回响应。"""
|
||||
repo = MagicMock()
|
||||
repo.count_pending_by_user.return_value = 0
|
||||
repo.count_pending_total.return_value = 0
|
||||
|
||||
task = _make_task()
|
||||
from fastapi import HTTPException
|
||||
|
||||
def _set_failed_limit(error_message="", **_kwargs):
|
||||
def _set_failed(reason="", **_kw):
|
||||
task.status = GenerationTaskStatus.FAILED
|
||||
task.error_message = error_message or "待处理任务超限"
|
||||
task.error_message = reason
|
||||
|
||||
task.mark_failed.side_effect = _set_failed_limit
|
||||
|
||||
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
|
||||
MockUC.return_value.execute.return_value = task
|
||||
with patch(
|
||||
"app.api.routes.generation_preview.safe_enqueue_generation_task",
|
||||
side_effect=UserPendingLimitExceeded(user_id="u1", pending_count=4, limit=3),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
# 全部变体入队失败且错误消息含"待处理任务" → 429
|
||||
assert exc_info.value.status_code == 429
|
||||
task.mark_failed.side_effect = _set_failed
|
||||
repo.create.return_value = task
|
||||
with patch(
|
||||
"app.api.routes.generation_preview.safe_enqueue_generation_task",
|
||||
side_effect=UserPendingLimitExceeded(user_id="u1", pending_count=4, limit=20),
|
||||
):
|
||||
resp = create_preview_generation_task(
|
||||
self._make_request(),
|
||||
authenticated_user=_make_user(),
|
||||
generation_task_repository=repo,
|
||||
db=MagicMock(),
|
||||
)
|
||||
assert resp.total == 1
|
||||
assert resp.items[0].status == "failed"
|
||||
|
||||
def test_enqueue_raises_global_queue_full(self):
|
||||
"""safe_enqueue 抛出 GlobalQueueFull → 503"""
|
||||
|
||||
@@ -0,0 +1,771 @@
|
||||
"""GPU 直连渲染管线:模板配置透传单测。
|
||||
|
||||
覆盖:标题样式(font/size/color/position/borderw/shadow)、静态字幕+ASR、BGM(volume/afade/adelay)、
|
||||
额外音轨音量,以及不传 config 时的默认兼容行为。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
# 路径对齐(同其它 unit tests)
|
||||
APP_ROOT = Path(__file__).resolve().parents[2] / "apps" / "worker"
|
||||
sys.path.insert(0, str(APP_ROOT))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stub helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass
|
||||
class _StubSeg:
|
||||
text: str
|
||||
start: float
|
||||
end: float
|
||||
|
||||
|
||||
class _StubClip:
|
||||
"""最小可用 stub:只包含 build_direct_render 需要的属性。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
local_path: str = "/tmp/_stub_clip.mp4",
|
||||
duration: float = 2.0,
|
||||
trim_start: float = 0.0,
|
||||
trim_end: float = 0.0,
|
||||
speed: float = 1.0,
|
||||
transition_type: str = "cut",
|
||||
config: dict | None = None,
|
||||
storage_key: str = "",
|
||||
):
|
||||
self.local_path = local_path
|
||||
self.duration = duration
|
||||
self.trim_start = trim_start
|
||||
self.trim_end = trim_end
|
||||
self.speed = speed
|
||||
self.transition_type = transition_type
|
||||
_cfg = dict(config or {"volume": 1.0})
|
||||
if storage_key:
|
||||
_cfg["_storage_key"] = storage_key
|
||||
elif "_storage_key" not in _cfg:
|
||||
_cfg["_storage_key"] = "test/clip.mp4"
|
||||
self.config = _cfg
|
||||
self._width = 1280
|
||||
self._height = 720
|
||||
|
||||
|
||||
def _make_clips(n: int = 2, dur: float = 2.0) -> list[_StubClip]:
|
||||
return [_StubClip(duration=dur) for _ in range(n)]
|
||||
|
||||
|
||||
def _patch_pipeline_helpers(monkeypatch):
|
||||
"""屏蔽 oss 上传和签名,避免依赖真实存储/网络;clip_has_audio/clip_volumes 通过参数传入。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
monkeypatch.setattr(
|
||||
gdp,
|
||||
"sign_asset_url",
|
||||
lambda sk, expires=3600: f"https://oss.example.com/{sk}",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
gdp,
|
||||
"upload_local_audio_and_sign",
|
||||
lambda p: (f"https://oss.example.com/{Path(p).name}", f"osskey/{Path(p).name}"),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 基线:不传 config 保持旧默认行为
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestNoConfigBackwardCompat:
|
||||
def test_default_title_drawtext_white_top(self, monkeypatch):
|
||||
"""不传 title_config 时:白字、top、48@720 按 width 缩放到 85、y=89(50@720p baseline)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(2, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
title_text="默认标题",
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True, True],
|
||||
clip_volumes=[1.0, 1.0],
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "fontcolor=0xffffff" in fc
|
||||
# 默认 position=top → y=89(scale(50)=89,与 CPU/vfb 一致),不含 h-th
|
||||
assert "y=89" in fc
|
||||
assert "h-th" not in fc
|
||||
# 默认字号 48@720 经 1280/720 缩放 = 85
|
||||
assert "fontsize=85" in fc
|
||||
assert "text='默认标题'" in fc
|
||||
# 默认粗体:黑色细描边 borderw=2@720(scale=4),与 CPU vfb 一致避免重影
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
|
||||
def test_no_subtitle_when_none(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 无标题无字幕时,应直接 format=yuv420p[vfinal]
|
||||
assert "format=yuv420p[vfinal]" in fc
|
||||
assert "drawtext=" not in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 标题样式:font/size/color/position/borderw/shadow
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestTitleStylePassthrough:
|
||||
def test_title_color_hex_converted_to_bgr(self, monkeypatch):
|
||||
"""#ff0000(红) → 0xff0000;注意我们直接按 RRGGBB 透传给 drawtext。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "红色标题", "color": "#ff0000", "position": "top", "size": 60},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "fontcolor=0xff0000" in fc
|
||||
# 60@720 按 1280/720 缩放 = 107
|
||||
assert "fontsize=107" in fc
|
||||
# top 位置默认 margin 50@720p → scale=89(未传 margin_top 用默认)
|
||||
assert "y=89" in fc
|
||||
assert "h-th" not in fc
|
||||
|
||||
def test_title_position_center(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "居中标题", "position": "center"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "(h-text_h)/2" in fc
|
||||
|
||||
def test_title_stroke_borderw(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={
|
||||
"text": "描边标题",
|
||||
"stroke": {"enabled": True, "width": 4, "color": "#0000ff"},
|
||||
},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# stroke width 4@720 经 1280/720 缩放 = 7
|
||||
assert "borderw=7" in fc
|
||||
assert "bordercolor=0x0000ff" in fc
|
||||
|
||||
def test_title_shadow_produces_two_drawtext_layers(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={
|
||||
"text": "阴影标题",
|
||||
"shadow": {"enabled": True, "offset_x": 3, "offset_y": 3, "color": "#000000@0.5"},
|
||||
},
|
||||
)
|
||||
# filter_complex 是 list[str]
|
||||
drawtext_count = sum(1 for f in plan.filter_complex if "drawtext=" in f)
|
||||
# 阴影层 + 主字层 = 2 条 drawtext
|
||||
assert drawtext_count == 2
|
||||
joined = " ".join(plan.filter_complex)
|
||||
# shadow offset 3@720 经 1280/720 缩放 = 5;默认 position=top → y=89
|
||||
assert "x=(w-text_w)/2+5" in joined
|
||||
assert "y=89+5" in joined
|
||||
assert "h-th" not in joined
|
||||
|
||||
def test_title_font_override(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "自定义字体", "font": "Noto Serif CJK SC"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "font=Noto Serif CJK SC" in fc
|
||||
|
||||
def test_title_disabled_hides_title(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
title_text="被禁用的标题",
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"enabled": False, "text": "被禁用的标题"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "drawtext=" not in fc
|
||||
|
||||
def test_title_custom_position_uses_pct_xy(self, monkeypatch):
|
||||
"""position=custom + pos_x/pos_y 百分比 → (w-text_w)*pct, (h-text_h)*pct。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={
|
||||
"text": "拖拽标题",
|
||||
"position": "custom",
|
||||
"pos_x": 30,
|
||||
"pos_y": 60,
|
||||
},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "(w-text_w)*0.3000" in fc
|
||||
assert "(h-text_h)*0.6000" in fc
|
||||
assert "h-th-" not in fc
|
||||
assert "y=89" not in fc
|
||||
|
||||
def test_title_margin_top_respected(self, monkeypatch):
|
||||
"""margin_top 透传:100@720,叠加默认 50 → 150@720 → scale=267@1280。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "远离顶部", "position": "top", "margin_top": 100},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 默认 50 + margin_top100 = 150@720 → scale=267
|
||||
assert "y=267" in fc
|
||||
|
||||
def test_title_default_bold_true(self, monkeypatch):
|
||||
"""不传 bold 时默认粗体:drawtext 用同色描边 borderw 模拟。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "粗体"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# bold=True 默认:黑色细描边 width=2@720 → scale=4,与 CPU vfb 一致
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
assert "text='粗体'" in fc
|
||||
|
||||
def test_title_bold_false_disables_faux_bold(self, monkeypatch):
|
||||
"""显式 bold=False 时不加粗描边。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "细体", "bold": False},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 无描边
|
||||
assert "borderw=" not in fc
|
||||
|
||||
def test_title_position_bottom(self, monkeypatch):
|
||||
"""position=bottom → y=h-th-{scaled margin}。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "底部", "position": "bottom"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# bottom margin 50@720 → scale=89
|
||||
assert "y=h-th-89" in fc
|
||||
|
||||
def test_title_size_scales_by_width_not_height(self, monkeypatch):
|
||||
"""不同分辨率下同 @720 基准的 size 等比缩放:1080x1920 下 36→54。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1080,
|
||||
output_height=1920,
|
||||
output_fps=30,
|
||||
title_text="竖屏标题",
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# default size 48@720 → 1080w = 72
|
||||
assert "fontsize=72" in fc
|
||||
# top margin 50@720 → 75
|
||||
assert "y=75" in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 字幕:静态 subtitle_text + ASR segments
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestSubtitlePassthrough:
|
||||
def test_static_subtitle_spans_full_duration(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(2, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True, True],
|
||||
clip_volumes=[1.0, 1.0],
|
||||
subtitle_config={"enabled": True, "text": "这是静态字幕", "color": "#00ff00"},
|
||||
static_subtitle_text="这是静态字幕",
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 应出现 static 文本,且 enable 范围 0 → 4.0
|
||||
assert "text='这是静态字幕'" in fc
|
||||
assert "between(t,0.000,4.000)" in fc
|
||||
assert "fontcolor=0x00ff00" in fc
|
||||
|
||||
def test_asr_segments_use_subtitle_style(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
segs = [
|
||||
_StubSeg("第一句", 0.0, 1.5),
|
||||
_StubSeg("第二句", 1.5, 3.0),
|
||||
]
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(2, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True, True],
|
||||
clip_volumes=[1.0, 1.0],
|
||||
subtitle_segments=segs,
|
||||
subtitle_config={"enabled": True, "color": "#0000ff", "size": 28, "position": "bottom"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "text='第一句'" in fc
|
||||
assert "text='第二句'" in fc
|
||||
assert "fontcolor=0x0000ff" in fc
|
||||
# sub size 28@720 经 1280/720 缩放 = 50
|
||||
assert "fontsize=50" in fc
|
||||
# bottom margin 50@720 → 89
|
||||
assert "y=h-th-89" in fc
|
||||
|
||||
def test_subtitle_disabled_hides_subs(self, monkeypatch):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
subtitle_config={"enabled": False, "text": "我被关了"},
|
||||
static_subtitle_text="我被关了",
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "drawtext=" not in fc
|
||||
|
||||
def test_subtitle_short_hex_color(self, monkeypatch):
|
||||
"""#fff → 0xffffff(缩写展开)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
subtitle_config={"text": "短色", "color": "#fff"},
|
||||
static_subtitle_text="短色",
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "fontcolor=0xffffff" in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BGM:volume / afade / adelay
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestBGMConfigPassthrough:
|
||||
def test_bgm_volume_from_config(self, monkeypatch, tmp_path):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"volume": 0.15, "enabled": True},
|
||||
)
|
||||
# BGM 音频滤镜链必须含 volume=0.15(挑输出 label 为 [au_bgm] 的那条)
|
||||
bgm_chain = [f for f in plan.filter_complex if f.rstrip().endswith("[au_bgm]")]
|
||||
assert len(bgm_chain) == 1, bgm_chain
|
||||
assert "volume=0.150" in bgm_chain[0]
|
||||
|
||||
def test_bgm_fade_in_out(self, monkeypatch, tmp_path):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(2, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True, True],
|
||||
clip_volumes=[1.0, 1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"volume": 0.3, "fade_in": 1.0, "fade_out": 1.5},
|
||||
)
|
||||
bgm_chain = [f for f in plan.filter_complex if f.rstrip().endswith("[au_bgm]")][0]
|
||||
assert "afade=t=in:st=0:d=1.00" in bgm_chain
|
||||
# fade_out 起点 = total_duration - fade_out = 2.5
|
||||
assert "afade=t=out:st=2.50:d=1.50" in bgm_chain
|
||||
|
||||
def test_bgm_audio_offset_adelay(self, monkeypatch, tmp_path):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"audio_offset": 2.5},
|
||||
)
|
||||
bgm_chain = [f for f in plan.filter_complex if f.rstrip().endswith("[au_bgm]")][0]
|
||||
# adelay 毫秒(2.5s → 2500),立体声双声道
|
||||
assert "adelay=2500|2500" in bgm_chain
|
||||
|
||||
def test_bgm_volume_adjust_db(self, monkeypatch, tmp_path):
|
||||
"""volume_adjust_db=-6dB → 增益 0.5,最终 volume 约 0.3*0.5=0.15。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"volume": 0.3, "volume_adjust_db": -6.0},
|
||||
)
|
||||
bgm_chain = [f for f in plan.filter_complex if f.rstrip().endswith("[au_bgm]")][0]
|
||||
# 0.3 * 10^(-6/20) ≈ 0.3 * 0.501 ≈ 0.150
|
||||
assert "volume=0.150" in bgm_chain
|
||||
|
||||
def test_bgm_disabled_drops_bgm_even_if_path_present(self, monkeypatch, tmp_path):
|
||||
"""bgm_config.enabled=False 时即使传 bgm_audio 也不挂载。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"enabled": False, "volume": 0.3},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "au_bgm" not in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# extra_audio_tracks 音量透传
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestExtraAudioVolume:
|
||||
def test_extra_audio_uses_passed_volume(self, monkeypatch, tmp_path):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
tts = tmp_path / "tts.m4a"
|
||||
tts.write_bytes(b"fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
extra_audio_tracks=[(tts, 0.7)],
|
||||
)
|
||||
# extra 音轨链应带 volume=0.7
|
||||
extras = [f for f in plan.filter_complex if "aex" in f and "volume" in f]
|
||||
assert any("volume=0.70" in e for e in extras)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 端到端:多配置组合 → filter_complex 无语法碎片
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestCombinedConfig:
|
||||
def test_title_static_sub_bgm_combined(self, monkeypatch, tmp_path):
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(2, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=4.0,
|
||||
clip_has_audio=[True, True],
|
||||
clip_volumes=[1.0, 1.0],
|
||||
bgm_audio=bgm,
|
||||
title_config={
|
||||
"text": "主标题",
|
||||
"color": "#ffff00",
|
||||
"position": "top",
|
||||
"size": 50,
|
||||
"stroke": {"enabled": True, "width": 2, "color": "#000000"},
|
||||
},
|
||||
subtitle_config={"text": "成片全字幕", "color": "#ffffff", "size": 24, "position": "bottom"},
|
||||
static_subtitle_text="成片全字幕",
|
||||
bgm_config={"volume": 0.2, "fade_in": 0.5, "fade_out": 1.0},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 标题:size 50@720 → 89,top margin 50@720 → 89,stroke 2@720 → 4
|
||||
assert "text='主标题'" in fc
|
||||
assert "fontcolor=0xffff00" in fc
|
||||
assert "fontsize=89" in fc
|
||||
assert "y=89" in fc
|
||||
assert "borderw=4" in fc
|
||||
# 字幕:size 24@720 → 43,bottom margin 50@720 → 89(默认不加粗)
|
||||
assert "text='成片全字幕'" in fc
|
||||
assert "fontsize=43" in fc
|
||||
assert "y=h-th-89" in fc
|
||||
assert "fontcolor=0xffffff" in fc
|
||||
# subtitle 默认 bold=False,无额外描边(用户未开 stroke)
|
||||
# BGM
|
||||
assert "volume=0.200" in fc
|
||||
assert "afade=t=in:st=0:d=0.50" in fc
|
||||
assert "afade=t=out:st=3.00:d=1.00" in fc
|
||||
# vfinal 存在
|
||||
assert "[vfinal]" in fc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 额外边界用例
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestEdgeCases:
|
||||
def test_invalid_color_falls_back_to_white(self, monkeypatch):
|
||||
"""非法色值回退 white,不抛异常。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "T", "color": "not-a-color"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 非法颜色不是 # 开头且不是命名,会被当命名色直接返回,不报错;确保至少 drawtext 有
|
||||
assert "drawtext=" in fc
|
||||
|
||||
def test_bold_title_increases_borderw(self, monkeypatch):
|
||||
"""bold=True 时若原无描边,自动加 borderw 黑色细描边模拟加粗(与 CPU vfb 一致,避免重影)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "粗体", "bold": True, "color": "#ff0000"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
# 仿粗用黑色细描边 2@720→scale=4,不跟文字色(避免同色描边重影)
|
||||
assert "borderw=4" in fc
|
||||
assert "bordercolor=0x000000" in fc
|
||||
|
||||
def test_asr_and_static_subtitle_asr_wins(self, monkeypatch):
|
||||
"""同时传 static_subtitle_text 和 ASR segments 时,ASR 优先(不插入静态全文)。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
segs = [_StubSeg("ASR1", 0.0, 1.0), _StubSeg("ASR2", 1.0, 2.0)]
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
subtitle_segments=segs,
|
||||
subtitle_config={"text": "静态全文", "color": "#ffffff"},
|
||||
static_subtitle_text="静态全文",
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "text='ASR1'" in fc
|
||||
assert "text='ASR2'" in fc
|
||||
# 静态全文不应该单独存在
|
||||
assert "between(t,0.000,2.000)" not in fc or "text='静态全文'" not in fc
|
||||
|
||||
def test_hex_color_with_alpha(self, monkeypatch):
|
||||
"""#rrggbbaa → 0xrrggbb@A 格式。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
title_config={"text": "半透明", "color": "#ff000080"},
|
||||
)
|
||||
fc = " ".join(plan.filter_complex)
|
||||
assert "fontcolor=0xff0000@" in fc
|
||||
|
||||
def test_bgm_invalid_volume_clamped(self, monkeypatch, tmp_path):
|
||||
"""volume 非法值回退默认 0.3;负值 clamp 到 0。"""
|
||||
import video_processing.gpu_direct_pipeline as gdp
|
||||
|
||||
_patch_pipeline_helpers(monkeypatch)
|
||||
bgm = tmp_path / "bgm.mp3"
|
||||
bgm.write_bytes(b"ID3fake")
|
||||
plan = gdp.build_direct_render(
|
||||
resolved_clips=_make_clips(1, 2.0),
|
||||
output_width=1280,
|
||||
output_height=720,
|
||||
output_fps=30,
|
||||
total_duration=2.0,
|
||||
clip_has_audio=[True],
|
||||
clip_volumes=[1.0],
|
||||
bgm_audio=bgm,
|
||||
bgm_config={"volume": -999},
|
||||
)
|
||||
bgm_chain = [f for f in plan.filter_complex if f.rstrip().endswith("[au_bgm]")][0]
|
||||
assert "volume=0.000" in bgm_chain
|
||||
@@ -666,7 +666,7 @@ class TestDownloadAssets:
|
||||
_make_clip("c2", order=1, asset_id="asset_002"),
|
||||
]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 2
|
||||
assert "asset_001" in asset_path_map
|
||||
@@ -693,7 +693,7 @@ class TestDownloadAssets:
|
||||
_make_clip("c2", order=1, asset_id="asset_002"),
|
||||
]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 1
|
||||
assert "asset_002" in asset_path_map
|
||||
@@ -714,7 +714,7 @@ class TestDownloadAssets:
|
||||
_make_clip("c1", order=0, asset_id="asset_001"),
|
||||
]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 0
|
||||
assert len(rendered_ids) == 0
|
||||
@@ -750,7 +750,7 @@ class TestDownloadAssets:
|
||||
_make_clip("c3", order=2, asset_id="asset_003"),
|
||||
]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 2
|
||||
assert "c1" in rendered_ids
|
||||
@@ -771,7 +771,7 @@ class TestDownloadAssets:
|
||||
_make_clip("c2", order=1, asset_id="asset_shared"),
|
||||
]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 1
|
||||
assert mock_download.call_count == 1
|
||||
@@ -795,7 +795,7 @@ class TestDownloadAssets:
|
||||
adapter = RenderAdapter(mock_db)
|
||||
clips = [_make_clip("c1", order=0, asset_id="asset_fallback")]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 1
|
||||
assert "c1" in rendered_ids
|
||||
@@ -820,7 +820,7 @@ class TestDownloadAssets:
|
||||
adapter = RenderAdapter(mock_db)
|
||||
clips = [_make_clip("c1", order=0, asset_id="asset_no_key")]
|
||||
|
||||
asset_path_map, rendered_ids, failed_ids = adapter._download_assets(clips, tmp_path)
|
||||
asset_path_map, rendered_ids, failed_ids, _storage_map = adapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(asset_path_map) == 0
|
||||
assert "c1" in failed_ids
|
||||
|
||||
+46
-119
@@ -1,4 +1,4 @@
|
||||
"""task_enqueue 单测 — 队列限流 + 安全入队逻辑."""
|
||||
"""task_enqueue 单测 — 队列限流 + 安全入队逻辑 (#2098)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -14,12 +14,8 @@ from app.core.task_enqueue import (
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
|
||||
# ── Fixtures / Helpers ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class MockRepository:
|
||||
"""Mock 任务仓储,用计数器模拟 pending 数量."""
|
||||
|
||||
def __init__(self, global_count: int = 0, user_count: int = 0):
|
||||
self._global = global_count
|
||||
self._user = user_count
|
||||
@@ -43,19 +39,12 @@ def make_mock_task(task_id: str = "task-1"):
|
||||
return task
|
||||
|
||||
|
||||
# ── check_queue_limits ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCheckQueueLimits:
|
||||
"""check_queue_limits 预检查限流."""
|
||||
|
||||
def test_below_limits_passes(self):
|
||||
repo = MockRepository(global_count=5, user_count=1)
|
||||
# 不抛异常就是通过
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_at_limit_raises(self):
|
||||
"""达到全局上限即拒绝."""
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=1)
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
@@ -67,210 +56,145 @@ class TestCheckQueueLimits:
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_at_limit_raises(self):
|
||||
"""达到用户上限即拒绝."""
|
||||
def test_user_at_limit_no_longer_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == USER_PENDING_LIMIT
|
||||
assert exc_info.value.limit == USER_PENDING_LIMIT
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_over_limit_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 1)
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_priority_over_user(self):
|
||||
"""全局和用户都超限时,优先抛全局异常."""
|
||||
repo = MockRepository(
|
||||
global_count=GLOBAL_PENDING_LIMIT + 1,
|
||||
user_count=USER_PENDING_LIMIT + 1,
|
||||
)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
def test_user_over_limit_no_longer_raises(self):
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_empty_user_id_skips_user_check(self):
|
||||
"""user_id 为空时跳过用户级检查."""
|
||||
repo = MockRepository(global_count=5, user_count=999)
|
||||
# 不抛异常 = 通过(只检查全局)
|
||||
check_queue_limits("", repo)
|
||||
|
||||
def test_custom_limits(self):
|
||||
"""支持自定义限流阈值."""
|
||||
repo = MockRepository(global_count=5, user_count=5)
|
||||
# 默认阈值下 user 5 > 3 会被拒
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("u1", repo)
|
||||
|
||||
# 自定义更高阈值就能通过
|
||||
check_queue_limits("u1", repo, user_pending_limit=10, global_pending_limit=10)
|
||||
|
||||
|
||||
# ── safe_enqueue_generation_task ──────────────────────────────────────────
|
||||
def test_custom_global_limit_still_honored(self):
|
||||
repo = MockRepository(global_count=15, user_count=999)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=10)
|
||||
check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=20)
|
||||
|
||||
|
||||
class TestSafeEnqueueGenerationTask:
|
||||
"""safe_enqueue_generation_task 安全入队."""
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_success_path(self, mock_celery):
|
||||
"""正常路径:入队前检查通过 → 发送Celery → 入队后检查通过."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once_with("worker.generate_video", args=[task.id])
|
||||
task.mark_failed.assert_not_called()
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_no_user_id_skips_user_check(self, mock_celery):
|
||||
"""不传 user_id 跳过用户级限流."""
|
||||
repo = MockRepository(global_count=1, user_count=999)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
assert result is True
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_precheck_global_over_marks_failed(self, mock_celery):
|
||||
"""入队前全局超限:标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT + 1, user_count=0)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
mock_celery.send_task.assert_not_called()
|
||||
assert repo.update_called == 1
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_precheck_user_over_marks_failed(self, mock_celery):
|
||||
"""入队前用户超限:标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 1)
|
||||
def test_precheck_user_over_still_enqueues(self, mock_celery, caplog):
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
mock_celery.send_task.assert_not_called()
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
task.mark_failed.assert_not_called()
|
||||
assert any("超过软上限" in r.message for r in caplog.records)
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_celery_send_false_returns_false(self, mock_celery):
|
||||
"""Celery 发送失败:返回 False,任务标记 failed."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
mock_celery.send_task.side_effect = Exception("celery down")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is False
|
||||
task.mark_failed.assert_called_once()
|
||||
assert "入队失败" in task.mark_failed.call_args[0][0]
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_celery_send_failure_update_also_fails(self, mock_celery):
|
||||
"""Celery 发送失败 + mark_failed 更新也失败:不崩溃."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
repo.update = MagicMock(side_effect=Exception("db down"))
|
||||
task = make_mock_task()
|
||||
mock_celery.send_task.side_effect = Exception("celery down")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is False
|
||||
# 不抛异常就是胜利
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_postcheck_global_over_rollback(self, mock_celery):
|
||||
"""入队后全局超限(并发竞态):回滚标记 failed,抛异常."""
|
||||
# 入队前刚好通过,但入队后再查发现超限
|
||||
call_count = [0]
|
||||
|
||||
def count_pending_total_side_effect():
|
||||
call_count[0] += 1
|
||||
if call_count[0] == 1: # 入队前检查
|
||||
return GLOBAL_PENDING_LIMIT # 等于上限,用 > 判断所以通过
|
||||
return GLOBAL_PENDING_LIMIT + 1 # 入队后再查,超限
|
||||
if call_count[0] == 1:
|
||||
return GLOBAL_PENDING_LIMIT
|
||||
return GLOBAL_PENDING_LIMIT + 1
|
||||
|
||||
repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=0)
|
||||
repo.count_pending_total = MagicMock(side_effect=count_pending_total_side_effect)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# 异常是 GlobalQueueFull 类型,且任务已被标记为 failed(含"入队后"原因)
|
||||
task.mark_failed.assert_called_once()
|
||||
assert "入队后" in task.mark_failed.call_args[0][0]
|
||||
mock_celery.send_task.assert_called_once()
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_postcheck_user_over_rollback(self, mock_celery):
|
||||
"""入队后用户超限:回滚标记 failed,抛异常."""
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
# 入队前用 > 判断,等于上限通过;入队后模拟并发超限
|
||||
original_user_count = repo.count_pending_by_user
|
||||
def test_postcheck_user_over_does_not_rollback(self, mock_celery, caplog):
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT)
|
||||
call_count = [0]
|
||||
|
||||
def count_by_user_side_effect(user_id):
|
||||
call_count[0] += 1
|
||||
if call_count[0] <= 1: # 入队前
|
||||
return USER_PENDING_LIMIT # 用 > 判断,等于时通过
|
||||
return USER_PENDING_LIMIT + 1 # 入队后,超限
|
||||
if call_count[0] <= 1:
|
||||
return USER_PENDING_LIMIT
|
||||
return USER_PENDING_LIMIT + 1
|
||||
|
||||
repo.count_pending_by_user = MagicMock(side_effect=count_by_user_side_effect)
|
||||
task = make_mock_task()
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
task.mark_failed.assert_called_once()
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
task.mark_failed.assert_not_called()
|
||||
assert any("超软上限(入队后)" in r.message for r in caplog.records)
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_log_task_status_enabled(self, mock_celery):
|
||||
"""log_task_status=True 时日志中包含状态."""
|
||||
repo = MockRepository(global_count=1, user_count=1)
|
||||
task = make_mock_task()
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1", log_task_status=True)
|
||||
assert result is True
|
||||
|
||||
@patch("app.core.task_enqueue.celery_app")
|
||||
def test_custom_limits_in_enqueue(self, mock_celery):
|
||||
"""自定义限流阈值用于入队检查."""
|
||||
repo = MockRepository(global_count=5, user_count=5)
|
||||
def test_custom_global_limit_in_enqueue(self, mock_celery):
|
||||
repo = MockRepository(global_count=15, user_count=999)
|
||||
task = make_mock_task()
|
||||
|
||||
# 默认阈值下用户 5 > 3 会被拒
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# 重置 mock 计数
|
||||
task.mark_failed.reset_mock()
|
||||
|
||||
# 调大阈值后通过
|
||||
result = safe_enqueue_generation_task(
|
||||
task,
|
||||
repo,
|
||||
user_id="user-1",
|
||||
user_pending_limit=10,
|
||||
global_pending_limit=10,
|
||||
)
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
|
||||
|
||||
# ── 异常类 ────────────────────────────────────────────────────────────────
|
||||
task2 = make_mock_task("task-2")
|
||||
mock_celery.reset_mock()
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task2, repo, user_id="user-1", user_pending_limit=999, global_pending_limit=10)
|
||||
|
||||
|
||||
class TestExceptionClasses:
|
||||
"""异常类消息格式."""
|
||||
|
||||
def test_user_pending_limit_message(self):
|
||||
exc = UserPendingLimitExceeded("u1", 5, 3)
|
||||
assert "u1" in str(exc)
|
||||
@@ -281,3 +205,6 @@ class TestExceptionClasses:
|
||||
exc = GlobalQueueFull(25, 20)
|
||||
assert "25" in str(exc)
|
||||
assert "20" in str(exc)
|
||||
|
||||
def test_user_pending_limit_constant(self):
|
||||
assert USER_PENDING_LIMIT == 20
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""任务队列限流防护单元测试。"""
|
||||
"""任务队列限流防护单元测试 (#2098: 用户级改为软上限,仅全局硬拒)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -19,17 +19,8 @@ from app.core.task_enqueue import (
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MockRepository:
|
||||
"""支持 pending 计数的 mock repository。
|
||||
|
||||
支持通过 set_pending 动态修改计数,用于模拟入队后计数变化的并发场景。
|
||||
"""
|
||||
|
||||
def __init__(self, user_pending: int = 0, global_pending: int = 0):
|
||||
self._user_pending = user_pending
|
||||
self._global_pending = global_pending
|
||||
@@ -47,7 +38,6 @@ class MockRepository:
|
||||
return task
|
||||
|
||||
def set_pending(self, *, user_pending: int | None = None, global_pending: int | None = None):
|
||||
"""动态修改 pending 计数,模拟并发场景。"""
|
||||
if user_pending is not None:
|
||||
self._user_pending = user_pending
|
||||
if global_pending is not None:
|
||||
@@ -67,58 +57,39 @@ class MockTask:
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_celery(monkeypatch):
|
||||
"""mock 掉 celery_app.send_task,避免真实发送。"""
|
||||
mock_send = MagicMock()
|
||||
monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send)
|
||||
return mock_send
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 常量导出测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_limit_constants_are_exported():
|
||||
"""限流阈值常量已导出,供业务代码引用。"""
|
||||
assert USER_PENDING_LIMIT == 3
|
||||
"""#2098: USER_PENDING_LIMIT 从 3 提到 20 作为软上限;GLOBAL_PENDING_LIMIT 保持 20 为硬上限。"""
|
||||
assert USER_PENDING_LIMIT == 20
|
||||
assert GLOBAL_PENDING_LIMIT == 20
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# check_queue_limits 单元测试(预检查用,>= 边界)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCheckQueueLimits:
|
||||
"""队列限流检查函数测试(预检查语义,>= 上限即拒绝)。"""
|
||||
"""check_queue_limits 预检查:仅全局硬上限拒绝,用户级改为软提示。"""
|
||||
|
||||
def test_normal_passes_through(self):
|
||||
"""正常范围内的任务不受限制。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_limit_exceeded_raises(self):
|
||||
"""用户 pending 超过上限抛 UserPendingLimitExceeded。"""
|
||||
repo = MockRepository(user_pending=4, global_pending=5)
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == 4
|
||||
assert exc_info.value.limit == 3
|
||||
def test_user_limit_exceeded_no_longer_raises(self):
|
||||
"""#2098: 用户 pending 超过软上限不再抛异常。"""
|
||||
repo = MockRepository(user_pending=100, global_pending=5)
|
||||
check_queue_limits("user-1", repo) # 不抛即通过
|
||||
|
||||
def test_user_at_limit_also_raises(self):
|
||||
"""用户 pending 刚好等于上限也拒绝(>= 边界)。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5)
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
check_queue_limits("user-1", repo)
|
||||
def test_user_at_limit_no_longer_raises(self):
|
||||
"""#2098: 用户 pending 等于软上限也不拒绝。"""
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_user_below_limit_passes(self):
|
||||
"""用户 pending 比上限少 1,通过。"""
|
||||
repo = MockRepository(user_pending=2, global_pending=5)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_limit_exceeded_raises(self):
|
||||
"""全局 pending 超过上限抛 GlobalQueueFull。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=21)
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
check_queue_limits("user-1", repo)
|
||||
@@ -126,82 +97,52 @@ class TestCheckQueueLimits:
|
||||
assert exc_info.value.limit == 20
|
||||
|
||||
def test_global_at_limit_also_raises(self):
|
||||
"""全局 pending 刚好等于上限也拒绝(>= 边界)。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_below_limit_passes(self):
|
||||
"""全局 pending 比上限少 1,通过。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=19)
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_global_takes_priority_over_user(self):
|
||||
"""全局和用户都超限时,优先抛全局异常。"""
|
||||
repo = MockRepository(user_pending=5, global_pending=25)
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
check_queue_limits("user-1", repo)
|
||||
|
||||
def test_empty_user_id_skips_user_check(self):
|
||||
"""不传 user_id 时跳过用户级检查,只做全局检查。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
# 用户超限但不传 user_id → 全局未超限,应该通过
|
||||
check_queue_limits("", repo)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# safe_enqueue_generation_task 限流集成测试(入队前用 >,包含当前任务)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSafeEnqueueWithLimits:
|
||||
"""安全入队函数的限流功能测试。"""
|
||||
"""safe_enqueue_generation_task:用户超限仅 warning 仍入队;全局超限硬拒。"""
|
||||
|
||||
def test_normal_task_enqueues_successfully(self, mock_celery):
|
||||
"""正常任务入队成功,返回 True。"""
|
||||
repo = MockRepository(user_pending=0, global_pending=0)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
|
||||
# 成功入队后持久化 celery 消息 ID(#1714:孤儿清理据此 revoke/清队列)
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_user_limit_rejected_with_failed_status(self, mock_celery):
|
||||
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
|
||||
repo = MockRepository(user_pending=5, global_pending=5)
|
||||
def test_user_limit_exceeded_still_enqueues(self, mock_celery, caplog):
|
||||
"""#2098: 用户远超软上限仍入队,任务不被标记 failed。"""
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(user_pending=100, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
assert task.status == "failed"
|
||||
assert "限流" in task.error_message
|
||||
assert len(repo.updated_tasks) == 1
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending" # 没被标记 failed
|
||||
assert any("超过软上限" in r.message for r in caplog.records)
|
||||
|
||||
def test_user_at_limit_still_passes(self, mock_celery):
|
||||
"""用户 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。
|
||||
|
||||
与预检查的 >= 语义一致:预检查时 pending=3 拒绝(不能再加新的),
|
||||
但 safe_enqueue 被调用时任务已是 pending(就是第3个),
|
||||
pending=3 不满足 >3,所以通过。
|
||||
"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5)
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_user_one_over_limit_rejected(self, mock_celery):
|
||||
"""用户 pending = limit + 1:超限被拒。"""
|
||||
repo = MockRepository(user_pending=4, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
|
||||
def test_global_limit_rejected_with_failed_status(self, mock_celery):
|
||||
"""全局超限:任务标记为 failed,抛 GlobalQueueFull。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=21)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
@@ -211,7 +152,6 @@ class TestSafeEnqueueWithLimits:
|
||||
assert len(repo.updated_tasks) == 1
|
||||
|
||||
def test_global_at_limit_still_passes(self, mock_celery):
|
||||
"""全局 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
@@ -219,7 +159,6 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_no_user_id_skips_user_limit(self, mock_celery):
|
||||
"""不传 user_id 时跳过用户级限流,只做全局检查。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
@@ -227,7 +166,6 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_called_once()
|
||||
|
||||
def test_no_user_id_still_checks_global(self, mock_celery):
|
||||
"""不传 user_id 时全局超限仍然被拦。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=25)
|
||||
task = MockTask("task-1")
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
@@ -235,121 +173,74 @@ class TestSafeEnqueueWithLimits:
|
||||
mock_celery.assert_not_called()
|
||||
|
||||
def test_default_limits_match_constants(self, mock_celery):
|
||||
"""默认配置与导出常量一致。"""
|
||||
# 刚好在默认限制内(limit - 1)
|
||||
repo = MockRepository(user_pending=2, global_pending=19)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
|
||||
def test_update_failure_does_not_crash(self, mock_celery):
|
||||
"""repository.update 失败也不崩溃,异常继续向上抛。"""
|
||||
|
||||
class BadRepo(MockRepository):
|
||||
def update(self, task):
|
||||
raise RuntimeError("db down")
|
||||
|
||||
repo = BadRepo(user_pending=5, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
# 仍然抛 UserPendingLimitExceeded,不会被 update 失败掩盖
|
||||
with pytest.raises(UserPendingLimitExceeded):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
mock_celery.assert_not_called()
|
||||
# 任务状态还是变了(内存里改了)
|
||||
assert task.status == "failed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 入队后最终校验(并发竞态兜底)测试
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPostEnqueueFinalCheck:
|
||||
"""入队后最终校验:模拟并发场景,Celery发送后计数增加被兜住。"""
|
||||
"""入队后校验:仅全局超限回滚;用户超限仅 warning。"""
|
||||
|
||||
def test_post_enqueue_global_overflow_rollback(self, mock_celery):
|
||||
"""并发场景:入队前检查通过,但发送Celery后全局计数超限 → 回滚为failed。
|
||||
|
||||
模拟两个请求同时通过入队前检查(都查到 global=19),
|
||||
都创建了任务(DB里变成 21),先发送Celery的那个在最终校验时被兜住。
|
||||
"""
|
||||
repo = MockRepository(user_pending=1, global_pending=20) # 入队前:20 > 20?否
|
||||
repo = MockRepository(user_pending=1, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
|
||||
# 模拟发送Celery后,另一个并发请求也创建了任务,全局变成21
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(global_pending=21)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(GlobalQueueFull) as exc_info:
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
# Celery 确实发出去了(兜底不撤销 Celery,只回滚 DB 状态)
|
||||
mock_celery.assert_called_once()
|
||||
# 任务被标记为 failed
|
||||
assert task.status == "failed"
|
||||
assert "入队后" in task.error_message
|
||||
assert exc_info.value.pending_count == 21
|
||||
assert len(repo.updated_tasks) == 1
|
||||
|
||||
def test_post_enqueue_user_overflow_rollback(self, mock_celery):
|
||||
"""并发场景:入队前检查通过,但发送Celery后用户计数超限 → 回滚为failed。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=5) # 入队前:3 > 3?否
|
||||
def test_post_enqueue_user_overflow_no_rollback(self, mock_celery, caplog):
|
||||
"""#2098: 入队后用户超软上限仅 warning,不回滚。"""
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.WARNING)
|
||||
repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=4)
|
||||
repo.set_pending(user_pending=USER_PENDING_LIMIT + 1)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(UserPendingLimitExceeded) as exc_info:
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "failed"
|
||||
assert "入队后" in task.error_message
|
||||
assert exc_info.value.user_id == "user-1"
|
||||
assert exc_info.value.pending_count == 4
|
||||
|
||||
def test_post_enqueue_global_priority_over_user(self, mock_celery):
|
||||
"""入队后校验:全局和用户都超限时,优先抛全局异常。"""
|
||||
repo = MockRepository(user_pending=3, global_pending=20)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=5, global_pending=22)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
with pytest.raises(GlobalQueueFull):
|
||||
safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert task.status == "failed"
|
||||
|
||||
def test_post_enqueue_no_change_still_passes(self, mock_celery):
|
||||
"""入队后计数没变 → 正常通过,不回滚。"""
|
||||
repo = MockRepository(user_pending=2, global_pending=10)
|
||||
task = MockTask("task-1")
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending" # 状态没变
|
||||
# 入队成功后持久化 celery_task_id(#1714),业务状态不变
|
||||
assert task.status == "pending" # 不回滚
|
||||
assert any("超软上限(入队后)" in r.message for r in caplog.records)
|
||||
|
||||
def test_post_enqueue_no_change_still_passes(self, mock_celery):
|
||||
repo = MockRepository(user_pending=2, global_pending=10)
|
||||
task = MockTask("task-1")
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
|
||||
assert result is True
|
||||
mock_celery.assert_called_once()
|
||||
assert task.status == "pending"
|
||||
assert len(repo.updated_tasks) == 1
|
||||
assert task.celery_task_id
|
||||
|
||||
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
|
||||
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
|
||||
repo = MockRepository(user_pending=10, global_pending=5)
|
||||
task = MockTask("task-1")
|
||||
|
||||
def side_effect(*args, **kwargs):
|
||||
repo.set_pending(user_pending=15, global_pending=5) # 用户超限但全局没超
|
||||
repo.set_pending(user_pending=15, global_pending=5)
|
||||
|
||||
mock_celery.side_effect = side_effect
|
||||
|
||||
result = safe_enqueue_generation_task(task, repo, user_id="")
|
||||
assert result is True # 用户级不检查,全局没超限 → 通过
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_user_pending_limit_exceeded_class_still_exists():
|
||||
"""UserPendingLimitExceeded 保留用于兼容历史 import/except(#2098 后不再主动 raise)。"""
|
||||
exc = UserPendingLimitExceeded("u1", 5, 3)
|
||||
assert exc.user_id == "u1"
|
||||
assert exc.pending_count == 5
|
||||
assert exc.limit == 3
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from video_processing.thumbnail_generator import _format_seek_time
|
||||
|
||||
@@ -198,10 +200,23 @@ class TestExtractAndUploadCoverFramesFallback:
|
||||
lambda: FakeClient(),
|
||||
)
|
||||
|
||||
# Mock ffmpeg 抽帧
|
||||
# Mock 单次 ffmpeg 抽帧直接返回 dummy 帧,避免真调用 ffmpeg
|
||||
def _fake_single_pass(video_path, seek_points, out_dir, prefix="frame", **kw):
|
||||
results = []
|
||||
for i, st in enumerate(seek_points):
|
||||
fp = Path(out_dir) / f"{prefix}_{i + 1:02d}.jpg"
|
||||
fp.write_bytes(b"\xff\xd8\xff\xe0") # 最小 jpeg 头
|
||||
results.append((st, str(fp)))
|
||||
return results
|
||||
|
||||
monkeypatch.setattr(
|
||||
"video_processing.thumbnail_generator.extract_first_frame",
|
||||
lambda video_path, output_path, **kw: output_path,
|
||||
"video_processing.thumbnail_generator._extract_frames_single_pass",
|
||||
_fake_single_pass,
|
||||
)
|
||||
# blackdetect 直接返回空
|
||||
monkeypatch.setattr(
|
||||
"video_processing.thumbnail_generator._detect_black_intervals",
|
||||
lambda *a, **kw: [],
|
||||
)
|
||||
# Mock upload
|
||||
monkeypatch.setattr(
|
||||
@@ -221,3 +236,176 @@ class TestExtractAndUploadCoverFramesFallback:
|
||||
assert len(result) == 2
|
||||
assert all("url" in item for item in result)
|
||||
assert all("position" in item for item in result)
|
||||
|
||||
|
||||
class TestSeekPointBlackAvoidance:
|
||||
"""_adjust_seek_points_avoid_black 纯逻辑测试."""
|
||||
|
||||
def test_no_black_intervals_returns_unchanged(self):
|
||||
from video_processing.thumbnail_generator import _adjust_seek_points_avoid_black
|
||||
|
||||
pts = [2.0, 5.0, 8.0]
|
||||
out = _adjust_seek_points_avoid_black(pts, [], duration=10.0)
|
||||
assert out == [2.0, 5.0, 8.0]
|
||||
|
||||
def test_point_in_black_shifts_forward(self):
|
||||
from video_processing.thumbnail_generator import _adjust_seek_points_avoid_black
|
||||
|
||||
# 黑屏 [4, 6],点在 5.0,向前偏移到 4-0.25=3.75
|
||||
pts = [5.0]
|
||||
out = _adjust_seek_points_avoid_black(pts, [(4.0, 6.0)], duration=10.0)
|
||||
assert out[0] == pytest.approx(3.75, abs=0.01)
|
||||
|
||||
def test_point_at_start_shifts_backward(self):
|
||||
from video_processing.thumbnail_generator import _adjust_seek_points_avoid_black
|
||||
|
||||
# 黑屏 [0, 3],点在 1.0,向前偏移 -0.25 会 <0 → 向后偏移到 3+0.25=3.25
|
||||
pts = [1.0]
|
||||
out = _adjust_seek_points_avoid_black(pts, [(0.0, 3.0)], duration=10.0)
|
||||
assert out[0] == pytest.approx(3.25, abs=0.01)
|
||||
|
||||
def test_all_black_keeps_point(self):
|
||||
from video_processing.thumbnail_generator import _adjust_seek_points_avoid_black
|
||||
|
||||
# 全黑,偏移都无效,保留原点
|
||||
pts = [5.0]
|
||||
out = _adjust_seek_points_avoid_black(pts, [(0.0, 10.0)], duration=10.0)
|
||||
assert out[0] == pytest.approx(5.0, abs=0.01)
|
||||
|
||||
def test_multiple_points_decouple(self):
|
||||
from video_processing.thumbnail_generator import _adjust_seek_points_avoid_black
|
||||
|
||||
pts = [2.0, 5.0, 8.0]
|
||||
black = [(4.5, 5.5)] # 只有中点在黑屏
|
||||
out = _adjust_seek_points_avoid_black(pts, black, duration=10.0)
|
||||
assert out[0] == 2.0
|
||||
assert out[2] == 8.0
|
||||
# 中点必须不在黑屏内
|
||||
assert not (4.5 <= out[1] <= 5.5)
|
||||
|
||||
|
||||
class TestScorerRunsInsideTempDir:
|
||||
"""P1 修复:scorer 必须在 TemporaryDirectory 块内调用(帧文件还在时)。"""
|
||||
|
||||
def _setup_mocks(self, monkeypatch, tmp_path, *, scorer_should_read=True):
|
||||
from video_processing.thumbnail_generator import extract_and_upload_cover_frames
|
||||
|
||||
class FakeClient:
|
||||
is_available = False
|
||||
|
||||
monkeypatch.setattr(
|
||||
"packages.shared.mediakit_client.get_mediakit_client",
|
||||
lambda: FakeClient(),
|
||||
)
|
||||
|
||||
self._frames_on_disk_when_called = []
|
||||
|
||||
def _fake_single_pass(video_path, seek_points, out_dir, prefix="frame", **kw):
|
||||
results = []
|
||||
for i, st in enumerate(seek_points):
|
||||
fp = Path(out_dir) / f"{prefix}_{i + 1:02d}.jpg"
|
||||
fp.write_bytes(b"\xff\xd8\xff\xe0" + b"X" * 200)
|
||||
results.append((st, str(fp)))
|
||||
return results
|
||||
|
||||
monkeypatch.setattr(
|
||||
"video_processing.thumbnail_generator._extract_frames_single_pass",
|
||||
_fake_single_pass,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"video_processing.thumbnail_generator._detect_black_intervals",
|
||||
lambda *a, **kw: [],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"video_processing.ffmpeg_utils.probe_duration",
|
||||
lambda path: 60.0,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"video_processing.oss_helpers.upload_to_oss",
|
||||
lambda path, key: f"https://oss.example.com/{key}",
|
||||
)
|
||||
# 标题叠加 no-op
|
||||
monkeypatch.setattr(
|
||||
"video_processing.thumbnail_generator.apply_title_overlay",
|
||||
lambda *a, **kw: None,
|
||||
)
|
||||
|
||||
# 记录 scorer 被调用时各 image_path 是否存在
|
||||
def _fake_scorer(candidates):
|
||||
for c in candidates:
|
||||
self._frames_on_disk_when_called.append(Path(c["image_path"]).exists())
|
||||
# 给个假评分:倒序排,验证顺序被应用
|
||||
scored = list(candidates)
|
||||
for i, c in enumerate(scored):
|
||||
c["score"] = float(len(scored) - i)
|
||||
scored.sort(key=lambda c: c["score"], reverse=True)
|
||||
return scored
|
||||
|
||||
monkeypatch.setattr(
|
||||
"packages.shared.cover_frame_scorer.score_frames",
|
||||
_fake_scorer,
|
||||
)
|
||||
return extract_and_upload_cover_frames
|
||||
|
||||
def test_scorer_reads_files_while_they_exist(self, tmp_path, monkeypatch):
|
||||
"""核心 P1:评分时帧文件必须还在磁盘上(在 TemporaryDirectory 内调用)。"""
|
||||
extract = self._setup_mocks(monkeypatch, tmp_path)
|
||||
video_file = tmp_path / "t.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
result = extract(str(video_file), "plan1", num_frames=3)
|
||||
# scorer 看到的 3 个文件都必须存在
|
||||
assert len(self._frames_on_disk_when_called) == 3
|
||||
assert all(self._frames_on_disk_when_called), f"scorer 调用时有文件已被删除: {self._frames_on_disk_when_called}"
|
||||
# 结果按评分降序排列(is_best 在第一个)
|
||||
assert len(result) == 3
|
||||
assert result[0].get("is_best") is True
|
||||
# 结果中不应该再暴露 image_path
|
||||
assert all("image_path" not in c for c in result)
|
||||
|
||||
def test_scorer_failure_falls_back_gracefully(self, tmp_path, monkeypatch):
|
||||
"""评分抛异常时不应中断上传,仍返回所有候选帧。"""
|
||||
from video_processing.thumbnail_generator import extract_and_upload_cover_frames
|
||||
|
||||
class FakeClient:
|
||||
is_available = False
|
||||
|
||||
monkeypatch.setattr("packages.shared.mediakit_client.get_mediakit_client", lambda: FakeClient())
|
||||
|
||||
def _fake_single_pass(video_path, seek_points, out_dir, prefix="frame", **kw):
|
||||
results = []
|
||||
for i, st in enumerate(seek_points):
|
||||
fp = Path(out_dir) / f"{prefix}_{i + 1:02d}.jpg"
|
||||
fp.write_bytes(b"\xff\xd8\xff\xe0" + b"X" * 100)
|
||||
results.append((st, str(fp)))
|
||||
return results
|
||||
|
||||
monkeypatch.setattr("video_processing.thumbnail_generator._extract_frames_single_pass", _fake_single_pass)
|
||||
monkeypatch.setattr("video_processing.thumbnail_generator._detect_black_intervals", lambda *a, **kw: [])
|
||||
monkeypatch.setattr("video_processing.ffmpeg_utils.probe_duration", lambda p: 60.0)
|
||||
monkeypatch.setattr(
|
||||
"video_processing.oss_helpers.upload_to_oss",
|
||||
lambda path, key: f"https://oss/{key}",
|
||||
)
|
||||
monkeypatch.setattr("video_processing.thumbnail_generator.apply_title_overlay", lambda *a, **kw: None)
|
||||
|
||||
def _boom(candidates):
|
||||
raise RuntimeError("cv2 crashed")
|
||||
|
||||
monkeypatch.setattr("packages.shared.cover_frame_scorer.score_frames", _boom)
|
||||
|
||||
video_file = tmp_path / "t.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
# 不应抛出
|
||||
result = extract_and_upload_cover_frames(str(video_file), "plan1", num_frames=3)
|
||||
assert len(result) == 3
|
||||
assert all("url" in c for c in result)
|
||||
|
||||
def test_best_frame_is_first_after_scoring(self, tmp_path, monkeypatch):
|
||||
"""评分后 best 帧(score 最高)在 candidates[0],is_best=True。"""
|
||||
extract = self._setup_mocks(monkeypatch, tmp_path)
|
||||
video_file = tmp_path / "t.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
result = extract(str(video_file), "plan1", num_frames=5)
|
||||
assert result[0]["is_best"] is True
|
||||
scores = [c.get("score", 0.0) for c in result]
|
||||
assert scores == sorted(scores, reverse=True)
|
||||
|
||||
@@ -128,6 +128,25 @@ class StubIngestJobRepository:
|
||||
def get(self, job_id: str) -> IngestJob | None:
|
||||
return self._jobs.get(job_id)
|
||||
|
||||
def find_by_asset_id(self, asset_id: str) -> IngestJob | None:
|
||||
"""返回 asset 最近一条在跑/已完成 job(FAILED 不返回,允许重提)。"""
|
||||
from packages.domain.classification import IngestJobStatus as S
|
||||
|
||||
if not asset_id:
|
||||
return None
|
||||
running = None
|
||||
completed = None
|
||||
for job in self._jobs.values():
|
||||
if getattr(job, "asset_id", "") != asset_id:
|
||||
continue
|
||||
if job.status in (S.PENDING, S.PROCESSING):
|
||||
if running is None or job.created_at > running.created_at:
|
||||
running = job
|
||||
elif job.status == S.COMPLETED:
|
||||
if completed is None or job.created_at > completed.created_at:
|
||||
completed = job
|
||||
return running or completed
|
||||
|
||||
def update(self, job: IngestJob) -> IngestJob:
|
||||
self._jobs[job.id] = job
|
||||
return job
|
||||
@@ -420,3 +439,172 @@ class TestMultipartUploadIdempotency:
|
||||
assert ingest_repo.created_count == 1
|
||||
# OSS 上传只发生一次(第二次在幂等检查处直接返回)
|
||||
assert storage.upload_file.call_count == 1
|
||||
|
||||
|
||||
# ── P1 修复(#2092):占位 asset 被误判为 duplicate → 素材永久 processing ──
|
||||
|
||||
|
||||
class TestPlaceholderAssetNotTreatedAsDuplicate:
|
||||
"""prepare 建了 PROCESSING 占位但还没 ingest,complete 必须补提 ingest 而不是短路返 duplicated。"""
|
||||
|
||||
def test_complete_hits_prepare_placeholder_without_job_submits_ingest(self):
|
||||
"""场景:/direct/prepare 建了 PROCESSING 占位(同 client_upload_id),complete 命中后应补提 ingest。"""
|
||||
from packages.domain.classification import IngestJobStatus
|
||||
|
||||
# 预建占位 asset(prepare 建的,PROCESSING,无 ingest job)
|
||||
placeholder = Asset(
|
||||
id="asset-placeholder",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="IMG_9999.MOV",
|
||||
storage_key="uploads/prepare/IMG_9999.MOV",
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.PROCESSING,
|
||||
client_upload_id="tok-prep",
|
||||
)
|
||||
# 注意:占位的 storage_key 是 prepare 生成的 key,complete 传入的是用户实际上传的 key
|
||||
client, asset_repo, ingest_repo, storage = _client(
|
||||
asset_repo=StubAssetRepository([placeholder]),
|
||||
)
|
||||
storage._normalize_storage_key = lambda k: k # complete 用自己的 key
|
||||
storage.get_url = lambda k: f"https://oss/{k}"
|
||||
|
||||
body = {
|
||||
"project_id": "proj-1",
|
||||
"library_id": "lib-1",
|
||||
"storage_key": "uploads/actual/IMG_9999.MOV", # complete 用真实上传 key
|
||||
"client_upload_id": "tok-prep", # 命中占位
|
||||
"file_size": 1024,
|
||||
}
|
||||
r = client.post("/api/v1/direct/complete", json=body)
|
||||
assert r.status_code == 200, r.text
|
||||
b = r.json()
|
||||
# 关键:不是 duplicate;返回 job_id;占位被复用(不新建 asset)
|
||||
assert b["duplicated"] is False, f"占位被误判为 duplicate: {b}"
|
||||
assert b["ingest_job_id"], "应补提 ingest job"
|
||||
assert b["asset_id"] == "asset-placeholder"
|
||||
# 不新建 asset(占位复用)
|
||||
assert len(asset_repo.created) == 0
|
||||
# job 被提交
|
||||
assert ingest_repo.created_count == 1
|
||||
|
||||
def test_ready_asset_treated_as_true_duplicate(self):
|
||||
"""READY 素材命中 → 真重复,返 duplicated 且不提交新 job。"""
|
||||
ready = Asset(
|
||||
id="asset-ready",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="done.MOV",
|
||||
storage_key="uploads/done/done.MOV",
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.READY,
|
||||
client_upload_id="tok-done",
|
||||
)
|
||||
client, asset_repo, ingest_repo, _storage = _client(asset_repo=StubAssetRepository([ready]))
|
||||
body = {
|
||||
"project_id": "proj-1",
|
||||
"library_id": "lib-1",
|
||||
"storage_key": "uploads/done/done.MOV",
|
||||
"client_upload_id": "tok-done",
|
||||
}
|
||||
r = client.post("/api/v1/direct/complete", json=body)
|
||||
assert r.status_code == 200
|
||||
b = r.json()
|
||||
assert b["duplicated"] is True
|
||||
assert b["asset_id"] == "asset-ready"
|
||||
assert ingest_repo.created_count == 0
|
||||
|
||||
def test_processing_asset_with_existing_job_is_idempotent_duplicate(self):
|
||||
"""PROCESSING 但已有在跑 job → 幂等重试,返 duplicated + 已有 job_id,不重复提交。"""
|
||||
processing = Asset(
|
||||
id="asset-running",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="running.MOV",
|
||||
storage_key="uploads/run/running.MOV",
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.PROCESSING,
|
||||
client_upload_id="tok-run",
|
||||
)
|
||||
client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([processing]))
|
||||
# 预置一个在跑 job
|
||||
existing_job = IngestJob(
|
||||
id="job-existing",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
storage_key="uploads/run/running.MOV",
|
||||
asset_id="asset-running",
|
||||
)
|
||||
ingest_repo.create(existing_job)
|
||||
before = ingest_repo.created_count
|
||||
|
||||
body = {
|
||||
"project_id": "proj-1",
|
||||
"library_id": "lib-1",
|
||||
"storage_key": "uploads/run/running.MOV",
|
||||
"client_upload_id": "tok-run",
|
||||
}
|
||||
r = client.post("/api/v1/direct/complete", json=body)
|
||||
assert r.status_code == 200
|
||||
b = r.json()
|
||||
assert b["duplicated"] is True
|
||||
assert b["ingest_job_id"] == "job-existing"
|
||||
# 没有新建 job
|
||||
assert ingest_repo.created_count == before
|
||||
assert len(asset_repo.created) == 0
|
||||
|
||||
def test_error_asset_allows_reingest(self):
|
||||
"""ERROR 状态素材命中 → 不视为 duplicate,重新走 ingest。"""
|
||||
errored = Asset(
|
||||
id="asset-err",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="err.MOV",
|
||||
storage_key="uploads/err/err.MOV",
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.ERROR,
|
||||
client_upload_id="tok-err",
|
||||
)
|
||||
client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([errored]))
|
||||
body = {
|
||||
"project_id": "proj-1",
|
||||
"library_id": "lib-1",
|
||||
"storage_key": "uploads/err/err.MOV",
|
||||
"client_upload_id": "tok-err",
|
||||
}
|
||||
r = client.post("/api/v1/direct/complete", json=body)
|
||||
assert r.status_code == 200
|
||||
b = r.json()
|
||||
assert b["duplicated"] is False
|
||||
assert b["ingest_job_id"]
|
||||
assert ingest_repo.created_count == 1
|
||||
|
||||
def test_multipart_placeholder_without_job_submits_ingest(self):
|
||||
"""multipart 上传命中 PROCESSING 占位且无 job → 补提 ingest(不短路返 duplicated)。"""
|
||||
placeholder = Asset(
|
||||
id="asset-mp-placeholder",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="mp.MOV",
|
||||
storage_key="uploads/mp/mp.MOV",
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.PROCESSING,
|
||||
client_upload_id="tok-mp",
|
||||
)
|
||||
client, asset_repo, ingest_repo, storage = _client(asset_repo=StubAssetRepository([placeholder]))
|
||||
storage.upload_file = MagicMock(return_value="https://oss/mp.MOV")
|
||||
storage.get_url = MagicMock(return_value="https://oss/mp.MOV")
|
||||
r = client.post(
|
||||
"/api/v1",
|
||||
data={
|
||||
"project_id": "proj-1",
|
||||
"library_id": "lib-1",
|
||||
"client_upload_id": "tok-mp",
|
||||
},
|
||||
files={"file": ("mp.MOV", b"data", "video/quicktime")},
|
||||
)
|
||||
assert r.status_code == 200, r.text
|
||||
b = r.json()
|
||||
assert b["duplicated"] is False, f"multipart 占位被误判为 duplicate: {b}"
|
||||
assert b["ingest_job_id"]
|
||||
assert ingest_repo.created_count == 1
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
|
||||
|
||||
验证:
|
||||
- 3 张新表可在内存 SQLite 上创建
|
||||
- 默认值与基本 CRUD 正常
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
Base,
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
|
||||
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
|
||||
|
||||
class TestViralVideoJobModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(
|
||||
id="job-001",
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_default_values(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(id="job-002", user_id="user-002")
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoJobModel, "job-002")
|
||||
assert fetched.images == []
|
||||
assert fetched.fusion_level == "ai_polish"
|
||||
assert fetched.style_strength == "medium"
|
||||
assert fetched.status == "pending"
|
||||
assert fetched.credits_cost == 0
|
||||
assert fetched.retry_count == 0
|
||||
assert fetched.style_guide is None
|
||||
assert fetched.intent_result is None
|
||||
|
||||
|
||||
class TestViralVideoStyleTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
style_config={"cut_speed": "fast"},
|
||||
sort_order=1,
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
|
||||
assert fetched.name == "快节奏"
|
||||
assert fetched.style_config == {"cut_speed": "fast"}
|
||||
assert fetched.sort_order == 1
|
||||
|
||||
|
||||
class TestViralVideoPromptTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoPromptTemplateModel(
|
||||
id="pt-001",
|
||||
prompt_type="image_analysis",
|
||||
name="图片分析模板",
|
||||
content="请分析图片:{image_url}",
|
||||
variables=["image_url"],
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
|
||||
assert fetched.prompt_type == "image_analysis"
|
||||
assert fetched.content == "请分析图片:{image_url}"
|
||||
assert fetched.variables == ["image_url"]
|
||||
assert fetched.version == 1
|
||||
assert fetched.is_active is True
|
||||
@@ -65,10 +65,24 @@ def test_generate_video_preserves_original_function():
|
||||
assert original.__name__ == "generate_video", f"expected __name__='generate_video', got '{original.__name__}'"
|
||||
|
||||
|
||||
def test_sync_task_config_to_plan_is_plain_function():
|
||||
"""Helper must NOT be registered as a Celery task."""
|
||||
from worker_app.tasks.generation import _sync_task_config_to_plan
|
||||
def test_build_task_config_override_is_plain_function():
|
||||
"""#2098: 原 _sync_task_config_to_plan 已拆分为 _build_task_config_override + _download_voice_for_task,
|
||||
均为普通函数,不应被注册为 Celery task。"""
|
||||
from worker_app.tasks.generation import _build_task_config_override, _download_voice_for_task
|
||||
|
||||
assert not hasattr(
|
||||
_sync_task_config_to_plan, "run"
|
||||
), "_sync_task_config_to_plan must be a plain function, not a Celery task"
|
||||
for fn in (_build_task_config_override, _download_voice_for_task):
|
||||
assert not hasattr(fn, "run"), f"{fn.__name__} must be a plain function, not a Celery task"
|
||||
|
||||
# Bug A: override 对 title_config 做 key 归一化 (font_size→size, font_color→color)
|
||||
override = _build_task_config_override(
|
||||
{
|
||||
"title_config": {"font_size": 48, "font_color": "#ff0000", "text": "hi"},
|
||||
"bgm_config": {"url": "http://x/bgm.mp3"},
|
||||
"output_width": 1080,
|
||||
"output_height": 1920,
|
||||
}
|
||||
)
|
||||
assert override["title"]["size"] == 48
|
||||
assert override["title"]["color"] == "#ff0000"
|
||||
assert override["bgm"]["url"] == "http://x/bgm.mp3"
|
||||
assert override["export"]["resolution"] == "1080x1920"
|
||||
|
||||
Reference in New Issue
Block a user