"""PointsService 单元测试 (#1895) — 使用 SQLite 内存数据库""" from __future__ import annotations import uuid from datetime import datetime, timezone from unittest.mock import patch import pytest from sqlalchemy import create_engine, event from sqlalchemy.orm import Session, sessionmaker from packages.domain.points_service import PointsService @pytest.fixture() def db_session(): """创建 SQLite 内存数据库 session,包含所有积分相关表。""" from packages.adapters.sqlalchemy_impl.models import Base engine = create_engine("sqlite://", echo=False) # SQLite 不支持 WITH FOR UPDATE,mock 掉 @event.listens_for(engine, "connect") def _disable_for_update(dbapi_conn, connection_record): pass Base.metadata.create_all(engine) SessionLocal = sessionmaker(bind=engine) session = SessionLocal() yield session session.close() @pytest.fixture() def service(): return PointsService() @pytest.fixture() def user_id(): return uuid.uuid4().hex class TestGetOrCreateAccount: def test_creates_new_account(self, service, db_session, user_id): data = service.get_or_create_account(user_id, db_session) assert data["user_id"] == user_id assert data["balance"] == 0 assert data["total_earned"] == 0 assert data["total_spent"] == 0 def test_returns_existing_account(self, service, db_session, user_id): service.get_or_create_account(user_id, db_session) data = service.get_or_create_account(user_id, db_session) assert data["user_id"] == user_id assert data["balance"] == 0 class TestCheckBalance: def test_sufficient_when_zero(self, service, db_session, user_id): result = service.check_balance(user_id, 0, db_session) assert result["sufficient"] is True def test_insufficient_when_new_account(self, service, db_session, user_id): result = service.check_balance(user_id, 10, db_session) assert result["sufficient"] is False assert result["remaining_after"] == -10 class TestDeductPoints: def test_deduct_fails_insufficient_balance(self, service, db_session, user_id): result = service.deduct_points(user_id, 100, "voice_clone_synth", db_session) assert result["success"] is False assert result["transaction_id"] is None def test_deduct_after_recharge(self, service, db_session, user_id): # 先充值 service.add_points(user_id, 50, "recharge", db_session) # 再扣减 result = service.deduct_points(user_id, 20, "voice_clone_synth", db_session) assert result["success"] is True assert result["balance"] == 30 def test_deduct_creates_transaction(self, service, db_session, user_id): service.add_points(user_id, 100, "recharge", db_session) result = service.deduct_points(user_id, 30, "voice_clone_synth", db_session) assert result["success"] is True txns = service.get_transactions(user_id, db_session) assert txns["total"] == 2 # 1 add + 1 deduct deduct_txn = [t for t in txns["items"] if t["type"] == "deduct"][0] assert deduct_txn["amount"] == 30 assert deduct_txn["balance_after"] == 70 class TestAddPoints: def test_add_new_account(self, service, db_session, user_id): result = service.add_points(user_id, 100, "recharge:starter_pack", db_session) assert result["success"] is True assert result["balance"] == 100 def test_add_accumulates(self, service, db_session, user_id): service.add_points(user_id, 50, "recharge", db_session) result = service.add_points(user_id, 30, "bonus", db_session) assert result["balance"] == 80 class TestRefundPoints: def test_refund_adds_back(self, service, db_session, user_id): service.add_points(user_id, 100, "recharge", db_session) service.deduct_points(user_id, 20, "voice_clone_synth", db_session) result = service.refund_points(user_id, 20, "voice_clone_synth", db_session) assert result["success"] is True assert result["balance"] == 100 def test_refund_creates_refund_transaction(self, service, db_session, user_id): service.add_points(user_id, 100, "recharge", db_session) service.refund_points(user_id, 10, "voice_clone_synth", db_session) txns = service.get_transactions(user_id, db_session) refund_txns = [t for t in txns["items"] if t["type"] == "add" and "refund" in t["source"]] assert len(refund_txns) == 1 assert "refund:" in refund_txns[0]["source"] class TestGetTransactions: def test_empty_for_new_user(self, service, db_session, user_id): result = service.get_transactions(user_id, db_session) assert result["total"] == 0 assert result["items"] == [] def test_pagination(self, service, db_session, user_id): for i in range(5): service.add_points(user_id, 10, f"batch_{i}", db_session) result = service.get_transactions(user_id, db_session, page=1, page_size=3) assert result["total"] == 5 assert len(result["items"]) == 3 result2 = service.get_transactions(user_id, db_session, page=2, page_size=3) assert len(result2["items"]) == 2 class TestGetDailyUsage: """智能混剪已免费,get_daily_usage 返回 unlimited(-1)占位。""" def test_returns_unlimited(self, service, db_session, user_id): result = service.get_daily_usage(user_id, db_session) assert result["free_clips_used"] == 0 assert result["free_clips_limit"] == -1 # -1 表示 unlimited assert result["free_clips_remaining"] == -1 assert "reset_at" in result class TestCreateOrder: def test_points_order(self, service, db_session, user_id): result = service.create_order(user_id, "points", "starter_pack", db_session) assert result["order_type"] == "points" assert result["product_code"] == "starter_pack" assert result["amount_cents"] == 990 assert result["status"] == "pending" def test_membership_order(self, service, db_session, user_id): result = service.create_order(user_id, "membership", "monthly", db_session) assert result["order_type"] == "membership" assert result["amount_cents"] == 1990 def test_unknown_package_raises(self, service, db_session, user_id): with pytest.raises(ValueError, match="Unknown points package"): service.create_order(user_id, "points", "nonexistent", db_session) def test_unknown_membership_raises(self, service, db_session, user_id): with pytest.raises(ValueError, match="Unknown membership type"): service.create_order(user_id, "membership", "lifetime", db_session) 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 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,不抛异常。""" 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。""" 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。""" 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 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。""" 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。""" 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。""" 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。""" 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