diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 72a4de854..c42cbbdb4 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -380,31 +380,37 @@ def _download_library_assets( session = SessionLocal() try: - # 构建查询:根据模式选择不同的过滤条件 + # 构建查询 query = session.query(AssetModel).filter( AssetModel.status == "ready", AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), ) - if asset_library_id: - # 素材库模式 - query = query.filter(AssetModel.asset_library_id == asset_library_id) + if asset_ids: + # 明确指定了 asset_ids:直接按 ID 查,不预先按 library/project 过滤 + # 避免项目级素材或跨库素材因为 library_id 不匹配而查不到 + # 归属安全由后面的归属校验保证 + query = query.filter(AssetModel.id.in_(asset_ids)) logger.info( - "下载素材库视频: asset_library_id=%s asset_ids=%s", - asset_library_id, - asset_ids or "all", + "下载指定素材: asset_ids=%d 个, asset_library_id=%s, project_id=%s", + len(asset_ids), + asset_library_id or "none", + project_id or "none", ) else: - # 项目级模式 - query = query.filter(AssetModel.project_id == project_id) - logger.info( - "下载项目级视频: project_id=%s asset_ids=%s", - project_id, - asset_ids or "all", - ) - - if asset_ids: - query = query.filter(AssetModel.id.in_(asset_ids)) + # 未指定 asset_ids:按 library 或 project 下载全部 ready 视频 + if asset_library_id: + query = query.filter(AssetModel.asset_library_id == asset_library_id) + logger.info( + "下载素材库全部视频: asset_library_id=%s", + asset_library_id, + ) + else: + query = query.filter(AssetModel.project_id == project_id) + logger.info( + "下载项目全部视频: project_id=%s", + project_id, + ) assets = query.order_by(AssetModel.created_at).all() @@ -421,13 +427,15 @@ def _download_library_assets( if missing_ids: raise ValueError(f"素材不存在: asset_ids={sorted(missing_ids)}") for asset in assets: + # 校验素材库归属(只要传了 asset_library_id 就校验) if asset_library_id and asset.asset_library_id != asset_library_id: raise ValueError( f"素材不属于指定素材库: asset_id={asset.id}, " f"expected_asset_library_id={asset_library_id}, " f"actual_asset_library_id={asset.asset_library_id}" ) - if not asset_library_id and project_id and asset.project_id != project_id: + # 校验项目归属(只要传了 project_id 就校验) + if project_id and asset.project_id != project_id: raise ValueError( f"素材不属于指定项目: asset_id={asset.id}, " f"expected_project_id={project_id}, " diff --git a/tests/unit/test_oneclick_gen_p0_fixes.py b/tests/unit/test_oneclick_gen_p0_fixes.py index 7bfd58ff0..efde21b16 100644 --- a/tests/unit/test_oneclick_gen_p0_fixes.py +++ b/tests/unit/test_oneclick_gen_p0_fixes.py @@ -132,10 +132,8 @@ class TestDownloadLibraryAssets: session.query.return_value = query filter_result = MagicMock() query.filter.return_value = filter_result - in_filter = MagicMock() - filter_result.filter.return_value = in_filter id_filter = MagicMock() - in_filter.filter.return_value = id_filter + filter_result.filter.return_value = id_filter assets = [self._make_asset("a1", "video/a1.mp4")] id_filter.order_by.return_value.all.return_value = assets