"""#1970 原子片段级选片核心单元测试(纯函数,不依赖 DB).""" from __future__ import annotations import random from packages.domain.asset_atom_clip import AssetAtomClip from packages.domain.atom_clip_selector import ( clips_to_segments, estimate_required_clip_count, reselect_clips_from_atoms, score_atom_clip, select_atom_clips, ) from packages.domain.atom_clip_service import compute_atom_clips def _clip(asset_id: str, start: float, end: float, clip_id: str = "") -> AssetAtomClip: return ( AssetAtomClip.create( asset_id=asset_id, start_time=start, end_time=end, clip_index=int(start), ) if not clip_id else AssetAtomClip( id=clip_id, asset_id=asset_id, start_time=start, end_time=end, duration=round(end - start, 3), clip_index=0, ) ) class TestEstimateCount: def test_basic(self): assert estimate_required_clip_count(30.0, 4.5) == 7 assert estimate_required_clip_count(18.0, 4.0) == round(18 / 4) def test_invalid_inputs_returns_one(self): assert estimate_required_clip_count(0) == 1 assert estimate_required_clip_count(10, 0) == 1 assert estimate_required_clip_count(-1) == 1 class TestScore: def test_unused_beats_used(self): c = _clip("a1", 0, 4) s_unused = score_atom_clip(c, target_duration=4.0, used_in_video=set()) s_used = score_atom_clip(c, target_duration=4.0, used_in_video={c.id}) assert s_unused > s_used def test_duration_fit_better_when_closer(self): target = 4.0 exact = score_atom_clip(_clip("a", 0, 4.0), target_duration=target) short = score_atom_clip(_clip("b", 0, 1.5), target_duration=target) assert exact > short def test_history_penalty(self): c = _clip("a1", 0, 4) normal = score_atom_clip(c, target_duration=4.0) penalized = score_atom_clip(c, target_duration=4.0, recently_used={c.id}) assert normal > penalized def test_asset_balance_penalizes_repeated_asset(self): c1 = _clip("a", 0, 4) first = score_atom_clip(c1, target_duration=4.0, asset_usage_counts={}) third = score_atom_clip(c1, target_duration=4.0, asset_usage_counts={"a": 2}) assert first > third class TestSelect: def test_no_duplicate_atom_within_video(self): pool = compute_atom_clips("a", 30.0, rng=random.Random(1)) used: set[str] = set() usage: dict[str, int] = {} chosen = [] rng = random.Random(5) for _ in range(4): ranked = select_atom_clips( pool, target_duration=4.0, used_atom_clip_ids=used, asset_usage_counts=usage, required_count=4, limit=1, rng=rng, ) assert ranked pick = ranked[0] assert pick.atom_clip_id not in used chosen.append(pick) used.add(pick.atom_clip_id) usage[pick.asset_id] = usage.get(pick.asset_id, 0) + 1 assert len(used) == 4 def test_same_asset_different_clips_allowed(self): pool = compute_atom_clips("a", 30.0, rng=random.Random(2)) used: set[str] = set() usage: dict[str, int] = {} rng = random.Random(7) picked_assets = set() for _ in range(3): pick = select_atom_clips( pool, target_duration=4.0, used_atom_clip_ids=used, asset_usage_counts=usage, limit=1, rng=rng, )[0] used.add(pick.atom_clip_id) usage[pick.asset_id] = usage.get(pick.asset_id, 0) + 1 picked_assets.add(pick.asset_id) # 单素材池允许同素材多片段 assert picked_assets == {"a"} assert len(used) == 3 def test_exhausted_pool_returns_empty(self): pool = [_clip("a", 0, 4)] ranked = select_atom_clips(pool, used_atom_clip_ids={pool[0].id}, target_duration=4.0) assert ranked == [] def test_recently_used_deprioritized_not_hard_blocked(self): # 两个片段,recent 中包含更合适的那个;它应被降权但不会从候选中消失 fresh = _clip("a", 0, 2.0, clip_id="fresh") recent = _clip("b", 0, 4.0, clip_id="recent") ranked = select_atom_clips( [fresh, recent], target_duration=4.0, recently_used_atom_ids={"recent"}, limit=2, rng=random.Random(0), # 噪声 0 不影响 ) ids = [r.atom_clip_id for r in ranked] assert set(ids) == {"fresh", "recent"} # 降权 + 噪声可能导致排序不稳定,只验证 recent 仍在候选中(不硬禁) def test_limit(self): pool = compute_atom_clips("a", 40.0, rng=random.Random(4)) ranked = select_atom_clips(pool, target_duration=4.0, limit=3) assert len(ranked) == 3 scores = [r.score for r in ranked] assert scores == sorted(scores, reverse=True) class TestClipsToSegments: def test_grouped_by_asset_sorted(self): clips = [ _clip("a", 10, 14), _clip("a", 0, 4), _clip("b", 2, 6), ] segs = clips_to_segments(clips) assert segs["a"] == [(0, 4), (10, 14)] assert segs["b"] == [(2, 6)] class TestReselectFromAtoms: def _src(self, n): return [{"order": i, "clip_type": "main", "duration": 4.0, "start_time": 0.0} for i in range(n)] def test_skeleton_preserved_and_unique(self): pool = compute_atom_clips("a", 30.0, rng=random.Random(11)) + compute_atom_clips( "b", 30.0, rng=random.Random(12) ) out = reselect_clips_from_atoms(self._src(5), pool, rng=random.Random(13)) assert out is not None assert len(out) == 5 ids = [c["atom_clip_id"] for c in out] assert len(set(ids)) == 5 for c in out: assert c["asset_id"] assert c["start_time"] >= 0 assert c["duration"] > 0 def test_insufficient_candidates_returns_none(self): pool = compute_atom_clips("a", 10.0, rng=random.Random(1)) assert reselect_clips_from_atoms(self._src(20), pool) is None def test_non_main_clips_left_untouched(self): pool = compute_atom_clips("a", 30.0, rng=random.Random(8)) src = [ {"order": 0, "clip_type": "intro", "duration": 2.0, "asset_id": "fixed"}, {"order": 1, "clip_type": "main", "duration": 4.0}, ] out = reselect_clips_from_atoms(src, pool, rng=random.Random(3)) assert out is not None assert out[0]["asset_id"] == "fixed" assert "atom_clip_id" not in out[0] assert out[1].get("atom_clip_id") def test_empty_inputs(self): assert reselect_clips_from_atoms([], [_clip("a", 0, 4)]) is None assert reselect_clips_from_atoms(self._src(2), []) is None def test_batch_used_excluded(self): pool = compute_atom_clips("a", 30.0, rng=random.Random(21)) batch_used = {pool[0].id} out = reselect_clips_from_atoms(self._src(3), pool, batch_used_atom_ids=batch_used, rng=random.Random(22)) assert out is not None assert pool[0].id not in {c["atom_clip_id"] for c in out}