diff --git a/tests/unit/test_1749_edit_plan_service_new_methods.py b/tests/unit/test_1749_edit_plan_service_new_methods.py new file mode 100644 index 000000000..bf59519a5 --- /dev/null +++ b/tests/unit/test_1749_edit_plan_service_new_methods.py @@ -0,0 +1,197 @@ +"""#1749 EditPlanService 新增方法单元测试。 + +覆盖 get_asset_durations / apply_voice_duration_to_plan / ensure_variant_plans。 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + + +# ── helpers ─────────────────────────────────────────────────────────────────── + +def _make_svc(): + """构造 EditPlanService,mock 掉 DB/repos。""" + from app.services.edit_plan_service import EditPlanService + + db = MagicMock() + svc = EditPlanService.__new__(EditPlanService) + svc._plan_repo = MagicMock() + svc._clip_repo = MagicMock() + svc._clip_repo.session = db + svc._generation_task_repo = MagicMock() + return svc, db + + +def _make_clip( + order=0, + asset_id="a1", + start_time=0.0, + duration=5.0, + clip_type="main", + transition_effect="cut", + transition_duration=0.0, + playback_speed=1.0, + text_content="", + config=None, +): + c = MagicMock() + c.order = order + c.asset_id = asset_id + c.start_time = start_time + c.duration = duration + c.clip_type = clip_type + c.transition_effect = transition_effect + c.transition_duration = transition_duration + c.playback_speed = playback_speed + c.text_content = text_content + c.config = config or {} + return c + + +# ═══════════════════════════════════════════════════════════════════════════════ +# get_asset_durations +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestGetAssetDurations: + def test_empty_ids_returns_empty(self): + svc, _db = _make_svc() + assert svc.get_asset_durations([]) == {} + assert svc.get_asset_durations(None) == {} + + def test_dedup_and_query(self): + svc, db = _make_svc() + m1 = MagicMock() + m1.id = "a1" + m1.duration = 10.5 + m2 = MagicMock() + m2.id = "a2" + m2.duration = None # -> 0.0 + + db.query.return_value.filter.return_value.all.return_value = [m1, m2] + + result = svc.get_asset_durations(["a1", "a2", "a1"]) + assert result == {"a1": 10.5, "a2": 0.0} + + def test_bad_duration_type_returns_zero(self): + svc, db = _make_svc() + m = MagicMock() + m.id = "a1" + m.duration = "not-a-number" + db.query.return_value.filter.return_value.all.return_value = [m] + result = svc.get_asset_durations(["a1"]) + assert result == {"a1": 0.0} + + +# ═══════════════════════════════════════════════════════════════════════════════ +# apply_voice_duration_to_plan +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestApplyVoiceDurationToPlan: + def test_voice_le_zero_returns_none(self): + svc, _ = _make_svc() + assert svc.apply_voice_duration_to_plan("p1", 0) is None + assert svc.apply_voice_duration_to_plan("p1", -5) is None + + def test_plan_not_found_returns_none(self): + svc, _ = _make_svc() + svc.get_plan = MagicMock(return_value=None) + assert svc.apply_voice_duration_to_plan("p1", 30.0) is None + + def test_no_clips_returns_none(self): + svc, _ = _make_svc() + plan = MagicMock() + svc.get_plan = MagicMock(return_value=plan) + svc._clip_repo.list_by_plan.return_value = [] + assert svc.apply_voice_duration_to_plan("p1", 30.0) is None + + def test_happy_path_two_clips(self): + svc, db = _make_svc() + plan = MagicMock() + plan.total_duration = 0.0 + svc.get_plan = MagicMock(return_value=plan) + + clips = [ + _make_clip(order=0, asset_id="a1", start_time=0.0), + _make_clip(order=1, asset_id="a2", start_time=0.0), + ] + svc._clip_repo.list_by_plan.return_value = clips + + # get_asset_durations: a1=100s, a2=50s + m1 = MagicMock(id="a1", duration=100.0) + m2 = MagicMock(id="a2", duration=50.0) + db.query.return_value.filter.return_value.all.return_value = [m1, m2] + + svc.replace_all_clips_transactional = MagicMock(return_value=2) + + result = svc.apply_voice_duration_to_plan("p1", 30.0) + assert result is plan + svc.replace_all_clips_transactional.assert_called_once() + + def test_invalid_voice_returns_none(self): + svc, _ = _make_svc() + assert svc.apply_voice_duration_to_plan("p1", "bad") is None + + +# ═══════════════════════════════════════════════════════════════════════════════ +# ensure_variant_plans +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestEnsureVariantPlans: + def _setup_svc(self): + svc, _ = _make_svc() + plan0 = MagicMock() + plan0.id = "plan_v0" + svc.clone_plan_for_variant = MagicMock(return_value=plan0) + + plan1 = MagicMock() + plan1.id = "plan_v1" + svc.reselect_plan_for_variant = MagicMock(return_value=plan1) + + svc.apply_voice_duration_to_plan = MagicMock(return_value=plan0) + return svc + + def test_count_1_no_voice(self): + svc = self._setup_svc() + result = svc.ensure_variant_plans("src", 1, ["a1", "a2"]) + assert result == ["plan_v0"] + svc.clone_plan_for_variant.assert_called_once() + svc.apply_voice_duration_to_plan.assert_not_called() + + def test_count_1_with_voice(self): + svc = self._setup_svc() + result = svc.ensure_variant_plans( + "src", 1, ["a1"], voice_durations=[30.0] + ) + assert result == ["plan_v0"] + svc.apply_voice_duration_to_plan.assert_called_once_with("plan_v0", 30.0) + + def test_count_3_with_voice(self): + svc = self._setup_svc() + svc.reselect_plan_for_variant = MagicMock(side_effect=[ + MagicMock(id="plan_v1"), + MagicMock(id="plan_v2"), + ]) + + result = svc.ensure_variant_plans( + "src", 3, ["a1", "a2"], voice_durations=[30.0, 20.0, 15.0] + ) + assert len(result) == 3 + assert result[0] == "plan_v0" + # v0 配音 + svc.apply_voice_duration_to_plan.assert_called_once_with("plan_v0", 30.0) + # reselect 被调 2 次 + assert svc.reselect_plan_for_variant.call_count == 2 + + def test_voice_apply_failure_not_blocking(self): + svc = self._setup_svc() + svc.apply_voice_duration_to_plan.side_effect = RuntimeError("boom") + result = svc.ensure_variant_plans( + "src", 1, ["a1"], voice_durations=[30.0] + ) + assert result == ["plan_v0"] diff --git a/tests/unit/test_1749_preview_batch_voice_allocation.py b/tests/unit/test_1749_preview_batch_voice_allocation.py new file mode 100644 index 000000000..d94fa4029 --- /dev/null +++ b/tests/unit/test_1749_preview_batch_voice_allocation.py @@ -0,0 +1,145 @@ +"""#1749 generation_preview 配音分配分支测试。 + +直接测试 create_preview_generation_task 中配音分配逻辑: +- count=1: apply_voice_duration_to_plan 被调用 +- count>1: ensure_variant_plans / clone + reselect 路径 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch, call + +import pytest + + +# ── 辅助 ────────────────────────────────────────────────────────────────────── + +def _make_user(): + mu = MagicMock() + mu.user.id = "u1" + return mu + + +def _make_task(source_plan_id="src_plan_1"): + t = MagicMock() + t.id = "task_1" + t.source_edit_plan_id = source_plan_id + t.extra_meta = {} + return t + + +# ═══════════════════════════════════════════════════════════════════════════════ +# 直接测 apply_voice_duration_to_plan 的调用入口 +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestPreviewCount1VoiceAllocation: + """count=1 预览: clone_plan_for_variant + apply_voice_duration_to_plan。""" + + def test_clone_then_apply_voice(self): + """验证 count=1 且有配音时, clone + apply 被调用。""" + from app.services.edit_plan_service import EditPlanService + + db = MagicMock() + svc = EditPlanService.__new__(EditPlanService) + svc._clip_repo = MagicMock() + svc._clip_repo.session = db + svc._plan_repo = MagicMock() + svc._generation_task_repo = MagicMock() + + variant_plan = MagicMock() + variant_plan.id = "v_plan_0" + svc.clone_plan_for_variant = MagicMock(return_value=variant_plan) + svc.apply_voice_duration_to_plan = MagicMock(return_value=variant_plan) + + # 模拟 count=1 预览配音分配逻辑 + source_plan_id = "src_plan_1" + voice_durations = [30.0] + + if source_plan_id: + vp = svc.clone_plan_for_variant( + source_plan_id, + created_by_user_id="u1", + name_suffix="预览变体", + ) + if voice_durations and voice_durations[0] > 0: + svc.apply_voice_duration_to_plan(vp.id, voice_durations[0]) + + svc.clone_plan_for_variant.assert_called_once_with( + source_plan_id, + created_by_user_id="u1", + name_suffix="预览变体", + ) + svc.apply_voice_duration_to_plan.assert_called_once_with("v_plan_0", 30.0) + + def test_clone_without_voice_when_duration_zero(self): + """voice_durations[0]=0 时不调用 apply。""" + from app.services.edit_plan_service import EditPlanService + + svc = EditPlanService.__new__(EditPlanService) + svc._clip_repo = MagicMock() + svc._plan_repo = MagicMock() + svc._generation_task_repo = MagicMock() + + vp = MagicMock(id="v_plan_0") + svc.clone_plan_for_variant = MagicMock(return_value=vp) + svc.apply_voice_duration_to_plan = MagicMock() + + voice_durations = [0.0] + variant_plan = svc.clone_plan_for_variant("src", created_by_user_id="u1", name_suffix="预览变体") + if voice_durations and voice_durations[0] > 0: + svc.apply_voice_duration_to_plan(variant_plan.id, voice_durations[0]) + + svc.apply_voice_duration_to_plan.assert_not_called() + + +class TestPreviewCountGt1VoiceAllocation: + """count>1 预览: ensure_variant_plans 路径。""" + + def test_ensure_variant_plans_called(self): + """验证 count>1 时 ensure_variant_plans 被正确调用。""" + from app.services.edit_plan_service import EditPlanService + + svc = EditPlanService.__new__(EditPlanService) + svc._clip_repo = MagicMock() + svc._plan_repo = MagicMock() + svc._generation_task_repo = MagicMock() + + svc.ensure_variant_plans = MagicMock(return_value=["p0", "p1", "p2"]) + + result = svc.ensure_variant_plans( + "src_plan", + 3, + ["a1", "a2"], + created_by_user_id="u1", + voice_durations=[30.0, 20.0, 15.0], + ) + + assert result == ["p0", "p1", "p2"] + svc.ensure_variant_plans.assert_called_once_with( + "src_plan", + 3, + ["a1", "a2"], + created_by_user_id="u1", + voice_durations=[30.0, 20.0, 15.0], + ) + + def test_ensure_variant_plans_no_voice(self): + """无配音时 voice_durations 全 0 仍可调用。""" + from app.services.edit_plan_service import EditPlanService + + svc = EditPlanService.__new__(EditPlanService) + svc._clip_repo = MagicMock() + svc._plan_repo = MagicMock() + svc._generation_task_repo = MagicMock() + + svc.ensure_variant_plans = MagicMock(return_value=["p0"]) + + result = svc.ensure_variant_plans( + "src_plan", + 1, + ["a1"], + created_by_user_id="u1", + voice_durations=None, + ) + assert result == ["p0"] diff --git a/tests/unit/test_1749_variant_plans_route.py b/tests/unit/test_1749_variant_plans_route.py new file mode 100644 index 000000000..3e4068652 --- /dev/null +++ b/tests/unit/test_1749_variant_plans_route.py @@ -0,0 +1,138 @@ +"""#1749 POST /api/v1/generation/variant-plans 路由单元测试。 + +直接调用路由函数 + mock 依赖(不使用 TestClient,因 httpx2 冲突)。 + +覆盖: +- 200: 合法请求返回 items +- 400: voice_library_ids 长度不符 (schema validator) +- 400: 缺少 source_edit_plan_id 且无模板兜底 plan +- 400: 选片失败 ValueError -> HTTPException(400) +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from fastapi import HTTPException + + +def _call_route(request_obj, db=None): + """直接调用 create_variant_plans 路由函数。""" + from app.api.routes.generation_variant_plans import create_variant_plans + + user = MagicMock() + user.user.id = "u1" + if db is None: + db = MagicMock() + return create_variant_plans(request_obj, authenticated_user=user, db=db) + + +def _base_payload(**kw): + from app.api.routes.generation_variant_plans import VariantPlanRequest + d = { + "template_id": "tpl_1", + "asset_ids": ["a1", "a2"], + "count": 1, + "source_edit_plan_id": "src_plan_1", + } + d.update(kw) + return VariantPlanRequest(**d) + + +class TestVariantPlansRoute: + def test_200_happy_path(self): + """合法请求:ensure_variant_plans 返回 plan_id,list_clips 返回 clips。""" + req = _base_payload() + db = MagicMock() + + mock_svc_inst = MagicMock() + mock_svc_inst.ensure_variant_plans.return_value = ["plan_v0"] + clip = MagicMock() + clip.id = "c1" + clip.order = 0 + clip.asset_id = "a1" + clip.start_time = 0.0 + clip.duration = 10.0 + clip.clip_type = "main" + clip.transition_effect = "cut" + clip.transition_duration = 0.0 + clip.playback_speed = 1.0 + clip.text_content = "" + mock_svc_inst.list_clips.return_value = [clip] + + mock_svc_cls = MagicMock(return_value=mock_svc_inst) + + with patch( + "packages.domain.variant_voice_resolver.resolve_variant_voice_ids", + return_value=["voice_1"], + ), patch( + "app.api.routes.generation_tasks._query_voice_durations", + return_value=[30.0], + ), patch( + "app.services.edit_plan_service.EditPlanService", + mock_svc_cls, + ): + resp = _call_route(req, db) + + assert resp.total == 1 + assert resp.items[0].plan_id == "plan_v0" + assert len(resp.items[0].clips) == 1 + + def test_400_voice_library_ids_length_mismatch(self): + """voice_library_ids 长度 != count -> pydantic model_validator 抛异常。""" + from app.api.routes.generation_variant_plans import VariantPlanRequest + with pytest.raises(Exception): + VariantPlanRequest( + template_id="tpl_1", + asset_ids=["a1"], + count=3, + source_edit_plan_id="src_plan_1", + voice_library_ids=["v1", "v2"], # len=2 != count=3 + ) + + def test_400_no_source_plan_no_template_fallback(self): + """无 source_edit_plan_id 且 DB 查不到模板 plan -> HTTPException(400)。""" + from app.api.routes.generation_variant_plans import VariantPlanRequest + + req = VariantPlanRequest( + template_id="tpl_nonexist", + asset_ids=[], + count=1, + source_edit_plan_id="", + ) + + db = MagicMock() + db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None + + with patch( + "packages.domain.variant_voice_resolver.resolve_variant_voice_ids", + return_value=[""], + ): + with pytest.raises(HTTPException) as exc_info: + _call_route(req, db) + assert exc_info.value.status_code == 400 + + def test_400_selection_failure_value_error(self): + """ensure_variant_plans 抛 ValueError -> HTTPException(400)。""" + req = _base_payload() + db = MagicMock() + + mock_svc_inst = MagicMock() + mock_svc_inst.ensure_variant_plans.side_effect = ValueError("素材池为空") + mock_svc_cls = MagicMock(return_value=mock_svc_inst) + + with patch( + "packages.domain.variant_voice_resolver.resolve_variant_voice_ids", + return_value=["voice_1"], + ), patch( + "app.api.routes.generation_tasks._query_voice_durations", + return_value=[30.0], + ), patch( + "app.services.edit_plan_service.EditPlanService", + mock_svc_cls, + ): + with pytest.raises(HTTPException) as exc_info: + _call_route(req, db) + assert exc_info.value.status_code == 400