diff --git a/apps/api/app/api/routes/templates_editor/clips.py b/apps/api/app/api/routes/templates_editor/clips.py index 44dcdd16e..8189f29eb 100755 --- a/apps/api/app/api/routes/templates_editor/clips.py +++ b/apps/api/app/api/routes/templates_editor/clips.py @@ -44,6 +44,7 @@ from packages.adapters.sqlalchemy_impl.template_repository import ( SQLAlchemyTemplateRepository, ) from packages.domain.plan_generator_utils import _calc_random_start_time +from packages.domain.smart_match import score_asset from packages.shared.mediakit_client import get_mediakit_client from .dependencies import get_draft_plan_id, get_editor_services @@ -667,10 +668,14 @@ def create_clips_from_assets_editor( # 2. 获取素材实际时长(去重查询) unique_asset_ids = list(dict.fromkeys(asset_ids)) asset_durations: dict[str, float] = {} + asset_smart_scores: dict[str, float] = {} for asset_id in unique_asset_ids: asset = asset_repo.get(asset_id) if asset and hasattr(asset, "duration"): asset_durations[asset_id] = float(asset.duration or 0.0) + # 计算 smart_match 综合评分,用于候选排序 + smart_score, _ = score_asset(asset) + asset_smart_scores[asset_id] = smart_score # 3. 在内存中计算所有片段数据(使用随机起始时间,不调用MediaKit) # 读取素材 metadata 中持久化的历史已用区间(跨任务/跨调用去重), @@ -725,7 +730,11 @@ def create_clips_from_assets_editor( } sorted_candidates = sorted( asset_ids, - key=lambda aid: (asset_use_counts.get(aid, 0), random.random()), + key=lambda aid: ( + -asset_smart_scores.get(aid, 0.0), + asset_use_counts.get(aid, 0), + random.random(), + ), ) for candidate in sorted_candidates: candidate_total = asset_durations.get(candidate, 0.0) diff --git a/apps/api/app/services/plan_generator_service.py b/apps/api/app/services/plan_generator_service.py index 12b1c0015..f000493dd 100755 --- a/apps/api/app/services/plan_generator_service.py +++ b/apps/api/app/services/plan_generator_service.py @@ -32,6 +32,7 @@ from packages.domain.plan_generator_utils import ( generate_default_clips, map_clip_types_for_mode, ) +from packages.domain.smart_match import score_asset from packages.domain.template_clip_config import TemplateClipConfig logger = logging.getLogger(__name__) @@ -218,8 +219,13 @@ class PlanGeneratorService: ) -> None: """按 editing_mode 将素材分配到 clips(就地修改,未持久化). - 委托给 plan_generator_utils.distribute_assets 纯函数。 + 先用 smart_match 评分对素材排序(高分优先),再委托给 + plan_generator_utils.distribute_assets 纯函数完成分配。 """ + # 用 smart_match 评分排序素材:高分(质量好/时长合适/新鲜/未使用)优先 + if self._asset_repo and not random_selection: + asset_ids = self._sort_assets_by_smart_score(asset_ids) + distribute_assets( clips, asset_ids, @@ -228,6 +234,23 @@ class PlanGeneratorService: asset_durations=asset_durations, ) + def _sort_assets_by_smart_score(self, asset_ids: List[str]) -> List[str]: + """按 smart_match 综合评分降序排列素材 ID。 + + 评分高的素材(质量好、时长合适、新鲜、使用次数少)排在前面。 + """ + scored: list[tuple[str, float]] = [] + for asset_id in asset_ids: + asset = self._asset_repo.get(asset_id) + if asset: + score, _ = score_asset(asset) + scored.append((asset_id, score)) + else: + scored.append((asset_id, 0.0)) + # 按评分降序排列 + scored.sort(key=lambda x: x[1], reverse=True) + return [aid for aid, _ in scored] + def _fetch_asset_durations(self, asset_ids: List[str]) -> dict[str, float]: """从数据库获取素材时长信息. diff --git a/tests/unit/test_editor_clips_random_start.py b/tests/unit/test_editor_clips_random_start.py index d3fb9a1ee..cae16fcf6 100644 --- a/tests/unit/test_editor_clips_random_start.py +++ b/tests/unit/test_editor_clips_random_start.py @@ -58,6 +58,10 @@ def _make_mock_asset(asset_id, duration): asset = MagicMock() asset.id = asset_id asset.duration = duration + # score_asset 所需的属性(避免 MagicMock 导致类型比较错误) + asset.quality_score = None + asset.created_at = None + asset.metadata = {} return asset diff --git a/tests/unit/test_mediakit_smart_clips.py b/tests/unit/test_mediakit_smart_clips.py index edf3d18a6..5178fe36d 100644 --- a/tests/unit/test_mediakit_smart_clips.py +++ b/tests/unit/test_mediakit_smart_clips.py @@ -347,6 +347,10 @@ def _make_rich_asset(asset_id, duration, storage_key="v.mp4", mime="video/mp4"): asset.duration = duration asset.storage_key = storage_key asset.mime_type = mime + # score_asset 所需的属性(避免 MagicMock 导致类型比较错误) + asset.quality_score = None + asset.created_at = None + asset.metadata = {} return asset diff --git a/tests/unit/test_plan_generator.py b/tests/unit/test_plan_generator.py index 10d02f62c..37544c067 100755 --- a/tests/unit/test_plan_generator.py +++ b/tests/unit/test_plan_generator.py @@ -908,6 +908,10 @@ class TestAssetDurationsAlwaysFetched: def fake_get(asset_id): mock_asset = MagicMock() mock_asset.duration = 30.0 # 每个素材 30 秒 + # score_asset 所需的属性 + mock_asset.quality_score = None + mock_asset.created_at = None + mock_asset.metadata = {} return mock_asset asset_repo.get = MagicMock(side_effect=fake_get) diff --git a/tests/unit/test_smart_match_integration.py b/tests/unit/test_smart_match_integration.py new file mode 100644 index 000000000..6ef23efad --- /dev/null +++ b/tests/unit/test_smart_match_integration.py @@ -0,0 +1,324 @@ +"""测试 smart_match 评分集成到素材选取路径。 + +验证: +- 使用次数多的素材评分低于使用次数少的(unused 维度降权生效) +- from-assets 路径中 sorted_candidates 按 smart_match 评分排序 +- 一键生成路径中 _sort_assets_by_smart_score 按评分降序 +""" + +from __future__ import annotations + +import os +import sys +from dataclasses import dataclass, field +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from unittest.mock import MagicMock, patch + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api")) + +import pytest + +from packages.domain.smart_match import score_asset, smart_select_assets + +# ── 辅助工厂 ────────────────────────────────────────────────────────────────── + + +@dataclass +class FakeAsset: + """Minimal Asset-like object.""" + + id: str + quality_score: float | None = 70.0 + duration: float = 15.0 + created_at: datetime | None = None + metadata: dict[str, Any] = field(default_factory=dict) + status: str = "ready" + file_type: str = "video" + + +def _asset_with_use_count(asset_id: str, use_count: int) -> FakeAsset: + """创建指定使用次数的素材,其他维度保持一致。""" + return FakeAsset( + id=asset_id, + quality_score=70.0, + duration=15.0, # 最优区间 5-30s + created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + metadata={"generation_use_count": use_count}, + ) + + +# ── score_asset 单元测试:unused 维度降权 ───────────────────────────────────── + + +class TestScoreAssetUnusedDiminsh: + """验证 unused 维度:使用次数越多,评分越低。""" + + def test_unused_scores_higher_than_used(self): + """use_count=0 的素材评分高于 use_count>0 的。""" + fresh = _asset_with_use_count("fresh", 0) + used = _asset_with_use_count("used", 1) + fresh_score, _ = score_asset(fresh) + used_score, _ = score_asset(used) + assert fresh_score > used_score + + def test_high_use_count_scores_lower_than_low(self): + """use_count=5 的素材评分低于 use_count=1 的。""" + low_use = _asset_with_use_count("low", 1) + high_use = _asset_with_use_count("high", 5) + low_score, _ = score_asset(low_use) + high_score, _ = score_asset(high_use) + assert low_score > high_score + + def test_unused_breakdown_values(self): + """验证 unused 维度的具体分值。""" + fresh = _asset_with_use_count("fresh", 0) + low = _asset_with_use_count("low", 2) + high = _asset_with_use_count("high", 10) + + _, fresh_bd = score_asset(fresh) + _, low_bd = score_asset(low) + _, high_bd = score_asset(high) + + # use_count=0 → unused_score=100 → component=10.0 + assert fresh_bd["unused"] == 10.0 + # use_count=2 → unused_score=70 → component=7.0 + assert low_bd["unused"] == 7.0 + # use_count=10 → unused_score=30 → component=3.0 + assert high_bd["unused"] == 3.0 + + def test_monotonically_decreasing_scores(self): + """使用次数递增时,总评分单调不增。""" + scores = [] + for count in [0, 1, 2, 3, 5, 10, 50]: + a = _asset_with_use_count(f"a{count}", count) + s, _ = score_asset(a) + scores.append(s) + # 验证非递增 + for i in range(len(scores) - 1): + assert ( + scores[i] >= scores[i + 1] + ), f"use_count 递增时评分应不增: scores[{i}]={scores[i]} < scores[{i+1}]={scores[i+1]}" + + +# ── smart_select_assets 排序测试 ───────────────────────────────────────────── + + +class TestSmartSelectAssetsOrdering: + """验证 smart_select_assets 返回结果按评分降序。""" + + def test_less_used_assets_ranked_higher(self): + """使用次数少的素材在结果中排名更高。""" + assets = [ + _asset_with_use_count("heavily_used", 10), + _asset_with_use_count("never_used", 0), + _asset_with_use_count("lightly_used", 2), + ] + results = smart_select_assets(assets) + ids = [r.asset.id for r in results] + # never_used 排第一,heavily_used 排最后 + assert ids[0] == "never_used" + assert ids[-1] == "heavily_used" + + def test_same_quality_different_use_count(self): + """质量相同时,使用次数少的排名更高。""" + assets = [ + _asset_with_use_count("used_5", 5), + _asset_with_use_count("used_0", 0), + ] + results = smart_select_assets(assets) + assert results[0].asset.id == "used_0" + assert results[1].asset.id == "used_5" + + +# ── from-assets 路径集成测试 ───────────────────────────────────────────────── + + +def _make_mock_asset_for_clips(aid, duration, use_count=0): + """创建带 score_asset 所需属性的 mock 素材。""" + asset = MagicMock() + asset.id = aid + asset.duration = duration + asset.quality_score = None + asset.created_at = None + asset.metadata = {"generation_use_count": use_count} + return asset + + +def _make_auth_user(): + auth = MagicMock() + auth.user.id = "user-001" + auth.user.email = "test@example.com" + auth.user.display_name = "测试用户" + auth.user_id = "user-001" + return auth + + +class TestFromAssetsSmartMatchIntegration: + """验证 clips.py 中 sorted_candidates 使用 smart_match 评分。""" + + def test_sorted_candidates_prefers_high_score_low_use_count(self): + """在 from-assets 路径中,smart_match 分高且使用次数少的素材排在前面。""" + from app.api.routes.templates_editor.clips import create_clips_from_assets_editor + from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest + + mock_asset_repo = MagicMock() + + def _get_asset(aid): + use_count = {"a_heavy": 10, "a_fresh": 0}[aid] + return _make_mock_asset_for_clips(aid, 30.0, use_count) + + mock_asset_repo.get = MagicMock(side_effect=_get_asset) + + mock_plan_svc = MagicMock() + mock_plan_svc.replace_all_clips_transactional = MagicMock(return_value=2) + + segments = [(0, 3.0, 5.0), (1, 3.0, 5.0)] + + with ( + patch( + "app.api.routes.templates_editor.clips._get_template_segments", + return_value=segments, + ), + patch( + "app.api.routes.templates_editor.clips.get_used_segments", + return_value={}, + ), + patch( + "app.api.routes.templates_editor.clips.record_used_segments", + return_value=None, + ), + ): + body = ClipsFromAssetsRequest( + asset_ids=["a_heavy", "a_fresh"], + required_clips_count=2, + ) + create_clips_from_assets_editor( + template_id="tmpl-1", + body=body, + background_tasks=MagicMock(), + plan_id="test-plan-001", + services=(MagicMock(), mock_plan_svc), + asset_repo=mock_asset_repo, + db=MagicMock(), + current_user=_make_auth_user(), + ) + + # 验证 replace_all_clips_transactional 被调用 + assert mock_plan_svc.replace_all_clips_transactional.called + call_args = mock_plan_svc.replace_all_clips_transactional.call_args + clips_data = call_args.args[1] + + # 第一个片段应该分配给 a_fresh(smart_match 分更高) + first_clip_asset = clips_data[0]["asset_id"] + assert ( + first_clip_asset == "a_fresh" + ), f"第一个片段应分配给 smart_match 分更高的 a_fresh,实际是 {first_clip_asset}" + + +# ── 一键生成路径集成测试 ───────────────────────────────────────────────────── + + +class TestPlanGeneratorSmartMatchIntegration: + """验证 PlanGeneratorService._sort_assets_by_smart_score 排序正确。""" + + def test_sort_assets_by_smart_score_descending(self): + """_sort_assets_by_smart_score 返回按评分降序排列的素材 ID。""" + from app.services.plan_generator_service import PlanGeneratorService + + mock_asset_repo = MagicMock() + + def _get_asset(aid): + use_count = {"high_use": 10, "low_use": 0, "mid_use": 3}[aid] + asset = MagicMock() + asset.id = aid + asset.duration = 15.0 + asset.quality_score = None + asset.created_at = None + asset.metadata = {"generation_use_count": use_count} + return asset + + mock_asset_repo.get = MagicMock(side_effect=_get_asset) + + db = MagicMock() + svc = PlanGeneratorService(db, asset_repo=mock_asset_repo) + + sorted_ids = svc._sort_assets_by_smart_score(["high_use", "low_use", "mid_use"]) + + # low_use (0次) 应排第一,high_use (10次) 应排最后 + assert sorted_ids[0] == "low_use" + assert sorted_ids[-1] == "high_use" + assert sorted_ids[1] == "mid_use" + + def test_distribute_assets_uses_smart_score_ordering(self): + """_distribute_assets 在非随机模式下按 smart_match 评分排序素材。""" + from app.services.plan_generator_service import PlanGeneratorService + + from packages.domain.edit_plan_clip import EditPlanClip + from packages.domain.editing_mode import EditingMode + + mock_asset_repo = MagicMock() + + def _get_asset(aid): + use_count = {"old_asset": 10, "new_asset": 0}[aid] + asset = MagicMock() + asset.id = aid + asset.duration = 30.0 + asset.quality_score = None + asset.created_at = None + asset.metadata = {"generation_use_count": use_count} + return asset + + mock_asset_repo.get = MagicMock(side_effect=_get_asset) + + db = MagicMock() + svc = PlanGeneratorService(db, asset_repo=mock_asset_repo) + + # 创建 2 个 main clips(需要提供 id 参数) + clips = [ + EditPlanClip(id="c1", plan_id="p1", clip_type="main", duration=5.0, order=0), + EditPlanClip(id="c2", plan_id="p1", clip_type="main", duration=5.0, order=1), + ] + + with patch("app.services.plan_generator_service.distribute_assets") as mock_dist: + svc._distribute_assets( + clips, + ["old_asset", "new_asset"], + EditingMode.ONE_TAKE.value, + random_selection=False, + ) + # 验证传给 distribute_assets 的 asset_ids 按 smart_match 排序 + call_args = mock_dist.call_args + passed_ids = call_args.args[1] + # new_asset (0次使用) 应排在 old_asset (10次使用) 前面 + assert passed_ids[0] == "new_asset" + assert passed_ids[1] == "old_asset" + + def test_random_selection_skips_smart_score_sort(self): + """random_selection=True 时不执行 smart_match 排序。""" + from app.services.plan_generator_service import PlanGeneratorService + + from packages.domain.edit_plan_clip import EditPlanClip + from packages.domain.editing_mode import EditingMode + + mock_asset_repo = MagicMock() + db = MagicMock() + svc = PlanGeneratorService(db, asset_repo=mock_asset_repo) + + clips = [ + EditPlanClip(id="c1", plan_id="p1", clip_type="main", duration=5.0, order=0), + ] + + with patch("app.services.plan_generator_service.distribute_assets") as mock_dist: + svc._distribute_assets( + clips, + ["a1", "a2"], + EditingMode.ONE_TAKE.value, + random_selection=True, + ) + # random_selection=True 时不应调用 asset_repo.get(不执行排序) + mock_asset_repo.get.assert_not_called()