Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f1a19acc61 | |||
| 382e7a24fc |
@@ -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)
|
||||
|
||||
@@ -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]:
|
||||
"""从数据库获取素材时长信息.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user