diff --git a/tests/unit/test_points_rules.py b/tests/unit/test_points_rules.py index 7c1392dc6..6dc7682f3 100644 --- a/tests/unit/test_points_rules.py +++ b/tests/unit/test_points_rules.py @@ -158,10 +158,18 @@ class TestResolveVideoDimensions: 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), + ("普清", 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") @@ -307,7 +315,9 @@ class TestCalculateViralVideoCredits: 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) + 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 @@ -373,7 +383,6 @@ class TestCalculateViralVideoCredits: 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 diff --git a/tests/unit/test_points_service.py b/tests/unit/test_points_service.py index f100156a8..4bfe239c4 100644 --- a/tests/unit/test_points_service.py +++ b/tests/unit/test_points_service.py @@ -189,7 +189,7 @@ class TestDeductViralVideo: 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 + from unittest.mock import MagicMock expected = {"success": True, "balance": 50.0, "transaction_id": "t1"} with patch.object(service, "deduct_points", return_value=expected) as mock_dp: @@ -207,9 +207,10 @@ class TestDeductViralVideo: 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: + 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 @@ -219,7 +220,6 @@ class TestSettleViralVideo: 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, @@ -235,7 +235,6 @@ class TestSettleViralVideo: 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: @@ -284,7 +281,6 @@ class TestSettleViralVideo: 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 ( @@ -302,7 +298,6 @@ class TestSettleViralVideo: 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) @@ -317,7 +312,6 @@ class TestRefundViralVideo: 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) @@ -331,7 +325,6 @@ class TestRefundViralVideo: 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: @@ -348,7 +341,6 @@ class TestRefundViralVideo: 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) diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index 5a0d19231..d0d9a9629 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -412,10 +412,9 @@ class TestConfirmCopyPointsDeduction: 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 fastapi import HTTPException from packages.domain.viral_video import ViralVideoStatus @@ -458,10 +457,9 @@ class TestConfirmCopyPointsDeduction: 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 fastapi import HTTPException from packages.domain.viral_video import ViralVideoStatus