From 1ea8fd39892bf27a9998b2f2d30881a040d7c524 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=81=B5=E5=BA=94?= Date: Tue, 7 Jul 2026 15:13:05 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20P2=20=E7=B4=A0=E6=9D=90=E5=BA=93?= =?UTF-8?q?=E8=87=AA=E5=8A=A8=E5=8C=B9=E9=85=8D=20-=20=E6=94=AF=E6=8C=81?= =?UTF-8?q?=20all/random/smart=20=E4=B8=89=E7=A7=8D=E7=B4=A0=E6=9D=90?= =?UTF-8?q?=E9=80=89=E5=8F=96=E6=A8=A1=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Schema: 新增 asset_select_mode + asset_select_count 字段 - Domain: GenerationTask 新增 asset_select_mode 字段 - Application: Command/UseCase 透传 asset_select_mode - Route: _select_assets_from_library() 辅助函数 + 创建任务集成 - SQLAlchemy: Model/Repository 映射 asset_select_mode - Alembic: 032 号迁移 - Worker: _download_library_assets() 支持 asset_ids 过滤 - 单测: 15 个测试覆盖三种模式 + 边界情况 --- ...d_asset_select_mode_to_generation_tasks.py | 28 +++ apps/api/app/api/routes/generation_tasks.py | 58 +++++- apps/api/app/schemas/generation_task.py | 9 + apps/worker/worker_app/tasks/generation.py | 26 +-- .../generation_task_repository.py | 3 + packages/adapters/sqlalchemy_impl/models.py | 1 + packages/application/generation_tasks.py | 2 + packages/domain/generation_task.py | 3 + tests/unit/test_asset_select_mode.py | 176 ++++++++++++++++++ 9 files changed, 293 insertions(+), 13 deletions(-) create mode 100644 alembic/versions/032_add_asset_select_mode_to_generation_tasks.py create mode 100644 tests/unit/test_asset_select_mode.py diff --git a/alembic/versions/032_add_asset_select_mode_to_generation_tasks.py b/alembic/versions/032_add_asset_select_mode_to_generation_tasks.py new file mode 100644 index 000000000..842863817 --- /dev/null +++ b/alembic/versions/032_add_asset_select_mode_to_generation_tasks.py @@ -0,0 +1,28 @@ +"""Add asset_select_mode to generation_tasks + +Revision ID: 032 +Revises: 031 +Create Date: 2026-07-07 + +素材库自动匹配功能:为 generation_tasks 表添加 asset_select_mode 字段。 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "032" +down_revision = "031" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "generation_tasks", + sa.Column("asset_select_mode", sa.String(20), nullable=False, server_default=""), + ) + + +def downgrade() -> None: + op.drop_column("generation_tasks", "asset_select_mode") diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 73239b54a..c68c82fd5 100644 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -1,3 +1,4 @@ +import random from typing import Any from app.auth import AuthenticatedUser, get_current_user @@ -52,6 +53,7 @@ def _to_generation_task_response(task) -> GenerationTaskResponse: title_ids=task.title_ids, voice_ids=task.voice_ids, source_edit_plan_id=task.source_edit_plan_id or "", + asset_select_mode=getattr(task, "asset_select_mode", ""), status=task.status, progress=task.progress, result_count=task.result_count, @@ -86,6 +88,49 @@ def _ensure_library_has_ready_video_assets(assets) -> None: ) +def _select_assets_from_library( + assets: list, + mode: str, + count: int, +) -> list[str]: + """根据选取模式从素材库中选取 ready 状态的视频素材 ID。 + + Args: + assets: 素材库中所有素材(Asset 实体列表) + mode: 选取模式 — all=全部, random=随机, smart=按质量评分 + count: 选取数量,0 表示全部(仅 random/smart 模式有效) + + Returns: + 选中的素材 ID 列表 + """ + ready_video_assets = [a for a in assets if a.status.value == "ready" and a.mime_type.startswith("video")] + + if not ready_video_assets: + return [] + + if mode == "random": + selected = ( + ready_video_assets if count <= 0 else random.sample(ready_video_assets, min(count, len(ready_video_assets))) + ) + return [a.id for a in selected] + + if mode == "smart": + # 按质量分降序排列(质量分高的优先),质量分相同时按时长降序 + sorted_assets = sorted( + ready_video_assets, + key=lambda a: ( + a.quality_score if a.quality_score is not None else 0.0, + a.duration if a.duration is not None else 0.0, + ), + reverse=True, + ) + selected = sorted_assets if count <= 0 else sorted_assets[:count] + return [a.id for a in selected] + + # 默认 all 模式:返回全部 ready 视频素材 + return [a.id for a in ready_video_assets] + + def _resolve_project_and_library( request: CreateGenerationTaskRequest, project_repository: Any, @@ -137,6 +182,7 @@ def create_generation_task( ) # asset_library 存在性校验(仅在提供了 asset_library_id 时) + resolved_asset_ids: list[str] = list(request.asset_ids) if asset_library_id: library = asset_library_repository.get(asset_library_id) if library is None or (project_id and library.project_id != project_id): @@ -145,6 +191,14 @@ def create_generation_task( assets = asset_repository.find_by_library(asset_library_id) _ensure_library_has_ready_video_assets(assets) + # 素材库自动匹配:当未显式指定 asset_ids 时,按模式自动选取 + if not resolved_asset_ids: + resolved_asset_ids = _select_assets_from_library( + assets, + mode=request.asset_select_mode, + count=request.asset_select_count, + ) + use_case = CreateGenerationTaskUseCase(generation_task_repository) count = request.count created_tasks = [] @@ -157,11 +211,12 @@ def create_generation_task( strategy_id=request.strategy_id, voice_library_id=request.voice_library_id, template_id=request.template_id, - asset_ids=request.asset_ids, + asset_ids=resolved_asset_ids, title_ids=request.title_ids, voice_ids=request.voice_ids, created_by_user_id=authenticated_user.user.id, source_edit_plan_id=request.source_edit_plan_id, + asset_select_mode=request.asset_select_mode, ) ) celery_app.send_task("worker.generate_video", args=[task.id]) @@ -245,6 +300,7 @@ def retry_generation_task( voice_ids=task.voice_ids, created_by_user_id=authenticated_user.user.id, source_edit_plan_id=task.source_edit_plan_id or "", + asset_select_mode=getattr(task, "asset_select_mode", ""), ) ) celery_app.send_task("worker.generate_video", args=[retried.id]) diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index a3aab1007..ccd2ad45b 100644 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -23,6 +23,14 @@ class CreateGenerationTaskRequest(BaseModel): source_edit_plan_id: str = "" # ── 批量生成 ── count: int = Field(default=1, ge=1, le=50, description="批量生成数量,默认1,最大50") + # ── 素材库自动匹配 ── + asset_select_mode: str = Field( + default="all", + description="素材选取模式:all=全部ready视频, random=随机选取, smart=智能匹配(按质量/时长评分)", + ) + asset_select_count: int = Field( + default=0, ge=0, le=100, description="选取数量,0表示全部(仅 random/smart 模式有效)" + ) @model_validator(mode="after") def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest": @@ -48,6 +56,7 @@ class GenerationTaskResponse(BaseModel): title_ids: list[str] = Field(default_factory=list) voice_ids: list[str] = Field(default_factory=list) source_edit_plan_id: str = "" + asset_select_mode: str = "" status: str progress: float result_count: int diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 373a29913..e36191c5c 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -139,14 +139,16 @@ def _download_library_assets( asset_library_id: str, temp_path: Path, video_extensions: tuple = (".mp4", ".mov", ".avi", ".mkv", ".webm"), + asset_ids: list[str] | None = None, ) -> list[str]: """ - 从素材库下载所有视频素材 + 从素材库下载视频素材 Args: asset_library_id: 素材库 ID temp_path: 临时目录路径 video_extensions: 支持的视频扩展名 + asset_ids: 指定素材 ID 列表,为空则下载全部 ready 视频素材 Returns: 下载成功的视频文件路径列表 @@ -161,16 +163,15 @@ def _download_library_assets( try: # 查询素材库中的视频素材 - assets = ( - session.query(AssetModel) - .filter( - AssetModel.asset_library_id == asset_library_id, - AssetModel.status == "ready", - AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), - ) - .order_by(AssetModel.created_at) - .all() + query = session.query(AssetModel).filter( + AssetModel.asset_library_id == asset_library_id, + AssetModel.status == "ready", + AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), ) + # 如果指定了 asset_ids,则只下载这些素材 + if asset_ids: + query = query.filter(AssetModel.id.in_(asset_ids)) + assets = query.order_by(AssetModel.created_at).all() if not assets: logger.info(f"No video assets found in library {asset_library_id}") @@ -259,6 +260,7 @@ def generate_video(self, task_id: str) -> dict: asset_library_id = gen_task.asset_library_id voice_library_id = gen_task.voice_library_id or "" mode = gen_task.strategy_id or "one_take" + task_asset_ids = list(gen_task.asset_ids or []) finally: session.close() @@ -275,8 +277,8 @@ def generate_video(self, task_id: str) -> dict: temp_path = Path(temp_dir) output_path = temp_path / output_name - # 从素材库下载视频素材 - downloaded_videos = _download_library_assets(asset_library_id, temp_path) + # 从素材库下载视频素材(如果任务指定了 asset_ids 则只下载这些) + downloaded_videos = _download_library_assets(asset_library_id, temp_path, asset_ids=task_asset_ids or None) audio_path = None if voice_library_id: diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 12c439292..ea0aa984d 100644 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -25,6 +25,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask: completed_at=model.completed_at, created_by_user_id=model.created_by_user_id, source_edit_plan_id=model.source_edit_plan_id or "", + asset_select_mode=model.asset_select_mode or "", created_at=model.created_at, ) @@ -52,6 +53,7 @@ class SQLAlchemyGenerationTaskRepository: completed_at=task.completed_at, created_by_user_id=task.created_by_user_id, source_edit_plan_id=task.source_edit_plan_id or None, + asset_select_mode=task.asset_select_mode or "", created_at=task.created_at, ) self.session.add(model) @@ -123,5 +125,6 @@ class SQLAlchemyGenerationTaskRepository: model.started_at = task.started_at model.completed_at = task.completed_at model.source_edit_plan_id = task.source_edit_plan_id or None + model.asset_select_mode = task.asset_select_mode or "" self.session.commit() return task diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 508dc5081..60a1e6b9a 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -253,6 +253,7 @@ class GenerationTaskModel(Base): completed_at = Column(DateTime, nullable=True) created_by_user_id = Column(String(32), nullable=False, default="", index=True) source_edit_plan_id = Column(String(32), nullable=True, index=True) + asset_select_mode = Column(String(20), nullable=False, default="") extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) diff --git a/packages/application/generation_tasks.py b/packages/application/generation_tasks.py index b3a20a30e..3ed4bb8fb 100644 --- a/packages/application/generation_tasks.py +++ b/packages/application/generation_tasks.py @@ -19,6 +19,7 @@ class CreateGenerationTaskCommand: voice_ids: list[str] = field(default_factory=list) created_by_user_id: str = "" source_edit_plan_id: str = "" + asset_select_mode: str = "" class CreateGenerationTaskUseCase: @@ -44,6 +45,7 @@ class CreateGenerationTaskUseCase: completed_at=None, created_by_user_id=command.created_by_user_id, source_edit_plan_id=command.source_edit_plan_id, + asset_select_mode=command.asset_select_mode, ) return self.generation_task_repository.create(task) diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index 7df475d27..cb7f1e84c 100644 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -43,6 +43,7 @@ class GenerationTask: completed_at: datetime | None = None source_edit_plan_id: str = "" created_by_user_id: str = "" + asset_select_mode: str = "" created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) @classmethod @@ -59,6 +60,7 @@ class GenerationTask: voice_ids: list[str] | None = None, created_by_user_id: str = "", source_edit_plan_id: str = "", + asset_select_mode: str = "", ) -> "GenerationTask": if not project_id.strip() and not template_id.strip(): raise ValueError("project_id 或 template_id 至少需要提供一个") @@ -76,4 +78,5 @@ class GenerationTask: voice_ids=list(voice_ids) if voice_ids else [], created_by_user_id=created_by_user_id.strip(), source_edit_plan_id=source_edit_plan_id.strip(), + asset_select_mode=asset_select_mode, ) diff --git a/tests/unit/test_asset_select_mode.py b/tests/unit/test_asset_select_mode.py new file mode 100644 index 000000000..98469e7fa --- /dev/null +++ b/tests/unit/test_asset_select_mode.py @@ -0,0 +1,176 @@ +""" +素材库自动匹配 单元测试 + +覆盖: +- all 模式:返回全部 ready 视频素材 ID +- random 模式:随机选取 N 个 +- smart 模式:按质量分/时长评分降序选取 +- 无 ready 视频素材时返回空列表 +- count=0 时返回全部(random/smart 模式) +- 非视频素材和非 ready 状态素材被过滤 +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +from app.api.routes.generation_tasks import _select_assets_from_library + +from packages.domain import Asset, AssetStatus + + +def _asset( + id: str, + name: str, + mime_type: str = "video/mp4", + status: AssetStatus = AssetStatus.READY, + quality_score: float | None = None, + duration: float | None = None, +) -> Asset: + a = Asset.create( + project_id="proj-1", + library_id="lib-1", + name=name, + storage_key=f"uploads/{name}", + mime_type=mime_type, + file_size=1024, + status=status, + quality_score=quality_score, + duration=duration, + ) + # create() 会覆盖 id,手动设置 + a.id = id + return a + + +class TestSelectAssetsAllMode: + """all 模式:返回全部 ready 视频素材。""" + + def test_returns_all_ready_video_assets(self): + assets = [ + _asset("a1", "v1.mp4"), + _asset("a2", "v2.mp4"), + _asset("a3", "v3.mp4"), + ] + result = _select_assets_from_library(assets, mode="all", count=0) + assert sorted(result) == ["a1", "a2", "a3"] + + def test_ignores_count_in_all_mode(self): + assets = [ + _asset("a1", "v1.mp4"), + _asset("a2", "v2.mp4"), + ] + result = _select_assets_from_library(assets, mode="all", count=1) + assert len(result) == 2 + + def test_filters_non_video_assets(self): + assets = [ + _asset("a1", "v1.mp4", mime_type="video/mp4"), + _asset("a2", "img.jpg", mime_type="image/jpeg"), + _asset("a3", "v2.mov", mime_type="video/quicktime"), + ] + result = _select_assets_from_library(assets, mode="all", count=0) + assert sorted(result) == ["a1", "a3"] + + def test_filters_non_ready_assets(self): + assets = [ + _asset("a1", "v1.mp4", status=AssetStatus.READY), + _asset("a2", "v2.mp4", status=AssetStatus.UPLOADING), + _asset("a3", "v3.mp4", status=AssetStatus.PROCESSING), + ] + result = _select_assets_from_library(assets, mode="all", count=0) + assert result == ["a1"] + + def test_empty_library_returns_empty(self): + result = _select_assets_from_library([], mode="all", count=0) + assert result == [] + + def test_no_ready_video_returns_empty(self): + assets = [ + _asset("a1", "v1.mp4", status=AssetStatus.UPLOADING), + _asset("a2", "img.jpg", mime_type="image/jpeg", status=AssetStatus.READY), + ] + result = _select_assets_from_library(assets, mode="all", count=0) + assert result == [] + + +class TestSelectAssetsRandomMode: + """random 模式:随机选取 N 个。""" + + def test_random_selects_exact_count(self): + assets = [_asset(f"a{i}", f"v{i}.mp4") for i in range(10)] + result = _select_assets_from_library(assets, mode="random", count=3) + assert len(result) == 3 + assert all(rid in [a.id for a in assets] for rid in result) + + def test_random_count_zero_returns_all(self): + assets = [_asset(f"a{i}", f"v{i}.mp4") for i in range(5)] + result = _select_assets_from_library(assets, mode="random", count=0) + assert len(result) == 5 + + def test_random_count_exceeds_total_returns_all(self): + assets = [_asset(f"a{i}", f"v{i}.mp4") for i in range(3)] + result = _select_assets_from_library(assets, mode="random", count=100) + assert len(result) == 3 + + +class TestSelectAssetsSmartMode: + """smart 模式:按质量分/时长评分降序选取。""" + + def test_smart_sorts_by_quality_score_desc(self): + assets = [ + _asset("low", "low.mp4", quality_score=0.3), + _asset("high", "high.mp4", quality_score=0.9), + _asset("mid", "mid.mp4", quality_score=0.6), + ] + result = _select_assets_from_library(assets, mode="smart", count=0) + assert result == ["high", "mid", "low"] + + def test_smart_tiebreak_by_duration_desc(self): + assets = [ + _asset("short", "short.mp4", quality_score=0.8, duration=10.0), + _asset("long", "long.mp4", quality_score=0.8, duration=60.0), + ] + result = _select_assets_from_library(assets, mode="smart", count=0) + assert result == ["long", "short"] + + def test_smart_with_count_limits_results(self): + assets = [ + _asset("a1", "v1.mp4", quality_score=0.9), + _asset("a2", "v2.mp4", quality_score=0.7), + _asset("a3", "v3.mp4", quality_score=0.5), + ] + result = _select_assets_from_library(assets, mode="smart", count=2) + assert result == ["a1", "a2"] + + def test_smart_null_quality_treated_as_zero(self): + assets = [ + _asset("scored", "scored.mp4", quality_score=0.5), + _asset("unscored", "unscored.mp4", quality_score=None), + ] + result = _select_assets_from_library(assets, mode="smart", count=0) + assert result == ["scored", "unscored"] + + def test_smart_count_zero_returns_all_sorted(self): + assets = [ + _asset("a1", "v1.mp4", quality_score=0.1), + _asset("a2", "v2.mp4", quality_score=0.9), + _asset("a3", "v3.mp4", quality_score=0.5), + ] + result = _select_assets_from_library(assets, mode="smart", count=0) + assert result == ["a2", "a3", "a1"] + + +class TestSelectAssetsDefaultMode: + """默认模式(未知 mode 字符串)应回退到 all。""" + + def test_unknown_mode_falls_back_to_all(self): + assets = [ + _asset("a1", "v1.mp4"), + _asset("a2", "v2.mp4"), + ] + result = _select_assets_from_library(assets, mode="unknown", count=0) + assert len(result) == 2