"""积分系统 Repository 层单元测试 (#1895)""" from __future__ import annotations import uuid from datetime import UTC, date, datetime, timezone import pytest from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from packages.adapters.sqlalchemy_impl.daily_usage_repository import SQLAlchemyDailyUsageRepository from packages.adapters.sqlalchemy_impl.models import ( Base, DailyUsageRecordModel, PointsAccountModel, PointsOrderModel, PointsTransactionModel, UserModel, ) from packages.adapters.sqlalchemy_impl.points_account_repository import SQLAlchemyPointsAccountRepository from packages.adapters.sqlalchemy_impl.points_order_repository import SQLAlchemyPointsOrderRepository from packages.adapters.sqlalchemy_impl.points_transaction_repository import SQLAlchemyPointsTransactionRepository from packages.domain.daily_usage_record import DailyUsageRecord from packages.domain.points_account import PointsAccount from packages.domain.points_order import PointsOrder from packages.domain.points_transaction import PointsTransaction @pytest.fixture() def db_session(): engine = create_engine("sqlite://", echo=False) Base.metadata.create_all(engine) SessionLocal = sessionmaker(bind=engine) session = SessionLocal() # 创建测试用户 user = UserModel( id="test-user-1", email="test@example.com", username="testuser", display_name="Test User", password_hash="xxx", ) session.add(user) session.commit() yield session session.close() @pytest.fixture() def account_repo(db_session): return SQLAlchemyPointsAccountRepository(db_session) @pytest.fixture() def txn_repo(db_session): return SQLAlchemyPointsTransactionRepository(db_session) @pytest.fixture() def order_repo(db_session): return SQLAlchemyPointsOrderRepository(db_session) @pytest.fixture() def daily_repo(db_session): return SQLAlchemyDailyUsageRepository(db_session) # ── PointsAccountRepository ── class TestPointsAccountRepository: def test_create_and_get(self, account_repo): account = PointsAccount.create(user_id="test-user-1") result = account_repo.create(account) assert result.user_id == "test-user-1" fetched = account_repo.get_by_user_id("test-user-1") assert fetched is not None assert fetched.id == account.id assert fetched.balance == 0 def test_get_nonexistent(self, account_repo): result = account_repo.get_by_user_id("nonexistent") assert result is None def test_update_balance(self, account_repo): account = PointsAccount.create(user_id="test-user-1") account_repo.create(account) account.balance = 100 account.total_earned = 150 account.total_spent = 50 updated = account_repo.update_balance(account) assert updated.balance == 100 fetched = account_repo.get_by_user_id("test-user-1") assert fetched.balance == 100 assert fetched.total_earned == 150 assert fetched.total_spent == 50 # ── PointsTransactionRepository ── class TestPointsTransactionRepository: def test_create_and_list(self, txn_repo): txn = PointsTransaction.create( user_id="test-user-1", account_id="acc-1", type="add", source="recharge", amount=100, balance_after=100, description="充值", ) txn_repo.create(txn) items, total = txn_repo.list_by_user("test-user-1") assert total == 1 assert items[0].amount == 100 assert items[0].type == "add" def test_list_with_type_filter(self, txn_repo): for t in ["add", "deduct", "add"]: txn = PointsTransaction.create( user_id="test-user-1", account_id="acc-1", type=t, source="test", amount=10, balance_after=10, ) txn_repo.create(txn) items, total = txn_repo.list_by_user("test-user-1", type="add") assert total == 2 def test_list_with_source_filter(self, txn_repo): for s in ["recharge", "ai_voice", "recharge"]: txn = PointsTransaction.create( user_id="test-user-1", account_id="acc-1", type="add", source=s, amount=10, balance_after=10, ) txn_repo.create(txn) items, total = txn_repo.list_by_user("test-user-1", source="recharge") assert total == 2 def test_list_pagination(self, txn_repo): for _i in range(5): txn = PointsTransaction.create( user_id="test-user-1", account_id="acc-1", type="add", source="test", amount=10, balance_after=10, ) txn_repo.create(txn) items, total = txn_repo.list_by_user("test-user-1", page=1, page_size=3) assert total == 5 assert len(items) == 3 items2, _ = txn_repo.list_by_user("test-user-1", page=2, page_size=3) assert len(items2) == 2 # ── PointsOrderRepository ── class TestPointsOrderRepository: def test_create_and_get(self, order_repo): order = PointsOrder.create( user_id="test-user-1", order_type="points", product_code="starter_pack", amount_cents=990, points_amount=100, ) order_repo.create(order) fetched = order_repo.get(order.id) assert fetched is not None assert fetched.product_code == "starter_pack" assert fetched.amount_cents == 990 def test_get_nonexistent(self, order_repo): assert order_repo.get("nonexistent") is None def test_update_status(self, order_repo): order = PointsOrder.create( user_id="test-user-1", order_type="points", product_code="starter_pack", amount_cents=990, ) order_repo.create(order) now = datetime.now(UTC) updated = order_repo.update_status(order.id, "paid", payment_id="pay-123", paid_at=now) assert updated is not None assert updated.status == "paid" assert updated.payment_id == "pay-123" def test_update_status_nonexistent(self, order_repo): result = order_repo.update_status("nonexistent", "paid") assert result is None def test_list_by_user(self, order_repo): for ot in ["points", "membership", "points"]: order = PointsOrder.create( user_id="test-user-1", order_type=ot, product_code="test", amount_cents=100, ) order_repo.create(order) items, total = order_repo.list_by_user("test-user-1", order_type="points") assert total == 2 def test_list_by_user_with_status(self, order_repo): order = PointsOrder.create( user_id="test-user-1", order_type="points", product_code="test", amount_cents=100, ) order_repo.create(order) items, total = order_repo.list_by_user("test-user-1", status="pending") assert total == 1 items2, total2 = order_repo.list_by_user("test-user-1", status="paid") assert total2 == 0 # ── DailyUsageRepository ── class TestDailyUsageRepository: def _today(self): # The model column is DateTime, so use datetime for comparison from datetime import datetime now = datetime.now(UTC) return now.replace(hour=0, minute=0, second=0, microsecond=0) def test_create_and_get(self, daily_repo): today = self._today() record = DailyUsageRecord.create(user_id="test-user-1", usage_date=today) record.count = 1 daily_repo.create(record) fetched = daily_repo.get_by_user_and_date("test-user-1", today) assert fetched is not None assert fetched.count == 1 def test_get_nonexistent(self, daily_repo): from datetime import timedelta tomorrow = self._today() + timedelta(days=1) result = daily_repo.get_by_user_and_date("test-user-1", tomorrow) assert result is None def test_update_count(self, daily_repo): today = self._today() record = DailyUsageRecord.create(user_id="test-user-1", usage_date=today) record.count = 1 daily_repo.create(record) record.count = 3 daily_repo.update_count(record) fetched = daily_repo.get_by_user_and_date("test-user-1", today) assert fetched.count == 3 def test_upsert_create(self, daily_repo): today = self._today() result = daily_repo.upsert("test-user-1", today, "free_clip") assert result.count == 1 def test_upsert_increment(self, daily_repo): today = self._today() daily_repo.upsert("test-user-1", today, "free_clip") result = daily_repo.upsert("test-user-1", today, "free_clip") assert result.count == 2