diff --git a/apps/api/app/services/plan_generator_service.py b/apps/api/app/services/plan_generator_service.py index 758510c0f..619c8abd0 100755 --- a/apps/api/app/services/plan_generator_service.py +++ b/apps/api/app/services/plan_generator_service.py @@ -131,6 +131,7 @@ class PlanGeneratorService: editing_mode, random_selection=random_preview, asset_durations=asset_durations, + user_id=created_by_user_id, ) # 5. 持久化所有 clips 并计算总时长 @@ -218,6 +219,7 @@ class PlanGeneratorService: *, random_selection: bool = False, asset_durations: dict[str, float] | None = None, + user_id: str = "", ) -> None: """按 editing_mode 将素材分配到 clips(就地修改,未持久化). @@ -239,6 +241,14 @@ class PlanGeneratorService: asset_ids = list(asset_ids) # 复制避免修改调用方原列表 random.shuffle(asset_ids) + # 查询已有视频的已用区间(跨视频避让) + external_used_segments = None + if user_id and self._clip_repo: + try: + external_used_segments = self._clip_repo.list_used_segments_by_user(user_id, limit_recent=50) + except Exception: + logger.warning("跨视频避让查询失败,回退到纯随机", exc_info=True) + distribute_assets( clips, asset_ids, @@ -246,6 +256,7 @@ class PlanGeneratorService: random_selection=random_selection, asset_durations=asset_durations, asset_scene_points=asset_scene_points, + external_used_segments=external_used_segments, ) def _fetch_asset_scene_points(self, asset_ids: List[str]) -> dict[str, list[float]]: diff --git a/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py b/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py index 4c241269f..16819092e 100755 --- a/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py +++ b/packages/adapters/sqlalchemy_impl/edit_plan_clip_repository.py @@ -131,3 +131,65 @@ class SQLAlchemyEditPlanClipRepository: created_at=model.created_at, updated_at=model.updated_at, ) + + def list_used_segments_by_user( + self, + user_id: str, + *, + limit_recent: int = 50, + ) -> dict[str, list[tuple[float, float]]]: + """查询用户已有视频中已使用的素材区间(跨视频避让). + + JOIN edit_plans 表,按 created_by_user_id 过滤,只查 status='completed' + 的 plan 下 status='rendered' 且 asset_id 非空的 clips。按 plan 的 + created_at DESC 取最近 limit_recent 个 plan。 + + Returns: + {asset_id: [(start_time, start_time + duration), ...]} + 空结果返回空 dict。 + """ + from packages.adapters.sqlalchemy_impl.models import EditPlanModel + + if not user_id: + return {} + + # 1. 查出最近 limit_recent 个已完成 plan 的 ID + recent_plan_ids = [ + row[0] + for row in self.session.query(EditPlanModel.id) + .filter( + EditPlanModel.created_by_user_id == user_id, + EditPlanModel.status == "completed", + ) + .order_by(EditPlanModel.created_at.desc()) + .limit(limit_recent) + .all() + ] + + if not recent_plan_ids: + return {} + + # 2. 查这些 plan 下已渲染、有素材的 clips + clips = ( + self.session.query( + EditPlanClipModel.asset_id, + EditPlanClipModel.start_time, + EditPlanClipModel.duration, + ) + .filter( + EditPlanClipModel.plan_id.in_(recent_plan_ids), + EditPlanClipModel.status == "rendered", + EditPlanClipModel.asset_id != "", + EditPlanClipModel.asset_id.isnot(None), + ) + .all() + ) + + # 3. 聚合为 {asset_id: [(start, start+duration), ...]} + result: dict[str, list[tuple[float, float]]] = {} + for asset_id, start_time, duration in clips: + if asset_id not in result: + result[asset_id] = [] + result[asset_id].append((start_time or 0.0, (start_time or 0.0) + (duration or 0.0))) + + return result diff --git a/packages/domain/plan_generator_utils.py b/packages/domain/plan_generator_utils.py index 446f6da21..e7530472f 100755 --- a/packages/domain/plan_generator_utils.py +++ b/packages/domain/plan_generator_utils.py @@ -169,6 +169,7 @@ def distribute_assets( random_selection: bool = False, asset_durations: dict[str, float] | None = None, asset_scene_points: dict[str, list[float]] | None = None, + external_used_segments: dict[str, list[tuple[float, float]]] | None = None, ) -> None: """按 editing_mode 将素材分配到 clips(就地修改). @@ -188,6 +189,7 @@ def distribute_assets( random_selection: 是否随机选择素材(用于预览生成) asset_durations: 素材 ID -> 时长(秒)映射,用于设置 start_time asset_scene_points: 素材 ID -> 场景切换点列表(metadata 缓存) + external_used_segments: 跨视频已用区间(来自其他视频的 clips),注入到分配逻辑中避让 """ if not asset_ids or not clips: return @@ -198,16 +200,16 @@ def distribute_assets( random.shuffle(asset_ids) if editing_mode == EditingMode.ONE_TAKE.value: - _distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points) + _distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments) elif editing_mode == EditingMode.PIP.value: - _distribute_pip(clips, asset_ids, asset_durations, asset_scene_points) + _distribute_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments) elif editing_mode == EditingMode.VOICE_OVER.value: - _distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points) + _distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments) elif editing_mode == EditingMode.VOICE_PIP.value: - _distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points) + _distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments) else: # 未知模式,退化为 one_take - _distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points) + _distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments) def _resolve_start_time( @@ -248,9 +250,12 @@ def _distribute_one_take( asset_ids: List[str], asset_durations: dict[str, float] | None = None, asset_scene_points: dict[str, list[float]] | None = None, + external_used_segments: dict[str, list[tuple[float, float]]] | None = None, ) -> None: """ONE_TAKE: 素材按顺序依次分配给 main 类型 clips.""" - used_segments: dict[str, list[tuple[float, float]]] = {} + used_segments: dict[str, list[tuple[float, float]]] = ( + {k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {} + ) main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value] for i, clip in enumerate(main_clips): if i < len(asset_ids): @@ -271,9 +276,12 @@ def _distribute_pip( asset_ids: List[str], asset_durations: dict[str, float] | None = None, asset_scene_points: dict[str, list[float]] | None = None, + external_used_segments: dict[str, list[tuple[float, float]]] | None = None, ) -> None: """PIP: 第1个素材→main(全屏背景),其余→overlay clips.""" - used_segments: dict[str, list[tuple[float, float]]] = {} + used_segments: dict[str, list[tuple[float, float]]] = ( + {k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {} + ) # 第1个素材 → main clip main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value] if main_clips and asset_ids: @@ -310,9 +318,12 @@ def _distribute_voice_over( asset_ids: List[str], asset_durations: dict[str, float] | None = None, asset_scene_points: dict[str, list[float]] | None = None, + external_used_segments: dict[str, list[tuple[float, float]]] | None = None, ) -> None: """VOICE_OVER: 素材→main clips (B-roll).""" - used_segments: dict[str, list[tuple[float, float]]] = {} + used_segments: dict[str, list[tuple[float, float]]] = ( + {k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {} + ) main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value] for i, clip in enumerate(main_clips): if i < len(asset_ids): @@ -333,9 +344,12 @@ def _distribute_voice_pip( asset_ids: List[str], asset_durations: dict[str, float] | None = None, asset_scene_points: dict[str, list[float]] | None = None, + external_used_segments: dict[str, list[tuple[float, float]]] | None = None, ) -> None: """VOICE_PIP: 第1个→background, 第2个→corner_voice, 其余→b_roll.""" - used_segments: dict[str, list[tuple[float, float]]] = {} + used_segments: dict[str, list[tuple[float, float]]] = ( + {k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {} + ) bg_clips = [c for c in clips if c.clip_type == "background"] voice_clips = [c for c in clips if c.clip_type == "corner_voice"] broll_clips = [c for c in clips if c.clip_type == "b_roll"] diff --git a/tests/unit/test_cross_video_avoidance.py b/tests/unit/test_cross_video_avoidance.py new file mode 100644 index 000000000..316e105ee --- /dev/null +++ b/tests/unit/test_cross_video_avoidance.py @@ -0,0 +1,335 @@ +"""Tests for Issue #1670 — 跨视频片段避让(生成前注入已用区间).""" + +from __future__ import annotations + +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest + +from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import ( + SQLAlchemyEditPlanClipRepository, +) +from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus +from packages.domain.plan_generator_utils import ( + _distribute_one_take, + distribute_assets, +) + +# ── Repository 层测试 ───────────────────────────────────────────────────────── + + +class TestListUsedSegmentsByUser: + """测试 list_used_segments_by_user 方法.""" + + def _make_repo(self, session_mock): + return SQLAlchemyEditPlanClipRepository(session_mock) + + def test_empty_user_id_returns_empty_dict(self): + """空 user_id 直接返回空 dict,不查 DB.""" + session = MagicMock() + repo = self._make_repo(session) + result = repo.list_used_segments_by_user("") + assert result == {} + session.query.assert_not_called() + + def test_no_completed_plans_returns_empty_dict(self): + """用户没有已完成的 plan 时返回空 dict.""" + session = MagicMock() + # Mock plan query returns empty + plan_query = MagicMock() + plan_query.filter.return_value = plan_query + plan_query.order_by.return_value = plan_query + plan_query.limit.return_value = plan_query + plan_query.all.return_value = [] + session.query.return_value = plan_query + + repo = self._make_repo(session) + result = repo.list_used_segments_by_user("user_123") + assert result == {} + + def test_aggregates_clips_from_multiple_plans(self): + """从多个已完成 plan 的 clips 聚合已用区间.""" + session = MagicMock() + + # Mock plan query: 2 completed plans + plan_query = MagicMock() + plan_query.filter.return_value = plan_query + plan_query.order_by.return_value = plan_query + plan_query.limit.return_value = plan_query + plan_query.all.return_value = [("plan_1",), ("plan_2",)] + session.query.return_value = plan_query + + # Mock clip query: clips from both plans + clip_query = MagicMock() + clip_query.filter.return_value = clip_query + clip_query.all.return_value = [ + ("asset_A", 0.0, 5.0), # plan_1, asset A: 0~5s + ("asset_A", 10.0, 3.0), # plan_1, asset A: 10~13s + ("asset_B", 2.0, 4.0), # plan_2, asset B: 2~6s + ] + # Second session.query call is for clips + session.query.side_effect = [plan_query, clip_query] + + repo = self._make_repo(session) + result = repo.list_used_segments_by_user("user_123") + + assert "asset_A" in result + assert len(result["asset_A"]) == 2 + assert result["asset_A"][0] == (0.0, 5.0) + assert result["asset_A"][1] == (10.0, 13.0) + assert "asset_B" in result + assert result["asset_B"][0] == (2.0, 6.0) + + def test_respects_limit_recent_parameter(self): + """limit_recent 参数限制查询的 plan 数量.""" + session = MagicMock() + + plan_query = MagicMock() + plan_query.filter.return_value = plan_query + plan_query.order_by.return_value = plan_query + plan_query.limit.return_value = plan_query + plan_query.all.return_value = [("plan_1",)] + session.query.return_value = plan_query + + clip_query = MagicMock() + clip_query.filter.return_value = clip_query + clip_query.all.return_value = [("asset_X", 1.0, 2.0)] + session.query.side_effect = [plan_query, clip_query] + + repo = self._make_repo(session) + result = repo.list_used_segments_by_user("user_123", limit_recent=10) + + # Verify limit was called with the parameter + plan_query.limit.assert_called_once_with(10) + assert "asset_X" in result + + +# ── Domain 层测试 ───────────────────────────────────────────────────────────── + + +class TestDistributeAssetsWithExternalSegments: + """测试 distribute_assets 传入 external_used_segments 的行为.""" + + def _make_clips(self, count: int, duration: float = 3.0) -> list[EditPlanClip]: + """创建指定数量的 MAIN 类型 clips.""" + return [ + EditPlanClip( + id=f"clip_{i}", + plan_id="plan_1", + clip_type="main", + order=i, + template_clip_config_id="", + asset_id="", + text_content="", + start_time=0.0, + duration=duration, + status=EditPlanClipStatus.PENDING, + ) + for i in range(count) + ] + + def test_external_used_segments_none_backward_compatible(self): + """external_used_segments=None 时行为不变(向后兼容).""" + clips = self._make_clips(3) + asset_ids = ["asset_1", "asset_2", "asset_3"] + asset_durations = {aid: 30.0 for aid in asset_ids} + + # Should not raise + distribute_assets( + clips, + asset_ids, + "one_take", + asset_durations=asset_durations, + external_used_segments=None, + ) + + # All clips should have assets assigned + for clip in clips: + assert clip.asset_id != "" + + def test_external_used_segments_avoids_existing_ranges(self): + """传入 external_used_segments 后,新分配的 start_time 避开已有区间.""" + clips = self._make_clips(2, duration=3.0) + asset_ids = ["asset_1"] + asset_durations = {"asset_1": 30.0} + + # Pretend asset_1 0~10s is already used by another video + external = {"asset_1": [(0.0, 10.0)]} + + # Run multiple times to check that start_time always avoids 0~10s + # (with some randomness, but the avoidance should be consistent) + for _ in range(10): + test_clips = self._make_clips(1, duration=3.0) + distribute_assets( + test_clips, + asset_ids, + "one_take", + asset_durations=asset_durations, + external_used_segments=external, + ) + start = test_clips[0].start_time + # Start time + duration (3s) should not overlap with 0~10 + # i.e., start >= 10.0 or start + 3 <= 0.0 (impossible since start >= 0) + assert ( + start >= 10.0 or start + 3.0 <= 0.0 or start >= 10.0 + ), f"start_time {start} overlaps with existing segment 0~10" + + def test_external_used_segments_deep_copy(self): + """external_used_segments 会被深拷贝,不会修改外部数据.""" + external = {"asset_1": [(0.0, 5.0)]} + original = {"asset_1": [(0.0, 5.0)]} + + clips = self._make_clips(1, duration=2.0) + asset_ids = ["asset_1"] + asset_durations = {"asset_1": 20.0} + + distribute_assets( + clips, + asset_ids, + "one_take", + asset_durations=asset_durations, + external_used_segments=external, + ) + + # External dict should be unchanged + assert external == original + + def test_empty_external_used_segments_same_as_none(self): + """空 dict 的 external_used_segments 行为与 None 相同.""" + clips = self._make_clips(2, duration=3.0) + asset_ids = ["asset_1", "asset_2"] + asset_durations = {aid: 30.0 for aid in asset_ids} + + # Should not raise and should assign assets normally + distribute_assets( + clips, + asset_ids, + "one_take", + asset_durations=asset_durations, + external_used_segments={}, + ) + for clip in clips: + assert clip.asset_id != "" + + +# ── Service 层测试 ──────────────────────────────────────────────────────────── + + +class TestServiceLayerIntegration: + """测试 _distribute_assets 在 service 层的查询逻辑.""" + + def _make_service(self, clip_repo_mock, asset_repo_mock=None): + """创建 PlanGeneratorService 并注入 mock repos.""" + from unittest.mock import MagicMock, patch + + from apps.api.app.services.plan_generator_service import PlanGeneratorService + + with ( + patch("apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanRepository"), + patch( + "apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanClipRepository", + return_value=clip_repo_mock, + ), + ): + db = MagicMock() + svc = PlanGeneratorService(db, asset_repo=asset_repo_mock) + svc._clip_repo = clip_repo_mock + return svc + + def _make_clip(self): + return EditPlanClip( + id="clip_1", + plan_id="plan_1", + clip_type="main", + order=0, + template_clip_config_id="", + asset_id="", + text_content="", + start_time=0.0, + duration=3.0, + status=EditPlanClipStatus.PENDING, + ) + + def test_query_called_with_user_id(self): + """有 user_id 时调用 list_used_segments_by_user.""" + clip_repo = MagicMock() + clip_repo.list_used_segments_by_user.return_value = {"asset_A": [(0.0, 5.0)]} + asset_repo = MagicMock() + asset_repo.get.return_value = None # smart_match fallback + + svc = self._make_service(clip_repo, asset_repo) + clips = [self._make_clip()] + + svc._distribute_assets( + clips, + ["asset_A"], + "one_take", + asset_durations={"asset_A": 30.0}, + user_id="user_123", + ) + + clip_repo.list_used_segments_by_user.assert_called_once_with("user_123", limit_recent=50) + + def test_query_not_called_without_user_id(self): + """无 user_id 时不调用查询.""" + clip_repo = MagicMock() + asset_repo = MagicMock() + asset_repo.get.return_value = None + + svc = self._make_service(clip_repo, asset_repo) + clips = [self._make_clip()] + + svc._distribute_assets( + clips, + ["asset_A"], + "one_take", + asset_durations={"asset_A": 30.0}, + user_id="", + ) + + clip_repo.list_used_segments_by_user.assert_not_called() + + def test_query_failure_does_not_block_generation(self): + """查询失败时不阻塞生成,回退到纯随机.""" + clip_repo = MagicMock() + clip_repo.list_used_segments_by_user.side_effect = Exception("DB error") + asset_repo = MagicMock() + asset_repo.get.return_value = None + + svc = self._make_service(clip_repo, asset_repo) + clips = [self._make_clip()] + + # Should not raise + svc._distribute_assets( + clips, + ["asset_A"], + "one_take", + asset_durations={"asset_A": 30.0}, + user_id="user_123", + ) + + # Clip should still get an asset assigned (fallback to random) + assert clips[0].asset_id == "asset_A" + + def test_preview_and_final_both_query(self): + """预览和正式生成都触发查询.""" + for random_selection in [True, False]: + clip_repo = MagicMock() + clip_repo.list_used_segments_by_user.return_value = {} + asset_repo = MagicMock() + asset_repo.get.return_value = None + + svc = self._make_service(clip_repo, asset_repo) + clips = [self._make_clip()] + + svc._distribute_assets( + clips, + ["asset_A"], + "one_take", + random_selection=random_selection, + asset_durations={"asset_A": 30.0}, + user_id="user_123", + ) + + clip_repo.list_used_segments_by_user.assert_called_once()