From 9ee55dd695dd81581dcb6453ab2a4f2df6f73420 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Fri, 2 Oct 2026 23:54:24 +0800 Subject: [PATCH] =?UTF-8?q?test(viral-video):=20=E8=A1=A5=E5=8D=95?= =?UTF-8?q?=E6=B5=8B=E8=A6=86=E7=9B=96=E5=8A=A8=E6=80=81=E7=A7=AF=E5=88=86?= =?UTF-8?q?=E5=AE=9A=E4=BB=B7=E5=85=A8=E9=83=A8=E5=88=86=E6=94=AF=EF=BC=88?= =?UTF-8?q?diff=E8=A6=86=E7=9B=96=E7=8E=8746%=E2=86=9278%=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - points_rules: resolve_video_dimensions/_match_model_prefix/_infer_resolution_key /calculate_viral_video_credits 覆盖别名/has_video_input/actual_tokens/兜底/防御分支(+28 tests,覆盖94%) - points_service: deduct/settle/refund_viral_video 覆盖none/refund/deduct/余额不足/异常分支(+11 tests,覆盖71%) - viral_video routes: confirm-copy 预扣/402/already_paid/estimate-credits端点(+8 tests,覆盖53%,全部缺失行已覆盖) --- tests/unit/test_points_rules.py | 266 ++++++++++++++++++++++++++ tests/unit/test_points_service.py | 177 +++++++++++++++++ tests/unit/test_viral_video_routes.py | 247 ++++++++++++++++++++++++ 3 files changed, 690 insertions(+) diff --git a/tests/unit/test_points_rules.py b/tests/unit/test_points_rules.py index 17294bfc1..7c1392dc6 100644 --- a/tests/unit/test_points_rules.py +++ b/tests/unit/test_points_rules.py @@ -117,3 +117,269 @@ class TestCalculatePointsCost: def test_retired_scenes_return_zero(self, scene): assert calculate_points_cost(scene, is_member=False) == 0 assert calculate_points_cost(scene, is_member=True, duration_minutes=10) == 0 + + +# ============ 爆款视频动态定价 (#2151) ============ + + +class TestResolveVideoDimensions: + """resolve_video_dimensions(): 分辨率别名、比例、默认兜底。""" + + def test_1080p_16_9(self): + """1080p + 16:9 → w=1920, h=1080。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("1080p", "16:9") + assert (w, h) == (1920, 1080) + + def test_480p_16_9(self): + """480p + 16:9 → h=480, w 按 16//9 计算。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("480p", "16:9") + assert h == 480 + assert w == 480 * 16 // 9 + + def test_720p_1_1(self): + """1:1 正方形 → w == h。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("720p", "1:1") + assert (w, h) == (720, 720) + + def test_1080p_1_1(self): + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("1080p", "1:1") + assert (w, h) == (1080, 1080) + + def test_resolution_aliases(self): + """中文/英文别名应正确映射到对应高度。""" + from packages.domain.points_rules import resolve_video_dimensions + + cases = [ + ("普清", 480), ("sd", 480), ("low", 480), ("default", 480), + ("高清", 720), ("medium", 720), ("hd", 720), + ("超清", 1080), ("fhd", 1080), ("ultra", 1080), + ("全能", 1080), ("high", 1080), + ] + for alias, expected_h in cases: + _, h = resolve_video_dimensions(alias, "1:1") + assert h == expected_h, f"{alias} -> h={h}, expected {expected_h}" + + def test_unknown_resolution_falls_back_to_720p(self): + """未知分辨率字符串兜底到 720p。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("2160p", "1:1") + assert h == 720 + assert w == 720 + + def test_empty_resolution_defaults_to_720p_9_16(self): + """空 resolution + 空 ratio → 默认 720p + 9:16。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("", "") + assert h == 720 + assert w == 720 * 9 // 16 + + def test_none_resolution_default_ratio(self): + """None resolution + None ratio → 720p + 9:16 默认。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions(None, None) + assert h == 720 + assert w == 720 * 9 // 16 + + def test_whitespace_resolution_case_insensitive(self): + """前后空格 + 大写应被规范化处理。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions(" 1080P ", " 16:9 ") + assert (w, h) == (1920, 1080) + + +class TestMatchModelPrefix: + """_match_model_prefix() 前缀匹配 + 兜底。""" + + def test_seedance_2_5_exact(self): + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("seedance-2.5") == "seedance-2.5" + + def test_seedance_2_5_with_variant(self): + """带后缀版本号(如 seedance-2.5-pro)仍匹配 seedance-2.5。""" + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("seedance-2.5-pro") == "seedance-2.5" + + def test_seedance_2_0_exact(self): + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("seedance-2.0") == "seedance-2.0" + + def test_seedance_2_0_with_variant(self): + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("seedance-2.0-lite") == "seedance-2.0" + + def test_unknown_model_falls_back_to_2_5(self): + """未知模型前缀兜底 seedance-2.5。""" + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("kling-v1") == "seedance-2.5" + assert _match_model_prefix("") == "seedance-2.5" + assert _match_model_prefix(None) == "seedance-2.5" + + def test_case_insensitive(self): + from packages.domain.points_rules import _match_model_prefix + + assert _match_model_prefix("SEEDANCE-2.0") == "seedance-2.0" + + +class TestInferResolutionKey: + """_infer_resolution_key(): 1000+/650-999/<650 三个分支。""" + + def test_height_ge_1000_is_1080p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(1000) == "1080p" + assert _infer_resolution_key(1080) == "1080p" + assert _infer_resolution_key(2160) == "1080p" + + def test_height_650_to_999_is_720p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(650) == "720p" + assert _infer_resolution_key(720) == "720p" + assert _infer_resolution_key(999) == "720p" + + def test_height_lt_650_is_480p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(480) == "480p" + assert _infer_resolution_key(649) == "480p" + assert _infer_resolution_key(0) == "480p" + + +class TestCalculateViralVideoCredits: + """calculate_viral_video_credits():爆款视频动态定价核心函数。""" + + def test_default_args_returns_float(self): + """默认参数返回 float。""" + from packages.domain.points_rules import calculate_viral_video_credits + + credits = calculate_viral_video_credits(15, 1280, 720) + assert isinstance(credits, float) + + def test_return_is_rounded_to_two_decimals(self): + """round(..., 2) 后值本身就是两位小数(再 round 不变化)。""" + from packages.domain.points_rules import calculate_viral_video_credits + + for dur, w, h in [(15, 1280, 720), (5, 854, 480), (30, 1920, 1080), (10, 720, 720)]: + credits = calculate_viral_video_credits(dur, w, h) + assert round(credits, 2) == credits + + def test_has_video_input_uses_lower_price(self): + """has_video_input=True 时使用参考视频价格(有视频输入便宜)。""" + from packages.domain.points_rules import calculate_viral_video_credits + + no_input = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5", has_video_input=False) + with_input = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5", has_video_input=True) + assert with_input < no_input + + def test_unknown_model_falls_back_to_seedance_2_5(self): + """未知 model 前缀兜底到 seedance-2.5 价格,与默认等价。""" + from packages.domain.points_rules import calculate_viral_video_credits + + unknown = calculate_viral_video_credits(15, 1280, 720, model="unknown-model") + default = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5") + assert unknown == default + + def test_actual_tokens_overrides_calculation(self): + """传入 actual_tokens>0 时用它替代公式计算的 tokens。""" + from packages.domain.points_rules import ( + VIRAL_VIDEO_FIXED_COST, + VIRAL_VIDEO_MODEL_PRICES, + VIRAL_VIDEO_PROFIT_MULTIPLIER, + calculate_viral_video_credits, + ) + + price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)] + actual_tokens = 2_000_000 + expected = round((actual_tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2) + credits = calculate_viral_video_credits(15, 1280, 720, actual_tokens=actual_tokens) + assert credits == expected + + def test_zero_duration_width_height_defensive_max1(self): + """duration/width/height 为 0/None 时 max(1,...) 防御,结果>0。""" + from packages.domain.points_rules import calculate_viral_video_credits + + c_zero = calculate_viral_video_credits(0, 0, 0) + assert c_zero > 0 + c_none = calculate_viral_video_credits(None, None, None) + assert c_none > 0 + c_one = calculate_viral_video_credits(1, 1, 1) + assert c_none == c_one + + def test_non_default_fps_affects_tokens(self): + """fps 非默认值(30) 应比默认(24) 积分高。""" + from packages.domain.points_rules import calculate_viral_video_credits + + c24 = calculate_viral_video_credits(15, 1280, 720, fps=24) + c30 = calculate_viral_video_credits(15, 1280, 720, fps=30) + assert c30 > c24 + + def test_seedance_2_0_priced_lower_than_2_5_at_1080p(self): + """seedance-2.0 在 1080p 无视频输入时定价低于 seedance-2.5。""" + from packages.domain.points_rules import calculate_viral_video_credits + + c20 = calculate_viral_video_credits(15, 1920, 1080, model="seedance-2.0", has_video_input=False) + c25 = calculate_viral_video_credits(15, 1920, 1080, model="seedance-2.5", has_video_input=False) + assert c20 < c25 + + def test_formula_includes_fixed_cost_and_multiplier(self): + """手算公式结果应与函数返回一致(固定成本 + 利润系数)。""" + from packages.domain.points_rules import ( + VIRAL_VIDEO_FIXED_COST, + VIRAL_VIDEO_FPS, + VIRAL_VIDEO_MODEL_PRICES, + VIRAL_VIDEO_PROFIT_MULTIPLIER, + calculate_viral_video_credits, + ) + + dur, w, h = 10, 1280, 720 + price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)] + tokens = dur * w * h * VIRAL_VIDEO_FPS / 1024.0 + expected = round((tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2) + assert calculate_viral_video_credits(dur, w, h) == expected + + def test_seedance_2_0_with_video_input_falls_back_to_seedance_2_5_price(self): + """seedance-2.0 + has_video_input=True 组合不在价格表,走 line 111 fallback 到 seedance-2.5 的 720p False 价格。""" + from packages.domain.points_rules import ( + VIRAL_VIDEO_FIXED_COST, + VIRAL_VIDEO_FPS, + VIRAL_VIDEO_MODEL_PRICES, + VIRAL_VIDEO_PROFIT_MULTIPLIER, + calculate_viral_video_credits, + ) + + dur, w, h = 10, 1280, 720 + # 兜底价格 = seedance-2.5/720p/False = 70.0 + price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)] + assert price == 70.0 + tokens = dur * w * h * VIRAL_VIDEO_FPS / 1024.0 + expected = round((tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2) + credits = calculate_viral_video_credits(dur, w, h, model="seedance-2.0", has_video_input=True) + assert credits == expected + + + def test_fps_zero_or_none_falls_back_to_default(self): + """fps=0/None 时 int(fps or 24) 兜底到默认 24,结果与 fps=24 一致。""" + from packages.domain.points_rules import calculate_viral_video_credits + + c_default = calculate_viral_video_credits(10, 1280, 720, fps=24) + c_zero = calculate_viral_video_credits(10, 1280, 720, fps=0) + c_none = calculate_viral_video_credits(10, 1280, 720, fps=None) + assert c_zero == c_default + assert c_none == c_default diff --git a/tests/unit/test_points_service.py b/tests/unit/test_points_service.py index 3d744801d..f100156a8 100644 --- a/tests/unit/test_points_service.py +++ b/tests/unit/test_points_service.py @@ -179,3 +179,180 @@ class TestCreateOrder: def test_unknown_order_type_raises(self, service, db_session, user_id): with pytest.raises(ValueError, match="Unknown order type"): service.create_order(user_id, "insurance", "basic", db_session) + + +# ============ 爆款视频(viral_video)动态定价方法 ============ + + +class TestDeductViralVideo: + """deduct_viral_video(): 预扣积分,委托给 deduct_points。""" + + def test_delegates_to_deduct_points_with_correct_args(self, service, db_session, user_id): + """deduct_viral_video 应以 source='viral_video', ref_id=job_id 调用 deduct_points。""" + from unittest.mock import MagicMock, patch + + expected = {"success": True, "balance": 50.0, "transaction_id": "t1"} + with patch.object(service, "deduct_points", return_value=expected) as mock_dp: + result = service.deduct_viral_video(user_id, 10.5, "job-abc", db_session) + + assert result == expected + mock_dp.assert_called_once() + kwargs = mock_dp.call_args.kwargs + assert kwargs["user_id"] == user_id + assert kwargs["amount"] == 10.5 + assert kwargs["source"] == "viral_video" + assert kwargs["db"] is db_session + assert kwargs["description"] == "爆款视频生成" + assert kwargs["ref_id"] == "job-abc" + + def test_none_credits_coerced_to_zero(self, service, db_session, user_id): + """credits=None 时应被 float(credits or 0) 转为 0,不抛异常。""" + from unittest.mock import patch + + with patch.object(service, "deduct_points", return_value={"success": True, "balance": 0, "transaction_id": "t"}) as mock_dp: + service.deduct_viral_video(user_id, None, "job-nil", db_session) + assert mock_dp.call_args.kwargs["amount"] == 0.0 + + +class TestSettleViralVideo: + """settle_viral_video(): 多退少补结算。""" + + def test_no_action_when_diff_below_epsilon(self, service, db_session, user_id): + """|diff|<0.01 时返回 action=none,不调 refund/deduct。""" + from unittest.mock import patch + + with ( + patch.object(service, "refund_points") as mock_refund, + patch.object(service, "deduct_points") as mock_deduct, + ): + result = service.settle_viral_video(user_id, estimated=10.00, actual=10.001, txn_id="t1", db=db_session) + + assert result["success"] is True + assert result["action"] == "none" + assert result["diff"] == 0.0 + mock_refund.assert_not_called() + mock_deduct.assert_not_called() + + def test_refund_when_actual_less_than_estimated(self, service, db_session, user_id): + """actualestimated 且补扣成功 → action=deduct, success=True。""" + from unittest.mock import patch + + deduct_res = {"success": True, "balance": 40.0, "transaction_id": "td-1"} + with patch.object(service, "deduct_points", return_value=deduct_res) as mock_deduct: + result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t4", db=db_session) + + assert result["success"] is True + assert result["action"] == "deduct" + assert result["amount"] == 5.0 + assert result["diff"] == 5.0 + mock_deduct.assert_called_once() + dk = mock_deduct.call_args.kwargs + assert dk["amount"] == 5.0 + assert dk["source"] == "viral_video" + assert dk["ref_id"] == "t4" + + def test_deduct_insufficient_balance_returns_success_false_not_raise(self, service, db_session, user_id): + """actual>estimated 补扣时余额不足(success=False)应记录 warning 但不抛异常。""" + import logging + from unittest.mock import patch + + deduct_res = {"success": False, "balance": 2.0, "transaction_id": None} + with ( + patch.object(service, "deduct_points", return_value=deduct_res), + patch("packages.domain.points_service.logger") as mock_logger, + ): + result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t5", db=db_session) + + # 即使补扣失败,函数也返回 action=deduct 但 success=False(不阻塞任务完成) + assert result["success"] is False + assert result["action"] == "deduct" + assert result["amount"] == 5.0 + # 应打印 warning + mock_logger.warning.assert_called_once() + + def test_deduct_exception_returns_failure(self, service, db_session, user_id): + """deduct_points 抛异常时应捕获并返回 success=False。""" + from unittest.mock import patch + + with patch.object(service, "deduct_points", side_effect=RuntimeError("db boom")): + result = service.settle_viral_video(user_id, estimated=10.0, actual=20.0, txn_id="t6", db=db_session) + + assert result["success"] is False + assert result["action"] == "deduct" + assert result["diff"] == 10.0 + + +class TestRefundViralVideo: + """refund_viral_video(): 爆款视频失败全额退款。""" + + def test_zero_amount_returns_none_action(self, service, db_session, user_id): + """amount<=0 直接返回 none action,不调 refund_points。""" + from unittest.mock import patch + + with patch.object(service, "refund_points") as mock_refund: + r1 = service.refund_viral_video(user_id, 0, "t0", db_session) + r2 = service.refund_viral_video(user_id, None, "t0", db_session) + r3 = service.refund_viral_video(user_id, -1.5, "t0", db_session) + + assert r1 == {"success": True, "action": "none", "amount": 0.0} + assert r2 == {"success": True, "action": "none", "amount": 0.0} + assert r3["action"] == "none" + mock_refund.assert_not_called() + + def test_success_path_delegates_to_refund_points(self, service, db_session, user_id): + """成功路径:透传 user_id/amount/ref_id=txn_id/source=viral_video。""" + from unittest.mock import patch + + expected = {"success": True, "balance": 80.0, "transaction_id": "rf-1"} + with patch.object(service, "refund_points", return_value=expected) as mock_refund: + result = service.refund_viral_video(user_id, 30.0, "txn-xyz", db_session) + + assert result == expected + mock_refund.assert_called_once() + rk = mock_refund.call_args.kwargs + assert rk["user_id"] == user_id + assert rk["amount"] == 30.0 + assert rk["source"] == "viral_video" + assert rk["ref_id"] == "txn-xyz" + assert rk["description"] == "爆款视频失败退款" + + def test_exception_returns_failure(self, service, db_session, user_id): + """refund_points 抛异常时返回 success=False/action=refund。""" + from unittest.mock import patch + + with patch.object(service, "refund_points", side_effect=RuntimeError("conn lost")): + result = service.refund_viral_video(user_id, 25.0, "txn-err", db_session) + + assert result["success"] is False + assert result["action"] == "refund" + assert result["amount"] == 25.0 diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index 1a4dd2499..5a0d19231 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -402,3 +402,250 @@ class TestConfirmCopy: with pytest.raises(HTTPException) as exc: vv_mod.confirm_copy("job-cc2", ConfirmCopyRequest(), authenticated_user=user, session=session) assert exc.value.status_code == 409 + + +# ── confirm-copy 积分预扣 + estimate-credits 端点 (#2151) ────────────── + + +class TestConfirmCopyPointsDeduction: + """confirm_copy 中积分预扣分支(points_enabled=True)。""" + + def test_points_enabled_deducts_successfully(self): + """points_enabled=True + 未预付 → 计算预估积分 → deduct_viral_video → 写入 credits_prepaid。""" + from fastapi import HTTPException + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-pay", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + # 默认 credits_prepaid=0, credits_cost=0 → 触发预扣 + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest(edited_copy="改好的文案") + + mock_svc = MagicMock() + mock_svc.deduct_viral_video.return_value = {"success": True, "balance": 100.0, "transaction_id": "txn-1"} + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch.object(vv_mod, "_settings", None, create=True), # ensure not cached + patch("app.config.settings") as mock_settings, + patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=5.2), + patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)), + patch("packages.domain.points_service.PointsService", return_value=mock_svc), + patch.object(vv_mod.celery_app, "send_task"), + ): + mock_settings.points_enabled = True + resp = vv_mod.confirm_copy("job-pay", req, authenticated_user=user, session=session) + + # deduct_viral_video 被调用 + mock_svc.deduct_viral_video.assert_called_once() + call_args = mock_svc.deduct_viral_video.call_args + assert call_args.args[0] == "u1" # user_id + assert call_args.args[1] == 5.2 # credits + assert call_args.args[2] == "job-pay" # job_id + # credits_prepaid / credits_transaction_id 被写入 + assert job.credits_prepaid == 5.2 + assert job.credits_transaction_id == "txn-1" + assert resp.id == "job-pay" + # resume + repo.update 至少调用过(其中一次是 credits 字段更新,一次是 resume 后) + job.resume_from_copy_generated.assert_called_once_with(edited_copy="改好的文案") + + def test_points_enabled_insufficient_balance_raises_402(self): + """余额不足(deduct_viral_video 返回 success=False)→ HTTP 402。""" + import pytest + from fastapi import HTTPException + + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-402", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest() + + mock_svc = MagicMock() + mock_svc.deduct_viral_video.return_value = {"success": False, "balance": 1.5, "transaction_id": None} + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch("app.config.settings") as mock_settings, + patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=10.0), + patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)), + patch("packages.domain.points_service.PointsService", return_value=mock_svc), + patch.object(vv_mod.celery_app, "send_task"), + ): + mock_settings.points_enabled = True + with pytest.raises(HTTPException) as exc: + vv_mod.confirm_copy("job-402", req, authenticated_user=user, session=session) + + assert exc.value.status_code == 402 + detail = exc.value.detail + assert detail["code"] == "INSUFFICIENT_POINTS" + assert detail["required"] == 10.0 + assert detail["balance"] == 1.5 + # 预扣失败不应调用 resume 或 send_task + job.resume_from_copy_generated.assert_not_called() + + def test_already_paid_skips_deduction(self): + """credits_prepaid>0(已经扣过费/重试场景) → 跳过预扣,不调用 PointsService。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job( + job_id="job-paid", + user_id="u1", + status=ViralVideoStatus.COPY_GENERATED, + credits_prepaid=8.5, + credits_transaction_id="txn-old", + ) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest(edited_copy="继续") + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch("app.config.settings") as mock_settings, + patch("packages.domain.points_service.PointsService") as MockSvc, + patch.object(vv_mod.celery_app, "send_task") as mock_send, + ): + mock_settings.points_enabled = True + resp = vv_mod.confirm_copy("job-paid", req, authenticated_user=user, session=session) + + # PointsService 不应被实例化(没预扣) + MockSvc.assert_not_called() + job.resume_from_copy_generated.assert_called_once_with(edited_copy="继续") + mock_send.assert_called_once_with("worker.run_viral_video_render", args=["job-paid"]) + assert resp.id == "job-paid" + # credits_prepaid 保持不变 + assert job.credits_prepaid == 8.5 + + def test_already_paid_via_credits_cost_skips_deduction(self): + """credits_cost>0 也算已付费(兼容旧字段),跳过预扣。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job( + job_id="job-paid2", + user_id="u1", + status=ViralVideoStatus.COPY_GENERATED, + credits_prepaid=0, + credits_cost=7.0, + ) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest() + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch("app.config.settings") as mock_settings, + patch("packages.domain.points_service.PointsService") as MockSvc, + patch.object(vv_mod.celery_app, "send_task"), + ): + mock_settings.points_enabled = True + vv_mod.confirm_copy("job-paid2", req, authenticated_user=user, session=session) + + MockSvc.assert_not_called() + + def test_points_disabled_skips_deduction(self): + """points_enabled=False 时不进入预扣逻辑,保持原流程。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import ConfirmCopyRequest + + from packages.domain.viral_video import ViralVideoStatus + + user = _auth_user("u1") + session = MagicMock() + job = _make_job(job_id="job-free", user_id="u1", status=ViralVideoStatus.COPY_GENERATED) + repo = MagicMock() + repo.get.return_value = job + req = ConfirmCopyRequest() + + with ( + patch.object(vv_mod, "_get_job_repo", return_value=repo), + patch("app.config.settings") as mock_settings, + patch("packages.domain.points_service.PointsService") as MockSvc, + patch.object(vv_mod.celery_app, "send_task"), + ): + mock_settings.points_enabled = False + vv_mod.confirm_copy("job-free", req, authenticated_user=user, session=session) + + MockSvc.assert_not_called() + job.resume_from_copy_generated.assert_called_once() + + +class TestEstimateCredits: + """POST /estimate-credits: 纯计算预估积分。""" + + def test_estimate_returns_float(self): + """正常参数应返回 estimated_credits 为 float 且>0。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import EstimateCreditsRequest + + req = EstimateCreditsRequest(model="seedance-2.5", resolution="720p", ratio="9:16", duration=15) + # 不需要 db / user 之外的依赖;authenticated_user 仍要传 + user = _auth_user("u1") + + resp = vv_mod.estimate_credits(req, authenticated_user=user) + assert isinstance(resp.estimated_credits, float) + assert resp.estimated_credits > 0 + # 应保留两位小数 + assert round(resp.estimated_credits, 2) == resp.estimated_credits + + def test_estimate_uses_dimensions_resolver(self): + """estimate_credits 应调用 resolve_video_dimensions 和 calculate_viral_video_credits。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import EstimateCreditsRequest + + req = EstimateCreditsRequest(model="seedance-2.5", resolution="1080p", ratio="16:9", duration=20) + user = _auth_user("u1") + + with ( + patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)) as mock_res, + patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=8.88) as mock_calc, + ): + resp = vv_mod.estimate_credits(req, authenticated_user=user) + + mock_res.assert_called_once_with("1080p", "16:9") + mock_calc.assert_called_once() + # 传给 calculate 的参数应包含 duration=20, w=1920, h=1080, model="seedance-2.5" + args, kwargs = mock_calc.call_args + assert args[0] == 20 + assert args[1] == 1920 + assert args[2] == 1080 + assert args[3] == "seedance-2.5" + assert resp.estimated_credits == 8.88 + + def test_estimate_empty_model_defaults_to_seedance_2_5(self): + """model 为空字符串时,传入 calculate 的 model 参数应为 'seedance-2.5'。""" + from app.api.routes import viral_video as vv_mod + from app.schemas.viral_video import EstimateCreditsRequest + + req = EstimateCreditsRequest(model="", resolution="720p", ratio="9:16", duration=10) + user = _auth_user("u1") + + with ( + patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)), + patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=3.5) as mock_calc, + ): + resp = vv_mod.estimate_credits(req, authenticated_user=user) + + args, kwargs = mock_calc.call_args + assert args[3] == "seedance-2.5" + assert resp.estimated_credits == 3.5