from __future__ import annotations from collections import Counter from datetime import datetime from decimal import Decimal from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.enums.credit_balance import CreditAllocationAction, CreditBalanceStatus, CreditScope from app.enums.credit_record import CreditRecordBillingScene, CreditRecordType from app.models.credit.allocation import CreditRecordAllocation from app.models.credit.balance import UserCreditBalance from app.models.credit.subscription import UserCreditSubscription from app.models.credit_record import CreditRecord from app.services.credit.locking import acquire_user_credit_lock from app.services.credit.query_service import get_available_credits from app.services.credit.utils import to_credit_decimal, utc_now from app.utils.id_gen import generate_id async def archive_expired_user_balances( db: AsyncSession, *, user_id: str, request_time: datetime | None = None, limit: int = 500, ) -> int: checked_at = request_time or utc_now() await acquire_user_credit_lock(db, user_id) result = await db.execute( select(UserCreditBalance) .where( UserCreditBalance.user_id == user_id, UserCreditBalance.expires_at <= checked_at, UserCreditBalance.unspent_amount > 0, UserCreditBalance.expired_processed_at.is_(None), UserCreditBalance.revoked_at.is_(None), ) .order_by(UserCreditBalance.expires_at.asc(), UserCreditBalance.id.asc()) .limit(max(1, limit)) .with_for_update() ) balances = list(result.scalars().all()) if not balances: return 0 current_available = await get_available_credits(db, user_id, request_time=checked_at) subscription_ids = [item.subscription_id for item in balances if item.subscription_id] manager_map: dict[str, str | None] = {} if subscription_ids: sub_result = await db.execute( select(UserCreditSubscription.id, UserCreditSubscription.team_manager_id_snapshot) .where(UserCreditSubscription.id.in_(subscription_ids)) ) manager_map = {str(row.id): row.team_manager_id_snapshot for row in sub_result.all()} for balance in balances: amount = to_credit_decimal(balance.unspent_amount) if amount <= 0: continue before_consumed = to_credit_decimal(balance.consumed_amount) record = CreditRecord( id=generate_id(), user_id=user_id, type=CreditRecordType.EXPIRE.value, amount=-amount, balance_delta=Decimal("0.00"), expired_amount=amount, balance_after=current_available, description=f"积分到期:{float(amount):.2f}", related_id=balance.id, request_time=checked_at, biz_key=f"credit-balance:{balance.id}:expire", billing_scene=CreditRecordBillingScene.CREDIT_EXPIRE.value, credit_level_snapshot=balance.credit_level, ) db.add(record) await db.flush() db.add( CreditRecordAllocation( id=generate_id(), credit_record_id=record.id, credit_balance_id=balance.id, user_id=user_id, allocation_action=CreditAllocationAction.EXPIRE.value, amount=amount, request_time=checked_at, credit_level_snapshot=balance.credit_level, credit_scope_snapshot=balance.credit_scope or CreditScope.PERSONAL.value, source_type_snapshot=balance.source_type, source_id_snapshot=balance.source_id, team_id_snapshot=balance.team_id, team_manager_id_snapshot=manager_map.get(balance.subscription_id or ""), subscription_id_snapshot=balance.subscription_id, subscription_period_id_snapshot=balance.subscription_period_id, seat_id_snapshot=None, valid_from_snapshot=balance.valid_from, expires_at_snapshot=balance.expires_at, unspent_before=amount, unspent_after=Decimal("0.00"), consumed_before=before_consumed, consumed_after=before_consumed, ) ) balance.unspent_amount = Decimal("0.00") balance.expired_amount = to_credit_decimal(balance.expired_amount) + amount balance.expired_processed_at = checked_at balance.status = CreditBalanceStatus.EXPIRED.value await db.flush() return len(balances) async def list_expired_balance_user_limits( db: AsyncSession, *, request_time: datetime, batch_size: int = 500, ) -> list[tuple[str, int]]: result = await db.execute( select(UserCreditBalance.id, UserCreditBalance.user_id) .where( UserCreditBalance.expires_at <= request_time, UserCreditBalance.unspent_amount > 0, UserCreditBalance.expired_processed_at.is_(None), UserCreditBalance.revoked_at.is_(None), ) .order_by(UserCreditBalance.expires_at.asc(), UserCreditBalance.id.asc()) .limit(max(1, batch_size)) ) selected = [(str(row.id), str(row.user_id)) for row in result.all()] per_user_limit = Counter(user_id for _, user_id in selected) return sorted(per_user_limit.items(), key=lambda item: item[0]) async def archive_expired_balances_batch( db: AsyncSession, *, request_time: datetime | None = None, batch_size: int = 500, ) -> int: checked_at = request_time or utc_now() user_limits = await list_expired_balance_user_limits( db, request_time=checked_at, batch_size=batch_size ) processed = 0 for user_id, user_limit in user_limits: processed += await archive_expired_user_balances( db, user_id=user_id, request_time=checked_at, limit=user_limit, ) return processed