feat: implement POST /assets/smart-match endpoint #1242

Merged
auto-approve-bot merged 5 commits from feat/smart-match-assets into develop 2026-08-05 09:51:01 +08:00
9 changed files with 1623 additions and 6 deletions
+222
View File
@@ -0,0 +1,222 @@
---
AIGC:
Label: "1"
ContentProducer: 001191110102MACQD9K64018705
ProduceID: 15868733686388_0/project_7655981463858544923-files/docs/1197_preview_generation_proposal.md
ReservedCode1: ""
ContentPropagator: 001191110102MACQD9K64028705
PropagateID: 15868733686388#1785468313901
ReservedCode2: ""
---
# #1197 预览生成接口方案评估
## 背景
智能剪辑「一键生成」流程中,第3步预览生成当前被跳过,直接进入下一步。需要实现真正的预览生成功能,让用户在正式生成前能看到效果预览。
## 现状分析
### 现有生成链路
```
API 触发生成 → GenerationTask入库 → Celery异步任务 → UnifiedRenderService渲染 → OSS上传 → 更新状态
```
**关键节点:**
1. **API层**`POST /generation-tasks``POST /templates/{id}/generate` 触发生成
2. **任务调度**Celery task `worker.generate_video`
3. **渲染引擎**`UnifiedRenderService`(统一渲染引擎,已接入9个效果层)
4. **输出配置**:默认 720p (1280x720),支持 `resolution` 字段自定义
5. **产物存储**`GeneratedVideo` 表记录,OSS 存储视频文件
### 已有可复用能力
| 能力 | 位置 | 是否可复用 |
|------|------|-----------|
| 任务创建与状态管理 | `GenerationTask` + `CreateGenerationTaskUseCase` | ✅ 是 |
| 素材下载与预处理 | `_download_video_assets` / `_download_voice_asset` | ✅ 是 |
| 统一渲染引擎 | `UnifiedRenderService` | ✅ 是 |
| 分辨率配置 | `resolution` 字段已支持 | ✅ 是 |
| 混音与后处理 | `_render_video` 内流程 | ✅ 是 |
| OSS 上传与查重 | `_upload_and_dedup` | ✅ 是 |
| 进度追踪 | `append_log` / `progress` 字段 | ✅ 是 |
## 方案对比
### 方案A:复用现有生成链路 + is_preview 标记(推荐)
**思路**:在现有 GenerationTask 上加 `is_preview` 标记,预览生成走完整链路但参数降级。
**改动点:**
1. **数据模型**`GenerationTask``is_preview: bool` 字段(默认 false);`GeneratedVideo``is_preview: bool`
2. **API 层**:生成接口加 `is_preview` 参数,预览任务不计入配额
3. **渲染参数**:预览模式下自动调整
- 分辨率:480p (854x480)
- 时长:限制前 15 秒(或模板第一个片段)
- 码率:降低至 1.5Mbps(正式 4Mbps
- 效果层:跳过高级转场/粒子特效等耗时效果
4. **任务调度**:预览任务走低优先级队列(或复用现有队列,标记优先级)
5. **前端对接**:预览生成结果带 `is_preview=true` 标记,前端展示"预览"标签
**优点:**
- 代码复用率 90%+,改动最小
- 与正式生成逻辑一致,预览效果真实可信
- 进度查询、结果展示等功能直接复用
- 后续可平滑升级:预览满意后一键转正式生成
**缺点:**
- 需要区分预览和正式任务,避免数据混淆
- 预览任务和正式任务竞争同一队列资源(可后续优化为独立队列)
**开发量估算**2-3 天
- 数据模型 + 迁移:0.5 天
- API 层改造:0.5 天
- 渲染参数降级:1 天
- 测试 + 联调:1 天
---
### 方案B:新建独立预览接口 + 轻量渲染逻辑
**思路**:新建独立的预览生成接口,使用简化的渲染逻辑(如只拼接素材+基础配音,跳过大部分效果)。
**改动点:**
1. 新增 `PreviewTask` 数据模型
2. 新增 `POST /api/v1/preview/generate` 接口
3. 新增独立的 Celery task `worker.generate_preview`
4. 简化渲染流程:只做素材裁剪+拼接+配音,跳过转场/滤镜/字幕特效等
**优点:**
- 完全隔离,不影响正式生成链路
- 可以做极致优化,预览生成速度快
- 数据模型清晰,不会混淆
**缺点:**
- 代码重复率高,两套生成逻辑维护成本翻倍
- 预览效果与正式生成可能不一致(效果层差异)
- 前端需要对接两套接口
- 无法从预览升级为正式生成(需重新走完整流程)
**开发量估算**4-5 天
- 数据模型 + 接口:1 天
- 简化渲染逻辑:2 天
- 测试 + 联调:1-2 天
---
### 方案C:图片预览(首帧/关键帧截图)
**思路**:不生成视频,只生成几张关键帧的预览图片。
**优点:**
- 生成速度极快(秒级)
- 资源消耗小
**缺点:**
- 预览效果差,用户无法感知动态效果
- 无法验证配音、转场、节奏等时间维度的效果
- 用户体验不佳,不如"真预览"有说服力
**开发量估算**1-2 天
---
## 推荐方案:方案A(复用现有生成链路)
### 核心理由
1. **效果保真**:预览和正式生成用同一套渲染引擎,效果一致,用户信任度高
2. **开发效率**90% 代码复用,2-3 天可上线
3. **可扩展性强**:后续可加「预览转正式」「低分辨率快速预览」等增强功能
4. **维护成本低**:一套生成逻辑,bug 修复和新功能同时生效
### 详细设计
#### 1. 数据模型变更
```python
# GenerationTask 新增字段
is_preview: bool = False
"""是否为预览生成"""
preview_of: str = ""
"""预览对应的正式任务 ID(或反向关联)"""
# GeneratedVideo 新增字段
is_preview: bool = False
"""是否为预览视频"""
```
**迁移**alembic 新增 migration,两个表各加 1-2 个字段。
#### 2. API 层
```
POST /api/v1/generation-tasks
Body 增加 is_preview: bool = false
POST /api/v1/templates/{id}/generate
Query 增加 is_preview: bool = false
```
**配额处理**:预览生成不计入用户配额,不占用生成次数限制。
#### 3. 渲染参数降级
| 参数 | 正式生成 | 预览生成 |
|------|---------|---------|
| 分辨率 | 720p (1280x720) | 480p (854x480) |
| 码率 | 4 Mbps | 1.5 Mbps |
| 时长 | 完整时长 | 前 15 秒(或第一段) |
| 帧率 | 30 fps | 24 fps |
| 转场效果 | 完整转场 | 仅淡入淡出(或简单切) |
| 特效滤镜 | 全部启用 | 跳过粒子/光效等高级效果 |
| 字幕 | 完整渲染 | 正常渲染(字幕是核心信息) |
| 配音 | 完整混音 | 正常混音(配音是核心信息) |
**实现方式**:在 `_render_video` 或 UnifiedRenderService 入口处,根据 `is_preview` 标记调整渲染配置。
#### 4. 任务调度
- 初期复用现有队列,预览任务正常排队
- 后续如需优化,可拆分独立预览队列(低优先级)
- 预览任务可设置较短超时时间
#### 5. 前端对接
- 调用生成接口时传 `is_preview=true`
- 结果列表中预览视频带「预览」标签
- 预览满意后可一键「升级为正式生成」(重新触发全分辨率生成,可复用素材下载缓存)
### 实施步骤
**Phase 1MVP2天):**
1. 数据模型 + 迁移
2. API 层支持 is_preview 参数
3. 渲染分辨率降级(480p
4. 不计入配额
5. 基础测试
**Phase 2(优化,1-2天):**
1. 时长限制(前15秒)
2. 效果层降级(跳高级效果)
3. 预览任务低优先级队列
4. 预览转正式生成功能
## 与前端对齐点
1. 预览生成的触发时机(第3步自动生成?用户点击才生成?)
2. 预览时长是固定15秒还是完整但低清?
3. 是否需要「预览转正式生成」功能
4. 预览视频的展示形态(和正式视频一样还是有特殊UI)
## 风险与注意事项
1. **数据混淆**:确保统计、计费、列表展示时正确区分预览和正式任务
2. **存储成本**:预览视频也占 OSS 空间,可设置自动清理(7天后自动删除)
3. **用户预期**:要明确告诉用户这是预览,效果和正式生成一致但清晰度低
4. **并发压力**:如果用户频繁生成预览,可能增加系统负载,需要限流
---
> 本内容由 Coze AI 生成,请遵循相关法律法规及《人工智能生成合成内容标识办法》使用与传播。
+382
View File
@@ -0,0 +1,382 @@
# #1197 预览生成接口技术方案(v2)
> 更新说明:v2 新增「多版本预览生成」能力,支持一个模板生成多个不重复的预览视频,左侧列表展示,用户可挑选满意的版本转正式生成。
## 1. 背景与目标
**现状**:智能剪辑「一键生成」第3步预览生成被跳过,用户直接进入正式生成,缺少效果预览环节。
**目标**
1. ✅ 实现真正的预览生成(低分辨率快速出片)
2.**支持生成 1~N 个不重复的预览版本**(默认 3 个),左侧列表展示
3. ✅ 预览满意后可一键转正式生成(复用素材下载缓存)
4. ✅ 不计入用户配额,不占用正式生成次数
---
## 2. 现有生成链路分析
### 2.1 链路总览
```
API 触发生成 → GenerationTask入库 → Celery异步任务
→ 下载素材 → 构建plan/clips → UnifiedRenderService渲染
→ 混音后处理 → OSS上传 + 查重 → 更新状态
```
### 2.2 决定视频差异的变量
要做"多个不重复版本",先分析哪些环节可以引入变化:
| 变量 | 当前行为 | 能否引入变化 | 影响程度 |
|------|---------|------------|---------|
| 素材选择 | 按 asset_ids 顺序全用 | ✅ 可随机选择子集/不同组合 | 大 |
| 素材排序 | 按 asset_ids 顺序 | ✅ 可 shuffle 重排 | 大 |
| 配音选择 | 固定 voice_library_id | ✅ 可选不同音色 | 中 |
| 标题选择 | 固定 title_ids 或随机选 | ✅ 可选不同标题 | 中 |
| BGM | 固定 bgm_config | ✅ 可选不同BGM | 小 |
| 转场效果 | 模板固定 | ✅ 可随机化转场类型 | 小 |
| 播放速度 | 模板固定 | ✅ 可微调速度 | 小 |
| 分辨率/码率 | 固定 | ✅ 预览可降级 | 不影响内容 |
### 2.3 可复用能力
- 任务创建与状态管理:`GenerationTask` + `CreateGenerationTaskUseCase`
- 素材下载与预处理:`_download_all_assets`
- 统一渲染引擎:`UnifiedRenderService`
- 分辨率配置:`resolution` 字段已支持
- 批量任务:`batch_id` 字段已存在(可用于预览组)
---
## 3. 总体方案:复用现有链路 + 多变体引擎
**核心思路**:沿用 v1 的"复用现有生成链路 + is_preview 标记"方案,在此基础上增加「多版本生成」能力。
**架构**
```
预览生成请求(count=N
创建预览批次(preview_batch
变体引擎生成 N 个变体参数(variation seed + 参数组合)
为每个变体创建 1 个 GenerationTaskis_preview=true
N 个 Celery 任务并行执行(走现有生成链路,参数降级)
N 个结果汇聚,前端左侧列表展示
```
---
## 4. 详细设计
### 4.1 数据模型变更
#### 4.1.1 GenerationTask 新增字段
```python
# 现有字段保留,新增:
is_preview: bool = False
"""是否为预览生成"""
preview_batch_id: str = ""
"""预览批次 ID(同批次的 N 个预览共享一个 batch)"""
variant_seed: int = 0
"""变体种子,用于控制随机化行为(素材选择、排序、转场等)"""
variant_params: dict = field(default_factory=dict)
"""变体参数快照(记录本次使用了哪些素材、标题、配音等,可追溯)
{
"asset_ids": [...], # 实际选用的素材子集
"title_id": "", # 选用的标题
"voice_id": "", # 选用的配音
"transition_style": "", # 转场风格
"bgm_track": "", # BGM 音轨
}
"""
```
#### 4.1.2 GeneratedVideo 新增字段
```python
is_preview: bool = False
"""是否为预览视频"""
preview_batch_id: str = ""
"""所属预览批次"""
variant_index: int = 0
"""在批次中的序号(0, 1, 2..."""
```
#### 4.1.3 迁移方案
alembic 新增 migration,两个表各加 4 个字段,默认值为空/false,无数据回填成本。
---
### 4.2 变体引擎(Variant Engine
**核心组件**:根据 count 和 seed,生成 N 组互不相同的生成参数。
#### 4.2.1 变纬度设计
| 维度 | 策略 | 说明 |
|------|------|------|
| **素材子集选择** | 从素材池中随机选 M 个(M=min(素材数, 模板clip数*2)) | 版本差异最大的来源 |
| **素材排序** | 随机打乱顺序 | 影响叙事节奏 |
| **标题选择** | 从 title_ids 中随机选 1 个 | 影响文案内容 |
| **配音选择** | 从 voice_ids 中随机选 1 个(如有多个) | 影响听觉体验 |
| **转场风格** | 从预设转场池中随机选 1 种 | 影响视觉过渡 |
| **BGM 选择** | 从 bgm 列表中随机选 1 首(如有配置) | 影响氛围 |
#### 4.2.2 去重机制
- 同一批次内,变体参数必须两两不同(至少素材组合或排序不同)
- 使用 `variant_seed` 保证可复现(相同 seed → 相同变体)
- 如果素材数量不足导致无法生成 N 个不同版本,按实际能生成的数量返回
#### 4.2.3 接口设计
```python
def generate_variants(
count: int,
seed: int,
asset_pool: list[str], # 可用素材 ID 列表
title_pool: list[str] = [], # 可用标题 ID 列表
voice_pool: list[str] = [], # 可用配音 ID 列表
template_id: str = "",
) -> list[dict]:
"""
生成 count 组变体参数。
每组参数包含:asset_ids(选用的素材+排序)、title_id、voice_id、
transition_style 等,确保两两不同。
"""
```
---
### 4.3 API 层设计
#### 4.3.1 预览生成接口
```
POST /api/v1/templates/{template_id}/generate-preview
```
**请求体**
```json
{
"asset_library_id": "lib_xxx",
"asset_ids": ["asset_1", "asset_2", ...],
"title_ids": ["title_1", "title_2"],
"voice_ids": ["voice_1", "voice_2"],
"bgm_config": {},
"count": 3,
"seed": 0
}
```
| 参数 | 类型 | 必填 | 默认 | 说明 |
|------|------|------|------|------|
| template_id | path | ✅ | - | 模板 ID |
| asset_library_id | body | ✅ | - | 素材库 ID |
| asset_ids | body | ✅ | - | 素材池(从中选子集/排序) |
| title_ids | body | - | [] | 标题池(可选,不传则不用标题) |
| voice_ids | body | - | [] | 配音池(可选) |
| bgm_config | body | - | {} | BGM 配置 |
| count | body | - | 3 | 生成几个预览版本(1~10) |
| seed | body | - | 0 | 随机种子,0 表示随机 |
**响应**
```json
{
"preview_batch_id": "pb_xxx",
"count": 3,
"tasks": [
{
"task_id": "gen_xxx_0",
"variant_index": 0,
"status": "processing"
},
{
"task_id": "gen_xxx_1",
"variant_index": 1,
"status": "processing"
},
...
]
}
```
#### 4.3.2 预览批次查询接口
```
GET /api/v1/preview-batches/{batch_id}
```
返回批次内所有预览任务的状态、结果(已完成的带 video_url)。
**响应**
```json
{
"preview_batch_id": "pb_xxx",
"count": 3,
"completed_count": 2,
"tasks": [
{
"task_id": "gen_xxx_0",
"variant_index": 0,
"status": "completed",
"video_url": "https://oss.xxx/preview/xxx.mp4",
"duration": 15.5,
"thumbnail_url": "https://oss.xxx/preview/xxx.jpg"
},
...
]
}
```
#### 4.3.3 预览转正式生成
```
POST /api/v1/preview-batches/{batch_id}/tasks/{task_id}/promote
```
将某个预览版本升级为正式生成(复用素材缓存,重新全分辨率渲染)。
---
### 4.4 渲染参数降级
预览模式下自动调整以下参数:
| 参数 | 正式生成 | 预览生成 |
|------|---------|---------|
| 分辨率 | 720p (1280x720) | 480p (854x480) |
| 码率 | 4 Mbps | 1.5 Mbps |
| 帧率 | 30 fps | 24 fps |
| 时长 | 完整时长 | 前 15 秒(或第一段完整clip) |
| 转场效果 | 完整转场 | 仅淡入淡出 |
| 高级特效 | 全部启用 | 跳过粒子/光效等 |
| 字幕 | 完整渲染 | 正常渲染 |
| 配音 | 完整混音 | 正常混音 |
| 输出质量 | high | medium |
**实现位置**`_render_video` 函数入口处,根据 `is_preview` 标记调整渲染配置。
---
### 4.5 任务调度
- **并行执行**:N 个预览任务并行提交到 Celery,不排队等待
- **低优先级**:预览任务走独立队列(`preview_queue`),不抢占正式生成资源
- **超时控制**:预览任务超时时间 5 分钟(正式 30 分钟)
- **自动清理**:预览视频 7 天后自动从 OSS 删除,任务记录标记为 archived
---
## 5. 前端对接要点
### 5.1 交互流程
```
第2步选素材 → 第3步点击"生成预览"
→ 显示 loading + 进度
→ 预览陆续完成,左侧列表逐张出现
→ 用户点击左侧不同版本,右侧预览区切换
→ 用户选中满意版本 → 点击"正式生成"
```
### 5.2 需要对齐的接口
1. **预览创建**`POST /templates/{id}/generate-preview`
2. **批次状态轮询**`GET /preview-batches/{id}`(建议 2s 轮询,或走 SSE
3. **预览转正式**`POST /preview-batches/{id}/tasks/{task_id}/promote`
### 5.3 数据格式对齐
预览视频条目结构:
```json
{
"id": "gen_xxx",
"variant_index": 0,
"status": "completed",
"video_url": "https://...",
"duration": 15.5,
"file_size": 2850000,
"thumbnail_url": "https://...",
"is_preview": true
}
```
---
## 6. 配额与计费
- 预览生成**不计入**用户配额
- 同一模板 + 同一素材池,每天最多生成 3 次多版本预览(防滥用)
- 单个预览批次最多 10 个版本
---
## 7. 实施步骤
### Phase 1:单版本预览(MVP2 天)
1. 数据模型 + 迁移(is_preview 字段)
2. API 层支持 is_preview 参数
3. 渲染分辨率降级(480p
4. 不计入配额
5. 基础测试
### Phase 2:多版本预览(3 天)
1. 变体引擎实现(素材随机选择 + 排序 + 去重)
2. preview_batch 批次管理
3. 批量创建 N 个预览任务
4. 批次查询接口
5. 前端联调
### Phase 3:预览转正式 + 优化(2 天)
1. 预览转正式生成接口(promote)
2. 素材下载缓存复用
3. 独立预览队列(低优先级)
4. 自动清理机制
5. 完整测试 + 压测
---
## 8. 风险与注意事项
| 风险 | 影响 | 应对 |
|------|------|------|
| 并发预览任务过多打满 worker | 正式生成被阻塞 | 独立预览队列 + 限流 |
| 变体生成的视频差异不够大 | 用户觉得"都一样" | 优先素材子集+排序差异,保证视觉差异 |
| 预览视频占用 OSS 存储 | 存储成本上升 | 7 天自动清理 + 低码率 |
| N 个版本同时下载重复素材 | 带宽浪费 | 批次内共享一次下载(Phase 3 优化) |
| 用户预期管理 | 以为预览就是最终效果 | 明确标注"预览版",说明分辨率差异 |
---
## 9. 开发量估算
| 阶段 | 后端 | 前端 | 合计 |
|------|------|------|------|
| Phase 1 单版本预览 | 2 天 | 1 天 | 3 天 |
| Phase 2 多版本预览 | 3 天 | 2 天 | 5 天 |
| Phase 3 转正式+优化 | 2 天 | 1 天 | 3 天 |
| **总计** | **7 天** | **4 天** | **~7 天(并行)** |
---
## 10. 与 v1 方案的差异总结
1. **新增多版本能力**:从"生成1个预览"升级为"生成N个不重复预览"
2. **新增变体引擎**:负责素材选择/排序/配音/标题的随机化
3. **新增批次概念**preview_batch 管理一组预览任务
4. **新增 promote 接口**:预览转正式生成
5. **独立队列**:预览不抢占正式生成资源
6. **开发量**:从 2-3 天增加到约 7 天(后端)
+49
View File
@@ -19,6 +19,9 @@ from app.schemas.asset import (
BatchTagRequest,
CreateAssetRequest,
ListAssetsResponse,
SmartMatchItem,
SmartMatchRequest,
SmartMatchResponse,
UpdateAssetRequest,
UpdateAssetReviewRequest,
)
@@ -30,6 +33,7 @@ from packages.application import (
CreateAssetUseCase,
)
from packages.domain import AssetStatus, ClassificationStatus
from packages.domain.smart_match import smart_select_assets
logger = logging.getLogger(__name__)
@@ -519,6 +523,51 @@ def batch_mark_assets(
)
@router.post("/smart-match", response_model=SmartMatchResponse)
def smart_match_assets(
request: SmartMatchRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
asset_library_repository: Any = Depends(get_asset_library_repository),
project_repository: Any = Depends(get_project_repository),
) -> SmartMatchResponse:
"""智能选素材:根据素材库内容,按质量分+时长均衡+新鲜度+未使用偏好综合评分,返回 Top N 素材。"""
library = asset_library_repository.get(request.library_id)
if library is None:
raise HTTPException(status_code=404, detail=f"AssetLibrary {request.library_id} not found")
check_project_access(library.project_id, authenticated_user.user.id, project_repository)
# 获取素材库中所有 ready 素材(DB 层按 kind 过滤,避免加载不必要的数据到内存)
# kind → file_type 映射:schema 已校验只允许 video/image/audio,与 file_type 一致
if request.kind:
filtered_assets = asset_repository.find_by_library_and_file_type(
request.library_id, request.kind, status=["ready"], limit=10000
)
else:
filtered_assets = asset_repository.find_by_library(
request.library_id, status=["ready"], limit=10000
)
total_candidates = len(filtered_assets)
# 调用统一智能选素材算法(kind 已在 DB 层过滤,无需重复过滤)
results = smart_select_assets(
filtered_assets,
limit=request.limit,
kind=None,
)
items = [
SmartMatchItem(
asset=_to_asset_response(r.asset),
score=r.score,
breakdown=r.breakdown,
)
for r in results
]
return SmartMatchResponse(items=items, total_candidates=total_candidates)
@router.get("/{asset_id}", response_model=AssetResponse)
def get_asset(
asset_id: str,
+27
View File
@@ -101,3 +101,30 @@ class ListAssetsResponse(BaseModel):
total: int = Field(default=0, ge=0)
skip: int = Field(default=0, ge=0)
limit: int = Field(default=100, ge=1)
class SmartMatchRequest(BaseModel):
"""智能选素材请求。"""
library_id: str = Field(..., min_length=1, description="素材库 ID")
limit: int | None = Field(default=None, ge=1, le=200, description="最大返回数量,不传则返回全部匹配素材")
kind: str | None = Field(
default=None,
pattern="^(video|image|audio)$",
description="按文件类型过滤,不传则返回所有类型",
)
class SmartMatchItem(BaseModel):
"""智能选素材结果条目。"""
asset: AssetResponse
score: float = Field(..., ge=0, le=100, description="综合得分 0-100")
breakdown: dict[str, float] = Field(default_factory=dict, description="各维度得分明细")
class SmartMatchResponse(BaseModel):
"""智能选素材响应。"""
items: list[SmartMatchItem]
total_candidates: int = Field(default=0, ge=0, description="参与评分的候选素材总数")
-6
View File
@@ -1441,12 +1441,6 @@
/* textarea removed in Q5 */
.xx-smart-match-tip {
font-size: 12px;
color: var(--text-tertiary, #94a3b8);
+208
View File
@@ -0,0 +1,208 @@
# #1197 预览生成接口方案评估
## 背景
智能剪辑「一键生成」流程中,第3步预览生成当前被跳过,直接进入下一步。需要实现真正的预览生成功能,让用户在正式生成前能看到效果预览。
## 现状分析
### 现有生成链路
```
API 触发生成 → GenerationTask入库 → Celery异步任务 → UnifiedRenderService渲染 → OSS上传 → 更新状态
```
**关键节点:**
1. **API层**`POST /generation-tasks``POST /templates/{id}/generate` 触发生成
2. **任务调度**Celery task `worker.generate_video`
3. **渲染引擎**`UnifiedRenderService`(统一渲染引擎,已接入9个效果层)
4. **输出配置**:默认 720p (1280x720),支持 `resolution` 字段自定义
5. **产物存储**`GeneratedVideo` 表记录,OSS 存储视频文件
### 已有可复用能力
| 能力 | 位置 | 是否可复用 |
|------|------|-----------|
| 任务创建与状态管理 | `GenerationTask` + `CreateGenerationTaskUseCase` | ✅ 是 |
| 素材下载与预处理 | `_download_video_assets` / `_download_voice_asset` | ✅ 是 |
| 统一渲染引擎 | `UnifiedRenderService` | ✅ 是 |
| 分辨率配置 | `resolution` 字段已支持 | ✅ 是 |
| 混音与后处理 | `_render_video` 内流程 | ✅ 是 |
| OSS 上传与查重 | `_upload_and_dedup` | ✅ 是 |
| 进度追踪 | `append_log` / `progress` 字段 | ✅ 是 |
## 方案对比
### 方案A:复用现有生成链路 + is_preview 标记(推荐)
**思路**:在现有 GenerationTask 上加 `is_preview` 标记,预览生成走完整链路但参数降级。
**改动点:**
1. **数据模型**`GenerationTask``is_preview: bool` 字段(默认 false);`GeneratedVideo``is_preview: bool`
2. **API 层**:生成接口加 `is_preview` 参数,预览任务不计入配额
3. **渲染参数**:预览模式下自动调整
- 分辨率:480p (854x480)
- 时长:限制前 15 秒(或模板第一个片段)
- 码率:降低至 1.5Mbps(正式 4Mbps
- 效果层:跳过高级转场/粒子特效等耗时效果
4. **任务调度**:预览任务走低优先级队列(或复用现有队列,标记优先级)
5. **前端对接**:预览生成结果带 `is_preview=true` 标记,前端展示"预览"标签
**优点:**
- 代码复用率 90%+,改动最小
- 与正式生成逻辑一致,预览效果真实可信
- 进度查询、结果展示等功能直接复用
- 后续可平滑升级:预览满意后一键转正式生成
**缺点:**
- 需要区分预览和正式任务,避免数据混淆
- 预览任务和正式任务竞争同一队列资源(可后续优化为独立队列)
**开发量估算**2-3 天
- 数据模型 + 迁移:0.5 天
- API 层改造:0.5 天
- 渲染参数降级:1 天
- 测试 + 联调:1 天
---
### 方案B:新建独立预览接口 + 轻量渲染逻辑
**思路**:新建独立的预览生成接口,使用简化的渲染逻辑(如只拼接素材+基础配音,跳过大部分效果)。
**改动点:**
1. 新增 `PreviewTask` 数据模型
2. 新增 `POST /api/v1/preview/generate` 接口
3. 新增独立的 Celery task `worker.generate_preview`
4. 简化渲染流程:只做素材裁剪+拼接+配音,跳过转场/滤镜/字幕特效等
**优点:**
- 完全隔离,不影响正式生成链路
- 可以做极致优化,预览生成速度快
- 数据模型清晰,不会混淆
**缺点:**
- 代码重复率高,两套生成逻辑维护成本翻倍
- 预览效果与正式生成可能不一致(效果层差异)
- 前端需要对接两套接口
- 无法从预览升级为正式生成(需重新走完整流程)
**开发量估算**4-5 天
- 数据模型 + 接口:1 天
- 简化渲染逻辑:2 天
- 测试 + 联调:1-2 天
---
### 方案C:图片预览(首帧/关键帧截图)
**思路**:不生成视频,只生成几张关键帧的预览图片。
**优点:**
- 生成速度极快(秒级)
- 资源消耗小
**缺点:**
- 预览效果差,用户无法感知动态效果
- 无法验证配音、转场、节奏等时间维度的效果
- 用户体验不佳,不如"真预览"有说服力
**开发量估算**1-2 天
---
## 推荐方案:方案A(复用现有生成链路)
### 核心理由
1. **效果保真**:预览和正式生成用同一套渲染引擎,效果一致,用户信任度高
2. **开发效率**90% 代码复用,2-3 天可上线
3. **可扩展性强**:后续可加「预览转正式」「低分辨率快速预览」等增强功能
4. **维护成本低**:一套生成逻辑,bug 修复和新功能同时生效
### 详细设计
#### 1. 数据模型变更
```python
# GenerationTask 新增字段
is_preview: bool = False
"""是否为预览生成"""
preview_of: str = ""
"""预览对应的正式任务 ID(或反向关联)"""
# GeneratedVideo 新增字段
is_preview: bool = False
"""是否为预览视频"""
```
**迁移**alembic 新增 migration,两个表各加 1-2 个字段。
#### 2. API 层
```
POST /api/v1/generation-tasks
Body 增加 is_preview: bool = false
POST /api/v1/templates/{id}/generate
Query 增加 is_preview: bool = false
```
**配额处理**:预览生成不计入用户配额,不占用生成次数限制。
#### 3. 渲染参数降级
| 参数 | 正式生成 | 预览生成 |
|------|---------|---------|
| 分辨率 | 720p (1280x720) | 480p (854x480) |
| 码率 | 4 Mbps | 1.5 Mbps |
| 时长 | 完整时长 | 前 15 秒(或第一段) |
| 帧率 | 30 fps | 24 fps |
| 转场效果 | 完整转场 | 仅淡入淡出(或简单切) |
| 特效滤镜 | 全部启用 | 跳过粒子/光效等高级效果 |
| 字幕 | 完整渲染 | 正常渲染(字幕是核心信息) |
| 配音 | 完整混音 | 正常混音(配音是核心信息) |
**实现方式**:在 `_render_video` 或 UnifiedRenderService 入口处,根据 `is_preview` 标记调整渲染配置。
#### 4. 任务调度
- 初期复用现有队列,预览任务正常排队
- 后续如需优化,可拆分独立预览队列(低优先级)
- 预览任务可设置较短超时时间
#### 5. 前端对接
- 调用生成接口时传 `is_preview=true`
- 结果列表中预览视频带「预览」标签
- 预览满意后可一键「升级为正式生成」(重新触发全分辨率生成,可复用素材下载缓存)
### 实施步骤
**Phase 1MVP2天):**
1. 数据模型 + 迁移
2. API 层支持 is_preview 参数
3. 渲染分辨率降级(480p
4. 不计入配额
5. 基础测试
**Phase 2(优化,1-2天):**
1. 时长限制(前15秒)
2. 效果层降级(跳高级效果)
3. 预览任务低优先级队列
4. 预览转正式生成功能
## 与前端对齐点
1. 预览生成的触发时机(第3步自动生成?用户点击才生成?)
2. 预览时长是固定15秒还是完整但低清?
3. 是否需要「预览转正式生成」功能
4. 预览视频的展示形态(和正式视频一样还是有特殊UI)
## 风险与注意事项
1. **数据混淆**:确保统计、计费、列表展示时正确区分预览和正式任务
2. **存储成本**:预览视频也占 OSS 空间,可设置自动清理(7天后自动删除)
3. **用户预期**:要明确告诉用户这是预览,效果和正式生成一致但清晰度低
4. **并发压力**:如果用户频繁生成预览,可能增加系统负载,需要限流
+4
View File
@@ -24,6 +24,7 @@ from .entities import (
from .generated_video import GeneratedVideo
from .generation_task import GenerationTask, GenerationTaskStatus
from .job import Job, JobStatus, JobType
from .smart_match import SmartMatchResult, score_asset, smart_select_assets
from .tag import Tag
from .template_clip_config import ClipType, TemplateClipConfig, TransitionEffect
from .title_library import TitleLibraryItem
@@ -61,6 +62,9 @@ __all__ = [
"TemplateClipConfig",
"TransitionEffect",
"User",
"SmartMatchResult",
"TitleLibraryItem",
"VoiceLibraryItem",
"score_asset",
"smart_select_assets",
]
+207
View File
@@ -0,0 +1,207 @@
"""统一智能选素材算法 — 合并 _helpers / generation_tasks / auto_clip_service 的重叠逻辑。
设计目标:
- 单一入口,替代 3 套分散的选素材代码
- 多维度加权评分:质量分 + 时长均衡 + 新鲜度 + 未使用偏好
- 多样性保障:按时长分桶(短/中/长)均衡选取,避免同质化
- 可扩展:后续接入 AI 模型时只需替换 score_asset()
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any
@dataclass
class SmartMatchResult:
"""单条素材的匹配结果。"""
asset: Any # Asset entity
score: float # 综合得分 0-100
breakdown: dict[str, float] = field(default_factory=dict) # 各维度得分明细
def _get_enum_value(obj: Any, attr: str) -> str:
"""安全获取属性值,兼容 StrEnum / 普通字符串。"""
val = getattr(obj, attr, None)
if val is None:
return ""
return val.value if hasattr(val, "value") else str(val)
def _duration_bucket(duration: float | None) -> str:
"""将素材时长分为 3 档:short(<10s) / medium(10-30s) / long(>30s)。"""
if duration is None or duration <= 0:
return "unknown"
if duration < 10:
return "short"
if duration <= 30:
return "medium"
return "long"
def score_asset(
asset: Any,
now: datetime | None = None,
) -> tuple[float, dict[str, float]]:
"""为单个素材计算综合得分(0-100)。
维度权重:
- quality_score (40%):素材质量分(0-100),无质量分按 50 计
- duration_fitness (30%):时长适配度,5-30s 为最优区间
- recency (20%):新鲜度,30 天内衰减
- unused_bonus (10%):未被使用过的素材加分
Returns:
(total_score, breakdown_dict)
"""
if now is None:
now = datetime.now(timezone.utc)
breakdown: dict[str, float] = {}
# 1. 质量分 (0-100) → 权重 40%
raw_quality = asset.quality_score if asset.quality_score is not None else 50.0
quality_component = raw_quality * 0.4
breakdown["quality"] = round(quality_component, 2)
# 2. 时长适配度 (0-100) → 权重 30%
# 最优区间 5-30s 得满分,越偏离越低
duration = getattr(asset, "duration", None) or 0.0
if duration <= 0:
duration_fitness = 30.0 # 未知时长给中等分
elif 5 <= duration <= 30:
duration_fitness = 100.0
elif duration < 5:
# 0-5s: 线性增长 20→100
duration_fitness = 20.0 + (duration / 5) * 80
else:
# >30s: 指数衰减,60s 时约 50 分
duration_fitness = 100.0 * math.exp(-0.02 * (duration - 30))
duration_fitness = max(duration_fitness, 10.0)
duration_component = duration_fitness * 0.3
breakdown["duration"] = round(duration_component, 2)
# 3. 新鲜度 (0-100) → 权重 20%
# 30 天半衰期
created_at = getattr(asset, "created_at", None)
if created_at is None:
recency = 50.0
else:
if created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=timezone.utc)
age_days = max(0, (now - created_at).total_seconds() / 86400)
recency = 100.0 * math.exp(-0.05 * age_days) # ~14天半衰期
recency_component = recency * 0.2
breakdown["recency"] = round(recency_component, 2)
# 4. 未使用偏好 (0-100) → 权重 10%
metadata = getattr(asset, "metadata", None) or {}
try:
use_count = int(metadata.get("generation_use_count") or 0)
except (ValueError, TypeError):
use_count = 0 # 脏数据时按未使用处理(保守策略:给未使用加分)
if use_count == 0:
unused_score = 100.0
elif use_count <= 3:
unused_score = 70.0
else:
unused_score = 30.0
unused_component = unused_score * 0.1
breakdown["unused"] = round(unused_component, 2)
total = quality_component + duration_component + recency_component + unused_component
return round(total, 2), breakdown
def smart_select_assets(
assets: list[Any],
*,
limit: int | None = None,
kind: str | None = None,
now: datetime | None = None,
) -> list[SmartMatchResult]:
"""从素材列表中智能选取素材。
Args:
assets: 候选素材列表(Asset 实体)
limit: 最大返回数量,None 表示不限制
kind: 按文件类型过滤(video/image/audio),None 表示不过滤
now: 当前时间(用于测试注入)
Returns:
按得分降序排列的 SmartMatchResult 列表
"""
# Step 1: 过滤 ready 状态
ready_assets = [a for a in assets if _get_enum_value(a, "status") == "ready"]
# Step 2: 按 kind 过滤
if kind:
ready_assets = [a for a in ready_assets if a.file_type == kind]
if not ready_assets:
return []
# Step 3: 评分
scored: list[SmartMatchResult] = []
for a in ready_assets:
total, breakdown = score_asset(a, now=now)
scored.append(SmartMatchResult(asset=a, score=total, breakdown=breakdown))
# Step 4: 按得分降序排序
scored.sort(key=lambda r: r.score, reverse=True)
# Step 5: 多样性保障 — 时长分桶均衡选取
if limit and limit > 0 and len(scored) > limit:
scored = _diversity_select(scored, limit)
elif limit and limit > 0:
scored = scored[:limit]
return scored
def _diversity_select(scored: list[SmartMatchResult], limit: int) -> list[SmartMatchResult]:
"""从已排序的候选中按分桶均衡选取,避免全选中同一时长档。
策略:轮流从 short/medium/long 桶中按得分顺序取,直到凑满 limit。
"""
buckets: dict[str, list[SmartMatchResult]] = {
"short": [],
"medium": [],
"long": [],
"unknown": [],
}
for r in scored:
bucket = _duration_bucket(getattr(r.asset, "duration", None))
buckets.setdefault(bucket, []).append(r)
selected: list[SmartMatchResult] = []
selected_ids: set[str] = set()
bucket_order = ["medium", "short", "long", "unknown"] # medium 优先
bucket_idx = {b: 0 for b in bucket_order}
while len(selected) < limit:
added = False
for b in bucket_order:
if len(selected) >= limit:
break
items = buckets.get(b, [])
idx = bucket_idx[b]
while idx < len(items):
candidate = items[idx]
idx += 1
if candidate.asset.id not in selected_ids:
selected.append(candidate)
selected_ids.add(candidate.asset.id)
added = True
break
bucket_idx[b] = idx
if not added:
break
# 按原始得分降序输出
selected.sort(key=lambda r: r.score, reverse=True)
return selected
+524
View File
@@ -0,0 +1,524 @@
"""Tests for packages/domain/smart_match.py — 统一智能选素材算法。"""
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import Any
import pytest
from packages.domain.smart_match import (
SmartMatchResult,
_diversity_select,
_duration_bucket,
score_asset,
smart_select_assets,
)
# ── Fixtures ─────────────────────────────────────────────────────────────────
@dataclass
class FakeAsset:
"""Minimal Asset-like object for testing."""
id: str
project_id: str = "proj-1"
library_id: str = "lib-1"
name: str = "test"
storage_key: str = "key"
mime_type: str = "video/mp4"
file_size: int = 1000
duration: float | None = None
width: int | None = 1080
height: int | None = 1920
quality_score: float | None = None
status: str = "ready"
metadata: dict = field(default_factory=dict)
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@property
def file_type(self) -> str:
if "/" in self.mime_type:
return self.mime_type.split("/")[0]
return self.mime_type
NOW = datetime(2026, 8, 5, 12, 0, 0, tzinfo=timezone.utc)
# ── score_asset tests ────────────────────────────────────────────────────────
class TestScoreAsset:
def test_high_quality_asset_scores_high(self):
asset = FakeAsset(id="a1", quality_score=95, duration=15)
score, breakdown = score_asset(asset, now=NOW)
assert score > 70
assert "quality" in breakdown
assert "duration" in breakdown
assert "recency" in breakdown
assert "unused" in breakdown
def test_low_quality_asset_scores_lower(self):
asset_good = FakeAsset(id="a1", quality_score=95, duration=15)
asset_bad = FakeAsset(
id="a2",
quality_score=20,
duration=15,
created_at=NOW - timedelta(days=60),
metadata={"generation_use_count": 10},
)
score_good, _ = score_asset(asset_good, now=NOW)
score_bad, _ = score_asset(asset_bad, now=NOW)
assert score_bad < score_good
def test_no_quality_score_defaults_to_50(self):
asset = FakeAsset(id="a1", quality_score=None, duration=15)
score, breakdown = score_asset(asset, now=NOW)
# quality component should be 50 * 0.4 = 20
assert breakdown["quality"] == pytest.approx(20.0, abs=0.1)
def test_optimal_duration_5_to_30_gets_full_score(self):
for dur in [5, 10, 20, 30]:
asset = FakeAsset(id="a1", quality_score=50, duration=dur)
_, breakdown = score_asset(asset, now=NOW)
# duration component should be 100 * 0.3 = 30
assert breakdown["duration"] == pytest.approx(30.0, abs=0.1)
def test_short_duration_below_5s_penalized(self):
asset = FakeAsset(id="a1", quality_score=50, duration=2)
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["duration"] < 30.0
def test_long_duration_above_30s_penalized(self):
asset = FakeAsset(id="a1", quality_score=50, duration=120)
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["duration"] < 30.0
def test_zero_duration_gives_moderate_score(self):
asset = FakeAsset(id="a1", quality_score=50, duration=0)
_, breakdown = score_asset(asset, now=NOW)
# duration_fitness = 30.0, component = 30 * 0.3 = 9
assert breakdown["duration"] == pytest.approx(9.0, abs=0.1)
def test_unused_asset_gets_full_bonus(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={})
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["unused"] == pytest.approx(10.0, abs=0.1)
def test_used_asset_gets_reduced_bonus(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={"generation_use_count": 5})
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["unused"] == pytest.approx(3.0, abs=0.1)
def test_dirty_metadata_use_count_string_does_not_crash(self):
"""int() conversion of non-numeric metadata should not raise, should default to 0."""
asset = FakeAsset(id="a1", quality_score=50, duration=15, metadata={"generation_use_count": "high"})
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["unused"] == pytest.approx(10.0, abs=0.1) # use_count=0 → unused_score=100 → 100*0.1=10
def test_recent_asset_scores_higher_recency(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=1))
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["recency"] > 15 # > 75% of max 20
def test_old_asset_scores_lower_recency(self):
asset = FakeAsset(id="a1", quality_score=50, duration=15, created_at=NOW - timedelta(days=60))
_, breakdown = score_asset(asset, now=NOW)
assert breakdown["recency"] < 5 # heavily decayed
# ── _duration_bucket tests ───────────────────────────────────────────────────
class TestDurationBucket:
def test_short(self):
assert _duration_bucket(5) == "short"
assert _duration_bucket(9.9) == "short"
def test_medium(self):
assert _duration_bucket(10) == "medium"
assert _duration_bucket(30) == "medium"
def test_long(self):
assert _duration_bucket(31) == "long"
assert _duration_bucket(120) == "long"
def test_unknown(self):
assert _duration_bucket(None) == "unknown"
assert _duration_bucket(0) == "unknown"
assert _duration_bucket(-1) == "unknown"
# ── smart_select_assets tests ────────────────────────────────────────────────
class TestSmartSelectAssets:
def test_filters_non_ready_assets(self):
assets = [
FakeAsset(id="a1", status="ready", quality_score=80, duration=15),
FakeAsset(id="a2", status="uploading", quality_score=90, duration=15),
FakeAsset(id="a3", status="error", quality_score=70, duration=15),
]
results = smart_select_assets(assets)
assert len(results) == 1
assert results[0].asset.id == "a1"
def test_filters_by_kind(self):
assets = [
FakeAsset(id="a1", mime_type="video/mp4", quality_score=80, duration=15),
FakeAsset(id="a2", mime_type="image/png", quality_score=90, duration=0),
FakeAsset(id="a3", mime_type="audio/mp3", quality_score=70, duration=30),
]
results = smart_select_assets(assets, kind="video")
assert len(results) == 1
assert results[0].asset.id == "a1"
def test_respects_limit(self):
assets = [FakeAsset(id=f"a{i}", quality_score=50 + i, duration=15) for i in range(20)]
results = smart_select_assets(assets, limit=5)
assert len(results) == 5
def test_returns_sorted_by_score_descending(self):
assets = [
FakeAsset(id="low", quality_score=20, duration=15),
FakeAsset(id="high", quality_score=95, duration=15),
FakeAsset(id="mid", quality_score=60, duration=15),
]
results = smart_select_assets(assets)
scores = [r.score for r in results]
assert scores == sorted(scores, reverse=True)
assert results[0].asset.id == "high"
def test_empty_list_returns_empty(self):
assert smart_select_assets([]) == []
def test_all_non_ready_returns_empty(self):
assets = [FakeAsset(id="a1", status="uploading")]
assert smart_select_assets(assets) == []
def test_diversity_select_balances_duration_buckets(self):
"""When limit is less than total, diversity select should pick from multiple buckets."""
assets = []
# 10 short clips
for i in range(10):
assets.append(FakeAsset(id=f"s{i}", quality_score=80, duration=5))
# 10 medium clips
for i in range(10):
assets.append(FakeAsset(id=f"m{i}", quality_score=80, duration=20))
# 10 long clips
for i in range(10):
assets.append(FakeAsset(id=f"l{i}", quality_score=80, duration=60))
results = smart_select_assets(assets, limit=6)
assert len(results) == 6
# Should have items from multiple buckets
buckets = {_duration_bucket(r.asset.duration) for r in results}
assert len(buckets) >= 2 # at least 2 different duration buckets
def test_no_limit_returns_all(self):
assets = [FakeAsset(id=f"a{i}", quality_score=50 + i, duration=15) for i in range(10)]
results = smart_select_assets(assets, limit=None)
assert len(results) == 10
def test_score_includes_breakdown(self):
asset = FakeAsset(id="a1", quality_score=80, duration=15, metadata={})
results = smart_select_assets([asset])
assert len(results) == 1
r = results[0]
assert r.score > 0
assert set(r.breakdown.keys()) == {"quality", "duration", "recency", "unused"}
def test_image_assets_can_be_selected(self):
assets = [
FakeAsset(id="img1", mime_type="image/jpeg", quality_score=90, duration=None),
FakeAsset(id="img2", mime_type="image/png", quality_score=70, duration=None),
]
results = smart_select_assets(assets, kind="image")
assert len(results) == 2
assert results[0].asset.id == "img1"
def test_str_enum_status_handled(self):
"""Test that StrEnum-like status objects are handled correctly."""
class StrEnumLike:
def __init__(self, value):
self.value = value
asset = FakeAsset(id="a1", quality_score=80, duration=15)
asset.status = StrEnumLike("ready")
results = smart_select_assets([asset])
assert len(results) == 1
# ── _diversity_select tests ──────────────────────────────────────────────────
class TestDiversitySelect:
def test_picks_from_all_buckets(self):
results = [
SmartMatchResult(asset=FakeAsset(id="s1", duration=5), score=90),
SmartMatchResult(asset=FakeAsset(id="s2", duration=3), score=85),
SmartMatchResult(asset=FakeAsset(id="m1", duration=20), score=80),
SmartMatchResult(asset=FakeAsset(id="l1", duration=60), score=75),
]
selected = _diversity_select(results, limit=3)
assert len(selected) == 3
ids = {r.asset.id for r in selected}
# Should have at least one from short, medium, long
assert "s1" in ids or "s2" in ids
assert "m1" in ids
assert "l1" in ids
def test_limit_larger_than_input_returns_all(self):
results = [
SmartMatchResult(asset=FakeAsset(id="a1", duration=5), score=90),
]
selected = _diversity_select(results, limit=10)
assert len(selected) == 1
# ── API endpoint tests ───────────────────────────────────────────────────────
import os
import sys
from pathlib import Path
from unittest.mock import MagicMock
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
from app.api.routes.assets import router
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_project_repository,
)
from packages.domain import (
Asset,
AssetLibrary,
AssetLibraryKind,
AssetStatus,
ClassificationStatus,
Project,
User,
)
class _StubProjectRepo:
def __init__(self, projects):
self._projects = projects
def find_by_id(self, pid):
return self._projects.get(pid)
class _StubAssetLibraryRepo:
def __init__(self, libraries):
self._libraries = libraries
def get(self, lid):
return self._libraries.get(lid)
class _StubAssetRepo:
def __init__(self, assets):
self._assets = assets
def find_by_library(self, lid, skip=0, limit=100, status=None):
result = [a for a in self._assets if a.library_id == lid]
if status:
result = [a for a in result if (a.status.value if hasattr(a.status, "value") else a.status) in status]
return result[skip : skip + limit]
def find_by_library_and_file_type(self, lid, file_type, skip=0, limit=100, status=None):
result = [a for a in self._assets if a.library_id == lid and a.file_type == file_type]
if status:
result = [a for a in result if (a.status.value if hasattr(a.status, "value") else a.status) in status]
return result[skip : skip + limit]
def _make_app(asset_repo, lib_repo, proj_repo):
app = FastAPI()
app.include_router(router, prefix="/assets")
fake_user = MagicMock()
fake_user.user = User(id="user-1", email="test@test.com", display_name="Test")
app.dependency_overrides[get_current_user] = lambda: AuthenticatedUser(user=fake_user.user)
app.dependency_overrides[get_asset_repository] = lambda: asset_repo
app.dependency_overrides[get_asset_library_repository] = lambda: lib_repo
app.dependency_overrides[get_project_repository] = lambda: proj_repo
app.dependency_overrides[get_storage_service] = lambda: MagicMock()
return app
def _make_test_data():
project = Project(id="proj-1", name="Test", owner_user_id="user-1")
library = AssetLibrary(
id="lib-1",
project_id="proj-1",
name="Videos",
kind=AssetLibraryKind.VIDEO,
)
assets = [
Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v1.mp4",
storage_key="k1",
mime_type="video/mp4",
quality_score=90,
duration=15,
status=AssetStatus.READY,
),
Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v2.mp4",
storage_key="k2",
mime_type="video/mp4",
quality_score=50,
duration=25,
status=AssetStatus.READY,
),
Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v3.mp4",
storage_key="k3",
mime_type="video/mp4",
quality_score=30,
duration=60,
status=AssetStatus.READY,
),
]
return project, library, assets
class TestSmartMatchEndpoint:
def test_returns_scored_items(self):
project, library, assets = _make_test_data()
app = _make_app(
_StubAssetRepo(assets),
_StubAssetLibraryRepo({"lib-1": library}),
_StubProjectRepo({"proj-1": project}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "lib-1"})
assert resp.status_code == 200, f"Got {resp.status_code}: {resp.text}"
data = resp.json()
assert len(data["items"]) == 3
assert data["total_candidates"] == 3
# Sorted by score descending
scores = [item["score"] for item in data["items"]]
assert scores == sorted(scores, reverse=True)
# Each item has breakdown
for item in data["items"]:
assert "quality" in item["breakdown"]
assert "duration" in item["breakdown"]
def test_limit_parameter(self):
project, library, assets = _make_test_data()
app = _make_app(
_StubAssetRepo(assets),
_StubAssetLibraryRepo({"lib-1": library}),
_StubProjectRepo({"proj-1": project}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "lib-1", "limit": 2})
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 2
assert data["total_candidates"] == 3
def test_kind_filter(self):
project, library, assets = _make_test_data()
# Add an image asset
img_asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="img.png",
storage_key="k4",
mime_type="image/png",
quality_score=95,
status=AssetStatus.READY,
)
assets.append(img_asset)
app = _make_app(
_StubAssetRepo(assets),
_StubAssetLibraryRepo({"lib-1": library}),
_StubProjectRepo({"proj-1": project}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "lib-1", "kind": "image"})
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 1
assert data["items"][0]["asset"]["mime_type"] == "image/png"
# total_candidates should only count filtered-by-kind assets (1 image, not 3 videos)
assert data["total_candidates"] == 1
def test_kind_filter_video_total_candidates(self):
"""Verify total_candidates reflects kind filtering, not total assets."""
project, library, assets = _make_test_data()
img_asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="img.png",
storage_key="k4",
mime_type="image/png",
quality_score=95,
status=AssetStatus.READY,
)
assets.append(img_asset)
app = _make_app(
_StubAssetRepo(assets),
_StubAssetLibraryRepo({"lib-1": library}),
_StubProjectRepo({"proj-1": project}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "lib-1", "kind": "video"})
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 3
# total_candidates = 3 videos only, not 4 (3 videos + 1 image)
assert data["total_candidates"] == 3
def test_library_not_found_returns_404(self):
app = _make_app(
_StubAssetRepo([]),
_StubAssetLibraryRepo({}),
_StubProjectRepo({}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "nonexistent"})
assert resp.status_code == 404
def test_empty_library_returns_empty_items(self):
project = Project(id="proj-1", name="Test", owner_user_id="user-1")
library = AssetLibrary(
id="lib-1",
project_id="proj-1",
name="Empty",
kind=AssetLibraryKind.VIDEO,
)
app = _make_app(
_StubAssetRepo([]),
_StubAssetLibraryRepo({"lib-1": library}),
_StubProjectRepo({"proj-1": project}),
)
client = TestClient(app)
resp = client.post("/assets/smart-match", json={"library_id": "lib-1"})
assert resp.status_code == 200
data = resp.json()
assert data["items"] == []
assert data["total_candidates"] == 0